|
10 | 10 | import torch |
11 | 11 | import torch.nn.functional as F |
12 | 12 |
|
| 13 | +import torchtitan.models.llama3.flavors as llama3_flavors |
| 14 | + |
13 | 15 | 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 | +) |
15 | 21 | from torchtitan.models.common.feed_forward import FeedForward |
16 | 22 | from torchtitan.models.common.linear import Linear |
17 | 23 |
|
@@ -59,6 +65,62 @@ def test_feed_forward_uses_one_physical_gate_up_linear(): |
59 | 65 | torch.testing.assert_close(w13_2HD[1], 5 * torch.ones_like(w13_2HD[1])) |
60 | 66 |
|
61 | 67 |
|
| 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 | + |
62 | 124 | def test_feed_forward_requires_two_w13_projections(): |
63 | 125 | config = FeedForward.Config( |
64 | 126 | w13=Linear.Config(in_features=4, out_features=8), |
|
0 commit comments