From e61106ad05cf73407be0ca291fbd3aaa5d72bcaf Mon Sep 17 00:00:00 2001 From: jeromeku Date: Mon, 31 Mar 2025 17:01:51 -0700 Subject: [PATCH] rename model registration methods --- unsloth/registry/__init__.py | 5 ++++ unsloth/registry/_gemma.py | 25 +++++++++++++----- unsloth/registry/_llama.py | 51 ++++++++++++++++++------------------ unsloth/registry/_qwen.py | 36 +++++++++++++------------ 4 files changed, 68 insertions(+), 49 deletions(-) diff --git a/unsloth/registry/__init__.py b/unsloth/registry/__init__.py index e69de29bb2..dd5b45c4ee 100644 --- a/unsloth/registry/__init__.py +++ b/unsloth/registry/__init__.py @@ -0,0 +1,5 @@ +# from ._deepseek import register_deepseek_models, register +# from ._llama import register_llama_models, register_llama_vision_models +# from ._mistral import register_mistral_models +# from ._openai import register_openai_models +# from ._qwen import register_qwen_models diff --git a/unsloth/registry/_gemma.py b/unsloth/registry/_gemma.py index 4fef26d533..8c47e7e69d 100644 --- a/unsloth/registry/_gemma.py +++ b/unsloth/registry/_gemma.py @@ -1,6 +1,7 @@ from unsloth.registry.registry import ModelInfo, ModelMeta, QuantType, _register_models -_IS_GEMMA_REGISTERED = False +_IS_GEMMA_3_BASE_REGISTERED = False +_IS_GEMMA_3_INSTRUCT_REGISTERED = False class GemmaModelInfo(ModelInfo): @classmethod @@ -32,17 +33,27 @@ GemmaMeta3Instruct = ModelMeta( quant_types=[QuantType.NONE, QuantType.BNB, QuantType.UNSLOTH, QuantType.GGUF], ) -def register_gemma_models(include_original_model: bool = False): - global _IS_GEMMA_REGISTERED - if _IS_GEMMA_REGISTERED: +def register_gemma_3_base_models(include_original_model: bool = False): + global _IS_GEMMA_3_BASE_REGISTERED + if _IS_GEMMA_3_BASE_REGISTERED: return _register_models(GemmaMeta3Base, include_original_model=include_original_model) - _register_models(GemmaMeta3Instruct, include_original_model=include_original_model) - _IS_GEMMA_REGISTERED = True + _IS_GEMMA_3_BASE_REGISTERED = True + +def register_gemma_3_instruct_models(include_original_model: bool = False): + global _IS_GEMMA_3_INSTRUCT_REGISTERED + if _IS_GEMMA_3_INSTRUCT_REGISTERED: + return + _register_models(GemmaMeta3Instruct, include_original_model=include_original_model) + _IS_GEMMA_3_INSTRUCT_REGISTERED = True + +def register_gemma_models(include_original_model: bool = False): + register_gemma_3_base_models(include_original_model=include_original_model) + register_gemma_3_instruct_models(include_original_model=include_original_model) -register_gemma_models(include_original_model=True) if __name__ == "__main__": + register_gemma_models(include_original_model=True) from unsloth.registry.registry import MODEL_REGISTRY, _check_model_info for model_id, model_info in MODEL_REGISTRY.items(): model_info = _check_model_info(model_id) diff --git a/unsloth/registry/_llama.py b/unsloth/registry/_llama.py index dbf7c8a9d6..c84c5b8d30 100644 --- a/unsloth/registry/_llama.py +++ b/unsloth/registry/_llama.py @@ -1,7 +1,7 @@ from unsloth.registry.registry import ModelInfo, ModelMeta, QuantType, _register_models -_IS_LLAMA_REGISTERED = False -_IS_LLAMA_VISION_REGISTERED = False +_IS_LLAMA_3_REGISTERED = False +_IS_LLAMA_3_2_VISION_REGISTERED = False class LlamaModelInfo(ModelInfo): @@ -19,7 +19,7 @@ class LlamaVisionModelInfo(ModelInfo): # Llama 3.1 -LlamaMeta3_1 = ModelMeta( +LlamaMeta_3_1 = ModelMeta( org="meta-llama", base_name="Llama", instruct_tags=[None, "Instruct"], @@ -31,7 +31,7 @@ LlamaMeta3_1 = ModelMeta( ) # Llama 3.2 Base Models -LlamaMeta3_2_Base = ModelMeta( +LlamaMeta_3_2_Base = ModelMeta( org="meta-llama", base_name="Llama", instruct_tags=[None], @@ -43,7 +43,7 @@ LlamaMeta3_2_Base = ModelMeta( ) # Llama 3.2 Instruction Tuned Models -LlamaMeta3_2_Instruct = ModelMeta( +LlamaMeta_3_2_Instruct = ModelMeta( org="meta-llama", base_name="Llama", instruct_tags=["Instruct"], @@ -55,7 +55,7 @@ LlamaMeta3_2_Instruct = ModelMeta( ) # Llama 3.2 Vision -LlamaMeta3_2_Vision = ModelMeta( +LlamaMeta_3_2_Vision = ModelMeta( org="meta-llama", base_name="Llama", instruct_tags=[None, "Instruct"], @@ -70,28 +70,29 @@ LlamaMeta3_2_Vision = ModelMeta( ) +def register_llama_3_models(include_original_model: bool = False): + global _IS_LLAMA_3_REGISTERED + if _IS_LLAMA_3_REGISTERED: + return + _register_models(LlamaMeta_3_1, include_original_model=include_original_model) + _register_models(LlamaMeta_3_2_Base, include_original_model=include_original_model) + _register_models(LlamaMeta_3_2_Instruct, include_original_model=include_original_model) + _IS_LLAMA_3_REGISTERED = True + +def register_llama_3_2_vision_models(include_original_model: bool = False): + global _IS_LLAMA_3_2_VISION_REGISTERED + if _IS_LLAMA_3_2_VISION_REGISTERED: + return + _register_models(LlamaMeta_3_2_Vision, include_original_model=include_original_model) + _IS_LLAMA_3_2_VISION_REGISTERED = True + + def register_llama_models(include_original_model: bool = False): - global _IS_LLAMA_REGISTERED - if _IS_LLAMA_REGISTERED: - return - _register_models(LlamaMeta3_1, 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 - - -def register_llama_vision_models(include_original_model: bool = False): - global _IS_LLAMA_VISION_REGISTERED - if _IS_LLAMA_VISION_REGISTERED: - return - _register_models(LlamaMeta3_2_Vision, include_original_model=include_original_model) - _IS_LLAMA_VISION_REGISTERED = True - - -register_llama_models(include_original_model=True) -#register_llama_vision_models(include_original_model=True) + register_llama_3_models(include_original_model=include_original_model) + register_llama_3_2_vision_models(include_original_model=include_original_model) if __name__ == "__main__": + register_llama_models(include_original_model=True) from unsloth.registry.registry import MODEL_REGISTRY, _check_model_info for model_id, model_info in MODEL_REGISTRY.items(): diff --git a/unsloth/registry/_qwen.py b/unsloth/registry/_qwen.py index c9a0a4d4ec..c364f9b099 100644 --- a/unsloth/registry/_qwen.py +++ b/unsloth/registry/_qwen.py @@ -1,7 +1,7 @@ from unsloth.registry.registry import ModelInfo, ModelMeta, QuantType, _register_models -_IS_QWEN_REGISTERED = False -_IS_QWEN_VL_REGISTERED = False +_IS_QWEN_2_5_REGISTERED = False +_IS_QWEN_2_5_VL_REGISTERED = False _IS_QWEN_QWQ_REGISTERED = False class QwenModelInfo(ModelInfo): @classmethod @@ -29,7 +29,7 @@ class QwenQVQPreviewModelInfo(ModelInfo): return super().construct_model_name(base_name, version, size, quant_type, instruct_tag, key) # Qwen2.5 Model Meta -QwenMeta = ModelMeta( +Qwen_2_5_Meta = ModelMeta( org="Qwen", base_name="Qwen", instruct_tags=[None, "Instruct"], @@ -41,7 +41,7 @@ QwenMeta = ModelMeta( ) # Qwen2.5 VL Model Meta -QwenVLMeta = ModelMeta( +Qwen_2_5_VLMeta = ModelMeta( org="Qwen", base_name="Qwen", instruct_tags=["Instruct"], # No base, only instruction tuned @@ -76,19 +76,19 @@ QwenQVQPreviewMeta = ModelMeta( quant_types=[QuantType.NONE, QuantType.BNB], ) -def register_qwen_models(include_original_model: bool = False): - global _IS_QWEN_REGISTERED - if _IS_QWEN_REGISTERED: +def register_qwen_2_5_models(include_original_model: bool = False): + global _IS_QWEN_2_5_REGISTERED + if _IS_QWEN_2_5_REGISTERED: return - _register_models(QwenMeta, include_original_model=include_original_model) - _IS_QWEN_REGISTERED = True + _register_models(Qwen_2_5_Meta, include_original_model=include_original_model) + _IS_QWEN_2_5_REGISTERED = True -def register_qwen_vl_models(include_original_model: bool = False): - global _IS_QWEN_VL_REGISTERED - if _IS_QWEN_VL_REGISTERED: +def register_qwen_2_5_vl_models(include_original_model: bool = False): + global _IS_QWEN_2_5_VL_REGISTERED + if _IS_QWEN_2_5_VL_REGISTERED: return - _register_models(QwenVLMeta, include_original_model=include_original_model) - _IS_QWEN_VL_REGISTERED = True + _register_models(Qwen_2_5_VLMeta, include_original_model=include_original_model) + _IS_QWEN_2_5_VL_REGISTERED = True def register_qwen_qwq_models(include_original_model: bool = False): global _IS_QWEN_QWQ_REGISTERED @@ -98,11 +98,13 @@ def register_qwen_qwq_models(include_original_model: bool = False): _register_models(QwenQVQPreviewMeta, include_original_model=include_original_model) _IS_QWEN_QWQ_REGISTERED = True -# register_qwen_models() -# register_qwen_vl_models() -register_qwen_qwq_models(include_original_model=True) +def register_qwen_models(include_original_model: bool = False): + register_qwen_2_5_models(include_original_model=include_original_model) + register_qwen_2_5_vl_models(include_original_model=include_original_model) + register_qwen_qwq_models(include_original_model=include_original_model) if __name__ == "__main__": + register_qwen_models(include_original_model=True) from unsloth.registry.registry import MODEL_REGISTRY, _check_model_info for model_id, model_info in MODEL_REGISTRY.items(): model_info = _check_model_info(model_id)