diff --git a/tests/python/test_revision_forwarding.py b/tests/python/test_revision_forwarding.py new file mode 100644 index 0000000000..b42604103c --- /dev/null +++ b/tests/python/test_revision_forwarding.py @@ -0,0 +1,879 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. +"""`revision` must reach the config, weight and tokenizer loads (issue #3544). + +FastLlamaModel.from_pretrained took a `revision` argument and never read it, so the +config, weights and tokenizer silently came from the repo's default branch. These are +AST-structural so they need no GPU, no network and no gated checkpoint; importing +unsloth on a CPU runner is what tests/conftest.py exists to work around. +""" + +import ast +import types +from pathlib import Path + +import pytest + + +REPO = Path(__file__).parents[2] +LLAMA = REPO / "unsloth" / "models" / "llama.py" +LOADER = REPO / "unsloth" / "models" / "loader.py" +VISION = REPO / "unsloth" / "models" / "vision.py" +TOKENIZER_UTILS = REPO / "unsloth" / "tokenizer_utils.py" +LOADER_UTILS = REPO / "unsloth" / "models" / "loader_utils.py" +SAVE = REPO / "unsloth" / "save.py" + + +def _tree(path): + return ast.parse(path.read_text(encoding = "utf-8")) + + +def _function( + tree, + name, + class_name = None, +): + body = tree.body + if class_name is not None: + classes = [n for n in body if isinstance(n, ast.ClassDef) and n.name == class_name] + assert classes, f"{class_name} not found" + body = classes[0].body + for node in body: + if isinstance(node, ast.FunctionDef) and node.name == name: + return node + raise AssertionError(f"{class_name or ''}.{name} not found") + + +def _params(function): + return [a.arg for a in function.args.args + function.args.kwonlyargs] + + +def _calls(function, callee): + """Every Call whose dotted name ends with `callee`.""" + return [ + node + for node in ast.walk(function) + if isinstance(node, ast.Call) and ast.unparse(node.func).split(".")[-1] == callee + ] + + +def _revision_kwarg(call): + for keyword in call.keywords: + if keyword.arg == "revision": + return keyword + return None + + +def test_fast_llama_model_reads_its_revision_argument(): + """The whole of #3544: the parameter existed but had zero reads.""" + function = _function(_tree(LLAMA), "from_pretrained", "FastLlamaModel") + assert "revision" in _params(function) + loads = [ + n + for n in ast.walk(function) + if isinstance(n, ast.Name) and n.id == "revision" and isinstance(n.ctx, ast.Load) + ] + assert loads, "revision is accepted but never read" + + +@pytest.mark.parametrize( + "callee, minimum", + [ + ("AutoConfig", 2), # checkpoint probe + main config + ("AutoModelForCausalLM", 2), # user-config and plain branches + ("AutoModelForSequenceClassification", 1), + ("load_correct_tokenizer", 1), + ], +) +def test_llama_loads_forward_revision(callee, minimum): + function = _function(_tree(LLAMA), "from_pretrained", "FastLlamaModel") + calls = ( + _calls(function, "from_pretrained") + if callee != "load_correct_tokenizer" + else _calls(function, "load_correct_tokenizer") + ) + if callee != "load_correct_tokenizer": + calls = [c for c in calls if ast.unparse(c.func).startswith(callee)] + assert len(calls) >= minimum, f"expected >= {minimum} {callee} loads, found {len(calls)}" + for call in calls: + assert _revision_kwarg(call) is not None, f"{callee} at line {call.lineno} drops revision" + + +def test_llama_does_not_pass_revision_to_load_vllm(): + """load_vllm has no `revision` parameter and load_vllm_kwargs is not filtered, + so putting one in that dict is an unconditional TypeError on the vLLM path.""" + function = _function(_tree(LLAMA), "from_pretrained", "FastLlamaModel") + dicts = [ + node.value + for node in ast.walk(function) + if isinstance(node, ast.Assign) + and any(getattr(t, "id", None) == "load_vllm_kwargs" for t in node.targets) + and isinstance(node.value, ast.Call) + ] + assert dicts, "load_vllm_kwargs assignment not found" + for call in dicts: + assert _revision_kwarg(call) is None, "revision is not a load_vllm argument" + + +def test_fast_base_model_does_not_bind_revision(): + """vision.py's weight load forwards **kwargs, so binding `revision` as a named + parameter would silently drop it from there and from kwargs.get('revision').""" + function = _function(_tree(VISION), "from_pretrained", "FastBaseModel") + assert "revision" not in _params(function) + assert function.args.kwarg is not None, "**kwargs is what carries revision here" + + +@pytest.mark.parametrize( + "callee, minimum", + [("AutoConfig", 4), ("auto_processor", 2), ("_AutoTokenizer", 2)], +) +def test_vision_loads_forward_revision(callee, minimum): + function = _function(_tree(VISION), "from_pretrained", "FastBaseModel") + calls = [ + c for c in _calls(function, "from_pretrained") if ast.unparse(c.func).startswith(callee) + ] + assert len(calls) >= minimum, f"expected >= {minimum} {callee} loads, found {len(calls)}" + for call in calls: + assert _revision_kwarg(call) is not None, f"{callee} at line {call.lineno} drops revision" + + +@pytest.mark.parametrize("name", ["load_correct_tokenizer", "_load_correct_tokenizer"]) +def test_tokenizer_helpers_accept_revision(name): + assert "revision" in _params(_function(_tree(TOKENIZER_UTILS), name)) + + +def test_tokenizer_helpers_forward_revision(): + tree = _tree(TOKENIZER_UTILS) + public = _function(tree, "load_correct_tokenizer") + inner = _calls(public, "_load_correct_tokenizer") + assert len(inner) == 1 and _revision_kwarg(inner[0]) is not None + + private = _function(tree, "_load_correct_tokenizer") + loads = _calls(private, "from_pretrained") + assert len(loads) >= 2, "expected the slow and fast tokenizer loads" + for call in loads: + assert ( + _revision_kwarg(call) is not None + ), f"tokenizer load at line {call.lineno} drops revision" + + +def _load_gate(): + """Exec just _revision_for_resolved_repo, so no GPU-bound import is needed.""" + source = LOADER.read_text(encoding = "utf-8") + function = _function(ast.parse(source), "_revision_for_resolved_repo") + namespace = {"logger": types.SimpleNamespace(warning_once = lambda *a, **k: None)} + module = ast.Module(body = [function], type_ignores = []) + ast.fix_missing_locations(module) + exec(compile(module, str(LOADER), "exec"), namespace) + return namespace["_revision_for_resolved_repo"] + + +def test_revision_survives_when_the_repo_is_unchanged(): + # The reported case: a user's own repo is never in the mapper tables. + gate = _load_gate() + assert gate("my-branch", "myorg/my-ft", "myorg/my-ft") == "my-branch" + + +@pytest.mark.parametrize( + "model_name, old_model_name", + [ + ("unsloth/llama-3-8b-bnb-4bit", "meta-llama/Meta-Llama-3-8B"), # prequant mirror + ("unsloth/Qwen3-30B-A3B", "unsloth/Qwen3-30B-A3B-bnb-4bit"), # suffix strip + ("/tmp/unsloth-fp8-cache/model", "meta-llama/Meta-Llama-3-8B"), # fp8 temp dir + ], +) +def test_revision_is_dropped_once_the_repo_is_remapped(model_name, old_model_name): + # The ref only exists on the repo the caller named, so pinning it elsewhere + # would 404 or, worse, resolve a same-named branch on a different repo. + assert _load_gate()("abc123", model_name, old_model_name) is None + + +def test_no_revision_stays_none_even_when_remapped(): + gate = _load_gate() + assert gate(None, "unsloth/llama-3-8b-bnb-4bit", "meta-llama/Meta-Llama-3-8B") is None + + +def _gate_with_warnings(): + source = LOADER.read_text(encoding = "utf-8") + function = _function(ast.parse(source), "_revision_for_resolved_repo") + warnings = [] + namespace = {"logger": types.SimpleNamespace(warning_once = lambda m: warnings.append(m))} + module = ast.Module(body = [function], type_ignores = []) + ast.fix_missing_locations(module) + exec(compile(module, str(LOADER), "exec"), namespace) + return namespace["_revision_for_resolved_repo"], warnings + + +def test_the_gate_warns_exactly_once_when_it_drops_a_revision(): + gate, warnings = _gate_with_warnings() + gate("abc123", "unsloth/x-bnb-4bit", "org/x", True) + assert len(warnings) == 1 + message = warnings[0] + # Both repos have to be named or the user cannot tell which load was silently redirected. + assert "abc123" in message and "org/x" in message and "unsloth/x-bnb-4bit" in message + + +def test_exact_name_mode_is_only_offered_when_it_would_help(): + """It gates the mapper substitution alone. The ModelScope download, the + ALLOW_PREQUANTIZED_MODELS strip and fast_inference_setup all ignore it, so + recommending it there sends the caller round the same loop.""" + gate, warnings = _gate_with_warnings() + gate("abc123", "unsloth/x-bnb-4bit", "org/x", True) + assert "use_exact_model_name" in warnings[0] + + gate, warnings = _gate_with_warnings() + gate("abc123", "/tmp/modelscope/x", "org/x", False) + assert "use_exact_model_name" not in warnings[0] + + +def test_both_loader_paths_pass_the_mapper_flag(): + tree = _tree(LOADER) + for class_name in ("FastLanguageModel", "FastModel"): + function = _function(tree, "from_pretrained", class_name) + assert [ + n + for n in ast.walk(function) + if isinstance(n, ast.Assign) + and any(getattr(t, "id", None) == "mapper_moved_name" for t in n.targets) + ], f"{class_name} must record whether the mapper moved the name" + for call in _calls(function, "_revision_for_resolved_repo"): + names = [getattr(a, "id", None) for a in call.args] + assert "mapper_moved_name" in names, "the gate needs the flag to tailor its remedy" + + +def test_both_loader_paths_gate_before_and_after_resolution(): + """The gate has to run before the AutoConfig / PeftConfig probes, or a pinned 4bit + load fails against the mirror instead of warning, and again after the last remap.""" + tree = _tree(LOADER) + for class_name in ("FastLanguageModel", "FastModel"): + function = _function(tree, "from_pretrained", class_name) + gates = _calls(function, "_revision_for_resolved_repo") + assert len(gates) == 2, f"{class_name} needs an early and a late gate, found {len(gates)}" + early, late = sorted(gates, key = lambda c: c.lineno) + + probes = [ + c + for c in _calls(function, "from_pretrained") + if ast.unparse(c.func).split(".")[0] in ("AutoConfig", "PeftConfig") + ] + assert probes, f"{class_name} has no config probe" + gated = 0 + for probe in probes: + assert probe.lineno > early.lineno, "the gate must precede the config probes" + keyword = _revision_kwarg(probe) + if keyword is None: + continue # the PEFT base-model probe deliberately pins nothing + # adapter_revision is the same gated value, taken before the vLLM drop that + # only the base model's config and weights answer to. + assert getattr(keyword.value, "id", None) in ( + "base_revision", + "adapter_revision", + ), f"probe at line {probe.lineno} uses the ungated revision" + gated += 1 + assert gated >= 2, f"{class_name} must gate its AutoConfig and PeftConfig probes" + + # The late gate feeds on base_revision so an already-dropped one warns only once. + assert getattr(late.args[0], "id", None) == "base_revision" + + +def test_the_late_gate_is_skipped_for_peft(): + """On a PEFT load model_name is necessarily the base model, so the remap warning + would fire for every versioned adapter while PeftModel loads the ref correctly.""" + tree = _tree(LOADER) + for class_name in ("FastLanguageModel", "FastModel"): + function = _function(tree, "from_pretrained", class_name) + late = sorted(_calls(function, "_revision_for_resolved_repo"), key = lambda c: c.lineno)[-1] + guards = [ + n + for n in ast.walk(function) + if isinstance(n, ast.If) + and ast.unparse(n.test).replace(" ", "") == "notis_peft" + and n.lineno <= late.lineno <= n.end_lineno + ] + assert guards, f"{class_name}'s late gate must sit under `if not is_peft`" + + +def test_the_adapter_load_keeps_the_callers_revision(): + """`revision` names the adapter repo, so PeftModel must get the ungated value.""" + tree = _tree(LOADER) + for class_name in ("FastLanguageModel", "FastModel"): + function = _function(tree, "from_pretrained", class_name) + peft_loads = [ + c + for c in _calls(function, "from_pretrained") + if ast.unparse(c.func).startswith("PeftModel") + ] + assert peft_loads, f"{class_name} has no PeftModel load" + for call in peft_loads: + keyword = _revision_kwarg(call) + assert keyword is not None and getattr(keyword.value, "id", None) == "revision" + + +@pytest.mark.parametrize("path, flag", [(LLAMA, "revision"), (VISION, "_revision")]) +def test_a_pinned_load_does_not_mix_refs_with_vllm(path, flag): + """load_vllm takes no revision, so vLLM fetches the default branch. Pinning only the + config and tokenizer would put two refs in one model, so the pin is dropped instead.""" + source = path.read_text(encoding = "utf-8") + tree = ast.parse(source) + name = "FastLlamaModel" if path is LLAMA else "FastBaseModel" + function = _function(tree, "from_pretrained", name) + clears = [ + n + for n in ast.walk(function) + if isinstance(n, ast.Assign) + and any(getattr(t, "id", None) == flag for t in n.targets) + and isinstance(n.value, ast.Constant) + and n.value.value is None + ] + assert clears, f"{path.name} never drops the revision on the vLLM path" + # It must happen before the config load, or the config is pinned and the weights are not. + configs = [ + c + for c in _calls(function, "from_pretrained") + if ast.unparse(c.func).startswith("AutoConfig") + ] + assert configs + assert min(c.lineno for c in clears) < min(c.lineno for c in configs) + + +def test_local_snapshot_resolution_takes_the_revision(): + """A local snapshot dir cannot be re-pointed by a revision handed to from_pretrained, + so the cache resolution itself has to select the requested ref.""" + tree = _tree(LOADER_UTILS) + for name in ("_resolve_hub_repo_local_dir", "_hub_repo_or_local_path"): + function = _function(tree, name) + assert "revision" in _params(function), f"{name} must accept revision" + resolver = _function(tree, "_resolve_hub_repo_local_dir") + downloads = _calls(resolver, "hf_hub_download") + assert downloads, "expected the cache probe download" + for call in downloads: + assert _revision_kwarg(call) is not None + wrapper = _function(tree, "_hub_repo_or_local_path") + inner = _calls(wrapper, "_resolve_hub_repo_local_dir") + assert inner and all(_revision_kwarg(c) is not None for c in inner) + + +def test_the_vllm_drop_only_fires_when_vllm_owns_the_weights(): + """fast_inference is turned off in that same block when vLLM is missing or the GPU + is too old, and a num_labels load goes through transformers regardless. Both of + those can honour the pin, so the drop must not be unconditional.""" + function = _function(_tree(LLAMA), "from_pretrained", "FastLlamaModel") + clears = [ + n + for n in ast.walk(function) + if isinstance(n, ast.Assign) + and any(getattr(t, "id", None) == "revision" for t in n.targets) + and isinstance(n.value, ast.Constant) + and n.value.value is None + ] + assert clears, "the vLLM revision drop is gone" + for clear in clears: + guards = [ + n + for n in ast.walk(function) + if isinstance(n, ast.If) + and n.lineno <= clear.lineno <= n.end_lineno + and "revision" in ast.unparse(n.test) + ] + assert guards, "the drop needs its own condition" + test = ast.unparse(guards[0].test) + assert "fast_inference" in test, "must re-check fast_inference" + assert "num_labels" in test, "a num_labels load runs in-process and can be pinned" + + +def test_the_tokenizer_revision_is_resolved_by_the_loader(): + """The tokenizer repo is not always the base model's: a PEFT load whose + tokenizer_name is the adapter keeps the caller's ref, which the base model cannot.""" + tree = _tree(LOADER) + helper = _function(tree, "_revision_for_tokenizer_repo") + assert helper, "the loader must resolve the tokenizer repo's revision" + for class_name in ("FastLanguageModel", "FastModel"): + function = _function(tree, "from_pretrained", class_name) + dispatches = [ + c + for c in _calls(function, "from_pretrained") + if any(k.arg == "tokenizer_revision" for k in c.keywords) + ] + assert dispatches, f"{class_name} must dispatch a tokenizer_revision" + for call in dispatches: + keyword = next(k for k in call.keywords if k.arg == "tokenizer_revision") + assert isinstance(keyword.value, ast.Call), "it has to be the resolved value" + + +def test_llama_uses_the_dispatched_tokenizer_revision(): + function = _function(_tree(LLAMA), "from_pretrained", "FastLlamaModel") + assert "tokenizer_revision" in _params(function) + loads = _calls(function, "load_correct_tokenizer") + assert loads + for call in loads: + keyword = _revision_kwarg(call) + assert keyword is not None + assert getattr(keyword.value, "id", None) == "tokenizer_revision" + + +def test_vision_pops_the_tokenizer_revision_before_the_weight_load(): + """FastBaseModel forwards **kwargs to the weight load, and transformers has no + tokenizer_revision argument, so it must be popped rather than read.""" + function = _function(_tree(VISION), "from_pretrained", "FastBaseModel") + pops = [ + c + for c in ast.walk(function) + if isinstance(c, ast.Call) + and ast.unparse(c.func).endswith("kwargs.pop") + and c.args + and getattr(c.args[0], "value", None) == "tokenizer_revision" + ] + assert pops, "tokenizer_revision must be popped from kwargs" + weight_loads = [ + c + for c in _calls(function, "from_pretrained") + if ast.unparse(c.func).startswith("auto_model") + ] + assert weight_loads + assert pops[0].lineno < min(c.lineno for c in weight_loads) + + +def _load_tokenizer_gate(): + source = LOADER.read_text(encoding = "utf-8") + function = _function(ast.parse(source), "_revision_for_tokenizer_repo") + namespace = {} + module = ast.Module(body = [function], type_ignores = []) + ast.fix_missing_locations(module) + exec(compile(module, str(LOADER), "exec"), namespace) + return namespace["_revision_for_tokenizer_repo"] + + +def test_an_adapter_ref_never_reaches_the_base_tokenizer(): + """On a PEFT load the late gate is skipped, so the gated value still names the + adapter. The base repo's tokenizer must take the model load's ref, which is None.""" + gate = _load_tokenizer_gate() + # Remote adapter, no explicit tokenizer_name: the tokenizer follows the base model. + assert gate(None, "org/base", "org/adapter", "v2", None) is None + + +def test_an_adapter_hosted_tokenizer_keeps_the_callers_ref(): + """An adapter is a separate repo with its own history, so the caller's ref still + names it even though the base model it sits on cannot answer to it.""" + gate = _load_tokenizer_gate() + assert gate("org/adapter", "org/base", "org/adapter", "v2", None, True) == "v2" + + +def test_a_remapped_plain_load_drops_the_tokenizer_pin_too(): + """Naming the requested repo as tokenizer_name must not smuggle the ref back in: the + weights now come off a mirror's default branch, and a pinned tokenizer beside them is + the ref mismatch the gate exists to prevent. Only a PEFT adapter is a separate repo.""" + gate = _load_tokenizer_gate() + assert gate("org/model", "unsloth/model-bnb-4bit", "org/model", "v2", None) is None + + +def test_a_plain_load_gives_the_tokenizer_the_model_ref(): + gate = _load_tokenizer_gate() + assert gate(None, "org/model", "org/model", "v2", "v2") == "v2" + # A third-party tokenizer repo is pinned by neither. + assert gate("other/tok", "org/model", "org/model", "v2", "v2") is None + + +def test_both_dispatches_share_one_model_revision(): + """The value handed to the base load and the one the tokenizer resolution sees have + to be the same, or the PEFT case leaks the adapter ref into the base repo.""" + tree = _tree(LOADER) + for class_name in ("FastLanguageModel", "FastModel"): + function = _function(tree, "from_pretrained", class_name) + assert [ + n + for n in ast.walk(function) + if isinstance(n, ast.Assign) + and any(getattr(t, "id", None) == "model_revision" for t in n.targets) + ], f"{class_name} must derive one model_revision" + dispatch = next( + c + for c in _calls(function, "from_pretrained") + if any(k.arg == "tokenizer_revision" for k in c.keywords) + ) + model_kw = _revision_kwarg(dispatch) + assert getattr(model_kw.value, "id", None) == "model_revision" + tok_kw = next(k for k in dispatch.keywords if k.arg == "tokenizer_revision") + assert "model_revision" in ast.unparse(tok_kw.value) + + +def test_a_direct_llama_call_still_pins_its_tokenizer(): + """FastLlamaModel is exported and the architecture wrappers forward only `revision`, + so tokenizer_revision has to fall back to it when the repos are the same.""" + function = _function(_tree(LLAMA), "from_pretrained", "FastLlamaModel") + fallbacks = [ + n + for n in ast.walk(function) + if isinstance(n, ast.Assign) + and any(getattr(t, "id", None) == "tokenizer_revision" for t in n.targets) + and getattr(n.value, "id", None) == "revision" + ] + assert fallbacks, "no fallback from revision to tokenizer_revision" + warms = _calls(function, "maybe_prefetch_hf_snapshot") + tokenizer_warms = [ + c + for c in warms + if any( + k.arg == "revision" and getattr(k.value, "id", None) == "tokenizer_revision" + for k in c.keywords + ) + ] + assert tokenizer_warms, "the tokenizer warm should use the same pin" + # The fallback must precede the warm, or the warm fetches the wrong ref. + assert fallbacks[0].lineno < min(c.lineno for c in tokenizer_warms) + + +@pytest.mark.parametrize( + "path, cls, name", + [ + (LLAMA, "FastLlamaModel", "tokenizer_revision"), + (VISION, "FastBaseModel", "_tokenizer_revision_arg"), + ], +) +def test_the_vllm_drop_clears_the_tokenizer_pin_too(path, cls, name): + """Clearing only the model pin left vLLM on the default branch while the tokenizer + stayed on the requested ref.""" + function = _function(ast.parse(path.read_text(encoding = "utf-8")), "from_pretrained", cls) + clears = [ + n + for n in ast.walk(function) + if isinstance(n, ast.Assign) + and any(getattr(t, "id", None) == name for t in n.targets) + and isinstance(n.value, ast.Constant) + and n.value.value is None + ] + assert clears, f"{path.name} never clears {name} on the vLLM path" + + +def _simulate_loader(): + """Run the loader's two revision decisions the way from_pretrained sequences them.""" + tree = ast.parse(LOADER.read_text(encoding = "utf-8")) + namespace = {"logger": types.SimpleNamespace(warning_once = lambda *a, **k: None)} + functions = [ + n + for n in tree.body + if isinstance(n, ast.FunctionDef) + and n.name in ("_revision_for_resolved_repo", "_revision_for_tokenizer_repo") + ] + module = ast.Module(body = functions, type_ignores = []) + ast.fix_missing_locations(module) + exec(compile(module, str(LOADER), "exec"), namespace) + gate = namespace["_revision_for_resolved_repo"] + tokenizer_gate = namespace["_revision_for_tokenizer_repo"] + + def run(old_model_name, model_name, is_peft, tokenizer_name, revision, mapper_moved_name): + base_revision = gate(revision, model_name, old_model_name, mapper_moved_name) + if not is_peft: + base_revision = gate(base_revision, model_name, old_model_name, mapper_moved_name) + model_revision = base_revision if not is_peft else None + return model_revision, tokenizer_gate( + tokenizer_name, model_name, old_model_name, revision, model_revision, is_peft + ) + + return run + + +@pytest.mark.parametrize( + "label, old_model_name, model_name, is_peft, tokenizer_name, revision, mapper_moved_name," + " expected_model, expected_tokenizer", + [ + ("plain pinned load", "org/m", "org/m", False, None, "v2", False, "v2", "v2"), + ( + "remapped to a prequant mirror", + "org/m", + "unsloth/m-bnb-4bit", + False, + None, + "v2", + True, + None, + None, + ), + # The adapter's ref is not the base repo's, and the tokenizer follows the base. + ("PEFT, remote adapter", "org/ad", "org/base", True, None, "v2", False, None, None), + # ... unless the tokenizer is the adapter itself, which the caller did pin. + ( + "PEFT, adapter-hosted tokenizer", + "org/ad", + "org/base", + True, + "org/ad", + "v2", + False, + None, + "v2", + ), + ( + "plain load, third-party tokenizer", + "org/m", + "org/m", + False, + "other/tok", + "v2", + False, + "v2", + None, + ), + # Naming the requested repo back does not survive the remap: the weights moved. + ( + "remapped, tokenizer named as the requested repo", + "org/m", + "unsloth/m-bnb-4bit", + False, + "org/m", + "v2", + True, + None, + None, + ), + ("no revision at all", "org/m", "unsloth/m-bnb-4bit", False, None, None, True, None, None), + ], + ids = lambda v: v if isinstance(v, str) and " " in v else None, +) +def test_the_revision_decision_matrix( + label, + old_model_name, + model_name, + is_peft, + tokenizer_name, + revision, + mapper_moved_name, + expected_model, + expected_tokenizer, +): + """One table for the whole contract: which repo each pin is allowed to reach.""" + run = _simulate_loader() + model_revision, tokenizer_revision = run( + old_model_name, model_name, is_peft, tokenizer_name, revision, mapper_moved_name + ) + assert model_revision == expected_model, label + assert tokenizer_revision == expected_tokenizer, label + + +def test_the_processor_fallback_carries_the_tokenizer_revision(): + """get_auto_processor runs when AutoProcessor raises, so it is a real load path: an + unpinned one there hands back a default-branch processor beside pinned weights.""" + function = _function(_tree(VISION), "from_pretrained", "FastBaseModel") + fallbacks = _calls(function, "get_auto_processor") + assert fallbacks, "the processor fallback must still exist" + for call in fallbacks: + keyword = _revision_kwarg(call) + assert keyword is not None, "the fallback needs the revision too" + assert getattr(keyword.value, "id", None) == "_tokenizer_revision" + + +def test_the_fp8_quantizer_takes_the_requested_revision(): + """Its output path replaces model_name, so the gate downstream drops the pin. If it + did not quantize the pinned ref itself, that ref never reaches the weights at all.""" + function = _function(_tree(LOADER_UTILS), "_offline_quantize_to_fp8") + assert "revision" in _params(function) + for callee in ("from_pretrained",): + loads = _calls(function, callee) + assert loads + for call in loads: + assert _revision_kwarg(call) is not None, ast.unparse(call.func) + + +def test_the_fp8_cache_name_is_revision_specific(): + """A shared temp dir keyed only on the repo name would serve one ref's artifact to + another, and the artifact outlives the process that built it.""" + function = _function(_tree(LOADER_UTILS), "_offline_quantize_to_fp8") + writes = [ + n + for n in ast.walk(function) + if isinstance(n, ast.AugAssign) and getattr(n.target, "id", None) == "cache_name" + ] + assert writes + guarded = [ + n + for n in ast.walk(function) + if isinstance(n, ast.If) + and "revision" in ast.unparse(n.test) + and any(n.lineno <= w.lineno <= n.end_lineno for w in writes) + ] + assert guarded, "two revisions of one repo would share a cache entry" + + +def test_both_loaders_hand_the_fp8_quantizer_the_revision(): + tree = _tree(LOADER) + for class_name in ("FastLanguageModel", "FastModel"): + function = _function(tree, "from_pretrained", class_name) + calls = _calls(function, "_offline_quantize_to_fp8") + assert calls, f"{class_name} must still quantize on the fly" + for call in calls: + keyword = _revision_kwarg(call) + assert keyword is not None + assert ( + getattr(keyword.value, "id", None) == "revision" + ), "the fp8 source is still the caller's own repo here" + + +def test_the_vllm_drop_happens_before_the_config_probe(): + """model_types, auto_model and the text-only decision all come off the probed config. + Reading it at a ref vLLM will not fetch picks the dispatch for a different model, so + the pin has to be gone before the probe, not just before the dispatch.""" + function = _function(_tree(LOADER), "from_pretrained", "FastModel") + drops = [ + n + for n in ast.walk(function) + if isinstance(n, ast.If) + and "is_vLLM_available" in ast.unparse(n.test) + and any( + isinstance(b, ast.Assign) + and any(getattr(t, "id", None) == "base_revision" for t in b.targets) + and isinstance(b.value, ast.Constant) + and b.value.value is None + for b in n.body + ) + ] + assert drops, "FastModel never drops base_revision for the vLLM path" + probes = [ + c + for c in _calls(function, "from_pretrained") + if ast.unparse(c.func).split(".")[0] in ("AutoConfig", "PeftConfig") + ] + assert probes + assert drops[0].end_lineno < min( + c.lineno for c in probes + ), "the probe would read a ref the weights will not be at" + # The probed config goes down untouched again, so nothing may re-gate it at dispatch. + dispatches = [ + c + for c in _calls(function, "from_pretrained") + if ast.unparse(c.func).startswith("FastBaseModel") + ] + assert dispatches + for call in dispatches: + keyword = next((k for k in call.keywords if k.arg == "auto_config"), None) + assert keyword is not None + assert getattr(keyword.value, "id", None) == "model_config" + + +def test_the_fp8_cache_key_survives_a_lossy_sanitization(): + """The readable half replaces every unsafe character with the same one, so `a/b` and + `a.b` collapse together. Only a digest of the raw ref keeps them apart.""" + function = _function(_tree(LOADER_UTILS), "_offline_quantize_to_fp8") + source = ast.unparse(function) + assert "sha256" in source or "blake2" in source, "the sanitized name alone collides" + digests = [ + n for n in ast.walk(function) if isinstance(n, ast.Call) and "sha256" in ast.unparse(n.func) + ] + assert digests + assert any("revision" in ast.unparse(n) for n in digests), "hash the ref, not the repo" + + +def test_the_peft_probe_keeps_the_adapter_ref_under_vllm(): + """The vLLM drop runs before is_peft is known. An adapter is loaded in-process by peft, + so zeroing its probe would read the default branch and either miss PEFT entirely or + resolve a different base model before attaching the pinned adapter.""" + tree = _tree(LOADER) + for class_name in ("FastLanguageModel", "FastModel"): + function = _function(tree, "from_pretrained", class_name) + probes = [ + c + for c in _calls(function, "from_pretrained") + if ast.unparse(c.func).startswith("PeftConfig") + ] + assert probes, f"{class_name} must still probe for an adapter" + for call in probes: + keyword = _revision_kwarg(call) + assert keyword is not None + assert ( + getattr(keyword.value, "id", None) == "adapter_revision" + ), "the adapter probe must not take the base model's gated ref" + + +def test_both_loaders_drop_the_vllm_pin_before_the_probe(): + """model_types picks the architecture class off the probed config, so reading it at a + ref vLLM will not fetch dispatches the wrong one.""" + tree = _tree(LOADER) + for class_name in ("FastLanguageModel", "FastModel"): + function = _function(tree, "from_pretrained", class_name) + drops = [ + n + for n in ast.walk(function) + if isinstance(n, ast.If) + and any( + isinstance(b, ast.Assign) + and any(getattr(t, "id", None) == "base_revision" for t in b.targets) + and isinstance(b.value, ast.Constant) + and b.value.value is None + for b in n.body + ) + ] + assert drops, f"{class_name} never drops base_revision for the vLLM path" + probes = [ + c + for c in _calls(function, "from_pretrained") + if ast.unparse(c.func).split(".")[0] in ("AutoConfig", "PeftConfig") + ] + assert probes + assert drops[0].end_lineno < min( + c.lineno for c in probes + ), f"{class_name} probes at a ref the weights will not be at" + + +def test_llama_owns_the_vllm_predicate_the_loader_gates_on(): + """FastLanguageModel also falls back in-process on pre-Volta GPUs and for a num_labels + load, so the loader cannot gate on `fast_inference and is_vLLM_available()` the way the + FastModel path does. One helper, used by both, or the two drift apart.""" + tree = _tree(LLAMA) + helper = _function(tree, "_vllm_will_load_weights") + source = ast.unparse(helper) + for token in ("is_vLLM_available", "get_device_capability", "hip", "num_labels"): + assert token in source, token + + guard = _function(tree, "from_pretrained", "FastLlamaModel") + assert _calls(guard, "_vllm_will_load_weights"), "llama.py must use its own helper" + loader = _function(_tree(LOADER), "from_pretrained", "FastLanguageModel") + assert _calls(loader, "_vllm_will_load_weights"), "the loader must gate on the same one" + + +def test_a_pinned_tokenizer_is_stamped_for_the_save_path(): + """save.py restores tokenizer.model from tokenizer.name_or_path, which names the repo + but not the branch, so a merged export would copy the default branch's asset.""" + stamps = _calls( + _function(_tree(TOKENIZER_UTILS), "load_correct_tokenizer"), "_mark_loaded_revision" + ) + assert stamps, "the loaded ref has to travel with the tokenizer" + assert any( + any(getattr(a, "id", None) == "revision" for a in c.args) for c in stamps + ), "stamp the ref that was actually loaded" + + tree = _tree(LOADER_UTILS) + assert _function(tree, "_mark_loaded_revision") + assert _function(tree, "_tokenizer_revision") + assert "revision" in _params(_function(tree, "_resolve_hub_repo_cached_file")) + + +@pytest.mark.parametrize( + "callee", ["_resolve_hub_repo_cached_file", "hf_hub_download", "model_info"] +) +def test_the_sentencepiece_restore_reads_the_stamped_ref(callee): + tree = _tree(SAVE) + functions = [ + n + for n in ast.walk(tree) + if isinstance(n, ast.FunctionDef) + and n.name in ("_has_tokenizer_model", "_preserve_sentencepiece_tokenizer_assets") + ] + assert functions + calls = [c for f in functions for c in _calls(f, callee)] + assert calls, f"{callee} not found on the save path" + for call in calls: + assert _revision_kwarg(call) is not None, f"{callee} at line {call.lineno} drops the ref" + + +def test_the_vision_path_stamps_its_pinned_tokenizer_too(): + """FastBaseModel builds its processor without load_correct_tokenizer, so the stamp the + save path reads has to be applied here as well or a pinned VLM load saves the default + branch's tokenizer.model. At the return, so a patch fallback cannot lose it.""" + function = _function(_tree(VISION), "from_pretrained", "FastBaseModel") + stamps = _calls(function, "_mark_loaded_revision") + assert stamps, "the vision path never stamps its loaded ref" + for call in stamps: + assert any( + getattr(a, "id", None) == "_tokenizer_revision" for a in call.args + ), "stamp the ref the tokenizer was actually read at" + returns = [n for n in ast.walk(function) if isinstance(n, ast.Return)] + assert returns + assert max(c.lineno for c in stamps) < max(r.lineno for r in returns) diff --git a/unsloth/models/llama.py b/unsloth/models/llama.py index 4ae55d976b..d4688d30a1 100644 --- a/unsloth/models/llama.py +++ b/unsloth/models/llama.py @@ -2244,6 +2244,26 @@ def unsloth_fast_generate(self, *args, **kwargs): return output +def _vllm_will_load_weights(fast_inference, num_labels = None): + """Whether vLLM, which takes no revision, ends up owning the weight load. + + The loader has to answer this before it probes the config, since that probe's ref + decides which architecture class the load is dispatched to. Mirrors the checks at the + top of from_pretrained below, which calls this too so the two cannot drift. + """ + if not fast_inference or num_labels is not None: + return False + # from_pretrained clears fast_inference when vLLM is missing and then re-enables it on + # hip, so hip ends up True either way. + if DEVICE_TYPE == "hip": + return True + if not is_vLLM_available(): + return False + if DEVICE_TYPE == "cuda" and torch.cuda.get_device_capability()[0] < 7: + return False + return True + + class FastLlamaModel: @staticmethod def _prepare_for_qat(model, qat_scheme): @@ -2299,6 +2319,7 @@ def from_pretrained( tokenizer_name = None, trust_remote_code = False, revision = None, + tokenizer_revision = None, fast_inference = False, # uses vLLM gpu_memory_utilization = 0.5, float8_kv_cache = False, @@ -2338,6 +2359,24 @@ def from_pretrained( raise RuntimeError( "Unsloth: `unsloth_vllm_standby` is True, but environment variable `UNSLOTH_VLLM_STANDBY` is not set to 1!" ) + # Only vLLM cannot take a revision. fast_inference may have just been turned + # off above, and a num_labels load goes in-process regardless; both of those + # can honour the pin, so use the same predicate as the prefetch warm below. + # Through the helper, which the loader also uses to gate its config probe. + if _vllm_will_load_weights(fast_inference, num_labels) and revision is not None: + # load_vllm takes no revision, so vLLM fetches the default branch. Pinning + # only the config and tokenizer would mix two refs in one model. + logger.warning_once( + f"Unsloth: Ignoring revision = `{revision}` since vLLM loads weights from " + "the default branch. Use `fast_inference = False` to load a pinned revision." + ) + revision = None + tokenizer_revision = None + + if tokenizer_revision is None and tokenizer_name in (None, model_name): + # A direct FastLlamaModel call, or an architecture wrapper forwarding only + # `revision`, leaves this unset while the config and weights are pinned. + tokenizer_revision = revision token = hf_login(token) if model_patcher is None: @@ -2407,6 +2446,7 @@ def from_pretrained( model_name, token = token, attn_implementation = "sdpa", + revision = revision, ) _checkpoint_quant = getattr(_checkpoint_config, "quantization_config", None) if _checkpoint_quant is not None: @@ -2416,6 +2456,7 @@ def from_pretrained( model_name, token = token, attn_implementation = "sdpa", + revision = revision, ) model_config.model_name = model_name model_max_seq_length = model_config.max_position_embeddings @@ -2429,11 +2470,12 @@ def from_pretrained( preferred_attn_impl = resolve_attention_implementation(model_function, model_config) # Prefetch the repo (killable child) so the weight load is a cache hit. Runs after the - # AutoConfig/model-class check so an unsupported repo fails on its small config fetch. No - # revision: the load resolves model_name (maybe a remapped prequant repo) on its default branch. + # AutoConfig/model-class check so an unsupported repo fails on its small config fetch. + # Warm the same revision the load uses, or the repo downloads twice. _prefetched = maybe_prefetch_hf_snapshot( model_name, token = token, + revision = revision, cache_dir = kwargs.get("cache_dir"), local_files_only = kwargs.get("local_files_only", False), # Skip the warm only for a real vLLM load; a num_labels classification load still goes @@ -2493,6 +2535,7 @@ def from_pretrained( cache_dir = _tokenizer_cache_dir, local_files_only = kwargs.get("local_files_only", False), tokenizer_only = True, + revision = tokenizer_revision, ) has_rope_scaling = False @@ -2626,6 +2669,7 @@ def from_pretrained( token = token, trust_remote_code = trust_remote_code, attn_implementation = preferred_attn_impl, + revision = revision, **kwargs, ) # Defensive: ensure the task head is in a floating dtype, guarding @@ -2654,8 +2698,8 @@ def from_pretrained( model_name, local_files_only = kwargs.get("local_files_only", False), token = token, - # Weights load from the default branch (revision not forwarded), so read scales from there too. - revision = None, + # Read scales from the same revision as the weights. + revision = revision, subfolder = kwargs.get("subfolder"), cache_dir = kwargs.get("cache_dir"), variant = kwargs.get("variant"), @@ -2674,6 +2718,7 @@ def from_pretrained( token = token, trust_remote_code = trust_remote_code, attn_implementation = preferred_attn_impl, + revision = revision, **kwargs, ) else: @@ -2686,6 +2731,7 @@ def from_pretrained( max_position_embeddings = max_position_embeddings, trust_remote_code = trust_remote_code, attn_implementation = preferred_attn_impl, + revision = revision, **kwargs, ) # Attach dispatch hooks for bnb multi-device loads. @@ -2704,8 +2750,8 @@ def from_pretrained( model_name, local_files_only = kwargs.get("local_files_only", False), token = token, - # Weights load from the default branch (revision not forwarded), so read scales from there too. - revision = None, + # Read scales from the same revision as the weights. + revision = revision, subfolder = kwargs.get("subfolder"), cache_dir = kwargs.get("cache_dir"), variant = kwargs.get("variant"), @@ -2781,6 +2827,7 @@ def from_pretrained( token = token, trust_remote_code = trust_remote_code, fix_tokenizer = fix_tokenizer, + revision = tokenizer_revision, **_tokenizer_cache_kwargs, ) diff --git a/unsloth/models/loader.py b/unsloth/models/loader.py index ec979f811d..533e74681e 100644 --- a/unsloth/models/loader.py +++ b/unsloth/models/loader.py @@ -27,7 +27,7 @@ DISABLE_SDPA_MODEL_NAMES, ) from .granite import FastGraniteModel -from .llama import FastLlamaModel, logger +from .llama import FastLlamaModel, logger, _vllm_will_load_weights from .mistral import FastMistralModel from .qwen2 import FastQwen2Model from .qwen3 import FastQwen3Model @@ -149,6 +149,59 @@ def _strip_unsloth_bnb_4bit_suffix(model_name: str) -> str: return s +def _revision_for_resolved_repo( + revision, + model_name, + old_model_name, + mapper_moved_name = False, +): + """Drop `revision` once the requested repo has been remapped to another one. + + A revision names a branch/tag/SHA on the repo the caller asked for, but from_pretrained + may resolve model_name to a different repo (a pre-quantized mirror, an fp8 temp dir, a + ModelScope snapshot, a -bnb-4bit strip), where that ref does not exist. Only the mapper + substitution answers to use_exact_model_name, so only suggest it when it would help. + """ + if revision is None or model_name == old_model_name: + return revision + remedy = ( + " Pass `use_exact_model_name = True` to load your repo as-is." if mapper_moved_name else "" + ) + logger.warning_once( + f"Unsloth: Ignoring revision = `{revision}` since `{old_model_name}` resolved to " + f"`{model_name}`, which does not have that revision.{remedy}" + ) + return None + + +def _revision_for_tokenizer_repo( + tokenizer_name, + model_name, + old_model_name, + revision, + model_revision, + is_peft = False, +): + """Pick the revision for whichever repo the tokenizer is actually read from. + + It is not always the base model's: an adapter-hosted tokenizer is a separate repo with + its own history, so it keeps the caller's ref even though the base model does not. An + unset tokenizer_name follows the resolved model_name and so takes whatever the model + load itself uses (None on a PEFT load, whose ref belongs to the adapter). + + On a plain load the tokenizer belongs to the same model as the weights, so it follows + model_revision even when the caller named its repo directly: a remap has already dropped + the pin off the weights, and a pinned tokenizer beside a mirror's default-branch weights + is the ref mismatch this whole gate exists to avoid. + """ + repo = tokenizer_name if tokenizer_name else model_name + if is_peft and repo == old_model_name: + return revision + if repo == model_name: + return model_revision + return None + + def _config_get( config, field_name, @@ -529,7 +582,10 @@ def from_pretrained( load_in_8bit, load_in_16bit, ) - model_name = _offline_quantize_to_fp8(model_name, fp8_mode, text_only = text_only) + # Still the caller's repo here, so their ref is the one to quantize from. + model_name = _offline_quantize_to_fp8( + model_name, fp8_mode, text_only = text_only, revision = revision + ) else: assert new_model_name is not None model_name = new_model_name @@ -537,6 +593,8 @@ def from_pretrained( # on-the-fly quantization to avoid double quantization if load_in_fp8 != False and new_model_name != old_model_name: load_in_fp8 = False + # Only this block honours use_exact_model_name; the transforms below do not. + mapper_moved_name = model_name != old_model_name # Check if pre-quantized models are allowed # AMD Instinct GPUs need blocksize = 128 on bitsandbytes < 0.49.2 (our pre-quants use blocksize = 64) @@ -557,6 +615,28 @@ def from_pretrained( from modelscope import snapshot_download model_name = snapshot_download(model_name) + # Gate before the probe below, or a pinned 4bit load fails against the mirror + # instead of warning. Kept separate: `revision` still names the adapter repo. + base_revision = _revision_for_resolved_repo( + revision, model_name, old_model_name, mapper_moved_name + ) + # The PeftConfig probe below reads the adapter repo, which peft loads in-process, so + # it keeps the ref even when vLLM takes the base model's away just after. + adapter_revision = base_revision + # vLLM takes no revision and fetches the default branch, so this pin is already dead + # for the weights. Drop it before the probe: model_types picks the architecture class + # off that config, and reading it at a ref the weights will not be at dispatches the + # wrong one. The predicate lives in llama.py, which also falls back in-process on + # pre-Volta GPUs and for a num_labels load; both of those can still honour the pin. + if base_revision is not None and _vllm_will_load_weights( + fast_inference, kwargs.get("num_labels") + ): + logger.warning_once( + f"Unsloth: Ignoring revision = `{base_revision}` since vLLM loads weights " + "from the default branch. Use `fast_inference = False` to load a pinned revision." + ) + base_revision = None + # First check if it's a normal model via AutoConfig from huggingface_hub.utils import ( disable_progress_bars, @@ -579,7 +659,7 @@ def from_pretrained( model_config = AutoConfig.from_pretrained( model_name, token = token, - revision = revision, + revision = base_revision, trust_remote_code = trust_remote_code, local_files_only = local_files_only, ) @@ -606,7 +686,7 @@ def from_pretrained( peft_config = PeftConfig.from_pretrained( model_name, token = token, - revision = revision, + revision = adapter_revision, trust_remote_code = trust_remote_code, local_files_only = local_files_only, ) @@ -850,6 +930,15 @@ def from_pretrained( if fast_inference: fast_inference, model_name = fast_inference_setup(model_name, model_config) + # model_name can move once more here. Skip for PEFT: model_name is then the base + # model, and `revision` names the adapter that PeftModel.from_pretrained loads below. + if not is_peft: + base_revision = _revision_for_resolved_repo( + base_revision, model_name, old_model_name, mapper_moved_name + ) + # On a PEFT load model_name is the base model, which the caller's ref is not for. + model_revision = base_revision if not is_peft else None + load_in_4bit_kwargs = load_in_4bit load_in_8bit_kwargs = load_in_8bit if quantization_config is not None and not fast_inference: @@ -876,7 +965,10 @@ def from_pretrained( model_patcher = dispatch_model, tokenizer_name = tokenizer_name, trust_remote_code = trust_remote_code, - revision = revision if not is_peft else None, + revision = model_revision, + tokenizer_revision = _revision_for_tokenizer_repo( + tokenizer_name, model_name, old_model_name, revision, model_revision, is_peft + ), fast_inference = fast_inference, gpu_memory_utilization = gpu_memory_utilization, float8_kv_cache = float8_kv_cache, @@ -1247,7 +1339,10 @@ def from_pretrained( load_in_8bit, load_in_16bit, ) - model_name = _offline_quantize_to_fp8(model_name, fp8_mode, text_only = text_only) + # Still the caller's repo here, so their ref is the one to quantize from. + model_name = _offline_quantize_to_fp8( + model_name, fp8_mode, text_only = text_only, revision = revision + ) else: assert new_model_name is not None model_name = new_model_name @@ -1255,6 +1350,8 @@ def from_pretrained( # on-the-fly quantization to avoid double quantization if load_in_fp8 != False and new_model_name != old_model_name: load_in_fp8 = False + # Only this block honours use_exact_model_name; the transforms below do not. + mapper_moved_name = model_name != old_model_name # Check if pre-quantized models are allowed # AMD Instinct GPUs need blocksize = 128 on bitsandbytes < 0.49.2 (our pre-quants use blocksize = 64) @@ -1276,6 +1373,26 @@ def from_pretrained( from modelscope import snapshot_download model_name = snapshot_download(model_name) + # Gate before the probe below, or a pinned 4bit load fails against the mirror + # instead of warning. Kept separate: `revision` still names the adapter repo. + base_revision = _revision_for_resolved_repo( + revision, model_name, old_model_name, mapper_moved_name + ) + # The PeftConfig probe below reads the adapter repo, which peft loads in-process, so + # it keeps the ref even when vLLM takes the base model's away just after. + adapter_revision = base_revision + # vLLM takes no revision and fetches the default branch, so this pin is already dead + # for the weights. Drop it here rather than at the dispatch: model_types, auto_model + # and the text-only decision all come off the config probed below, and reading that + # at a ref the weights will not be at picks the dispatch for the wrong model. Same + # predicate FastBaseModel uses, so its own guard is a no-op on this path. + if base_revision is not None and fast_inference and is_vLLM_available(): + logger.warning_once( + f"Unsloth: Ignoring revision = `{base_revision}` since vLLM loads weights " + "from the default branch. Use `fast_inference = False` to load a pinned revision." + ) + base_revision = None + # First check if it's a normal model via AutoConfig from huggingface_hub.utils import ( disable_progress_bars, @@ -1309,7 +1426,7 @@ def _dispatch_diffusion(): token = token, device_map = device_map, trust_remote_code = trust_remote_code, - revision = revision, + revision = base_revision, **kwargs, ) @@ -1319,7 +1436,7 @@ def _dispatch_diffusion(): model_config = AutoConfig.from_pretrained( model_name, token = token, - revision = revision, + revision = base_revision, trust_remote_code = trust_remote_code, local_files_only = local_files_only, ) @@ -1352,7 +1469,7 @@ def _dispatch_diffusion(): peft_config = PeftConfig.from_pretrained( model_name, token = token, - revision = revision, + revision = adapter_revision, trust_remote_code = trust_remote_code, local_files_only = local_files_only, ) @@ -1798,6 +1915,15 @@ def _dispatch_diffusion(): load_in_4bit_kwargs = False load_in_8bit_kwargs = False + # FastBaseModel remaps again via fast_inference_setup. Skip for PEFT: model_name is + # then the base model, and `revision` names the adapter PeftModel loads below. + if not is_peft: + base_revision = _revision_for_resolved_repo( + base_revision, model_name, old_model_name, mapper_moved_name + ) + # On a PEFT load model_name is the base model, which the caller's ref is not for. + model_revision = base_revision if not is_peft else None + model, tokenizer = FastBaseModel.from_pretrained( model_name = model_name, max_seq_length = max_seq_length, @@ -1809,7 +1935,10 @@ def _dispatch_diffusion(): token = token, device_map = device_map, trust_remote_code = trust_remote_code, - revision = revision if not is_peft else None, + revision = model_revision, + tokenizer_revision = _revision_for_tokenizer_repo( + tokenizer_name, model_name, old_model_name, revision, model_revision, is_peft + ), model_types = model_types, tokenizer_name = tokenizer_name, auto_model = auto_model, diff --git a/unsloth/models/loader_utils.py b/unsloth/models/loader_utils.py index e8d0fa657c..8ea3ed8ad4 100644 --- a/unsloth/models/loader_utils.py +++ b/unsloth/models/loader_utils.py @@ -13,6 +13,7 @@ # limitations under the License. from ..device_type import DEVICE_TYPE_TORCH +import hashlib import importlib import os import torch @@ -314,11 +315,16 @@ def _offline_quantize_to_fp8( fp8_mode: str, *, text_only: bool = False, + revision: str = None, ) -> str: """Quantize the model to fp8 via torchao, save to a temp dir, return its path. For vllm >= 0.12.0, prefer dynamic quantization in vllm instead (via hf_overrides={"quantization_config_file": "torchao_config.json"}). + + The caller's revision has to reach the source loads, and the cache name has to name it + too: the returned path replaces model_name, so the revision gate downstream drops the + pin, and two refs of one repo would otherwise share (and reuse) a single artifact. """ from transformers import ( AutoModelForCausalLM, @@ -329,7 +335,7 @@ def _offline_quantize_to_fp8( AutoConfig, ) - config = AutoConfig.from_pretrained(model_name) + config = AutoConfig.from_pretrained(model_name, revision = revision) is_vlm = any( x.endswith(("ForConditionalGeneration", "ForVisionText2Text")) for x in (getattr(config, "architectures", None) or []) @@ -356,6 +362,13 @@ def _offline_quantize_to_fp8( temp_dir = tempfile.gettempdir() # Cache text-only and full-VLM artifacts separately so neither reuses the other. #5816 cache_name = model_name.split("/")[-1] + "-fp8-" + fp8_mode + if revision is not None: + # Slashes and dots would escape the temp dir, so the readable half is sanitized and + # therefore lossy: `release/v1` and `release.v1` collapse to one name. A digest of + # the raw ref rides along so two refs never share (and silently reuse) an artifact. + digest = hashlib.sha256(revision.encode("utf-8")).hexdigest()[:12] + readable = re.sub(r"[^0-9A-Za-z_-]", "_", revision)[:40] + cache_name += "-rev-" + readable + "-" + digest if text_config is not None: cache_name += "-text-only" new_model_name = os.path.join(temp_dir, cache_name) @@ -375,9 +388,10 @@ def _offline_quantize_to_fp8( model = auto_model.from_pretrained( model_name, config = config, + revision = revision, **load_kwargs, ) - tokenizer = auto_processor.from_pretrained(model_name) + tokenizer = auto_processor.from_pretrained(model_name, revision = revision) model.save_pretrained(new_model_name, safe_serialization = False) del model for _ in range(2): @@ -884,6 +898,37 @@ def _get_effective_local_files_only(kwargs): # The load's cache_dir travels with it too: saving derives one from HF_HUB_CACHE / # HF_HOME, which does not see a caller-supplied cache. _LOADED_CACHE_DIR_ATTR = "_unsloth_loaded_cache_dir" +# So does the ref it was read at. Saving restores sentencepiece assets from +# tokenizer.name_or_path, which names the repo but not the branch, so without this stamp a +# merged export copies the default branch's tokenizer.model beside pinned metadata, or +# misses the file when it only exists on the pinned ref. +_LOADED_REVISION_ATTR = "_unsloth_loaded_revision" + + +def _mark_loaded_revision(result, revision): + """Stamp the ref a tokenizer/processor was loaded at onto the returned objects.""" + if revision is None: + return result + for obj in result if isinstance(result, (tuple, list)) else (result,): + try: + targets = (obj, getattr(obj, "tokenizer", None)) + except Exception: + targets = (obj,) + for target in targets: + if target is None: + continue + # Objects that reject new attributes (__slots__) are skipped. + try: + setattr(target, _LOADED_REVISION_ATTR, str(revision)) + except Exception: + pass + return result + + +def _tokenizer_revision(tokenizer): + """The ref this tokenizer was loaded at, or None for the default branch.""" + tokenizer = tokenizer.tokenizer if hasattr(tokenizer, "tokenizer") else tokenizer + return getattr(tokenizer, _LOADED_REVISION_ATTR, None) def _mark_loaded_local_files_only(result, cache_dir = None): @@ -1214,6 +1259,7 @@ def _resolve_hub_repo_local_dir( *, token = None, cache_dir = None, + revision = None, # Default closed: a "resolve local dir" helper must not download. False here # means five filenames each retried with backoff before it gives up. local_files_only = True, @@ -1249,6 +1295,7 @@ def _resolve_hub_repo_local_dir( token = token, cache_dir = cache_dir, local_files_only = local_files_only, + revision = revision, ) if path and os.path.isfile(path): return os.path.dirname(path) @@ -1264,6 +1311,7 @@ def _resolve_hub_repo_cached_file( token = None, cache_dir = None, local_files_only = True, + revision = None, ): """Return a cached file path under a Hub snapshot, or None if absent.""" local_dir = _resolve_hub_repo_local_dir( @@ -1271,6 +1319,7 @@ def _resolve_hub_repo_cached_file( token = token, cache_dir = cache_dir, local_files_only = local_files_only, + revision = revision, filenames = (filename,), ) if local_dir is None: @@ -1286,6 +1335,7 @@ def _hub_repo_or_local_path( cache_dir = None, local_files_only = False, filenames = None, + revision = None, ): """Prefer a cached snapshot path over a Hub repo id when offline or ``local_files_only``.""" if isinstance(repo_id, str) and os.path.isdir(repo_id): @@ -1298,6 +1348,7 @@ def _hub_repo_or_local_path( token = token, cache_dir = cache_dir, local_files_only = True, + revision = revision, filenames = filenames or ( "tokenizer_config.json", @@ -1318,6 +1369,7 @@ def _load_pretrained_tokenizer_fast( trust_remote_code = False, cache_dir = None, local_files_only = False, + revision = None, ): """Load ``PreTrainedTokenizerFast`` without Hub metadata probes when cached/offline. @@ -1331,6 +1383,7 @@ def _load_pretrained_tokenizer_fast( token = token, cache_dir = cache_dir, local_files_only = lfo, + revision = revision, filenames = ( "tokenizer_config.json", "tokenizer.json", @@ -1344,6 +1397,7 @@ def _load_pretrained_tokenizer_fast( trust_remote_code = trust_remote_code, cache_dir = cache_dir, local_files_only = lfo, + revision = revision, ) diff --git a/unsloth/models/vision.py b/unsloth/models/vision.py index 13cd582437..6ee89f06e9 100644 --- a/unsloth/models/vision.py +++ b/unsloth/models/vision.py @@ -636,6 +636,7 @@ def unsloth_base_fast_generate(self, *args, **kwargs): _hub_repo_or_local_path, _is_offline_related_error, _load_pretrained_tokenizer_fast, + _mark_loaded_revision, _offline_aware_load, ) @@ -664,6 +665,7 @@ def _construct_vlm_processor_fallback( trust_remote_code, cache_dir = None, local_files_only = False, + revision = None, ): """Build a VLM processor manually when AutoProcessor.from_pretrained fails (some VLMs have unresolvable tokenizer_class entries): load the image processor + tokenizer @@ -680,6 +682,7 @@ def _construct_vlm_processor_fallback( token = token, cache_dir = cache_dir, local_files_only = local_files_only, + revision = revision, ) # Load image processor image_processor = AutoImageProcessor.from_pretrained( @@ -688,6 +691,7 @@ def _construct_vlm_processor_fallback( trust_remote_code = trust_remote_code, cache_dir = cache_dir, local_files_only = local_files_only, + revision = revision, ) # Load tokenizer via PreTrainedTokenizerFast (bypasses tokenizer_class check). # Resolve the cached snapshot first so transformers does not call model_info (#7481). @@ -698,6 +702,7 @@ def _construct_vlm_processor_fallback( trust_remote_code = trust_remote_code, cache_dir = cache_dir, local_files_only = local_files_only, + revision = revision, ) # Read tokenizer_config.json for special tokens: prefer the local file (offline # / local checkpoint dir), else hf_hub_download with local_files_only forwarded. @@ -723,6 +728,7 @@ def _construct_vlm_processor_fallback( "tokenizer_config.json", token = token, cache_dir = cache_dir, + revision = revision, local_files_only = local_files_only, ) with open(config_path, "r", encoding = "utf-8") as f: @@ -756,6 +762,7 @@ def _construct_vlm_processor_fallback( trust_remote_code = trust_remote_code, cache_dir = cache_dir, local_files_only = local_files_only, + revision = revision, ) proc_class_name = PROCESSOR_MAPPING_NAMES.get(config.model_type) except Exception as _e: @@ -862,6 +869,22 @@ def from_pretrained( if os.environ.get("UNSLOTH_MODEL_NAME", "") == "": os.environ["UNSLOTH_MODEL_NAME"] = model_name.lower() + # Read revision from kwargs, not the signature: the weight load below forwards + # **kwargs, so binding it would drop it there. Pin its repo before any remap. + _revision = kwargs.get("revision") + _tokenizer_revision_arg = kwargs.pop("tokenizer_revision", None) + if _revision is not None and fast_inference and is_vLLM_available(): + # load_vllm takes no revision, so vLLM fetches the default branch. Pinning only + # the config and tokenizer would mix two refs in one model. + logger.warning_once( + f"Unsloth: Ignoring revision = `{_revision}` since vLLM loads weights from " + "the default branch. Use `fast_inference = False` to load a pinned revision." + ) + _revision = None + _tokenizer_revision_arg = None + kwargs.pop("revision", None) + _revision_repo = model_name + # Resolve text-only before the is_vlm / vLLM checks so is_vlm stays consistent; # skip the vision tower only for families with their own text decoder (Gemma 3). #5816 if text_only and auto_config is None: @@ -870,6 +893,7 @@ def from_pretrained( token = token, trust_remote_code = trust_remote_code, local_files_only = local_files_only, + revision = _revision, ) if text_only and hasattr(auto_config, "vision_config"): parent_config = auto_config @@ -1027,6 +1051,7 @@ def from_pretrained( token = token, trust_remote_code = trust_remote_code, local_files_only = local_files_only, + revision = _revision, ) model_class = resolve_model_class(auto_model, auto_config) attn_impl = resolve_attention_implementation( @@ -1108,6 +1133,7 @@ def from_pretrained( and _tokenizer_repo != model_name ) if _warm_tokenizer_repo: + # No revision: this only runs when the repo differs from the revision's repo. maybe_prefetch_hf_snapshot( _tokenizer_repo, token = token, @@ -1182,6 +1208,7 @@ def from_pretrained( token = token, trust_remote_code = trust_remote_code, local_files_only = local_files_only, + revision = _revision, ) if hasattr(auto_config, "quantization_config"): from transformers.quantizers.auto import ( @@ -1236,6 +1263,7 @@ def from_pretrained( token = token, trust_remote_code = trust_remote_code, local_files_only = local_files_only, + revision = _revision, ) _set_attn_impl(auto_config, config_attn_impl) model_config = auto_config @@ -1433,6 +1461,11 @@ def from_pretrained( # Counteract saved tokenizers tokenizer_name = model_name if tokenizer_name is None else tokenizer_name + # Resolved by the loader, which knows whether the tokenizer repo is the caller's + # (a PEFT adapter) or the resolved base model. Falls back for a direct call. + _tokenizer_revision = _tokenizer_revision_arg + if _tokenizer_revision is None and tokenizer_name == _revision_repo: + _tokenizer_revision = _revision # On the vLLM path the tokenizer warm was deferred (fast_inference_setup may remap model_name). # Warm the now-final tokenizer repo so the load below hits the cache (a cached/local repo is a no-op). @@ -1440,7 +1473,8 @@ def from_pretrained( maybe_prefetch_hf_snapshot( tokenizer_name, token = token, - revision = kwargs.get("revision"), + # Match the tokenizer load below, which only pins its own repo. + revision = _tokenizer_revision, cache_dir = kwargs.get("cache_dir"), local_files_only = kwargs.get("local_files_only", False), tokenizer_only = True, @@ -1485,6 +1519,7 @@ def _acquire_processor(lfo): trust_remote_code = trust_remote_code, cache_dir = kwargs.get("cache_dir"), local_files_only = lfo, + revision = _tokenizer_revision, ) except Exception as _e: _tok = None @@ -1498,6 +1533,7 @@ def _acquire_processor(lfo): trust_remote_code = trust_remote_code, cache_dir = kwargs.get("cache_dir"), local_files_only = lfo, + revision = _tokenizer_revision, ) except Exception as _e: _err = _e @@ -1509,6 +1545,7 @@ def _acquire_processor(lfo): trust_remote_code = trust_remote_code, cache_dir = kwargs.get("cache_dir"), local_files_only = lfo, + revision = _tokenizer_revision, ) except Exception: # Swallow so the manual fallback / entry-point retry can run. @@ -1528,6 +1565,7 @@ def _acquire_processor(lfo): trust_remote_code, cache_dir = kwargs.get("cache_dir"), local_files_only = lfo, + revision = _tokenizer_revision, ) except Exception as _fe: _fallback, _fb_err = None, _fe @@ -1614,6 +1652,7 @@ def _is_degraded_vlm(_t): trust_remote_code = trust_remote_code, cache_dir = kwargs.get("cache_dir"), local_files_only = local_files_only, + revision = _tokenizer_revision, ) model, _fallback_tok = patch_tokenizer(model, _fallback_tok) # Re-attach as processor wrapper if original was a processor @@ -1641,6 +1680,7 @@ def _last_resort_tokenizer(lfo): token = token, cache_dir = kwargs.get("cache_dir"), local_files_only = lfo, + revision = _tokenizer_revision, ) try: return _AutoTokenizer.from_pretrained( @@ -1650,6 +1690,7 @@ def _last_resort_tokenizer(lfo): trust_remote_code = trust_remote_code, cache_dir = kwargs.get("cache_dir"), local_files_only = lfo, + revision = _tokenizer_revision, ) except Exception: return _load_pretrained_tokenizer_fast( @@ -1659,6 +1700,7 @@ def _last_resort_tokenizer(lfo): trust_remote_code = trust_remote_code, cache_dir = kwargs.get("cache_dir"), local_files_only = lfo, + revision = _tokenizer_revision, ) _last_resort_err = None @@ -1734,6 +1776,10 @@ def _last_resort_tokenizer(lfo): for _ in range(3): gc.collect() clean_gpu_cache() + # Saving restores sentencepiece assets from the repo name alone, which does not + # carry the branch this was read at. Stamped here rather than at each of the + # processor branches above, so a patch fallback cannot lose it. + _mark_loaded_revision(tokenizer, _tokenizer_revision) return model, tokenizer @staticmethod diff --git a/unsloth/save.py b/unsloth/save.py index dd0fceb235..99afc7db76 100644 --- a/unsloth/save.py +++ b/unsloth/save.py @@ -68,6 +68,7 @@ class Peft_Linear4bit: get_model_name, _resolve_hub_repo_cached_file, _tokenizer_cache_dir, + _tokenizer_revision, _tokenizer_wants_local_only, ) from .models._utils import _convert_torchao_model @@ -460,8 +461,11 @@ def _has_tokenizer_model(tokenizer, token = None): return False if os.path.isdir(source): return os.path.isfile(os.path.join(source, "tokenizer.model")) - if source in _TOKENIZER_MODEL_CACHE: - return _TOKENIZER_MODEL_CACHE[source] + # Refs of one repo can differ in whether they ship the asset, so memoize per ref. + revision = _tokenizer_revision(tokenizer) + cache_key = (source, revision) + if cache_key in _TOKENIZER_MODEL_CACHE: + return _TOKENIZER_MODEL_CACHE[cache_key] # Hub repo id: probe local cache before model_info (issue #7481). cache_dir = _tokenizer_cache_dir(tokenizer) or os.environ.get("HF_HUB_CACHE") @@ -476,23 +480,24 @@ def _has_tokenizer_model(tokenizer, token = None): token = token, local_files_only = True, cache_dir = cache_dir, + revision = revision, ) if cached_path is not None: - _TOKENIZER_MODEL_CACHE[source] = True + _TOKENIZER_MODEL_CACHE[cache_key] = True return True if _tokenizer_wants_local_only(tokenizer): return False try: - repo_info = HfApi(token = token).model_info(source, files_metadata = False) + repo_info = HfApi(token = token).model_info(source, revision = revision, files_metadata = False) except Exception: return False has_tokenizer_model = any( sibling.rfilename == "tokenizer.model" for sibling in (repo_info.siblings or []) ) - _TOKENIZER_MODEL_CACHE[source] = has_tokenizer_model + _TOKENIZER_MODEL_CACHE[cache_key] = has_tokenizer_model return has_tokenizer_model @@ -555,6 +560,7 @@ def _preserve_sentencepiece_tokenizer_assets( token = token, local_files_only = True, cache_dir = cache_dir, + revision = _tokenizer_revision(tokenizer), ) if cached_path is not None: downloaded_path = cached_path @@ -567,6 +573,7 @@ def _preserve_sentencepiece_tokenizer_assets( token = token, local_files_only = _tokenizer_wants_local_only(tokenizer), cache_dir = cache_dir, + revision = _tokenizer_revision(tokenizer), ) except Exception: downloaded_path = None diff --git a/unsloth/tokenizer_utils.py b/unsloth/tokenizer_utils.py index c7f61288d5..9dd58484ca 100644 --- a/unsloth/tokenizer_utils.py +++ b/unsloth/tokenizer_utils.py @@ -568,6 +568,7 @@ def _load_correct_tokenizer( trust_remote_code = False, cache_dir = "huggingface_tokenizers_cache", fix_tokenizer = True, + revision = None, ): if IS_COLAB_ENVIRONMENT: cache_dir = cache_dir @@ -596,6 +597,7 @@ def _load_correct_tokenizer( legacy = False, from_slow = True, cache_dir = cache_dir, + revision = revision, ) except: slow_tokenizer = None @@ -614,6 +616,7 @@ def _load_correct_tokenizer( token = token, trust_remote_code = trust_remote_code, cache_dir = cache_dir, + revision = revision, ) if not fix_tokenizer or tokenizer_name.lower() in IGNORED_TOKENIZER_NAMES: @@ -669,6 +672,7 @@ def load_correct_tokenizer( trust_remote_code = False, cache_dir = "huggingface_tokenizers_cache", fix_tokenizer = True, + revision = None, ): tokenizer = _load_correct_tokenizer( tokenizer_name = tokenizer_name, @@ -678,6 +682,7 @@ def load_correct_tokenizer( trust_remote_code = trust_remote_code, cache_dir = cache_dir, fix_tokenizer = fix_tokenizer, + revision = revision, ) if fix_tokenizer: @@ -711,6 +716,12 @@ def load_correct_tokenizer( pass tokenizer.chat_template = chat_template + # Saving restores sentencepiece assets from the repo name alone, which does not carry + # the branch this was read at, so stamp it for the save path to find. Imported here: + # models.loader_utils pulls in models._utils, which imports this module at load time. + from .models.loader_utils import _mark_loaded_revision + + _mark_loaded_revision(tokenizer, revision) return tokenizer