diff --git a/docs/cookbook/diffusion/MiniMax/MiniMax-H3.mdx b/docs/cookbook/diffusion/MiniMax/MiniMax-H3.mdx index a8da1857b..f7f70d640 100644 --- a/docs/cookbook/diffusion/MiniMax/MiniMax-H3.mdx +++ b/docs/cookbook/diffusion/MiniMax/MiniMax-H3.mdx @@ -977,6 +977,59 @@ fails closed when the selected format, projection, or topology is incompatible. + + +Cube sparse attention applies TopK sparsity only to H3's 3D visual streams. +Text, audio, standalone reference images, and the text-only token refiner stay +dense. It runs on pure PyTorch plus FlexAttention, so it has no third-party +kernel dependency. + +Select it for the H3 transformer with +`--component-attention-backends transformer=cube_sparse_attn` and pass +`--attention-backend-config` with both `local_cube_size` and +`topk_ratio_list`. Scoping the backend leaves the text encoder on its native +backend: + +```bash Command +sglang serve \ + --model-path MiniMaxAI/MiniMax-H3 \ + --model-variant ref2va \ + --num-gpus 8 \ + --ulysses-degree 8 \ + --performance-mode speed \ + --component-attention-backends transformer=cube_sparse_attn \ + --attention-backend-config '{"local_cube_size": [4, 4, 4], "topk_ratio_list": [1.0, 1.0, 0.8, 0.7, 0.6, 0.5, 0.5]}' \ + --port 30010 +``` + +- `local_cube_size` is the `(T, H, W)` cube that groups neighboring latent + tokens into one attention block. It must have exactly three entries. +- `topk_ratio_list` sets the per-step keep ratio and must have exactly one + entry per denoise step, each in `(0, 1]`. A ratio of `1.0` keeps a step + on H3's native dense attention path; smaller values select the FlexAttention + sparse path and drop more blocks. The example above matches a request with + `num_inference_steps: 8`, whose endpoint-inclusive sigma schedule has seven + denoise updates. + +Cube labeling is coordinate-driven. FL2VA keyframes share the target video's +position grid, so a keyframe token and a target token at the same `(T, H, W)` +coordinate receive the same semantic cube label. Duplicate coordinates do not +extend the temporal grid; a semantic cube can therefore span multiple physical +attention blocks. In Ref2VA, standalone reference images remain dense, while +reference videos and the target video contribute to one global TopK candidate +pool rather than receiving separate per-stream quotas. + + +Cube sparse attention is an approximate backend and is not a consistency +ground-truth mode. `topk_ratio_list` length must equal the denoise step count +or the server rejects the request. Cube sparse attention does not support Ring +parallelism; use `--ulysses-degree` without `--ring-degree`. FlexAttention's +routing overhead can outweigh sparse-kernel savings on short sequences, so +benchmark latency as well as visual and audio quality on the target workload. + + + + The picker explicitly writes `--encoder-parallel auto` in every single-node diff --git a/docs/docs/sglang-diffusion/compatibility_matrix.mdx b/docs/docs/sglang-diffusion/compatibility_matrix.mdx index c31487257..ffdb56e1a 100644 --- a/docs/docs/sglang-diffusion/compatibility_matrix.mdx +++ b/docs/docs/sglang-diffusion/compatibility_matrix.mdx @@ -910,3 +910,9 @@ The entries below simply reflect configurations that have been manually validate ### Sliding Tile Attention - Currently, only Hopper GPUs (H100s) are supported. + +### Cube Sparse Attention + +- Available only for MiniMax-H3 (`--component-attention-backends transformer=cube_sparse_attn`). It sparsifies only the packed sequence's 3D visual streams; text, audio, standalone reference images, and the text-only token refiner remain dense. +- Requires `--attention-backend-config` with both `local_cube_size` and `topk_ratio_list`. `topk_ratio_list` must have one entry per denoise step, each in `(0, 1]`. See the [MiniMax-H3 cookbook](/cookbook/diffusion/MiniMax/MiniMax-H3) for a worked example. +- Runs on pure PyTorch plus FlexAttention, so it has no third-party kernel dependency. diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/utils.py b/python/sglang/multimodal_gen/runtime/entrypoints/utils.py index 97a7d3588..4b6da3235 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/utils.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/utils.py @@ -887,9 +887,13 @@ def prepare_request( """ Create a Req object with sampling_params as a parameter. """ + attention_backend_config = server_args.attention_backend_config or {} + vsa_sparsity = attention_backend_config.get( + "VSA_sparsity", attention_backend_config.get("sparsity", 0.0) + ) req = Req( sampling_params=sampling_params, - VSA_sparsity=server_args.attention_backend_config.VSA_sparsity, + VSA_sparsity=vsa_sparsity, ) sampling_params.apply_request_extra(req) if getattr(sampling_params, "max_sequence_length", None) is not None: diff --git a/python/sglang/multimodal_gen/runtime/layers/attention/backends/cube_sparse_attn/__init__.py b/python/sglang/multimodal_gen/runtime/layers/attention/backends/cube_sparse_attn/__init__.py new file mode 100644 index 000000000..7d2441373 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/layers/attention/backends/cube_sparse_attn/__init__.py @@ -0,0 +1,16 @@ +# SPDX-License-Identifier: Apache-2.0 +from sglang.multimodal_gen.runtime.layers.attention.backends.cube_sparse_attn.backend import ( + CubeSparseAttentionBackend, + CubeSparseAttentionImpl, + CubeSparseAttentionMetadata, + CubeSparseAttentionMetadataBuilder, + cube_sparse_attention, +) + +__all__ = [ + "CubeSparseAttentionBackend", + "CubeSparseAttentionImpl", + "CubeSparseAttentionMetadata", + "CubeSparseAttentionMetadataBuilder", + "cube_sparse_attention", +] diff --git a/python/sglang/multimodal_gen/runtime/layers/attention/backends/cube_sparse_attn/backend.py b/python/sglang/multimodal_gen/runtime/layers/attention/backends/cube_sparse_attn/backend.py new file mode 100644 index 000000000..26f2be21d --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/layers/attention/backends/cube_sparse_attn/backend.py @@ -0,0 +1,344 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Cube sparse attention backend (correctness-first FlexAttention kernel). + +The mask/metadata layer is kernel-agnostic (see ``mask.py``); the kernel call +is confined to ``_run_block_sparse_attention`` so faster block-sparse kernels +can be swapped in without touching mask semantics. A semantic cube label can +occupy multiple physical FlexAttention blocks when an embedded keyframe and +target frame share coordinates. +""" + +import functools +from dataclasses import dataclass +from typing import Any + +import torch +import torch.nn.functional as F +from torch.nn.attention.flex_attention import BlockMask, flex_attention + +from sglang.multimodal_gen.runtime.layers.attention.backends.attention_backend import ( + AttentionBackend, + AttentionImpl, + AttentionMetadata, + AttentionMetadataBuilder, +) +from sglang.multimodal_gen.runtime.layers.attention.backends.cube_sparse_attn.mask import ( + CubePrecomputed, + PackedStreams, + cube_topk_block_indices, + precompute_cube_attention, +) +from sglang.multimodal_gen.runtime.layers.attention.backends.flash_attn import ( + FlashAttentionImpl, +) +from sglang.multimodal_gen.runtime.managers.forward_context import get_forward_context +from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum + + +class CubeSparseAttentionBackend(AttentionBackend): + + @staticmethod + def get_enum() -> AttentionBackendEnum: + return AttentionBackendEnum.CUBE_SPARSE_ATTN + + @staticmethod + def get_impl_cls() -> type["CubeSparseAttentionImpl"]: + return CubeSparseAttentionImpl + + @staticmethod + def get_metadata_cls() -> type["CubeSparseAttentionMetadata"]: + return CubeSparseAttentionMetadata + + @staticmethod + def get_builder_cls() -> type["CubeSparseAttentionMetadataBuilder"]: + return CubeSparseAttentionMetadataBuilder + + +@dataclass +class CubeSparseAttentionMetadata(AttentionMetadata): + # Per-denoise-step top-k keep ratios; indexed by current_timestep. + topk_ratio_list: list[float] + # precompute_cube_attention(...) result for the positive packed layout. + precomputed: CubePrecomputed + + +class CubeSparseAttentionMetadataBuilder(AttentionMetadataBuilder): + + def __init__(self): + pass + + def prepare(self): + pass + + def build( # type: ignore[override] + self, + *, + packed: dict[str, Any], + local_cube_size: list[int] | tuple[int, ...], + topk_ratio_list: list[float], + num_steps: int, + device: torch.device, + **kwargs: dict[str, Any], + ) -> CubeSparseAttentionMetadata: + """Build cube metadata from a minimax_h3_packed_sequence(...) layout. + + The five per-stream token index tensors are derived from the packed + dict; segment shapes come from its ``stream_layout`` entry. + """ + topk_ratio_list = [float(ratio) for ratio in topk_ratio_list] + if len(topk_ratio_list) != num_steps: + raise ValueError( + f"topk_ratio_list has {len(topk_ratio_list)} entries for " + f"{num_steps} denoise steps" + ) + for ratio in topk_ratio_list: + if not 0.0 < ratio <= 1.0: + raise ValueError( + f"topk_ratio_list entries must be in (0, 1], got {ratio}" + ) + if "stream_layout" not in packed: + raise ValueError( + "packed layout has no stream_layout entry; cube sparse " + "attention requires the packed_sequence stream_layout export" + ) + layout = packed["stream_layout"] + + img_pos = packed["img_pos"].view(-1).to(torch.long) + update_mask = packed["update_mask"].view(-1).to(torch.bool) + audio_pos = packed["audio_pos"].view(-1).to(torch.long) + if "audio_update_mask" in packed: + audio_update_mask = packed["audio_update_mask"].view(-1).to(torch.bool) + else: + audio_update_mask = torch.ones(audio_pos.shape[0], dtype=torch.bool) + text_index = packed["text_pos"].view(-1).to(torch.long) + cond_image_index = img_pos[~update_mask] + latent_index = img_pos[update_mask] + cond_audio_index = audio_pos[~audio_update_mask] + audio_index = audio_pos[audio_update_mask] + + seq_len = int(packed["seq_len"]) + used = int(packed["cu_seqlens"].view(-1)[1]) + sparse_ratios = [ratio for ratio in topk_ratio_list if ratio < 1.0] + precomputed = precompute_cube_attention( + [tuple(layout["target_shape"])], + torch.tensor([0, used], dtype=torch.long), + seq_len, + tuple(local_cube_size), + device, + PackedStreams( + text=text_index, + cond_image=cond_image_index, + latent=latent_index, + cond_audio=cond_audio_index, + audio=audio_index, + cond_image_shapes=[tuple(layout["cond_image_shapes"])], + cond_image_roles=[tuple(layout["cond_image_roles"])], + cond_event_orders=[tuple(layout["cond_event_orders"])], + cond_audio_stream_lens=[tuple(layout["cond_audio_stream_lens"])], + ), + packed["img_position_ids"], + max(sparse_ratios, default=0.0), + ) + precomputed.runtime.pad_score_mod = _make_pad_score_mod( + precomputed.layout.is_real + ) + return CubeSparseAttentionMetadata( + current_timestep=0, + topk_ratio_list=topk_ratio_list, + precomputed=precomputed, + ) + + +@functools.cache +def _compiled_flex_attention(): + return torch.compile(flex_attention, mode="max-autotune-no-cudagraphs") + + +def _make_pad_score_mod(is_real: torch.Tensor): + is_real_bool = is_real.to(torch.bool) + + def _pad_score_mod(score, b, h, q_idx, kv_idx): + valid = is_real_bool[q_idx] & is_real_bool[kv_idx] + return torch.where(valid, score, float("-inf")) + + return _pad_score_mod + + +def _run_block_sparse_attention( + padded_q: torch.Tensor, + padded_k: torch.Tensor, + padded_v: torch.Tensor, + block_layout: dict[str, torch.Tensor | None], + precomputed: CubePrecomputed, + softmax_scale: float, +) -> torch.Tensor: + """Run block-sparse attention on the cube-padded layout. + + ``padded_q/k/v`` are ``[padded_seqlen, heads, dim]``. ``block_layout`` + contains physical KV rows produced directly by semantic TopK. + Per-head sparse buffers are reused. Sparse steps must omit the full-KV + tensors entirely rather than pass zero-count ones: on CUDA a present but + empty full-KV pair still steers FlexAttention into its slower mixed-layout + specialization, while omitting the pair selects the partial-only kernel. + Single kernel entry point; swap here for faster kernels. + """ + num_heads = padded_q.shape[1] + + def expand_heads(value): + if value.shape[1] == num_heads: + return value + if value.shape[1] != 1: + raise ValueError( + f"cube BlockMask has {value.shape[1]} heads for {num_heads} Q heads" + ) + return value.expand(1, num_heads, *value.shape[2:]) + + kv_num_blocks = block_layout["kv_num_blocks"] + kv_indices = block_layout["kv_indices"] + if kv_num_blocks is None or kv_indices is None: + raise ValueError("cube BlockMask requires compact KV block tensors") + kv_num_blocks = expand_heads(kv_num_blocks) + kv_indices = expand_heads(kv_indices) + full_kv_num_blocks = block_layout.get("full_kv_num_blocks") + full_kv_indices = block_layout.get("full_kv_indices") + if (full_kv_num_blocks is None) != (full_kv_indices is None): + raise ValueError( + "cube BlockMask full_kv_num_blocks and full_kv_indices must " + "both be present or both be omitted" + ) + if full_kv_num_blocks is not None: + full_kv_num_blocks = expand_heads(full_kv_num_blocks) + full_kv_indices = expand_heads(full_kv_indices) + + block_mask = BlockMask.from_kv_blocks( + kv_num_blocks, + kv_indices, + full_kv_num_blocks=full_kv_num_blocks, + full_kv_indices=full_kv_indices, + BLOCK_SIZE=precomputed.layout.cube_token_size, + seq_lengths=(padded_q.shape[0], padded_k.shape[0]), + compute_q_blocks=False, + ) + out = _compiled_flex_attention()( + padded_q.permute(1, 0, 2)[None], + padded_k.permute(1, 0, 2)[None], + padded_v.permute(1, 0, 2)[None], + score_mod=precomputed.runtime.pad_score_mod, + block_mask=block_mask, + scale=softmax_scale, + ) + return out[0].permute(1, 0, 2) + + +def cube_sparse_attention( + query: torch.Tensor, + key: torch.Tensor, + value: torch.Tensor, + attn_metadata: CubeSparseAttentionMetadata, + softmax_scale: float, +) -> torch.Tensor: + """Cube sparse attention over packed [total, heads, dim] q/k/v. + + Builds a per-step label top-k mask from pooled q/k, applies the + cube-contiguous reorder and padding, runs the block-sparse kernel, then + scatters the result back to packed-token order. + Rows past real_total_len (packing pad) get zero output. + """ + precomputed = attn_metadata.precomputed + topk_ratio = attn_metadata.topk_ratio_list[attn_metadata.current_timestep] + layout = precomputed.layout + real_total_len = layout.real_total_len + gather_idx = layout.gather_indices + + block_layout = cube_topk_block_indices( + query[:real_total_len], key[:real_total_len], precomputed, topk_ratio + ) + + padded_q = query[:real_total_len].index_select(0, gather_idx) + padded_k = key[:real_total_len].index_select(0, gather_idx) + padded_v = value[:real_total_len].index_select(0, gather_idx) + pad_indices = layout.pad_indices + if pad_indices.numel() > 0: + padded_q.index_fill_(0, pad_indices, 0) + padded_k.index_fill_(0, pad_indices, 0) + padded_v.index_fill_(0, pad_indices, 0) + + out = _run_block_sparse_attention( + padded_q, + padded_k, + padded_v, + block_layout, + precomputed, + softmax_scale, + ) + + output = out.index_select(0, layout.expand_indices) + if real_total_len < query.shape[0]: + output = F.pad(output, (0, 0, 0, 0, 0, query.shape[0] - real_total_len)) + return output + + +class CubeSparseAttentionImpl(AttentionImpl): + + def __init__( + self, + num_heads: int, + head_size: int, + softmax_scale: float, + causal: bool = False, + num_kv_heads: int | None = None, + prefix: str = "", + **extra_impl_args, + ) -> None: + if causal: + raise ValueError("cube sparse attention is non-causal only") + self.softmax_scale = softmax_scale + # Preserve H3's exact dense baseline on schedule entries that disable sparsity + self._dense_impl = FlashAttentionImpl( + num_heads=num_heads, + head_size=head_size, + causal=False, + softmax_scale=softmax_scale, + num_kv_heads=num_kv_heads, + prefix=prefix, + ) + + def forward( + self, + query: torch.Tensor, + key: torch.Tensor, + value: torch.Tensor, + attn_metadata: CubeSparseAttentionMetadata, + ) -> torch.Tensor: + return cube_sparse_attention( + query, key, value, attn_metadata, self.softmax_scale + ) + + def forward_varlen( + self, + query: torch.Tensor, + key: torch.Tensor, + value: torch.Tensor, + *, + cu_seqlens: torch.Tensor, + max_seqlen: int, + cu_seqlens_host: tuple[int, ...] | None = None, + ) -> torch.Tensor: + """Run the packed H3 layout carried by the active forward context.""" + + metadata = get_forward_context().attn_metadata + if not isinstance(metadata, CubeSparseAttentionMetadata): + raise ValueError( + "cube sparse attention requires CubeSparseAttentionMetadata " + "in the active forward context" + ) + if metadata.topk_ratio_list[metadata.current_timestep] == 1.0: + return self._dense_impl.forward_varlen( + query, + key, + value, + cu_seqlens=cu_seqlens, + max_seqlen=max_seqlen, + cu_seqlens_host=cu_seqlens_host, + ) + return cube_sparse_attention(query, key, value, metadata, self.softmax_scale) diff --git a/python/sglang/multimodal_gen/runtime/layers/attention/backends/cube_sparse_attn/mask.py b/python/sglang/multimodal_gen/runtime/layers/attention/backends/cube_sparse_attn/mask.py new file mode 100644 index 000000000..7633a22e4 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/layers/attention/backends/cube_sparse_attn/mask.py @@ -0,0 +1,960 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Cube sparse attention mask machinery. + +This module owns cube label layout, precomputation, condition-event +validation, and per-step top-k masks. Kernel-specific code is intentionally +kept in the backend; the padded layout gives every cube label one or more +``cube_token_size`` physical blocks for block-sparse kernels to consume +directly. +""" + +import math +from dataclasses import dataclass, field +from typing import Any, Callable + +import torch + + +@dataclass(frozen=True) +class PackedStreams: + """The per-modality token streams of one packed batch. + + These five index tensors plus the four per-sample condition metadata lists + always travel together and are only ever read as a set: they are the caller + describing *what is in the packed sequence*, as distinct from ``cube_size`` + and ``position_ids``, which describe how to grid it. Grouping them keeps + the modality contract in one place, so adding a stream is a change to this + class rather than to every signature along the path. + + ``text``, ``cond_image``, ``latent``, ``cond_audio`` and ``audio`` are flat + token-index tensors into the packed sequence. The four ``cond_*`` lists + carry one entry per sample and default to empty via + :func:`_normalize_cube_sample_metadata`. + """ + + text: torch.Tensor + cond_image: torch.Tensor + latent: torch.Tensor + cond_audio: torch.Tensor + audio: torch.Tensor + cond_image_shapes: list | None = None + cond_image_roles: list | None = None + cond_event_orders: list | None = None + cond_audio_stream_lens: list | None = None + + def as_index_tuple(self): + """The five streams in packed order, for uniform per-stream handling.""" + return (self.text, self.cond_image, self.latent, self.cond_audio, self.audio) + + +@dataclass(frozen=True) +class CubeLayout: + """Request-static cube layout, fully determined by ``precompute_cube_attention``. + + Every field here is an invariant of the packed sequence: it is computed once + per request and never mutated afterwards. Per-forward mutable state lives in + :class:`CubeRuntime` instead, so a reader can tell at a glance which + values exist before the first forward and which appear only during it. + + Three index spaces meet in this object, and conflating the first two is the + main bug risk. A tensor's *leading dimension* tells you which space it + lives in, so that is called out per field below: + + * **semantic label** (``L = num_labels``) -- one entry per occupied + ``(t, x, y)`` cube bin. TopK selection happens entirely in this space. + * **physical block** (``B = num_blocks``) -- one entry per FlexAttention + block of ``cube_token_size`` tokens. A semantic label owns + ``label_block_counts[label]`` physical blocks, which exceeds one whenever + a joint keyframe shares a target cube, so ``B >= L`` always. + * **token** -- one entry per packed or cube-padded token. + + The naming rule: ``label_*`` fields are *indexed by* semantic label, + ``block_labels`` maps the other way (physical block to owning label), and + ``*_physical_layout`` dicts are FlexAttention KV descriptors in block space. + ``num_labels`` and ``num_blocks`` are both plain counts, which is exactly + why they read alike -- check which space a value came from before using one + to size the other. + """ + + # ── semantic-label space (leading dim L, unless noted) ─────────────── + num_labels: int + """``L`` -- the number of semantic cubes. Not interchangeable with ``num_blocks``.""" + topk_mask: torch.Tensor + """``[L, L]`` bool: per-sample TopK candidate pool. Dense labels excluded.""" + base_block_mask: torch.Tensor + """``[L, L]`` bool: always-visible edges (self-diagonal + dense row/col).""" + sparse_label_mask: torch.Tensor + """``[L]`` bool: which labels take the sparse path at all.""" + label_lengths: torch.Tensor + """``[L]`` int: real token count per label, the ``segment_reduce`` lengths.""" + label_block_counts: torch.Tensor + """``[L]`` int: physical blocks owned by each label; ``sum() == num_blocks``.""" + label_block_indices: torch.Tensor + """``[L, max_label_block_count]`` int: label to its physical block ids, padded.""" + max_label_block_count: int + """Row width of ``label_block_indices``; ``1`` unless a label spans blocks.""" + topk_semantic_capacity: int + """Upper bound on selected labels per query label, for buffer sizing.""" + + # ── physical-block space (leading dim B) ───────────────────────────── + num_blocks: int + """``B`` -- the number of FlexAttention blocks. Not interchangeable with ``num_labels``.""" + block_labels: torch.Tensor + """``[B]`` int: physical block to owning semantic label (inverse of ``label_block_indices``).""" + base_physical_layout: dict[str, torch.Tensor] + """FlexAttention KV descriptor for ``base_block_mask``, in block space.""" + dense_physical_layout: dict[str, torch.Tensor] + """All-visible KV descriptor, returned directly when ``topk_ratio == 1.0``.""" + + # ── token space (packed and cube-padded) ──────────────────────────── + real_total_len: int + """Packed token count before cube padding.""" + padded_seqlen: int + """``num_blocks * cube_token_size`` -- token count after padding.""" + cube_token_size: int + """``prod(cube_size)`` -- tokens per physical block; the FlexAttention BLOCK_SIZE.""" + is_real: torch.Tensor + """``[padded_seqlen]`` int32 (0/1): real token vs. cube padding.""" + pad_indices: torch.Tensor + """Positions of the padding slots, for zeroing padded Q/K/V rows.""" + sorted_real_indices: torch.Tensor + """``[real_total_len]``: packed index of each token in cube-sorted order.""" + expand_indices: torch.Tensor + """``[real_total_len]``: packed position to its cube-padded position.""" + gather_indices: torch.Tensor + """``[padded_seqlen]``: cube-padded position to packed position (0 on pads).""" + + +@dataclass +class CubeRuntime: + """Per-forward mutable state, absent until the first ``forward`` call. + + ``pad_score_mod`` is installed by the backend once the layout is known; + the KV buffers are allocated lazily on the first sparse step and then + reused, so their shapes double as a guard against a changing head count. + """ + + pad_score_mod: Callable[..., Any] | None = None + kv_num_blocks_buffer: torch.Tensor | None = None + kv_indices_buffer: torch.Tensor | None = None + + +@dataclass +class CubePrecomputed: + """The cube metadata handed to the kernel: static layout + live buffers.""" + + layout: CubeLayout + runtime: CubeRuntime = field(default_factory=CubeRuntime) + + +def normalize_condition_event_order( + events, visual_count, audio_count, *, allow_audio_subset=False +): + normalized = [] + for event in events: + event_type, index = event[0], event[1] + event_hash = event[2] if len(event) > 2 else "" + normalized.append((event_type, index, event_hash)) + event_types = {event_type for event_type, _, _ in normalized} + unsupported = event_types - {"imgvid", "audio"} + if unsupported: + raise ValueError(f"unsupported condition event types: {sorted(unsupported)}") + imgvid_indices = [ + index for item_type, index, _ in normalized if item_type == "imgvid" + ] + if imgvid_indices != list(range(visual_count)): + raise ValueError( + f"condition imgvid indices {imgvid_indices} do not cover " + f"{visual_count} tensors" + ) + audio_indices = [ + index for item_type, index, _ in normalized if item_type == "audio" + ] + if allow_audio_subset: + out_of_range = [ + index for index in audio_indices if not 0 <= index < audio_count + ] + if out_of_range: + raise ValueError( + f"condition audio indices {out_of_range} out of range for " + f"{audio_count} tensors" + ) + if len(set(audio_indices)) != len(audio_indices): + raise ValueError( + f"condition audio indices {audio_indices} contain duplicates" + ) + elif audio_indices != list(range(audio_count)): + raise ValueError( + f"condition audio indices {audio_indices} do not cover " + f"{audio_count} tensors" + ) + return normalized + + +def _ceil_div(value, divisor): + return (value + divisor - 1) // divisor + + +def _cube_token_size(cube_size): + cube_size = tuple(int(value) for value in cube_size) + if len(cube_size) != 3 or any(value <= 0 for value in cube_size): + raise ValueError(f"local_cube_size must be a positive 3D size, got {cube_size}") + token_size = math.prod(cube_size) + if token_size & (token_size - 1): + raise ValueError( + "local_cube_size product must be a power of two for FlexAttention, " + f"got {cube_size} ({token_size} tokens)" + ) + return token_size + + +def _normalize_cube_visual_shape(shape): + shape = tuple(int(value) for value in shape) + if len(shape) != 3 or any(value <= 0 for value in shape): + raise ValueError( + f"cube attention visual shapes must be positive 3D shapes, got {shape}" + ) + return shape + + +_COND_IMAGE_ROLES = frozenset({"joint_cube", "independent_cube", "dense_prefix"}) + + +def _normalize_cond_image_roles(roles, visual_count, sample_idx): + roles = tuple(roles) + if len(roles) != visual_count: + raise ValueError( + f"cube attention sample {sample_idx} has {visual_count} condition " + f"visual streams but {len(roles)} condition roles" + ) + unsupported = sorted(set(roles) - _COND_IMAGE_ROLES) + if unsupported: + raise ValueError( + f"cube attention sample {sample_idx} has unsupported condition " + f"visual roles: {unsupported}" + ) + return roles + + +def _cube_sample_segments(indices, cu_seqlens): + bounds = torch.searchsorted(indices, cu_seqlens).tolist() + return [indices[start:end] for start, end in zip(bounds[:-1], bounds[1:])] + + +def _group_cube_visual_segment(indices, shape, cube_size, sample_idx, segment_name): + expected = math.prod(shape) + if indices.numel() != expected: + raise ValueError( + f"cube attention {segment_name} for sample {sample_idx} has " + f"{indices.numel()} tokens, expected {expected} for shape {shape}" + ) + cube_counts = tuple(_ceil_div(dim, extent) for dim, extent in zip(shape, cube_size)) + linear_indices = torch.arange( + expected, + dtype=torch.long, + device=indices.device, + ) + block_labels = torch.zeros_like(linear_indices) + for coordinate, extent, count in zip( + torch.unravel_index(linear_indices, shape), + cube_size, + cube_counts, + ): + block_labels = block_labels * count + coordinate // extent + ordered_labels, order = torch.sort(block_labels, stable=True) + return indices.index_select(0, order), ordered_labels, math.prod(cube_counts) + + +def _rank_position_axis(values): + """Map one floating position axis to stable, zero-based unique ranks.""" + order = torch.argsort(values, stable=True) + sorted_values = values.index_select(0, order) + changed = torch.empty_like(sorted_values, dtype=torch.bool) + changed[0] = True + changed[1:] = sorted_values[1:] != sorted_values[:-1] + sorted_ranks = changed.to(torch.long).cumsum(0) - 1 + ranks = torch.empty_like(sorted_ranks) + ranks[order] = sorted_ranks + return ranks + + +def _group_joint_cube_visual_segment( + indices, + position_ids, + cube_size, + sample_idx, +): + """Group target + embedded keyframes on their shared position grid. + + Floating RoPE coordinates are stably ranked per axis, so a keyframe with + the same temporal/spatial position as a target frame receives the same + semantic cube label. A semantic label may consequently contain more than + one physical block. + """ + if indices.numel() == 0: + raise ValueError( + f"cube attention joint visual stream for sample {sample_idx} is empty" + ) + positions = position_ids.index_select(0, indices) + if positions.ndim != 2 or positions.shape[1] != 3: + raise ValueError( + "cube attention img_position_ids must have shape [sequence, 3], " + f"got {tuple(position_ids.shape)}" + ) + ranked_axes = tuple(_rank_position_axis(positions[:, axis]) for axis in range(3)) + cube_counts = tuple( + int(axis.max().item()) // extent + 1 + for axis, extent in zip(ranked_axes, cube_size) + ) + block_labels = torch.zeros(indices.numel(), dtype=torch.long, device=indices.device) + for coordinate, extent, count in zip(ranked_axes, cube_size, cube_counts): + block_labels = block_labels * count + coordinate // extent + ordered_labels, order = torch.sort(block_labels, stable=True) + return indices.index_select(0, order), ordered_labels, math.prod(cube_counts) + + +def _group_cube_1d_segment(indices, cube_token_size): + local_labels = ( + torch.arange( + indices.numel(), + dtype=torch.long, + device=indices.device, + ) + // cube_token_size + ) + return indices, local_labels, _ceil_div(indices.numel(), cube_token_size) + + +def _normalize_cube_sample_metadata(values, num_samples, name): + if values is None: + return [()] * num_samples + if len(values) != num_samples: + raise ValueError( + f"cube attention received {len(values)} {name} entries for " + f"{num_samples} samples" + ) + return [() if value is None else value for value in values] + + +def _split_cube_streams(indices, stream_sizes, sample_idx, stream_name): + streams = [] + offset = 0 + for stream_size in stream_sizes: + stream_size = int(stream_size) + if stream_size < 0: + raise ValueError( + f"cube attention {stream_name} stream sizes must be nonnegative" + ) + streams.append(indices[offset : offset + stream_size]) + offset += stream_size + if offset != indices.numel(): + raise ValueError( + f"cube attention {stream_name} streams for sample {sample_idx} cover " + f"{offset} tokens, but packing contains {indices.numel()} tokens" + ) + return streams + + +def _pack_block_rows(block_mask): + """Pack a head-independent ``[Q, KV]`` mask without sorting.""" + counts = block_mask.sum(dim=-1, dtype=torch.int32) + capacity = int(counts.max().item()) if counts.numel() else 0 + indices = torch.zeros( + block_mask.shape[0], capacity, dtype=torch.int32, device=block_mask.device + ) + if capacity: + positions = block_mask.to(torch.int32).cumsum(dim=-1) - 1 + row_ids, column_ids = block_mask.nonzero(as_tuple=True) + indices[row_ids, positions[row_ids, column_ids].to(torch.long)] = column_ids.to( + torch.int32 + ) + return counts[None, None], indices[None, None] + + +def _build_physical_base_layouts( + base_block_mask, + dense_attention_mask, + block_labels, +): + def pack_as_full(semantic_mask): + physical = semantic_mask[block_labels[:, None], block_labels[None, :]] + full_counts, full_indices = _pack_block_rows(physical) + empty_counts = torch.zeros_like(full_counts) + # FlexAttention's CUDA kernel still expects a valid KV-index pointer + # when every non-full count is zero. Keep one unread placeholder slot + # instead of passing a zero-storage tensor. + empty_indices = torch.zeros( + *full_indices.shape[:-1], + 1, + dtype=torch.int32, + device=full_indices.device, + ) + return { + "full_kv_num_blocks": full_counts, + "full_kv_indices": full_indices, + "kv_num_blocks": empty_counts, + "kv_indices": empty_indices, + } + + return pack_as_full(base_block_mask), pack_as_full(dense_attention_mask) + + +def _raise_for_unoccupied_labels(occupied_labels, sample_label_ranges): + """Reject cube-label allocations that left grid cells without tokens. + + Joint grouping sizes each sample's label range from the full 3D grid, so + an unoccupied label means a ``joint_cube`` condition stream introduced + coordinates outside the target's densely tiled grid. That is an upstream + role-assignment error, not a supported layout: such a stream must be + declared ``independent_cube`` (or ``dense_prefix``) instead. + """ + occupied = set(occupied_labels.tolist()) + for sample_idx, (start, end) in enumerate(sample_label_ranges): + missing = [label for label in range(start, end) if label not in occupied] + if missing: + raise ValueError( + f"cube attention sample {sample_idx} allocated labels " + f"[{start}, {end}) but {len(missing)} of them received no " + f"tokens (first missing: {missing[0]}). A joint_cube " + "condition visual stream does not share the target's " + "position grid; declare it independent_cube or dense_prefix " + "instead." + ) + raise ValueError( + "cube attention allocated labels outside every sample range; " + f"occupied {len(occupied)} labels for ranges {sample_label_ranges}" + ) + + +def _build_cube_segment_layout( + sample_shapes, + cu_seqlens, + real_total_len, + cube_size, + device, + streams, + position_ids, +): + cube_token_size = math.prod(cube_size) + num_samples = len(sample_shapes) + if cu_seqlens.numel() != num_samples + 1: + raise ValueError( + f"cube attention received {cu_seqlens.numel() - 1} packed sequences " + f"for {num_samples} target shapes" + ) + cond_image_shapes = _normalize_cube_sample_metadata( + streams.cond_image_shapes, + num_samples, + "condition-shape", + ) + cond_image_roles = _normalize_cube_sample_metadata( + streams.cond_image_roles, + num_samples, + "condition-role", + ) + cond_event_orders = _normalize_cube_sample_metadata( + streams.cond_event_orders, + num_samples, + "condition-event", + ) + cond_audio_stream_lens = _normalize_cube_sample_metadata( + streams.cond_audio_stream_lens, + num_samples, + "condition-audio", + ) + ( + text_segments, + cond_image_segments, + latent_segments, + cond_audio_segments, + audio_segments, + ) = [ + _cube_sample_segments( + indices.to(device=device, dtype=torch.long), + cu_seqlens, + ) + for indices in streams.as_index_tuple() + ] + + cube_labels = torch.full( + (real_total_len,), + -1, + dtype=torch.int32, + device=device, + ) + label_offset = 0 + sample_label_ranges = [] + sparse_labels = [] + ordered_segments = [] + + def add_segment(ordered_indices, local_block_labels, num_blocks, *, sparse): + nonlocal label_offset + if ordered_indices.numel() == 0: + return + cube_labels[ordered_indices] = (local_block_labels + label_offset).to( + torch.int32 + ) + ordered_segments.append(ordered_indices) + sparse_labels.extend([sparse] * num_blocks) + label_offset += num_blocks + + for sample_idx in range(num_samples): + text = text_segments[sample_idx] + cond_image = cond_image_segments[sample_idx] + latent = latent_segments[sample_idx] + cond_audio = cond_audio_segments[sample_idx] + target_audio = audio_segments[sample_idx] + sample_label_start = label_offset + + cond_audio_streams = _split_cube_streams( + cond_audio, + cond_audio_stream_lens[sample_idx], + sample_idx, + "condition audio", + ) + visual_shapes = [ + _normalize_cube_visual_shape(shape) + for shape in cond_image_shapes[sample_idx] + ] + raw_visual_streams = _split_cube_streams( + cond_image, + [math.prod(shape) for shape in visual_shapes], + sample_idx, + "condition visual", + ) + visual_roles = _normalize_cond_image_roles( + cond_image_roles[sample_idx], len(visual_shapes), sample_idx + ) + + events = normalize_condition_event_order( + cond_event_orders[sample_idx], + visual_count=len(raw_visual_streams), + audio_count=len(cond_audio_streams), + allow_audio_subset=True, + ) + add_segment( + *_group_cube_1d_segment(text, cube_token_size), + sparse=False, + ) + listed_audio = set() + for event_type, event_idx, _ in events: + event_idx = int(event_idx) + if event_type == "audio": + listed_audio.add(event_idx) + add_segment( + *_group_cube_1d_segment( + cond_audio_streams[event_idx], + cube_token_size, + ), + sparse=False, + ) + else: + shape = visual_shapes[event_idx] + role = visual_roles[event_idx] + if role == "joint_cube": + continue + if role == "independent_cube" and (len(shape) != 3 or shape[0] <= 1): + raise ValueError( + "independent_cube condition visual streams must have " + f"a genuine 3D shape, got {shape}" + ) + if role == "independent_cube": + grouped = _group_cube_visual_segment( + raw_visual_streams[event_idx], + shape, + cube_size, + sample_idx, + f"condition visual stream {event_idx}", + ) + else: + grouped = _group_cube_1d_segment( + raw_visual_streams[event_idx], cube_token_size + ) + add_segment(*grouped, sparse=role == "independent_cube") + for stream_idx, stream in enumerate(cond_audio_streams): + if stream_idx not in listed_audio: + add_segment( + *_group_cube_1d_segment(stream, cube_token_size), + sparse=False, + ) + add_segment( + *_group_cube_1d_segment(target_audio, cube_token_size), + sparse=False, + ) + + target_shape = sample_shapes[sample_idx] + expected_target_tokens = math.prod(target_shape) + if latent.numel() != expected_target_tokens: + raise ValueError( + f"cube attention target visual for sample {sample_idx} has " + f"{latent.numel()} tokens, expected {expected_target_tokens} " + f"for shape {target_shape}" + ) + joint_streams = [ + stream + for stream, role in zip(raw_visual_streams, visual_roles) + if role == "joint_cube" + ] + joint_indices = torch.cat([*joint_streams, latent]) + add_segment( + *_group_joint_cube_visual_segment( + joint_indices, + position_ids, + cube_size, + sample_idx, + ), + sparse=True, + ) + sample_label_ranges.append((sample_label_start, label_offset)) + + sort_idx = torch.cat(ordered_segments) + if sort_idx.numel() != real_total_len: + unassigned = torch.nonzero(cube_labels < 0, as_tuple=False).flatten() + first_token = int(unassigned[0]) + first_sample = int(torch.searchsorted(cu_seqlens, first_token, right=True)) - 1 + raise ValueError( + f"cube attention modality segments cover {sort_idx.numel()} of " + f"{real_total_len} packed tokens; first unassigned token {first_token} " + f"in sample {first_sample}" + ) + + return ( + cube_labels, + sort_idx, + sample_label_ranges, + torch.tensor(sparse_labels, dtype=torch.bool, device=device), + ) + + +def precompute_cube_attention( + sample_shapes, + cu_seqlens, + total_len, + cube_size, + device, + streams, + position_ids, + max_sparse_topk_ratio, +): + """Build the request-static cube layout for one packed batch. + + ``streams`` is a :class:`PackedStreams` describing the per-modality token + streams; ``cube_size`` and ``position_ids`` describe the 3D grid they are + ranked on. + """ + sample_shapes = [_normalize_cube_visual_shape(shape) for shape in sample_shapes] + cu_seqlens = cu_seqlens.to(device=device, dtype=torch.long) + real_total_len = int(cu_seqlens[-1].item()) + if real_total_len > total_len: + raise ValueError( + f"cube attention real length {real_total_len} exceeds total length {total_len}" + ) + cube_token_size = _cube_token_size(cube_size) + cube_size = tuple(int(value) for value in cube_size) + position_ids = position_ids.to(device=device) + if position_ids.ndim != 2 or position_ids.shape != (total_len, 3): + raise ValueError( + "cube attention img_position_ids must have shape " + f"[{total_len}, 3], got {tuple(position_ids.shape)}" + ) + cube_labels, sort_idx, sample_label_ranges, sparse_label_mask = ( + _build_cube_segment_layout( + sample_shapes, + cu_seqlens, + real_total_len, + cube_size, + device, + streams, + position_ids, + ) + ) + num_labels = sum(end - start for start, end in sample_label_ranges) + dead_label = num_labels + + topk_mask = torch.zeros(num_labels, num_labels, dtype=torch.bool, device=device) + base_block_mask = torch.zeros( + num_labels, num_labels, dtype=torch.bool, device=device + ) + dense_attention_mask = torch.zeros( + num_labels, num_labels, dtype=torch.bool, device=device + ) + base_block_mask.fill_diagonal_(True) + for start, end in sample_label_ranges: + dense_attention_mask[start:end, start:end] = True + sample_sparse = sparse_label_mask[start:end] + topk_mask[start:end, start:end] = sample_sparse.unsqueeze( + 1 + ) & sample_sparse.unsqueeze(0) + sample_dense = ~sample_sparse + base_block_mask[start:end, start:end] |= sample_dense.unsqueeze( + 1 + ) | sample_dense.unsqueeze(0) + + sorted_labels = cube_labels[sort_idx] + occupied_labels, counts_per_label = sorted_labels.unique_consecutive( + return_counts=True + ) + if occupied_labels.numel() != num_labels: + _raise_for_unoccupied_labels(occupied_labels, sample_label_ranges) + padded_counts = ( + (counts_per_label + cube_token_size - 1) // cube_token_size + ) * cube_token_size + + padded_offsets = torch.zeros( + len(padded_counts) + 1, dtype=torch.long, device=device + ) + padded_offsets[1:] = padded_counts.cumsum(0) + group_starts = torch.zeros( + len(counts_per_label) + 1, dtype=torch.long, device=device + ) + group_starts[1:] = counts_per_label.cumsum(0) + + sorted_positions = torch.arange(real_total_len, device=device) + group_idx = torch.bucketize(sorted_positions, group_starts[1:], right=True) + padded_pos_sorted = padded_offsets[group_idx] + ( + sorted_positions - group_starts[group_idx] + ) + + expand_indices = torch.empty(real_total_len, dtype=torch.long, device=device) + expand_indices[sort_idx] = padded_pos_sorted + + padded_seqlen = int(padded_offsets[-1].item()) + label_block_counts = padded_counts // cube_token_size + label_block_offsets = torch.zeros(num_labels + 1, dtype=torch.long, device=device) + label_block_offsets[1:] = label_block_counts.cumsum(0) + num_blocks = int(label_block_offsets[-1].item()) + block_labels = torch.repeat_interleave( + torch.arange(num_labels, dtype=torch.long, device=device), + label_block_counts, + ) + max_label_block_count = int(label_block_counts.max().item()) + label_block_slots = torch.arange( + max_label_block_count, dtype=torch.long, device=device + ) + label_block_indices = label_block_offsets[:-1, None] + label_block_slots + label_block_indices = torch.where( + label_block_slots < label_block_counts[:, None], + label_block_indices, + torch.zeros_like(label_block_indices), + ) + gather_indices = torch.zeros(padded_seqlen, dtype=torch.long, device=device) + gather_indices[padded_pos_sorted] = sort_idx + + padded_cube_labels = torch.full( + (padded_seqlen,), dead_label, dtype=torch.int32, device=device + ) + padded_cube_labels[padded_pos_sorted] = sorted_labels + + is_real = (padded_cube_labels != dead_label).to(torch.int32) + pad_indices = torch.nonzero(is_real == 0, as_tuple=False).squeeze(1) + label_lengths = counts_per_label.to(torch.long) + base_physical_layout, dense_physical_layout = _build_physical_base_layouts( + base_block_mask, + dense_attention_mask, + block_labels, + ) + sparse_sizes = topk_mask.sum(dim=-1) + if max_sparse_topk_ratio > 0: + max_selected_counts = ( + sparse_sizes.to(torch.float32) * float(max_sparse_topk_ratio) + ).to(torch.long) + max_selected_counts.clamp_(min=1) + max_selected_counts = torch.minimum(max_selected_counts, sparse_sizes) + topk_semantic_capacity = int(max_selected_counts.max().item()) + else: + topk_semantic_capacity = 0 + + return CubePrecomputed( + layout=CubeLayout( + num_labels=num_labels, + topk_mask=topk_mask, + base_block_mask=base_block_mask, + sparse_label_mask=sparse_label_mask, + label_lengths=label_lengths, + label_block_counts=label_block_counts, + label_block_indices=label_block_indices, + max_label_block_count=max_label_block_count, + topk_semantic_capacity=topk_semantic_capacity, + num_blocks=num_blocks, + block_labels=block_labels, + base_physical_layout=base_physical_layout, + dense_physical_layout=dense_physical_layout, + real_total_len=real_total_len, + padded_seqlen=padded_seqlen, + cube_token_size=cube_token_size, + is_real=is_real, + pad_indices=pad_indices, + sorted_real_indices=sort_idx, + expand_indices=expand_indices, + gather_indices=gather_indices, + ) + ) + + +def _cube_topk_selection(q_real, k_real, precomputed, topk_ratio): + layout = precomputed.layout + dim = q_real.shape[-1] + + qk_sorted = torch.cat((q_real, k_real), dim=-1)[layout.sorted_real_indices] + label_lengths = layout.label_lengths + qk_pool = torch.segment_reduce( + qk_sorted, "sum", lengths=label_lengths, axis=0, unsafe=True + ) + qk_pool /= label_lengths.float().view(-1, 1, 1) + + q_pool, k_pool = torch.split(qk_pool, dim, dim=-1) + scores = torch.einsum("lhd,mhd->hlm", q_pool, k_pool) * (dim**-0.5) + + candidate_mask = layout.topk_mask + sparse_sizes = candidate_mask.sum(dim=-1) + sparse_labels = sparse_sizes > 0 + scores.masked_fill_(~candidate_mask.unsqueeze(0), float("-inf")) + + selected_counts = (sparse_sizes.to(torch.float32) * topk_ratio).to(torch.long) + selected_counts.clamp_(min=1) + selected_counts = torch.minimum(selected_counts, sparse_sizes) + selected_counts = torch.where( + sparse_labels, selected_counts, torch.zeros_like(selected_counts) + ) + + selected_order = torch.argsort( + scores, + dim=-1, + descending=True, + stable=True, + ) + return selected_order, selected_counts + + +def cube_topk_block_indices(q_real, k_real, precomputed, topk_ratio): + """Build physical FlexAttention KV rows directly from semantic TopK. + + Semantic TopK expands directly to physical block ids without materializing + a per-head BxB boolean mask. Static base ids and selected sparse ids are + merged into one ordered per-head KV prefix; padding is handled by the + backend score modifier. Sparse steps intentionally omit full-KV metadata + entirely so FlexAttention selects its partial-only kernel. + + This is the only production expansion of :func:`_cube_topk_selection`. The + test suite expands the same selection into a semantic ``[H, L, L]`` mask + and asserts the two agree at bool level, because a divergence confined to a + few blocks can hide under an attention numeric tolerance. + """ + layout = precomputed.layout + runtime = precomputed.runtime + if topk_ratio == 1.0: + return layout.dense_physical_layout + + selected_order, selected_counts = _cube_topk_selection( + q_real, k_real, precomputed, topk_ratio + ) + semantic_capacity = layout.topk_semantic_capacity + if semantic_capacity <= 0: + raise ValueError( + "cube sparse metadata has no compact TopK capacity for a sparse ratio" + ) + selected_semantic = selected_order[..., :semantic_capacity] + semantic_rank = torch.arange(semantic_capacity, device=selected_order.device).view( + 1, 1, -1 + ) + selected_valid = semantic_rank < selected_counts.view(1, -1, 1) + + block_labels = layout.block_labels + selected_semantic = selected_semantic.index_select(1, block_labels) + selected_valid = selected_valid.index_select(1, block_labels).expand_as( + selected_semantic + ) + q_semantic = block_labels.view(1, -1, 1) + # Deduplicate against the static base set. This is not just an + # optimization: the truncation to num_blocks further down is only lossless + # because base ids and selected ids are disjoint (this line) and + # label->physical-blocks is a partition, so each row holds at most + # num_blocks distinct valid ids. Removing this intersection would make + # the sort below silently drop real KV blocks whenever + # base_capacity + selected blocks exceeds num_blocks. + selected_valid = ( + selected_valid & ~layout.base_block_mask[q_semantic, selected_semantic] + ) + + label_block_indices = layout.label_block_indices + label_block_counts = layout.label_block_counts + max_label_blocks = layout.max_label_block_count + selected_physical = label_block_indices[selected_semantic] + physical_rank = torch.arange(max_label_blocks, device=selected_order.device).view( + 1, 1, 1, -1 + ) + selected_physical_valid = selected_valid.unsqueeze(-1) & ( + physical_rank < label_block_counts[selected_semantic].unsqueeze(-1) + ) + selected_physical = selected_physical.flatten(-2) + selected_physical_valid = selected_physical_valid.flatten(-2) + + base_layout = layout.base_physical_layout + num_heads = q_real.shape[1] + num_blocks = layout.num_blocks + base_counts = base_layout["full_kv_num_blocks"].expand(1, num_heads, -1)[0] + base_indices = base_layout["full_kv_indices"].expand(1, num_heads, -1, -1)[0] + base_rank = torch.arange(base_indices.shape[-1], device=q_real.device).view( + 1, 1, -1 + ) + base_valid = base_rank < base_counts.unsqueeze(-1) + + # Block ids are bounded by num_blocks (< 2**31), so int32 is wide enough and + # halves the transient footprint of the sort below, which is the largest + # per-step allocation on the sparse path. + candidate_indices = torch.cat( + (base_indices.to(torch.int32), selected_physical.to(torch.int32)), dim=-1 + ) + candidate_valid = torch.cat((base_valid, selected_physical_valid), dim=-1) + packed_kv = torch.where( + candidate_valid, + candidate_indices, + torch.full_like(candidate_indices, num_blocks), + ) + # Sorting sends the num_blocks sentinel to the tail, so the first + # num_blocks slots hold every valid id. The candidate width + # (base_capacity + topk_semantic_capacity * max_label_block_count) may + # exceed num_blocks, but the valid count per row cannot: base and selected + # ids are disjoint and label->blocks is a partition. That invariant is + # what makes this truncation lossless rather than a silent block drop; it + # is asserted in test_compact_kv_rows_never_exceed_block_count. + packed_kv = torch.sort(packed_kv, dim=-1, stable=True).values[..., :num_blocks] + packed_kv = torch.where( + packed_kv < num_blocks, + packed_kv, + torch.zeros_like(packed_kv), + ) + + buffer = runtime.kv_indices_buffer + if buffer is None: + runtime.kv_num_blocks_buffer = torch.empty( + 1, num_heads, num_blocks, dtype=torch.int32, device=q_real.device + ) + runtime.kv_indices_buffer = torch.empty( + 1, + num_heads, + num_blocks, + num_blocks, + dtype=torch.int32, + device=q_real.device, + ) + else: + expected = (1, num_heads, num_blocks, num_blocks) + if tuple(buffer.shape) != expected: + raise ValueError( + "cube attention KV buffer shape changed across calls: " + f"{tuple(buffer.shape)} vs {expected}" + ) + + kv_counts = runtime.kv_num_blocks_buffer + kv_indices = runtime.kv_indices_buffer + selected_physical_counts = selected_physical_valid.sum(dim=-1, dtype=torch.int32) + kv_counts[0].copy_(base_counts + selected_physical_counts) + kv_indices[0].copy_(packed_kv.to(torch.int32)) + + return { + "kv_num_blocks": kv_counts, + "kv_indices": kv_indices, + "full_kv_num_blocks": None, + "full_kv_indices": None, + } diff --git a/python/sglang/multimodal_gen/runtime/models/dits/minimax_h3.py b/python/sglang/multimodal_gen/runtime/models/dits/minimax_h3.py index b6459693d..fbc8b6a73 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/minimax_h3.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/minimax_h3.py @@ -612,6 +612,7 @@ def _minimax_h3_attention_core_impl( get_attn_backend( attention.head_dim, q.dtype, + selected_attention_backend=attention._selected_attention_backend, attention_requirements=AttentionRequirements(packed_varlen=True), ) ) @@ -691,6 +692,7 @@ class MiniMaxH3Attention(nn.Module): *, prefix: str, bcg_breakpoint: bool = True, + cube_sparse_capable: bool = True, ) -> None: super().__init__() self.bcg_breakpoint = bcg_breakpoint @@ -709,6 +711,13 @@ class MiniMaxH3Attention(nn.Module): self.prefix = prefix self._attention_impl = None self._attention_backend_enum: AttentionBackendEnum | None = None + # attention initializes on the first real QKV tensors, after the + # component-loading context has ended; retain the transformer-scoped + # selection so a component override is not silently lost at runtime + self._selected_attention_backend = get_component_forced_attn_backend() + # Cube metadata describes only the packed multimodal sequence. The + # text-only token refiner must preserve the exact dense FA baseline. + self._cube_sparse_capable = cube_sparse_capable # The checkpoint stores one fused qkv tensor. Each logical Q/K/V # matrix must be sharded independently; a plain ColumnParallelLinear # would instead slice across the concatenated tensor and is incorrect @@ -759,6 +768,15 @@ class MiniMaxH3Attention(nn.Module): ) def _set_attention_backend(self, backend) -> None: + if ( + backend.get_enum() is AttentionBackendEnum.CUBE_SPARSE_ATTN + and not self._cube_sparse_capable + ): + backend = get_attn_backend( + self.head_dim, + _BF16_DTYPE, + selected_attention_backend=AttentionBackendEnum.FA, + ) impl_cls = backend.get_impl_cls() self._attention_impl = impl_cls( num_heads=self.num_heads, @@ -1466,6 +1484,7 @@ class MiniMaxH3TokenRefinerBlock(nn.Module): quant_config, prefix=f"{prefix}.attn", bcg_breakpoint=False, + cube_sparse_capable=False, ) self.mlp = MiniMaxH3MLP(arch, quant_config, prefix=f"{prefix}.mlp") diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/minimax_h3/denoise_loop.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/minimax_h3/denoise_loop.py index 0ca97cb20..b04643e61 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/minimax_h3/denoise_loop.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/minimax_h3/denoise_loop.py @@ -20,6 +20,12 @@ from sglang.multimodal_gen.runtime.distributed.parallel_state import ( get_ring_ctx, get_ulysses_ctx, ) +from sglang.multimodal_gen.runtime.layers.attention.backends.attention_backend import ( + AttentionMetadata, +) +from sglang.multimodal_gen.runtime.managers.forward_context import ( + set_forward_context, +) MINIMAX_H3_IMGVID_COND_TIMESTEP = 0.999 # ref2va audio reference anchor timestep @@ -447,6 +453,7 @@ def minimax_h3_denoise_loop( device: torch.device, imgvid_cond_noise_aug_for_inference: float = MINIMAX_H3_IMGVID_COND_TIMESTEP, audio_cond_noise_aug_for_inference: float = MINIMAX_H3_AUDIO_REF_COND_TIMESTEP, + attn_metadata: AttentionMetadata | None = None, on_step: Callable[[int, torch.Tensor, torch.Tensor], None] | None = None, step_profiler: Callable[[int], AbstractContextManager] | None = None, ) -> tuple[torch.Tensor, torch.Tensor]: @@ -555,7 +562,16 @@ def minimax_h3_denoise_loop( audio_rows=audio_rows, step_timesteps=timestep_plan[step], ) - with torch.inference_mode(): + if attn_metadata is not None: + attn_metadata.current_timestep = step + if model_forward is None and attn_metadata is not None: + forward_cm: AbstractContextManager = set_forward_context( + current_timestep=step, + attn_metadata=attn_metadata, + ) + else: + forward_cm = nullcontext() + with forward_cm, torch.inference_mode(): if model_forward is None: v_video, v_audio = model(**fk) else: diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/minimax_h3/packed_sequence.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/minimax_h3/packed_sequence.py index cd7fd66b5..4bd41d4e4 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/minimax_h3/packed_sequence.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/minimax_h3/packed_sequence.py @@ -208,6 +208,19 @@ def minimax_h3_packed_sequence( token_tags[img_pos] = 0 # VISUAL (condition images + target video) cu = torch.tensor([0, used, seq_len], dtype=torch.int32) + # Cube-sparse-attention segment shapes; every other field is per-token. + stream_layout = { + "target_shape": (latent_t, ph, pw), + "cond_image_shapes": tuple((1, ph, pw) for _ in resolved_cond_frame_indices), + # FL2VA keyframes live on the target timeline. The attention + # metadata builder must fold them into the target's 3D cube grid + # instead of inferring a dense image role from their T=1 shape. + "cond_image_roles": tuple("joint_cube" for _ in resolved_cond_frame_indices), + "cond_event_orders": tuple( + ("imgvid", index) for index in range(len(resolved_cond_frame_indices)) + ), + "cond_audio_stream_lens": (), + } packed = { "seq_len": seq_len, "img_pos": img_pos, @@ -217,6 +230,7 @@ def minimax_h3_packed_sequence( "img_position_ids": g, "token_tags": token_tags, "cu_seqlens": cu, + "stream_layout": stream_layout, } if include_video_pos: # Conditioning keyframes are images. Only generated video rows are @@ -526,6 +540,47 @@ def minimax_h3_packed_sequence_ref2va_blocks( token_tags[img_pos] = 0 # VISUAL (reference images/videos + target video) cu = torch.tensor([0, used, seq_len], dtype=torch.int32) + # Cube-sparse-attention segment shapes; streams listed in ref-block order, + # matching the img_pos/audio_pos concatenation above (audio rows precede + # video rows within a video-bearing block). + cond_image_shapes: list[tuple[int, int, int]] = [] + cond_image_roles: list[str] = [] + cond_event_orders: list[tuple[str, int]] = [] + cond_audio_stream_lens: list[int] = [] + for item in parsed: + kind = str(item["kind"]) + if kind == "image": + cond_event_orders.append(("imgvid", len(cond_image_shapes))) + cond_image_shapes.append( + ( + 1, + int(item["latent_h"]) // _PATCH_H, + int(item["latent_w"]) // _PATCH_W, + ) + ) + cond_image_roles.append("dense_prefix") + elif kind == "audio": + cond_event_orders.append(("audio", len(cond_audio_stream_lens))) + cond_audio_stream_lens.append(int(item["audio_rows"])) + else: + cond_event_orders.append(("audio", len(cond_audio_stream_lens))) + cond_audio_stream_lens.append(int(item["audio_rows"])) + cond_event_orders.append(("imgvid", len(cond_image_shapes))) + cond_image_shapes.append( + ( + int(item["latent_t"]), + int(item["latent_h"]) // _PATCH_H, + int(item["latent_w"]) // _PATCH_W, + ) + ) + cond_image_roles.append("independent_cube") + stream_layout = { + "target_shape": (latent_t, ph, pw), + "cond_image_shapes": tuple(cond_image_shapes), + "cond_image_roles": tuple(cond_image_roles), + "cond_event_orders": tuple(cond_event_orders), + "cond_audio_stream_lens": tuple(cond_audio_stream_lens), + } packed = { "seq_len": seq_len, "img_pos": img_pos, @@ -536,6 +591,7 @@ def minimax_h3_packed_sequence_ref2va_blocks( "img_position_ids": g, "token_tags": token_tags, "cu_seqlens": cu, + "stream_layout": stream_layout, } if ref_video_pos_parts is not None: # Reference image blocks remain dense; reference videos and the diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/minimax_h3/stages/denoising.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/minimax_h3/stages/denoising.py index 26b817b8b..99f61de19 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/minimax_h3/stages/denoising.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/minimax_h3/stages/denoising.py @@ -16,6 +16,10 @@ from sglang.multimodal_gen.runtime.cache.cache_dit_integration import ( CacheDitConfig, disable_cache_on_transformer, ) +from sglang.multimodal_gen.runtime.layers.attention.backends.cube_sparse_attn import ( + CubeSparseAttentionMetadata, + CubeSparseAttentionMetadataBuilder, +) from sglang.multimodal_gen.runtime.managers.memory_managers.component_residency_strategies import ( is_fsdp_managed_module, ) @@ -335,6 +339,46 @@ def _resolve_denoise_model( return model.to(device).eval() +def _build_cube_attn_metadata( + server_args: ServerArgs, + *, + packed: dict[str, Any], + num_steps: int, + device: torch.device, +) -> CubeSparseAttentionMetadata | None: + """Build cube sparse attention metadata when that backend is selected.""" + transformer_backend = (server_args.component_attention_backends or {}).get( + "transformer", server_args.attention_backend + ) + if str(transformer_backend).lower() != "cube_sparse_attn": + return None + + config = server_args.attention_backend_config or {} + local_cube_size = config.get("local_cube_size") + topk_ratio_list = config.get("topk_ratio_list") + if not local_cube_size or not topk_ratio_list: + raise ValueError( + "cube_sparse_attn requires --attention-backend-config with " + "local_cube_size and topk_ratio_list" + ) + metadata = CubeSparseAttentionMetadataBuilder().build( + packed=packed, + local_cube_size=local_cube_size, + topk_ratio_list=topk_ratio_list, + num_steps=num_steps, + device=device, + ) + logger.info( + "cube sparse attention enabled: local_cube_size=%s " + "topk_ratio_list(len=%d, min=%.4f, max=%.4f)", + list(local_cube_size), + len(metadata.topk_ratio_list), + min(metadata.topk_ratio_list), + max(metadata.topk_ratio_list), + ) + return metadata + + def _precompute_refined_prompt_embeds( model: Any, positive: Any, @@ -671,6 +715,12 @@ class MiniMaxH3DenoisingStage(DenoisingStage): imgvid_noise_aug=imgvid_noise_aug, audio_noise_aug=audio_noise_aug, ) + attn_metadata = _build_cube_attn_metadata( + server_args, + packed=packed, + num_steps=len(sigmas_video) - 1, + device=device, + ) placement_managed = self._component_residency_manager is not None if placement_managed: @@ -716,7 +766,11 @@ class MiniMaxH3DenoisingStage(DenoisingStage): video_rows, audio_rows = minimax_h3_denoise_loop( model=model, - model_forward=partial(self._forward_dit, batch=batch), + model_forward=partial( + self._forward_dit, + batch=batch, + attn_metadata=attn_metadata, + ), positive=positive, initial_video_rows=initial_video, initial_audio_rows=initial_audio, @@ -727,6 +781,7 @@ class MiniMaxH3DenoisingStage(DenoisingStage): device=device, imgvid_cond_noise_aug_for_inference=float(imgvid_noise_aug), audio_cond_noise_aug_for_inference=float(audio_noise_aug), + attn_metadata=attn_metadata, on_step=on_step, step_profiler=partial( self._profile_denoising_step, @@ -767,6 +822,7 @@ class MiniMaxH3DenoisingStage(DenoisingStage): step_index: int, *, batch: Req, + attn_metadata: CubeSparseAttentionMetadata | None = None, ) -> tuple[torch.Tensor, torch.Tensor]: """Route the custom full loop through the native denoising runner.""" @@ -776,7 +832,7 @@ class MiniMaxH3DenoisingStage(DenoisingStage): with set_forward_context( current_timestep=step_index, - attn_metadata=None, + attn_metadata=attn_metadata, forward_batch=batch, ): runner = self._maybe_get_bcg_runner(model) diff --git a/python/sglang/multimodal_gen/runtime/platforms/cuda.py b/python/sglang/multimodal_gen/runtime/platforms/cuda.py index 46ad5b686..2f6b95df5 100644 --- a/python/sglang/multimodal_gen/runtime/platforms/cuda.py +++ b/python/sglang/multimodal_gen/runtime/platforms/cuda.py @@ -252,6 +252,22 @@ class _VideoSparseAttentionBackendResolver(_CudaAttentionBackendResolver): raise ImportError("Video Sparse Attention backend is not installed.") from e +class _CubeSparseAttentionBackendResolver(_CudaAttentionBackendResolver): + backend = AttentionBackendEnum.CUBE_SPARSE_ATTN + + @classmethod + def resolve(cls, platform) -> str: + # MiniMax H3's text-only token refiner deliberately stays on the exact + # FA baseline when the packed multimodal blocks use cube attention. + # Initialize the Blackwell FA generation on the cube selection path as + # well, otherwise the refiner can fall into an unavailable FA2 package. + if not platform._prepare_flash_attention_for_blackwell(): + raise RuntimeError( + "cube sparse attention requires FlashAttention for H3's dense paths" + ) + return "sglang.multimodal_gen.runtime.layers.attention.backends.cube_sparse_attn.CubeSparseAttentionBackend" + + class _SparseVideoGen2AttentionBackendResolver(_CudaAttentionBackendResolver): backend = AttentionBackendEnum.SPARSE_VIDEO_GEN_2_ATTN @@ -419,6 +435,7 @@ _CUDA_ATTENTION_BACKEND_RESOLVERS = { _SageAttention3BackendResolver, _SpargeAttentionBackendResolver, _VideoSparseAttentionBackendResolver, + _CubeSparseAttentionBackendResolver, _SparseVideoGen2AttentionBackendResolver, _SolAttnBackendResolver, _VMOBAAttentionBackendResolver, diff --git a/python/sglang/multimodal_gen/runtime/platforms/interface.py b/python/sglang/multimodal_gen/runtime/platforms/interface.py index a7bd1c8a8..17627e3f8 100644 --- a/python/sglang/multimodal_gen/runtime/platforms/interface.py +++ b/python/sglang/multimodal_gen/runtime/platforms/interface.py @@ -46,6 +46,7 @@ class AttentionBackendEnum(enum.Enum): RAIN_FUSION_ATTN = enum.auto() SOL_ATTN = enum.auto() SUBBLOCK_SPARSE_ATTN = enum.auto() + CUBE_SPARSE_ATTN = enum.auto() NO_ATTENTION = enum.auto() def __str__(self): @@ -66,6 +67,7 @@ class AttentionBackendEnum(enum.Enum): AttentionBackendEnum.RAIN_FUSION_ATTN, AttentionBackendEnum.SOL_ATTN, AttentionBackendEnum.SUBBLOCK_SPARSE_ATTN, + AttentionBackendEnum.CUBE_SPARSE_ATTN, } diff --git a/python/sglang/multimodal_gen/test/unit/test_cube_sparse_attn_backend.py b/python/sglang/multimodal_gen/test/unit/test_cube_sparse_attn_backend.py new file mode 100644 index 000000000..4555212ed --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_cube_sparse_attn_backend.py @@ -0,0 +1,722 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Hermetic tests for the cube sparse attention backend. + +Mask semantics are locked against a token-level masked-SDPA oracle and the +reference policy: only 3D visual streams participate in global top-k. +""" + +import unittest +from types import SimpleNamespace +from unittest import mock + +import torch + +from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams +from sglang.multimodal_gen.runtime.entrypoints.utils import prepare_request +from sglang.multimodal_gen.runtime.layers.attention.backends.cube_sparse_attn import ( + backend as cube_backend, +) +from sglang.multimodal_gen.runtime.layers.attention.backends.cube_sparse_attn.backend import ( + CubeSparseAttentionImpl, + CubeSparseAttentionMetadataBuilder, + cube_sparse_attention, +) +from sglang.multimodal_gen.runtime.layers.attention.backends.cube_sparse_attn.mask import ( + _cube_topk_selection, + cube_topk_block_indices, +) +from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_context +from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.minimax_h3.packed_sequence import ( + minimax_h3_packed_sequence, + minimax_h3_packed_sequence_ref2va_blocks, +) + +_CUBE_SIZE = (2, 2, 2) +_NUM_HEADS = 3 +_HEAD_DIM = 8 + + +def _build_t2va_packed(*, include_keyframe_cond=False): + return minimax_h3_packed_sequence( + # Span multiple cube labels so an accidental sparse text policy is + # observable even when top-k clamps each non-empty group to one label. + text_len=17, + latent_t=2, + latent_h=8, + latent_w=12, + audio_t=5, + include_keyframe_cond=include_keyframe_cond, + keyframe_frame_indices=[0] if include_keyframe_cond else None, + frame_count=5 if include_keyframe_cond else None, + ) + + +def _build_metadata(packed, topk_ratio_list): + return CubeSparseAttentionMetadataBuilder().build( + packed=packed, + local_cube_size=_CUBE_SIZE, + topk_ratio_list=topk_ratio_list, + num_steps=len(topk_ratio_list), + device=torch.device("cpu"), + ) + + +def _random_qkv(seq_len): + generator = torch.Generator().manual_seed(0) + return tuple( + torch.randn( + seq_len, _NUM_HEADS, _HEAD_DIM, dtype=torch.float32, generator=generator + ) + for _ in range(3) + ) + + +def _token_labels(precomputed): + lengths = precomputed.layout.label_lengths + device = lengths.device + labels = torch.empty( + precomputed.layout.real_total_len, dtype=torch.long, device=device + ) + labels[precomputed.layout.sorted_real_indices] = torch.repeat_interleave( + torch.arange(precomputed.layout.num_labels, device=device), lengths + ) + return labels + + +def _reference_policy_block_mask(q_real, k_real, precomputed, topk_ratio): + """Hermetic label-mask oracle for the reference sparse policy.""" + dim = q_real.shape[-1] + qk_sorted = torch.cat((q_real, k_real), dim=-1)[ + precomputed.layout.sorted_real_indices + ] + label_lengths = precomputed.layout.label_lengths + qk_pool = torch.segment_reduce( + qk_sorted, "sum", lengths=label_lengths, axis=0, unsafe=True + ) + qk_pool /= label_lengths.float().view(-1, 1, 1) + + q_pool, k_pool = torch.split(qk_pool, dim, dim=-1) + scores = torch.einsum("lhd,mhd->hlm", q_pool, k_pool) * (dim**-0.5) + candidate_mask = precomputed.layout.topk_mask + sparse_sizes = candidate_mask.sum(dim=-1) + sparse_labels = sparse_sizes > 0 + scores.masked_fill_(~candidate_mask.unsqueeze(0), float("-inf")) + + top_k_per_label = (sparse_sizes.to(torch.float32) * topk_ratio).to(torch.long) + top_k_per_label.clamp_(min=1) + top_k_per_label = torch.minimum(top_k_per_label, sparse_sizes) + top_k_per_label = torch.where( + sparse_labels, top_k_per_label, torch.zeros_like(top_k_per_label) + ) + + sorted_idx = torch.argsort(scores, dim=-1, descending=True, stable=True) + selected_rank = torch.arange(scores.shape[-1], device=scores.device).view( + 1, 1, -1 + ) < top_k_per_label.view(1, -1, 1) + block_mask = torch.zeros_like(scores, dtype=torch.bool) + block_mask.scatter_(-1, sorted_idx, selected_rank.expand_as(sorted_idx)) + block_mask &= candidate_mask.unsqueeze(0) + block_mask |= precomputed.layout.base_block_mask.unsqueeze(0) + return block_mask + + +def _assert_tokens_are_dense(test_case, allowed, token_indices): + token_indices = token_indices.to(torch.long) + test_case.assertGreater(token_indices.numel(), 0) + test_case.assertTrue(allowed.index_select(1, token_indices).all()) + test_case.assertTrue(allowed.index_select(2, token_indices).all()) + + +def _naive_masked_attention(q, k, v, allowed, scale): + scores = torch.einsum("ihd,jhd->hij", q, k) * scale + scores = scores.masked_fill(~allowed, float("-inf")) + probs = torch.softmax(scores, dim=-1) + return torch.einsum("hij,jhd->ihd", probs, v) + + +def cube_topk_block_mask(q_real, k_real, precomputed, topk_ratio): + """Semantic ``[H, num_labels, num_labels]`` TopK mask -- test oracle. + + Production runs ``cube_topk_block_indices``, which expands + ``_cube_topk_selection`` straight into physical KV rows. This function + expands the same selection into a semantic boolean mask instead: it is the + readable form the tests assert against, and it lives here rather than in + the production module because nothing outside the test suite reads it. + + ``test_compact_indices_match_semantic_mask`` locks the two expansions + together, so this oracle drifting away from production is itself a test + failure rather than a silent numeric shift. + """ + layout = precomputed.layout + selected_order, selected_counts = _cube_topk_selection( + q_real, k_real, precomputed, topk_ratio + ) + candidate_mask = layout.topk_mask + selected_rank = torch.arange( + selected_order.shape[-1], device=selected_order.device + ).view(1, 1, -1) < selected_counts.view(1, -1, 1) + block_mask = torch.zeros_like(selected_order, dtype=torch.bool) + block_mask.scatter_(-1, selected_order, selected_rank.expand_as(selected_order)) + block_mask &= candidate_mask.unsqueeze(0) + block_mask |= layout.base_block_mask.unsqueeze(0) + return block_mask + + +def _physical_mask_from_block_layout(block_layout, layout, num_heads): + """Expand compact KV rows back into a dense ``[H, B, B]`` physical mask. + + This is the inverse of what ``cube_topk_block_indices`` packs, so it lets a + test compare the production compact path against the semantic mask that + the ``cube_topk_block_mask`` oracle above returns. + """ + num_blocks = layout.num_blocks + physical_allowed = torch.zeros(num_heads, num_blocks, num_blocks, dtype=torch.bool) + for count_key, index_key in ( + ("kv_num_blocks", "kv_indices"), + ("full_kv_num_blocks", "full_kv_indices"), + ): + if block_layout[count_key] is None: + continue + counts = block_layout[count_key].expand(1, num_heads, -1)[0] + indices = block_layout[index_key].expand(1, num_heads, -1, -1)[0] + ranks = torch.arange(indices.shape[-1]).view(1, 1, -1) + valid = ranks < counts.unsqueeze(-1) + head_ids, query_ids, slot_ids = valid.nonzero(as_tuple=True) + physical_allowed[ + head_ids, query_ids, indices[head_ids, query_ids, slot_ids].long() + ] = True + return physical_allowed + + +def _reference_block_sparse_attention( + padded_q, padded_k, padded_v, block_layout, precomputed, softmax_scale +): + """Dense reference honoring the _run_block_sparse_attention contract. + + Uncompiled flex_attention ignores BlockMask block structure (mask_mod + only), so CPU tests substitute this oracle; the real compiled-flex kernel + is exercised by the CUDA test class. + """ + cube_token_size = precomputed.layout.cube_token_size + is_real = precomputed.layout.is_real + num_heads = padded_q.shape[1] + physical_allowed = _physical_mask_from_block_layout( + block_layout, precomputed.layout, num_heads + ) + real = is_real.to(torch.bool) + allowed = ( + physical_allowed.repeat_interleave(cube_token_size, dim=1).repeat_interleave( + cube_token_size, dim=2 + ) + & real[None, :, None] + & real[None, None, :] + ) + scores = torch.einsum("ihd,jhd->hij", padded_q, padded_k) * softmax_scale + scores = scores.masked_fill(~allowed, float("-inf")) + probs = torch.nan_to_num(torch.softmax(scores, dim=-1), nan=0.0) + return torch.einsum("hij,jhd->ihd", probs, padded_v) + + +class TestCubeSparsePrecompute(unittest.TestCase): + def test_layout_invariants(self): + packed = _build_t2va_packed() + metadata = _build_metadata(packed, [0.5]) + pre = metadata.precomputed + used = int(packed["cu_seqlens"].view(-1)[1]) + cube_token_size = pre.layout.cube_token_size + + self.assertEqual(pre.layout.real_total_len, used) + self.assertEqual(cube_token_size, 8) + self.assertEqual( + pre.layout.padded_seqlen, pre.layout.num_labels * cube_token_size + ) + self.assertEqual(int(pre.layout.label_lengths.sum()), used) + self.assertTrue((pre.layout.label_lengths <= cube_token_size).all()) + # expand/gather round-trip: every real token has a unique padded slot. + expand = pre.layout.expand_indices + self.assertEqual(expand.unique().numel(), used) + self.assertTrue( + torch.equal(pre.layout.gather_indices[expand], torch.arange(used)) + ) + self.assertEqual(int(pre.layout.is_real.sum()), used) + + def test_full_ratio_mask_is_dense_within_sample(self): + packed = _build_t2va_packed() + metadata = _build_metadata(packed, [1.0]) + pre = metadata.precomputed + q, k, _ = _random_qkv(int(packed["seq_len"])) + mask = cube_topk_block_mask( + q[: pre.layout.real_total_len], k[: pre.layout.real_total_len], pre, 1.0 + ) + self.assertTrue(mask.all()) + + def test_builder_rejects_bad_topk_list(self): + packed = _build_t2va_packed() + with self.assertRaisesRegex(ValueError, "denoise steps"): + CubeSparseAttentionMetadataBuilder().build( + packed=packed, + local_cube_size=_CUBE_SIZE, + topk_ratio_list=[0.5, 0.5], + num_steps=3, + device=torch.device("cpu"), + ) + with self.assertRaisesRegex(ValueError, r"\(0, 1\]"): + CubeSparseAttentionMetadataBuilder().build( + packed=packed, + local_cube_size=_CUBE_SIZE, + topk_ratio_list=[0.0], + num_steps=1, + device=torch.device("cpu"), + ) + + def test_builder_requires_stream_layout(self): + packed = _build_t2va_packed() + del packed["stream_layout"] + with self.assertRaisesRegex(ValueError, "stream_layout"): + _build_metadata(packed, [1.0]) + + def test_builder_requires_a_3d_cube_size(self): + packed = _build_t2va_packed() + with self.assertRaisesRegex(ValueError, "positive 3D size"): + CubeSparseAttentionMetadataBuilder().build( + packed=packed, + local_cube_size=(8, 8), + topk_ratio_list=[0.5], + num_steps=1, + device=torch.device("cpu"), + ) + anisotropic = CubeSparseAttentionMetadataBuilder().build( + packed=packed, + local_cube_size=(2, 4, 8), + topk_ratio_list=[0.5], + num_steps=1, + device=torch.device("cpu"), + ) + self.assertEqual(anisotropic.precomputed.layout.cube_token_size, 64) + + +class TestCubeSparseModalityPolicy(unittest.TestCase): + def _allowed_mask(self, packed, topk_ratio): + metadata = _build_metadata(packed, [topk_ratio]) + pre = metadata.precomputed + q, k, _ = _random_qkv(int(packed["seq_len"])) + block_mask = cube_topk_block_mask( + q[: pre.layout.real_total_len], + k[: pre.layout.real_total_len], + pre, + topk_ratio, + ) + labels = _token_labels(pre) + allowed = block_mask[:, labels[:, None], labels[None, :]] + return pre, q, k, block_mask, allowed + + def _assert_has_3d_sparse_labels(self, pre): + self.assertGreater(int(pre.layout.sparse_label_mask.sum()), 0) + + def test_t2va_text_and_audio_are_dense(self): + packed = _build_t2va_packed() + pre, q, k, block_mask, allowed = self._allowed_mask(packed, 0.25) + self._assert_has_3d_sparse_labels(pre) + + _assert_tokens_are_dense(self, allowed, packed["text_pos"].view(-1)) + _assert_tokens_are_dense(self, allowed, packed["audio_pos"].view(-1)) + + img_pos = packed["img_pos"].view(-1).to(torch.long) + target_video = img_pos[packed["update_mask"].view(-1).to(torch.bool)] + target_allowed = allowed.index_select(1, target_video).index_select( + 2, target_video + ) + self.assertFalse(target_allowed.all()) + self.assertTrue( + torch.equal( + block_mask, + _reference_policy_block_mask( + q[: pre.layout.real_total_len], + k[: pre.layout.real_total_len], + pre, + 0.25, + ), + ) + ) + + def test_keyframe_and_target_share_joint_cube_labels(self): + packed = _build_t2va_packed(include_keyframe_cond=True) + pre, _, _, _, allowed = self._allowed_mask(packed, 0.25) + self._assert_has_3d_sparse_labels(pre) + + img_pos = packed["img_pos"].view(-1).to(torch.long) + update_mask = packed["update_mask"].view(-1).to(torch.bool) + keyframe = img_pos[~update_mask] + target = img_pos[update_mask] + frame_rows = 4 * 6 + labels = _token_labels(pre) + self.assertTrue( + torch.equal( + labels.index_select(0, keyframe), + labels.index_select(0, target[:frame_rows]), + ) + ) + # A keyframe duplicates the first temporal plane inside each joint + # cube, so a semantic cube can span more than one physical block. + self.assertGreater( + int(pre.layout.label_lengths.max()), pre.layout.cube_token_size + ) + self.assertFalse(allowed.index_select(1, keyframe).all()) + _assert_tokens_are_dense(self, allowed, packed["text_pos"].view(-1)) + _assert_tokens_are_dense(self, allowed, packed["audio_pos"].view(-1)) + + def test_ref_image_and_audio_are_dense_but_all_videos_share_topk(self): + packed = minimax_h3_packed_sequence_ref2va_blocks( + text_len=17, + latent_t=2, + latent_h=8, + latent_w=12, + audio_t=4, + ref_blocks=[ + {"kind": "image", "latent_h": 8, "latent_w": 12}, + { + "kind": "video_audio", + "ref_audio_t": 3, + "latent_t": 2, + "latent_h": 8, + "latent_w": 12, + }, + ], + ) + pre, q, k, block_mask, allowed = self._allowed_mask(packed, 0.4) + self._assert_has_3d_sparse_labels(pre) + + img_pos = packed["img_pos"].view(-1).to(torch.long) + update_mask = packed["update_mask"].view(-1).to(torch.bool) + condition_visual = img_pos[~update_mask] + ref_image_tokens = 1 * 4 * 6 + ref_image = condition_visual[:ref_image_tokens] + _assert_tokens_are_dense(self, allowed, ref_image) + _assert_tokens_are_dense(self, allowed, packed["text_pos"].view(-1)) + _assert_tokens_are_dense(self, allowed, packed["audio_pos"].view(-1)) + + # Ref video and target video each contain six 2x2x2 cubes. They must + # occupy one shared global candidate pool rather than separate quotas. + self.assertEqual(int(pre.layout.sparse_label_mask.sum()), 12) + self.assertTrue( + torch.equal( + block_mask, + _reference_policy_block_mask( + q[: pre.layout.real_total_len], + k[: pre.layout.real_total_len], + pre, + 0.4, + ), + ) + ) + + +class TestCubeSparseAttentionOutput(unittest.TestCase): + def test_compact_layout_uses_fewer_active_indices_than_full_mask(self): + packed = _build_t2va_packed(include_keyframe_cond=True) + metadata = _build_metadata(packed, [0.4]) + pre = metadata.precomputed + q, k, _ = _random_qkv(int(packed["seq_len"])) + layout = cube_topk_block_indices( + q[: pre.layout.real_total_len], + k[: pre.layout.real_total_len], + pre, + 0.4, + ) + + active_indices = int(layout["kv_num_blocks"].sum()) + full_boolean_cells = _NUM_HEADS * pre.layout.num_blocks**2 + self.assertLess(active_indices, full_boolean_cells) + self.assertGreater(int(layout["kv_num_blocks"].sum()), 0) + self.assertIsNone(layout["full_kv_num_blocks"]) + self.assertIsNone(layout["full_kv_indices"]) + + def test_compact_indices_match_semantic_mask(self): + """Lock the two expansions of ``_cube_topk_selection`` together. + + Production attention consumes ``cube_topk_block_indices`` (compact KV + rows) while the parity and topology tests assert against + ``cube_topk_block_mask`` (a semantic ``[H, L, L]`` mask). They share a + selection helper but expand it independently, so without this test a + divergence in either expansion would only surface as a small numeric + drift inside an attention tolerance -- or not at all, since no + production code path reads the semantic mask. + """ + for include_keyframe_cond in (False, True): + packed = _build_t2va_packed(include_keyframe_cond=include_keyframe_cond) + for topk_ratio in (0.25, 0.4, 1.0): + with self.subTest(keyframe=include_keyframe_cond, ratio=topk_ratio): + metadata = _build_metadata(packed, [topk_ratio]) + pre = metadata.precomputed + real_total_len = pre.layout.real_total_len + q, k, _ = _random_qkv(int(packed["seq_len"])) + q_real = q[:real_total_len] + k_real = k[:real_total_len] + + semantic = cube_topk_block_mask(q_real, k_real, pre, topk_ratio) + compact = cube_topk_block_indices(q_real, k_real, pre, topk_ratio) + + block_labels = pre.layout.block_labels + expected = semantic[:, block_labels[:, None], block_labels[None, :]] + actual = _physical_mask_from_block_layout( + compact, pre.layout, _NUM_HEADS + ) + self.assertTrue( + torch.equal(expected, actual), + f"{int((expected ^ actual).sum())} physical block " + "cells differ between the semantic mask and the " + "compact KV rows", + ) + + def test_compact_kv_rows_never_exceed_block_count(self): + """Pin the invariant that makes the KV truncation lossless. + + ``cube_topk_block_indices`` builds a candidate row of width + ``base_capacity + topk_semantic_capacity * max_label_block_count``, + which can exceed ``num_blocks``, then sorts and truncates to + ``num_blocks``. That is only lossless because base ids and selected + ids are disjoint and ``label -> physical blocks`` is a partition, so no + row can hold more than ``num_blocks`` distinct valid ids. If the + deduplication against ``base_block_mask`` were dropped, the truncation + would silently discard real KV blocks -- visible only as slight output + drift, never as an error. Assert the count directly, and cross-check + it against the semantic mask so an over-count cannot hide behind an + equally wrong row width. + """ + for include_keyframe_cond in (False, True): + packed = _build_t2va_packed(include_keyframe_cond=include_keyframe_cond) + for topk_ratio in (0.25, 0.4): + with self.subTest(keyframe=include_keyframe_cond, ratio=topk_ratio): + metadata = _build_metadata(packed, [topk_ratio]) + pre = metadata.precomputed + real_total_len = pre.layout.real_total_len + q, k, _ = _random_qkv(int(packed["seq_len"])) + q_real = q[:real_total_len] + k_real = k[:real_total_len] + + compact = cube_topk_block_indices(q_real, k_real, pre, topk_ratio) + num_blocks = pre.layout.num_blocks + counts = compact["kv_num_blocks"] + self.assertLessEqual( + int(counts.max()), + num_blocks, + "a compact KV row claims more blocks than exist, so " + "the sort-and-truncate step dropped valid ids", + ) + + # The row width really is over-provisioned relative to + # num_blocks; without that the assertion above is vacuous. + self.assertGreaterEqual( + pre.layout.base_physical_layout["full_kv_indices"].shape[-1] + + pre.layout.topk_semantic_capacity + * pre.layout.max_label_block_count, + num_blocks, + ) + + semantic = cube_topk_block_mask(q_real, k_real, pre, topk_ratio) + block_labels = pre.layout.block_labels + expected_counts = semantic[ + :, block_labels[:, None], block_labels[None, :] + ].sum(dim=-1, dtype=torch.int32) + self.assertTrue( + torch.equal(counts[0], expected_counts), + "compact KV counts disagree with the semantic mask", + ) + + def _run_and_compare(self, packed, topk_ratio): + metadata = _build_metadata(packed, [topk_ratio]) + pre = metadata.precomputed + seq_len = int(packed["seq_len"]) + real_total_len = pre.layout.real_total_len + q, k, v = _random_qkv(seq_len) + scale = _HEAD_DIM**-0.5 + + with mock.patch.object( + cube_backend, + "_run_block_sparse_attention", + _reference_block_sparse_attention, + ): + out = cube_sparse_attention(q, k, v, metadata, softmax_scale=scale) + + label_mask = cube_topk_block_mask( + q[:real_total_len], k[:real_total_len], pre, topk_ratio + ) + labels = _token_labels(pre) + allowed = label_mask[:, labels[:, None], labels[None, :]] + expected = _naive_masked_attention( + q[:real_total_len], + k[:real_total_len], + v[:real_total_len], + allowed, + scale, + ) + + torch.testing.assert_close(out[:real_total_len], expected, atol=2e-4, rtol=2e-4) + self.assertEqual(out.shape, (seq_len, _NUM_HEADS, _HEAD_DIM)) + self.assertTrue((out[real_total_len:] == 0).all()) + + def test_sparse_output_matches_masked_sdpa(self): + self._run_and_compare(_build_t2va_packed(), topk_ratio=0.4) + + def test_keyframe_joint_layout_matches_masked_sdpa(self): + self._run_and_compare( + _build_t2va_packed(include_keyframe_cond=True), topk_ratio=0.4 + ) + + def test_full_ratio_matches_dense_attention(self): + packed = _build_t2va_packed() + metadata = _build_metadata(packed, [1.0]) + real_total_len = metadata.precomputed.layout.real_total_len + q, k, v = _random_qkv(int(packed["seq_len"])) + scale = _HEAD_DIM**-0.5 + + with mock.patch.object( + cube_backend, + "_run_block_sparse_attention", + _reference_block_sparse_attention, + ): + out = cube_sparse_attention(q, k, v, metadata, softmax_scale=scale) + + dense = _naive_masked_attention( + q[:real_total_len], + k[:real_total_len], + v[:real_total_len], + torch.ones(_NUM_HEADS, real_total_len, real_total_len, dtype=torch.bool), + scale, + ) + torch.testing.assert_close(out[:real_total_len], dense, atol=2e-4, rtol=2e-4) + + def test_ref2va_layout_end_to_end(self): + packed = minimax_h3_packed_sequence_ref2va_blocks( + text_len=5, + latent_t=2, + latent_h=8, + latent_w=12, + audio_t=4, + ref_blocks=[ + {"kind": "image", "latent_h": 8, "latent_w": 12}, + { + "kind": "video_audio", + "ref_audio_t": 3, + "latent_t": 2, + "latent_h": 8, + "latent_w": 12, + }, + ], + ) + layout = packed["stream_layout"] + self.assertEqual(layout["cond_image_shapes"], ((1, 4, 6), (2, 4, 6))) + self.assertEqual(layout["cond_audio_stream_lens"], (6,)) + self.assertEqual( + layout["cond_event_orders"], + (("imgvid", 0), ("audio", 0), ("imgvid", 1)), + ) + self._run_and_compare(packed, topk_ratio=0.4) + + +class TestCubeSparseRequestConfig(unittest.TestCase): + def test_full_ratio_uses_native_dense_backend(self): + packed = _build_t2va_packed() + metadata = _build_metadata(packed, [1.0]) + q, k, v = _random_qkv(int(packed["seq_len"])) + expected = torch.empty_like(q) + impl = CubeSparseAttentionImpl( + num_heads=_NUM_HEADS, + head_size=_HEAD_DIM, + causal=False, + softmax_scale=_HEAD_DIM**-0.5, + ) + impl._dense_impl = mock.Mock() + impl._dense_impl.forward_varlen.return_value = expected + + with set_forward_context(0, metadata): + actual = impl.forward_varlen( + q, + k, + v, + cu_seqlens=packed["cu_seqlens"].to(torch.int32), + max_seqlen=int(packed["cu_seqlens"][1]), + ) + + self.assertIs(actual, expected) + impl._dense_impl.forward_varlen.assert_called_once() + + def test_cube_config_defaults_unrelated_vsa_sparsity(self): + server_args = SimpleNamespace( + attention_backend_config={ + "local_cube_size": [4, 4, 4], + "topk_ratio_list": [1.0], + }, + enable_trace=False, + ) + + request = prepare_request(server_args, SamplingParams(prompt="test")) + + self.assertEqual(request.VSA_sparsity, 0.0) + + +class TestCubeSparseAttentionCuda(unittest.TestCase): + """Compiled-flex path on CUDA with the production cube size (64 tokens).""" + + @classmethod + def setUpClass(cls): + if not torch.cuda.is_available(): + raise unittest.SkipTest("requires CUDA") + + def _run(self, topk_ratio): + device = torch.device("cuda") + packed = minimax_h3_packed_sequence( + text_len=33, + latent_t=4, + latent_h=32, + latent_w=48, + audio_t=17, + include_keyframe_cond=False, + ) + metadata = CubeSparseAttentionMetadataBuilder().build( + packed=packed, + local_cube_size=(4, 4, 4), + topk_ratio_list=[topk_ratio], + num_steps=1, + device=device, + ) + pre = metadata.precomputed + real_total_len = pre.layout.real_total_len + seq_len = int(packed["seq_len"]) + generator = torch.Generator(device="cpu").manual_seed(1) + q, k, v = ( + torch.randn(seq_len, 4, 128, dtype=torch.float32, generator=generator).to( + device=device, dtype=torch.bfloat16 + ) + for _ in range(3) + ) + scale = 128**-0.5 + + out = cube_sparse_attention(q, k, v, metadata, softmax_scale=scale) + + label_mask = cube_topk_block_mask( + q[:real_total_len], k[:real_total_len], pre, topk_ratio + ) + labels = _token_labels(pre) + allowed = label_mask[:, labels[:, None], labels[None, :]] + expected = _naive_masked_attention( + q[:real_total_len].float(), + k[:real_total_len].float(), + v[:real_total_len].float(), + allowed, + scale, + ) + torch.testing.assert_close( + out[:real_total_len].float(), expected, atol=2.6e-2, rtol=2.6e-2 + ) + self.assertTrue((out[real_total_len:] == 0).all()) + + def test_compiled_flex_sparse(self): + self._run(topk_ratio=0.25) + + def test_compiled_flex_full_ratio(self): + self._run(topk_ratio=1.0) + + +if __name__ == "__main__": + unittest.main() diff --git a/python/sglang/multimodal_gen/test/unit/test_cuda_attention_backend.py b/python/sglang/multimodal_gen/test/unit/test_cuda_attention_backend.py index fbf7522b7..b53cd246b 100644 --- a/python/sglang/multimodal_gen/test/unit/test_cuda_attention_backend.py +++ b/python/sglang/multimodal_gen/test/unit/test_cuda_attention_backend.py @@ -1,7 +1,8 @@ import sys import types import unittest -from unittest.mock import patch +from types import ModuleType +from unittest.mock import Mock, patch import torch @@ -86,6 +87,26 @@ class TestCudaAttentionBackendSelection(unittest.TestCase): "sglang.multimodal_gen.runtime.layers.attention.backends.aiter.AITerBackend", ) + def test_direct_cube_sparse_selection(self): + self.assertEqual( + self.resolve(AttentionBackendEnum.CUBE_SPARSE_ATTN), + "sglang.multimodal_gen.runtime.layers.attention.backends." + "cube_sparse_attn.CubeSparseAttentionBackend", + ) + + def test_blackwell_cube_selection_initializes_fa4_for_token_refiner(self): + FakeCudaPlatform.is_blackwell_device = True + module_name = ( + "sglang.multimodal_gen.runtime.layers.attention.backends.flash_attn" + ) + fake_flash_attn = ModuleType(module_name) + fake_flash_attn.set_fa_ver = Mock() + + with patch.dict("sys.modules", {module_name: fake_flash_attn}): + self.resolve(AttentionBackendEnum.CUBE_SPARSE_ATTN) + + fake_flash_attn.set_fa_ver.assert_called_once_with(4) + def test_default_backend_uses_torch_sdpa_on_sm120(self): FakeCudaPlatform.is_sm120_device = True diff --git a/python/sglang/multimodal_gen/test/unit/test_minimax_h3_admission.py b/python/sglang/multimodal_gen/test/unit/test_minimax_h3_admission.py index ec74ed9cd..ef9a18ab4 100644 --- a/python/sglang/multimodal_gen/test/unit/test_minimax_h3_admission.py +++ b/python/sglang/multimodal_gen/test/unit/test_minimax_h3_admission.py @@ -465,6 +465,15 @@ def test_validate_server_args_requires_packed_varlen_backend(): with pytest.raises(ValueError, match="does not implement packed varlen"): MiniMaxH3PipelineConfig.validate_server_args(config, server_args) + server_args.component_attention_backends = {"transformer": "cube_sparse_attn"} + server_args.resolve_component_attention_backend = lambda *_names: ( + AttentionBackendEnum.CUBE_SPARSE_ATTN, + "transformer", + ) + server_args.ring_degree = 2 + with pytest.raises(ValueError, match="ring parallelism requires"): + MiniMaxH3PipelineConfig.validate_server_args(config, server_args) + def test_validate_server_args_accepts_transformer_backend_override(): config = MiniMaxH3PipelineConfig() diff --git a/python/sglang/multimodal_gen/test/unit/test_minimax_h3_denoise_loop.py b/python/sglang/multimodal_gen/test/unit/test_minimax_h3_denoise_loop.py index 2a57c638d..4b364b582 100644 --- a/python/sglang/multimodal_gen/test/unit/test_minimax_h3_denoise_loop.py +++ b/python/sglang/multimodal_gen/test/unit/test_minimax_h3_denoise_loop.py @@ -1,13 +1,16 @@ # SPDX-License-Identifier: Apache-2.0 """Numerical contract for request-static H3 denoise metadata.""" +from types import SimpleNamespace from unittest.mock import patch +import pytest import torch from sglang.multimodal_gen.configs.models.dits.minimax_h3 import ( MINIMAX_H3_ADALN_MODALITY_NUM, ) +from sglang.multimodal_gen.runtime.managers.forward_context import get_forward_context from sglang.multimodal_gen.runtime.models.schedulers.scheduling_minimax_h3_euler_ancestral import ( _minimax_h3_euler_eta0_step, _minimax_h3_rf_v_to_x0, @@ -16,12 +19,15 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.m MiniMaxH3DenoiseBranch, _build_local_embedding_layout, _minimax_h3_update_target_rows_, + minimax_h3_denoise_loop, ) from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.minimax_h3.packed_sequence import ( minimax_h3_packed_sequence, minimax_h3_packed_sequence_ref2va_blocks, ) from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.minimax_h3.stages.denoising import ( + MiniMaxH3DenoisingStage, + _build_cube_attn_metadata, _precompute_refined_prompt_embeds, ) @@ -198,6 +204,138 @@ def test_rank_local_token_tags_match_reference_slice(): ) +def test_cube_metadata_builder_uses_packed_layout_and_validates_step_count(): + packed = minimax_h3_packed_sequence( + text_len=3, + latent_t=2, + latent_h=8, + latent_w=8, + audio_t=3, + include_keyframe_cond=False, + ) + server_args = SimpleNamespace( + attention_backend="cube_sparse_attn", + component_attention_backends={}, + attention_backend_config={ + "local_cube_size": [4, 4, 4], + "topk_ratio_list": [1.0, 0.5], + }, + ) + + metadata = _build_cube_attn_metadata( + server_args, + packed=packed, + num_steps=2, + device=torch.device("cpu"), + ) + assert metadata.topk_ratio_list == [1.0, 0.5] + assert metadata.precomputed.layout.cube_token_size == 64 + + with pytest.raises(ValueError, match="denoise steps"): + _build_cube_attn_metadata( + server_args, + packed=packed, + num_steps=3, + device=torch.device("cpu"), + ) + + +def test_cube_metadata_follows_transformer_backend_override(): + packed = minimax_h3_packed_sequence( + text_len=3, + latent_t=2, + latent_h=8, + latent_w=8, + audio_t=3, + include_keyframe_cond=False, + ) + server_args = SimpleNamespace( + attention_backend="fa", + component_attention_backends={"transformer": "cube_sparse_attn"}, + attention_backend_config={ + "local_cube_size": [4, 4, 4], + "topk_ratio_list": [0.5], + }, + ) + assert ( + _build_cube_attn_metadata( + server_args, + packed=packed, + num_steps=1, + device=torch.device("cpu"), + ) + is not None + ) + + server_args.attention_backend = "cube_sparse_attn" + server_args.component_attention_backends["transformer"] = "fa" + assert ( + _build_cube_attn_metadata( + server_args, + packed=packed, + num_steps=1, + device=torch.device("cpu"), + ) + is None + ) + + +def test_cube_metadata_is_updated_per_step(): + branch = _branch("t2va") + metadata = SimpleNamespace(current_timestep=-1, topk_ratio_list=[1.0, 0.25]) + seen = [] + + def model_forward(_model, _kwargs, step): + seen.append((step, metadata.current_timestep)) + return ( + torch.zeros(int(branch.update_mask.sum()), 96), + torch.zeros(branch.audio_pos.numel(), 32), + ) + + minimax_h3_denoise_loop( + model=SimpleNamespace(prepare_adaln_plans=lambda _: None), + model_forward=model_forward, + positive=branch, + initial_video_rows=torch.zeros(branch.img_pos.numel(), 96), + initial_audio_rows=torch.zeros(branch.audio_pos.numel(), 32), + keyframe_cond_rows=None, + sigmas_video=[1.0, 0.5, 0.0], + sigmas_audio=[1.0, 0.5, 0.0], + device=torch.device("cpu"), + attn_metadata=metadata, + ) + + assert seen == [(0, 0), (1, 1)] + + +def test_native_dit_forward_publishes_cube_metadata_in_forward_context(): + metadata = SimpleNamespace(current_timestep=0, topk_ratio_list=[0.5]) + batch = SimpleNamespace() + + def model(**_kwargs): + context = get_forward_context() + assert context.current_timestep == 0 + assert context.attn_metadata is metadata + assert context.forward_batch is batch + return torch.zeros(1, 96), torch.zeros(1, 32) + + stage = MiniMaxH3DenoisingStage.__new__(MiniMaxH3DenoisingStage) + with patch.object( + MiniMaxH3DenoisingStage, + "_maybe_get_bcg_runner", + return_value=None, + ): + video, audio = stage._forward_dit( + model, + {}, + 0, + batch=batch, + attn_metadata=metadata, + ) + assert video.shape == (1, 96) + assert audio.shape == (1, 32) + + def test_grouped_outputs_share_prompt_refinement(): class Refiner: calls = 0 diff --git a/python/sglang/multimodal_gen/test/unit/test_minimax_h3_dit_contract.py b/python/sglang/multimodal_gen/test/unit/test_minimax_h3_dit_contract.py index 613d0b518..54a1b8f0f 100644 --- a/python/sglang/multimodal_gen/test/unit/test_minimax_h3_dit_contract.py +++ b/python/sglang/multimodal_gen/test/unit/test_minimax_h3_dit_contract.py @@ -18,6 +18,9 @@ from sglang.multimodal_gen.runtime.distributed.parallel_state import ( from sglang.multimodal_gen.runtime.layers.attention.backends.attention_backend import ( AttentionRequirements, ) +from sglang.multimodal_gen.runtime.layers.attention.backends.cube_sparse_attn import ( + CubeSparseAttentionBackend, +) from sglang.multimodal_gen.runtime.layers.attention.backends.sdpa import SDPAImpl from sglang.multimodal_gen.runtime.layers.attention.selector import ( component_attn_backend_context_manager, @@ -35,6 +38,7 @@ from sglang.multimodal_gen.runtime.loader.utils import get_param_names_mapping from sglang.multimodal_gen.runtime.models.dits.minimax_h3 import ( MINIMAX_H3_FP32_BUFFER_NAMES, MINIMAX_H3_FP32_PARAM_NAMES, + MiniMaxH3Attention, MiniMaxH3DiTBlock, MiniMaxH3DiTModel, _copy_grouped_qkv_tp_shard, @@ -409,6 +413,109 @@ def test_tp_and_ulysses_admission_uses_tp_local_shapes(): ) +def test_cube_backend_advertises_packed_varlen_capability(): + assert CubeSparseAttentionBackend.supports_packed_varlen() + + +def test_attention_retains_transformer_scoped_backend(): + _ensure_single_process_parallel_runtime() + with ( + patch( + "sglang.multimodal_gen.runtime.models.dits.minimax_h3." + "get_component_forced_attn_backend", + return_value=AttentionBackendEnum.CUBE_SPARSE_ATTN, + ), + torch.device("meta"), + ): + attention = MiniMaxH3Attention( + MiniMaxH3DiTArchConfig(), None, prefix="blocks.0.attn" + ) + + assert ( + attention._selected_attention_backend is AttentionBackendEnum.CUBE_SPARSE_ATTN + ) + + +def test_model_lazy_resolver_keeps_transformer_scoped_backend(): + _ensure_single_process_parallel_runtime() + with ( + patch( + "sglang.multimodal_gen.runtime.models.dits.minimax_h3." + "get_component_forced_attn_backend", + return_value=AttentionBackendEnum.CUBE_SPARSE_ATTN, + ), + torch.device("meta"), + ): + model = MiniMaxH3DiTModel( + config=MiniMaxH3DiTConfig(), hf_config={}, quant_config=None + ) + + class FakeBackend: + @staticmethod + def get_enum(): + return AttentionBackendEnum.CUBE_SPARSE_ATTN + + with ( + patch( + "sglang.multimodal_gen.runtime.models.dits.minimax_h3.get_attn_backend", + return_value=FakeBackend, + ) as resolve, + patch.object(MiniMaxH3Attention, "_set_attention_backend") as install, + ): + model._resolve_attention_backend_once() + + resolve.assert_called_once_with( + model.arch.attention_head_dim, + torch.bfloat16, + selected_attention_backend=AttentionBackendEnum.CUBE_SPARSE_ATTN, + attention_requirements=AttentionRequirements(packed_varlen=True), + ) + attention_count = sum( + isinstance(module, MiniMaxH3Attention) for module in model.modules() + ) + assert install.call_count == attention_count + assert model._resolved_attention_backend is AttentionBackendEnum.CUBE_SPARSE_ATTN + + +def test_token_refiner_routes_cube_selection_to_exact_fa(): + class FakeImpl: + def __init__(self, **kwargs): + self.kwargs = kwargs + + class FakeBackend: + def __init__(self, enum): + self.enum = enum + + def get_enum(self): + return self.enum + + def get_impl_cls(self): + return FakeImpl + + cube = FakeBackend(AttentionBackendEnum.CUBE_SPARSE_ATTN) + fa = FakeBackend(AttentionBackendEnum.FA) + attention = MiniMaxH3Attention.__new__(MiniMaxH3Attention) + torch.nn.Module.__init__(attention) + attention.head_dim = 128 + attention.num_heads = 8 + attention.softmax_scale = 128**-0.5 + attention.prefix = "test.attn" + attention._cube_sparse_capable = False + + with patch( + "sglang.multimodal_gen.runtime.models.dits.minimax_h3.get_attn_backend", + return_value=fa, + ) as resolve: + attention._set_attention_backend(cube) + + assert attention._attention_backend_enum is AttentionBackendEnum.FA + resolve.assert_called_once_with( + 128, + torch.bfloat16, + selected_attention_backend=AttentionBackendEnum.FA, + ) + + def test_meta_model_captures_component_attention_override(): _ensure_single_process_parallel_runtime() with ( diff --git a/python/sglang/multimodal_gen/test/unit/test_minimax_h3_packed_sequence.py b/python/sglang/multimodal_gen/test/unit/test_minimax_h3_packed_sequence.py index 58af1079a..302707562 100644 --- a/python/sglang/multimodal_gen/test/unit/test_minimax_h3_packed_sequence.py +++ b/python/sglang/multimodal_gen/test/unit/test_minimax_h3_packed_sequence.py @@ -23,6 +23,16 @@ class TestMiniMaxH3PackedSequence(unittest.TestCase): self.assertEqual(int(built["img_pos"].shape[0]), 62 * 24 * 38) self.assertEqual(int(built["seq_len"]) % 64, 0) self.assertEqual(built["token_tags"][built["audio_pos"]].unique().tolist(), [2]) + self.assertEqual( + built["stream_layout"], + { + "target_shape": (62, 24, 38), + "cond_image_shapes": (), + "cond_image_roles": (), + "cond_event_orders": (), + "cond_audio_stream_lens": (), + }, + ) def test_fl2va_first_last_cond_blocks_use_exact_rope_span(self): text_len = 11 @@ -61,6 +71,16 @@ class TestMiniMaxH3PackedSequence(unittest.TestCase): ) self.assertFalse(built["update_mask"][:cond_rows].any()) self.assertTrue(built["update_mask"][cond_rows:].all()) + self.assertEqual( + built["stream_layout"], + { + "target_shape": (37, 24, 38), + "cond_image_shapes": ((1, 24, 38), (1, 24, 38)), + "cond_image_roles": ("joint_cube", "joint_cube"), + "cond_event_orders": (("imgvid", 0), ("imgvid", 1)), + "cond_audio_stream_lens": (), + }, + ) def test_i2va_and_l2va_single_cond_blocks_use_endpoint_rope(self): text_len = 11 @@ -154,6 +174,21 @@ class TestMiniMaxH3PackedSequence(unittest.TestCase): target_video_t0 = built["img_position_ids"][built["img_pos"][12], 0] target_audio_t0 = built["img_position_ids"][built["audio_pos"][8], 0] self.assertEqual(float(target_audio_t0), float(target_video_t0)) + self.assertEqual( + built["stream_layout"], + { + "target_shape": (2, 2, 2), + "cond_image_shapes": ((1, 2, 2), (2, 2, 2)), + "cond_image_roles": ("dense_prefix", "independent_cube"), + "cond_event_orders": ( + ("imgvid", 0), + ("audio", 0), + ("imgvid", 1), + ("audio", 1), + ), + "cond_audio_stream_lens": (6, 2), + }, + ) def test_ref2va_hybrid_packs_keyframes_before_references(self): built = minimax_h3_packed_sequence_ref2va_blocks(