From c8875f8125b16b145fdbf9c61e4ce64450fc13a4 Mon Sep 17 00:00:00 2001 From: Pian Pawakapan Date: Thu, 6 Aug 2026 16:35:11 -0700 Subject: [PATCH 1/7] Update (base update) [ghstack-poisoned] --- .ci/docker/requirements.txt | 2 +- pyproject.toml | 2 +- tests/unit_tests/test_parallel_dims.py | 58 +++++--- torchtitan/distributed/fsdp.py | 7 +- torchtitan/distributed/parallel_dims.py | 29 ++-- torchtitan/distributed/spmd_types.py | 37 ++--- torchtitan/models/common/vision_encoder.py | 16 ++- torchtitan/models/qwen3_5/__init__.py | 18 ++- torchtitan/models/qwen3_5/model.py | 138 +++++++++++++----- torchtitan/models/qwen3_5/parallelize.py | 45 ++++-- torchtitan/models/qwen3_5/rope.py | 1 - torchtitan/models/qwen3_5/sharding.py | 147 ++++++++++++++++---- torchtitan/models/qwen3_5/vision_encoder.py | 29 ++-- torchtitan/protocols/sharding.py | 19 ++- 14 files changed, 396 insertions(+), 152 deletions(-) diff --git a/.ci/docker/requirements.txt b/.ci/docker/requirements.txt index 4838bf60e8..5e4ea5f319 100644 --- a/.ci/docker/requirements.txt +++ b/.ci/docker/requirements.txt @@ -7,4 +7,4 @@ tokenizers >= 0.15.0 safetensors einops pillow -spmd_types==0.2.1 +spmd_types==0.2.3 diff --git a/pyproject.toml b/pyproject.toml index c45098ef2c..dff1133956 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -24,7 +24,7 @@ dependencies = [ "wandb", "einops", "pillow", - "spmd_types==0.2.1", + "spmd_types==0.2.3", ] dynamic = ["version"] diff --git a/tests/unit_tests/test_parallel_dims.py b/tests/unit_tests/test_parallel_dims.py index ad68dcac94..affa18990d 100644 --- a/tests/unit_tests/test_parallel_dims.py +++ b/tests/unit_tests/test_parallel_dims.py @@ -12,7 +12,7 @@ import torch import torch.distributed as dist from torch.distributed.device_mesh import init_device_mesh -from torch.distributed.tensor import Shard +from torch.distributed.tensor import Replicate from torch.distributed.tensor.debug import CommDebugMode from torch.testing._internal.distributed._tensor.common_dtensor import ( DTensorTestBase, @@ -28,7 +28,6 @@ ) from torchtitan.distributed.spmd_types import ( spmd_distribute_tensor, - spmd_layout_to_dtensor_placements, spmd_redistribute_per_axis, spmd_validate_redistributions, ) @@ -37,7 +36,7 @@ dense_sequence_parallel_placement, ) from torchtitan.models.llama3 import model_registry -from torchtitan.protocols.sharding import ShardingConfig +from torchtitan.protocols.sharding import resolve_placements, ShardingConfig class TestParallelDimsValidation(unittest.TestCase): @@ -236,19 +235,6 @@ class TestSpmdLayout(DTensorTestBase): def world_size(self): return 4 - def test_converts_partition_spec_to_dtensor_shard(self): - """PartitionSpec refines V into concrete DTensor Shard placement.""" - layout = SpmdLayout( - {MeshAxisName.TP: spmd.V}, - partition_spec=spmd.PartitionSpec(MeshAxisName.TP), - ) - - self.assertEqual(layout.per_axis_spmd_types(), {MeshAxisName.TP: spmd.S(0)}) - self.assertEqual( - spmd_layout_to_dtensor_placements(layout), - {MeshAxisName.TP: Shard(0)}, - ) - def test_seq_parallel_activation_per_axis_spmd_types(self): """PartitionSpec can map multiple mesh axes to one tensor dim.""" layout = SpmdLayout( @@ -280,6 +266,21 @@ def test_unfold_dp_axes(self): ["dp_replicate", "dp_shard", "cp", "tp"], ) + @with_comms + def test_resolve_placements_ignores_extra_untranslatable_axes(self): + """Extra layout axes are ignored before converting to DTensor placements.""" + mesh = init_device_mesh( + self.device_type, (self.world_size,), mesh_dim_names=("tp",) + ) + layout = SpmdLayout( + { + MeshAxisName.DP: spmd.V, + MeshAxisName.TP: spmd.I, + } + ) + + self.assertEqual(resolve_placements(layout, mesh), (Replicate(),)) + def test_rejects_partition_spec_reorder_redistribute(self): """((DP, CP), None) -> ((CP, DP), None) not supported by a single redistribute call.""" with self.assertRaises(ValueError) as cm: @@ -332,6 +333,31 @@ def test_rejects_multi_axis_redistribute(self): ) ) + def test_rejects_redistribute_from_varying(self): + for src_dp, dst_dp in ((spmd.V, spmd.R), (spmd.R, spmd.V)): + with self.subTest(src_dp=src_dp, dst_dp=dst_dp): + with self.assertRaisesRegex( + ValueError, + "output: SpmdLayout-based redistribution changes mesh axis " + "'dp' with spmd.V as the source or destination type", + ): + spmd_validate_redistributions( + ShardingConfig( + out_src_shardings=SpmdLayout( + { + MeshAxisName.DP: src_dp, + MeshAxisName.TP: spmd.I, + } + ), + out_dst_shardings=SpmdLayout( + { + MeshAxisName.DP: dst_dp, + MeshAxisName.TP: spmd.I, + } + ), + ) + ) + @with_comms def test_partition_spec_order_controls_state_shard(self): """Test spmd_distribute_tensor follows PartitionSpec order. diff --git a/torchtitan/distributed/fsdp.py b/torchtitan/distributed/fsdp.py index 9b2d47e093..ddce16cf25 100644 --- a/torchtitan/distributed/fsdp.py +++ b/torchtitan/distributed/fsdp.py @@ -85,6 +85,7 @@ def apply_fsdp_to_vision_encoder( reduce_dtype: torch.dtype, reshard_after_forward_policy: str = "default", pp_enabled: bool = False, + dp_mesh_dims: "DataParallelMeshDims | None" = None, ) -> None: """FSDP a VLM vision encoder as a single unit. @@ -96,10 +97,12 @@ def apply_fsdp_to_vision_encoder( reshard_after_forward = get_fsdp_reshard_after_forward_policy( reshard_after_forward_policy, pp_enabled=pp_enabled ) + fsdp_config: dict[str, Any] = {"mesh": dp_mesh, "mp_policy": mp_policy} + if dp_mesh_dims is not None: + fsdp_config["dp_mesh_dims"] = dp_mesh_dims fully_shard( vision_encoder, - mesh=dp_mesh, - mp_policy=mp_policy, + **fsdp_config, reshard_after_forward=reshard_after_forward, ) diff --git a/torchtitan/distributed/parallel_dims.py b/torchtitan/distributed/parallel_dims.py index 34d7b58009..7c5fedcfd6 100644 --- a/torchtitan/distributed/parallel_dims.py +++ b/torchtitan/distributed/parallel_dims.py @@ -19,7 +19,13 @@ from torchtitan.tools.utils import device_type -__all__ = ["MeshAxisName", "ParallelDims", "SpmdLayout", "unfold_dp_axes"] +__all__ = [ + "MeshAxisName", + "ParallelDims", + "SpmdLayout", + "unfold_dp_axis", + "unfold_dp_axes", +] class StrEnum(str, Enum): @@ -115,16 +121,21 @@ def per_axis_spmd_types(self) -> dict[MeshAxisName, spmd.PerMeshAxisSpmdType]: return result +def unfold_dp_axis(axis: MeshAxisName | str) -> tuple[MeshAxisName, ...]: + """Expand logical ``dp`` into concrete dense storage mesh axes.""" + axis_name = MeshAxisName(axis) + if axis_name == MeshAxisName.DP: + return (MeshAxisName.DP_REPLICATE, MeshAxisName.DP_SHARD) + return (axis_name,) + + def unfold_dp_axes(axes: Iterable[MeshAxisName | str]) -> list[str]: """Expand logical ``dp`` into concrete dense storage mesh axes.""" - result: list[str] = [] - for axis in axes: - axis_value = axis.value if isinstance(axis, MeshAxisName) else axis - if axis_value == "dp": - result.extend(("dp_replicate", "dp_shard")) - else: - result.append(axis_value) - return result + return [ + concrete_axis.value + for axis in axes + for concrete_axis in unfold_dp_axis(axis) + ] @dataclass diff --git a/torchtitan/distributed/spmd_types.py b/torchtitan/distributed/spmd_types.py index e0e2b50e62..35d2f3a0e2 100644 --- a/torchtitan/distributed/spmd_types.py +++ b/torchtitan/distributed/spmd_types.py @@ -16,7 +16,6 @@ import spmd_types as spmd import torch from torch.distributed.device_mesh import DeviceMesh -from torch.distributed.tensor import Partial, Placement, Replicate, Shard from torchtitan.distributed.utils import get_spmd_backend @@ -40,7 +39,6 @@ "set_current_spmd_mesh", "set_spmd_meshes", "maybe_set_sparse_mesh", - "spmd_layout_to_dtensor_placements", ] @@ -134,30 +132,6 @@ def maybe_set_sparse_mesh() -> Iterator[None]: yield -def spmd_layout_to_dtensor_placements( - layout: "SpmdLayout", -) -> dict["MeshAxisName", Placement]: - """Convert an SPMD layout to DTensor placements keyed by mesh axis name.""" - from torchtitan.distributed.parallel_dims import MeshAxisName - - result: dict[MeshAxisName, Placement] = {} - for axis_name, axis_type in layout.per_axis_spmd_types().items(): - if axis_type == spmd.R or axis_type == spmd.I: - dtensor_placement: Placement = Replicate() - elif axis_type == spmd.P: - dtensor_placement = Partial() - else: - assert isinstance(axis_type, spmd.Shard) - dtensor_placement = Shard(axis_type.dim) - - if axis_name == MeshAxisName.DP: - result[MeshAxisName.DP_REPLICATE] = dtensor_placement - result[MeshAxisName.DP_SHARD] = dtensor_placement - else: - result[axis_name] = dtensor_placement - return result - - def annotate_input_spmd_types( parallel_dims: "ParallelDims", inputs: torch.Tensor, @@ -259,6 +233,17 @@ def _validate_redistribute_spmd_pair( "spmd_redistribute_per_axis only supports one single-axis " "redistribution." ) + if changed_axes and ( + src_types[changed_axes[0]] is spmd.V + or dst_types[changed_axes[0]] is spmd.V + ): + axis = changed_axes[0] + raise ValueError( + f"{name}: SpmdLayout-based redistribution changes mesh axis " + f"{axis.value!r} with spmd.V as the source or destination type. " + "Config-based redistribution requires non-V types; write an " + "explicit collective when the value semantics are unclear." + ) # 2) If neither has PartitionSpec, comparing per_axis_spmd_types() is sufficient. if src.partition_spec is None and dst.partition_spec is None: diff --git a/torchtitan/models/common/vision_encoder.py b/torchtitan/models/common/vision_encoder.py index 64ac8ec89a..41c37035be 100644 --- a/torchtitan/models/common/vision_encoder.py +++ b/torchtitan/models/common/vision_encoder.py @@ -24,6 +24,7 @@ from collections.abc import Callable from dataclasses import dataclass, field +import spmd_types as spmd import torch from torch.nn.attention.flex_attention import BlockMask, create_block_mask @@ -40,6 +41,15 @@ ] +@spmd.local_map( + in_types=(spmd.PartitionSpec("dp", None, "tp"), None), + out_types=spmd.PartitionSpec("dp", None, "tp", None), +) +def _local_head_split(t: torch.Tensor, head_dim: int) -> torch.Tensor: + # TODO(pianpwk): Remove once spmd_types tracks sharding evenness. + return t.view(t.shape[0], t.shape[1], -1, head_dim) + + def get_vision_block_mask_mod(num_patches: torch.Tensor) -> Callable: """Block-diagonal mask: each visual item attends only to its own patches. @@ -122,9 +132,9 @@ def forward( # -1 infers the head count locally (= num_heads / TP under tensor # parallelism, where wq/wk/wv are colwise-sharded). - q_NPHDh = self.wq(x).view(N, P, -1, self.head_dim) - k_NPHDh = self.wk(x).view(N, P, -1, self.head_dim) - v_NPHDh = self.wv(x).view(N, P, -1, self.head_dim) + q_NPHDh = _local_head_split(self.wq(x), self.head_dim) + k_NPHDh = _local_head_split(self.wk(x), self.head_dim) + v_NPHDh = _local_head_split(self.wv(x), self.head_dim) q_NPHDh, k_NPHDh = rope_apply(q_NPHDh, k_NPHDh, rope_cache) diff --git a/torchtitan/models/qwen3_5/__init__.py b/torchtitan/models/qwen3_5/__init__.py index 363cbdcc20..75481b9448 100644 --- a/torchtitan/models/qwen3_5/__init__.py +++ b/torchtitan/models/qwen3_5/__init__.py @@ -16,6 +16,7 @@ Conv1d, Embedding, Linear, + ScaledBiasRowwiseLinear, SigmoidGatedFeedForward, ) from torchtitan.models.common.config_utils import ( @@ -117,6 +118,17 @@ def _linear(in_features: int, out_features: int) -> Linear.Config: ) +def _scaled_bias_rowwise_linear( + in_features: int, out_features: int +) -> ScaledBiasRowwiseLinear.Config: + return ScaledBiasRowwiseLinear.Config( + in_features=in_features, + out_features=out_features, + bias=True, + param_init=_LINEAR_INIT, + ) + + def _offset_norm(dim: int) -> OffsetRMSNorm.Config: return OffsetRMSNorm.Config(dim=dim, eps=_EPS, param_init=_OFFSET_NORM_INIT) @@ -178,11 +190,11 @@ def _qwen35_vision_encoder_config( wq=_linear(dim, dim), wk=_linear(dim, dim), wv=_linear(dim, dim), - proj=_linear(dim, dim), + proj=_scaled_bias_rowwise_linear(dim, dim), ), mlp=VisionMLP.Config( fc1=_linear(dim, ffn_dim), - fc2=_linear(ffn_dim, dim), + fc2=_scaled_bias_rowwise_linear(ffn_dim, dim), ), ), rotary_pos_emb=VisionRotaryEmbedding.Config( @@ -193,7 +205,7 @@ def _qwen35_vision_encoder_config( merged_hidden_size=merged_hidden_size, norm=LayerNorm.Config(normalized_shape=dim, eps=layer_norm_eps), fc1=_linear(merged_hidden_size, merged_hidden_size), - fc2=_linear(merged_hidden_size, out_hidden_size), + fc2=_scaled_bias_rowwise_linear(merged_hidden_size, out_hidden_size), ), param_init=_POS_EMBED_INIT, ) diff --git a/torchtitan/models/qwen3_5/model.py b/torchtitan/models/qwen3_5/model.py index 64f001841f..9a381b7283 100644 --- a/torchtitan/models/qwen3_5/model.py +++ b/torchtitan/models/qwen3_5/model.py @@ -5,21 +5,27 @@ # LICENSE file in the root directory of this source tree. +import contextlib from collections.abc import Callable from dataclasses import dataclass from typing import Literal +import spmd_types as spmd import torch import torch.nn.functional as F +from fla.modules.conv.triton.ops import CausalConv1dFunction from fla.ops.gated_delta_rule import ( chunk_gated_delta_rule as _fla_chunk_gated_delta_rule, fused_recurrent_gated_delta_rule as _fla_fused_recurrent_gated_delta_rule, ) +from fla.ops.gated_delta_rule.chunk import ChunkGatedDeltaRuleFunction +from fla.ops.gated_delta_rule.fused_recurrent import FusedRecurrentFunction from torch import nn from torch.distributed.tensor import DTensor from torch.distributed.tensor.experimental import local_map +from torchtitan.distributed.utils import get_spmd_backend from torchtitan.models.common import Conv1d, Linear from torchtitan.models.common.attention import ( AttentionMasksType, @@ -37,12 +43,25 @@ from torchtitan.protocols.module import Module from .rope import MRoPE -from .sharding import set_qwen35_sharding_config +from .sharding import annotate_multimodal_input_spmd_types, set_qwen35_sharding_config from .vision_encoder import Qwen35VisionEncoder GatedDeltaBackend = Literal["fla_chunked", "fla_fused_recurrent"] +spmd.register_local_autograd_function(ChunkGatedDeltaRuleFunction) +spmd.register_local_autograd_function(FusedRecurrentFunction) +spmd.register_local_autograd_function(CausalConv1dFunction) + +@spmd.local_map( + in_types=( + {"dp": spmd.S(0), "tp": spmd.S(2)}, + {"dp": spmd.R, "tp": spmd.S(0)}, + {"dp": spmd.V, "tp": spmd.R}, + {"dp": spmd.V, "tp": spmd.R}, + ), + out_types={"dp": spmd.S(0), "tp": spmd.S(2)}, +) def _causal_conv1d_varlen( x_BTD: torch.Tensor, weight: torch.Tensor, @@ -206,6 +225,15 @@ def forward( return result[0] +@spmd.local_map( + in_types=(spmd.PartitionSpec("dp", None, "tp"), None), + out_types=spmd.PartitionSpec("dp", None, "tp", None), +) +def _local_head_split(t: torch.Tensor, head_dim: int) -> torch.Tensor: + # TODO(pianpwk): this should be doable once spmd_types tracks sharding evenness. + return t.view(t.shape[0], t.shape[1], -1, head_dim) + + class GatedDeltaNet(Module): """Gated DeltaNet linear attention. @@ -316,6 +344,13 @@ def _conv_varlen( ) return self._local_map_conv(x_BLD, conv, _conv_varlen, cu_seqlens) + if get_spmd_backend() == "spmd_types": + return _causal_conv1d_varlen( + x_BLD, + conv.weight, + cu_seqlens, + cu_seqlens_cpu, + ) return _causal_conv1d_varlen( x_BLD, conv.weight, @@ -324,22 +359,32 @@ def _conv_varlen( ) x_BDL = F.pad(x_BLD.transpose(1, 2), [self.conv_kernel_size - 1, 0]) + + def _conv(x_local_BDL: torch.Tensor, w_local: torch.Tensor) -> torch.Tensor: + # groups == local out-channels for depthwise channel-sharded conv. + return F.conv1d( + x_local_BDL, + w_local, + None, + conv.stride, + conv.padding, + conv.dilation, + w_local.size(0), + ) + if isinstance(x_BDL, DTensor): # TODO: Remove once the DTensor Conv1d dispatch fix for sharded # groups lands in a released torch. - def _conv(x_local_BDL: torch.Tensor, w_local: torch.Tensor) -> torch.Tensor: - # groups == local out-channels (depthwise, channel-sharded) - return F.conv1d( - x_local_BDL, - w_local, - None, - conv.stride, - conv.padding, - conv.dilation, - w_local.size(0), - ) - x_BDL = self._local_map_conv(x_BDL, conv, _conv) + elif get_spmd_backend() == "spmd_types": + conv_spmd = spmd.local_map( + in_types=( + {"dp": spmd.S(0), "tp": spmd.S(1)}, + {"dp": spmd.R, "tp": spmd.S(0)}, + ), + out_types={"dp": spmd.S(0), "tp": spmd.S(1)}, + )(_conv) + x_BDL = conv_spmd(x_BDL, conv.weight) else: x_BDL = conv(x_BDL) return F.silu(x_BDL).transpose(1, 2) @@ -369,11 +414,6 @@ def forward( device="cpu", ) - if cu_seqlens is None: - kernel_B, kernel_L = B, L - else: - kernel_B, kernel_L = 1, B * L - def _maybe_flatten(tensor: torch.Tensor) -> torch.Tensor: if cu_seqlens is None: return tensor @@ -388,21 +428,24 @@ def _maybe_flatten(tensor: torch.Tensor) -> torch.Tensor: self.conv_q, cu_seqlens, cu_seqlens_cpu, - ).view(kernel_B, kernel_L, -1, self.key_head_dim) + ) + xq_BLNK = _local_head_split(xq_BLNK, self.key_head_dim) xk_BLNK = self._causal_conv( _maybe_flatten(self.in_proj_k(x_BLD)), self.conv_k, cu_seqlens, cu_seqlens_cpu, - ).view(kernel_B, kernel_L, -1, self.key_head_dim) + ) + xk_BLNK = _local_head_split(xk_BLNK, self.key_head_dim) xv_BLNV = self._causal_conv( _maybe_flatten(self.in_proj_v(x_BLD)), self.conv_v, cu_seqlens, cu_seqlens_cpu, - ).view(kernel_B, kernel_L, -1, self.value_head_dim) - xz_BLNV = _maybe_flatten(self.in_proj_z(x_BLD)).view( - kernel_B, kernel_L, -1, self.value_head_dim + ) + xv_BLNV = _local_head_split(xv_BLNV, self.value_head_dim) + xz_BLNV = _local_head_split( + _maybe_flatten(self.in_proj_z(x_BLD)), self.value_head_dim ) xa_BLN = _maybe_flatten(self.in_proj_a(x_BLD)) xb_BLN = _maybe_flatten(self.in_proj_b(x_BLD)) @@ -493,10 +536,10 @@ def forward( B, L, _ = x_BLD.shape # wq is 2x wider: produces query + gate - xq_gate_BLN2H = self.wq(x_BLD).view(B, L, -1, self.head_dim * 2) + xq_gate_BLN2H = _local_head_split(self.wq(x_BLD), self.head_dim * 2) xq_BLNH, gate_BLNH = xq_gate_BLN2H.chunk(2, dim=-1) - xk_BLNH = self.wk(x_BLD).view(B, L, -1, self.head_dim) - xv_BLNH = self.wv(x_BLD).view(B, L, -1, self.head_dim) + xk_BLNH = _local_head_split(self.wk(x_BLD), self.head_dim) + xv_BLNH = _local_head_split(self.wv(x_BLD), self.head_dim) # QK norm (before RoPE) xq_BLNH = self.q_norm(xq_BLNH) @@ -695,6 +738,12 @@ def __init__(self, config: Config): self.vision_encoder = config.vision_encoder.build() self.spatial_merge_size = config.vision_encoder.spatial_merge_size + def multimodal_context(self) -> contextlib.AbstractContextManager[None]: + """Use local DP typechecking while preparing multimodal inputs.""" + if get_spmd_backend() == "spmd_types": + return spmd.set_current_mesh(local_axes=("dp",)) + return contextlib.nullcontext() + def get_attention_masks( self, positions: torch.Tensor, @@ -841,17 +890,32 @@ def forward( # pyrefly: ignore [bad-override] mrope_positions: torch.Tensor | None = None, special_tokens: dict[str, int] | None = None, ): - if self.tok_embeddings is not None: - x = self._prepare_multimodal_embeds( - tokens, - pixel_values=pixel_values, - pixel_values_videos=pixel_values_videos, - grid_thw=grid_thw, - grid_thw_videos=grid_thw_videos, - special_tokens=special_tokens, # pyrefly: ignore [bad-argument-type] - ) - else: - x = tokens + with self.multimodal_context(): + if get_spmd_backend() == "spmd_types": + annotate_multimodal_input_spmd_types( + mrope_positions=mrope_positions, + pixel_values=pixel_values, + pixel_values_videos=pixel_values_videos, + grid_thw=grid_thw, + grid_thw_videos=grid_thw_videos, + ) + + if self.tok_embeddings is not None: + x = self._prepare_multimodal_embeds( + tokens, + pixel_values=pixel_values, + pixel_values_videos=pixel_values_videos, + grid_thw=grid_thw, + grid_thw_videos=grid_thw_videos, + special_tokens=special_tokens, # pyrefly: ignore [bad-argument-type] + ) + else: + x = tokens + + if get_spmd_backend() == "spmd_types": + # The scatter restores a token-aligned tensor, so text-model DP + # resumes as global batch sharding after the multimodal region. + spmd.assert_type(x, {"dp": spmd.S(0), "tp": spmd.R}) # 3D MRoPE positions for multimodal batches, else 2D text positions. rope_positions = mrope_positions if mrope_positions is not None else positions diff --git a/torchtitan/models/qwen3_5/parallelize.py b/torchtitan/models/qwen3_5/parallelize.py index 666c22db48..872bca974e 100644 --- a/torchtitan/models/qwen3_5/parallelize.py +++ b/torchtitan/models/qwen3_5/parallelize.py @@ -27,6 +27,11 @@ apply_fsdp_to_decoder, apply_fsdp_to_vision_encoder, ) +from torchtitan.distributed.full_dtensor import ( + resolve_fsdp_mesh, + resolve_sparse_fsdp_mesh, + validate_config, +) from torchtitan.distributed.tensor_parallel import maybe_enable_async_tp @@ -65,6 +70,11 @@ def parallelize_qwen3_5( if parallelism.enable_async_tensor_parallel and not model_compile_enabled: raise RuntimeError("Async TP requires torch.compile") + if parallelism.spmd_backend == "spmd_types": + validate_config(parallel_dims, model) + # pyrefly: ignore [not-callable] + model.parallelize(parallel_dims) + elif parallel_dims.tp_enabled or parallel_dims.ep_enabled: # pyrefly: ignore [not-callable] model.parallelize(parallel_dims) @@ -83,10 +93,24 @@ def parallelize_qwen3_5( # pyrefly: ignore [bad-argument-type] apply_compile(model.vision_encoder, compile_config) - dp_mesh_names = ( - ["dp_replicate", "fsdp"] if parallel_dims.dp_replicate_enabled else ["fsdp"] - ) - dp_mesh = parallel_dims.get_mesh(dp_mesh_names) + if parallelism.spmd_backend == "spmd_types": + dp_mesh, dp_mesh_dims = resolve_fsdp_mesh(parallel_dims) + edp_mesh, edp_mesh_dims = resolve_sparse_fsdp_mesh(parallel_dims) + else: + dp_mesh_names = ( + ["dp_replicate", "fsdp"] if parallel_dims.dp_replicate_enabled else ["fsdp"] + ) + dp_mesh = parallel_dims.get_mesh(dp_mesh_names) + dp_mesh_dims = None + edp_mesh = None + edp_mesh_dims = None + if parallel_dims.ep_enabled: + edp_mesh_names = ( + ["dp_replicate", "efsdp"] + if parallel_dims.dp_replicate_enabled + else ["efsdp"] + ) + edp_mesh = parallel_dims.get_optional_mesh(edp_mesh_names) if model.vision_encoder is not None: apply_fsdp_to_vision_encoder( @@ -96,17 +120,9 @@ def parallelize_qwen3_5( reduce_dtype=TORCH_DTYPE_MAP[training.mixed_precision_reduce], reshard_after_forward_policy=parallelism.fsdp_reshard_after_forward, pp_enabled=parallel_dims.pp_enabled, + dp_mesh_dims=dp_mesh_dims, ) - edp_mesh = None - if parallel_dims.ep_enabled: - edp_mesh_names = ( - ["dp_replicate", "efsdp"] - if parallel_dims.dp_replicate_enabled - else ["efsdp"] - ) - edp_mesh = parallel_dims.get_optional_mesh(edp_mesh_names) - apply_fsdp_to_decoder( model, # pyrefly: ignore [bad-argument-type] dp_mesh, @@ -117,6 +133,9 @@ def parallelize_qwen3_5( reshard_after_forward_policy=parallelism.fsdp_reshard_after_forward, ep_degree=parallel_dims.ep, edp_mesh=edp_mesh, + dp_mesh_dims=dp_mesh_dims, + edp_mesh_dims=edp_mesh_dims, + enable_symm_mem=parallelism.enable_fsdp_symm_mem, ) return model diff --git a/torchtitan/models/qwen3_5/rope.py b/torchtitan/models/qwen3_5/rope.py index e2577c29b6..c3d558c6ed 100644 --- a/torchtitan/models/qwen3_5/rope.py +++ b/torchtitan/models/qwen3_5/rope.py @@ -77,7 +77,6 @@ def _compute_mrope_cache(self, position_ids: torch.Tensor) -> torch.Tensor: if isinstance(position_ids, DTensor) else position_ids ) - pos = pos.to(device=rope_cache.device) _maybe_check_max_pos(pos, max_valid_pos=rope_cache.shape[0] - 1) head_dim = rope_cache.shape[-1] // 2 diff --git a/torchtitan/models/qwen3_5/sharding.py b/torchtitan/models/qwen3_5/sharding.py index 7f3aaa5964..2c568715d3 100644 --- a/torchtitan/models/qwen3_5/sharding.py +++ b/torchtitan/models/qwen3_5/sharding.py @@ -19,7 +19,9 @@ from typing import TYPE_CHECKING import spmd_types as spmd +import torch +from torchtitan.distributed.parallel_dims import MeshAxisName from torchtitan.models.common.decoder_sharding import ( colwise_config, dense_activation_placement, @@ -34,6 +36,9 @@ from torchtitan.models.common.moe_sharding import set_moe_sharding_config from torchtitan.protocols.sharding import LocalMapConfig, ShardingConfig, SpmdLayout +DP = MeshAxisName.DP +TP = MeshAxisName.TP + if TYPE_CHECKING: from torchtitan.models.common import SigmoidGatedFeedForward from torchtitan.models.qwen3_5.model import ( @@ -45,17 +50,79 @@ from torchtitan.models.qwen3_5.vision_encoder import Qwen35VisionEncoder +def annotate_multimodal_input_spmd_types( + *, + mrope_positions: torch.Tensor | None, + pixel_values: torch.Tensor | None, + pixel_values_videos: torch.Tensor | None, + grid_thw: torch.Tensor | None, + grid_thw_videos: torch.Tensor | None, +) -> None: + """Annotate Qwen3.5 multimodal inputs with their local SPMD types.""" + token_type = { + MeshAxisName.DP: spmd.S(0), + MeshAxisName.TP: spmd.R, + } + multimodal_type = { + MeshAxisName.DP: spmd.V, + MeshAxisName.TP: spmd.I, + } + + if mrope_positions is not None: + spmd.assert_type(mrope_positions, token_type) + for tensor in ( + pixel_values, + pixel_values_videos, + grid_thw, + grid_thw_videos, + ): + if tensor is not None: + spmd.assert_type(tensor, multimodal_type) + + def _replicate_norm() -> ShardingConfig: """Replicate norm (weight/bias and activations) — used by the vision encoder, which runs without sequence parallelism.""" + activation = dense_activation_placement(tp=spmd.I) return ShardingConfig( state_shardings={ - "weight": dense_param_placement(tp=spmd.R), - "bias": dense_param_placement(tp=spmd.R), + "weight": dense_param_placement(tp=spmd.I), + "bias": dense_param_placement(tp=spmd.I), }, - in_src_shardings={"input": dense_activation_placement(tp=spmd.R)}, + in_src_shardings={"input": activation}, + in_dst_shardings={"input": activation}, + out_src_shardings=activation, + out_dst_shardings=activation, + ) + + +def _vision_colwise_config( + *, input_tp: spmd.PerMeshAxisSpmdType = spmd.I +) -> ShardingConfig: + activation = dense_activation_placement(tp=input_tp) + return ShardingConfig( + state_shardings={ + "weight": dense_param_placement(tp=spmd.S(0)), + "bias": dense_param_placement(tp=spmd.S(0)), + }, + in_src_shardings={"input": activation}, in_dst_shardings={"input": dense_activation_placement(tp=spmd.R)}, - out_dst_shardings=dense_activation_placement(tp=spmd.R), + out_src_shardings=dense_activation_placement(tp=spmd.S(-1)), + ) + + +def _vision_scaled_bias_rowwise_config() -> ShardingConfig: + input_layout = dense_activation_placement(tp=spmd.S(2)) + return ShardingConfig( + state_shardings={ + "weight": dense_param_placement(tp=spmd.S(1)), + "bias": dense_param_placement(tp=spmd.R), + }, + in_src_shardings={"input": input_layout}, + in_dst_shardings={"input": input_layout}, + out_src_shardings=dense_activation_placement(tp=spmd.P), + out_dst_shardings=dense_activation_placement(tp=spmd.I), + local_map=LocalMapConfig(in_grad_placements=(input_layout,)), ) @@ -66,6 +133,7 @@ def _qk_norm_sharding() -> ShardingConfig: state_shardings={"weight": dense_param_placement(tp=spmd.R)}, in_src_shardings={"input": head_plc}, in_dst_shardings={"input": head_plc}, + out_src_shardings=head_plc, out_dst_shardings=head_plc, ) @@ -117,15 +185,22 @@ def set_qwen35_sharding_config( ) _set_vision_encoder_sharding(config.vision_encoder) # The embedding path stays replicated through multimodal vision scatter. - # The first attention block restores SP; later decoder block inputs are SP. - first_layer_input_layout = dense_activation_placement(tp=spmd.R) + # Layer 0 restores SP at the block boundary; later decoder blocks are SP. + decoder_input_layout = dense_activation_placement(tp=spmd.R) layer_input_layout = dense_sequence_parallel_placement() for layer_idx, layer_cfg in enumerate(config.layers): + layer_cfg.sharding_config = ShardingConfig( + in_src_shardings={ + "x_BLD": ( + decoder_input_layout if layer_idx == 0 else layer_input_layout + ) + }, + in_dst_shardings={"x_BLD": layer_input_layout}, + out_src_shardings=layer_input_layout, + ) _set_qwen35_layer_sharding( layer_cfg, - attention_input_layout=( - first_layer_input_layout if layer_idx == 0 else layer_input_layout - ), + attention_input_layout=layer_input_layout, enable_ep=enable_ep, ) @@ -191,6 +266,7 @@ def _set_shared_expert_gate_sharding( "weight": dense_param_placement(tp=spmd.R), "bias": dense_param_placement(tp=spmd.R), }, + out_src_shardings=dense_activation_placement(tp=spmd.R), out_dst_shardings=dense_activation_placement(tp=spmd.R), ) @@ -203,22 +279,26 @@ def _set_vision_encoder_sharding(ve_cfg: "Qwen35VisionEncoder.Config") -> None: Norms are Replicate. pos_embed is Replicate via state_shardings. """ ve_cfg.sharding_config = ShardingConfig( - state_shardings={"pos_embed": dense_param_placement(tp=spmd.R)}, + state_shardings={"pos_embed": dense_param_placement(tp=spmd.I)}, + # I->R convert to scatter into text embeddings. + out_src_shardings=SpmdLayout({DP: spmd.V, TP: spmd.I}), + out_dst_shardings=SpmdLayout({DP: spmd.V, TP: spmd.R}), ) ve_cfg.rotary_pos_emb.sharding_config = ShardingConfig( state_shardings={"inv_freq": dense_param_placement(tp=spmd.I)}, out_src_shardings=dense_param_placement(tp=spmd.I), ) - # patch_embed receives plain pixel_values — wrap as DTensor(Replicate) + patch_activation = dense_activation_placement(tp=spmd.I) ve_cfg.patch_embed_proj.sharding_config = ShardingConfig( state_shardings={ - "weight": dense_param_placement(tp=spmd.R), - "bias": dense_param_placement(tp=spmd.R), + "weight": dense_param_placement(tp=spmd.I), + "bias": dense_param_placement(tp=spmd.I), }, - in_src_shardings={"input": dense_activation_placement(tp=spmd.R)}, - in_dst_shardings={"input": dense_activation_placement(tp=spmd.R)}, - out_dst_shardings=dense_activation_placement(tp=spmd.R), + in_src_shardings={"input": patch_activation}, + in_dst_shardings={"input": patch_activation}, + out_src_shardings=patch_activation, + out_dst_shardings=patch_activation, ) # Block sub-modules @@ -227,23 +307,29 @@ def _set_vision_encoder_sharding(ve_cfg: "Qwen35VisionEncoder.Config") -> None: block.norm2.sharding_config = _replicate_norm() block.attn.sharding_config = ShardingConfig( - in_src_shardings={"rope_cache": dense_activation_placement(tp=spmd.R)}, - in_dst_shardings={"rope_cache": dense_activation_placement(tp=spmd.R)}, + in_src_shardings={ + "x": dense_activation_placement(tp=spmd.I), + "rope_cache": SpmdLayout({DP: spmd.R, TP: spmd.I}), + }, + in_dst_shardings={ + "x": dense_activation_placement(tp=spmd.R), + "rope_cache": SpmdLayout({DP: spmd.R, TP: spmd.R}), + }, ) - block.attn.wq.sharding_config = colwise_config() - block.attn.wk.sharding_config = colwise_config() - block.attn.wv.sharding_config = colwise_config() - block.attn.proj.sharding_config = rowwise_config(output_sp=False) + block.attn.wq.sharding_config = _vision_colwise_config(input_tp=spmd.R) + block.attn.wk.sharding_config = _vision_colwise_config(input_tp=spmd.R) + block.attn.wv.sharding_config = _vision_colwise_config(input_tp=spmd.R) + block.attn.proj.sharding_config = _vision_scaled_bias_rowwise_config() set_gqa_inner_attention_local_map(block.attn.inner_attention) - block.mlp.fc1.sharding_config = colwise_config() - block.mlp.fc2.sharding_config = rowwise_config(output_sp=False) + block.mlp.fc1.sharding_config = _vision_colwise_config() + block.mlp.fc2.sharding_config = _vision_scaled_bias_rowwise_config() # Merger sub-modules merger = ve_cfg.merger merger.norm.sharding_config = _replicate_norm() - merger.fc1.sharding_config = colwise_config() - merger.fc2.sharding_config = rowwise_config(output_sp=False) + merger.fc1.sharding_config = _vision_colwise_config() + merger.fc2.sharding_config = _vision_scaled_bias_rowwise_config() def _set_full_attention_sharding( @@ -313,12 +399,20 @@ def _set_deltanet_sharding( state_shardings={"weight": dense_param_placement(tp=spmd.R)}, in_src_shardings={"x": _norm_plc, "gate": _norm_plc}, in_dst_shardings={"x": _norm_plc, "gate": _norm_plc}, + out_src_shardings=_norm_plc, out_dst_shardings=_norm_plc, ) # GatedDeltaKernel: local_map converts DTensor q/k/v/g/beta to local. _kernel_plc = dense_activation_placement(tp=spmd.S(2)) deltanet_cfg.kernel.sharding_config = ShardingConfig( + in_src_shardings={ + "xq_BLNK": _kernel_plc, + "xk_BLNK": _kernel_plc, + "xv_BLNV": _kernel_plc, + "g_BLN": _kernel_plc, + "beta_BLN": _kernel_plc, + }, in_dst_shardings={ "xq_BLNK": _kernel_plc, "xk_BLNK": _kernel_plc, @@ -339,5 +433,6 @@ def _set_deltanet_sharding( }, in_src_shardings={"x_BLD": attention_input_layout}, in_dst_shardings={"x_BLD": dense_activation_placement(tp=spmd.R)}, + out_src_shardings=dense_sequence_parallel_placement(), out_dst_shardings=dense_sequence_parallel_placement(), ) diff --git a/torchtitan/models/qwen3_5/vision_encoder.py b/torchtitan/models/qwen3_5/vision_encoder.py index 2fd11120cd..c9668c5466 100644 --- a/torchtitan/models/qwen3_5/vision_encoder.py +++ b/torchtitan/models/qwen3_5/vision_encoder.py @@ -6,12 +6,14 @@ from dataclasses import dataclass, field +import spmd_types as spmd import torch import torch.nn as nn import torch.nn.functional as F from torch.distributed.tensor import DTensor from torch.distributed.tensor.experimental import local_map +from torchtitan.distributed.utils import get_spmd_backend from torchtitan.models.common import Linear from torchtitan.models.common.nn_modules import GELU, LayerNorm from torchtitan.models.common.rope import _maybe_wrap_positions, CosSinRoPE @@ -52,6 +54,8 @@ def _compute_learned_pos_embeds( merge_size = spatial_merge_size pos_embeds = learned_pos_embed.new_zeros(len(grids), max_num_patch, dim) + if get_spmd_backend() == "spmd_types" and spmd.is_type_checking(): + pos_embeds = spmd.mutate_type(pos_embeds, "tp", src=spmd.R, dst=spmd.I) # Group images by (h, w) to batch compute position embeddings hw_to_indices: dict[tuple[int, int], list[int]] = {} @@ -141,6 +145,8 @@ def _compute_2d_rope_cache( rope_embeds = torch.zeros( len(grids), max_num_patch, head_dim // 2, device=device, dtype=torch.float32 ) + if get_spmd_backend() == "spmd_types" and spmd.is_type_checking(): + rope_embeds = spmd.mutate_type(rope_embeds, "tp", src=spmd.R, dst=spmd.I) # Group images by (h, w) to batch compute RoPE embeddings hw_to_indices: dict[tuple[int, int], list[int]] = {} @@ -180,6 +186,9 @@ def _compute_2d_rope_cache( .expand(merged_h, merged_w, merge_size, merge_size) .reshape(-1) ) + if get_spmd_backend() == "spmd_types" and spmd.is_type_checking(): + row_idx = spmd.mutate_type(row_idx, "tp", src=spmd.R, dst=spmd.I) + col_idx = spmd.mutate_type(col_idx, "tp", src=spmd.R, dst=spmd.I) # 2D RoPE: row and col each get separate frequency sets, concatenated # (not interleaved). freq_table shape: (max_hw, head_dim//4) @@ -242,6 +251,8 @@ def forward(self, seqlen: int) -> torch.Tensor: seqlen, device=self.inv_freq.device, dtype=self.inv_freq.dtype ) seq = _maybe_wrap_positions(seq, self.inv_freq) + if get_spmd_backend() == "spmd_types" and spmd.is_type_checking(): + seq = spmd.mutate_type(seq, "tp", src=spmd.R, dst=spmd.I) return torch.outer(seq, self.inv_freq) # pyrefly: ignore @@ -432,14 +443,16 @@ def forward( x = x + learned_pos mask_mod = get_vision_block_mask_mod(num_patch) - attention_mask = compiled_create_block_mask( - mask_mod, - num_vision, - None, - max_num_patch, - max_num_patch, - device=x.device, - ) + # BlockMask creation and use in FlexAttention are blackboxed by typechecking. + with spmd.no_typecheck(): + attention_mask = compiled_create_block_mask( + mask_mod, + num_vision, + None, + max_num_patch, + max_num_patch, + device=x.device, + ) for layer in self.layers.values(): x = layer( diff --git a/torchtitan/protocols/sharding.py b/torchtitan/protocols/sharding.py index 118d43d964..4f3af43cfd 100644 --- a/torchtitan/protocols/sharding.py +++ b/torchtitan/protocols/sharding.py @@ -14,12 +14,15 @@ from dataclasses import dataclass, field +import spmd_types as spmd from torch.distributed.device_mesh import DeviceMesh from torch.distributed.tensor import Partial, Placement, Replicate, Shard -from torchtitan.distributed.parallel_dims import MeshAxisName, SpmdLayout - -from torchtitan.distributed.spmd_types import spmd_layout_to_dtensor_placements +from torchtitan.distributed.parallel_dims import ( + MeshAxisName, + SpmdLayout, + unfold_dp_axis, +) __all__ = [ @@ -143,18 +146,22 @@ def resolve_placements( # TODO(fegin): remove the size-1 ``Shard(d)``/``Partial`` to ``Replicate()`` # conversion once FlexShard replaces ``fully_shard``. assert mesh.mesh_dim_names is not None, "DeviceMesh must have named axes" - placements = spmd_layout_to_dtensor_placements(layout) + axis_types = {} + for axis_name, axis_type in layout.per_axis_spmd_types().items(): + for concrete_axis_name in unfold_dp_axis(axis_name): + axis_types[concrete_axis_name] = axis_type + result = [] for i, axis_name in enumerate(mesh.mesh_dim_names): key = MeshAxisName(axis_name) - if key not in placements: + if key not in axis_types: raise ValueError( f"ShardingConfig does not declare a placement for mesh axis " f"{axis_name!r}. Declared: " f"{sorted(k.value for k in layout.axes())}; " f"required: {list(mesh.mesh_dim_names)}." ) - p = placements[key] + p = spmd.spmd_type_to_dtensor_placement(axis_types[key]) if isinstance(p, (Shard, Partial)) and mesh.size(i) == 1: p = Replicate() result.append(p) From 2072ce556c5f9065eeee0f4e5c70765e262d775d Mon Sep 17 00:00:00 2001 From: Pian Pawakapan Date: Thu, 6 Aug 2026 16:35:11 -0700 Subject: [PATCH 2/7] Update [ghstack-poisoned] --- torchtitan/models/common/decoder_sharding.py | 12 +- torchtitan/models/common/moe_sharding.py | 56 +++--- torchtitan/models/deepseek_v3/sharding.py | 10 +- torchtitan/models/kimi_k2_7/__init__.py | 18 +- torchtitan/models/kimi_k2_7/model.py | 50 +++-- torchtitan/models/kimi_k2_7/parallelize.py | 45 +++-- torchtitan/models/kimi_k2_7/sharding.py | 174 +++++++++++++----- torchtitan/models/kimi_k2_7/vision_encoder.py | 50 +++-- 8 files changed, 299 insertions(+), 116 deletions(-) diff --git a/torchtitan/models/common/decoder_sharding.py b/torchtitan/models/common/decoder_sharding.py index 0fef653d42..51fac461fc 100644 --- a/torchtitan/models/common/decoder_sharding.py +++ b/torchtitan/models/common/decoder_sharding.py @@ -62,13 +62,21 @@ def dense_sequence_parallel_placement() -> SpmdLayout: ) -def colwise_config() -> ShardingConfig: - """ColwiseParallel: weight S(0), output S(-1).""" +def colwise_config(*, input_layout: SpmdLayout | None = None) -> ShardingConfig: + """ColwiseParallel: optional input -> R, weight S(0), output S(-1).""" return ShardingConfig( state_shardings={ "weight": dense_param_placement(tp=spmd.S(0)), "bias": dense_param_placement(tp=spmd.S(0)), }, + in_src_shardings=( + {"input": input_layout} if input_layout is not None else None + ), + in_dst_shardings=( + {"input": dense_activation_placement(tp=spmd.R)} + if input_layout is not None + else None + ), out_src_shardings=dense_activation_placement(tp=spmd.S(-1)), ) diff --git a/torchtitan/models/common/moe_sharding.py b/torchtitan/models/common/moe_sharding.py index a4f1743ba4..66c908a447 100644 --- a/torchtitan/models/common/moe_sharding.py +++ b/torchtitan/models/common/moe_sharding.py @@ -135,12 +135,15 @@ def _router_gate_config(*, enable_ep: bool, enable_sp: bool) -> ShardingConfig: ) else: input_layout = dense_activation_placement(tp=spmd.R) + output_layout = dense_activation_placement(tp=spmd.I) return ShardingConfig( state_shardings=state, in_src_shardings={"input": input_layout}, in_dst_shardings={"input": input_layout}, + # Router values are identical across TP ranks. Keep that invariant + # type until a consumer explicitly enters replicated computation. out_src_shardings=input_layout, - out_dst_shardings=input_layout, + out_dst_shardings=output_layout, ) @@ -150,14 +153,14 @@ def _tokens_per_expert_placement(*, enable_ep: bool) -> SpmdLayout: Each DP/CP rank processes different data and accumulates partial token counts, so DP/CP axes are ``Partial``. TP is ``Partial`` when EP is enabled (MoE reuses the mesh axis named TP for sequence-token sharding, so - each rank sees different tokens) or ``Replicate`` when EP is disabled (all + each rank sees different tokens) or ``Invariant`` when EP is disabled (all TP ranks see the same tokens). """ return SpmdLayout( { DP: spmd.P, CP: spmd.P, - TP: spmd.P if enable_ep else spmd.R, + TP: spmd.P if enable_ep else spmd.I, } ) @@ -181,7 +184,9 @@ def _moe_sharding_config(*, enable_ep: bool, enable_sp: bool) -> ShardingConfig: return ShardingConfig( state_shardings={ - "expert_bias_E": dense_param_placement(tp=spmd.R), + # Without EP, every TP rank sees the same tokens and applies the + # same load-balancing update, so the bias remains invariant. + "expert_bias_E": dense_param_placement(tp=spmd.R if enable_ep else spmd.I), "tokens_per_expert_E": _tokens_per_expert_placement(enable_ep=enable_ep), }, in_src_shardings={"x_BLD": sp_layout}, @@ -278,8 +283,23 @@ def set_moe_sharding_config( } experts_in_layout = dense_sequence_parallel_placement() experts_in_grad_layout = dense_sequence_parallel_placement() + pre_experts_metadata_layout = experts_in_layout + pre_tokens_per_expert_layout = _tokens_per_expert_placement(enable_ep=True) + experts_tokens_per_expert_layout = pre_tokens_per_expert_layout else: pre_experts_in_layout = dense_activation_placement(tp=spmd.R) + pre_experts_metadata_layout = dense_activation_placement(tp=spmd.I) + pre_tokens_per_expert_layout = _tokens_per_expert_placement(enable_ep=False) + # Counts are invariant across TP before entering the local expert + # implementation. Replicate them explicitly for local grouped kernels, + # which consume replicated offsets alongside TP-sharded expert weights. + experts_tokens_per_expert_layout = SpmdLayout( + { + DP: spmd.P, + CP: spmd.P, + TP: spmd.R, + } + ) state_shardings = { name: expert_param_placement_dense(tp_placement=placement) for name, placement in expert_param_layout.items() @@ -291,34 +311,26 @@ def set_moe_sharding_config( moe_cfg.routed_experts.sharding_config = ShardingConfig( in_src_shardings={ "x_BLD": pre_experts_in_layout, - "topk_scores_BLK": experts_in_layout, - "topk_expert_ids_BLK": experts_in_layout, - "num_local_tokens_per_expert_E": _tokens_per_expert_placement( - enable_ep=enable_ep - ), + "topk_scores_BLK": pre_experts_metadata_layout, + "topk_expert_ids_BLK": pre_experts_metadata_layout, + "num_local_tokens_per_expert_E": pre_tokens_per_expert_layout, }, in_dst_shardings={ "x_BLD": experts_in_layout, "topk_scores_BLK": experts_in_layout, "topk_expert_ids_BLK": experts_in_layout, - "num_local_tokens_per_expert_E": _tokens_per_expert_placement( - enable_ep=enable_ep - ), + "num_local_tokens_per_expert_E": experts_tokens_per_expert_layout, }, out_src_shardings=experts_out_layout, out_dst_shardings=experts_out_layout, local_map=LocalMapConfig( in_grad_placements=( - ( - experts_in_grad_layout, - experts_in_grad_layout, - experts_in_grad_layout, - # num_local_tokens_per_expert_E is routing metadata, but it is - # still a DTensor input to local_map and must have placements. - _tokens_per_expert_placement(enable_ep=enable_ep), - ) - if enable_ep - else None + experts_in_grad_layout, + experts_in_grad_layout, + experts_in_grad_layout, + # num_local_tokens_per_expert_E is routing metadata, but it is + # still a DTensor input to local_map and must have placements. + experts_tokens_per_expert_layout, ), ), ) diff --git a/torchtitan/models/deepseek_v3/sharding.py b/torchtitan/models/deepseek_v3/sharding.py index 22ec5ab403..8ce044cf69 100644 --- a/torchtitan/models/deepseek_v3/sharding.py +++ b/torchtitan/models/deepseek_v3/sharding.py @@ -106,9 +106,15 @@ def _set_deepseek_v3_layer_sharding( state_shardings={"weight": dense_param_placement(tp=spmd.R)}, ) attention.wkv_a.sharding_config = replicate_weight - attention.kv_norm.sharding_config = replicate_weight + attention.kv_norm.sharding_config = ShardingConfig( + state_shardings={"weight": dense_param_placement(tp=spmd.R)}, + out_src_shardings=dense_activation_placement(tp=spmd.R), + out_dst_shardings=dense_activation_placement(tp=spmd.I), + ) - attention.wkv_b.sharding_config = colwise_config() + attention.wkv_b.sharding_config = colwise_config( + input_layout=dense_activation_placement(tp=spmd.I) + ) attention.wo.sharding_config = rowwise_config(output_sp=enable_sp) set_gqa_inner_attention_local_map(attention.inner_attention) diff --git a/torchtitan/models/kimi_k2_7/__init__.py b/torchtitan/models/kimi_k2_7/__init__.py index 2373a5a0ef..5696c267f5 100644 --- a/torchtitan/models/kimi_k2_7/__init__.py +++ b/torchtitan/models/kimi_k2_7/__init__.py @@ -17,6 +17,7 @@ Embedding, Linear, RMSNorm, + ScaledBiasRowwiseLinear, TransformerBlock, ) from torchtitan.models.common.nn_modules import LayerNorm @@ -104,6 +105,17 @@ def _vl_linear(in_features: int, out_features: int) -> Linear.Config: ) +def _vl_scaled_bias_rowwise_linear( + in_features: int, out_features: int +) -> ScaledBiasRowwiseLinear.Config: + return ScaledBiasRowwiseLinear.Config( + in_features=in_features, + out_features=out_features, + bias=True, + param_init=_LINEAR_INIT, + ) + + def _vl_layernorm(dim: int, eps: float = 1e-5) -> LayerNorm.Config: return LayerNorm.Config(normalized_shape=dim, eps=eps) @@ -138,11 +150,11 @@ def _vision_encoder_config( wq=_vl_linear(dim, dim), wk=_vl_linear(dim, dim), wv=_vl_linear(dim, dim), - proj=_vl_linear(dim, dim), + proj=_vl_scaled_bias_rowwise_linear(dim, dim), ), mlp=VisionMLP.Config( fc1=_vl_linear(dim, ffn_dim), - fc2=_vl_linear(ffn_dim, dim), + fc2=_vl_scaled_bias_rowwise_linear(ffn_dim, dim), ), ) @@ -168,7 +180,7 @@ def _vision_encoder_config( merged_dim=merged_dim, pre_norm=_vl_layernorm(dim), linear_1=_vl_linear(merged_dim, merged_dim), - linear_2=_vl_linear(merged_dim, text_hidden_size), + linear_2=_vl_scaled_bias_rowwise_linear(merged_dim, text_hidden_size), ), ) diff --git a/torchtitan/models/kimi_k2_7/model.py b/torchtitan/models/kimi_k2_7/model.py index bfa6ffed65..354d9b880d 100644 --- a/torchtitan/models/kimi_k2_7/model.py +++ b/torchtitan/models/kimi_k2_7/model.py @@ -8,10 +8,13 @@ https://github.com/sgl-project/sglang/blob/e0c0c0a45cb1bda90392bfa2bba4184f5b0638a0/python/sglang/srt/models/kimi_k25.py """ +import contextlib from dataclasses import dataclass +import spmd_types as spmd import torch +from torchtitan.distributed.utils import get_spmd_backend from torchtitan.models.common.attention import AttentionMasksType from torchtitan.models.common.decoder import Decoder from torchtitan.models.common.multimodal import ( @@ -20,7 +23,10 @@ ) from torchtitan.models.deepseek_v3.model import DeepSeekV3Model -from .sharding import set_kimi_k2_5_sharding_config +from .sharding import ( + annotate_multimodal_input_spmd_types, + set_kimi_k2_5_sharding_config, +) from .vision_encoder import KimiK25VisionEncoder @@ -76,6 +82,12 @@ def __init__(self, config: Config): config.vision_encoder.build() if config.vision_encoder is not None else None ) + def multimodal_context(self) -> contextlib.AbstractContextManager[None]: + """Use local DP typechecking while preparing multimodal inputs.""" + if get_spmd_backend() == "spmd_types": + return spmd.set_current_mesh(local_axes=("dp",)) + return contextlib.nullcontext() + def _prepare_multimodal_embeds( self, tokens: torch.Tensor, @@ -166,17 +178,31 @@ def forward( # pyrefly: ignore [bad-override] Returns: (batch, seq_len, vocab_size) logits. """ - if self.tok_embeddings is not None: - x = self._prepare_multimodal_embeds( - tokens, - pixel_values=pixel_values, - grid_thw=grid_thw, - pixel_values_videos=pixel_values_videos, - grid_thw_videos=grid_thw_videos, - special_tokens=special_tokens, # pyrefly: ignore [bad-argument-type] - ) - else: - x = tokens + with self.multimodal_context(): + if get_spmd_backend() == "spmd_types": + annotate_multimodal_input_spmd_types( + pixel_values=pixel_values, + grid_thw=grid_thw, + pixel_values_videos=pixel_values_videos, + grid_thw_videos=grid_thw_videos, + ) + + if self.tok_embeddings is not None: + x = self._prepare_multimodal_embeds( + tokens, + pixel_values=pixel_values, + grid_thw=grid_thw, + pixel_values_videos=pixel_values_videos, + grid_thw_videos=grid_thw_videos, + special_tokens=special_tokens, + ) + else: + x = tokens + + if get_spmd_backend() == "spmd_types": + # The scatter restores a token-aligned tensor, so text-model DP + # resumes as global batch sharding after the multimodal region. + spmd.assert_type(x, {"dp": spmd.S(0), "tp": spmd.R}) for layer in self.layers.values(): x = layer(x, attention_masks, positions) diff --git a/torchtitan/models/kimi_k2_7/parallelize.py b/torchtitan/models/kimi_k2_7/parallelize.py index d02d0b359c..47caea8887 100644 --- a/torchtitan/models/kimi_k2_7/parallelize.py +++ b/torchtitan/models/kimi_k2_7/parallelize.py @@ -29,6 +29,11 @@ apply_fsdp_to_decoder, apply_fsdp_to_vision_encoder, ) +from torchtitan.distributed.full_dtensor import ( + resolve_fsdp_mesh, + resolve_sparse_fsdp_mesh, + validate_config, +) from torchtitan.distributed.tensor_parallel import maybe_enable_async_tp @@ -67,6 +72,11 @@ def parallelize_kimi_k2_5( if parallel_dims.tp_enabled or parallel_dims.ep_enabled: if parallelism.enable_async_tensor_parallel and not model_compile_enabled: raise RuntimeError("Async TP requires torch.compile") + + if parallelism.spmd_backend == "spmd_types": + validate_config(parallel_dims, model) + model.parallelize(parallel_dims) # pyrefly: ignore [not-callable] + elif parallel_dims.tp_enabled or parallel_dims.ep_enabled: model.parallelize(parallel_dims) # pyrefly: ignore [not-callable] if parallel_dims.tp_enabled: @@ -84,10 +94,24 @@ def parallelize_kimi_k2_5( # pyrefly: ignore [bad-argument-type] apply_compile(model.vision_encoder, compile_config) - dp_mesh_names = ( - ["dp_replicate", "fsdp"] if parallel_dims.dp_replicate_enabled else ["fsdp"] - ) - dp_mesh = parallel_dims.get_mesh(dp_mesh_names) + if parallelism.spmd_backend == "spmd_types": + dp_mesh, dp_mesh_dims = resolve_fsdp_mesh(parallel_dims) + edp_mesh, edp_mesh_dims = resolve_sparse_fsdp_mesh(parallel_dims) + else: + dp_mesh_names = ( + ["dp_replicate", "fsdp"] if parallel_dims.dp_replicate_enabled else ["fsdp"] + ) + dp_mesh = parallel_dims.get_mesh(dp_mesh_names) + dp_mesh_dims = None + edp_mesh = None + edp_mesh_dims = None + if parallel_dims.ep_enabled: + edp_mesh_names = ( + ["dp_replicate", "efsdp"] + if parallel_dims.dp_replicate_enabled + else ["efsdp"] + ) + edp_mesh = parallel_dims.get_optional_mesh(edp_mesh_names) # FSDP the vision encoder as a single unit, before the decoder's FSDP. # @@ -105,17 +129,9 @@ def parallelize_kimi_k2_5( reduce_dtype=TORCH_DTYPE_MAP[training.mixed_precision_reduce], reshard_after_forward_policy=parallelism.fsdp_reshard_after_forward, pp_enabled=parallel_dims.pp_enabled, + dp_mesh_dims=dp_mesh_dims, ) - edp_mesh = None - if parallel_dims.ep_enabled: - edp_mesh_names = ( - ["dp_replicate", "efsdp"] - if parallel_dims.dp_replicate_enabled - else ["efsdp"] - ) - edp_mesh = parallel_dims.get_optional_mesh(edp_mesh_names) - apply_fsdp_to_decoder( model, # pyrefly: ignore [bad-argument-type] dp_mesh, @@ -126,6 +142,9 @@ def parallelize_kimi_k2_5( reshard_after_forward_policy=parallelism.fsdp_reshard_after_forward, ep_degree=parallel_dims.ep, edp_mesh=edp_mesh, + dp_mesh_dims=dp_mesh_dims, + edp_mesh_dims=edp_mesh_dims, + enable_symm_mem=parallelism.enable_fsdp_symm_mem, ) return model diff --git a/torchtitan/models/kimi_k2_7/sharding.py b/torchtitan/models/kimi_k2_7/sharding.py index 058e061006..6835e391a9 100644 --- a/torchtitan/models/kimi_k2_7/sharding.py +++ b/torchtitan/models/kimi_k2_7/sharding.py @@ -12,25 +12,26 @@ - Decoder (MLA + MoE): reuses ``set_deepseek_v3_sharding_config``. Multimodal configs keep the token embedding ``Replicate`` for the vision scatter and resume SP at layer 0 (see ``_shard_decoder_after_embedding_scatter``). -- Vision encoder: activations flow ``Replicate`` (no SP -- the patch sequence is +- Vision encoder: activations flow ``Invariant`` (no SP -- the patch sequence is short, so sequence-sharding would add gather/scatter around the block-diagonal attention for little memory gain). Only the linear layers are Colwise/Rowwise - sharded for memory; norms and position embeddings stay ``Replicate``. + sharded for memory; norms and position embeddings stay ``Invariant``. """ from typing import TYPE_CHECKING import spmd_types as spmd +import torch +from torchtitan.distributed.parallel_dims import MeshAxisName from torchtitan.models.common.decoder_sharding import ( - colwise_config, dense_activation_placement, dense_param_placement, - rowwise_config, + dense_sequence_parallel_placement, set_gqa_inner_attention_local_map, ) from torchtitan.models.deepseek_v3.sharding import set_deepseek_v3_sharding_config -from torchtitan.protocols.sharding import LocalMapConfig, ShardingConfig +from torchtitan.protocols.sharding import LocalMapConfig, ShardingConfig, SpmdLayout if TYPE_CHECKING: from torchtitan.models.kimi_k2_7.model import KimiK25Model @@ -38,14 +39,73 @@ _REPLICATE_PARAM = dense_param_placement(tp=spmd.R) _REPLICATE_ACT = dense_activation_placement(tp=spmd.R) -_REPLICATE_NORM = ShardingConfig( - state_shardings={"weight": _REPLICATE_PARAM, "bias": _REPLICATE_PARAM}, - in_src_shardings={"input": _REPLICATE_ACT}, - in_dst_shardings={"input": _REPLICATE_ACT}, - out_dst_shardings=_REPLICATE_ACT, +_VISION_INVARIANT_PARAM = dense_param_placement(tp=spmd.I) +_VISION_INVARIANT_ACT = dense_activation_placement(tp=spmd.I) + +_VISION_INVARIANT_NORM = ShardingConfig( + state_shardings={ + "weight": _VISION_INVARIANT_PARAM, + "bias": _VISION_INVARIANT_PARAM, + }, + in_src_shardings={"input": _VISION_INVARIANT_ACT}, + in_dst_shardings={"input": _VISION_INVARIANT_ACT}, + out_src_shardings=_VISION_INVARIANT_ACT, + out_dst_shardings=_VISION_INVARIANT_ACT, ) +def _vision_colwise_config( + *, input_tp: spmd.PerMeshAxisSpmdType = spmd.I +) -> ShardingConfig: + input_layout = dense_activation_placement(tp=input_tp) + return ShardingConfig( + state_shardings={ + "weight": dense_param_placement(tp=spmd.S(0)), + "bias": dense_param_placement(tp=spmd.S(0)), + }, + in_src_shardings={"input": input_layout}, + in_dst_shardings={"input": dense_activation_placement(tp=spmd.R)}, + out_src_shardings=dense_activation_placement(tp=spmd.S(-1)), + ) + + +def _vision_scaled_bias_rowwise_config() -> ShardingConfig: + input_layout = dense_activation_placement(tp=spmd.S(2)) + return ShardingConfig( + state_shardings={ + "weight": dense_param_placement(tp=spmd.S(1)), + "bias": dense_param_placement(tp=spmd.R), + }, + in_src_shardings={"input": input_layout}, + in_dst_shardings={"input": input_layout}, + out_src_shardings=dense_activation_placement(tp=spmd.P), + out_dst_shardings=dense_activation_placement(tp=spmd.I), + local_map=LocalMapConfig(in_grad_placements=(input_layout,)), + ) + + +def annotate_multimodal_input_spmd_types( + *, + pixel_values: torch.Tensor | None, + grid_thw: torch.Tensor | None, + pixel_values_videos: torch.Tensor | None, + grid_thw_videos: torch.Tensor | None, +) -> None: + """Annotate Kimi K2.5 multimodal inputs with their local SPMD types.""" + multimodal_type = { + MeshAxisName.DP: spmd.V, + MeshAxisName.TP: spmd.I, + } + for tensor in ( + pixel_values, + grid_thw, + pixel_values_videos, + grid_thw_videos, + ): + if tensor is not None: + spmd.assert_type(tensor, multimodal_type) + + def set_kimi_k2_5_sharding_config( config: "KimiK25Model.Config", *, @@ -64,14 +124,13 @@ def set_kimi_k2_5_sharding_config( def _shard_decoder_after_embedding_scatter(config: "KimiK25Model.Config") -> None: - """Keep ``tok_embeddings`` ``Replicate`` and resume SP at layer 0's output. + """Keep ``tok_embeddings`` ``Replicate`` and resume SP at layer 0's input. The vision scatter writes features at arbitrary sequence positions, so it needs the full (``Replicate``) embedding -- a ``Shard(1)`` one cannot be - indexed by sequence position locally. Layer 0 then takes a ``Replicate`` - input and its rowwise ``wo`` reduce-scatters back to ``Shard(1)``, so the - residual is sequence-parallel from layer 0's output and layers ``1..N-1`` - are unchanged full SP. + indexed by sequence position locally. Explicitly shard the completed + multimodal embedding before layer 0 so its residual and attention output + use the same sequence-parallel layout. """ config.tok_embeddings.sharding_config = ShardingConfig( state_shardings={"weight": dense_param_placement(tp=spmd.S(0))}, @@ -83,61 +142,78 @@ def _shard_decoder_after_embedding_scatter(config: "KimiK25Model.Config") -> Non ) layer0 = config.layers[0] - layer0.attention_norm.sharding_config = ShardingConfig( - state_shardings={"weight": _REPLICATE_PARAM}, - in_src_shardings={"input": _REPLICATE_ACT}, - out_src_shardings=_REPLICATE_ACT, - ) - layer0.attention.sharding_config = ShardingConfig( + sequence_parallel = dense_sequence_parallel_placement() + # Enter standard SP before layer 0 instead of relying on DTensor to shard + # only the residual at the first attention add. + layer0.sharding_config = ShardingConfig( in_src_shardings={"x": _REPLICATE_ACT}, - in_dst_shardings={"x": _REPLICATE_ACT}, + in_dst_shardings={"x": sequence_parallel}, + out_src_shardings=sequence_parallel, ) def _set_vision_encoder_sharding(ve_cfg) -> None: - """Replicate-activation TP plan for the MoonViT3d vision encoder. + """Invariant-activation TP plan for the MoonViT3d vision encoder. Linear layers are Colwise/Rowwise sharded for memory; norms and the - learnable position table are Replicate. ``patch_embed`` wraps the plain - ``pixel_values`` input as ``DTensor(Replicate)`` so the rest of the encoder - runs in DTensor space. + learnable position table are Invariant. Colwise regions convert their + Invariant input to Replicate, and matching Rowwise regions reduce Partial + outputs back to Invariant. """ - # The encoder's own ``pos_embed`` table is Replicate (F.interpolate runs on it). ve_cfg.sharding_config = ShardingConfig( - state_shardings={"pos_embed": _REPLICATE_PARAM}, + state_shardings={"pos_embed": _VISION_INVARIANT_PARAM}, + # The surrounding multimodal scatter operates on TP-replicated values. + out_src_shardings=SpmdLayout( + {MeshAxisName.DP: spmd.V, MeshAxisName.TP: spmd.I} + ), + out_dst_shardings=SpmdLayout( + {MeshAxisName.DP: spmd.V, MeshAxisName.TP: spmd.R} + ), + ) + ve_cfg.rotary_pos_emb.sharding_config = ShardingConfig( + state_shardings={"inv_freq": _VISION_INVARIANT_PARAM}, + out_src_shardings=_VISION_INVARIANT_PARAM, ) - # patch_embed (Linear): receives plain pixel_values -> wrap as Replicate. ve_cfg.patch_embed_proj.sharding_config = ShardingConfig( - state_shardings={"weight": _REPLICATE_PARAM, "bias": _REPLICATE_PARAM}, - in_src_shardings={"input": _REPLICATE_ACT}, - in_dst_shardings={"input": _REPLICATE_ACT}, - out_dst_shardings=_REPLICATE_ACT, + state_shardings={ + "weight": _VISION_INVARIANT_PARAM, + "bias": _VISION_INVARIANT_PARAM, + }, + in_src_shardings={"input": _VISION_INVARIANT_ACT}, + in_dst_shardings={"input": _VISION_INVARIANT_ACT}, + out_src_shardings=_VISION_INVARIANT_ACT, + out_dst_shardings=_VISION_INVARIANT_ACT, ) # Transformer block sub-modules (shared VisionTransformerBlock: norm1/norm2). block = ve_cfg.block - block.norm1.sharding_config = _REPLICATE_NORM - block.norm2.sharding_config = _REPLICATE_NORM + block.norm1.sharding_config = _VISION_INVARIANT_NORM + block.norm2.sharding_config = _VISION_INVARIANT_NORM - # The stacked 2D rope_cache enters the attention as a plain (Replicate) - # tensor input so it is DTensor-wrapped before meeting head-sharded q/k. + # Gather x and rope_cache at the attention boundary before head-sharded Q/K/V. block.attn.sharding_config = ShardingConfig( - in_src_shardings={"rope_cache": _REPLICATE_ACT}, - in_dst_shardings={"rope_cache": _REPLICATE_ACT}, + in_src_shardings={ + "x": _VISION_INVARIANT_ACT, + "rope_cache": _VISION_INVARIANT_ACT, + }, + in_dst_shardings={ + "x": _REPLICATE_ACT, + "rope_cache": _REPLICATE_ACT, + }, ) - block.attn.wq.sharding_config = colwise_config() - block.attn.wk.sharding_config = colwise_config() - block.attn.wv.sharding_config = colwise_config() - block.attn.proj.sharding_config = rowwise_config(output_sp=False) + block.attn.wq.sharding_config = _vision_colwise_config(input_tp=spmd.R) + block.attn.wk.sharding_config = _vision_colwise_config(input_tp=spmd.R) + block.attn.wv.sharding_config = _vision_colwise_config(input_tp=spmd.R) + block.attn.proj.sharding_config = _vision_scaled_bias_rowwise_config() set_gqa_inner_attention_local_map(block.attn.inner_attention) - block.mlp.fc1.sharding_config = colwise_config() - block.mlp.fc2.sharding_config = rowwise_config(output_sp=False) + block.mlp.fc1.sharding_config = _vision_colwise_config() + block.mlp.fc2.sharding_config = _vision_scaled_bias_rowwise_config() # Final norm + projector. - ve_cfg.final_norm.sharding_config = _REPLICATE_NORM + ve_cfg.final_norm.sharding_config = _VISION_INVARIANT_NORM proj = ve_cfg.projector - proj.pre_norm.sharding_config = _REPLICATE_NORM - proj.linear_1.sharding_config = colwise_config() - proj.linear_2.sharding_config = rowwise_config(output_sp=False) + proj.pre_norm.sharding_config = _VISION_INVARIANT_NORM + proj.linear_1.sharding_config = _vision_colwise_config() + proj.linear_2.sharding_config = _vision_scaled_bias_rowwise_config() diff --git a/torchtitan/models/kimi_k2_7/vision_encoder.py b/torchtitan/models/kimi_k2_7/vision_encoder.py index 5accd91f85..94f4ae0bbb 100644 --- a/torchtitan/models/kimi_k2_7/vision_encoder.py +++ b/torchtitan/models/kimi_k2_7/vision_encoder.py @@ -18,14 +18,17 @@ """ from dataclasses import dataclass, field +from typing import cast +import spmd_types as spmd import torch import torch.nn as nn import torch.nn.functional as F +from torchtitan.distributed.utils import get_spmd_backend from torchtitan.models.common import Linear from torchtitan.models.common.nn_modules import GELU, LayerNorm -from torchtitan.models.common.rope import ComplexRoPE +from torchtitan.models.common.rope import _maybe_wrap_positions, ComplexRoPE from torchtitan.models.common.vision_encoder import ( compiled_create_block_mask, get_vision_block_mask_mod, @@ -86,6 +89,10 @@ def _compute_learned_pos_embeds( """ height, width, dim = pos_embed.shape pos = pos_embed.new_zeros(len(grids), max_num_patch, dim) + if get_spmd_backend() == "spmd_types" and spmd.is_type_checking(): + # The ragged batch varies across DP but is identical across TP ranks. + pos = spmd.mutate_type(pos, "dp", src=spmd.R, dst=spmd.V) + pos = spmd.mutate_type(pos, "tp", src=spmd.R, dst=spmd.I) # (dim, height, width) for F.interpolate; .float() for bicubic. grid_table = pos_embed.permute(2, 0, 1).unsqueeze(0).float() @@ -155,9 +162,12 @@ def _compute_2d_rope_cache( """ device = freq_table.device - angles = torch.zeros( - len(grids), max_num_patch, head_dim // 2, device=device, dtype=freq_table.dtype - ) + angles = freq_table.new_zeros(len(grids), max_num_patch, head_dim // 2) + if get_spmd_backend() == "spmd_types" and spmd.is_type_checking(): + # len(grids) and max_num_patch are rank-local shapes, derived from multimodal tensors, + # so the constructed cache varies across DP and is identical across TP. + angles = spmd.mutate_type(angles, "dp", src=spmd.R, dst=spmd.V) + angles = spmd.mutate_type(angles, "tp", src=spmd.R, dst=spmd.I) # Group by (h, w) so the per-resolution angle grid is built once. hw_to_indices: dict[tuple[int, int], list[int]] = {} @@ -168,6 +178,10 @@ def _compute_2d_rope_cache( # Raster order: position p -> (row = p // w, col = p % w). Gather each # axis's angles from the precomputed table (freq_table[pos] = pos*inv_freq). flat = torch.arange(h * w, device=device) + flat = cast(torch.Tensor, _maybe_wrap_positions(flat, freq_table)) + if get_spmd_backend() == "spmd_types" and spmd.is_type_checking(): + # Every TP rank constructs the same non-gradient patch indices. + flat = spmd.mutate_type(flat, "tp", src=spmd.R, dst=spmd.I) x_ang = freq_table[flat % w] # (h*w, head_dim/4) column y_ang = freq_table[flat // w] # (h*w, head_dim/4) row # Interleave x/y so pair 2k uses x-position, pair 2k+1 uses y-position. @@ -211,6 +225,11 @@ def _tpool_patch_merger( max_merged = max((h // kh) * (w // kw) for _, h, w in grids) merged = hidden_NPD.new_zeros(num_vision, max_merged, merged_dim) + if get_spmd_backend() == "spmd_types" and spmd.is_type_checking(): + # num_vision and max_merged are rank-local shapes, derived from multimodal tensors, + # so the constructed output varies across DP and is identical across TP. + merged = spmd.mutate_type(merged, "dp", src=spmd.R, dst=spmd.V) + merged = spmd.mutate_type(merged, "tp", src=spmd.R, dst=spmd.I) for i, (t, h, w) in enumerate(grids): seq = hidden_NPD[i, : t * h * w] @@ -271,7 +290,11 @@ def forward(self, seqlen: int) -> torch.Tensor: seq = torch.arange( seqlen, device=self.inv_freq.device, dtype=self.inv_freq.dtype ) - return torch.outer(seq, self.inv_freq) + seq = cast(torch.Tensor, _maybe_wrap_positions(seq, self.inv_freq)) + if get_spmd_backend() == "spmd_types" and spmd.is_type_checking(): + # vision rope interacts with tensors unsharded on TP (I@TP) + seq = spmd.mutate_type(seq, "tp", src=spmd.R, dst=spmd.I) + return torch.outer(seq, self.inv_freq) # pyrefly: ignore class VisionProjector(Module): @@ -442,14 +465,15 @@ def forward( x = self.patch_embed(pixel_values) + learned_pos mask_mod = get_vision_block_mask_mod(num_patch) - attention_mask = compiled_create_block_mask( - mask_mod, - num_vision, - None, - max_num_patch, - max_num_patch, - device=x.device, - ) + with spmd.no_typecheck(): + attention_mask = compiled_create_block_mask( + mask_mod, + num_vision, + None, + max_num_patch, + max_num_patch, + device=x.device, + ) for block in self.layers.values(): x = block( From f483d052c26ffea8fe0205a29ed5929b9ca887f2 Mon Sep 17 00:00:00 2001 From: Pian Pawakapan Date: Fri, 7 Aug 2026 11:33:58 -0700 Subject: [PATCH 3/7] Update (base update) [ghstack-poisoned] --- torchtitan/models/common/vision_sharding.py | 127 ++++++++++++++++++++ torchtitan/models/qwen3_5/model.py | 3 +- torchtitan/models/qwen3_5/sharding.py | 102 +++------------- 3 files changed, 148 insertions(+), 84 deletions(-) create mode 100644 torchtitan/models/common/vision_sharding.py diff --git a/torchtitan/models/common/vision_sharding.py b/torchtitan/models/common/vision_sharding.py new file mode 100644 index 0000000000..6ebfeed342 --- /dev/null +++ b/torchtitan/models/common/vision_sharding.py @@ -0,0 +1,127 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +"""Sharding configs for common vision encoder components.""" + +from typing import TYPE_CHECKING + +import spmd_types as spmd + +from torchtitan.distributed.parallel_dims import MeshAxisName +from torchtitan.models.common.decoder_sharding import set_gqa_inner_attention_local_map +from torchtitan.protocols.sharding import LocalMapConfig, ShardingConfig, SpmdLayout + +if TYPE_CHECKING: + from torchtitan.models.common.vision_encoder import VisionTransformerBlock + + +DP = MeshAxisName.DP +TP = MeshAxisName.TP + + +def invariant_norm_config() -> ShardingConfig: + """Norm whose state and activations are invariant across TP ranks.""" + return ShardingConfig( + state_shardings={ + "weight": SpmdLayout({DP: spmd.R, TP: spmd.I}), + "bias": SpmdLayout({DP: spmd.R, TP: spmd.I}), + }, + in_src_shardings={ + "input": SpmdLayout({DP: spmd.V, TP: spmd.I}), + }, + in_dst_shardings={ + "input": SpmdLayout({DP: spmd.V, TP: spmd.I}), + }, + out_src_shardings=SpmdLayout({DP: spmd.V, TP: spmd.I}), + out_dst_shardings=SpmdLayout({DP: spmd.V, TP: spmd.I}), + ) + + +def vision_invariant_linear_config() -> ShardingConfig: + """Unsharded linear whose state and activations are invariant at TP.""" + return ShardingConfig( + state_shardings={ + "weight": SpmdLayout({DP: spmd.R, TP: spmd.I}), + "bias": SpmdLayout({DP: spmd.R, TP: spmd.I}), + }, + in_src_shardings={ + "input": SpmdLayout({DP: spmd.V, TP: spmd.I}), + }, + in_dst_shardings={ + "input": SpmdLayout({DP: spmd.V, TP: spmd.I}), + }, + out_src_shardings=SpmdLayout({DP: spmd.V, TP: spmd.I}), + out_dst_shardings=SpmdLayout({DP: spmd.V, TP: spmd.I}), + ) + + +def vision_colwise_config( + *, input_tp: spmd.PerMeshAxisSpmdType = spmd.I +) -> ShardingConfig: + """Colwise vision linear with a TP-replicated local matmul input.""" + return ShardingConfig( + state_shardings={ + "weight": SpmdLayout({DP: spmd.R, TP: spmd.S(0)}), + "bias": SpmdLayout({DP: spmd.R, TP: spmd.S(0)}), + }, + in_src_shardings={ + "input": SpmdLayout({DP: spmd.V, TP: input_tp}), + }, + in_dst_shardings={ + "input": SpmdLayout({DP: spmd.V, TP: spmd.R}), + }, + out_src_shardings=SpmdLayout({DP: spmd.V, TP: spmd.S(-1)}), + ) + + +def vision_scaled_bias_rowwise_config() -> ShardingConfig: + """Scaled-bias rowwise vision linear returning a TP-invariant activation.""" + return ShardingConfig( + state_shardings={ + "weight": SpmdLayout({DP: spmd.R, TP: spmd.S(1)}), + "bias": SpmdLayout({DP: spmd.R, TP: spmd.R}), + }, + in_src_shardings={ + "input": SpmdLayout({DP: spmd.V, TP: spmd.S(2)}), + }, + in_dst_shardings={ + "input": SpmdLayout({DP: spmd.V, TP: spmd.S(2)}), + }, + out_src_shardings=SpmdLayout({DP: spmd.V, TP: spmd.P}), + out_dst_shardings=SpmdLayout({DP: spmd.V, TP: spmd.I}), + local_map=LocalMapConfig( + in_grad_placements=(SpmdLayout({DP: spmd.V, TP: spmd.S(2)}),) + ), + ) + + +def set_vision_transformer_block_sharding_config( + block: "VisionTransformerBlock.Config", + *, + rope_cache_dp: spmd.PerMeshAxisSpmdType, +) -> None: + """Set TP sharding for the common vision transformer block.""" + block.norm1.sharding_config = invariant_norm_config() + block.norm2.sharding_config = invariant_norm_config() + + block.attn.sharding_config = ShardingConfig( + in_src_shardings={ + "x": SpmdLayout({DP: spmd.V, TP: spmd.I}), + "rope_cache": SpmdLayout({DP: rope_cache_dp, TP: spmd.I}), + }, + in_dst_shardings={ + "x": SpmdLayout({DP: spmd.V, TP: spmd.R}), + "rope_cache": SpmdLayout({DP: rope_cache_dp, TP: spmd.R}), + }, + ) + block.attn.wq.sharding_config = vision_colwise_config(input_tp=spmd.R) + block.attn.wk.sharding_config = vision_colwise_config(input_tp=spmd.R) + block.attn.wv.sharding_config = vision_colwise_config(input_tp=spmd.R) + block.attn.proj.sharding_config = vision_scaled_bias_rowwise_config() + set_gqa_inner_attention_local_map(block.attn.inner_attention) + + block.mlp.fc1.sharding_config = vision_colwise_config() + block.mlp.fc2.sharding_config = vision_scaled_bias_rowwise_config() diff --git a/torchtitan/models/qwen3_5/model.py b/torchtitan/models/qwen3_5/model.py index 9a381b7283..36f225a749 100644 --- a/torchtitan/models/qwen3_5/model.py +++ b/torchtitan/models/qwen3_5/model.py @@ -25,6 +25,7 @@ from torch.distributed.tensor import DTensor from torch.distributed.tensor.experimental import local_map +from torchtitan.distributed.spmd_types import spmd_mesh_size from torchtitan.distributed.utils import get_spmd_backend from torchtitan.models.common import Conv1d, Linear from torchtitan.models.common.attention import ( @@ -740,7 +741,7 @@ def __init__(self, config: Config): def multimodal_context(self) -> contextlib.AbstractContextManager[None]: """Use local DP typechecking while preparing multimodal inputs.""" - if get_spmd_backend() == "spmd_types": + if get_spmd_backend() == "spmd_types" and spmd_mesh_size("dp") > 1: return spmd.set_current_mesh(local_axes=("dp",)) return contextlib.nullcontext() diff --git a/torchtitan/models/qwen3_5/sharding.py b/torchtitan/models/qwen3_5/sharding.py index 2c568715d3..b3507ecffd 100644 --- a/torchtitan/models/qwen3_5/sharding.py +++ b/torchtitan/models/qwen3_5/sharding.py @@ -34,6 +34,13 @@ set_gqa_inner_attention_local_map, ) from torchtitan.models.common.moe_sharding import set_moe_sharding_config +from torchtitan.models.common.vision_sharding import ( + invariant_norm_config, + set_vision_transformer_block_sharding_config, + vision_colwise_config, + vision_invariant_linear_config, + vision_scaled_bias_rowwise_config, +) from torchtitan.protocols.sharding import LocalMapConfig, ShardingConfig, SpmdLayout DP = MeshAxisName.DP @@ -80,52 +87,6 @@ def annotate_multimodal_input_spmd_types( spmd.assert_type(tensor, multimodal_type) -def _replicate_norm() -> ShardingConfig: - """Replicate norm (weight/bias and activations) — used by the vision - encoder, which runs without sequence parallelism.""" - activation = dense_activation_placement(tp=spmd.I) - return ShardingConfig( - state_shardings={ - "weight": dense_param_placement(tp=spmd.I), - "bias": dense_param_placement(tp=spmd.I), - }, - in_src_shardings={"input": activation}, - in_dst_shardings={"input": activation}, - out_src_shardings=activation, - out_dst_shardings=activation, - ) - - -def _vision_colwise_config( - *, input_tp: spmd.PerMeshAxisSpmdType = spmd.I -) -> ShardingConfig: - activation = dense_activation_placement(tp=input_tp) - return ShardingConfig( - state_shardings={ - "weight": dense_param_placement(tp=spmd.S(0)), - "bias": dense_param_placement(tp=spmd.S(0)), - }, - in_src_shardings={"input": activation}, - in_dst_shardings={"input": dense_activation_placement(tp=spmd.R)}, - out_src_shardings=dense_activation_placement(tp=spmd.S(-1)), - ) - - -def _vision_scaled_bias_rowwise_config() -> ShardingConfig: - input_layout = dense_activation_placement(tp=spmd.S(2)) - return ShardingConfig( - state_shardings={ - "weight": dense_param_placement(tp=spmd.S(1)), - "bias": dense_param_placement(tp=spmd.R), - }, - in_src_shardings={"input": input_layout}, - in_dst_shardings={"input": input_layout}, - out_src_shardings=dense_activation_placement(tp=spmd.P), - out_dst_shardings=dense_activation_placement(tp=spmd.I), - local_map=LocalMapConfig(in_grad_placements=(input_layout,)), - ) - - def _qk_norm_sharding() -> ShardingConfig: """Per-head QK-norm sharding: weight Replicate, activations Shard(2).""" head_plc = dense_activation_placement(tp=spmd.S(2)) @@ -279,57 +240,32 @@ def _set_vision_encoder_sharding(ve_cfg: "Qwen35VisionEncoder.Config") -> None: Norms are Replicate. pos_embed is Replicate via state_shardings. """ ve_cfg.sharding_config = ShardingConfig( - state_shardings={"pos_embed": dense_param_placement(tp=spmd.I)}, + state_shardings={ + "pos_embed": SpmdLayout({DP: spmd.R, TP: spmd.I}), + }, # I->R convert to scatter into text embeddings. out_src_shardings=SpmdLayout({DP: spmd.V, TP: spmd.I}), out_dst_shardings=SpmdLayout({DP: spmd.V, TP: spmd.R}), ) ve_cfg.rotary_pos_emb.sharding_config = ShardingConfig( - state_shardings={"inv_freq": dense_param_placement(tp=spmd.I)}, - out_src_shardings=dense_param_placement(tp=spmd.I), - ) - - patch_activation = dense_activation_placement(tp=spmd.I) - ve_cfg.patch_embed_proj.sharding_config = ShardingConfig( state_shardings={ - "weight": dense_param_placement(tp=spmd.I), - "bias": dense_param_placement(tp=spmd.I), + "inv_freq": SpmdLayout({DP: spmd.R, TP: spmd.I}), }, - in_src_shardings={"input": patch_activation}, - in_dst_shardings={"input": patch_activation}, - out_src_shardings=patch_activation, - out_dst_shardings=patch_activation, + out_src_shardings=SpmdLayout({DP: spmd.R, TP: spmd.I}), ) - # Block sub-modules - block = ve_cfg.block - block.norm1.sharding_config = _replicate_norm() - block.norm2.sharding_config = _replicate_norm() + ve_cfg.patch_embed_proj.sharding_config = vision_invariant_linear_config() - block.attn.sharding_config = ShardingConfig( - in_src_shardings={ - "x": dense_activation_placement(tp=spmd.I), - "rope_cache": SpmdLayout({DP: spmd.R, TP: spmd.I}), - }, - in_dst_shardings={ - "x": dense_activation_placement(tp=spmd.R), - "rope_cache": SpmdLayout({DP: spmd.R, TP: spmd.R}), - }, + set_vision_transformer_block_sharding_config( + ve_cfg.block, + rope_cache_dp=spmd.R, ) - block.attn.wq.sharding_config = _vision_colwise_config(input_tp=spmd.R) - block.attn.wk.sharding_config = _vision_colwise_config(input_tp=spmd.R) - block.attn.wv.sharding_config = _vision_colwise_config(input_tp=spmd.R) - block.attn.proj.sharding_config = _vision_scaled_bias_rowwise_config() - set_gqa_inner_attention_local_map(block.attn.inner_attention) - - block.mlp.fc1.sharding_config = _vision_colwise_config() - block.mlp.fc2.sharding_config = _vision_scaled_bias_rowwise_config() # Merger sub-modules merger = ve_cfg.merger - merger.norm.sharding_config = _replicate_norm() - merger.fc1.sharding_config = _vision_colwise_config() - merger.fc2.sharding_config = _vision_scaled_bias_rowwise_config() + merger.norm.sharding_config = invariant_norm_config() + merger.fc1.sharding_config = vision_colwise_config() + merger.fc2.sharding_config = vision_scaled_bias_rowwise_config() def _set_full_attention_sharding( From a92f18f2f8246a21e1c2f5bb47206b7e19a48421 Mon Sep 17 00:00:00 2001 From: Pian Pawakapan Date: Fri, 7 Aug 2026 11:56:47 -0700 Subject: [PATCH 4/7] Update (base update) [ghstack-poisoned] --- scripts/ci/pytorch_ci_test_runner.sh | 2 +- tests/integration_tests/models.py | 2 +- torchtitan/models/qwen3_5/model.py | 2 +- torchtitan/models/qwen3_5/parallelize.py | 2 -- 4 files changed, 3 insertions(+), 5 deletions(-) diff --git a/scripts/ci/pytorch_ci_test_runner.sh b/scripts/ci/pytorch_ci_test_runner.sh index a8cce192a9..4ccb827ac1 100755 --- a/scripts/ci/pytorch_ci_test_runner.sh +++ b/scripts/ci/pytorch_ci_test_runner.sh @@ -41,7 +41,7 @@ case "$COMMAND" in model_tests) python -m tests.integration_tests.run_tests \ --test_suite models \ - --exclude "qwen3_5_moe_fsdp+tp+ep+pp" \ + --exclude "qwen3_5_moe_fsdp+tp+ep+pp_spmd_types" \ --ngpu "$NGPU" \ "$OUTPUT_DIR" ;; diff --git a/tests/integration_tests/models.py b/tests/integration_tests/models.py index 119604b97e..a92173e5fd 100755 --- a/tests/integration_tests/models.py +++ b/tests/integration_tests/models.py @@ -13,7 +13,7 @@ def _enable_spmd_backend(t: OverrideDefinitions, backend: str) -> OverrideDefinitions: """Use ``backend`` for every variant, or return an unsupported test unchanged.""" if backend == "spmd_types" and any( - "--module qwen3_5" in arg or "--module kimi_k2_7" in arg + "--module kimi_k2_7" in arg for variant in t.override_args for arg in variant ): diff --git a/torchtitan/models/qwen3_5/model.py b/torchtitan/models/qwen3_5/model.py index 36f225a749..b9f457bea7 100644 --- a/torchtitan/models/qwen3_5/model.py +++ b/torchtitan/models/qwen3_5/model.py @@ -908,7 +908,7 @@ def forward( # pyrefly: ignore [bad-override] pixel_values_videos=pixel_values_videos, grid_thw=grid_thw, grid_thw_videos=grid_thw_videos, - special_tokens=special_tokens, # pyrefly: ignore [bad-argument-type] + special_tokens=special_tokens, ) else: x = tokens diff --git a/torchtitan/models/qwen3_5/parallelize.py b/torchtitan/models/qwen3_5/parallelize.py index 872bca974e..8af3a05e67 100644 --- a/torchtitan/models/qwen3_5/parallelize.py +++ b/torchtitan/models/qwen3_5/parallelize.py @@ -72,7 +72,6 @@ def parallelize_qwen3_5( if parallelism.spmd_backend == "spmd_types": validate_config(parallel_dims, model) - # pyrefly: ignore [not-callable] model.parallelize(parallel_dims) elif parallel_dims.tp_enabled or parallel_dims.ep_enabled: # pyrefly: ignore [not-callable] @@ -135,7 +134,6 @@ def parallelize_qwen3_5( edp_mesh=edp_mesh, dp_mesh_dims=dp_mesh_dims, edp_mesh_dims=edp_mesh_dims, - enable_symm_mem=parallelism.enable_fsdp_symm_mem, ) return model From 8b55ed242676a01f2ffc86f40c0d886f73ea3299 Mon Sep 17 00:00:00 2001 From: Pian Pawakapan Date: Tue, 11 Aug 2026 12:15:06 -0700 Subject: [PATCH 5/7] Update (base update) [ghstack-poisoned] --- tests/integration_tests/models.py | 3 +-- torchtitan/models/common/rope.py | 1 + torchtitan/models/qwen3_5/model.py | 2 +- torchtitan/models/qwen3_5/parallelize.py | 2 +- 4 files changed, 4 insertions(+), 4 deletions(-) diff --git a/tests/integration_tests/models.py b/tests/integration_tests/models.py index bccc6113c0..71196a5ca4 100755 --- a/tests/integration_tests/models.py +++ b/tests/integration_tests/models.py @@ -13,8 +13,7 @@ def _enable_spmd_backend(t: OverrideDefinitions, backend: str) -> OverrideDefinitions: """Use ``backend`` for every variant, or return an unsupported test unchanged.""" if backend == "spmd_types" and any( - "--module kimi_k2_7" in arg - or "--module muse_glimmer" in arg + "--module kimi_k2_7" in arg or "--module muse_glimmer" in arg for variant in t.override_args for arg in variant ): diff --git a/torchtitan/models/common/rope.py b/torchtitan/models/common/rope.py index d75998ec7d..64556187db 100644 --- a/torchtitan/models/common/rope.py +++ b/torchtitan/models/common/rope.py @@ -22,6 +22,7 @@ ] +# pyrefly: ignore [not-callable] @spmd.no_typecheck() def _maybe_check_max_pos(positions: torch.Tensor, *, max_valid_pos: int) -> None: """Async bounds check: verify all position values <= max_valid_pos. diff --git a/torchtitan/models/qwen3_5/model.py b/torchtitan/models/qwen3_5/model.py index 2434592ca5..f1137a0f6d 100644 --- a/torchtitan/models/qwen3_5/model.py +++ b/torchtitan/models/qwen3_5/model.py @@ -895,7 +895,7 @@ def forward( # pyrefly: ignore [bad-override] pixel_values_videos=pixel_values_videos, grid_thw=grid_thw, grid_thw_videos=grid_thw_videos, - special_tokens=special_tokens, + special_tokens=special_tokens, # pyrefly: ignore [bad-argument-type] ) else: x = tokens diff --git a/torchtitan/models/qwen3_5/parallelize.py b/torchtitan/models/qwen3_5/parallelize.py index 8af3a05e67..0371c8a859 100644 --- a/torchtitan/models/qwen3_5/parallelize.py +++ b/torchtitan/models/qwen3_5/parallelize.py @@ -72,7 +72,7 @@ def parallelize_qwen3_5( if parallelism.spmd_backend == "spmd_types": validate_config(parallel_dims, model) - model.parallelize(parallel_dims) + model.parallelize(parallel_dims) # pyrefly: ignore [not-callable] elif parallel_dims.tp_enabled or parallel_dims.ep_enabled: # pyrefly: ignore [not-callable] model.parallelize(parallel_dims) From 54946dbaf3f38aa19326f3d4609404deac979f83 Mon Sep 17 00:00:00 2001 From: Pian Pawakapan Date: Wed, 12 Aug 2026 17:15:41 -0700 Subject: [PATCH 6/7] Update (base update) [ghstack-poisoned] --- torchtitan/models/qwen3_5/model.py | 2 +- torchtitan/models/qwen3_5/sharding.py | 15 ++++++++------- 2 files changed, 9 insertions(+), 8 deletions(-) diff --git a/torchtitan/models/qwen3_5/model.py b/torchtitan/models/qwen3_5/model.py index 7df3025ff3..879b7f6b23 100644 --- a/torchtitan/models/qwen3_5/model.py +++ b/torchtitan/models/qwen3_5/model.py @@ -727,7 +727,7 @@ def update_from_config( set_qwen35_sharding_config( self, enable_ep=parallelism.expert_parallel_degree > 1, - varlen=isinstance( + use_deltanet_varlen_metadata=isinstance( self.first_attention.inner_attention, VarlenAttention.Config, ), diff --git a/torchtitan/models/qwen3_5/sharding.py b/torchtitan/models/qwen3_5/sharding.py index 714dbcc368..17eccbf196 100644 --- a/torchtitan/models/qwen3_5/sharding.py +++ b/torchtitan/models/qwen3_5/sharding.py @@ -125,7 +125,7 @@ def set_qwen35_sharding_config( config: "Qwen35Model.Config", *, enable_ep: bool, - varlen: bool, + use_deltanet_varlen_metadata: bool, ) -> None: """Fill ``sharding_config`` on all Qwen3.5 sub-configs. @@ -164,7 +164,7 @@ def set_qwen35_sharding_config( layer_cfg, attention_input_layout=layer_input_layout, enable_ep=enable_ep, - varlen=varlen, + use_deltanet_varlen_metadata=use_deltanet_varlen_metadata, ) @@ -173,7 +173,7 @@ def _set_qwen35_layer_sharding( *, attention_input_layout: SpmdLayout, enable_ep: bool, - varlen: bool, + use_deltanet_varlen_metadata: bool, ) -> None: layer_cfg.attention_norm.sharding_config = _decoder_norm_sharding( attention_input_layout @@ -190,7 +190,7 @@ def _set_qwen35_layer_sharding( _set_deltanet_sharding( layer_cfg.delta_net, attention_input_layout=attention_input_layout, - varlen=varlen, + use_varlen_metadata=use_deltanet_varlen_metadata, ) if layer_cfg.feed_forward is not None: @@ -302,7 +302,7 @@ def _set_deltanet_sharding( deltanet_cfg: "GatedDeltaNet.Config", *, attention_input_layout: SpmdLayout, - varlen: bool, + use_varlen_metadata: bool, ) -> None: """Sharding for GatedDeltaNet: head-sharded TP on projections. @@ -334,10 +334,11 @@ def _set_deltanet_sharding( # RowwiseParallel on output projection (reduce-scatter to SP) deltanet_cfg.out_proj.sharding_config = rowwise_config(output_sp=True) - # Varlen flattens (B, L) to (1, B*L), moving DP sharding to dim 1. + # DeltaNet flattens (B, L) to (1, B * L) when it receives varlen + # metadata, so the DP-sharded batch dimension moves from dim 0 to dim 1. deltanet_activation_layout = ( SpmdLayout({DP: spmd.S(1), TP: spmd.S(2)}) - if varlen + if use_varlen_metadata else dense_activation_placement(tp=spmd.S(2)) ) From d3a0699ef462978fc1cc642ba4f54999efad924b Mon Sep 17 00:00:00 2001 From: Pian Pawakapan Date: Thu, 13 Aug 2026 16:46:09 -0700 Subject: [PATCH 7/7] Update (base update) [ghstack-poisoned] --- torchtitan/models/qwen3_5/model.py | 65 +++++++++------------------ torchtitan/models/qwen3_5/sharding.py | 19 +++----- 2 files changed, 25 insertions(+), 59 deletions(-) diff --git a/torchtitan/models/qwen3_5/model.py b/torchtitan/models/qwen3_5/model.py index ab2fb1e6e7..ac233fb05e 100644 --- a/torchtitan/models/qwen3_5/model.py +++ b/torchtitan/models/qwen3_5/model.py @@ -33,7 +33,6 @@ AttentionMasksType, BaseAttention, create_varlen_metadata_for_document, - FlexAttention, local_head_split, VarlenAttention, VarlenMetadata, @@ -334,7 +333,6 @@ def _causal_conv( cu_seqlens: torch.Tensor | None = None, cu_seqlens_cpu: torch.Tensor | None = None, ) -> torch.Tensor: - # varlen attention path if cu_seqlens is not None: if isinstance(x_BLD, DTensor): @@ -364,10 +362,10 @@ def _conv_varlen( @spmd.local_map( in_types=( - {"dp": spmd.S(0), "tp": spmd.S(1)}, + {"dp": spmd.S(2), "tp": spmd.S(1)}, {"dp": spmd.R, "tp": spmd.S(0)}, ), - out_types={"dp": spmd.S(0), "tp": spmd.S(1)}, + out_types={"dp": spmd.S(2), "tp": spmd.S(1)}, ) def _local_depthwise_conv1d( x_local_BDL: torch.Tensor, w_local: torch.Tensor @@ -423,53 +421,43 @@ def forward( # device offsets when it is materialized as a new tensor. spmd.mutate_type(cu_seqlens_cpu, "dp", src=spmd.R, dst=spmd.V) - def _maybe_flatten(tensor: torch.Tensor) -> torch.Tensor: - if cu_seqlens is None: - return tensor + def fold_bl_dim(tensor: torch.Tensor) -> torch.Tensor: return tensor.reshape(1, B * L, *tensor.shape[2:]) - dp_shard_dim = 1 if cu_seqlens is not None else 0 - - # Shapes: - # xq_BLNK, xk_BLNK: (B, L, n_key_heads, key_head_dim) - # xv_BLNV, xz_BLNV: (B, L, n_value_heads, value_head_dim) - # xa_BLN, xb_BLN: (B, L, n_value_heads) + # Folded recurrence shapes: + # xq_BLNK, xk_BLNK: (1, B * L, n_key_heads, key_head_dim) + # xv_BLNV, xz_BLNV: (1, B * L, n_value_heads, value_head_dim) + # xa_BLN, xb_BLN: (1, B * L, n_value_heads) xq_BLNK = self._causal_conv( - _maybe_flatten(self.in_proj_q(x_BLD)), + fold_bl_dim(self.in_proj_q(x_BLD)), self.conv_q, cu_seqlens, cu_seqlens_cpu, ) - xq_BLNK = local_head_split( - xq_BLNK, self.key_head_dim, dp_shard_dim=dp_shard_dim - ) + xq_BLNK = local_head_split(xq_BLNK, self.key_head_dim, dp_shard_dim=1) xk_BLNK = self._causal_conv( - _maybe_flatten(self.in_proj_k(x_BLD)), + fold_bl_dim(self.in_proj_k(x_BLD)), self.conv_k, cu_seqlens, cu_seqlens_cpu, ) - xk_BLNK = local_head_split( - xk_BLNK, self.key_head_dim, dp_shard_dim=dp_shard_dim - ) + xk_BLNK = local_head_split(xk_BLNK, self.key_head_dim, dp_shard_dim=1) xv_BLNV = self._causal_conv( - _maybe_flatten(self.in_proj_v(x_BLD)), + fold_bl_dim(self.in_proj_v(x_BLD)), self.conv_v, cu_seqlens, cu_seqlens_cpu, ) - xv_BLNV = local_head_split( - xv_BLNV, self.value_head_dim, dp_shard_dim=dp_shard_dim - ) + xv_BLNV = local_head_split(xv_BLNV, self.value_head_dim, dp_shard_dim=1) xz_BLNV = local_head_split( - _maybe_flatten(self.in_proj_z(x_BLD)), + fold_bl_dim(self.in_proj_z(x_BLD)), self.value_head_dim, - dp_shard_dim=dp_shard_dim, + dp_shard_dim=1, ) - xa_BLN = _maybe_flatten(self.in_proj_a(x_BLD)) - xb_BLN = _maybe_flatten(self.in_proj_b(x_BLD)) + xa_BLN = fold_bl_dim(self.in_proj_a(x_BLD)) + xb_BLN = fold_bl_dim(self.in_proj_b(x_BLD)) - # Gating signals have shape (B, L, n_value_heads): + # Gating signals have shape (1, B * L, n_value_heads): # g_BLN: decay rate per head, always negative # beta_BLN: update gate in (0, 1) g_BLN = -torch.exp(self.A_log.float()) * F.softplus( @@ -489,13 +477,8 @@ def _maybe_flatten(tensor: torch.Tensor) -> torch.Tensor: out_BLNV = self.norm(out_BLNV, xz_BLNV) - # Merge value heads and restore (B, L); under varlen the kernel ran on a - # flattened (1, B*L) layout, so this also unpacks the batch. - out_BLD = ( - unflatten_to_bld(out_BLNV, B, L) - if cu_seqlens is not None - else out_BLNV.flatten(2) - ) + # Merge value heads and restore (B, L) from the folded (1, B * L) layout. + out_BLD = unflatten_to_bld(out_BLNV, B, L) return self.out_proj(out_BLD) @@ -736,17 +719,9 @@ def update_from_config( f"n_value_heads ({n_value_heads})." ) - first_attention = self.first_attention set_qwen35_sharding_config( self, enable_ep=parallelism.expert_parallel_degree > 1, - deltanet_inputs_flattened=( - first_attention is not None - and isinstance( - first_attention.inner_attention, - (FlexAttention.Config, VarlenAttention.Config), - ) - ), ) def get_nparams_and_flops( diff --git a/torchtitan/models/qwen3_5/sharding.py b/torchtitan/models/qwen3_5/sharding.py index da0d32416f..9df9fc8480 100644 --- a/torchtitan/models/qwen3_5/sharding.py +++ b/torchtitan/models/qwen3_5/sharding.py @@ -135,7 +135,6 @@ def set_qwen35_sharding_config( config: "Qwen35Model.Config", *, enable_ep: bool, - deltanet_inputs_flattened: bool, ) -> None: """Fill ``sharding_config`` on all Qwen3.5 sub-configs. @@ -174,7 +173,6 @@ def set_qwen35_sharding_config( layer_cfg, attention_input_layout=layer_input_layout, enable_ep=enable_ep, - deltanet_inputs_flattened=deltanet_inputs_flattened, ) @@ -183,7 +181,6 @@ def _set_qwen35_layer_sharding( *, attention_input_layout: SpmdLayout, enable_ep: bool, - deltanet_inputs_flattened: bool, ) -> None: layer_cfg.attention_norm.sharding_config = _decoder_norm_sharding( attention_input_layout @@ -200,7 +197,6 @@ def _set_qwen35_layer_sharding( _set_deltanet_sharding( layer_cfg.delta_net, attention_input_layout=attention_input_layout, - inputs_flattened=deltanet_inputs_flattened, ) if layer_cfg.feed_forward is not None: @@ -312,7 +308,6 @@ def _set_deltanet_sharding( deltanet_cfg: "GatedDeltaNet.Config", *, attention_input_layout: SpmdLayout, - inputs_flattened: bool, ) -> None: """Sharding for GatedDeltaNet: head-sharded TP on projections. @@ -344,15 +339,11 @@ def _set_deltanet_sharding( # RowwiseParallel on output projection (reduce-scatter to SP) deltanet_cfg.out_proj.sharding_config = rowwise_config(output_sp=True) - # Pretraining supplies document offsets for both Flex and Varlen attention, - # so DeltaNet flattens (B, L) to (1, B * L) and moves DP sharding from dim 0 - # to dim 1. CP is intentionally absent from this layout: Qwen3.5 currently - # rejects CP, and its sequence sharding would collide with DP on dim 1. - deltanet_activation_layout = ( - SpmdLayout({DP: spmd.S(1), TP: spmd.S(2)}) - if inputs_flattened - else dense_activation_placement(tp=spmd.S(2)) - ) + # Training folds (B, L) to (1, B * L), so tensor DP shards dim 1. Inference + # already supplies folded tokens and uses separate vLLM DP workers. CP is + # omitted because Qwen3.5 rejects CP, whose sequence sharding would collide + # with tensor DP on dim 1. + deltanet_activation_layout = SpmdLayout({DP: spmd.S(1), TP: spmd.S(2)}) # RMSNormGated: per-head norm, weight Replicate, activations Shard(2) deltanet_cfg.norm.sharding_config = ShardingConfig(