Skip to content
Draft
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
136 changes: 135 additions & 1 deletion tests/unit_tests/cpu/test_context_parallel_validation.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,11 +5,15 @@
# LICENSE file in the root directory of this source tree.

import unittest
from unittest import mock

import pytest
import torch
from torch.nn.attention.flex_attention import BlockMask

from torchtitan.config import ParallelismConfig
from torchtitan.distributed.context_parallel import validate_cp_backend
from torchtitan.distributed.context_parallel import cp_shard, validate_cp_backend
from torchtitan.distributed.pipeline_parallel import pipeline_vlm


class TestValidateCpBackend(unittest.TestCase):
Expand All @@ -31,6 +35,78 @@ def test_allows_partial_dtensor_without_cp(self):
validate_cp_backend(self._parallelism(spmd_backend="partial_dtensor", cp=1))


class TestContextParallelMaskSharding(unittest.TestCase):
def test_mixed_mask_mapping_preserves_non_block_mask_metadata(self):
input_T = torch.arange(8)
input_shard_T = input_T[:4]
block_mask = mock.Mock(spec=BlockMask)
sharded_block_mask = mock.Mock(spec=BlockMask)
varlen_metadata = mock.sentinel.varlen_metadata
cp_context = mock.sentinel.cp_context
attention_masks = {
"quadratic_attention": block_mask,
"deltanet": varlen_metadata,
"deltanet_cp_context": cp_context,
}
cp_mesh = mock.Mock()
cp_mesh.size.return_value = 2

with mock.patch(
"torchtitan.distributed.context_parallel.api._context_parallel_shard",
side_effect=[[input_shard_T], [sharded_block_mask]],
):
sharded_inputs, sharded_masks = cp_shard(
cp_mesh,
(input_T,),
attention_masks,
load_balancer_type=None,
)

self.assertIs(sharded_inputs[0], input_shard_T)
assert isinstance(sharded_masks, dict)
self.assertIs(sharded_masks["quadratic_attention"], sharded_block_mask)
self.assertIs(sharded_masks["deltanet"], varlen_metadata)
self.assertIs(sharded_masks["deltanet_cp_context"], cp_context)


class TestVlmPipelineInputModules(unittest.TestCase):
def test_post_scatter_reshard_stays_with_token_embeddings(self):
model = mock.Mock()
model.decoder_input_reshard = mock.Mock()
parallelism = ParallelismConfig(
module_fqns_per_model_part=[
["vision_encoder", "tok_embeddings", "layers.0"],
["layers.1", "norm", "lm_head"],
]
)
expected = mock.sentinel.pipeline_result

with mock.patch(
"torchtitan.distributed.pipeline_parallel.pipeline_llm",
return_value=expected,
) as pipeline_llm:
result = pipeline_vlm(
model,
parallel_dims=mock.sentinel.parallel_dims,
parallelism=parallelism,
model_config=mock.sentinel.model_config,
)

self.assertIs(result, expected)
stage_fqns = pipeline_llm.call_args.kwargs[
"parallelism"
].module_fqns_per_model_part
self.assertEqual(
stage_fqns[0],
[
"vision_encoder",
"tok_embeddings",
"decoder_input_reshard",
"layers.0",
],
)


class TestDecoderConfigCpValidation(unittest.TestCase):
"""``Decoder.Config.update_from_config`` applies the CP gates at config time."""

Expand Down Expand Up @@ -94,5 +170,63 @@ def test_allows_partial_dtensor_without_cp(self):
config.model_spec.model.update_from_config(config=config)


class TestQwen35ConfigCpValidation(unittest.TestCase):
@staticmethod
def _config():
try:
from torchtitan.models.qwen3_5.config_registry import qwen35_debugmodel
except ModuleNotFoundError as exc:
raise unittest.SkipTest(
f"Qwen3.5 optional dependency unavailable: {exc.name}"
) from exc

config = qwen35_debugmodel()
config.parallelism.spmd_backend = "spmd_types"
config.parallelism.context_parallel_degree = 2
config.training.max_context_length = 512
return config

def test_rejects_context_parallel_load_balancing(self):
config = self._config()
config.parallelism.context_parallel_load_balancer = "headtail"
with self.assertRaisesRegex(ValueError, "contiguous sequence shards"):
config.model_spec.model.update_from_config( # pyrefly: ignore[missing-attribute]
config=config
)

def test_allows_contiguous_context_parallel_sharding(self):
import spmd_types as spmd

from torchtitan.distributed.parallel_dims import MeshAxisName
from torchtitan.distributed.spmd_types import (
_per_axis_types,
spmd_validate_redistributions,
)

config = self._config()
config.parallelism.context_parallel_load_balancer = None
config.model_spec.model.update_from_config( # pyrefly: ignore[missing-attribute]
config=config
)

model_config = config.model_spec.model
reshard_config = model_config.decoder_input_reshard.sharding_config
assert reshard_config is not None
assert reshard_config.in_src_shardings is not None
assert reshard_config.in_dst_shardings is not None
self.assertEqual(
_per_axis_types(reshard_config.in_src_shardings["input"])[MeshAxisName.CP],
spmd.R,
)
self.assertEqual(
_per_axis_types(reshard_config.in_dst_shardings["input"])[MeshAxisName.CP],
spmd.S(0),
)
spmd_validate_redistributions(reshard_config)
first_layer_config = model_config.layers[0].sharding_config
assert first_layer_config is not None
spmd_validate_redistributions(first_layer_config)


if __name__ == "__main__":
unittest.main()
8 changes: 8 additions & 0 deletions tests/unit_tests/cpu/test_parallel_dims.py
Original file line number Diff line number Diff line change
Expand Up @@ -275,6 +275,14 @@ def test_decoder_layout_partition_spec_ranks(self):
attention_activation_placement().partition_spec,
((MeshAxisName.DP, MeshAxisName.CP), MeshAxisName.TP, None),
)
self.assertEqual(
token_id_placement(cp=spmd.R).partition_spec,
(MeshAxisName.DP,),
)
self.assertEqual(
_per_axis_types(token_id_placement(cp=spmd.R))[MeshAxisName.CP],
spmd.R,
)

def test_unfold_dp_axes(self):
"""Logical DP expands only when resolving concrete mesh axes."""
Expand Down
37 changes: 37 additions & 0 deletions tests/unit_tests/gpu/test_qwen3_5_deltanet.py
Original file line number Diff line number Diff line change
Expand Up @@ -144,7 +144,9 @@ def forward(
*,
cu_seqlens: torch.Tensor | None = None,
cu_seqlens_cpu: torch.Tensor | None = None,
cp_context: object | None = None,
) -> torch.Tensor:
assert cp_context is None
if xq_THK.shape[1] != xv_THV.shape[1]:
assert xv_THV.shape[1] % xq_THK.shape[1] == 0
repeat = xv_THV.shape[1] // xq_THK.shape[1]
Expand All @@ -168,6 +170,41 @@ def forward(


class TestQwen35DeltaNetVarlen(unittest.TestCase):
def test_chunk_kernel_forwards_cp_context(self):
try:
from torchtitan.models.qwen3_5 import GatedDeltaKernel
except ModuleNotFoundError as exc:
raise unittest.SkipTest(
f"Qwen3.5 optional dependency unavailable: {exc.name}"
) from exc

kernel = GatedDeltaKernel(GatedDeltaKernel.Config(backend="fla_chunked"))
xq_THK = torch.randn(8, 2, 4)
xk_THK = torch.randn(8, 2, 4)
xv_THV = torch.randn(8, 2, 4)
g_TH = torch.randn(8, 2)
beta_TH = torch.randn(8, 2)
cp_context = mock.sentinel.cp_context

with mock.patch(
"torchtitan.models.qwen3_5.gdn._fla_chunk_gated_delta_rule",
return_value=(xv_THV.unsqueeze(0), None),
) as chunk_gated_delta_rule:
output_THV = kernel(
xq_THK,
xk_THK,
xv_THV,
g_TH,
beta_TH,
cp_context=cp_context,
)

torch.testing.assert_close(output_THV, xv_THV)
kwargs = chunk_gated_delta_rule.call_args.kwargs
self.assertIs(kwargs["cp_context"], cp_context)
self.assertIsNone(kwargs["cu_seqlens"])
self.assertIsNone(kwargs["cu_seqlens_cpu"])

def test_flex_masks_ignore_padding_position_resets(self):
try:
from torchtitan.models.common.decoder import Decoder
Expand Down
80 changes: 80 additions & 0 deletions tests/unit_tests/test_qwen3_5_mrope_positions.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@
"""

import unittest
from unittest import mock

import torch
from torch import nn
Expand Down Expand Up @@ -129,6 +130,85 @@ def test_multimodal_batch_routes_mrope_to_layers(self):
batch["attention_masks"]["deltanet"].cu_seq_q_host, (0, 3, 5, 10)
)

def test_context_parallel_context_is_built_before_input_sharding(self):
import spmd_types as spmd

from torchtitan.distributed.parallel_dims import MeshAxisName
from torchtitan.distributed.spmd_types import _per_axis_types
from torchtitan.models.common.attention import VarlenMetadata

model, _sink, _parallel_dims, parallelism = self._build_stub_model()
positions = torch.arange(8, dtype=torch.int32)
cu_seqlens = torch.tensor([0, 8], dtype=torch.int32)
deltanet_metadata = VarlenMetadata(
cu_seq_q=cu_seqlens,
cu_seq_k=cu_seqlens,
max_q=8,
max_k=8,
cu_seq_q_host=(0, 8),
)
attention_masks = {
"quadratic_attention": None,
"deltanet": deltanet_metadata,
}
cp_mesh = mock.Mock()
cp_mesh.size.return_value = 2
cp_group = mock.sentinel.cp_group
cp_mesh.get_group.return_value = cp_group
parallel_dims = mock.Mock()
parallel_dims.cp_enabled = True
parallel_dims.get_mesh.return_value = cp_mesh
parallelism.context_parallel_load_balancer = None
cp_context = mock.sentinel.cp_context
input_dict = {
"input": torch.randint(0, 100, (8,)),
"positions": positions,
"labels": torch.zeros(8),
"pixel_values": torch.randn(4, 8),
"grid_thw": torch.tensor([[1, 2, 2]]),
}

with (
mock.patch.object(
model,
"get_attention_masks",
return_value=attention_masks,
),
mock.patch(
"torchtitan.models.qwen3_5.model.build_cp_context",
return_value=cp_context,
) as build_context,
mock.patch(
"torchtitan.distributed.context_parallel.api.prepare_context_parallel_input",
side_effect=lambda batch, *_args: batch,
) as prepare_cp_input,
):
_inputs, _labels, batch = model.preprocess_inputs(
input_dict,
parallel_dims=parallel_dims,
parallelism=parallelism,
)

self.assertIs(batch["attention_masks"]["deltanet_cp_context"], cp_context)
build_context.assert_called_once()
args, kwargs = build_context.call_args
self.assertIs(args[0], cu_seqlens)
self.assertIs(kwargs["group"], cp_group)
self.assertEqual(kwargs["conv1d_kernel_size"], model.gdn_conv_kernel_size)
torch.testing.assert_close(
kwargs["cu_seqlens_cpu"], torch.tensor([0, 8], dtype=torch.long)
)

input_sharding = prepare_cp_input.call_args.args[1]
self.assertEqual(
_per_axis_types(input_sharding["input"])[MeshAxisName.CP], spmd.R
)
self.assertEqual(
_per_axis_types(input_sharding["pixel_values"])[MeshAxisName.CP],
spmd.R,
)
self.assertIs(batch["pixel_values"], input_dict["pixel_values"])


if __name__ == "__main__":
unittest.main()
Loading
Loading