diff --git a/docs/design/feature/realtime_ar_diffusion.md b/docs/design/feature/realtime_ar_diffusion.md index 5963befe7f6..426ee5afd73 100644 --- a/docs/design/feature/realtime_ar_diffusion.md +++ b/docs/design/feature/realtime_ar_diffusion.md @@ -178,6 +178,24 @@ For LingBot World v2: - `lingbot_world/pipeline.py` constructs conditioning, owns small non-KV session state, and produces one block plus the standard metadata envelope. +## Ulysses sequence parallelism + +LingBot supports pure Ulysses sequence parallelism for both direct and +AR-Diffusion execution. With Ulysses degree greater than one, hidden tokens, +camera features, token-expanded timestep modulation, and RoPE tables are sharded +together. Without SP, timestep modulation retains the frame-broadcast path. +Self-attention performs the sequence-to-head all-to-all before reading or writing +paged KV, while static text K/V uses the same local head shard. Text K/V shards +own compact storage; cross-attention exchanges query/output layouts to use these +shards. This retains the shared cache geometry and reduces text K/V storage at +the cost of two all-to-all calls per layer. + +The output head projects local tokens before gathering the flow values. For +the 14B model this reduces the gathered width from 5120 to 64; frame modulation +uses each shard's global token offset, including shards that split a frame. Only +`ulysses_mode="strict"` is supported; `advanced_uaa`, Ring, and AllGather-KV modes +remain unsupported for this model. + ## Non-goals This contract does not currently provide: diff --git a/examples/offline_inference/diffusion/lingbot_world_v2.py b/examples/offline_inference/diffusion/lingbot_world_v2.py index 590442c4d84..4bbea2f4f69 100644 --- a/examples/offline_inference/diffusion/lingbot_world_v2.py +++ b/examples/offline_inference/diffusion/lingbot_world_v2.py @@ -1,5 +1,5 @@ # SPDX-License-Identifier: Apache-2.0 -# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project """Generate a LingBot-World v2 video from an image and camera trajectory. The official checkpoint is licensed separately under CC BY-NC-SA and is @@ -99,6 +99,7 @@ def parse_args(argv: Sequence[str] | None = None) -> argparse.Namespace: default=1, help="Number of GPUs used for tensor parallelism inside the DiT.", ) + parser.add_argument("--ulysses-degree", type=int, default=1, help="Pure Ulysses sequence parallel degree.") parser.add_argument("--flow-shift", type=float, default=5.0, help="Positive FlowUniPC scheduler shift.") parser.add_argument("--fps", type=int, default=16, help="Frames per second in the exported MP4.") parser.add_argument("--output", default="lingbot_world_v2.mp4", help="Output MP4 path.") @@ -158,6 +159,8 @@ def build_omni_kwargs( if args.tensor_parallel_size <= 0: raise ValueError("--tensor-parallel-size must be a positive integer.") + if args.ulysses_degree <= 0: + raise ValueError("--ulysses-degree must be a positive integer.") flow_shift = _positive_finite(args.flow_shift, "--flow-shift") model_path = Path(args.model).expanduser() model = str(model_path.resolve()) if model_path.exists() else args.model @@ -165,6 +168,7 @@ def build_omni_kwargs( "model": model, "flow_shift": flow_shift, "tensor_parallel_size": args.tensor_parallel_size, + "ulysses_degree": args.ulysses_degree, "enforce_eager": args.enforce_eager, "model_config": {"lingbot_action_root": str(paths.action_root)}, } diff --git a/tests/diffusion/models/lingbot_world/test_lingbot_world_attention.py b/tests/diffusion/models/lingbot_world/test_lingbot_world_attention.py index f5567a706d1..3d883414e1e 100644 --- a/tests/diffusion/models/lingbot_world/test_lingbot_world_attention.py +++ b/tests/diffusion/models/lingbot_world/test_lingbot_world_attention.py @@ -1,5 +1,5 @@ # SPDX-License-Identifier: Apache-2.0 -# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project from __future__ import annotations @@ -7,7 +7,7 @@ import math import sys from pathlib import Path -from types import ModuleType +from types import ModuleType, SimpleNamespace import pytest import torch @@ -33,6 +33,10 @@ "vllm_omni.diffusion", "vllm_omni.diffusion.attention", "vllm_omni.diffusion.attention.layer", + "vllm_omni.diffusion.distributed", + "vllm_omni.diffusion.distributed.comm", + "vllm_omni.diffusion.distributed.parallel_state", + "vllm_omni.diffusion.distributed.sp_plan", "vllm_omni.diffusion.layers", "vllm_omni.diffusion.layers.norm", "vllm_omni.diffusion.layers.rope", @@ -58,6 +62,47 @@ def _install_vllm_stubs() -> None: distributed.get_tensor_model_parallel_world_size = lambda: 1 distributed.tensor_model_parallel_all_reduce = lambda value: value + class _SeqAllToAll4D: + @staticmethod + def apply(group, value, scatter_idx, gather_idx, use_sync=False): + del group, scatter_idx, gather_idx, use_sync + return value + + setattr( + sys.modules["vllm_omni.diffusion.distributed.comm"], + "SeqAllToAll4D", + _SeqAllToAll4D, + ) + + def get_sp_group(): + return SimpleNamespace( + ulysses_world_size=1, + ulysses_rank=0, + ulysses_group=None, + ) + + setattr( + sys.modules["vllm_omni.diffusion.distributed.parallel_state"], + "get_sp_group", + get_sp_group, + ) + + class _SequenceParallelInput: + def __init__(self, split_dim, expected_dims=None, split_output=False, auto_pad=False): + self.split_dim = split_dim + self.expected_dims = expected_dims + self.split_output = split_output + self.auto_pad = auto_pad + + class _SequenceParallelOutput: + def __init__(self, gather_dim, expected_dims=None): + self.gather_dim = gather_dim + self.expected_dims = expected_dims + + sp_plan = sys.modules["vllm_omni.diffusion.distributed.sp_plan"] + setattr(sp_plan, "SequenceParallelInput", _SequenceParallelInput) + setattr(sp_plan, "SequenceParallelOutput", _SequenceParallelOutput) + def set_weight_attrs(weight: torch.Tensor, attrs: dict) -> None: for name, value in attrs.items(): setattr(weight, name, value) @@ -541,3 +586,38 @@ def test_tp_rmsnorm_weight_loader_selects_rank_shard(monkeypatch: pytest.MonkeyP norm.weight.weight_loader(norm.weight, torch.tensor([10.0, 20.0, 30.0, 40.0])) torch.testing.assert_close(norm.weight, torch.tensor([30.0, 40.0])) + + +def test_single_token_text_kv_shard_owns_compact_storage(monkeypatch): + module = _load_module() + monkeypatch.setattr( + module, + "get_sp_group", + lambda: SimpleNamespace( + ulysses_world_size=2, + ulysses_rank=1, + ulysses_group=None, + ), + ) + attention = module.LingBotCrossAttention(dim=8, num_heads=4) + full = torch.arange(8, dtype=torch.float32).reshape(1, 1, 4, 2) + shard = attention.shard_kv_heads(full) + torch.testing.assert_close(shard, full[:, :, 2:], rtol=0, atol=0) + assert shard.is_contiguous() + assert shard.untyped_storage().nbytes() == shard.numel() * shard.element_size() + + +@pytest.mark.parametrize("attention_class", ["LingBotSelfAttention", "LingBotCrossAttention"]) +def test_ulysses_rejects_non_divisible_head_count(monkeypatch, attention_class): + module = _load_module() + monkeypatch.setattr( + module, + "get_sp_group", + lambda: SimpleNamespace( + ulysses_world_size=3, + ulysses_rank=0, + ulysses_group=None, + ), + ) + with pytest.raises(ValueError, match="heads must be divisible"): + getattr(module, attention_class)(dim=8, num_heads=4) diff --git a/tests/diffusion/models/lingbot_world/test_lingbot_world_transformer.py b/tests/diffusion/models/lingbot_world/test_lingbot_world_transformer.py index e2e2a08b6a6..b4bf7805fe3 100644 --- a/tests/diffusion/models/lingbot_world/test_lingbot_world_transformer.py +++ b/tests/diffusion/models/lingbot_world/test_lingbot_world_transformer.py @@ -1,5 +1,5 @@ # SPDX-License-Identifier: Apache-2.0 -# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project from __future__ import annotations @@ -7,14 +7,17 @@ import inspect import json import math +from datetime import timedelta from pathlib import Path +from types import SimpleNamespace import pytest import torch from tests.diffusion.models.lingbot_world import test_lingbot_world_attention as attention_tests +from tests.helpers.mark import hardware_test -pytestmark = [pytest.mark.core_model, pytest.mark.diffusion, pytest.mark.cpu] +pytestmark = [pytest.mark.core_model, pytest.mark.diffusion] _FIXTURE_PATH = Path(__file__).with_name("fixtures") / "lingbot_world_weight_index_names.json.fixture" _SHAPE_FIXTURE_PATH = Path(__file__).with_name("fixtures") / "lingbot_world_official_shapes.json.fixture" @@ -117,6 +120,7 @@ def _assert_self_cache_unchanged(cache, snapshot) -> None: attention_tests._assert_cache_unchanged(layer_cache, layer_snapshot) +@pytest.mark.cpu def test_tiny_transformer_runs_four_chunks_with_explicit_cache_commit_and_camera_path() -> None: torch.manual_seed(7) module = attention_tests._load_module() @@ -191,6 +195,7 @@ def test_tiny_transformer_runs_four_chunks_with_explicit_cache_commit_and_camera assert not torch.equal(camera_output, alternate_output) +@pytest.mark.cpu def test_partial_cross_attention_cache_projects_text_only_for_missing_layers() -> None: module = attention_tests._load_module() model = _tiny_model(module, num_layers=2).eval() @@ -229,6 +234,7 @@ def test_partial_cross_attention_cache_projects_text_only_for_missing_layers() - @pytest.mark.parametrize("frames", [1, 6], ids=("partial", "multiple_blocks")) +@pytest.mark.cpu def test_forward_rejects_chunks_that_do_not_equal_configured_block_size(frames: int) -> None: module = attention_tests._load_module() model = _tiny_model( @@ -257,6 +263,7 @@ def test_forward_rejects_chunks_that_do_not_equal_configured_block_size(frames: assert cache.cross_attention == [None] +@pytest.mark.cpu def test_forward_accepts_exactly_one_configured_frame_block() -> None: module = attention_tests._load_module() model = _tiny_model( @@ -279,6 +286,7 @@ def test_forward_accepts_exactly_one_configured_frame_block() -> None: assert output.shape == (1, 2, 3, 4, 4) +@pytest.mark.cpu def test_transformer_allocates_request_cache_from_its_configured_geometry() -> None: module = attention_tests._load_module() model = _tiny_model( @@ -303,6 +311,7 @@ def test_transformer_allocates_request_cache_from_its_configured_geometry() -> N assert all(layer.value.shape == (2, 6 * 4 * 6, 2, 2) for layer in cache.self_attention) +@pytest.mark.cpu def test_video_patch_embedding_uses_temporal_height_width_token_order() -> None: module = attention_tests._load_module() model = _tiny_model(module, num_layers=1, num_frames_per_block=2) @@ -323,6 +332,7 @@ def test_video_patch_embedding_uses_temporal_height_width_token_order() -> None: torch.testing.assert_close(tokens[0, :, 0], torch.tensor([1.0, 3.0, 9.0, 11.0])) +@pytest.mark.cpu def test_unpatchify_restores_two_frame_channel_and_spatial_order() -> None: module = attention_tests._load_module() model = _tiny_model(module, num_layers=1, num_frames_per_block=2) @@ -363,6 +373,7 @@ def test_unpatchify_restores_two_frame_channel_and_spatial_order() -> None: torch.testing.assert_close(output, expected) +@pytest.mark.cpu def test_head_modulation_broadcasts_distinct_condition_per_frame() -> None: module = attention_tests._load_module() model = _tiny_model(module, num_layers=1, num_frames_per_block=2) @@ -381,6 +392,7 @@ def test_head_modulation_broadcasts_distinct_condition_per_frame() -> None: torch.testing.assert_close(output, expected) +@pytest.mark.cpu def test_constructor_defaults_match_official_checkpoint_config() -> None: module = attention_tests._load_module() parameters = inspect.signature(module.CausalLingBotWorldTransformer3DModel.__init__).parameters @@ -406,6 +418,7 @@ def test_constructor_defaults_match_official_checkpoint_config() -> None: assert {name: parameters[name].default for name in expected} == expected +@pytest.mark.cpu def test_transformer_exposes_parameter_dtype_for_pipeline_runtime() -> None: module = attention_tests._load_module() model = _tiny_model(module) @@ -413,12 +426,14 @@ def test_transformer_exposes_parameter_dtype_for_pipeline_runtime() -> None: assert model.dtype == next(model.parameters()).dtype +@pytest.mark.cpu def test_transformer_declares_regional_compile_block() -> None: module = attention_tests._load_module() assert module.CausalLingBotWorldTransformer3DModel._repeated_blocks == ["LingBotAttentionBlock"] +@pytest.mark.cpu def test_lingbot_rms_norm_uses_global_tp_square_mean(monkeypatch) -> None: module = attention_tests._load_module() reduced_values: list[torch.Tensor] = [] @@ -440,6 +455,7 @@ def all_reduce(value: torch.Tensor) -> torch.Tensor: torch.testing.assert_close(reduced_values[0], torch.tensor([[25.0]])) +@pytest.mark.cpu def test_attention_rejects_heads_not_divisible_by_tp_size(monkeypatch) -> None: module = attention_tests._load_module() monkeypatch.setattr(module, "get_tensor_model_parallel_world_size", lambda: 3) @@ -448,6 +464,7 @@ def test_attention_rejects_heads_not_divisible_by_tp_size(monkeypatch) -> None: module.LingBotSelfAttention(dim=4, num_heads=2) +@pytest.mark.cpu def test_constructor_rejects_unsupported_qk_norm_with_config_error() -> None: module = attention_tests._load_module() @@ -456,6 +473,7 @@ def test_constructor_rejects_unsupported_qk_norm_with_config_error() -> None: @pytest.mark.parametrize("field", ["image_dim", "added_kv_proj_dim", "pos_embed_seq_len"]) +@pytest.mark.cpu def test_constructor_rejects_non_null_image_embedding_fields(field: str) -> None: module = attention_tests._load_module() @@ -463,6 +481,7 @@ def test_constructor_rejects_non_null_image_embedding_fields(field: str) -> None module.CausalLingBotWorldTransformer3DModel(**{field: 4}) +@pytest.mark.cpu def test_constructor_rejects_unsupported_quantization_with_runtime_error() -> None: module = attention_tests._load_module() @@ -470,6 +489,7 @@ def test_constructor_rejects_unsupported_quantization_with_runtime_error() -> No module.CausalLingBotWorldTransformer3DModel(quant_config=object()) +@pytest.mark.cpu def test_from_config_accepts_diffusers_metadata_and_normalizes_patch_size() -> None: module = attention_tests._load_module() @@ -522,6 +542,7 @@ def __init__(self, **kwargs) -> None: ("sink_size", 0), ], ) +@pytest.mark.cpu def test_from_config_rejects_checkpoint_topology_drift(field: str, value: object) -> None: module = attention_tests._load_module() config = { @@ -554,6 +575,7 @@ def test_from_config_rejects_checkpoint_topology_drift(field: str, value: object module.CausalLingBotWorldTransformer3DModel.from_config(config) +@pytest.mark.cpu def test_from_config_ignores_non_semantic_checkpoint_metadata() -> None: module = attention_tests._load_module() config = { @@ -569,6 +591,7 @@ def test_from_config_ignores_non_semantic_checkpoint_metadata() -> None: assert model.config.num_layers == 40 +@pytest.mark.cpu def test_load_weights_uses_parameter_loaders_and_rejects_unknown_model_keys() -> None: module = attention_tests._load_module() model = _tiny_model(module, num_layers=1) @@ -606,6 +629,7 @@ def record_loader(param: torch.Tensor, loaded_weight: torch.Tensor, shard_id: st model.load_weights([("unexpected_model.weight", torch.ones(1))]) +@pytest.mark.cpu def test_load_weights_consumes_checkpoint_iterator_incrementally() -> None: module = attention_tests._load_module() model = _tiny_model(module, num_layers=1) @@ -619,6 +643,7 @@ def weights(): assert model.load_weights(weights()) == {first_name} +@pytest.mark.cpu def test_real_auto_weights_loader_covers_fused_model_parameter_namespace() -> None: utils = pytest.importorskip("vllm.model_executor.models.utils") module = attention_tests._load_module() @@ -638,6 +663,7 @@ def __init__(self) -> None: assert set(dict(pipeline.named_parameters())) <= loaded +@pytest.mark.cpu def test_checkpoint_weight_index_fixture_matches_model_namespaces() -> None: module = attention_tests._load_module() fixture = json.loads(_FIXTURE_PATH.read_text()) @@ -667,6 +693,7 @@ def test_checkpoint_weight_index_fixture_matches_model_namespaces() -> None: assert all(len(shard["parameter_names_sha256"]) == 64 for shard in fixture["shards"].values()) +@pytest.mark.cpu def test_official_default_shapes_match_public_safetensors_header_fixture() -> None: module = attention_tests._load_module() fixture = json.loads(_SHAPE_FIXTURE_PATH.read_text()) @@ -683,3 +710,284 @@ def test_official_default_shapes_match_public_safetensors_header_fixture() -> No shape, dtype = parameter_specs[name] assert tuple(shape) == tuple(expected_shape), name assert dtype == expected_dtype + + +@pytest.mark.cpu +def test_timestep_expansion_only_runs_with_ulysses(monkeypatch): + dtype = torch.bfloat16 + module = attention_tests._load_module() + hidden = torch.randn(2, 12, 4).to(dtype) + camera = torch.randn_like(hidden) + table = torch.randn(2, 3, 6, 4).to(dtype) + rotary = (torch.randn(12, 2), torch.randn(12, 2)) + prepare = module._LingBotSPPrepare() + assert prepare(hidden, camera, table, rotary)[2] is table + monkeypatch.setattr(module, "get_sp_group", lambda: SimpleNamespace(ulysses_world_size=2)) + expanded = prepare(hidden, camera, table, rotary)[2] + assert expanded.shape == (2, 12, 6, 4) + torch.testing.assert_close(expanded, table.repeat_interleave(4, dim=1), rtol=0, atol=0) + + +@pytest.mark.parametrize("unsharded", [None, 3], ids=["missing-hook", "unsharded-rope"]) +@pytest.mark.cpu +def test_missing_or_partial_sp_split_fails_before_attention(monkeypatch, unsharded): + module = attention_tests._load_module() + model = _tiny_model(module, num_frames_per_block=3, sliding_window_num_frames=6) + cache = _cache(module, model) + for block in model.blocks: + block.self_attn.ulysses_world_size = 2 + monkeypatch.setattr(module, "get_sp_group", lambda: SimpleNamespace(ulysses_world_size=2)) + original = model.sp_prepare.forward + + def partial_split(*args): + values = original(*args) + if unsharded is None: + return values + return tuple( + value if i == unsharded else value.chunk(2, dim=1 if i < 3 else 0)[0] for i, value in enumerate(values) + ) + + monkeypatch.setattr(model.sp_prepare, "forward", partial_split) + + def forbidden(*args, **kwargs): + raise AssertionError("attention/collectives must not run before shard validation") + + monkeypatch.setattr(model.blocks[0], "forward", forbidden) + with pytest.raises(RuntimeError, match="SP input hooks"): + model( + torch.randn(1, 36, 3, 4, 4), + torch.tensor([1.0]), + torch.randn(1, 3, 6), + torch.randn(1, 384, 3, 4, 4), + cache=cache, + start_frame=0, + update_cache=False, + ) + + +# Real multi-rank regression: synthetic weights, native SP/TP and attention. +_HEADS, _HEAD_DIM, _LAYERS = 8, 32, 2 +_FRAMES, _SIDE, _TOKENS_PER_FRAME = 3, 32, 256 + + +def _model(dtype): + from vllm_omni.diffusion.models.lingbot_world.transformer import CausalLingBotWorldTransformer3DModel + + model = CausalLingBotWorldTransformer3DModel( + num_attention_heads=_HEADS, + attention_head_dim=_HEAD_DIM, + num_layers=_LAYERS, + in_channels=4, + out_channels=4, + text_dim=8, + freq_dim=16, + ffn_dim=512, + patch_size=(1, 2, 2), + sink_size=1, + num_frames_per_block=_FRAMES, + sliding_window_num_frames=6, + rope_max_seq_len=32, + ).eval() + return model.to(device=torch.device("cuda", torch.accelerator.current_device_index()), dtype=dtype) + + +def _checkpoint(model): + weights = {} + generator = torch.Generator().manual_seed(17) + for name, param in sorted(model.named_parameters()): + data = torch.randn(param.shape, generator=generator) * 0.02 + if "norm" in name and name.endswith("weight"): + data.fill_(1) + if name.endswith("modulation") and data.shape[-2] == 6: + data[..., 2, :] = 1 + data[..., 5, :] = 1 + data = data.to(param.dtype) + param.copy_(data) + if ".self_attn.qkv." in name: + for component, chunk in zip(("q", "k", "v"), data.chunk(3), strict=True): + weights[name.replace(".self_attn.qkv.", f".self_attn.{component}.")] = chunk + else: + weights[name] = data + return weights + + +def _rollout(model, mode, dtype, batch): + from vllm_omni.diffusion.models.lingbot_world.transformer import LingBotTransformerCache + from vllm_omni.experimental.ar_diffusion.capability import ARDiffusionKVBranchSpec + from vllm_omni.experimental.ar_diffusion.kv_cache import ARDiffusionKVCache, ARDiffusionKVConfig + from vllm_omni.experimental.ar_diffusion.kv_cache.state import ARDiffusionKVState + + device = next(model.parameters()).device + state = None + if mode == "paged": + kv = ARDiffusionKVCache( + ARDiffusionKVConfig(enable=True, chunk_size=_TOKENS_PER_FRAME, window_chunks=5, sink_chunks=1), + num_layers=_LAYERS, + num_kv_heads=model.blocks[0].self_attn.num_sp_heads, + head_size=_HEAD_DIM, + dtype=dtype, + block_size=_TOKENS_PER_FRAME, + max_model_len=4096, + available_bytes=1 << 27, + kv_branches=(ARDiffusionKVBranchSpec("main", 0),), + session_capacity=1, + frames_per_block=_FRAMES, + max_scratch_tokens_per_branch=_FRAMES * _TOKENS_PER_FRAME, + cross_attention_lengths={"text": 5}, + device=device, + ) + state = ARDiffusionKVState(kv, "numeric", {"main": kv.begin_request("numeric")}, num_layers=_LAYERS) + cache = LingBotTransformerCache(self_attention=[], cross_attention=[None] * _LAYERS) + else: + cache = model.allocate_cache( + batch_size=batch, latent_height=_SIDE, latent_width=_SIDE, device=device, dtype=dtype + ) + generator = torch.Generator().manual_seed(123) + text = torch.randn(batch, 5, 8, generator=generator).to(device, dtype) + outputs = [] + try: + for start in range(0, 12, _FRAMES): + latent = torch.randn(batch, 4, _FRAMES, _SIDE, _SIDE, generator=generator).to(device, dtype) + camera = torch.randn(batch, 384, _FRAMES, _SIDE, _SIDE, generator=generator).to(device, dtype) + for commit in (False, False, True): + if state is not None: + cache.self_attention = state.get_kv_caches( + "main", seq_len=_FRAMES * _TOKENS_PER_FRAME, commit_current=commit + ) + output = model( + latent, + torch.tensor([100.0, 400.0, 900.0], device=device).expand(batch, -1), + text, + camera, + cache=cache, + start_frame=start, + update_cache=commit, + ) + outputs.append(output.float().cpu()) + if state is not None: + if not state.is_cross_attention_populated("main", "text"): + state.populate_cross_attention( + "main", "text", [(layer.key, layer.value) for layer in cache.cross_attention] + ) + for layer, pool in zip( + cache.cross_attention, state.get_cross_attention_kv("main", "text"), strict=True + ): + layer.key, layer.value = pool["k"], pool["v"] + state.commit_paged_context("main") + cross = [ + (layer.key.detach().float().cpu(), layer.value.detach().float().cpu()) for layer in cache.cross_attention + ] + return torch.stack(outputs), cross + finally: + if state is not None: + state.close() + + +def _worker(rank, world_size, sp_size, tp_size, mode, dtype, batch, rendezvous): + from vllm.config import VllmConfig + from vllm.config.vllm import set_current_vllm_config + from vllm.distributed import get_tensor_model_parallel_rank + + from vllm_omni.diffusion.config import set_current_diffusion_config + from vllm_omni.diffusion.data import AttentionConfig, AttentionSpec, DiffusionParallelConfig, OmniDiffusionConfig + from vllm_omni.diffusion.distributed.parallel_state import ( + destroy_distributed_env, + destroy_model_parallel, + get_sp_group, + init_distributed_environment, + initialize_model_parallel, + ) + from vllm_omni.diffusion.forward_context import set_forward_context + from vllm_omni.diffusion.registry import _apply_sequence_parallel_if_enabled + from vllm_omni.platforms import current_omni_platform + + current_omni_platform.set_device(torch.device("cuda", rank)) + torch.backends.cuda.matmul.allow_tf32 = False + torch.backends.cudnn.allow_tf32 = False + torch.distributed.init_process_group( + "nccl", init_method=f"file://{rendezvous}", world_size=world_size, rank=rank, timeout=timedelta(seconds=90) + ) + init_distributed_environment(world_size=world_size, rank=rank, local_rank=rank, backend="nccl") + try: + with torch.inference_mode(), set_current_vllm_config(VllmConfig()): + for baseline in (True, False): + sp, tp = (1, 1) if baseline else (sp_size, tp_size) + initialize_model_parallel( + data_parallel_size=world_size // (sp * tp), + sequence_parallel_size=sp, + ulysses_degree=sp, + tensor_parallel_size=tp, + ) + config = OmniDiffusionConfig( + model=str(Path(rendezvous).parent), + dtype=dtype, + enforce_eager=True, + parallel_config=DiffusionParallelConfig( + data_parallel_size=world_size // (sp * tp), + sequence_parallel_size=sp, + ulysses_degree=sp, + tensor_parallel_size=tp, + ), + diffusion_attention_config=AttentionConfig(default=AttentionSpec(backend="TORCH_SDPA")), + ) + with set_current_diffusion_config(config), set_forward_context(omni_diffusion_config=config): + model = _model(dtype) + if baseline: + weights = _checkpoint(model) + else: + model.load_weights(iter(weights.items())) + _apply_sequence_parallel_if_enabled(SimpleNamespace(transformer=model), config) + result, cross = _rollout(model, mode, dtype, batch) + if baseline: + expected, expected_cross = result, cross + else: + assert expected.abs().max() > 1e-3 + bound = 1e-5 if dtype == torch.float32 else 1e-2 + error = (result - expected).double().norm() / expected.double().norm() + assert error <= bound, f"rank={rank}, TP{tp} SP{sp}: relative L2 {error.item():.3g} > {bound}" + local_heads = _HEADS // tp // sp + first = ( + get_tensor_model_parallel_rank() * (_HEADS // tp) + + get_sp_group().ulysses_rank * local_heads + ) + for (k, v), (ek, ev) in zip(cross, expected_cross, strict=True): + tolerance = 1e-5 if dtype == torch.float32 else 2e-2 + torch.testing.assert_close( + k, ek[:, :, first : first + local_heads], rtol=tolerance, atol=tolerance + ) + torch.testing.assert_close( + v, ev[:, :, first : first + local_heads], rtol=tolerance, atol=tolerance + ) + if rank == 0: + print(f"{mode} {dtype} TP{tp} SP{sp}: relative L2={error.item():.3g}", flush=True) + del model + torch.distributed.barrier() + destroy_model_parallel() + finally: + destroy_distributed_env() + + +def _run(tmp_path: Path, sp, tp, mode, dtype, batch): + if torch.accelerator.device_count() < sp * tp: + pytest.skip(f"requires {sp * tp} GPUs") + torch.multiprocessing.spawn( + _worker, args=(sp * tp, sp, tp, mode, dtype, batch, str(tmp_path / "rendezvous")), nprocs=sp * tp + ) + + +@pytest.mark.parallel +@hardware_test(res={"cuda": "L4"}, num_cards=2) +def test_sp2_direct_matches_sp1_fp32(tmp_path): + _run(tmp_path, 2, 1, "direct", torch.float32, 2) + + +@pytest.mark.parallel +@hardware_test(res={"cuda": "L4"}, num_cards=4) +def test_sp4_direct_matches_sp1_bf16(tmp_path): + _run(tmp_path, 4, 1, "direct", torch.bfloat16, 1) + + +@pytest.mark.parallel +@hardware_test(res={"cuda": "L4"}, num_cards=4) +def test_tp2_sp2_paged_matches_sp1_bf16(tmp_path): + _run(tmp_path, 2, 2, "paged", torch.bfloat16, 1) diff --git a/tests/diffusion/models/lingbot_world/test_pipeline_lingbot_world.py b/tests/diffusion/models/lingbot_world/test_pipeline_lingbot_world.py index 52113208a74..a58eee935db 100644 --- a/tests/diffusion/models/lingbot_world/test_pipeline_lingbot_world.py +++ b/tests/diffusion/models/lingbot_world/test_pipeline_lingbot_world.py @@ -126,8 +126,6 @@ def __init__(self, *, raise_on_call: int | None = None, dtype: torch.dtype = tor sink_size=3, ) self.blocks = nn.ModuleList([nn.Identity(), nn.Identity()]) - for block in self.blocks: - block.self_attn = SimpleNamespace(num_local_heads=2, head_dim=4) self.calls: list[dict] = [] self.cache_allocations: list[dict] = [] self.raise_on_call = raise_on_call @@ -387,6 +385,9 @@ def _od_config(**overrides): "parallel_config": SimpleNamespace( pipeline_parallel_size=1, sequence_parallel_size=1, + ulysses_degree=1, + ring_degree=1, + allgather_degree=1, cfg_parallel_size=1, vae_patch_parallel_size=1, use_hsdp=False, @@ -438,15 +439,15 @@ def test_pipeline_respects_loader_managed_component_placement(offload_field: str assert getattr(pipeline.vae, "to_calls", []) == [] -def test_ar_diffusion_capability_uses_fixed_tp_local_lingbot_geometry() -> None: +def test_ar_diffusion_capability_uses_transformer_local_head_geometry() -> None: module = _load_pipeline_module() - module.get_tensor_model_parallel_world_size = lambda: 1 pipeline = _pipeline(module) - + # Use the constructed head count even when the config describes a different geometry. + pipeline.transformer.blocks[0].self_attn = SimpleNamespace(num_sp_heads=1) spec = pipeline.ar_diffusion_kv_cache_spec() assert spec.num_layers == 2 - assert spec.num_kv_heads == 2 + assert spec.num_kv_heads == 1 assert spec.head_size == 4 assert spec.tokens_per_frame == 1 assert spec.frames_per_block == 3 @@ -604,7 +605,6 @@ def test_component_discovery_uses_official_checkpoint_contract() -> None: ("field", "value", "feature"), [ ("pipeline_parallel_size", 2, "pipeline parallelism"), - ("sequence_parallel_size", 2, "sequence parallelism"), ("cfg_parallel_size", 2, "CFG parallelism"), ("vae_patch_parallel_size", 2, "VAE parallelism"), ("use_hsdp", True, "HSDP"), @@ -622,6 +622,43 @@ def test_unsupported_parallel_modes_fail_before_component_loading(field: str, va assert module._loader_state.prefetch_calls == [] +def test_pure_ulysses_parallel_config_is_supported() -> None: + module = _load_pipeline_module() + parallel_config = _od_config().parallel_config + parallel_config.sequence_parallel_size = 2 + parallel_config.ulysses_degree = 2 + + pipeline = module.LingBotWorldCausalDMDPipeline(od_config=_od_config(parallel_config=parallel_config)) + + assert pipeline.transformer is not None + + +@pytest.mark.parametrize( + "overrides", + [ + {"sequence_parallel_size": 4, "ring_degree": 2}, # Normalized hybrid. + {"ulysses_degree": 1, "allgather_degree": 2}, # Normalized AllGather-KV. + {"ring_degree": 2}, # Isolate each clause from the SP-size mismatch. + {"allgather_degree": 2}, + {"ulysses_mode": "advanced_uaa"}, + {"ulysses_a2a_permute": True}, + {"ulysses_degree": None}, # Missing degree must not imply pure Ulysses. + ], +) +def test_unsupported_sp_config_fails_before_component_loading(overrides): + module = _load_pipeline_module() + config = _od_config().parallel_config + config.sequence_parallel_size = config.ulysses_degree = 2 + for name, value in overrides.items(): + if value is None: + delattr(config, name) + else: + setattr(config, name, value) + with pytest.raises(NotImplementedError, match="pure Ulysses"): + module.LingBotWorldCausalDMDPipeline(od_config=_od_config(parallel_config=config)) + assert module._loader_state.prefetch_calls == [] + + def test_unsupported_quantization_fails_before_component_loading() -> None: module = _load_pipeline_module() diff --git a/tests/examples/offline_inference/test_lingbot_world_v2.py b/tests/examples/offline_inference/test_lingbot_world_v2.py new file mode 100644 index 00000000000..41d2e73b712 --- /dev/null +++ b/tests/examples/offline_inference/test_lingbot_world_v2.py @@ -0,0 +1,24 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project + +from pathlib import Path + +import pytest + +from examples.offline_inference.diffusion import lingbot_world_v2 as example + +pytestmark = [pytest.mark.core_model, pytest.mark.cpu, pytest.mark.diffusion] + + +def test_offline_ulysses_argument(tmp_path): + argv = ["--image", "frame.jpg", "--action-dir", "forward", "--prompt", "A lake"] + paths = example.LingBotPaths( + tmp_path / "frame.jpg", tmp_path / "forward", tmp_path, Path("forward"), 9, tmp_path / "out.mp4" + ) + assert example.parse_args(argv).ulysses_degree == 1 + args = example.parse_args([*argv, "--ulysses-degree", "4"]) + kwargs = example.build_omni_kwargs(args, paths) + assert kwargs["ulysses_degree"] == 4 and kwargs["tensor_parallel_size"] == 1 + args.ulysses_degree = 0 + with pytest.raises(ValueError, match="--ulysses-degree"): + example.build_omni_kwargs(args, paths) diff --git a/vllm_omni/diffusion/models/lingbot_world/pipeline.py b/vllm_omni/diffusion/models/lingbot_world/pipeline.py index b41abd861d6..d2dfde3d4e9 100644 --- a/vllm_omni/diffusion/models/lingbot_world/pipeline.py +++ b/vllm_omni/diffusion/models/lingbot_world/pipeline.py @@ -18,7 +18,6 @@ from diffusers.utils.torch_utils import randn_tensor from torch import nn from transformers import AutoTokenizer, UMT5EncoderModel -from vllm.distributed import get_tensor_model_parallel_world_size from vllm.model_executor.models.utils import AutoWeightsLoader from vllm_omni.diffusion.data import DiffusionOutput, OmniDiffusionConfig @@ -203,7 +202,6 @@ def _validate_parallel_config(od_config: OmniDiffusionConfig) -> None: return unsupported_sizes = { "pipeline_parallel_size": "pipeline parallelism", - "sequence_parallel_size": "sequence parallelism", "cfg_parallel_size": "CFG parallelism", "vae_patch_parallel_size": "VAE parallelism", } @@ -211,6 +209,25 @@ def _validate_parallel_config(od_config: OmniDiffusionConfig) -> None: size = getattr(parallel_config, field, 1) or 1 if size > 1: raise NotImplementedError(f"LingBot World v1 does not support {feature} ({field}={size}).") + sequence_parallel_size = getattr(parallel_config, "sequence_parallel_size", 1) or 1 + ulysses_degree = getattr(parallel_config, "ulysses_degree", 1) or 1 + ring_degree = getattr(parallel_config, "ring_degree", 1) or 1 + allgather_degree = getattr(parallel_config, "allgather_degree", 1) or 1 + ulysses_mode = getattr(parallel_config, "ulysses_mode", "strict") + ulysses_a2a_permute = bool(getattr(parallel_config, "ulysses_a2a_permute", False)) + if ( + sequence_parallel_size != ulysses_degree + or ring_degree != 1 + or allgather_degree != 1 + or ulysses_mode != "strict" + or ulysses_a2a_permute + ): + raise NotImplementedError( + "LingBot World sequence parallelism requires pure Ulysses with ulysses_a2a_permute disabled: " + f"sequence_parallel_size={sequence_parallel_size}, ulysses_degree={ulysses_degree}, " + f"ring_degree={ring_degree}, allgather_degree={allgather_degree}, ulysses_mode={ulysses_mode!r}, " + f"ulysses_a2a_permute={ulysses_a2a_permute}." + ) if getattr(parallel_config, "use_hsdp", False): raise NotImplementedError("LingBot World v1 does not support HSDP.") if getattr(parallel_config, "enable_expert_parallel", False): @@ -557,8 +574,6 @@ def ar_diffusion_kv_cache_spec(self) -> ARDiffusionKVCacheSpec: latent_height = self._ar_height // spatial latent_width = self._ar_width // spatial tokens_per_frame = (latent_height // patch_height) * (latent_width // patch_width) - tp_size = get_tensor_model_parallel_world_size() - num_local_heads = int(self.transformer.config.num_attention_heads) // tp_size total_window_frames = ( int(self.transformer.config.local_attn_size) if int(self.transformer.config.local_attn_size) != -1 @@ -579,7 +594,7 @@ def ar_diffusion_kv_cache_spec(self) -> ARDiffusionKVCacheSpec: ) return ARDiffusionKVCacheSpec( num_layers=int(self.transformer.config.num_layers), - num_kv_heads=num_local_heads, + num_kv_heads=int(self.transformer.blocks[0].self_attn.num_sp_heads), head_size=int(self.transformer.config.attention_head_dim), tokens_per_frame=tokens_per_frame, frames_per_block=int(self.transformer.config.num_frames_per_block), @@ -960,20 +975,11 @@ def _ar_text_caches( def layer_kv() -> Iterator[tuple[torch.Tensor, torch.Tensor]]: for block in self.transformer.blocks: cross_attention = block.cross_attn - key = cross_attention.norm_k(cross_attention.k(projected_text)).unflatten( - 2, - ( - cross_attention.num_local_heads, - cross_attention.head_dim, - ), - ) - value = cross_attention.v(projected_text).unflatten( - 2, - ( - cross_attention.num_local_heads, - cross_attention.head_dim, - ), + shape = (cross_attention.num_local_heads, cross_attention.head_dim) + key = cross_attention.shard_kv_heads( + cross_attention.norm_k(cross_attention.k(projected_text)).unflatten(2, shape) ) + value = cross_attention.shard_kv_heads(cross_attention.v(projected_text).unflatten(2, shape)) yield key, value state.populate_cross_attention( diff --git a/vllm_omni/diffusion/models/lingbot_world/transformer.py b/vllm_omni/diffusion/models/lingbot_world/transformer.py index 35ec1880112..32931db0c9e 100644 --- a/vllm_omni/diffusion/models/lingbot_world/transformer.py +++ b/vllm_omni/diffusion/models/lingbot_world/transformer.py @@ -1,5 +1,5 @@ # SPDX-License-Identifier: Apache-2.0 -# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project """Checkpoint-compatible causal DiT for the LingBot World v2 model package.""" from __future__ import annotations @@ -24,6 +24,9 @@ from vllm.model_executor.utils import set_weight_attrs from vllm_omni.diffusion.attention.layer import Attention +from vllm_omni.diffusion.distributed.comm import SeqAllToAll4D +from vllm_omni.diffusion.distributed.parallel_state import get_sp_group +from vllm_omni.diffusion.distributed.sp_plan import SequenceParallelInput, SequenceParallelOutput from vllm_omni.diffusion.layers.norm import LayerNorm from vllm_omni.diffusion.layers.rope import RotaryEmbeddingWan from vllm_omni.experimental.ar_diffusion.kv_cache.paged_attention import ( @@ -127,6 +130,37 @@ def _projection_prefix(prefix: str, name: str) -> str: return f"{prefix}.{name}" if prefix else name +def _ulysses_state() -> tuple[int, int, torch.distributed.ProcessGroup | None]: + coordinator = get_sp_group() + return ( + int(coordinator.ulysses_world_size), + int(coordinator.ulysses_rank), + coordinator.ulysses_group, + ) + + +class _LingBotSPPrepare(nn.Module): + """Expand frame conditioning only when Ulysses shards the token sequence.""" + + def forward( + self, + hidden_states: torch.Tensor, + camera_hidden_states: torch.Tensor, + timestep_projection: torch.Tensor, + rotary_emb: tuple[torch.Tensor, torch.Tensor], + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + num_frames = timestep_projection.shape[1] + if hidden_states.shape[1] % num_frames: + raise ValueError("LingBot token count must be divisible by the latent frame count.") + tokens_per_frame = hidden_states.shape[1] // num_frames + if get_sp_group().ulysses_world_size > 1: + timestep_projection = ( + timestep_projection.unsqueeze(2).expand(-1, -1, tokens_per_frame, -1, -1).flatten(1, 2) + ) + cosine, sine = rotary_emb + return hidden_states, camera_hidden_states, timestep_projection, cosine, sine + + class LingBotSelfAttention(nn.Module): """Block-causal self-attention over retained history and one full chunk.""" @@ -157,6 +191,13 @@ def __init__( prefix=_projection_prefix(prefix, "qkv"), ) self.num_local_heads = self.qkv.num_heads + self.ulysses_world_size, self.ulysses_rank, self.ulysses_group = _ulysses_state() + if self.num_local_heads % self.ulysses_world_size: + raise ValueError( + "LingBot local attention heads must be divisible by the Ulysses degree: " + f"heads={self.num_local_heads}, ulysses={self.ulysses_world_size}." + ) + self.num_sp_heads = self.num_local_heads // self.ulysses_world_size self.tp_inner_dim = self.num_local_heads * self.head_dim self.o = RowParallelLinear( dim, @@ -170,9 +211,9 @@ def __init__( self.norm_k = _LingBotRMSNorm(self.tp_inner_dim, eps) self.rotary_embedding = RotaryEmbeddingWan(is_neox_style=False, half_head_dim=True) self.attn = Attention( - num_heads=self.num_local_heads, + num_heads=self.num_sp_heads, head_size=self.head_dim, - num_kv_heads=self.num_local_heads, + num_kv_heads=self.num_sp_heads, softmax_scale=self.head_dim**-0.5, causal=False, role="self", @@ -287,6 +328,11 @@ def forward( query = self.rotary_embedding(query, cos, sin) key = self.rotary_embedding(key, cos, sin) + if self.ulysses_world_size > 1: + query = SeqAllToAll4D.apply(self.ulysses_group, query, 2, 1, False) + key = SeqAllToAll4D.apply(self.ulysses_group, key, 2, 1, False) + value = SeqAllToAll4D.apply(self.ulysses_group, value, 2, 1, False) + if isinstance(cache, ARDiffusionPagedLayerInputs): if query.shape[0] != 1: raise RuntimeError("LingBot AR-Diffusion paged attention requires batch_size=1.") @@ -347,6 +393,8 @@ def forward( ) else: output = self.attn(query, visible_key, visible_value) + if self.ulysses_world_size > 1: + output = SeqAllToAll4D.apply(self.ulysses_group, output, 1, 2, False) return self.o(output.flatten(2, 3)) @@ -373,6 +421,13 @@ def __init__( self.num_heads = num_heads self.head_dim = dim // num_heads self.num_local_heads = num_heads // tp_size + self.ulysses_world_size, self.ulysses_rank, self.ulysses_group = _ulysses_state() + if self.num_local_heads % self.ulysses_world_size: + raise ValueError( + "LingBot local cross-attention heads must be divisible by the Ulysses degree: " + f"heads={self.num_local_heads}, ulysses={self.ulysses_world_size}." + ) + self.num_sp_heads = self.num_local_heads // self.ulysses_world_size self.tp_inner_dim = self.num_local_heads * self.head_dim self.q = ColumnParallelLinear( @@ -410,9 +465,9 @@ def __init__( self.norm_q = _LingBotRMSNorm(self.tp_inner_dim, eps) self.norm_k = _LingBotRMSNorm(self.tp_inner_dim, eps) self.attn = Attention( - num_heads=self.num_local_heads, + num_heads=self.num_sp_heads, head_size=self.head_dim, - num_kv_heads=self.num_local_heads, + num_kv_heads=self.num_sp_heads, softmax_scale=self.head_dim**-0.5, causal=False, role="cross", @@ -422,6 +477,14 @@ def __init__( disable_kv_quant=True, ) + def shard_kv_heads(self, value: torch.Tensor) -> torch.Tensor: + """Select this Ulysses rank's projected K/V heads.""" + + if self.ulysses_world_size == 1: + return value + start = self.ulysses_rank * self.num_sp_heads + return value[:, :, start : start + self.num_sp_heads].clone(memory_format=torch.contiguous_format) + def forward( self, hidden_states: torch.Tensor, @@ -431,6 +494,8 @@ def forward( ) -> tuple[torch.Tensor, LingBotAttentionCache]: query = self.norm_q(self.q(hidden_states)) query = query.unflatten(2, (self.num_local_heads, self.head_dim)) + if self.ulysses_world_size > 1: + query = SeqAllToAll4D.apply(self.ulysses_group, query, 2, 1, False) # Text K/V is constant within a request and is projected once per layer. if cache is None: @@ -438,8 +503,8 @@ def forward( raise ValueError("encoder_hidden_states are required when the cross-attention cache is empty.") key = self.norm_k(self.k(encoder_hidden_states)) value = self.v(encoder_hidden_states) - key = key.unflatten(2, (self.num_local_heads, self.head_dim)) - value = value.unflatten(2, (self.num_local_heads, self.head_dim)) + key = self.shard_kv_heads(key.unflatten(2, (self.num_local_heads, self.head_dim))) + value = self.shard_kv_heads(value.unflatten(2, (self.num_local_heads, self.head_dim))) cache = LingBotAttentionCache( key=key, value=value, @@ -452,6 +517,8 @@ def forward( value = cache.value[:, : cache.end] output = self.attn(query, key, value) + if self.ulysses_world_size > 1: + output = SeqAllToAll4D.apply(self.ulysses_group, output, 1, 2, False) return self.o(output.flatten(2, 3)), cache @@ -540,7 +607,18 @@ def forward( rotary_emb: tuple[torch.Tensor, torch.Tensor], ) -> tuple[torch.Tensor, LingBotAttentionCache]: batch_size, token_count, dim = hidden_states.shape + if ( + timestep_projection.ndim != 4 + or timestep_projection.shape[0] != batch_size + or timestep_projection.shape[2:] != (6, dim) + ): + raise ValueError("LingBot timestep projection must have shape (batch, frames or tokens, 6, dim).") num_frames = timestep_projection.shape[1] + if num_frames == 0 or token_count % num_frames: + raise ValueError("LingBot timestep projection must align with hidden tokens.") + if self.self_attn.ulysses_world_size > 1 and num_frames != token_count: + raise ValueError("LingBot Ulysses requires one timestep projection per local token.") + # Without SP, broadcast per frame. With SP, each entry represents one token. tokens_per_frame = token_count // num_frames # Timestep, camera, and text remain separate conditioning paths. modulation = self.modulation.unsqueeze(1) + timestep_projection.float() @@ -635,13 +713,26 @@ def forward( self, hidden_states: torch.Tensor, timestep_embedding: torch.Tensor, + *, + tokens_per_frame: int | None = None, + token_offset: int = 0, ) -> torch.Tensor: num_frames = timestep_embedding.shape[1] - tokens_per_frame = hidden_states.shape[1] // num_frames modulation = self.modulation.unsqueeze(1) + timestep_embedding.unsqueeze(2).float() shift, scale = modulation.chunk(2, dim=2) - normalized = self.norm(hidden_states.float()).unflatten(1, (num_frames, tokens_per_frame)) - normalized = (normalized * (1 + scale) + shift).flatten(1, 2).to(hidden_states.dtype) + normalized = self.norm(hidden_states.float()) + if tokens_per_frame is None: + normalized = normalized.unflatten(1, (num_frames, -1)) + normalized = (normalized * (1 + scale) + shift).flatten(1, 2) + else: + # A sequence shard can begin/end inside a frame. + frames = ( + torch.arange(hidden_states.shape[1], device=hidden_states.device) + token_offset + ) // tokens_per_frame + scale = scale.squeeze(2).index_select(1, frames) + shift = shift.squeeze(2).index_select(1, frames) + normalized = normalized * (1 + scale) + shift + normalized = normalized.to(hidden_states.dtype) return self.head(normalized) @@ -678,6 +769,16 @@ class CausalLingBotWorldTransformer3DModel(nn.Module): _repeated_blocks = ["LingBotAttentionBlock"] packed_modules_mapping = {"qkv": ["q", "k", "v"]} _layerwise_offload_blocks_attrs = ["blocks"] + _sp_plan = { + "sp_prepare": { + 0: SequenceParallelInput(split_dim=1, expected_dims=3, split_output=True), + 1: SequenceParallelInput(split_dim=1, expected_dims=3, split_output=True), + 2: SequenceParallelInput(split_dim=1, expected_dims=4, split_output=True), + 3: SequenceParallelInput(split_dim=0, expected_dims=2, split_output=True), + 4: SequenceParallelInput(split_dim=0, expected_dims=2, split_output=True), + }, + "sp_output_gather": SequenceParallelOutput(gather_dim=1, expected_dims=3), + } def __init__( self, @@ -768,6 +869,8 @@ def __init__( stride=patch_size, ) self.patch_embedding_wancamctrl = _LingBotCameraPatchEmbedding(6 * 8 * 8, dim, patch_size) + self.sp_prepare = _LingBotSPPrepare() + self.sp_output_gather = nn.Identity() self.c2ws_hidden_states_layer1 = ColumnParallelLinear( dim, dim, @@ -904,8 +1007,7 @@ def allocate_cache( self.config.local_attn_size if self.config.local_attn_size != -1 else self.config.sliding_window_num_frames ) max_tokens = int(window_frames * post_patch_height * post_patch_width) - tp_size = get_tensor_model_parallel_world_size() - num_local_heads = self.config.num_attention_heads // tp_size + num_local_heads = self.blocks[0].self_attn.num_sp_heads return allocate_lingbot_cache( batch_size=batch_size, num_layers=self.config.num_layers, @@ -1138,6 +1240,33 @@ def forward( query_len=hidden_states.shape[1], ) cache.self_attention = [layer_context.to_layer_inputs() for layer_context in cache.self_attention] + sp_size = self.blocks[0].self_attn.ulysses_world_size + global_tokens = patched_frames * tokens_per_frame + hidden_dim = hidden_states.shape[-1] + rope_dim = rotary_emb[0].shape[-1] + if global_tokens % sp_size: + raise ValueError("LingBot Ulysses requires the token count to be divisible by its degree.") + hidden_states, camera_hidden_states, timestep_projection, cosine, sine = self.sp_prepare( + hidden_states, + camera_hidden_states, + timestep_projection, + rotary_emb, + ) + if sp_size > 1: + local_tokens = global_tokens // sp_size + expected_hidden = (batch_size, local_tokens, hidden_dim) + if ( + hidden_states.shape != expected_hidden + or camera_hidden_states.shape != expected_hidden + or timestep_projection.shape != (batch_size, local_tokens, 6, hidden_dim) + or cosine.shape != (local_tokens, rope_dim) + or sine.shape != (local_tokens, rope_dim) + ): + raise RuntimeError( + "LingBot SP input hooks did not shard all conditioning tensors consistently; " + f"expected {local_tokens} tokens per rank. Check SP hook registration and input dimensions." + ) + rotary_emb = (cosine, sine) # Phase 3: each layer receives its own cache entry. Text K/V is passed # only when absent; the returned cache is stored for subsequent DMD # steps and causal blocks in this request. @@ -1156,9 +1285,17 @@ def forward( ) cache.cross_attention[index] = cross_cache - # Phase 4: map tokens to per-patch 16-channel flow values and restore - # [B, C, F, H, W] for the Pipeline's sampler update. - hidden_states = self.head(hidden_states, timestep_embedding) + # Phase 4: project local tokens before gathering the much narrower flow values. + hidden_states = self.head( + hidden_states, + timestep_embedding, + tokens_per_frame=tokens_per_frame if sp_size > 1 else None, + token_offset=get_sp_group().ulysses_rank * (global_tokens // sp_size) if sp_size > 1 else 0, + ) + hidden_states = self.sp_output_gather(hidden_states) + projected_dim = self.config.out_channels * math.prod(self.config.patch_size) + if sp_size > 1 and hidden_states.shape != (batch_size, global_tokens, projected_dim): + raise RuntimeError("LingBot SP output hook did not gather the full projected token sequence.") return self._unpatchify( hidden_states, batch_size=batch_size,