From ea43e4fb674888c82bf791ed53454ceb67ee09ce Mon Sep 17 00:00:00 2001 From: Jin Soo Ihm Date: Mon, 31 Aug 2026 17:45:55 -0700 Subject: [PATCH 1/2] Fixed typecheck --- tests/integration_tests/models.py | 6 ++++++ torchtitan/models/kimi_k2_7/vision_encoder.py | 8 ++++++-- torchtitan/models/qwen3_5/gdn.py | 4 ++-- torchtitan/models/qwen3_5/model.py | 7 ++++--- torchtitan/models/qwen3_5/vision_encoder.py | 8 ++++++-- torchtitan_recipes/tests/models.py | 11 +++++++++++ 6 files changed, 35 insertions(+), 9 deletions(-) diff --git a/tests/integration_tests/models.py b/tests/integration_tests/models.py index 139340b3cb..4626bf922f 100755 --- a/tests/integration_tests/models.py +++ b/tests/integration_tests/models.py @@ -115,6 +115,12 @@ def build_model_tests_list() -> list[OverrideDefinitions]: ngpu=8, ), # Integration Test Cases for Qwen3.5 + OverrideDefinitions( + configs=[recipes.qwen35_debugmodel_fsdp2_tp2], + test_descr="Qwen3.5 FSDP+TP", + test_name="qwen3_5_fsdp+tp", + ngpu=4, + ), OverrideDefinitions( configs=[recipes.qwen35_debugmodel_moe_fsdp2_tp2_pp2_ep4], test_descr="Qwen3.5 MoE FSDP+TP+EP+PP", diff --git a/torchtitan/models/kimi_k2_7/vision_encoder.py b/torchtitan/models/kimi_k2_7/vision_encoder.py index 7041ed9182..d77d9c6812 100644 --- a/torchtitan/models/kimi_k2_7/vision_encoder.py +++ b/torchtitan/models/kimi_k2_7/vision_encoder.py @@ -121,7 +121,9 @@ def _compute_learned_pos_embeds( packed_pos = torch.cat([pos[i] for i in range(len(grids))], dim=0) if get_spmd_backend() == "spmd_types" and spmd.is_type_checking(): packed_pos = spmd.mutate_type( - packed_pos, src=spmd.R, dst={"dp": spmd.V, "tp": spmd.I} + packed_pos, + src={"dp": spmd.R, "tp": spmd.I}, + dst={"dp": spmd.V, "tp": spmd.I}, ) return packed_pos @@ -183,7 +185,9 @@ def _compute_2d_rope_cache( packed_angles = torch.cat([angles[i] for i in range(len(grids))], dim=0) if get_spmd_backend() == "spmd_types" and spmd.is_type_checking(): packed_angles = spmd.mutate_type( - packed_angles, src=spmd.R, dst={"dp": spmd.V, "tp": spmd.I} + packed_angles, + src={"dp": spmd.R, "tp": spmd.I}, + dst={"dp": spmd.V, "tp": spmd.I}, ) return torch.polar(torch.ones_like(packed_angles), packed_angles).unsqueeze(1) diff --git a/torchtitan/models/qwen3_5/gdn.py b/torchtitan/models/qwen3_5/gdn.py index 6223576e6d..4c2fdacf0f 100644 --- a/torchtitan/models/qwen3_5/gdn.py +++ b/torchtitan/models/qwen3_5/gdn.py @@ -21,7 +21,7 @@ from torchtitan.distributed.utils import is_in_batch_invariant_mode from torchtitan.models.common import Conv1d, Linear -from torchtitan.models.common.attention import VarlenMetadata +from torchtitan.models.common.attention import local_head_split, VarlenMetadata from torchtitan.protocols.module import Module @@ -442,7 +442,7 @@ def forward( key_head_dim=self.key_head_dim, value_head_dim=self.value_head_dim, ) - gate_THV = gate_TC.view(num_tokens, -1, self.value_head_dim) + gate_THV = local_head_split(gate_TC, self.value_head_dim) output_THV = self.norm(output_THV, gate_THV) out_TD = output_THV.reshape(num_tokens, -1) return self.out_proj(out_TD) diff --git a/torchtitan/models/qwen3_5/model.py b/torchtitan/models/qwen3_5/model.py index 1a88c714cc..e88887cb8e 100644 --- a/torchtitan/models/qwen3_5/model.py +++ b/torchtitan/models/qwen3_5/model.py @@ -27,6 +27,7 @@ BaseAttention, create_varlen_metadata_for_document, FlexAttention, + local_head_split, VarlenAttention, VarlenMetadata, ) @@ -144,10 +145,10 @@ def forward( num_tokens = x_TD.shape[0] # wq is 2x wider: produces query + gate - xq_gate_THC = self.wq(x_TD).view(num_tokens, -1, self.head_dim * 2) + xq_gate_THC = local_head_split(self.wq(x_TD), self.head_dim * 2) xq_THK, gate_THV = xq_gate_THC.chunk(2, dim=-1) - xk_THK = self.wk(x_TD).view(num_tokens, -1, self.head_dim) - xv_THV = self.wv(x_TD).view(num_tokens, -1, self.head_dim) + xk_THK = local_head_split(self.wk(x_TD), self.head_dim) + xv_THV = local_head_split(self.wv(x_TD), self.head_dim) # QK norm (before RoPE) xq_THK = self.q_norm(xq_THK) diff --git a/torchtitan/models/qwen3_5/vision_encoder.py b/torchtitan/models/qwen3_5/vision_encoder.py index 68477d3db8..9cab725650 100644 --- a/torchtitan/models/qwen3_5/vision_encoder.py +++ b/torchtitan/models/qwen3_5/vision_encoder.py @@ -118,7 +118,9 @@ def _compute_learned_pos_embeds( packed_pos_embeds = torch.cat([pos_embeds[i] for i in range(len(grids))], dim=0) if get_spmd_backend() == "spmd_types" and spmd.is_type_checking(): packed_pos_embeds = spmd.mutate_type( - packed_pos_embeds, src=spmd.R, dst={"dp": spmd.V, "tp": spmd.I} + packed_pos_embeds, + src={"dp": spmd.R, "tp": spmd.I}, + dst={"dp": spmd.V, "tp": spmd.I}, ) return packed_pos_embeds @@ -212,7 +214,9 @@ def _compute_2d_rope_cache( packed_rope_embeds = torch.cat([rope_embeds[i] for i in range(len(grids))], dim=0) if get_spmd_backend() == "spmd_types" and spmd.is_type_checking(): packed_rope_embeds = spmd.mutate_type( - packed_rope_embeds, src=spmd.R, dst={"dp": spmd.V, "tp": spmd.I} + packed_rope_embeds, + src={"dp": spmd.R, "tp": spmd.I}, + dst={"dp": spmd.V, "tp": spmd.I}, ) packed_rope_embeds = torch.cat((packed_rope_embeds, packed_rope_embeds), dim=-1) rope_cache = torch.cat( diff --git a/torchtitan_recipes/tests/models.py b/torchtitan_recipes/tests/models.py index e8e9135596..3532200148 100644 --- a/torchtitan_recipes/tests/models.py +++ b/torchtitan_recipes/tests/models.py @@ -185,6 +185,17 @@ def qwen3_debugmodel_non_fused_qkv_fsdp2_tp2_cp2() -> Trainer.Config: return config +def qwen35_debugmodel_fsdp2_tp2() -> Trainer.Config: + from torchtitan.models.qwen3_5.config_registry import qwen35_debugmodel + + config = qwen35_debugmodel() + _use_spmd_types(config, typechecking=True) + config.parallelism.data_parallel_shard_degree = 2 + config.parallelism.tensor_parallel_degree = 2 + config.training.disable_cuda_graphs = True + return config + + def qwen35_debugmodel_moe_fsdp2_tp2_pp2_ep4() -> Trainer.Config: from torchtitan.models.qwen3_5.config_registry import qwen35_debugmodel_moe From 7ff1e1c415f469aba81923d69da0739740fcbbc6 Mon Sep 17 00:00:00 2001 From: Jin Soo Ihm Date: Thu, 3 Sep 2026 13:09:51 -0700 Subject: [PATCH 2/2] update --- torchtitan/models/kimi_k2_7/vision_encoder.py | 10 ++++++---- torchtitan/models/qwen3_5/vision_encoder.py | 10 ++++++---- 2 files changed, 12 insertions(+), 8 deletions(-) diff --git a/torchtitan/models/kimi_k2_7/vision_encoder.py b/torchtitan/models/kimi_k2_7/vision_encoder.py index d77d9c6812..cca9fb580f 100644 --- a/torchtitan/models/kimi_k2_7/vision_encoder.py +++ b/torchtitan/models/kimi_k2_7/vision_encoder.py @@ -122,8 +122,9 @@ def _compute_learned_pos_embeds( if get_spmd_backend() == "spmd_types" and spmd.is_type_checking(): packed_pos = spmd.mutate_type( packed_pos, - src={"dp": spmd.R, "tp": spmd.I}, - dst={"dp": spmd.V, "tp": spmd.I}, + "dp", + src=spmd.R, + dst=spmd.V, ) return packed_pos @@ -186,8 +187,9 @@ def _compute_2d_rope_cache( if get_spmd_backend() == "spmd_types" and spmd.is_type_checking(): packed_angles = spmd.mutate_type( packed_angles, - src={"dp": spmd.R, "tp": spmd.I}, - dst={"dp": spmd.V, "tp": spmd.I}, + "dp", + src=spmd.R, + dst=spmd.V, ) return torch.polar(torch.ones_like(packed_angles), packed_angles).unsqueeze(1) diff --git a/torchtitan/models/qwen3_5/vision_encoder.py b/torchtitan/models/qwen3_5/vision_encoder.py index 9cab725650..ddf4edcecb 100644 --- a/torchtitan/models/qwen3_5/vision_encoder.py +++ b/torchtitan/models/qwen3_5/vision_encoder.py @@ -119,8 +119,9 @@ def _compute_learned_pos_embeds( if get_spmd_backend() == "spmd_types" and spmd.is_type_checking(): packed_pos_embeds = spmd.mutate_type( packed_pos_embeds, - src={"dp": spmd.R, "tp": spmd.I}, - dst={"dp": spmd.V, "tp": spmd.I}, + "dp", + src=spmd.R, + dst=spmd.V, ) return packed_pos_embeds @@ -215,8 +216,9 @@ def _compute_2d_rope_cache( if get_spmd_backend() == "spmd_types" and spmd.is_type_checking(): packed_rope_embeds = spmd.mutate_type( packed_rope_embeds, - src={"dp": spmd.R, "tp": spmd.I}, - dst={"dp": spmd.V, "tp": spmd.I}, + "dp", + src=spmd.R, + dst=spmd.V, ) packed_rope_embeds = torch.cat((packed_rope_embeds, packed_rope_embeds), dim=-1) rope_cache = torch.cat(