Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 6 additions & 0 deletions tests/integration_tests/models.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
10 changes: 8 additions & 2 deletions torchtitan/models/kimi_k2_7/vision_encoder.py
Original file line number Diff line number Diff line change
Expand Up @@ -121,7 +121,10 @@ 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,
"dp",
src=spmd.R,
dst=spmd.V,
)
return packed_pos

Expand Down Expand Up @@ -183,7 +186,10 @@ 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,
"dp",
src=spmd.R,
dst=spmd.V,
)
return torch.polar(torch.ones_like(packed_angles), packed_angles).unsqueeze(1)

Expand Down
4 changes: 2 additions & 2 deletions torchtitan/models/qwen3_5/gdn.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down Expand Up @@ -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)
7 changes: 4 additions & 3 deletions torchtitan/models/qwen3_5/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,7 @@
BaseAttention,
create_varlen_metadata_for_document,
FlexAttention,
local_head_split,
VarlenAttention,
VarlenMetadata,
)
Expand Down Expand Up @@ -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)
Expand Down
10 changes: 8 additions & 2 deletions torchtitan/models/qwen3_5/vision_encoder.py
Original file line number Diff line number Diff line change
Expand Up @@ -118,7 +118,10 @@ 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,
"dp",
src=spmd.R,
dst=spmd.V,
)
return packed_pos_embeds

Expand Down Expand Up @@ -212,7 +215,10 @@ 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,
"dp",
src=spmd.R,
dst=spmd.V,
)
packed_rope_embeds = torch.cat((packed_rope_embeds, packed_rope_embeds), dim=-1)
rope_cache = torch.cat(
Expand Down
11 changes: 11 additions & 0 deletions torchtitan_recipes/tests/models.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
Loading