* bug fix #2008 (#2039) * fix (#2051) * Update loader.py * Update pyproject.toml * Update pyproject.toml * Update vision.py * more prints * Update loader.py * LoRA 16bit fix * Update vision.py * Update vision.py * Update _utils.py * Update vision.py * move forced float32 * Update _utils.py * Update _utils.py * Update _utils.py * Update _utils.py * move print * Update _utils.py * disable bfloat16 * Fix forced float32 * move float32 * Ensure trust_remote_code propegates down to unsloth_compile_transformers (#2075) * Update _utils.py * Show both `peft_error` and `autoconfig_error`, not just `autoconfig_error` (#2080) When loading a PEFT model fails, only the `autoconfig_error` is shown. Instead of the `peft_error`, which is what really matters when we're trying to load a PEFT adapter, the user will see something like this: ``` RuntimeError: Unrecognized model in my_model. Should have a `model_type` key in its config.json, or contain one of the following strings in its name: albert, align, altclip, ... ``` This PR just changes it so `autoconfig_error` and `peft_error` are both displayed. * fix error message (#2046) * Update vision.py * Update _utils.py * Update pyproject.toml * Update __init__.py * Update __init__.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update rl_replacements.py * Update vision.py * Update rl_replacements.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Remove double generate patch * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update mapper.py * Update vision.py * fix: config.torch_dtype in LlamaModel_fast_forward_inference (#2091) * fix: config.torch_dtype in LlamaModel_fast_forward_inference * Update llama.py * update for consistency --------- Co-authored-by: Daniel Han <danielhanchen@gmail.com> * versioning * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * model_type_arch * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update loader.py * check * Update _utils.py * Update loader.py * Update loader.py * Remove prints * Update README.md typo * Update _utils.py * Update _utils.py * versioning * Update _utils.py * Update _utils.py * Update _utils.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update vision.py * HF Transfer * fix(utils): add missing importlib import to fix NameError (#2134) This commit fixes a NameError that occurs when `importlib` is referenced in _utils.py without being imported, especially when UNSLOTH_USE_MODELSCOPE=1 is enabled. By adding the missing import statement, the code will no longer throw a NameError. * Add QLoRA Train and Merge16bit Test (#2130) * add reference and unsloth lora merging tests * add test / dataset printing to test scripts * allow running tests from repo root * add qlora test readme * more readme edits * ruff formatting * additional readme comments * forgot to add actual tests * add apache license * Update pyproject.toml * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update loader.py * Update loader.py * Revert * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Bug fix * Update mapper.py * check SDPA for Mistral 3, Pixtral * Update vision.py * Versioning * Update rl_replacements.py * Update README.md * add model registry * move hf hub utils to unsloth/utils * refactor global model info dicts to dataclasses * fix dataclass init * fix llama registration * remove deprecated key function * start registry reog * add llama vision * quant types -> Enum * remap literal quant types to QuantType Enum * add llama model registration * fix quant tag mapping * add qwen2.5 models to registry * add option to include original model in registry * handle quant types per model size * separate registration of base and instruct llama3.2 * add QwenQVQ to registry * add gemma3 to registry * add phi * add deepseek v3 * add deepseek r1 base * add deepseek r1 zero * add deepseek distill llama * add deepseek distill models * remove redundant code when constructing model names * add mistral small to registry * rename model registration methods * rename deepseek registration methods * refactor naming for mistral and phi * add global register models * refactor model registration tests for new registry apis * add model search method * remove deprecated registration api * add quant type test * add registry readme * make llama registration more specific * clear registry when executing individual model registration file * more registry readme updates * Update _auto_install.py * Llama4 * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Synthetic data * Update mapper.py * Xet and Synthetic * Update synthetic.py * Update loader.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update pyproject.toml * Delete .gitignore --------- Co-authored-by: Mukkesh Ganesh <mukmckenzie@gmail.com> Co-authored-by: Kareem <81531392+KareemMusleh@users.noreply.github.com> Co-authored-by: Xander Hawthorne <167850078+CuppaXanax@users.noreply.github.com> Co-authored-by: Isaac Breen <isaac.breen@icloud.com> Co-authored-by: lurf21 <93976703+lurf21@users.noreply.github.com> Co-authored-by: Jack Shi Wei Lun <87535974+jackswl@users.noreply.github.com> Co-authored-by: naliazheli <nalia0316@gmail.com> Co-authored-by: jeromeku <jerome.ku@gmail.com> Co-authored-by: Michael Han <107991372+shimmyshimmer@users.noreply.github.com>
91 lines
No EOL
2.9 KiB
Python
91 lines
No EOL
2.9 KiB
Python
"""
|
|
|
|
Test model registration methods
|
|
Checks that model registration methods work for respective models as well as all models
|
|
The check is performed
|
|
- by registering the models
|
|
- checking that the instantiated models can be found on huggingface hub by querying for the model id
|
|
|
|
"""
|
|
|
|
from dataclasses import dataclass
|
|
|
|
import pytest
|
|
from huggingface_hub import ModelInfo as HfModelInfo
|
|
|
|
from unsloth.registry import register_models, search_models
|
|
from unsloth.registry._deepseek import register_deepseek_models
|
|
from unsloth.registry._gemma import register_gemma_models
|
|
from unsloth.registry._llama import register_llama_models
|
|
from unsloth.registry._mistral import register_mistral_models
|
|
from unsloth.registry._phi import register_phi_models
|
|
from unsloth.registry._qwen import register_qwen_models
|
|
from unsloth.registry.registry import MODEL_REGISTRY, QUANT_TAG_MAP, QuantType
|
|
from unsloth.utils.hf_hub import get_model_info
|
|
|
|
MODEL_NAMES = [
|
|
"llama",
|
|
"qwen",
|
|
"mistral",
|
|
"phi",
|
|
"gemma",
|
|
"deepseek",
|
|
]
|
|
MODEL_REGISTRATION_METHODS = [
|
|
register_llama_models,
|
|
register_qwen_models,
|
|
register_mistral_models,
|
|
register_phi_models,
|
|
register_gemma_models,
|
|
register_deepseek_models,
|
|
]
|
|
|
|
|
|
@dataclass
|
|
class ModelTestParam:
|
|
name: str
|
|
register_models: callable
|
|
|
|
|
|
def _test_model_uploaded(model_ids: list[str]):
|
|
missing_models = []
|
|
for _id in model_ids:
|
|
model_info: HfModelInfo = get_model_info(_id)
|
|
if not model_info:
|
|
missing_models.append(_id)
|
|
|
|
return missing_models
|
|
|
|
|
|
TestParams = [
|
|
ModelTestParam(name, models)
|
|
for name, models in zip(MODEL_NAMES, MODEL_REGISTRATION_METHODS)
|
|
]
|
|
|
|
|
|
# Test that model registration methods register respective models
|
|
@pytest.mark.parametrize("model_test_param", TestParams, ids=lambda param: param.name)
|
|
def test_model_registration(model_test_param: ModelTestParam):
|
|
MODEL_REGISTRY.clear()
|
|
registration_method = model_test_param.register_models
|
|
registration_method()
|
|
registered_models = MODEL_REGISTRY.keys()
|
|
missing_models = _test_model_uploaded(registered_models)
|
|
assert not missing_models, (
|
|
f"{model_test_param.name} missing following models: {missing_models}"
|
|
)
|
|
|
|
|
|
def test_all_model_registration():
|
|
register_models()
|
|
registered_models = MODEL_REGISTRY.keys()
|
|
missing_models = _test_model_uploaded(registered_models)
|
|
assert not missing_models, f"Missing following models: {missing_models}"
|
|
|
|
def test_quant_type():
|
|
# Test that the quant_type is correctly set for model paths
|
|
# NOTE: for models registered under org="unsloth" with QuantType.NONE aliases QuantType.UNSLOTH
|
|
dynamic_quant_models = search_models(quant_types=[QuantType.UNSLOTH])
|
|
assert all(m.quant_type == QuantType.UNSLOTH for m in dynamic_quant_models)
|
|
quant_tag = QUANT_TAG_MAP[QuantType.UNSLOTH]
|
|
assert all(quant_tag in m.model_path for m in dynamic_quant_models) |