已合并
feat(trainer): support Qwen3-MoE attention activation swap #1177
feat(trainer): support Qwen3-MoE attention activation swap #1177
已合并
songjiaqi创建于 8月15日
共 10 个文件变更+373-2
@@ -92,6 +92,8 @@ mixed_precision:
92activation_checkpoint:92activation_checkpoint:
93 mode: off # off, full, selective93 mode: off # off, full, selective
94 94 
95+activation_swap: none # attention, none
96+ 
95compile:97compile:
96 enabled: false98 enabled: false
97 mode: default99 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,
atomgit-bot
atomgit-botatomgit-bot8月15日

🟠 High Priority

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

改动建议
277
- validate_placement=validate_placement,
277
+ validate_placement=validate_placement,
278
+ distributed_setup=distributed_setup,
279
+ activation_checkpoint=activation_checkpoint,
280
+ activation_swap=activation_swap,
281
+ )
应用建议
likedislike
不准确?
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
30from hyper_models.components.activation_checkpoint import (30from 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+)
33from hyper_models.components.compile import apply_compile37from hyper_models.components.compile import apply_compile
34from hyper_models.components.distributed.fsdp2 import FSDP2Manager, _instantiate_fsdp238from hyper_models.components.distributed.fsdp2 import FSDP2Manager, _instantiate_fsdp2
35from hyper_models.components.distributed.pipelining import _instantiate_pipeline39from 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 keeps568 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":
DavidFFFan
DavidFFFanDavidFFFan8月17日

非none就swap attention?

likedislike
songjiaqi
songjiaqi
8月17日 评论:
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.config376 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=ActivationCheckpointConfig700 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 # data705 # 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+@pytest.mark.parametrize(
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_setup94 assert trainer.model.distributed_setup is distributed_setup
93 assert trainer.model.peft_config is peft_config95 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.config99 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 
115def test_trainer_optimizer_stage_passes_runtime_context() -> None:119def 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.model225 _ = 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,