Skip to content

Commit 3ddb7ec

Browse files
committed
Fix FFN depth-scaled initialization target
1 parent 948d65c commit 3ddb7ec

17 files changed

Lines changed: 108 additions & 39 deletions

File tree

‎tests/unit_tests/cpu/test_feed_forward.py‎

Lines changed: 63 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -10,8 +10,14 @@
1010
import torch
1111
import torch.nn.functional as F
1212

13+
import torchtitan.models.llama3.flavors as llama3_flavors
14+
1315
from torchtitan.models.common.activation import SiTUGLU
14-
from torchtitan.models.common.config_utils import fused_gate_up_param_init
16+
from torchtitan.models.common.config_utils import (
17+
fused_gate_up_param_init,
18+
make_ffn_config,
19+
make_shared_expert_ffn_config,
20+
)
1521
from torchtitan.models.common.feed_forward import FeedForward
1622
from torchtitan.models.common.linear import Linear
1723

@@ -59,6 +65,62 @@ def test_feed_forward_uses_one_physical_gate_up_linear():
5965
torch.testing.assert_close(w13_2HD[1], 5 * torch.ones_like(w13_2HD[1]))
6066

6167

68+
def test_make_ffn_config_uses_input_init_for_gate_and_up():
69+
config = make_ffn_config(
70+
dim=4,
71+
hidden_dim=8,
72+
w1_param_init={"weight": _fill(1.0)},
73+
w2_param_init={"weight": _fill(2.0)},
74+
)
75+
feed_forward = config.build()
76+
feed_forward.init_states()
77+
78+
w13_2HD = feed_forward.w13.weight
79+
torch.testing.assert_close(w13_2HD[0], torch.ones_like(w13_2HD[0]))
80+
torch.testing.assert_close(w13_2HD[1], torch.ones_like(w13_2HD[1]))
81+
torch.testing.assert_close(
82+
feed_forward.w2.weight, 2 * torch.ones_like(feed_forward.w2.weight)
83+
)
84+
85+
86+
def test_make_shared_expert_ffn_config_uses_input_init_for_gate_and_up():
87+
config = make_shared_expert_ffn_config(
88+
dim=4,
89+
hidden_dim=8,
90+
w1_param_init={"weight": _fill(1.0)},
91+
w2_param_init={"weight": _fill(2.0)},
92+
)
93+
feed_forward = config.build()
94+
feed_forward.init_states()
95+
96+
w13_2HD = feed_forward.w13.weight
97+
torch.testing.assert_close(w13_2HD[0], torch.ones_like(w13_2HD[0]))
98+
torch.testing.assert_close(w13_2HD[1], torch.ones_like(w13_2HD[1]))
99+
torch.testing.assert_close(
100+
feed_forward.w2.weight, 2 * torch.ones_like(feed_forward.w2.weight)
101+
)
102+
103+
104+
def test_llama3_depth_scales_only_ffn_output(monkeypatch):
105+
linear_init = {"weight": _fill(1.0), "bias": _fill(0.0)}
106+
depth_init = {"weight": _fill(2.0), "bias": _fill(0.0)}
107+
monkeypatch.setattr(llama3_flavors, "_LINEAR_INIT", linear_init)
108+
monkeypatch.setattr(
109+
llama3_flavors, "_depth_init", lambda _layer_id: depth_init
110+
)
111+
112+
build_config, max_context_length = llama3_flavors.MODEL_FLAVORS["debugmodel"]
113+
model_config = build_config(attn_backend="flex", seq_len=max_context_length)
114+
feed_forward = model_config.layers[0].feed_forward.build()
115+
feed_forward.init_states()
116+
117+
w13_2HD = feed_forward.w13.weight
118+
torch.testing.assert_close(w13_2HD[0], torch.ones_like(w13_2HD[0]))
119+
torch.testing.assert_close(w13_2HD[1], torch.ones_like(w13_2HD[1]))
120+
torch.testing.assert_close(
121+
feed_forward.w2.weight, 2 * torch.ones_like(feed_forward.w2.weight)
122+
)
123+
62124
def test_feed_forward_requires_two_w13_projections():
63125
config = FeedForward.Config(
64126
w13=Linear.Config(in_features=4, out_features=8),

‎tests/unit_tests/cpu/test_lora.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -391,7 +391,7 @@ def test_stacked_lora_adapter_does_not_repeat_base_redistribution():
391391
dim=4,
392392
hidden_dim=8,
393393
w1_param_init=init,
394-
w2w3_param_init=init,
394+
w2_param_init=init,
395395
)
396396
config = LoRATransform(
397397
rank=2,

‎tests/unit_tests/cpu/test_moe.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -428,7 +428,7 @@ def test_shared_expert_uses_runtime_sp_aware_output_projection(self):
428428
dim=4,
429429
hidden_dim=8,
430430
w1_param_init={},
431-
w2w3_param_init={},
431+
w2_param_init={},
432432
)
433433

434434
self.assertIs(type(config.w13), ColumnParallelLinear.Config)

‎tests/unit_tests/cpu/test_tp_kv_heads_validation.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -93,7 +93,7 @@ def _make_llama3_config(n_heads: int, n_kv_heads: int | None) -> "Llama3Model.Co
9393
dim=_DIM,
9494
hidden_dim=compute_ffn_hidden_dim(_DIM, multiple_of=256),
9595
w1_param_init=_LINEAR_INIT,
96-
w2w3_param_init=_LINEAR_INIT,
96+
w2_param_init=_LINEAR_INIT,
9797
),
9898
)
9999
)

‎tests/unit_tests/gpu/test_async_linear.py‎

Lines changed: 5 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -204,7 +204,7 @@ def test_w13_tp_shards_the_matrix_row_dimension(self):
204204
dim=DIM,
205205
hidden_dim=hidden_dim,
206206
w1_param_init=init,
207-
w2w3_param_init=init,
207+
w2_param_init=init,
208208
)
209209
set_dense_ffn_sharding(
210210
ffn_config,
@@ -237,7 +237,7 @@ def test_stacked_w13_typechecks(self):
237237
dim=DIM,
238238
hidden_dim=128,
239239
w1_param_init=init,
240-
w2w3_param_init=init,
240+
w2_param_init=init,
241241
)
242242
set_dense_ffn_sharding(
243243
ffn_config,
@@ -429,7 +429,7 @@ def test_matches_standard_feed_forward(self):
429429
torch.manual_seed(0)
430430
standard = (
431431
make_ffn_config(
432-
dim=dim, hidden_dim=hidden, w1_param_init=init, w2w3_param_init=init
432+
dim=dim, hidden_dim=hidden, w1_param_init=init, w2_param_init=init
433433
)
434434
.build()
435435
.to(dev)
@@ -438,7 +438,7 @@ def test_matches_standard_feed_forward(self):
438438
dim=dim,
439439
hidden_dim=hidden,
440440
w1_param_init=init,
441-
w2w3_param_init=init,
441+
w2_param_init=init,
442442
)
443443
async_config = AsyncTensorParallelTransform(
444444
enable_sequence_parallel=True
@@ -534,7 +534,7 @@ def make():
534534
dim=dim,
535535
hidden_dim=hidden,
536536
w1_param_init=init,
537-
w2w3_param_init=init,
537+
w2_param_init=init,
538538
)
539539

540540
torch.manual_seed(0)

‎tests/unit_tests/gpu/test_fused_swiglu.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -147,7 +147,7 @@ def _dist_gemm_ffn_config(**kwargs):
147147

148148
init = {"weight": torch.nn.init.zeros_}
149149
return make_ffn_config(
150-
dim=_DIM, hidden_dim=_HIDDEN, w1_param_init=init, w2w3_param_init=init, **kwargs
150+
dim=_DIM, hidden_dim=_HIDDEN, w1_param_init=init, w2_param_init=init, **kwargs
151151
)
152152

153153

‎tests/unit_tests/gpu/test_tensor_parallel_feed_forward.py‎

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -49,7 +49,7 @@ def test_matches_unsharded_with_and_without_sequence_parallel(self):
4949
dim=dim,
5050
hidden_dim=hidden_dim,
5151
w1_param_init=init,
52-
w2w3_param_init=init,
52+
w2_param_init=init,
5353
)
5454
reference = copy.deepcopy(base_config).build().to(device)
5555
parallel_config = copy.deepcopy(base_config)
@@ -142,7 +142,7 @@ def test_shared_expert_without_sp_reduces_at_parent_boundary(self):
142142
dim=dim,
143143
hidden_dim=hidden_dim,
144144
w1_param_init=init,
145-
w2w3_param_init=init,
145+
w2_param_init=init,
146146
)
147147
reference = copy.deepcopy(base_config).build().to(device)
148148
parallel_config = copy.deepcopy(base_config)
@@ -220,7 +220,7 @@ def test_shared_expert_with_sp_reduce_scatters_at_w2_boundary(self):
220220
dim=dim,
221221
hidden_dim=hidden_dim,
222222
w1_param_init=init,
223-
w2w3_param_init=init,
223+
w2_param_init=init,
224224
)
225225
reference = copy.deepcopy(base_config).build().to(device)
226226
parallel_config = copy.deepcopy(base_config)

‎torchtitan/experiments/transformers_modeling_backend/moe_replacement.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -530,7 +530,7 @@ def _build_moe_config(params: dict, config) -> MoE.Config:
530530
dim=shared_info["dim"],
531531
hidden_dim=shared_info["hidden_dim"],
532532
w1_param_init=_LINEAR_INIT,
533-
w2w3_param_init=_LINEAR_INIT,
533+
w2_param_init=_LINEAR_INIT,
534534
)
535535
if shared_info["has_sigmoid_gate"]:
536536
# Import only for the Qwen3.5 topology so unrelated HF models do

‎torchtitan/models/common/config_utils.py‎

Lines changed: 16 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -265,20 +265,24 @@ def make_ffn_config(
265265
dim: int,
266266
hidden_dim: int,
267267
w1_param_init: dict[str, Callable],
268-
w2w3_param_init: dict[str, Callable],
268+
w2_param_init: dict[str, Callable],
269269
) -> FeedForward.Config:
270-
"""Build a fully-specified FeedForward.Config."""
270+
"""Build a fully-specified FeedForward.Config.
271+
272+
``w1`` and ``w3`` are the gate/up projections and share
273+
``w1_param_init``; ``w2`` is the residual output projection.
274+
"""
271275
return FeedForward.Config(
272276
w13=ColumnParallelLinear.Config(
273277
in_features=dim,
274278
out_features=hidden_dim,
275279
num_linears=2,
276-
param_init=fused_gate_up_param_init(w1_param_init, w2w3_param_init),
280+
param_init=fused_gate_up_param_init(w1_param_init, w1_param_init),
277281
),
278282
w2=RowParallelLinear.Config(
279283
in_features=hidden_dim,
280284
out_features=dim,
281-
param_init=w2w3_param_init,
285+
param_init=w2_param_init,
282286
),
283287
)
284288

@@ -288,24 +292,27 @@ def make_shared_expert_ffn_config(
288292
dim: int,
289293
hidden_dim: int,
290294
w1_param_init: dict[str, Callable],
291-
w2w3_param_init: dict[str, Callable],
295+
w2_param_init: dict[str, Callable],
292296
) -> FeedForward.Config:
293-
"""Build a shared FFN whose output reduction is selected at runtime."""
297+
"""Build a shared FFN whose output reduction is selected at runtime.
298+
299+
``w1`` and ``w3`` are the gate/up projections and share
300+
``w1_param_init``; ``w2`` is the residual output projection.
301+
"""
294302
return FeedForward.Config(
295303
w13=ColumnParallelLinear.Config(
296304
in_features=dim,
297305
out_features=hidden_dim,
298306
num_linears=2,
299-
param_init=fused_gate_up_param_init(w1_param_init, w2w3_param_init),
307+
param_init=fused_gate_up_param_init(w1_param_init, w1_param_init),
300308
),
301309
w2=SharedExpertRowParallelLinear.Config(
302310
in_features=hidden_dim,
303311
out_features=dim,
304-
param_init=w2w3_param_init,
312+
param_init=w2_param_init,
305313
),
306314
)
307315

308-
309316
def make_moe_config(
310317
*,
311318
num_experts: int = 8,

‎torchtitan/models/deepseek_v3/flavors.py‎

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -84,7 +84,7 @@ def _depth_experts_init(layer_id: int) -> dict[str, Callable]:
8484
return {
8585
"w1_EFD": partial(nn.init.trunc_normal_, std=0.02),
8686
"w2_EDF": partial(nn.init.trunc_normal_, std=depth_scaled_std(0.02, layer_id)),
87-
"w3_EFD": partial(nn.init.trunc_normal_, std=depth_scaled_std(0.02, layer_id)),
87+
"w3_EFD": partial(nn.init.trunc_normal_, std=0.02),
8888
}
8989

9090

@@ -276,7 +276,7 @@ def build_mla_moe_layers(
276276
dim=dim,
277277
hidden_dim=dense_hidden_dim,
278278
w1_param_init=linear_init,
279-
w2w3_param_init=depth_init(layer_id),
279+
w2_param_init=depth_init(layer_id),
280280
)
281281
moe_cfg = None
282282
else:
@@ -303,7 +303,7 @@ def build_mla_moe_layers(
303303
dim=dim,
304304
hidden_dim=moe_hidden_dim * num_shared_experts,
305305
w1_param_init=linear_init,
306-
w2w3_param_init=depth_init(layer_id),
306+
w2_param_init=depth_init(layer_id),
307307
),
308308
aux_loss_coeff=aux_loss_coeff,
309309
)

0 commit comments

Comments
 (0)