From e734f3504fea37179129f1f1451099ebe79e0f14 Mon Sep 17 00:00:00 2001 From: jeromeku Date: Mon, 31 Mar 2025 15:06:08 -0700 Subject: [PATCH] remove redundant code when constructing model names --- unsloth/registry/_deepseek.py | 8 ++------ unsloth/registry/_gemma.py | 4 +--- unsloth/registry/_llama.py | 8 ++------ unsloth/registry/_phi.py | 4 +--- unsloth/registry/_qwen.py | 32 ++++++++------------------------ unsloth/registry/registry.py | 8 +++++--- 6 files changed, 19 insertions(+), 45 deletions(-) diff --git a/unsloth/registry/_deepseek.py b/unsloth/registry/_deepseek.py index 8e87ba11dd..148093155c 100644 --- a/unsloth/registry/_deepseek.py +++ b/unsloth/registry/_deepseek.py @@ -10,9 +10,7 @@ class DeepseekV3ModelInfo(ModelInfo): @classmethod def construct_model_name(cls, base_name, version, size, quant_type, instruct_tag): key = f"{base_name}-V{version}" - key = cls.append_instruct_tag(key, instruct_tag) - key = cls.append_quant_type(key, quant_type) - return key + return super().construct_model_name(base_name, version, size, quant_type, instruct_tag, key) class DeepseekR1ModelInfo(ModelInfo): @classmethod @@ -20,9 +18,7 @@ class DeepseekR1ModelInfo(ModelInfo): key = f"{base_name}-{version}" if version else base_name if size: key = f"{key}-{size}B" - key = cls.append_instruct_tag(key, instruct_tag) - key = cls.append_quant_type(key, quant_type) - return key + return super().construct_model_name(base_name, version, size, quant_type, instruct_tag, key) # Deepseek V3 Model Meta DeepseekV3Meta = ModelMeta( diff --git a/unsloth/registry/_gemma.py b/unsloth/registry/_gemma.py index b9abb3737d..4fef26d533 100644 --- a/unsloth/registry/_gemma.py +++ b/unsloth/registry/_gemma.py @@ -6,9 +6,7 @@ class GemmaModelInfo(ModelInfo): @classmethod def construct_model_name(cls, base_name, version, size, quant_type, instruct_tag): key = f"{base_name}-{version}-{size}B" - key = cls.append_instruct_tag(key, instruct_tag) - key = cls.append_quant_type(key, quant_type) - return key + return super().construct_model_name(base_name, version, size, quant_type, instruct_tag, key) # Gemma3 Base Model Meta GemmaMeta3Base = ModelMeta( diff --git a/unsloth/registry/_llama.py b/unsloth/registry/_llama.py index 6ae838517f..dbf7c8a9d6 100644 --- a/unsloth/registry/_llama.py +++ b/unsloth/registry/_llama.py @@ -8,18 +8,14 @@ class LlamaModelInfo(ModelInfo): @classmethod def construct_model_name(cls, base_name, version, size, quant_type, instruct_tag): key = f"{base_name}-{version}-{size}B" - key = cls.append_instruct_tag(key, instruct_tag) - key = cls.append_quant_type(key, quant_type) - return key + return super().construct_model_name(base_name, version, size, quant_type, instruct_tag, key) class LlamaVisionModelInfo(ModelInfo): @classmethod def construct_model_name(cls, base_name, version, size, quant_type, instruct_tag): key = f"{base_name}-{version}-{size}B-Vision" - key = cls.append_instruct_tag(key, instruct_tag) - key = cls.append_quant_type(key, quant_type) - return key + return super().construct_model_name(base_name, version, size, quant_type, instruct_tag, key) # Llama 3.1 diff --git a/unsloth/registry/_phi.py b/unsloth/registry/_phi.py index a6d18cbd61..c69eaf83bb 100644 --- a/unsloth/registry/_phi.py +++ b/unsloth/registry/_phi.py @@ -7,9 +7,7 @@ class PhiModelInfo(ModelInfo): @classmethod def construct_model_name(cls, base_name, version, size, quant_type, instruct_tag): key = f"{base_name}-{version}" - key = cls.append_instruct_tag(key, instruct_tag) - key = cls.append_quant_type(key, quant_type) - return key + return super().construct_model_name(base_name, version, size, quant_type, instruct_tag, key) # Phi Model Meta PhiMeta = ModelMeta( diff --git a/unsloth/registry/_qwen.py b/unsloth/registry/_qwen.py index 0b902e3130..c9a0a4d4ec 100644 --- a/unsloth/registry/_qwen.py +++ b/unsloth/registry/_qwen.py @@ -5,44 +5,28 @@ _IS_QWEN_VL_REGISTERED = False _IS_QWEN_QWQ_REGISTERED = False class QwenModelInfo(ModelInfo): @classmethod - def construct_model_name( - cls, base_name, version, size, quant_type, instruct_tag - ): + def construct_model_name(cls, base_name, version, size, quant_type, instruct_tag): key = f"{base_name}{version}-{size}B" - key = cls.append_instruct_tag(key, instruct_tag) - key = cls.append_quant_type(key, quant_type) - return key + return super().construct_model_name(base_name, version, size, quant_type, instruct_tag, key) class QwenVLModelInfo(ModelInfo): @classmethod - def construct_model_name( - cls, base_name, version, size, quant_type, instruct_tag - ): + def construct_model_name(cls, base_name, version, size, quant_type, instruct_tag): key = f"{base_name}{version}-VL-{size}B" - key = cls.append_instruct_tag(key, instruct_tag) - key = cls.append_quant_type(key, quant_type) - return key + return super().construct_model_name(base_name, version, size, quant_type, instruct_tag, key) class QwenQwQModelInfo(ModelInfo): @classmethod - def construct_model_name( - cls, base_name, version, size, quant_type, instruct_tag - ): + def construct_model_name(cls, base_name, version, size, quant_type, instruct_tag): key = f"{base_name}-{size}B" - key = cls.append_instruct_tag(key, instruct_tag) - key = cls.append_quant_type(key, quant_type) - return key + return super().construct_model_name(base_name, version, size, quant_type, instruct_tag, key) class QwenQVQPreviewModelInfo(ModelInfo): @classmethod - def construct_model_name( - cls, base_name, version, size, quant_type, instruct_tag - ): + def construct_model_name(cls, base_name, version, size, quant_type, instruct_tag): key = f"{base_name}-{size}B-Preview" - key = cls.append_instruct_tag(key, instruct_tag) - key = cls.append_quant_type(key, quant_type) - return key + return super().construct_model_name(base_name, version, size, quant_type, instruct_tag, key) # Qwen2.5 Model Meta QwenMeta = ModelMeta( diff --git a/unsloth/registry/registry.py b/unsloth/registry/registry.py index 1e2c667e13..590beebeeb 100644 --- a/unsloth/registry/registry.py +++ b/unsloth/registry/registry.py @@ -36,7 +36,7 @@ class ModelInfo: instruct_tag: str = None quant_type: QuantType = None description: str = None - + def __post_init__(self): self.name = self.name or self.construct_model_name( self.base_name, @@ -61,8 +61,10 @@ class ModelInfo: return key @classmethod - def construct_model_name(cls, base_name, version, size, quant_type, instruct_tag): - raise NotImplementedError("Subclass must implement this method") + def construct_model_name(cls, base_name, version, size, quant_type, instruct_tag, key=""): + key = cls.append_instruct_tag(key, instruct_tag) + key = cls.append_quant_type(key, quant_type) + return key @property def model_path(