diff --git a/unsloth/registry/_llama.py b/unsloth/registry/_llama.py index b62491596c..6ae838517f 100644 --- a/unsloth/registry/_llama.py +++ b/unsloth/registry/_llama.py @@ -34,11 +34,23 @@ LlamaMeta3_1 = ModelMeta( quant_types=[QuantType.NONE, QuantType.BNB, QuantType.UNSLOTH], ) -# Llama 3.2 -LlamaMeta3_2 = ModelMeta( +# Llama 3.2 Base Models +LlamaMeta3_2_Base = ModelMeta( org="meta-llama", base_name="Llama", - instruct_tags=[None, "Instruct"], + instruct_tags=[None], + model_version="3.2", + model_sizes=["1", "3"], + model_info_cls=LlamaModelInfo, + is_multimodal=False, + quant_types=[QuantType.NONE, QuantType.BNB, QuantType.UNSLOTH], +) + +# Llama 3.2 Instruction Tuned Models +LlamaMeta3_2_Instruct = ModelMeta( + org="meta-llama", + base_name="Llama", + instruct_tags=["Instruct"], model_version="3.2", model_sizes=["1", "3"], model_info_cls=LlamaModelInfo, @@ -67,7 +79,8 @@ def register_llama_models(include_original_model: bool = False): if _IS_LLAMA_REGISTERED: return _register_models(LlamaMeta3_1, include_original_model=include_original_model) - _register_models(LlamaMeta3_2, include_original_model=include_original_model) + _register_models(LlamaMeta3_2_Base, include_original_model=include_original_model) + _register_models(LlamaMeta3_2_Instruct, include_original_model=include_original_model) _IS_LLAMA_REGISTERED = True @@ -79,8 +92,8 @@ def register_llama_vision_models(include_original_model: bool = False): _IS_LLAMA_VISION_REGISTERED = True -# register_llama_models(include_original_model=True) -register_llama_vision_models(include_original_model=True) +register_llama_models(include_original_model=True) +#register_llama_vision_models(include_original_model=True) if __name__ == "__main__": from unsloth.registry.registry import MODEL_REGISTRY, _check_model_info diff --git a/unsloth/registry/registry.py b/unsloth/registry/registry.py index 6f50f61d63..e7a2be0876 100644 --- a/unsloth/registry/registry.py +++ b/unsloth/registry/registry.py @@ -1,6 +1,6 @@ +import warnings from dataclasses import dataclass, field from enum import Enum -from typing import Literal class QuantType(Enum): @@ -127,7 +127,7 @@ def _check_model_info(model_id: str, properties: list[str] = ["lastModified"]): model_info: HfModelInfo = api.model_info(model_id, expand=properties) except Exception as e: if isinstance(e, RepositoryNotFoundError): - print(f"\u2718 {model_id} not found") + warnings.warn(f"{model_id} not found on Hugging Face") model_info = None else: raise e