diff --git a/.github/workflows/prek.yml b/.github/workflows/prek.yml index 17e6913da..07403f7cf 100644 --- a/.github/workflows/prek.yml +++ b/.github/workflows/prek.yml @@ -224,6 +224,7 @@ jobs: tests/unit/test_prefix_tree_grad_parity.py \ tests/unit/test_prefix_tree_packing.py \ tests/unit/test_trainer_rank_handoff_budget.py \ + tests/unit/test_qwen35_adapter_config.py \ tests/unit/test_trainer_rank_physical_reserve.py \ tests/unit/test_trainer_rank_validation.py \ tests/unit/test_trainer_rank_weird_shapes.py \ @@ -265,6 +266,7 @@ jobs: --ignore=tests/unit/test_prefix_tree_grad_parity.py \ --ignore=tests/unit/test_prefix_tree_packing.py \ --ignore=tests/unit/test_trainer_rank_handoff_budget.py \ + --ignore=tests/unit/test_qwen35_adapter_config.py \ --ignore=tests/unit/test_trainer_rank_physical_reserve.py \ --ignore=tests/unit/test_trainer_rank_validation.py \ --ignore=tests/unit/test_trainer_rank_weird_shapes.py \ diff --git a/src/art/megatron/model_support/handlers/qwen3_5.py b/src/art/megatron/model_support/handlers/qwen3_5.py index 309046311..8015da2e1 100644 --- a/src/art/megatron/model_support/handlers/qwen3_5.py +++ b/src/art/megatron/model_support/handlers/qwen3_5.py @@ -509,11 +509,12 @@ def _is_self_attn_q_proj_lora_b(key: str) -> bool: @lru_cache(maxsize=8) -def _qwen35_text_config(base_model_name_or_path: str) -> Any: +def _qwen35_text_config(base_model_name_or_path: str, revision: str | None) -> Any: from transformers import AutoConfig config = AutoConfig.from_pretrained( base_model_name_or_path, + revision=revision, local_files_only=True, trust_remote_code=True, ) @@ -521,19 +522,34 @@ def _qwen35_text_config(base_model_name_or_path: str) -> Any: def _qwen35_attention_dims(adapter_config: dict[str, Any]) -> tuple[int, int, int]: - num_heads = adapter_config.get("num_attention_heads") - num_groups = adapter_config.get("num_key_value_heads") - head_dim = adapter_config.get("head_dim") + dims = { + key: adapter_config.get(key) + for key in ("num_attention_heads", "num_key_value_heads", "head_dim") + } hidden_size = adapter_config.get("hidden_size") - if num_heads is None: + if None in dims.values(): + # Take each missing dimension from the base model's config rather than + # defaulting it: Qwen3.5 uses grouped queries and a head size that is + # not hidden_size / heads. base_model = adapter_config.get("base_model_name_or_path") if not base_model: raise RuntimeError("Qwen3.5 LoRA adapter config is missing base model path") - config = _qwen35_text_config(str(base_model)) - num_heads = getattr(config, "num_attention_heads") - num_groups = getattr(config, "num_key_value_heads", num_heads) - head_dim = getattr(config, "head_dim", None) - hidden_size = getattr(config, "hidden_size", None) + # Resolve the adapter's pinned snapshot, not whatever the name + # currently points at (an offline cache may hold only the pin). + revision = adapter_config.get("revision") or None + config = _qwen35_text_config( + str(base_model), None if revision is None else str(revision) + ) + for key, value in dims.items(): + if value is None: + dims[key] = getattr(config, key, None) + if hidden_size is None: + hidden_size = getattr(config, "hidden_size", None) + num_heads = dims["num_attention_heads"] + num_groups = dims["num_key_value_heads"] + head_dim = dims["head_dim"] + if num_heads is None: + raise RuntimeError("Qwen3.5 config is missing num_attention_heads") num_heads = int(num_heads) num_groups = int(num_groups if num_groups is not None else num_heads) if head_dim is None: diff --git a/src/art/megatron/model_support/lora_disk.py b/src/art/megatron/model_support/lora_disk.py index f0b01183b..dbf3f9846 100644 --- a/src/art/megatron/model_support/lora_disk.py +++ b/src/art/megatron/model_support/lora_disk.py @@ -19,6 +19,32 @@ safe_open = safetensors.safe_open +def model_attention_dimensions(provider: Any) -> dict[str, int]: + """The running model's attention shape, in adapter-config keys.""" + dimensions = { + "num_attention_heads": getattr(provider, "num_attention_heads", None), + "num_key_value_heads": getattr(provider, "num_query_groups", None), + "head_dim": getattr(provider, "kv_channels", None), + "hidden_size": getattr(provider, "hidden_size", None), + } + return {key: int(value) for key, value in dimensions.items() if value is not None} + + +def with_model_attention_dimensions( + adapter_config: dict[str, Any], provider: Any +) -> dict[str, Any]: + """Fill attention dimensions the adapter config omits or nulls from the model. + + Values the adapter sets win. The handler resolves anything still missing + from the base model's config at the adapter's revision. + """ + config = dict(adapter_config) + for key, value in model_attention_dimensions(provider).items(): + if config.get(key) is None: + config[key] = value + return config + + def _jsonable_config(value: Any) -> Any: if isinstance(value, dict): return {key: _jsonable_config(item) for key, item in value.items()} @@ -127,14 +153,24 @@ def load_lora_tensors_for_megatron( lora_path: str | Path, *, handler: ModelSupportHandler | None = None, + provider: Any = None, allow_unvalidated_arch: bool = False, ) -> dict[str, torch.Tensor]: + """Load an adapter in Megatron layout. + + With the running model's ``provider``, conversion uses its attention shape + wherever the adapter config lacks one, instead of looking the base model + up again by name. + """ resolved_handler = resolve_lora_handler( lora_path, handler, allow_unvalidated_arch=allow_unvalidated_arch, ) + adapter_config = load_adapter_config(lora_path) + if provider is not None: + adapter_config = with_model_attention_dimensions(adapter_config, provider) return resolved_handler.from_vllm_lora_tensors( load_vllm_lora_tensors(lora_path), - adapter_config=load_adapter_config(lora_path), + adapter_config=adapter_config, ) diff --git a/src/art/megatron/train.py b/src/art/megatron/train.py index cf7d9452a..e203bcdd5 100644 --- a/src/art/megatron/train.py +++ b/src/art/megatron/train.py @@ -1006,6 +1006,7 @@ def _prepare_rl_training_state( job.source_adapter_path, runtime.rank, handler=runtime.model_support_handler, + provider=runtime.provider, optimizer=runtime.optimizer, ) if runtime.optimizer is None: @@ -1069,10 +1070,13 @@ def _load_adapter_into_model( rank: int, *, handler: Any | None = None, + provider: Any = None, optimizer: Any | None = None, ) -> dict[str, torch.Tensor]: print0(rank, "Loading adapter model from", lora_path) - adapter_model = load_lora_tensors_for_megatron(lora_path, handler=handler) + adapter_model = load_lora_tensors_for_megatron( + lora_path, handler=handler, provider=provider + ) load_adapter_into_model( model_chunks, adapter_model, @@ -2319,6 +2323,7 @@ def _prepare_kl_reference_logprobs( ref_adapter_path, runtime.rank, handler=runtime.model_support_handler, + provider=runtime.provider, ) loaded_ref_adapter = True return _precompute_reference_logprobs( diff --git a/src/art/trainer_rank/_checkpoint.py b/src/art/trainer_rank/_checkpoint.py index 1c3b33d97..530a7f35e 100644 --- a/src/art/trainer_rank/_checkpoint.py +++ b/src/art/trainer_rank/_checkpoint.py @@ -1492,7 +1492,9 @@ def _load_adapter( ) loaded = load_lora_tensors_for_megatron( - source.path, handler=trainer.runtime.model_support_handler + source.path, + handler=trainer.runtime.model_support_handler, + provider=getattr(trainer.runtime, "provider", None), ) return {key: value for key, value in loaded.items() if key in set(keys)} safe_open = importlib.import_module("safetensors").safe_open diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index d4718cab0..b51d21605 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -2626,19 +2626,11 @@ def _validate_checkpoint_adapter_config( "adapter_config['base_model_name_or_path'] must be a string" ) if base_model.startswith(("Qwen/Qwen3.5-", "Qwen/Qwen3.6-", "Qwen/Qwen3.8-")): - dimensions = { - "num_attention_heads": getattr( - self.runtime.provider, "num_attention_heads", None - ), - "num_key_value_heads": getattr( - self.runtime.provider, "num_query_groups", None - ), - "head_dim": getattr(self.runtime.provider, "kv_channels", None), - "hidden_size": getattr(self.runtime.provider, "hidden_size", None), - } - for key, value in dimensions.items(): - if value is not None: - config[key] = int(value) + from art.megatron.model_support.lora_disk import ( + model_attention_dimensions, + ) + + config.update(model_attention_dimensions(self.runtime.provider)) if not isinstance(rank, int) or isinstance(rank, bool): raise TypeError("adapter_config['r'] must be an integer") if not isinstance(config_alpha_value, int | float) or isinstance( diff --git a/tests/integration/megatron/train_inf_mismatch/real_path.py b/tests/integration/megatron/train_inf_mismatch/real_path.py index e56ea0456..7eb69063b 100644 --- a/tests/integration/megatron/train_inf_mismatch/real_path.py +++ b/tests/integration/megatron/train_inf_mismatch/real_path.py @@ -1573,6 +1573,7 @@ def _configure_worker_bundle(bundle: Any) -> None: adapter_model = load_lora_tensors_for_megatron( str(adapter_path), handler=runtime.model_support_handler, + provider=runtime.provider, allow_unvalidated_arch=request.config.allow_unvalidated_arch, ) megatron_train.load_adapter_into_model( diff --git a/tests/unit/test_qwen35_adapter_config.py b/tests/unit/test_qwen35_adapter_config.py new file mode 100644 index 000000000..2880d073a --- /dev/null +++ b/tests/unit/test_qwen35_adapter_config.py @@ -0,0 +1,255 @@ +import json +from pathlib import Path +from types import SimpleNamespace +from typing import Any, cast + +import pytest +import torch + +qwen35 = pytest.importorskip("art.megatron.model_support.handlers.qwen3_5") + +REPO = "Qwen/Qwen3.6-35B-A3B" +PIN = "995ad96eacd98c81ed38be0c5b274b04031597b0" +OTHER = "1" * 40 +LAYER = "base_model.model.model.language_model.layers.0.self_attn.q_proj" + + +@pytest.fixture +def hub_cache(tmp_path, monkeypatch): + """An offline cache that holds only the snapshots a test adds.""" + import huggingface_hub.constants + + repo = tmp_path / "models--Qwen--Qwen3.6-35B-A3B" + + def add_snapshot(revision: str, groups: int, *, main: bool = False) -> None: + snapshot = repo / "snapshots" / revision + snapshot.mkdir(parents=True) + (snapshot / "config.json").write_text( + json.dumps( + { + "model_type": "llama", + "num_attention_heads": 16, + "num_key_value_heads": groups, + "head_dim": 256, + "hidden_size": 2048, + } + ) + ) + if main: + (repo / "refs").mkdir(exist_ok=True) + (repo / "refs" / "main").write_text(revision) + + monkeypatch.setattr(huggingface_hub.constants, "HF_HUB_CACHE", str(tmp_path)) + monkeypatch.setenv("HF_HUB_OFFLINE", "1") + qwen35._qwen35_text_config.cache_clear() + yield add_snapshot + qwen35._qwen35_text_config.cache_clear() + + +def _dims(revision: str | None = None) -> tuple[int, int, int]: + config: dict[str, object] = {"base_model_name_or_path": REPO} + if revision is not None: + config["revision"] = revision + return qwen35._qwen35_attention_dims(config) + + +def test_attention_dims_resolve_the_pinned_snapshot_offline(hub_cache): + hub_cache(PIN, groups=2) + assert _dims(PIN) == (16, 2, 256) + # By name alone, the lookup needs refs/main, which this cache lacks. + with pytest.raises(OSError): + _dims() + + +def test_a_different_main_revision_does_not_change_a_pinned_adapter(hub_cache): + hub_cache(PIN, groups=2) + hub_cache(OTHER, groups=4, main=True) + assert [_dims(PIN), _dims(OTHER), _dims(PIN)] == [ + (16, 2, 256), + (16, 4, 256), + (16, 2, 256), + ] + assert _dims() == (16, 4, 256) # unpinned adapters still follow main + assert _dims("") == (16, 4, 256) # an empty revision is unpinned + + +def test_missing_dimensions_come_from_the_pinned_config_not_defaults(hub_cache): + hub_cache(PIN, groups=2) + config = { + "base_model_name_or_path": REPO, + "revision": PIN, + "num_attention_heads": 16, + "hidden_size": 2048, + } + # Not one group per head, nor head_dim = hidden_size / heads = 128. + assert qwen35._qwen35_attention_dims(config) == (16, 2, 256) + + +def _write_adapter( + path: Path, rows: int, **dimensions: object +) -> dict[str, torch.Tensor]: + from art.megatron.model_support.lora_disk import save_vllm_lora_tensors + + tensors = { + f"{LAYER}.lora_A.weight": torch.randn(2, 32), + f"{LAYER}.lora_B.weight": torch.randn(rows, 2), + } + # As published for a pinned base model, without attention dimensions. + config = { + "base_model_name_or_path": REPO, + "revision": PIN, + "r": 2, + "lora_alpha": 32, + "target_modules": ["q_proj"], + "art_lora_format": "vllm", + **dimensions, + } + save_vllm_lora_tensors(path, tensors, config) + return tensors + + +# Four heads in two query groups, with head_dim 8: the runtime's shape. The +# hidden size is not heads x head_dim, as in Qwen3.5, so deriving head_dim from +# it would be caught. +PROVIDER = SimpleNamespace( + num_attention_heads=4, num_query_groups=2, kv_channels=8, hidden_size=48 +) +ROWS = 2 * 2 * 2 * 8 # groups x (query + gate) x heads per group x head_dim +ART_KEY = f"{LAYER}.lora_B.weight".replace(".language_model.layers.", ".layers.") + + +def _expected(tensors: dict[str, torch.Tensor]) -> torch.Tensor: + return qwen35._qwen35_q_proj_lora_b_from_vllm( + tensors[f"{LAYER}.lora_B.weight"], + {"num_attention_heads": 4, "num_key_value_heads": 2, "head_dim": 8}, + ) + + +def _forbid_lookup(monkeypatch) -> None: + def lookup(*_args): + raise AssertionError("adapter conversion looked the base model up again") + + monkeypatch.setattr(qwen35, "_qwen35_text_config", lookup) + + +@pytest.mark.parametrize("handler", ["QWEN3_5_DENSE_HANDLER", "QWEN3_5_MOE_HANDLER"]) +def test_checkpoint_load_converts_with_the_running_models_attention_shape( + tmp_path, monkeypatch, handler +): + from art.trainer_rank import _checkpoint + + _forbid_lookup(monkeypatch) + trainer = SimpleNamespace( + runtime=SimpleNamespace( + provider=PROVIDER, model_support_handler=getattr(qwen35, handler) + ) + ) + tensors = _write_adapter(tmp_path, ROWS) + loaded = _checkpoint._load_adapter( + cast(Any, trainer), + cast(Any, SimpleNamespace(manifest=None, path=tmp_path)), + [ART_KEY], + ) + torch.testing.assert_close(loaded[ART_KEY], _expected(tensors)) + + +def test_null_adapter_dimensions_are_filled_from_the_running_model( + tmp_path, monkeypatch +): + from art.megatron.model_support.lora_disk import load_lora_tensors_for_megatron + + _forbid_lookup(monkeypatch) + # A null group count must not fall back to one group per head. + tensors = _write_adapter(tmp_path, ROWS, num_key_value_heads=None, head_dim=None) + loaded = load_lora_tensors_for_megatron( + tmp_path, handler=qwen35.QWEN3_5_MOE_HANDLER, provider=PROVIDER + ) + torch.testing.assert_close(loaded[ART_KEY], _expected(tensors)) + + +@pytest.mark.parametrize("handler", ["QWEN3_5_DENSE_HANDLER", "QWEN3_5_MOE_HANDLER"]) +def test_adapter_and_model_dimensions_combine(tmp_path, monkeypatch, handler): + from art.megatron.model_support.lora_disk import load_lora_tensors_for_megatron + + _forbid_lookup(monkeypatch) + # The adapter has heads and head size, the model has query groups but no + # head size: together they are complete. + tensors = _write_adapter(tmp_path, ROWS, num_attention_heads=4, head_dim=8) + provider = SimpleNamespace(num_attention_heads=4, num_query_groups=2) + loaded = load_lora_tensors_for_megatron( + tmp_path, handler=getattr(qwen35, handler), provider=provider + ) + torch.testing.assert_close(loaded[ART_KEY], _expected(tensors)) + + +def test_adapter_dimensions_take_precedence_over_the_running_model(): + from art.megatron.model_support.lora_disk import with_model_attention_dimensions + + config = with_model_attention_dimensions( + {"revision": PIN, "num_attention_heads": 1, "num_key_value_heads": None}, + PROVIDER, + ) + assert config == { + "revision": PIN, + "num_attention_heads": 1, + "num_key_value_heads": 2, + "head_dim": 8, + "hidden_size": 48, + } + + +@pytest.mark.parametrize( + "provider", + [ + # Query groups must not default to one per head. + SimpleNamespace(num_attention_heads=4, kv_channels=8, hidden_size=48), + # The head size must not be derived from the hidden size. + SimpleNamespace(num_attention_heads=4, num_query_groups=2, hidden_size=48), + ], + ids=["no-query-groups", "no-head-size"], +) +def test_dimensions_the_model_lacks_come_from_the_pinned_lookup( + tmp_path, monkeypatch, provider +): + from art.megatron.model_support.lora_disk import load_lora_tensors_for_megatron + + looked_up = [] + + def lookup(name, revision): + looked_up.append((name, revision)) + return SimpleNamespace(num_attention_heads=4, num_key_value_heads=2, head_dim=8) + + monkeypatch.setattr(qwen35, "_qwen35_text_config", lookup) + tensors = _write_adapter(tmp_path, ROWS) + loaded = load_lora_tensors_for_megatron( + tmp_path, handler=qwen35.QWEN3_5_MOE_HANDLER, provider=provider + ) + assert looked_up == [(REPO, PIN)] + torch.testing.assert_close(loaded[ART_KEY], _expected(tensors)) + + +def test_export_of_a_pinned_adapter_resolves_its_revision(hub_cache): + hub_cache(PIN, groups=2) + tensor = torch.randn(2 * 2 * 8 * 256, 2) # 16 heads in 2 groups, head_dim 256 + config = {"base_model_name_or_path": REPO, "revision": PIN} + art_key = "base_model.model.model.layers.0.self_attn.q_proj.lora_B.weight" + exported, _ = qwen35.QWEN3_5_MOE_HANDLER.to_vllm_lora_tensors( + {art_key: tensor}, adapter_config=config + ) + expected = qwen35._qwen35_q_proj_lora_b_to_vllm( + tensor, {"num_attention_heads": 16, "num_key_value_heads": 2, "head_dim": 256} + ) + torch.testing.assert_close(exported[f"{LAYER}.lora_B.weight"], expected) + + +def test_megatron_service_adapter_load_passes_the_running_model(monkeypatch): + train = pytest.importorskip("art.megatron.train") + seen = [] + monkeypatch.setattr( + train, + "load_lora_tensors_for_megatron", + lambda path, **kwargs: seen.append(kwargs) or {}, + ) + monkeypatch.setattr(train, "load_adapter_into_model", lambda *a, **k: None) + train._load_adapter_into_model([], "adapter", 0, handler=None, provider=PROVIDER) + assert seen[0]["provider"] is PROVIDER