rename model registration methods
This commit is contained in:
parent
0236e05307
commit
e61106ad05
4 changed files with 64 additions and 45 deletions
|
|
@ -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
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue