handle quant types per model size
This commit is contained in:
parent
83eef36689
commit
f1258b164f
3 changed files with 29 additions and 18 deletions
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue