已合并
feat(trainer): support Qwen3-MoE attention activation swap #1177
songjiaqi创建于 8月15日
feat(trainer): support Qwen3-MoE attention activation swap #1177
已合并
共 10 个文件变更+373-2
| @@ -92,6 +92,8 @@ mixed_precision: | |||
| 92 | activation_checkpoint: | 92 | activation_checkpoint: |
| 93 | mode: off # off, full, selective | 93 | mode: off # off, full, selective |
| 94 | 94 | ||
| 95 | +activation_swap: none # attention, none | ||
| 96 | + | ||
| 95 | compile: | 97 | compile: |
| 96 | enabled: false | 98 | enabled: false |
| 97 | mode: default | 99 | mode: default |
| @@ -81,6 +81,7 @@ class _BaseHyperAutoModelClass: | |||
| 81 | compile_config=None, | 81 | compile_config=None, |
| 82 | freeze_config=None, | 82 | freeze_config=None, |
| 83 | activation_checkpoint: Optional[str] = None, | 83 | activation_checkpoint: Optional[str] = None, |
| 84 | + activation_swap: str = "none", | ||
| 84 | **kwargs, | 85 | **kwargs, |
| 85 | ) -> PreTrainedModel: | 86 | ) -> PreTrainedModel: |
| 86 | """HF-compatible from_pretrained entry point. | 87 | """HF-compatible from_pretrained entry point. |
| @@ -132,6 +133,7 @@ class _BaseHyperAutoModelClass: | |||
| 132 | compile_config=compile_config, | 133 | compile_config=compile_config, |
| 133 | freeze_config=freeze_config, | 134 | freeze_config=freeze_config, |
| 134 | activation_checkpoint=activation_checkpoint, | 135 | activation_checkpoint=activation_checkpoint, |
| 136 | + activation_swap=activation_swap, | ||
| 135 | **kwargs, | 137 | **kwargs, |
| 136 | ) | 138 | ) |
| 137 | 139 | ||
| @@ -146,6 +148,7 @@ class _BaseHyperAutoModelClass: | |||
| 146 | torch_dtype="auto", | 148 | torch_dtype="auto", |
| 147 | attn_implementation="sdpa", | 149 | attn_implementation="sdpa", |
| 148 | activation_checkpoint: Optional[str] = None, | 150 | activation_checkpoint: Optional[str] = None, |
| 151 | + activation_swap: str = "none", | ||
| 149 | **kwargs, | 152 | **kwargs, |
| 150 | ) -> PreTrainedModel: | 153 | ) -> PreTrainedModel: |
| 151 | """Build model from PretrainedConfig (no weight loading). | 154 | """Build model from PretrainedConfig (no weight loading). |
| @@ -180,6 +183,7 @@ class _BaseHyperAutoModelClass: | |||
| 180 | load_base_model=False, | 183 | load_base_model=False, |
| 181 | distributed_setup=distributed_setup, | 184 | distributed_setup=distributed_setup, |
| 182 | activation_checkpoint=activation_checkpoint, | 185 | activation_checkpoint=activation_checkpoint, |
| 186 | + activation_swap=activation_swap, | ||
| 183 | **kwargs, | 187 | **kwargs, |
| 184 | ) | 188 | ) |
| 185 | 189 | ||
| @@ -206,6 +210,7 @@ class _BaseHyperAutoModelClass: | |||
| 206 | compile_config=None, | 210 | compile_config=None, |
| 207 | freeze_config=None, | 211 | freeze_config=None, |
| 208 | activation_checkpoint: Optional[str] = None, | 212 | activation_checkpoint: Optional[str] = None, |
| 213 | + activation_swap: str = "none", | ||
| 209 | **kwargs, | 214 | **kwargs, |
| 210 | ) -> PreTrainedModel: | 215 | ) -> PreTrainedModel: |
| 211 | """Core model building orchestration. | 216 | """Core model building orchestration. |
| @@ -270,8 +275,8 @@ class _BaseHyperAutoModelClass: | |||
| 270 | load_base_model=load_base_model, | 275 | load_base_model=load_base_model, |
| 271 | pretrained_path=pretrained_model_name_or_path, | 276 | pretrained_path=pretrained_model_name_or_path, |
| 272 | validate_placement=validate_placement, | 277 | validate_placement=validate_placement, |
| 273 | - distributed_setup=distributed_setup, | ||
| 274 | activation_checkpoint=activation_checkpoint, | 278 | activation_checkpoint=activation_checkpoint, |
| 279 | + activation_swap=activation_swap, | ||
| 275 | ) | 280 | ) |
| 276 | 281 | ||
| 277 | model.train() | 282 | model.train() |
| @@ -30,6 +30,10 @@ from hyper_models._transformers.checkpoint_loader import CheckpointManager, Load | |||
| 30 | from hyper_models.components.activation_checkpoint import ( | 30 | from hyper_models.components.activation_checkpoint import ( |
| 31 | _apply_activation_checkpointing as _apply_activation_checkpointing_impl, | 31 | _apply_activation_checkpointing as _apply_activation_checkpointing_impl, |
| 32 | ) | 32 | ) |
| 33 | +from hyper_models.components.activation_swap.attention_swap import ( | ||
| 34 | + apply_qwen3_moe_attention_swap, | ||
| 35 | + validate_attention_swap, | ||
| 36 | +) | ||
| 33 | from hyper_models.components.compile import apply_compile | 37 | from hyper_models.components.compile import apply_compile |
| 34 | from hyper_models.components.distributed.fsdp2 import FSDP2Manager, _instantiate_fsdp2 | 38 | from hyper_models.components.distributed.fsdp2 import FSDP2Manager, _instantiate_fsdp2 |
| 35 | from hyper_models.components.distributed.pipelining import _instantiate_pipeline | 39 | from hyper_models.components.distributed.pipelining import _instantiate_pipeline |
| @@ -549,6 +553,7 @@ def apply_model_infrastructure( | |||
| 549 | freeze_config=None, | 553 | freeze_config=None, |
| 550 | compile_config=None, | 554 | compile_config=None, |
| 551 | activation_checkpoint: Optional[str] = None, | 555 | activation_checkpoint: Optional[str] = None, |
| 556 | + activation_swap: str = "none", | ||
| 552 | is_meta_device: bool = False, | 557 | is_meta_device: bool = False, |
| 553 | is_hf_model: bool = False, | 558 | is_hf_model: bool = False, |
| 554 | device=None, | 559 | device=None, |
| @@ -563,6 +568,7 @@ def apply_model_infrastructure( | |||
| 563 | materialization/loading -> per-layer compile. Placement validation keeps | 568 | materialization/loading -> per-layer compile. Placement validation keeps |
| 564 | the DTensor placement path and skips FSDP2 and compile. | 569 | the DTensor placement path and skips FSDP2 and compile. |
| 565 | """ | 570 | """ |
| 571 | + | ||
| 566 | distributed_setup = kwargs.get("distributed_setup") | 572 | distributed_setup = kwargs.get("distributed_setup") |
| 567 | 573 | ||
| 568 | if isinstance(compile_config, dict): | 574 | if isinstance(compile_config, dict): |
| @@ -607,7 +613,7 @@ def apply_model_infrastructure( | |||
| 607 | validate_placement, | 613 | validate_placement, |
| 608 | ) | 614 | ) |
| 609 | 615 | ||
| 610 | - # Step 9: activation checkpointing remains inside the FSDP boundary. | 616 | + # Step 9-1: activation checkpointing remains inside the FSDP boundary. |
| 611 | if activation_checkpoint not in (None, "off"): | 617 | if activation_checkpoint not in (None, "off"): |
| 612 | model = _apply_activation_checkpointing( | 618 | model = _apply_activation_checkpointing( |
| 613 | model, | 619 | model, |
| @@ -615,6 +621,17 @@ def apply_model_infrastructure( | |||
| 615 | enable_compile=compile_for_execution, | 621 | enable_compile=compile_for_execution, |
| 616 | ) | 622 | ) |
| 617 | 623 | ||
| 624 | + # Step 9-2:activation swap. | ||
| 625 | + validate_attention_swap( | ||
| 626 | + activation_swap, | ||
| 627 | + activation_checkpoint=activation_checkpoint, | ||
| 628 | + enable_compile=compile_for_execution, | ||
| 629 | + pp_size=getattr(mesh, "pp_size", 1), | ||
| 630 | + ) | ||
| 631 | + if activation_swap != "none": | ||
非none就swap attention? ![]() ![]() | |||
| 632 | + model = apply_qwen3_moe_attention_swap(model, activation_swap) | ||
| 633 | + | ||
| 634 | + | ||
| 618 | # Step 10: FSDP2 is execution-only; validation retains DTensor placement. | 635 | # Step 10: FSDP2 is execution-only; validation retains DTensor placement. |
| 619 | if validate_placement: | 636 | if validate_placement: |
| 620 | if fsdp2_manager is not None: | 637 | if fsdp2_manager is not None: |
| @@ -0,0 +1,27 @@ | |||
| 1 | +# Copyright 2026 Huawei Technologies Co., Ltd | ||
| 2 | +# | ||
| 3 | +# Licensed under the Apache License, Version 2.0 (the "License"); | ||
| 4 | +# you may not use this file except in compliance with the License. | ||
| 5 | +# You may obtain a copy of the License at | ||
| 6 | +# | ||
| 7 | +# http://www.apache.org/licenses/LICENSE-2.0 | ||
| 8 | +# | ||
| 9 | +# Unless required by applicable law or agreed to in writing, software | ||
| 10 | +# distributed under the License is distributed on an "AS IS" BASIS, | ||
| 11 | +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. | ||
| 12 | +# See the License for the specific language governing permissions and | ||
| 13 | +# limitations under the License. | ||
| 14 | +# ============================================================================ | ||
| 15 | +"""Model-specific activation swap support.""" | ||
| 16 | + | ||
| 17 | +from hyper_models.components.activation_swap.attention_swap import ( | ||
| 18 | + apply_qwen3_moe_attention_swap, | ||
| 19 | + qwen3_attention_swap_policy, | ||
| 20 | + validate_attention_swap, | ||
| 21 | +) | ||
| 22 | + | ||
| 23 | +__all__ = [ | ||
| 24 | + "apply_qwen3_moe_attention_swap", | ||
| 25 | + "qwen3_attention_swap_policy", | ||
| 26 | + "validate_attention_swap", | ||
| 27 | +] | ||
| @@ -0,0 +1,170 @@ | |||
| 1 | +# Copyright 2026 Huawei Technologies Co., Ltd | ||
| 2 | +# | ||
| 3 | +# Licensed under the Apache License, Version 2.0 (the "License"); | ||
| 4 | +# you may not use this file except in compliance with the License. | ||
| 5 | +# You may obtain a copy of the License at | ||
| 6 | +# | ||
| 7 | +# http://www.apache.org/licenses/LICENSE-2.0 | ||
| 8 | +# | ||
| 9 | +# Unless required by applicable law or agreed to in writing, software | ||
| 10 | +# distributed under the License is distributed on an "AS IS" BASIS, | ||
| 11 | +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. | ||
| 12 | +# See the License for the specific language governing permissions and | ||
| 13 | +# limitations under the License. | ||
| 14 | +# ============================================================================ | ||
| 15 | +"""Attention activation swap support for Hugging Face Qwen3-MoE models.""" | ||
| 16 | + | ||
| 17 | +import logging | ||
| 18 | +from typing import Any, Optional | ||
| 19 | + | ||
| 20 | +import torch | ||
| 21 | +from torch import nn | ||
| 22 | + | ||
| 23 | +from hyper_parallel.core.activation_checkpoint import CheckpointPolicy, SwapManager, swap_wrapper | ||
| 24 | +from hyper_parallel.platform import get_platform | ||
| 25 | +from hyper_parallel.platform.platform import PlatformType | ||
| 26 | + | ||
| 27 | +logger = logging.getLogger(__name__) | ||
| 28 | +platform = get_platform() | ||
| 29 | + | ||
| 30 | +_MIN_SWAP_TENSOR_BYTES = 1024 * 1024 | ||
| 31 | + | ||
| 32 | + | ||
| 33 | +def _qwen3_types() -> tuple[type[nn.Module], type[nn.Module], type[nn.Module]]: | ||
| 34 | + """Load Qwen3 classes only when the opt-in feature is requested.""" | ||
| 35 | + try: | ||
| 36 | + from transformers.models.qwen3_moe.modeling_qwen3_moe import ( # pylint: disable=C0415 | ||
| 37 | + Qwen3MoeAttention, | ||
| 38 | + Qwen3MoeDecoderLayer, | ||
| 39 | + Qwen3MoeForCausalLM, | ||
| 40 | + ) | ||
| 41 | + except (ImportError, ModuleNotFoundError) as exc: | ||
| 42 | + raise ValueError( | ||
| 43 | + "activation_swap='attention' requires a Transformers build with Qwen3-MoE support" | ||
| 44 | + ) from exc | ||
| 45 | + return Qwen3MoeForCausalLM, Qwen3MoeDecoderLayer, Qwen3MoeAttention | ||
| 46 | + | ||
| 47 | + | ||
| 48 | +def qwen3_attention_swap_policy(tensor: torch.Tensor) -> CheckpointPolicy: | ||
| 49 | + """Select large, independently-owned attention activations for swapping. | ||
| 50 | + | ||
| 51 | + Args: | ||
| 52 | + tensor: Tensor saved by autograd inside a Qwen3-MoE attention forward. | ||
| 53 | + | ||
| 54 | + Returns: | ||
| 55 | + ``MUST_SWAP`` for tensors that are safe and worthwhile to transfer; | ||
| 56 | + otherwise ``MUST_SAVE``. | ||
| 57 | + """ | ||
| 58 | + if not tensor.requires_grad or tensor.dim() < 2: | ||
| 59 | + return CheckpointPolicy.MUST_SAVE | ||
| 60 | + storage_bytes = tensor.untyped_storage().size() | ||
| 61 | + tensor_bytes = tensor.numel() * tensor.element_size() | ||
| 62 | + if storage_bytes != tensor_bytes or tensor_bytes < _MIN_SWAP_TENSOR_BYTES: | ||
| 63 | + return CheckpointPolicy.MUST_SAVE | ||
| 64 | + return CheckpointPolicy.MUST_SWAP | ||
| 65 | + | ||
| 66 | + | ||
| 67 | +def validate_attention_swap( | ||
| 68 | + activation_swap: str, | ||
| 69 | + *, | ||
| 70 | + activation_checkpoint: Optional[str] = None, | ||
| 71 | + enable_compile: bool = False, | ||
| 72 | + pp_size: int = 1, | ||
| 73 | +) -> None: | ||
| 74 | + """Validate attention swap mode and incompatible model-build features. | ||
| 75 | + | ||
| 76 | + Args: | ||
| 77 | + activation_swap: Requested activation swap mode. | ||
| 78 | + activation_checkpoint: Activation recomputation mode. | ||
| 79 | + enable_compile: Whether graph compilation is enabled. | ||
| 80 | + pp_size: Configured pipeline-parallel world size. | ||
| 81 | + | ||
| 82 | + Raises: | ||
| 83 | + ValueError: If the mode is invalid or an unsupported combination is enabled. | ||
| 84 | + """ | ||
| 85 | + if activation_swap not in ("none", "attention"): | ||
| 86 | + raise ValueError( | ||
| 87 | + "activation_swap must be one of ('none', 'attention'), " | ||
| 88 | + f"got {activation_swap!r}" | ||
| 89 | + ) | ||
| 90 | + if activation_swap == "none": | ||
| 91 | + return | ||
| 92 | + if enable_compile: | ||
| 93 | + raise ValueError("activation_swap='attention' is incompatible with torch.compile") | ||
| 94 | + if activation_checkpoint not in (None, "off"): | ||
| 95 | + raise ValueError("activation_swap='attention' is incompatible with activation checkpointing") | ||
| 96 | + if pp_size != 1: | ||
| 97 | + raise ValueError("activation_swap='attention' does not support pipeline parallelism") | ||
| 98 | + | ||
| 99 | + | ||
| 100 | +def _find_qwen3_moe_attentions(model: nn.Module) -> list[Any]: | ||
| 101 | + """Validate the supported HF model structure and return local attentions.""" | ||
| 102 | + qwen3_model_type, qwen3_layer_type, qwen3_attention_type = _qwen3_types() | ||
| 103 | + if type(model) is not qwen3_model_type: | ||
| 104 | + raise ValueError( | ||
| 105 | + "activation_swap='attention' only supports Hugging Face Qwen3MoeForCausalLM for now, " | ||
| 106 | + f"got {type(model).__name__}" | ||
| 107 | + ) | ||
| 108 | + layers = getattr(getattr(model, "model", None), "layers", None) | ||
| 109 | + if not isinstance(layers, nn.ModuleList) or len(layers) != model.config.num_hidden_layers: | ||
| 110 | + raise ValueError( | ||
| 111 | + "activation_swap='attention' requires model.layers to contain exactly " | ||
| 112 | + f"{model.config.num_hidden_layers} Qwen3MoeDecoderLayer instances" | ||
| 113 | + ) | ||
| 114 | + | ||
| 115 | + attentions = [] | ||
| 116 | + for layer_index, layer in enumerate(layers): | ||
| 117 | + if type(layer) is not qwen3_layer_type: | ||
| 118 | + raise ValueError( | ||
| 119 | + "activation_swap='attention' requires every model.layers entry to be " | ||
| 120 | + f"Qwen3MoeDecoderLayer, got {type(layer).__name__} at index {layer_index}" | ||
| 121 | + ) | ||
| 122 | + attention = getattr(layer, "self_attn", None) | ||
| 123 | + if type(attention) is not qwen3_attention_type: | ||
| 124 | + raise ValueError( | ||
| 125 | + "activation_swap='attention' requires every decoder layer self_attn to be " | ||
| 126 | + f"Qwen3MoeAttention, got {type(attention).__name__} at index {layer_index}" | ||
| 127 | + ) | ||
| 128 | + attentions.append(attention) | ||
| 129 | + return attentions | ||
| 130 | + | ||
| 131 | + | ||
| 132 | +def apply_qwen3_moe_attention_swap(model: nn.Module, activation_swap: str) -> nn.Module: | ||
| 133 | + """Wrap Qwen3-MoE attention modules and install layer-wise swap scheduling. | ||
| 134 | + | ||
| 135 | + Args: | ||
| 136 | + model: Model being prepared, before sharding and FSDP wrapping. | ||
| 137 | + activation_swap: Requested activation swap mode. | ||
| 138 | + | ||
| 139 | + Returns: | ||
| 140 | + The input model, patched in place when attention swapping is enabled. | ||
| 141 | + """ | ||
| 142 | + if activation_swap != "attention": | ||
| 143 | + raise ValueError( | ||
| 144 | + "activation_swap must be one of ('none', 'attention'), " | ||
| 145 | + f"got {activation_swap!r}" | ||
| 146 | + ) | ||
| 147 | + attentions = _find_qwen3_moe_attentions(model) | ||
| 148 | + wrapped_attentions = [] | ||
| 149 | + for layer, attention in zip(model.model.layers, attentions): | ||
| 150 | + wrapped_attention = swap_wrapper( | ||
| 151 | + attention, | ||
| 152 | + policy_fn=qwen3_attention_swap_policy, | ||
| 153 | + group_swap=True, | ||
| 154 | + ) | ||
| 155 | + layer.self_attn = wrapped_attention | ||
| 156 | + wrapped_attentions.append(wrapped_attention) | ||
| 157 | + | ||
| 158 | + swap_manager = SwapManager() | ||
| 159 | + for current_attention, next_attention in zip(wrapped_attentions, wrapped_attentions[1:]): | ||
| 160 | + swap_manager.set_forward_prefetch_layer(current_attention, next_attention) | ||
| 161 | + | ||
| 162 | + logger.info("Enabled attention activation swap for %d Qwen3-MoE layers", len(wrapped_attentions)) | ||
| 163 | + return model | ||
| 164 | + | ||
| 165 | + | ||
| 166 | +__all__ = [ | ||
| 167 | + "apply_qwen3_moe_attention_swap", | ||
| 168 | + "qwen3_attention_swap_policy", | ||
| 169 | + "validate_attention_swap", | ||
| 170 | +] | ||
| @@ -370,6 +370,7 @@ class BaseTrainer(Stateful, ABC): | |||
| 370 | distributed_setup=self.distributed_setup, | 370 | distributed_setup=self.distributed_setup, |
| 371 | peft_config=self.peft_config, | 371 | peft_config=self.peft_config, |
| 372 | activation_checkpoint=self.config.activation_checkpoint.mode, | 372 | activation_checkpoint=self.config.activation_checkpoint.mode, |
| 373 | + activation_swap=self.config.activation_swap, | ||
| 373 | compile_config=self.config.compile, | 374 | compile_config=self.config.compile, |
| 374 | ) | 375 | ) |
| 375 | self.model_config = self.model.config | 376 | self.model_config = self.model.config |
| @@ -699,6 +699,7 @@ class TrainerConfig: | |||
| 699 | activation_checkpoint: ActivationCheckpointConfig = field( | 699 | activation_checkpoint: ActivationCheckpointConfig = field( |
| 700 | default_factory=ActivationCheckpointConfig | 700 | default_factory=ActivationCheckpointConfig |
| 701 | ) | 701 | ) |
| 702 | + activation_swap: Literal["none", "attention"] = "none" | ||
| 702 | compile: CompileConfig = field(default_factory=CompileConfig) | 703 | compile: CompileConfig = field(default_factory=CompileConfig) |
| 703 | 704 | ||
| 704 | # data | 705 | # data |
| @@ -0,0 +1,130 @@ | |||
| 1 | +# Copyright 2026 Huawei Technologies Co., Ltd | ||
| 2 | +# | ||
| 3 | +# Licensed under the Apache License, Version 2.0 (the "License"); | ||
| 4 | +# you may not use this file except in compliance with the License. | ||
| 5 | +# You may obtain a copy of the License at | ||
| 6 | +# | ||
| 7 | +# http://www.apache.org/licenses/LICENSE-2.0 | ||
| 8 | +# | ||
| 9 | +# Unless required by applicable law or agreed to in writing, software | ||
| 10 | +# distributed under the License is distributed on an "AS IS" BASIS, | ||
| 11 | +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. | ||
| 12 | +# See the License for the specific language governing permissions and | ||
| 13 | +# limitations under the License. | ||
| 14 | +# ============================================================================ | ||
| 15 | +"""Tests for Qwen3-30B-A3B attention activation swapping.""" | ||
| 16 | + | ||
| 17 | +import os | ||
| 18 | +from types import SimpleNamespace | ||
| 19 | +from unittest.mock import MagicMock, patch | ||
| 20 | + | ||
| 21 | +os.environ.setdefault("HYPER_PARALLEL_PLATFORM", "torch") | ||
| 22 | + | ||
| 23 | +import pytest | ||
| 24 | +import torch | ||
| 25 | +from torch import nn | ||
| 26 | +from transformers.models.qwen3_moe.modeling_qwen3_moe import ( | ||
| 27 | + Qwen3MoeAttention, | ||
| 28 | + Qwen3MoeDecoderLayer, | ||
| 29 | + Qwen3MoeForCausalLM, | ||
| 30 | +) | ||
| 31 | + | ||
| 32 | +from hyper_models.components.activation_swap.attention_swap import ( | ||
| 33 | + apply_qwen3_moe_attention_swap, | ||
| 34 | + qwen3_attention_swap_policy, | ||
| 35 | + validate_attention_swap, | ||
| 36 | +) | ||
| 37 | +from hyper_parallel.core.activation_checkpoint import CheckpointPolicy | ||
| 38 | +from hyper_parallel.platform.torch.activation_checkpoint.activation_swap import SwapWrapper | ||
| 39 | + | ||
| 40 | + | ||
| 41 | +def _new_module(module_type: type[nn.Module]) -> nn.Module: | ||
| 42 | + """Construct a module instance without allocating its production weights.""" | ||
| 43 | + module = module_type.__new__(module_type) | ||
| 44 | + nn.Module.__init__(module) | ||
| 45 | + return module | ||
| 46 | + | ||
| 47 | + | ||
| 48 | +def _qwen3_30b_a3b_stub() -> Qwen3MoeForCausalLM: | ||
| 49 | + """Build the supported HF module hierarchy with lightweight layer contents.""" | ||
| 50 | + model = _new_module(Qwen3MoeForCausalLM) | ||
| 51 | + model.config = SimpleNamespace( | ||
| 52 | + attention_bias=False, | ||
| 53 | + attention_dropout=0.0, | ||
| 54 | + decoder_sparse_step=1, | ||
| 55 | + head_dim=128, | ||
| 56 | + hidden_size=2048, | ||
| 57 | + intermediate_size=6144, | ||
| 58 | + max_position_embeddings=40960, | ||
| 59 | + moe_intermediate_size=768, | ||
| 60 | + num_attention_heads=32, | ||
| 61 | + num_experts=128, | ||
| 62 | + num_experts_per_tok=8, | ||
| 63 | + num_hidden_layers=48, | ||
| 64 | + num_key_value_heads=4, | ||
| 65 | + norm_topk_prob=True, | ||
| 66 | + tie_word_embeddings=False, | ||
| 67 | + use_sliding_window=False, | ||
| 68 | + vocab_size=151936, | ||
| 69 | + ) | ||
| 70 | + model.model = nn.Module() | ||
| 71 | + model.model.layers = nn.ModuleList() | ||
| 72 | + for _ in range(model.config.num_hidden_layers): | ||
| 73 | + layer = _new_module(Qwen3MoeDecoderLayer) | ||
| 74 | + attention = _new_module(Qwen3MoeAttention) | ||
| 75 | + attention.proj = nn.Linear(2, 2) | ||
| 76 | + layer.self_attn = attention | ||
| 77 | + layer.input_layernorm = nn.Identity() | ||
| 78 | + layer.post_attention_layernorm = nn.Identity() | ||
| 79 | + layer.mlp = nn.Linear(2, 2) | ||
| 80 | + model.model.layers.append(layer) | ||
| 81 | + return model | ||
| 82 | + | ||
| 83 | + | ||
| 84 | +def test_attention_swap_policy_filters_unsafe_or_small_tensors() -> None: | ||
| 85 | + no_grad = torch.empty(1024, 1024) | ||
| 86 | + one_dimensional = torch.empty(1024 * 1024, requires_grad=True) | ||
| 87 | + small = torch.empty(16, 16, requires_grad=True) | ||
| 88 | + base = torch.empty(1024, 1024, requires_grad=True) | ||
| 89 | + shared_storage_view = base[:512] | ||
| 90 | + large = torch.empty(512, 512, requires_grad=True) | ||
| 91 | + | ||
| 92 | + assert qwen3_attention_swap_policy(no_grad) is CheckpointPolicy.MUST_SAVE | ||
| 93 | + assert qwen3_attention_swap_policy(one_dimensional) is CheckpointPolicy.MUST_SAVE | ||
| 94 | + assert qwen3_attention_swap_policy(small) is CheckpointPolicy.MUST_SAVE | ||
| 95 | + assert qwen3_attention_swap_policy(shared_storage_view) is CheckpointPolicy.MUST_SAVE | ||
| 96 | + assert qwen3_attention_swap_policy(large) is CheckpointPolicy.MUST_SWAP | ||
| 97 | + | ||
| 98 | + | ||
| 99 | + | ||
| 100 | + ("kwargs", "message"), | ||
| 101 | + [ | ||
| 102 | + ({"enable_compile": True}, "torch.compile"), | ||
| 103 | + ({"activation_checkpoint": "full"}, "activation checkpointing"), | ||
| 104 | + ({"pp_size": 2}, "pipeline parallelism"), | ||
| 105 | + ], | ||
| 106 | +) | ||
| 107 | +def test_attention_swap_rejects_incompatible_features(kwargs: dict, message: str) -> None: | ||
| 108 | + with pytest.raises(ValueError, match=message): | ||
| 109 | + validate_attention_swap("attention", **kwargs) | ||
| 110 | + | ||
| 111 | + | ||
| 112 | +def test_attention_swap_only_wraps_attention_and_schedules_local_layers() -> None: | ||
| 113 | + model = _qwen3_30b_a3b_stub() | ||
| 114 | + state_dict_keys = set(model.state_dict()) | ||
| 115 | + norms = [layer.input_layernorm for layer in model.model.layers] | ||
| 116 | + mlps = [layer.mlp for layer in model.model.layers] | ||
| 117 | + manager = MagicMock() | ||
| 118 | + | ||
| 119 | + with patch( | ||
| 120 | + "hyper_models.components.activation_swap.attention_swap.SwapManager", | ||
| 121 | + return_value=manager, | ||
| 122 | + ): | ||
| 123 | + result = apply_qwen3_moe_attention_swap(model, "attention") | ||
| 124 | + | ||
| 125 | + assert result is model | ||
| 126 | + assert all(isinstance(layer.self_attn, SwapWrapper) for layer in model.model.layers) | ||
| 127 | + assert [layer.input_layernorm for layer in model.model.layers] == norms | ||
| 128 | + assert [layer.mlp for layer in model.model.layers] == mlps | ||
| 129 | + assert set(model.state_dict()) == state_dict_keys | ||
| 130 | + assert manager.set_forward_prefetch_layer.call_count == 47 | ||
| @@ -49,6 +49,7 @@ def _model_target( | |||
| 49 | distributed_setup: object, | 49 | distributed_setup: object, |
| 50 | peft_config: object, | 50 | peft_config: object, |
| 51 | activation_checkpoint: object, | 51 | activation_checkpoint: object, |
| 52 | + activation_swap: str, | ||
| 52 | ) -> SimpleNamespace: | 53 | ) -> SimpleNamespace: |
| 53 | """Return a model object for delegation testing.""" | 54 | """Return a model object for delegation testing.""" |
| 54 | return SimpleNamespace( | 55 | return SimpleNamespace( |
| @@ -56,6 +57,7 @@ def _model_target( | |||
| 56 | distributed_setup=distributed_setup, | 57 | distributed_setup=distributed_setup, |
| 57 | peft_config=peft_config, | 58 | peft_config=peft_config, |
| 58 | activation_checkpoint=activation_checkpoint, | 59 | activation_checkpoint=activation_checkpoint, |
| 60 | + activation_swap=activation_swap, | ||
| 59 | ) | 61 | ) |
| 60 | 62 | ||
| 61 | 63 | ||
| @@ -92,6 +94,7 @@ def test_trainer_model_stage_delegates_to_config_target() -> None: | |||
| 92 | assert trainer.model.distributed_setup is distributed_setup | 94 | assert trainer.model.distributed_setup is distributed_setup |
| 93 | assert trainer.model.peft_config is peft_config | 95 | assert trainer.model.peft_config is peft_config |
| 94 | assert trainer.model.activation_checkpoint == "full" | 96 | assert trainer.model.activation_checkpoint == "full" |
| 97 | + assert trainer.model.activation_swap == "none" | ||
| 95 | assert trainer.model_parts == [trainer.model] | 98 | assert trainer.model_parts == [trainer.model] |
| 96 | assert trainer.model_config is trainer.model.config | 99 | assert trainer.model_config is trainer.model.config |
| 97 | assert trainer.hsdp_model_parts == [] | 100 | assert trainer.hsdp_model_parts == [] |
| @@ -110,6 +113,7 @@ def test_trainer_model_stage_passes_off_activation_checkpointing_mode() -> None: | |||
| 110 | trainer._build_model() | 113 | trainer._build_model() |
| 111 | 114 | ||
| 112 | assert trainer.model.activation_checkpoint == "off" | 115 | assert trainer.model.activation_checkpoint == "off" |
| 116 | + assert trainer.model.activation_swap == "none" | ||
| 113 | 117 | ||
| 114 | 118 | ||
| 115 | def test_trainer_optimizer_stage_passes_runtime_context() -> None: | 119 | def test_trainer_optimizer_stage_passes_runtime_context() -> None: |
| @@ -224,6 +224,20 @@ class TestTargetResolution(unittest.TestCase): | |||
| 224 | with self.assertRaises(AttributeError): | 224 | with self.assertRaises(AttributeError): |
| 225 | _ = config.optimizer.model | 225 | _ = config.optimizer.model |
| 226 | 226 | ||
| 227 | + def test_activation_swap_defaults_to_none_and_accepts_attention(self): | ||
| 228 | + default_config = resolve_root(_root()) | ||
| 229 | + attention_config = resolve_root(_root(activation_swap="attention")) | ||
| 230 | + | ||
| 231 | + self.assertEqual(default_config.activation_swap, "none") | ||
| 232 | + self.assertEqual(attention_config.activation_swap, "attention") | ||
| 233 | + | ||
| 234 | + def test_activation_swap_rejects_unknown_mode(self): | ||
| 235 | + with self.assertRaisesRegex( | ||
| 236 | + ConfigResolutionError, | ||
| 237 | + r"\$\.activation_swap: expected one of \('none', 'attention'\)", | ||
| 238 | + ): | ||
| 239 | + resolve_root(_root(activation_swap="all")) | ||
| 240 | + | ||
| 227 | def test_gradient_clipping_is_not_an_optimizer_argument(self): | 241 | def test_gradient_clipping_is_not_an_optimizer_argument(self): |
| 228 | with self.assertRaisesRegex( | 242 | with self.assertRaisesRegex( |
| 229 | ConfigResolutionError, | 243 | ConfigResolutionError, |


🟠 High Priority
建议:在
_build_model调用apply_model_infrastructure时恢复透传distributed_setup=distributed_setup,使其与from_pretrained/from_config已传入的该形参保持一致。