Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions .github/workflows/prek.yml
Original file line number Diff line number Diff line change
Expand Up @@ -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 \
Expand Down Expand Up @@ -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 \
Expand Down
36 changes: 26 additions & 10 deletions src/art/megatron/model_support/handlers/qwen3_5.py
Original file line number Diff line number Diff line change
Expand Up @@ -509,31 +509,47 @@ 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,
)
return getattr(config, "text_config", config)


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:
Expand Down
38 changes: 37 additions & 1 deletion src/art/megatron/model_support/lora_disk.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()}
Expand Down Expand Up @@ -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,
)
7 changes: 6 additions & 1 deletion src/art/megatron/train.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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(
Expand Down
4 changes: 3 additions & 1 deletion src/art/trainer_rank/_checkpoint.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
18 changes: 5 additions & 13 deletions src/art/trainer_rank/_impl.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
1 change: 1 addition & 0 deletions tests/integration/megatron/train_inf_mismatch/real_path.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
Loading
Loading