handle quant types per model size

This commit is contained in:
jeromeku 2025-03-31 09:27:43 -07:00
commit f1258b164f
3 changed files with 29 additions and 18 deletions

View file

@ -3,6 +3,7 @@ from unsloth.registry.registry import ModelInfo, ModelMeta, QuantType, _register
_IS_LLAMA_REGISTERED = False
_IS_LLAMA_VISION_REGISTERED = False
class LlamaModelInfo(ModelInfo):
@classmethod
def construct_model_name(cls, base_name, version, size, quant_type, instruct_tag):
@ -27,7 +28,7 @@ LlamaMeta3_1 = ModelMeta(
base_name="Llama",
instruct_tags=[None, "Instruct"],
model_version="3.1",
model_sizes=[8],
model_sizes=["8"],
model_info_cls=LlamaModelInfo,
is_multimodal=False,
quant_types=[QuantType.NONE, QuantType.BNB, QuantType.UNSLOTH],
@ -39,10 +40,10 @@ LlamaMeta3_2 = ModelMeta(
base_name="Llama",
instruct_tags=[None, "Instruct"],
model_version="3.2",
model_sizes=[1, 3],
model_sizes=["1", "3"],
model_info_cls=LlamaModelInfo,
is_multimodal=False,
quant_types=[QuantType.NONE, QuantType.BNB, QuantType.UNSLOTH],
quant_types=[QuantType.NONE, QuantType.BNB, QuantType.UNSLOTH, QuantType.GGUF],
)
# Llama 3.2 Vision
@ -51,37 +52,42 @@ LlamaMeta3_2_Vision = ModelMeta(
base_name="Llama",
instruct_tags=[None, "Instruct"],
model_version="3.2",
model_sizes=[11, 90],
model_sizes=["11", "90"],
model_info_cls=LlamaVisionModelInfo,
is_multimodal=True,
quant_types=[QuantType.NONE, QuantType.BNB, QuantType.UNSLOTH],
quant_types={
"11": [QuantType.NONE, QuantType.BNB, QuantType.UNSLOTH],
"90": [QuantType.NONE],
},
)
def register_llama_models():
def register_llama_models(include_original_model: bool = False):
global _IS_LLAMA_REGISTERED
if _IS_LLAMA_REGISTERED:
return
_register_models(LlamaMeta3_1)
_register_models(LlamaMeta3_2)
_register_models(LlamaMeta3_1, include_original_model=include_original_model)
_register_models(LlamaMeta3_2, include_original_model=include_original_model)
_IS_LLAMA_REGISTERED = True
def register_llama_vision_models():
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)
_register_models(LlamaMeta3_2_Vision, include_original_model=include_original_model)
_IS_LLAMA_VISION_REGISTERED = True
register_llama_models()
register_llama_vision_models()
# 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
for model_id, model_info in MODEL_REGISTRY.items():
model_info = _check_model_info(model_id)
if model_info is None:
print(f"\u2718 {model_id}")
else:
print(f"\u2713 {model_id}")
print(f"\u2713 {model_id}")

View file

@ -40,7 +40,7 @@ QwenMeta = ModelMeta(
base_name="Qwen",
instruct_tags=[None, "Instruct"],
model_version="2.5",
model_sizes=[3, 7],
model_sizes=["3", "7"],
model_info_cls=QwenModelInfo,
is_multimodal=False,
quant_types=[QuantType.NONE, QuantType.BNB, QuantType.UNSLOTH],
@ -52,7 +52,7 @@ QwenVLMeta = ModelMeta(
base_name="Qwen",
instruct_tags=["Instruct"], # No base, only instruction tuned
model_version="2.5",
model_sizes=[3, 7, 32, 72],
model_sizes=["3", "7", "32", "72"],
model_info_cls=QwenVLModelInfo,
is_multimodal=True,
quant_types=[QuantType.NONE, QuantType.BNB, QuantType.UNSLOTH],
@ -64,7 +64,7 @@ QwenQwQMeta = ModelMeta(
base_name="QwQ",
instruct_tags=[None],
model_version="",
model_sizes=[32],
model_sizes=["32"],
model_info_cls=QwenQwQModelInfo,
is_multimodal=False,
quant_types=[QuantType.NONE, QuantType.BNB, QuantType.UNSLOTH, QuantType.GGUF],

View file

@ -74,7 +74,7 @@ class ModelMeta:
model_info_cls: type[ModelInfo]
model_sizes: list[str] = field(default_factory=list)
instruct_tags: list[str] = field(default_factory=list)
quant_types: list[QuantType] = field(default_factory=list)
quant_types: list[QuantType] | dict[str, list[QuantType]] = field(default_factory=list)
is_multimodal: bool = False
@ -146,7 +146,12 @@ def _register_models(model_meta: ModelMeta, include_original_model: bool = False
for size in model_sizes:
for instruct_tag in instruct_tags:
for quant_type in quant_types:
# Handle quant types per model size
if isinstance(quant_types, dict):
_quant_types = quant_types[size]
else:
_quant_types = quant_types
for quant_type in _quant_types:
_org = "unsloth" # unsloth models -- these are all quantized versions of the original model
register_model(
model_info_cls=model_info_cls,