[diffusion] feat: support FastH3 (4-step VSA-distilled MiniMax-H3) with a VSA-H3 attention backend (#37480)
Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Fable 5
parent
99b9109553
commit
f586654518
+1
@@ -314,6 +314,7 @@ Use the preset categories this way:
|
||||
| `cosmos3-super-t2v-cfg2tp2` | `nvidia/Cosmos3-Super` | No | Explicit four-GPU TP2 x CFG2 throughput comparator. On H200 it was 48.00% faster end to end than TP2, but the topology changed the deterministic output (SSIM 0.914244, PSNR 29.469771 dB), so do not treat it as lossless-equivalent or select it automatically. |
|
||||
| `wan-i2v` | `Wan-AI/Wan2.2-I2V-A14B-Diffusers` | Yes: `wan22_i2v_a14b_720p` | Nightly cat image and motion prompt, 1280x720, 81 frames, 4 GPUs, CFG parallel, Ulysses degree 2, text encoder CPU offload and pinned CPU memory |
|
||||
| `minimax-h3-t2va` | `MiniMaxAI/MiniMax-H3` | Yes: `minimax_h3_t2va_5s` | H3 FL2VA-partition T2VA baseline: 1344x768 resolved canvas, 5 seconds / 124 frames at 24 fps, 50 joint video-audio steps, 4 GPUs, TP2 + Ulysses2, eager BF16/FP32. The helper writes H3's request contract to a generated config. |
|
||||
| `fasth3-t2va-vsa` | `FastVideo/FastVideo-FastH3-4-step-Preview-v1-VSA-DataFree` | No | FastH3 4-step distilled T2VA on the trained VSA-H3 backend: 1344x768, 10 seconds / 243 frames, five sigma points = four DiT forwards, 4 GPUs, Ulysses 4, eager, 2-step warmup request. Compare against `--attention-backend fa` on the same weights for the dense-fallback gap. |
|
||||
| `longcat-image` | `meituan-longcat/LongCat-Image` | No | Eager DiT baseline at 1024x1024, 50 steps, guidance 4.5; prompt rewrite is disabled so Qwen2.5-VL does not contaminate the DiT A/B. |
|
||||
| `longcat-image-edit` | `meituan-longcat/LongCat-Image-Edit` | No | Native edit baseline using the public SGLang edit fixture. Its 1536x1024 source resolves to 1264x848 under the checkpoint's roughly-one-megapixel aspect-ratio rule, and the BCG comparator captures that exact serving canvas; prompt rewrite is disabled to isolate the DiT. |
|
||||
| `longcat-image-edit-turbo` | `meituan-longcat/LongCat-Image-Edit-Turbo` | No | Matching distilled edit baseline using the same public fixture, prompt, and 1264x848 BCG canvas. Its registered sampling class owns the eight-step, guidance-1 schedule. |
|
||||
|
||||
+27
@@ -425,6 +425,33 @@ MODELS = {
|
||||
"num-inference-steps",
|
||||
},
|
||||
},
|
||||
# H3 rejects a 1-step warmup request, hence --warmup-steps=2.
|
||||
"fasth3-t2va-vsa": {
|
||||
"path": "FastVideo/FastVideo-FastH3-4-step-Preview-v1-VSA-DataFree",
|
||||
"prompt": (
|
||||
"A curious raccoon peers through a vibrant field of yellow "
|
||||
"sunflowers, its eyes wide with interest."
|
||||
),
|
||||
"seed": 1000,
|
||||
"config_overrides": {
|
||||
"task": "t2va",
|
||||
"conditions": [],
|
||||
"target": {
|
||||
"short_edge": 768,
|
||||
"aspect_ratio": "16:9",
|
||||
"duration_seconds": 10.0,
|
||||
},
|
||||
"num_inference_steps": 5,
|
||||
},
|
||||
"extra_args": [
|
||||
"--num-gpus=4",
|
||||
"--attention-backend=video_sparse_attn_h3",
|
||||
'--attention-backend-config={"VSA_sparsity": 0.9}',
|
||||
"--enable-torch-compile=false",
|
||||
"--warmup-steps=2",
|
||||
],
|
||||
"force_eager": True,
|
||||
},
|
||||
# Source-tracked extras from current registry / GPU test coverage.
|
||||
"longcat-image": {
|
||||
"path": "meituan-longcat/LongCat-Image",
|
||||
|
||||
@@ -9,7 +9,7 @@ SGLang diffusion features an end-to-end unified pipeline for accelerating diffus
|
||||
## Key Features
|
||||
|
||||
SGLang Diffusion has the following features:
|
||||
- Broad model support: Wan, FastWan, FLUX, Qwen-Image, LongCat-Image, Z-Image, Ideogram 4, Krea-2, Cosmos3, LTX-2/LTX-2.3/LTX-2.5, MiniMax-H3, LingBot Video MoE, LingBot World, SANA-Video/SANA-WM, JoyEcho, MOVA, GLM-Image, ERNIE-Image, Hunyuan3D, and more
|
||||
- Broad model support: Wan, FastWan, FLUX, Qwen-Image, LongCat-Image, Z-Image, Ideogram 4, Krea-2, Cosmos3, LTX-2/LTX-2.3/LTX-2.5, MiniMax-H3, FastH3, LingBot Video MoE, LingBot World, SANA-Video/SANA-WM, JoyEcho, MOVA, GLM-Image, ERNIE-Image, Hunyuan3D, and more
|
||||
- Fast inference speed: empowered by optimized `sgl-kernel` kernels, scheduler/runtime improvements, caching acceleration, and native diffusion hot-path optimizations
|
||||
- Ease of use: OpenAI-compatible api, CLI, and python sdk support
|
||||
- Multi-platform support:
|
||||
|
||||
@@ -46,6 +46,7 @@ class MiniMaxH3DiTArchConfig(DiTArchConfig):
|
||||
3,
|
||||
),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_out\.0\.(.*)$": r"blocks.\1.attn.out_proj.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_gate_compress\.(.*)$": r"blocks.\1.attn.to_gate_compress.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn\.norm_q\.(.*)$": r"blocks.\1.attn.q_norm.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn\.norm_k\.(.*)$": r"blocks.\1.attn.k_norm.\2",
|
||||
r"^transformer_blocks\.(\d+)\.ff\.net\.0\.proj\.(.*)$": r"blocks.\1.mlp.fc1.\2",
|
||||
@@ -99,6 +100,7 @@ class MiniMaxH3DiTArchConfig(DiTArchConfig):
|
||||
final_norm_eps: float = 1e-5
|
||||
checkpoint_uses_diffusers_layout: bool = False
|
||||
adaln_affine_input_dim: int | None = None
|
||||
has_gate_compress: bool = False
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
super().__post_init__()
|
||||
|
||||
@@ -24,6 +24,14 @@ class MiniMaxH3AudioVAEConfig(VAEConfig):
|
||||
load_encoder: bool = True
|
||||
load_decoder: bool = True
|
||||
|
||||
def update_model_arch(self, source_model_dict: dict) -> None:
|
||||
# Native-Diffusers AutoencoderKLMiniMaxH3Audio config field name.
|
||||
aliases = {"sampling_rate": "sample_rate"}
|
||||
model_dict = {
|
||||
aliases.get(key, key): value for key, value in source_model_dict.items()
|
||||
}
|
||||
super().update_model_arch(model_dict)
|
||||
|
||||
def post_init(self) -> None:
|
||||
validate_minimax_h3_vae_latent_stats(
|
||||
self.arch_config,
|
||||
|
||||
@@ -56,6 +56,17 @@ class MiniMaxH3VideoVAEConfig(VAEConfig):
|
||||
f"{self.parallel_decode_mode!r}"
|
||||
)
|
||||
|
||||
def update_model_arch(self, source_model_dict: dict) -> None:
|
||||
# Native-Diffusers AutoencoderKLMiniMaxH3 config field names.
|
||||
aliases = {
|
||||
"clip_length": "vae_clip_length",
|
||||
"token_drop": "vae_token_drop",
|
||||
}
|
||||
model_dict = {
|
||||
aliases.get(key, key): value for key, value in source_model_dict.items()
|
||||
}
|
||||
super().update_model_arch(model_dict)
|
||||
|
||||
def post_init(self) -> None:
|
||||
self.resolved_parallel_decode_mode()
|
||||
validate_minimax_h3_vae_latent_stats(
|
||||
|
||||
@@ -44,6 +44,7 @@ from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import (
|
||||
LTX23PipelineConfig,
|
||||
)
|
||||
from sglang.multimodal_gen.configs.pipeline_configs.minimax_h3 import (
|
||||
FastH3PipelineConfig,
|
||||
MiniMaxH3PipelineConfig,
|
||||
)
|
||||
from sglang.multimodal_gen.configs.pipeline_configs.mova import MOVAPipelineConfig
|
||||
@@ -97,5 +98,6 @@ __all__ = [
|
||||
"LingBotWorldCausalDMDConfig",
|
||||
"LingBotWorldV2CausalDMDConfig",
|
||||
"LingBotVideoMoEPipelineConfig",
|
||||
"FastH3PipelineConfig",
|
||||
"MiniMaxH3PipelineConfig",
|
||||
]
|
||||
|
||||
@@ -260,6 +260,21 @@ class MiniMaxH3PipelineConfig(PipelineConfig):
|
||||
)
|
||||
if selected_backend is None:
|
||||
return
|
||||
if selected_backend is AttentionBackendEnum.VIDEO_SPARSE_ATTN_H3:
|
||||
if server_args.ring_degree > 1:
|
||||
raise ValueError(
|
||||
"VSA-H3 does not support --ring-degree > 1; use Ulysses "
|
||||
"sequence parallelism."
|
||||
)
|
||||
if (
|
||||
server_args.enable_torch_compile
|
||||
or server_args.enable_breakable_cuda_graph
|
||||
):
|
||||
raise ValueError(
|
||||
"VSA-H3 builds per-step tile metadata eagerly and is not "
|
||||
"validated under torch.compile or the breakable CUDA "
|
||||
"graph; disable them or use --attention-backend fa."
|
||||
)
|
||||
get_attn_backend(
|
||||
self.dit_config.arch_config.attention_head_dim,
|
||||
torch.bfloat16,
|
||||
@@ -279,4 +294,28 @@ class MiniMaxH3PipelineConfig(PipelineConfig):
|
||||
return safetensors_list
|
||||
|
||||
|
||||
__all__ = ["MiniMaxH3PipelineConfig"]
|
||||
@dataclass
|
||||
class FastH3PipelineConfig(MiniMaxH3PipelineConfig):
|
||||
"""FastH3: 4-step VSA-distilled MiniMax-H3, t2va only."""
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.dit_config.arch_config.has_gate_compress = True
|
||||
|
||||
def validate_quality_deployment(self, server_args) -> None:
|
||||
raise ValueError(
|
||||
'quality="high" is audited only for the base MiniMax-H3 50-step '
|
||||
"4xH200 deployment; the FastH3 4-step distilled checkpoint has no "
|
||||
'audited high-quality deployment. Use quality="lossless".'
|
||||
)
|
||||
|
||||
def validate_server_args(self, server_args) -> None:
|
||||
if server_args.model_variant is not None:
|
||||
raise ValueError(
|
||||
"FastH3 ships one t2va-distilled weight partition; "
|
||||
"--model-variant does not apply. FL2VA and Ref2VA tasks were "
|
||||
"not distilled; use MiniMaxAI/MiniMax-H3 for those."
|
||||
)
|
||||
super().validate_server_args(server_args)
|
||||
|
||||
|
||||
__all__ = ["FastH3PipelineConfig", "MiniMaxH3PipelineConfig"]
|
||||
|
||||
@@ -302,4 +302,26 @@ class MiniMaxH3SamplingParams(SamplingParams):
|
||||
req.extra.update(self.build_request_extra())
|
||||
|
||||
|
||||
__all__ = ["MiniMaxH3SamplingParams"]
|
||||
@dataclass
|
||||
class FastH3SamplingParams(MiniMaxH3SamplingParams):
|
||||
"""FastH3: five sigma grid points, i.e. the four distilled DiT forwards."""
|
||||
|
||||
num_inference_steps: int = 5
|
||||
|
||||
def _validate(self) -> None:
|
||||
super()._validate()
|
||||
if self.num_inference_steps != 5:
|
||||
raise ValueError(
|
||||
"FastH3 is distilled for exactly five sigma grid points (four DiT "
|
||||
f"forwards); got num_inference_steps={self.num_inference_steps}. "
|
||||
"Use MiniMaxAI/MiniMax-H3 for other schedules."
|
||||
)
|
||||
if self.task is not None and self.task.strip().lower() != "t2va":
|
||||
raise ValueError(
|
||||
"FastH3 is distilled for t2va only; fl2va and ref2va were not "
|
||||
f"distilled (got task={self.task!r}). Use MiniMaxAI/MiniMax-H3 "
|
||||
"for those tasks."
|
||||
)
|
||||
|
||||
|
||||
__all__ = ["FastH3SamplingParams", "MiniMaxH3SamplingParams"]
|
||||
|
||||
@@ -29,6 +29,7 @@ if TYPE_CHECKING:
|
||||
|
||||
from sglang.multimodal_gen.configs.pipeline_configs import (
|
||||
Cosmos3Config,
|
||||
FastH3PipelineConfig,
|
||||
FastHunyuanConfig,
|
||||
FluxPipelineConfig,
|
||||
HeliosDistilledConfig,
|
||||
@@ -162,7 +163,10 @@ from sglang.multimodal_gen.configs.sample.ltx_2 import (
|
||||
LTX23SamplingParams,
|
||||
)
|
||||
from sglang.multimodal_gen.configs.sample.ltx_2_5 import LTX25SamplingParams
|
||||
from sglang.multimodal_gen.configs.sample.minimax_h3 import MiniMaxH3SamplingParams
|
||||
from sglang.multimodal_gen.configs.sample.minimax_h3 import (
|
||||
FastH3SamplingParams,
|
||||
MiniMaxH3SamplingParams,
|
||||
)
|
||||
from sglang.multimodal_gen.configs.sample.mova import (
|
||||
MOVA_360P_SamplingParams,
|
||||
MOVA_720P_SamplingParams,
|
||||
@@ -334,6 +338,7 @@ _MODEL_NAME_DETECTORS: List[Tuple[str, Callable[[str], bool]]] = []
|
||||
KNOWN_NON_DIFFUSERS_DIFFUSION_MODEL_PATTERNS: Dict[str, str] = {
|
||||
"minimaxai/minimax-h3": "MiniMaxH3Pipeline",
|
||||
"minimax/minimax-h3": "MiniMaxH3Pipeline",
|
||||
"fastvideo/fastvideo-fasth3-4-step-preview-v1-vsa-datafree": "FastH3Pipeline",
|
||||
"lerobot/pi05": "Pi05Pipeline",
|
||||
"pi05": "Pi05Pipeline",
|
||||
"pi0.5": "Pi05Pipeline",
|
||||
@@ -972,6 +977,17 @@ def _register_configs():
|
||||
in model_id.lower().replace("-", "").replace("_", "")
|
||||
],
|
||||
)
|
||||
register_configs(
|
||||
sampling_param_cls=FastH3SamplingParams,
|
||||
pipeline_config_cls=FastH3PipelineConfig,
|
||||
hf_model_paths=[
|
||||
"FastVideo/FastVideo-FastH3-4-step-Preview-v1-VSA-DataFree",
|
||||
],
|
||||
model_detectors=[
|
||||
lambda model_id: "fasth3"
|
||||
in model_id.lower().replace("-", "").replace("_", "")
|
||||
],
|
||||
)
|
||||
# FLUX
|
||||
register_configs(
|
||||
sampling_param_cls=FluxSamplingParams,
|
||||
|
||||
+488
@@ -0,0 +1,488 @@
|
||||
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
|
||||
# (the video_sparse_attn_h3 backend), rewritten for SGLang's packed-varlen
|
||||
# MiniMax-H3 attention contract.
|
||||
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""VSA for MiniMax-H3's packed mixed-modality self-attention.
|
||||
|
||||
H3 runs one joint bidirectional attention over
|
||||
``[text | condition keyframes | audio | generated video]``. Tiles are
|
||||
``[segment-pure prefix chunks] + [3D video tiles]``; prefix tiles never
|
||||
straddle segment boundaries. Selection is a top-k over pooled tile scores
|
||||
emitted directly as the per-query-tile index lists the vendored Triton
|
||||
tile-64 kernel consumes, with per-tile valid sizes so ragged interior tiles
|
||||
mask exactly.
|
||||
|
||||
Non-video queries are always dense. Non-video keys are either
|
||||
always-selected for every query ("exempt", default) or compete in top-k
|
||||
under a FLOP-matched budget ("compete"). The compression branch is gated by
|
||||
``to_gate_compress``: base H3 has no such weights and the gate loads as
|
||||
zeros (pure sparse); VSA-distilled students (FastH3) ship trained gates
|
||||
that activate the branch.
|
||||
"""
|
||||
|
||||
import functools
|
||||
import math
|
||||
import re
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.multimodal_gen.runtime.layers.attention.backends.attention_backend import (
|
||||
AttentionBackend,
|
||||
AttentionImpl,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.layers.attention.backends.video_sparse_attn import (
|
||||
construct_variable_block_sizes,
|
||||
get_non_pad_index,
|
||||
get_tile_partition_indices,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.layers.attention.backends.vsa_h3_kernels import (
|
||||
vsa_h3_block_sparse_attn_forward,
|
||||
vsa_h3_pack_tiles,
|
||||
vsa_h3_untile,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum
|
||||
|
||||
# The FastH3 checkpoints are trained and served at the 64-token (4, 4, 4)
|
||||
# tile geometry; the kernel block size matches it exactly.
|
||||
VSA_H3_TILE_SHAPE = (4, 4, 4)
|
||||
VSA_H3_TILE_ELEMS = math.prod(VSA_H3_TILE_SHAPE)
|
||||
|
||||
_DIT_BLOCK_PREFIX = re.compile(r"^blocks\.(\d+)\.")
|
||||
|
||||
|
||||
@functools.lru_cache(maxsize=8)
|
||||
def _h3_tile_geometry(
|
||||
prefix_segments: tuple[int, ...],
|
||||
dit_seq_shape: tuple[int, int, int],
|
||||
device: torch.device,
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, int, int]:
|
||||
"""Segment-pure prefix chunks, then video tiles.
|
||||
|
||||
Returns (variable_block_sizes int32 [n_tiles], pack_index int32 [S_pad]:
|
||||
padded position -> packed row or -1, unpack_index int32 [used]: packed row
|
||||
-> padded position, num_prefix_tiles, num_video_tiles).
|
||||
"""
|
||||
prefix_len = sum(prefix_segments)
|
||||
|
||||
prefix_sizes: list[int] = []
|
||||
for segment in prefix_segments:
|
||||
full, remainder = divmod(segment, VSA_H3_TILE_ELEMS)
|
||||
prefix_sizes.extend([VSA_H3_TILE_ELEMS] * full)
|
||||
if remainder:
|
||||
prefix_sizes.append(remainder)
|
||||
num_prefix_tiles = len(prefix_sizes)
|
||||
|
||||
num_tiles = (
|
||||
math.ceil(dit_seq_shape[0] / VSA_H3_TILE_SHAPE[0]),
|
||||
math.ceil(dit_seq_shape[1] / VSA_H3_TILE_SHAPE[1]),
|
||||
math.ceil(dit_seq_shape[2] / VSA_H3_TILE_SHAPE[2]),
|
||||
)
|
||||
video_sizes = construct_variable_block_sizes(dit_seq_shape, num_tiles, device)
|
||||
num_video_tiles = int(video_sizes.numel())
|
||||
|
||||
video_indices = (
|
||||
get_tile_partition_indices(dit_seq_shape, VSA_H3_TILE_SHAPE, device)
|
||||
+ prefix_len
|
||||
)
|
||||
tile_partition_indices = torch.cat(
|
||||
[
|
||||
torch.arange(prefix_len, device=device, dtype=torch.long),
|
||||
video_indices,
|
||||
]
|
||||
)
|
||||
variable_block_sizes = torch.cat(
|
||||
[
|
||||
torch.tensor(prefix_sizes, dtype=torch.long, device=device),
|
||||
video_sizes.to(torch.long),
|
||||
]
|
||||
)
|
||||
non_pad_index = get_non_pad_index(variable_block_sizes, VSA_H3_TILE_ELEMS)
|
||||
|
||||
total = prefix_len + math.prod(dit_seq_shape)
|
||||
sizes_sum = int(variable_block_sizes.sum())
|
||||
if sizes_sum != total or non_pad_index.numel() != total:
|
||||
raise ValueError(
|
||||
f"VSA-H3 tile geometry mismatch for prefix={prefix_segments}, "
|
||||
f"video={dit_seq_shape}: sizes sum {sizes_sum}, non-pad "
|
||||
f"{non_pad_index.numel()}, expected {total}."
|
||||
)
|
||||
|
||||
seq_pad = variable_block_sizes.numel() * VSA_H3_TILE_ELEMS
|
||||
pack_index = torch.full((seq_pad,), -1, dtype=torch.int32, device=device)
|
||||
pack_index[non_pad_index] = tile_partition_indices.to(torch.int32)
|
||||
unpack_index = torch.empty(total, dtype=torch.int32, device=device)
|
||||
unpack_index[tile_partition_indices] = non_pad_index.to(torch.int32)
|
||||
return (
|
||||
variable_block_sizes.to(torch.int32),
|
||||
pack_index,
|
||||
unpack_index,
|
||||
num_prefix_tiles,
|
||||
num_video_tiles,
|
||||
)
|
||||
|
||||
|
||||
class VideoSparseAttentionH3Backend(AttentionBackend):
|
||||
accept_output_buffer: bool = True
|
||||
|
||||
@staticmethod
|
||||
def get_supported_head_sizes() -> list[int]:
|
||||
return [64, 128]
|
||||
|
||||
@staticmethod
|
||||
def get_enum() -> AttentionBackendEnum:
|
||||
return AttentionBackendEnum.VIDEO_SPARSE_ATTN_H3
|
||||
|
||||
@staticmethod
|
||||
def get_impl_cls() -> type["VideoSparseAttentionH3Impl"]:
|
||||
return VideoSparseAttentionH3Impl
|
||||
|
||||
@staticmethod
|
||||
def get_metadata_cls() -> type["VideoSparseAttentionH3Metadata"]:
|
||||
return VideoSparseAttentionH3Metadata
|
||||
|
||||
@staticmethod
|
||||
def get_builder_cls() -> type["VideoSparseAttentionH3MetadataBuilder"]:
|
||||
return VideoSparseAttentionH3MetadataBuilder
|
||||
|
||||
|
||||
@dataclass
|
||||
class VideoSparseAttentionH3Metadata(AttentionMetadata):
|
||||
VSA_sparsity: float
|
||||
total_seq_length: int
|
||||
num_prefix_tiles: int
|
||||
num_video_tiles: int
|
||||
exempt: bool
|
||||
variable_block_sizes: torch.Tensor
|
||||
pack_index: torch.Tensor
|
||||
unpack_index: torch.Tensor
|
||||
dense_layers: tuple[int, ...] = ()
|
||||
workspace_cache: dict = field(default_factory=dict)
|
||||
|
||||
@property
|
||||
def num_tiles(self) -> int:
|
||||
return self.num_prefix_tiles + self.num_video_tiles
|
||||
|
||||
|
||||
class VideoSparseAttentionH3MetadataBuilder(AttentionMetadataBuilder):
|
||||
def __init__(self) -> None:
|
||||
self._workspace_cache: dict = {}
|
||||
|
||||
def prepare(self) -> None:
|
||||
pass
|
||||
|
||||
def build( # type: ignore[override]
|
||||
self,
|
||||
current_timestep: int,
|
||||
raw_latent_shape: tuple[int, int, int],
|
||||
patch_size: tuple[int, int, int],
|
||||
VSA_sparsity: float,
|
||||
prefix_segments: tuple[int, ...],
|
||||
device: torch.device,
|
||||
exempt: bool = True,
|
||||
dense_layers: tuple[int, ...] = (),
|
||||
dense_first_n_steps: int = 0,
|
||||
**kwargs: dict[str, Any],
|
||||
) -> VideoSparseAttentionH3Metadata:
|
||||
dit_seq_shape = (
|
||||
raw_latent_shape[0] // patch_size[0],
|
||||
raw_latent_shape[1] // patch_size[1],
|
||||
raw_latent_shape[2] // patch_size[2],
|
||||
)
|
||||
prefix_segments = tuple(int(s) for s in prefix_segments if s > 0)
|
||||
if current_timestep < dense_first_n_steps:
|
||||
VSA_sparsity = 0.0
|
||||
|
||||
(
|
||||
variable_block_sizes,
|
||||
pack_index,
|
||||
unpack_index,
|
||||
num_prefix_tiles,
|
||||
num_video_tiles,
|
||||
) = _h3_tile_geometry(prefix_segments, dit_seq_shape, device)
|
||||
|
||||
return VideoSparseAttentionH3Metadata(
|
||||
current_timestep=current_timestep,
|
||||
VSA_sparsity=float(VSA_sparsity),
|
||||
total_seq_length=sum(prefix_segments) + math.prod(dit_seq_shape),
|
||||
num_prefix_tiles=num_prefix_tiles,
|
||||
num_video_tiles=num_video_tiles,
|
||||
exempt=exempt,
|
||||
variable_block_sizes=variable_block_sizes,
|
||||
pack_index=pack_index,
|
||||
unpack_index=unpack_index,
|
||||
dense_layers=tuple(int(layer) for layer in dense_layers),
|
||||
workspace_cache=self._workspace_cache,
|
||||
)
|
||||
|
||||
|
||||
def _compute_topk(sparsity: float, num_video_tiles: int) -> int:
|
||||
keep = math.ceil((1.0 - sparsity) * num_video_tiles)
|
||||
return max(1, min(keep, num_video_tiles))
|
||||
|
||||
|
||||
def _topk_tile_lists(
|
||||
scores: torch.Tensor,
|
||||
num_prefix_tiles: int,
|
||||
num_video_tiles: int,
|
||||
sparsity: float,
|
||||
exempt: bool,
|
||||
) -> torch.Tensor:
|
||||
"""scores [H, n_tiles, n_tiles] -> ascending int32 kv-tile lists
|
||||
[H, num_video_tiles, width]; width = num_prefix_tiles + keep (exempt) or
|
||||
min(keep + num_prefix_tiles, n_tiles) (compete)."""
|
||||
prefix = num_prefix_tiles
|
||||
keep = _compute_topk(sparsity, num_video_tiles)
|
||||
video_rows = scores[:, prefix:, :]
|
||||
if exempt or prefix == 0:
|
||||
picked = video_rows[:, :, prefix:].topk(keep, dim=-1).indices + prefix
|
||||
picked = picked.sort(dim=-1).values
|
||||
if prefix == 0:
|
||||
return picked.to(torch.int32)
|
||||
prefix_cols = torch.arange(prefix, device=scores.device).expand(
|
||||
*picked.shape[:-1], prefix
|
||||
)
|
||||
return torch.cat([prefix_cols, picked], dim=-1).to(torch.int32)
|
||||
keep_total = min(keep + prefix, scores.shape[-1])
|
||||
return (
|
||||
video_rows.topk(keep_total, dim=-1).indices.sort(dim=-1).values.to(torch.int32)
|
||||
)
|
||||
|
||||
|
||||
def _workspace_key(
|
||||
meta: VideoSparseAttentionH3Metadata,
|
||||
heads: int,
|
||||
head_dim: int,
|
||||
has_gate: bool,
|
||||
dtype: torch.dtype,
|
||||
device: torch.device,
|
||||
) -> tuple:
|
||||
return (
|
||||
meta.num_tiles,
|
||||
meta.num_prefix_tiles,
|
||||
meta.exempt,
|
||||
heads,
|
||||
head_dim,
|
||||
has_gate,
|
||||
dtype,
|
||||
device,
|
||||
)
|
||||
|
||||
|
||||
class _Workspace:
|
||||
"""Per-geometry scratch: tiled q/k/v(/gate) [3|4, H, S_pad, D], pooled fp32
|
||||
tile means [3, H, n_tiles, D], and the kernel index lists (prefix rows and
|
||||
prefix columns are static; only the top-k video columns change per layer).
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
meta: VideoSparseAttentionH3Metadata,
|
||||
heads: int,
|
||||
head_dim: int,
|
||||
has_gate: bool,
|
||||
dtype: torch.dtype,
|
||||
device: torch.device,
|
||||
) -> None:
|
||||
n_tiles = meta.num_tiles
|
||||
seq_pad = n_tiles * VSA_H3_TILE_ELEMS
|
||||
self.key = _workspace_key(meta, heads, head_dim, has_gate, dtype, device)
|
||||
self.tiled = torch.empty(
|
||||
(3 + int(has_gate), heads, seq_pad, head_dim), dtype=dtype, device=device
|
||||
)
|
||||
self.pooled = torch.empty(
|
||||
(3, heads, n_tiles, head_dim), dtype=torch.float32, device=device
|
||||
)
|
||||
self.out_tiled = torch.empty(
|
||||
(heads, seq_pad, head_dim), dtype=dtype, device=device
|
||||
)
|
||||
all_tiles = torch.arange(n_tiles, dtype=torch.int32, device=device)
|
||||
self.dense_index = all_tiles.repeat(heads, n_tiles, 1)
|
||||
self.dense_num = torch.full(
|
||||
(heads, n_tiles), n_tiles, dtype=torch.int32, device=device
|
||||
)
|
||||
self.q2k_index = self.dense_index.clone()
|
||||
self.q2k_num = self.dense_num.clone()
|
||||
|
||||
def sparse_lists(
|
||||
self, video_lists: torch.Tensor, num_prefix_tiles: int
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
width = video_lists.shape[-1]
|
||||
self.q2k_index[:, num_prefix_tiles:, :width] = video_lists
|
||||
self.q2k_num[:, num_prefix_tiles:] = width
|
||||
return self.q2k_index, self.q2k_num
|
||||
|
||||
|
||||
def _get_workspace(
|
||||
meta: VideoSparseAttentionH3Metadata, query: torch.Tensor, has_gate: bool
|
||||
) -> _Workspace:
|
||||
heads, head_dim = query.shape[-2], query.shape[-1]
|
||||
key = _workspace_key(meta, heads, head_dim, has_gate, query.dtype, query.device)
|
||||
workspace = meta.workspace_cache.get("workspace")
|
||||
if workspace is None or workspace.key != key:
|
||||
workspace = _Workspace(
|
||||
meta, heads, head_dim, has_gate, query.dtype, query.device
|
||||
)
|
||||
meta.workspace_cache["workspace"] = workspace
|
||||
return workspace
|
||||
|
||||
|
||||
def _select_kv_lists(
|
||||
ws: _Workspace,
|
||||
meta: VideoSparseAttentionH3Metadata,
|
||||
scores: torch.Tensor | None,
|
||||
sparsity: float,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Kernel index lists for this layer: dense, or top-k video columns."""
|
||||
keep = _compute_topk(sparsity, meta.num_video_tiles)
|
||||
if sparsity <= 0.0 or keep >= meta.num_video_tiles:
|
||||
return ws.dense_index, ws.dense_num
|
||||
return ws.sparse_lists(
|
||||
_topk_tile_lists(
|
||||
scores,
|
||||
meta.num_prefix_tiles,
|
||||
meta.num_video_tiles,
|
||||
sparsity,
|
||||
meta.exempt,
|
||||
),
|
||||
meta.num_prefix_tiles,
|
||||
)
|
||||
|
||||
|
||||
class VideoSparseAttentionH3Impl(AttentionImpl):
|
||||
def __init__(
|
||||
self,
|
||||
num_heads: int,
|
||||
head_size: int,
|
||||
causal: bool,
|
||||
softmax_scale: float,
|
||||
num_kv_heads: int | None = None,
|
||||
prefix: str = "",
|
||||
**extra_impl_args,
|
||||
) -> None:
|
||||
self.num_heads = num_heads
|
||||
self.head_size = head_size
|
||||
self.softmax_scale = softmax_scale
|
||||
self.prefix = prefix
|
||||
match = _DIT_BLOCK_PREFIX.match(prefix)
|
||||
self.layer_idx = int(match.group(1)) if match else None
|
||||
# The token refiner and any other non-packed caller resolve the same
|
||||
# backend object; they run the exact dense kernel instead.
|
||||
from sglang.multimodal_gen.runtime.layers.attention.backends.flash_attn import (
|
||||
FlashAttentionImpl,
|
||||
)
|
||||
|
||||
self._dense_fallback = FlashAttentionImpl(
|
||||
num_heads=num_heads,
|
||||
head_size=head_size,
|
||||
causal=causal,
|
||||
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: AttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
raise NotImplementedError(
|
||||
"VSA-H3 serves MiniMax-H3's packed varlen attention; use " "forward_varlen."
|
||||
)
|
||||
|
||||
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,
|
||||
attn_metadata: VideoSparseAttentionH3Metadata | None = None,
|
||||
gate_compress: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""query/key/value: [T, H, D] packed rows (post-norm, post-RoPE)."""
|
||||
if self.layer_idx is None or attn_metadata is None:
|
||||
if attn_metadata is None and self.layer_idx is not None:
|
||||
raise RuntimeError(
|
||||
"VSA-H3 needs per-step attention metadata from the "
|
||||
"MiniMax-H3 denoising stage; none was set in the forward "
|
||||
"context."
|
||||
)
|
||||
return self._dense_fallback.forward_varlen(
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
cu_seqlens=cu_seqlens,
|
||||
max_seqlen=max_seqlen,
|
||||
cu_seqlens_host=cu_seqlens_host,
|
||||
)
|
||||
|
||||
meta = attn_metadata
|
||||
bounds = (
|
||||
cu_seqlens_host
|
||||
if cu_seqlens_host is not None
|
||||
else tuple(int(item) for item in cu_seqlens.tolist())
|
||||
)
|
||||
used = int(bounds[1])
|
||||
if used != meta.total_seq_length:
|
||||
raise ValueError(
|
||||
f"VSA-H3 metadata was built for {meta.total_seq_length} packed "
|
||||
f"rows, got {used}. The step metadata and the packed sequence "
|
||||
"layout have diverged."
|
||||
)
|
||||
|
||||
sparsity = 0.0 if self.layer_idx in meta.dense_layers else meta.VSA_sparsity
|
||||
has_gate = gate_compress is not None
|
||||
ws = _get_workspace(meta, query, has_gate)
|
||||
|
||||
vsa_h3_pack_tiles(
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
gate_compress,
|
||||
meta.pack_index,
|
||||
meta.variable_block_sizes,
|
||||
ws.tiled,
|
||||
ws.pooled,
|
||||
)
|
||||
q_pooled, k_pooled, v_pooled = ws.pooled
|
||||
|
||||
scores = None
|
||||
if sparsity > 0.0 or has_gate:
|
||||
scores = torch.matmul(q_pooled, k_pooled.transpose(-2, -1)) * (
|
||||
self.head_size**-0.5
|
||||
)
|
||||
|
||||
q2k_index, q2k_num = _select_kv_lists(ws, meta, scores, sparsity)
|
||||
vsa_h3_block_sparse_attn_forward(
|
||||
ws.tiled[0:1],
|
||||
ws.tiled[1:2],
|
||||
ws.tiled[2:3],
|
||||
q2k_index[None],
|
||||
q2k_num[None],
|
||||
meta.variable_block_sizes,
|
||||
out=ws.out_tiled[None],
|
||||
)
|
||||
|
||||
out_compress = None
|
||||
if has_gate:
|
||||
out_compress = torch.matmul(torch.softmax(scores, dim=-1), v_pooled)
|
||||
|
||||
result = torch.empty(query.shape, dtype=query.dtype, device=query.device)
|
||||
vsa_h3_untile(
|
||||
ws.out_tiled,
|
||||
ws.tiled[3] if has_gate else None,
|
||||
out_compress,
|
||||
meta.unpack_index,
|
||||
used,
|
||||
result,
|
||||
)
|
||||
return result
|
||||
@@ -0,0 +1,343 @@
|
||||
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
|
||||
# (fastvideo-kernel triton_kernels/block_sparse_attn_triton.py).
|
||||
# Inference-only subset: the block-sparse forward. Backward stays upstream;
|
||||
# SGLang serves no-grad forwards. The tile pack/unpack kernels are SGLang's.
|
||||
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Triton 64-token block-sparse attention for MiniMax-H3 VSA.
|
||||
|
||||
The attention kernel consumes an explicit per-query-block index list
|
||||
(``q2k_index`` / ``q2k_num``) plus per-key-block valid token counts
|
||||
(``variable_block_sizes``), so ragged interior tiles - segment-pure prefix
|
||||
chunks and 3D video tiles whose dimensions do not divide the tile shape - mask
|
||||
their pad columns exactly.
|
||||
|
||||
``vsa_h3_pack_tiles`` gathers packed ``[T, H, D]`` rows into the head-major
|
||||
padded tile layout the attention kernel reads and pools each tile in the same
|
||||
pass; ``vsa_h3_untile`` scatters the attention output back to packed rows and
|
||||
folds in the gated compression branch. Both replace chains of index copies
|
||||
and transposes that otherwise cost more than the attention kernel itself.
|
||||
"""
|
||||
|
||||
import math
|
||||
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
from triton.tools.tensor_descriptor import TensorDescriptor
|
||||
|
||||
# BLOCK_M / BLOCK_N are structural, not tunable: the kernel indexes the top-k
|
||||
# list per BLOCK_M q-tile and addresses keys as kv_idx * BLOCK_N, so both must
|
||||
# match the granularity q2k_index and variable_block_sizes were built at.
|
||||
VSA_H3_KERNEL_BLOCK = 64
|
||||
|
||||
# Pinned instead of autotuned: on B300 num_warps=4 wins at every sequence
|
||||
# length and num_stages=5 is within 0.5% of the best (7 spills); FastVideo
|
||||
# reports the same optimum for Blackwell.
|
||||
_ATTN_NUM_WARPS = 4
|
||||
_ATTN_NUM_STAGES = 5
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _attn_fwd_sparse(
|
||||
desc_q,
|
||||
desc_k,
|
||||
desc_v,
|
||||
desc_o,
|
||||
sm_scale,
|
||||
q2k_index,
|
||||
q2k_num,
|
||||
max_kv_blks,
|
||||
variable_block_sizes,
|
||||
H,
|
||||
N_CTX_Q,
|
||||
HEAD_DIM: tl.constexpr,
|
||||
BLOCK_M: tl.constexpr,
|
||||
BLOCK_N: tl.constexpr,
|
||||
):
|
||||
q_blk = tl.program_id(0)
|
||||
off_hz = tl.program_id(1)
|
||||
b = off_hz // H
|
||||
h = off_hz % H
|
||||
q_tiles = N_CTX_Q // BLOCK_M
|
||||
meta_base = off_hz * q_tiles + q_blk
|
||||
|
||||
kv_blocks = tl.load(q2k_num + meta_base)
|
||||
kv_ptr = q2k_index + meta_base.to(tl.int64) * max_kv_blks
|
||||
|
||||
q = desc_q.load([b, h, q_blk * BLOCK_M, 0]).reshape([BLOCK_M, HEAD_DIM])
|
||||
m_i = tl.full([BLOCK_M], -float("inf"), tl.float32)
|
||||
l_i = tl.zeros([BLOCK_M], dtype=tl.float32) + 1.0
|
||||
acc = tl.zeros([BLOCK_M, HEAD_DIM], dtype=tl.float32)
|
||||
qk_scale = sm_scale * 1.44269504 # 1/ln2
|
||||
|
||||
for i in range(0, kv_blocks):
|
||||
kv_idx = tl.load(kv_ptr + i).to(tl.int32)
|
||||
block_size = tl.load(variable_block_sizes + kv_idx)
|
||||
k = desc_k.load([b, h, kv_idx * BLOCK_N, 0]).reshape([BLOCK_N, HEAD_DIM])
|
||||
qk = tl.dot(q, tl.trans(k))
|
||||
mask = tl.arange(0, BLOCK_N) < block_size
|
||||
qk = tl.where(mask[None, :], qk, -float("inf"))
|
||||
|
||||
m_ij = tl.maximum(m_i, tl.max(qk, 1) * qk_scale)
|
||||
p = tl.math.exp2(qk * qk_scale - m_ij[:, None])
|
||||
l_ij = tl.sum(p, 1)
|
||||
|
||||
alpha = tl.math.exp2(m_i - m_ij)
|
||||
l_i = l_i * alpha + l_ij
|
||||
acc = acc * alpha[:, None]
|
||||
|
||||
v = desc_v.load([b, h, kv_idx * BLOCK_N, 0]).reshape([BLOCK_N, HEAD_DIM])
|
||||
acc = tl.dot(p.to(tl.bfloat16), v, acc)
|
||||
m_i = m_ij
|
||||
|
||||
acc = acc / l_i[:, None]
|
||||
desc_o.store(
|
||||
[b, h, q_blk * BLOCK_M, 0],
|
||||
acc.to(desc_o.dtype).reshape([1, 1, BLOCK_M, HEAD_DIM]),
|
||||
)
|
||||
|
||||
|
||||
def vsa_h3_block_sparse_attn_forward(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
q2k_index: torch.Tensor,
|
||||
q2k_num: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
out: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""q/k/v: contiguous [B, H, S_pad, D] bf16 with S_pad = n_tiles * 64; pad
|
||||
rows zero. q2k_index/q2k_num: contiguous [B, H, n_tiles, max_kv] /
|
||||
[B, H, n_tiles] int32."""
|
||||
batch, heads, seq_q, head_dim = q.shape
|
||||
seq_kv = k.shape[2]
|
||||
if seq_q % VSA_H3_KERNEL_BLOCK or seq_kv % VSA_H3_KERNEL_BLOCK:
|
||||
raise ValueError(
|
||||
f"VSA-H3 kernel needs 64-multiple sequence lengths, got q={seq_q}, "
|
||||
f"kv={seq_kv}"
|
||||
)
|
||||
if variable_block_sizes.numel() != seq_kv // VSA_H3_KERNEL_BLOCK:
|
||||
raise ValueError(
|
||||
"variable_block_sizes must have one entry per 64-token key block: "
|
||||
f"{variable_block_sizes.numel()} vs {seq_kv // VSA_H3_KERNEL_BLOCK}"
|
||||
)
|
||||
if out is None:
|
||||
out = torch.empty_like(q)
|
||||
block = [1, 1, VSA_H3_KERNEL_BLOCK, head_dim]
|
||||
desc_q, desc_k, desc_v, desc_o = (
|
||||
TensorDescriptor.from_tensor(t, block_shape=block) for t in (q, k, v, out)
|
||||
)
|
||||
grid = (seq_q // VSA_H3_KERNEL_BLOCK, batch * heads, 1)
|
||||
_attn_fwd_sparse[grid](
|
||||
desc_q,
|
||||
desc_k,
|
||||
desc_v,
|
||||
desc_o,
|
||||
1.0 / math.sqrt(head_dim),
|
||||
q2k_index,
|
||||
q2k_num,
|
||||
q2k_index.shape[-1],
|
||||
variable_block_sizes,
|
||||
heads,
|
||||
seq_q,
|
||||
HEAD_DIM=head_dim,
|
||||
BLOCK_M=VSA_H3_KERNEL_BLOCK,
|
||||
BLOCK_N=VSA_H3_KERNEL_BLOCK,
|
||||
num_warps=_ATTN_NUM_WARPS,
|
||||
num_stages=_ATTN_NUM_STAGES,
|
||||
)
|
||||
return out
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _pack_tiles_kernel(
|
||||
Q,
|
||||
K,
|
||||
V,
|
||||
G,
|
||||
src_index,
|
||||
variable_block_sizes,
|
||||
Tiled,
|
||||
Pooled,
|
||||
stride_q_row,
|
||||
stride_q_head,
|
||||
stride_k_row,
|
||||
stride_k_head,
|
||||
stride_v_row,
|
||||
stride_v_head,
|
||||
stride_g_row,
|
||||
stride_g_head,
|
||||
H,
|
||||
S_PAD,
|
||||
N_TILES,
|
||||
HAS_GATE: tl.constexpr,
|
||||
HEAD_DIM: tl.constexpr,
|
||||
BLOCK: tl.constexpr,
|
||||
):
|
||||
tile = tl.program_id(0)
|
||||
h = tl.program_id(1)
|
||||
rows = tile * BLOCK + tl.arange(0, BLOCK)
|
||||
cols = tl.arange(0, HEAD_DIM)
|
||||
src = tl.load(src_index + rows)
|
||||
valid = src >= 0
|
||||
src = tl.where(valid, src, 0).to(tl.int64)
|
||||
mask = valid[:, None]
|
||||
size = tl.load(variable_block_sizes + tile).to(tl.float32)
|
||||
|
||||
tensor_stride = H.to(tl.int64) * S_PAD * HEAD_DIM
|
||||
out_off = (
|
||||
h.to(tl.int64) * S_PAD * HEAD_DIM + rows[:, None] * HEAD_DIM + cols[None, :]
|
||||
)
|
||||
pool_off = (h * N_TILES + tile).to(tl.int64) * HEAD_DIM + cols
|
||||
|
||||
x = tl.load(
|
||||
Q + src[:, None] * stride_q_row + h * stride_q_head + cols[None, :],
|
||||
mask=mask,
|
||||
other=0.0,
|
||||
)
|
||||
tl.store(Tiled + out_off, x)
|
||||
tl.store(Pooled + pool_off, tl.sum(x.to(tl.float32), 0) / size)
|
||||
|
||||
x = tl.load(
|
||||
K + src[:, None] * stride_k_row + h * stride_k_head + cols[None, :],
|
||||
mask=mask,
|
||||
other=0.0,
|
||||
)
|
||||
tl.store(Tiled + tensor_stride + out_off, x)
|
||||
tl.store(
|
||||
Pooled + H * N_TILES * HEAD_DIM + pool_off,
|
||||
tl.sum(x.to(tl.float32), 0) / size,
|
||||
)
|
||||
|
||||
x = tl.load(
|
||||
V + src[:, None] * stride_v_row + h * stride_v_head + cols[None, :],
|
||||
mask=mask,
|
||||
other=0.0,
|
||||
)
|
||||
tl.store(Tiled + 2 * tensor_stride + out_off, x)
|
||||
tl.store(
|
||||
Pooled + 2 * H * N_TILES * HEAD_DIM + pool_off,
|
||||
tl.sum(x.to(tl.float32), 0) / size,
|
||||
)
|
||||
|
||||
if HAS_GATE:
|
||||
x = tl.load(
|
||||
G + src[:, None] * stride_g_row + h * stride_g_head + cols[None, :],
|
||||
mask=mask,
|
||||
other=0.0,
|
||||
)
|
||||
tl.store(Tiled + 3 * tensor_stride + out_off, x)
|
||||
|
||||
|
||||
def vsa_h3_pack_tiles(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
gate: torch.Tensor | None,
|
||||
src_index: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
tiled: torch.Tensor,
|
||||
pooled: torch.Tensor,
|
||||
) -> None:
|
||||
"""Gather packed [T, H, D] rows into ``tiled`` [3|4, H, S_pad, D] and write
|
||||
fp32 per-tile means of q/k/v into ``pooled`` [3, H, n_tiles, D].
|
||||
``src_index`` maps each padded position to its packed row, or -1 (pad -> 0).
|
||||
"""
|
||||
_, heads, seq_pad, head_dim = tiled.shape
|
||||
n_tiles = seq_pad // VSA_H3_KERNEL_BLOCK
|
||||
has_gate = gate is not None
|
||||
g = gate if has_gate else q
|
||||
assert all(t.stride(-1) == 1 for t in (q, k, v, g))
|
||||
_pack_tiles_kernel[(n_tiles, heads)](
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
g,
|
||||
src_index,
|
||||
variable_block_sizes,
|
||||
tiled,
|
||||
pooled,
|
||||
q.stride(0),
|
||||
q.stride(1),
|
||||
k.stride(0),
|
||||
k.stride(1),
|
||||
v.stride(0),
|
||||
v.stride(1),
|
||||
g.stride(0),
|
||||
g.stride(1),
|
||||
heads,
|
||||
seq_pad,
|
||||
n_tiles,
|
||||
HAS_GATE=has_gate,
|
||||
HEAD_DIM=head_dim,
|
||||
BLOCK=VSA_H3_KERNEL_BLOCK,
|
||||
)
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _untile_kernel(
|
||||
OutTiled,
|
||||
Gate,
|
||||
OutC,
|
||||
dst_index,
|
||||
Res,
|
||||
used,
|
||||
total,
|
||||
S_PAD,
|
||||
N_TILES,
|
||||
HAS_GATE: tl.constexpr,
|
||||
HEAD_DIM: tl.constexpr,
|
||||
BLOCK: tl.constexpr,
|
||||
):
|
||||
row_block = tl.program_id(0)
|
||||
h = tl.program_id(1)
|
||||
H = tl.num_programs(1)
|
||||
rows = row_block * BLOCK + tl.arange(0, BLOCK)
|
||||
cols = tl.arange(0, HEAD_DIM)
|
||||
in_used = rows < used
|
||||
pos = tl.load(dst_index + rows, mask=in_used, other=0).to(tl.int64)
|
||||
head_base = h.to(tl.int64) * S_PAD * HEAD_DIM
|
||||
off = head_base + pos[:, None] * HEAD_DIM + cols[None, :]
|
||||
o = tl.load(OutTiled + off, mask=in_used[:, None], other=0.0).to(tl.float32)
|
||||
if HAS_GATE:
|
||||
g = tl.load(Gate + off, mask=in_used[:, None], other=0.0).to(tl.float32)
|
||||
tile = pos // BLOCK
|
||||
c = tl.load(
|
||||
OutC + (h * N_TILES + tile)[:, None] * HEAD_DIM + cols[None, :],
|
||||
mask=in_used[:, None],
|
||||
other=0.0,
|
||||
)
|
||||
o = o + c * g
|
||||
res_off = (rows[:, None] * H + h).to(tl.int64) * HEAD_DIM + cols[None, :]
|
||||
tl.store(Res + res_off, o.to(Res.type.element_ty), mask=(rows < total)[:, None])
|
||||
|
||||
|
||||
def vsa_h3_untile(
|
||||
out_tiled: torch.Tensor,
|
||||
gate_tiled: torch.Tensor | None,
|
||||
out_compress: torch.Tensor | None,
|
||||
dst_index: torch.Tensor,
|
||||
used: int,
|
||||
result: torch.Tensor,
|
||||
) -> None:
|
||||
"""Scatter ``out_tiled`` [H, S_pad, D] to packed rows of ``result`` [T, H, D]
|
||||
(rows past ``used`` are zero), adding ``out_compress`` [H, n_tiles, D] fp32
|
||||
scaled by ``gate_tiled`` when given."""
|
||||
heads, seq_pad, head_dim = out_tiled.shape
|
||||
total = result.shape[0]
|
||||
has_gate = gate_tiled is not None
|
||||
_untile_kernel[(triton.cdiv(total, VSA_H3_KERNEL_BLOCK), heads)](
|
||||
out_tiled,
|
||||
gate_tiled if has_gate else out_tiled,
|
||||
out_compress if has_gate else out_tiled,
|
||||
dst_index,
|
||||
result,
|
||||
used,
|
||||
total,
|
||||
seq_pad,
|
||||
seq_pad // VSA_H3_KERNEL_BLOCK,
|
||||
HAS_GATE=has_gate,
|
||||
HEAD_DIM=head_dim,
|
||||
BLOCK=VSA_H3_KERNEL_BLOCK,
|
||||
)
|
||||
@@ -67,6 +67,7 @@ from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config impor
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.layers.usp import _ring_attention_varlen
|
||||
from sglang.multimodal_gen.runtime.loader.utils import get_param_names_mapping
|
||||
from sglang.multimodal_gen.runtime.managers.forward_context import get_forward_context
|
||||
from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import (
|
||||
LayerwiseOffloadableModuleMixin,
|
||||
is_layerwise_offloaded_module,
|
||||
@@ -86,6 +87,18 @@ logger = init_logger(__name__)
|
||||
|
||||
_ARCH_DEFAULTS = MiniMaxH3DiTArchConfig()
|
||||
|
||||
_NON_LORA_DELTA_SUFFIXES = (".diff", ".diff_b", ".set_weight")
|
||||
|
||||
|
||||
def _reject_non_lora_delta_tensors(adapter: dict[str, torch.Tensor]) -> None:
|
||||
offending = sorted(key for key in adapter if key.endswith(_NON_LORA_DELTA_SUFFIXES))
|
||||
if offending:
|
||||
raise ValueError(
|
||||
f"LoRA adapter carries {len(offending)} non-LoRA tensors "
|
||||
f"(.diff/.diff_b/.set_weight, e.g. {offending[0]}) that no MiniMax-H3 "
|
||||
"LoRA mapping rule applies; serve a checkpoint with them merged instead."
|
||||
)
|
||||
|
||||
|
||||
def _diffusers_h3_checkpoint(
|
||||
iterator: Iterable[tuple[str, torch.Tensor]],
|
||||
@@ -592,6 +605,7 @@ def _minimax_h3_attention_core_impl(
|
||||
ulysses_active: bool,
|
||||
subblock_sparse_query_block_mask: torch.Tensor | None = None,
|
||||
ring_active: bool = False,
|
||||
gate_compress: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""Dynamic varlen attention and Ulysses/Ring collectives.
|
||||
|
||||
@@ -602,11 +616,14 @@ def _minimax_h3_attention_core_impl(
|
||||
|
||||
if ulysses_active:
|
||||
from sglang.multimodal_gen.runtime.layers.usp import (
|
||||
_usp_input_all_to_all,
|
||||
_usp_input_all_to_all_packed_qkv,
|
||||
_usp_output_all_to_all,
|
||||
)
|
||||
|
||||
q, k, v = _usp_input_all_to_all_packed_qkv(q, k, v)
|
||||
if gate_compress is not None:
|
||||
gate_compress = _usp_input_all_to_all(gate_compress[None], head_dim=2)[0]
|
||||
|
||||
if attention._attention_impl is None:
|
||||
attention._set_attention_backend(
|
||||
@@ -618,6 +635,26 @@ def _minimax_h3_attention_core_impl(
|
||||
)
|
||||
)
|
||||
|
||||
if attention._attention_backend_enum is AttentionBackendEnum.VIDEO_SPARSE_ATTN_H3:
|
||||
attn_metadata = (
|
||||
get_forward_context().attn_metadata
|
||||
if attention.prefix.startswith("blocks.")
|
||||
else None
|
||||
)
|
||||
out = attention._attention_impl.forward_varlen(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
cu_seqlens=cu_seqlens,
|
||||
max_seqlen=max_seqlen,
|
||||
cu_seqlens_host=cu_seqlens_host,
|
||||
attn_metadata=attn_metadata,
|
||||
gate_compress=gate_compress,
|
||||
)
|
||||
if ulysses_active:
|
||||
out = _usp_output_all_to_all(out[None], head_dim=2)[0]
|
||||
return out
|
||||
|
||||
if ring_active:
|
||||
ring_ws, _ = get_ring_ctx()
|
||||
if attention._attention_backend_enum is not AttentionBackendEnum.FA:
|
||||
@@ -767,6 +804,18 @@ class MiniMaxH3Attention(nn.Module):
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.out_proj",
|
||||
)
|
||||
# VSA compression gate; stays bf16 and unquantized (zero gate == pure sparse).
|
||||
self.to_gate_compress: ColumnParallelLinear | None = None
|
||||
if arch.has_gate_compress and prefix.startswith("blocks."):
|
||||
self.to_gate_compress = ColumnParallelLinear(
|
||||
arch.hidden_size,
|
||||
self.inner_dim,
|
||||
bias=False,
|
||||
gather_output=False,
|
||||
params_dtype=_BF16_DTYPE,
|
||||
quant_config=None,
|
||||
prefix=f"{prefix}.to_gate_compress",
|
||||
)
|
||||
|
||||
def _set_attention_backend(self, backend) -> None:
|
||||
if (
|
||||
@@ -1032,6 +1081,14 @@ class MiniMaxH3Attention(nn.Module):
|
||||
)
|
||||
q, k = _apply_rope_qk(q, k, cos_sin_cache, positions)
|
||||
|
||||
gate_compress = None
|
||||
if (
|
||||
self._attention_backend_enum is AttentionBackendEnum.VIDEO_SPARSE_ATTN_H3
|
||||
and self.to_gate_compress is not None
|
||||
):
|
||||
gate_flat, _ = self.to_gate_compress(x)
|
||||
gate_compress = gate_flat.view(total, self.num_heads, self.head_dim)
|
||||
|
||||
attention_core = (
|
||||
_minimax_h3_attention_core_bcg
|
||||
if self.bcg_breakpoint
|
||||
@@ -1048,6 +1105,7 @@ class MiniMaxH3Attention(nn.Module):
|
||||
subblock_sparse_query_block_mask=subblock_sparse_query_block_mask,
|
||||
ulysses_active=ulysses_active,
|
||||
ring_active=ring_active,
|
||||
gate_compress=gate_compress,
|
||||
)
|
||||
out = out.reshape(total, self.num_heads * self.head_dim)
|
||||
out, _ = self.out_proj(out)
|
||||
@@ -1780,6 +1838,7 @@ class MiniMaxH3DiTModel(BaseDiT, LayerwiseOffloadableModuleMixin):
|
||||
self, adapter: dict[str, torch.Tensor]
|
||||
) -> dict[str, torch.Tensor]:
|
||||
"""Project released-checkpoint AdaLN LoRAs onto pruned coordinates."""
|
||||
_reject_non_lora_delta_tensors(adapter)
|
||||
full_width = self.arch.adaln_affine_input_dim
|
||||
if full_width is None:
|
||||
return adapter
|
||||
@@ -2102,6 +2161,16 @@ class MiniMaxH3DiTModel(BaseDiT, LayerwiseOffloadableModuleMixin):
|
||||
get_global_forced_attn_backend()
|
||||
or self._component_attention_backend_override
|
||||
)
|
||||
if selected_backend is None:
|
||||
selected_backend = next(
|
||||
(
|
||||
module._selected_attention_backend
|
||||
for module in self.modules()
|
||||
if isinstance(module, MiniMaxH3Attention)
|
||||
and module._selected_attention_backend is not None
|
||||
),
|
||||
None,
|
||||
)
|
||||
backend = get_attn_backend(
|
||||
self.arch.attention_head_dim,
|
||||
_BF16_DTYPE,
|
||||
|
||||
@@ -170,4 +170,17 @@ class MiniMaxH3Pipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
)
|
||||
|
||||
|
||||
EntryClass = MiniMaxH3Pipeline
|
||||
class FastH3Pipeline(MiniMaxH3Pipeline):
|
||||
"""FastH3: 4-step DMD2-distilled MiniMax-H3 (t2va only).
|
||||
|
||||
The flat single-partition repo is materialized into the base-H3 layout by
|
||||
the bundled model overlay (see model_overlays/), so every stage, loader,
|
||||
and admission path below is exactly the MiniMax-H3 one. There is no
|
||||
FL2VA/Ref2VA partition layout to default into.
|
||||
"""
|
||||
|
||||
pipeline_name = "FastH3Pipeline"
|
||||
default_model_subfolder = None
|
||||
|
||||
|
||||
EntryClass = [MiniMaxH3Pipeline, FastH3Pipeline]
|
||||
|
||||
+86
-3
@@ -5,7 +5,7 @@ single-branch execution, and payload validation.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping
|
||||
from collections.abc import Callable, Mapping
|
||||
from contextlib import contextmanager
|
||||
from functools import partial
|
||||
from typing import Any
|
||||
@@ -39,7 +39,10 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
|
||||
VerificationResult,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||
from sglang.multimodal_gen.runtime.platforms import (
|
||||
AttentionBackendEnum,
|
||||
current_platform,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.runtime.utils.nvtx_pytorch_hooks import maybe_nvtx_range
|
||||
@@ -731,6 +734,13 @@ class MiniMaxH3DenoisingStage(DenoisingStage):
|
||||
device,
|
||||
placement_managed=placement_managed,
|
||||
)
|
||||
build_vsa_h3_step_metadata = _maybe_prepare_vsa_h3_step_metadata(
|
||||
model=model,
|
||||
packed=packed,
|
||||
ctx=ctx,
|
||||
server_args=server_args,
|
||||
device=device,
|
||||
)
|
||||
positive = MiniMaxH3DenoiseBranch(
|
||||
packed=packed,
|
||||
text_embeddings=emb["hidden_states"],
|
||||
@@ -770,6 +780,7 @@ class MiniMaxH3DenoisingStage(DenoisingStage):
|
||||
self._forward_dit,
|
||||
batch=batch,
|
||||
attn_metadata=attn_metadata,
|
||||
build_vsa_h3_step_metadata=build_vsa_h3_step_metadata,
|
||||
),
|
||||
positive=positive,
|
||||
initial_video_rows=initial_video,
|
||||
@@ -823,6 +834,7 @@ class MiniMaxH3DenoisingStage(DenoisingStage):
|
||||
*,
|
||||
batch: Req,
|
||||
attn_metadata: CubeSparseAttentionMetadata | None = None,
|
||||
build_vsa_h3_step_metadata: Callable[[int], Any] | None = None,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Route the custom full loop through the native denoising runner."""
|
||||
|
||||
@@ -832,7 +844,11 @@ class MiniMaxH3DenoisingStage(DenoisingStage):
|
||||
|
||||
with set_forward_context(
|
||||
current_timestep=step_index,
|
||||
attn_metadata=attn_metadata,
|
||||
attn_metadata=(
|
||||
build_vsa_h3_step_metadata(step_index)
|
||||
if build_vsa_h3_step_metadata is not None
|
||||
else attn_metadata
|
||||
),
|
||||
forward_batch=batch,
|
||||
):
|
||||
runner = self._maybe_get_bcg_runner(model)
|
||||
@@ -964,6 +980,73 @@ def _assemble_condition_rows(ctx: _FullLoopContext) -> None:
|
||||
ctx.keyframe_frame_count = int(ctx.keyframe["frame_count"])
|
||||
|
||||
|
||||
def _maybe_prepare_vsa_h3_step_metadata(
|
||||
*,
|
||||
model: Any,
|
||||
packed: Mapping[str, torch.Tensor],
|
||||
ctx: _FullLoopContext,
|
||||
server_args: ServerArgs,
|
||||
device: torch.device,
|
||||
) -> Callable[[int], Any] | None:
|
||||
"""Per-step VSA-H3 metadata builder over the request-static packed layout,
|
||||
or None off the VSA path."""
|
||||
model._resolve_attention_backend_once()
|
||||
if (
|
||||
model._resolved_attention_backend
|
||||
is not AttentionBackendEnum.VIDEO_SPARSE_ATTN_H3
|
||||
):
|
||||
return None
|
||||
if ctx.is_ref2va:
|
||||
raise NotImplementedError(
|
||||
"VSA-H3 supports the t2va/fl2va packed layout; the ref2va "
|
||||
"reference-block layout is not tiled yet. Use --attention-backend "
|
||||
"fa for ref2va."
|
||||
)
|
||||
|
||||
config = server_args.attention_backend_config or {}
|
||||
tile_size = int(config.get("vsa_tile_size", 64))
|
||||
if tile_size != 64:
|
||||
raise ValueError(
|
||||
"VSA-H3 in SGLang serves the trained 64-token (4, 4, 4) tile "
|
||||
f"geometry; got vsa_tile_size={tile_size}."
|
||||
)
|
||||
sparsity = float(config.get("VSA_sparsity", config.get("sparsity", 0.9)))
|
||||
if not 0.0 <= sparsity < 1.0:
|
||||
raise ValueError(f"VSA sparsity must be in [0, 1), got {sparsity}")
|
||||
mode = str(config.get("vsa_mode", "exempt"))
|
||||
if mode not in ("exempt", "compete"):
|
||||
raise ValueError(f"vsa_mode must be 'exempt' or 'compete', got {mode!r}")
|
||||
dense_first_n_steps = int(config.get("vsa_dense_first_n_steps", 0))
|
||||
dense_layers = tuple(int(layer) for layer in config.get("vsa_dense_layers", ()))
|
||||
|
||||
text_len = int(packed["text_pos"].numel())
|
||||
video_rows = int(packed["update_mask"].sum())
|
||||
cond_rows = int(packed["img_pos"].numel()) - video_rows
|
||||
audio_rows = int(packed["audio_pos"].numel())
|
||||
patch_size = server_args.pipeline_config.dit_config.arch_config.patch_size
|
||||
|
||||
from sglang.multimodal_gen.runtime.layers.attention.backends.video_sparse_attn_h3 import (
|
||||
VideoSparseAttentionH3MetadataBuilder,
|
||||
)
|
||||
|
||||
builder = VideoSparseAttentionH3MetadataBuilder()
|
||||
|
||||
def build(step_index: int):
|
||||
return builder.build(
|
||||
current_timestep=step_index,
|
||||
raw_latent_shape=(ctx.latent_t, ctx.latent_h, ctx.latent_w),
|
||||
patch_size=patch_size,
|
||||
VSA_sparsity=sparsity,
|
||||
prefix_segments=(text_len, cond_rows, audio_rows),
|
||||
device=device,
|
||||
exempt=mode == "exempt",
|
||||
dense_layers=dense_layers,
|
||||
dense_first_n_steps=dense_first_n_steps,
|
||||
)
|
||||
|
||||
return build
|
||||
|
||||
|
||||
def _build_packed_layout(
|
||||
ctx: _FullLoopContext,
|
||||
emb: Mapping[str, Any],
|
||||
|
||||
@@ -252,6 +252,40 @@ class _VideoSparseAttentionBackendResolver(_CudaAttentionBackendResolver):
|
||||
raise ImportError("Video Sparse Attention backend is not installed.") from e
|
||||
|
||||
|
||||
class _VideoSparseAttentionH3BackendResolver(_CudaAttentionBackendResolver):
|
||||
backend = AttentionBackendEnum.VIDEO_SPARSE_ATTN_H3
|
||||
|
||||
# The vendored Triton tile-64 kernel is written against Hopper and
|
||||
# Blackwell block-sparse geometry; older architectures fail closed.
|
||||
supported_capabilities = {(9, 0), (10, 0), (10, 3)}
|
||||
|
||||
@classmethod
|
||||
def resolve(cls, platform) -> str:
|
||||
capability = platform.get_device_capability()
|
||||
capability_tuple = (
|
||||
(capability.major, capability.minor) if capability is not None else None
|
||||
)
|
||||
if capability_tuple not in cls.supported_capabilities:
|
||||
found = capability.as_version_str() if capability else "unknown"
|
||||
raise ValueError(
|
||||
"VSA-H3 (video_sparse_attn_h3) needs compute capability 9.0 "
|
||||
"(Hopper), 10.0 (B200 / GB200) or 10.3 (B300 / GB300); "
|
||||
f"this device reports {found}."
|
||||
)
|
||||
try:
|
||||
from sglang.multimodal_gen.runtime.layers.attention.backends.video_sparse_attn_h3 import ( # noqa: F401
|
||||
VideoSparseAttentionH3Backend,
|
||||
)
|
||||
|
||||
return "sglang.multimodal_gen.runtime.layers.attention.backends.video_sparse_attn_h3.VideoSparseAttentionH3Backend"
|
||||
except Exception as e:
|
||||
logger.error("Failed to import VSA-H3 attention backend: %s", str(e))
|
||||
raise ImportError(
|
||||
"VSA-H3 attention needs Triton and the in-tree tile-64 "
|
||||
"block-sparse kernel."
|
||||
) from e
|
||||
|
||||
|
||||
class _CubeSparseAttentionBackendResolver(_CudaAttentionBackendResolver):
|
||||
backend = AttentionBackendEnum.CUBE_SPARSE_ATTN
|
||||
|
||||
@@ -435,6 +469,7 @@ _CUDA_ATTENTION_BACKEND_RESOLVERS = {
|
||||
_SageAttention3BackendResolver,
|
||||
_SpargeAttentionBackendResolver,
|
||||
_VideoSparseAttentionBackendResolver,
|
||||
_VideoSparseAttentionH3BackendResolver,
|
||||
_CubeSparseAttentionBackendResolver,
|
||||
_SparseVideoGen2AttentionBackendResolver,
|
||||
_SolAttnBackendResolver,
|
||||
|
||||
@@ -35,6 +35,7 @@ class AttentionBackendEnum(enum.Enum):
|
||||
SAGE_ATTN_3 = enum.auto()
|
||||
SPARGE_ATTN = enum.auto()
|
||||
VIDEO_SPARSE_ATTN = enum.auto()
|
||||
VIDEO_SPARSE_ATTN_H3 = enum.auto()
|
||||
SPARSE_VIDEO_GEN_2_ATTN = enum.auto()
|
||||
VMOBA_ATTN = enum.auto()
|
||||
AITER = enum.auto()
|
||||
@@ -57,6 +58,7 @@ class AttentionBackendEnum(enum.Enum):
|
||||
return self in {
|
||||
AttentionBackendEnum.SLIDING_TILE_ATTN,
|
||||
AttentionBackendEnum.VIDEO_SPARSE_ATTN,
|
||||
AttentionBackendEnum.VIDEO_SPARSE_ATTN_H3,
|
||||
AttentionBackendEnum.SPARSE_VIDEO_GEN_2_ATTN,
|
||||
AttentionBackendEnum.VMOBA_ATTN,
|
||||
AttentionBackendEnum.SLA_ATTN,
|
||||
|
||||
@@ -42,6 +42,10 @@ BUILTIN_MODEL_OVERLAY_REGISTRY: dict[str, dict[str, Any]] = {
|
||||
"overlay_repo_id": "AgainstEntropy/SANA-WM_streaming-overlay",
|
||||
"overlay_revision": "62c6840871ecc3559189047513ba0670e1bf62e7",
|
||||
},
|
||||
"FastVideo/FastVideo-FastH3-4-step-Preview-v1-VSA-DataFree": {
|
||||
"overlay_repo_id": "kevin-mi/FastH3-4step-Preview-overlay",
|
||||
"overlay_revision": "f769cb8001dae335089de7b250364335bc7cb183",
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -688,6 +688,49 @@ MINIMAX_H3_FOUR_GPU_H100_CASES = [
|
||||
run_models_api_check=False,
|
||||
run_t2v_input_reference_check=False,
|
||||
),
|
||||
DiffusionTestCase(
|
||||
"fasth3_t2va_vsa_4gpu_h100",
|
||||
DiffusionServerArgs(
|
||||
model_path="FastVideo/FastVideo-FastH3-4-step-Preview-v1-VSA-DataFree",
|
||||
modality="video",
|
||||
num_gpus=4,
|
||||
extras=[
|
||||
"--attention-backend",
|
||||
"video_sparse_attn_h3",
|
||||
"--attention-backend-config",
|
||||
'{"VSA_sparsity": 0.9}',
|
||||
"--enable-torch-compile",
|
||||
"false",
|
||||
],
|
||||
),
|
||||
DiffusionSamplingParams(
|
||||
prompt=(
|
||||
"A curious raccoon peers through a vibrant field of yellow "
|
||||
"sunflowers, its eyes wide with interest."
|
||||
),
|
||||
output_size="1344x768",
|
||||
seconds=5,
|
||||
output_format="mp4",
|
||||
expect_audio_output=True,
|
||||
num_outputs_per_prompt=1,
|
||||
extras={
|
||||
"task": "t2va",
|
||||
"conditions": [],
|
||||
"target": {
|
||||
"short_edge": 768,
|
||||
"aspect_ratio": "16:9",
|
||||
"duration_seconds": 5.0,
|
||||
},
|
||||
"num_inference_steps": 5,
|
||||
"seed": 42,
|
||||
},
|
||||
),
|
||||
run_perf_check=False,
|
||||
run_consistency_check=False,
|
||||
run_component_accuracy_check=False,
|
||||
run_models_api_check=False,
|
||||
run_t2v_input_reference_check=False,
|
||||
),
|
||||
]
|
||||
|
||||
TWO_GPU_CASES = [
|
||||
|
||||
@@ -99,6 +99,7 @@ class TestDiffusionBenchmarkSkill(unittest.TestCase):
|
||||
"lingbot-world",
|
||||
"lingbot-world-v2",
|
||||
"fastwan21-t2v-1.3b",
|
||||
"fasth3-t2va-vsa",
|
||||
"wan22-t2v-nvfp4",
|
||||
"krea2-turbo",
|
||||
"krea2-raw",
|
||||
|
||||
@@ -0,0 +1,124 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""FastH3 (4-step VSA-distilled MiniMax-H3) registration and admission contracts."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from sglang.multimodal_gen.configs.pipeline_configs.minimax_h3 import (
|
||||
FastH3PipelineConfig,
|
||||
MiniMaxH3PipelineConfig,
|
||||
)
|
||||
from sglang.multimodal_gen.configs.sample.minimax_h3 import FastH3SamplingParams
|
||||
from sglang.multimodal_gen.registry import (
|
||||
get_model_info,
|
||||
get_non_diffusers_pipeline_name,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.distributed.parallel_state import (
|
||||
maybe_init_distributed_environment_and_model_parallel,
|
||||
model_parallel_is_initialized,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.layers.linear import UnquantizedLinearMethod
|
||||
from sglang.multimodal_gen.runtime.layers.quantization.fp8 import Fp8Config
|
||||
from sglang.multimodal_gen.runtime.models.dits.minimax_h3 import MiniMaxH3DiTModel
|
||||
from sglang.multimodal_gen.test.single_test_file.component_accuracy.utils import (
|
||||
ensure_distributed_env_defaults,
|
||||
)
|
||||
|
||||
FASTH3_MODEL_ID = "FastVideo/FastVideo-FastH3-4-step-Preview-v1-VSA-DataFree"
|
||||
|
||||
|
||||
def _ensure_single_process_parallel_runtime() -> None:
|
||||
if model_parallel_is_initialized():
|
||||
return
|
||||
ensure_distributed_env_defaults()
|
||||
maybe_init_distributed_environment_and_model_parallel(tp_size=1, sp_size=1)
|
||||
|
||||
|
||||
def test_registry_resolves_fasth3_configs() -> None:
|
||||
info = get_model_info(FASTH3_MODEL_ID)
|
||||
assert info.sampling_param_cls is FastH3SamplingParams
|
||||
assert info.pipeline_config_cls is FastH3PipelineConfig
|
||||
assert get_non_diffusers_pipeline_name(FASTH3_MODEL_ID) == "FastH3Pipeline"
|
||||
materialized = (
|
||||
"/cache/materialized_models/"
|
||||
"FastVideo__FastVideo-FastH3-4-step-Preview-v1-VSA-DataFree-0123abcd"
|
||||
)
|
||||
assert get_model_info(materialized).sampling_param_cls is FastH3SamplingParams
|
||||
|
||||
|
||||
def test_fasth3_sampling_defaults_and_task_rejection() -> None:
|
||||
params = FastH3SamplingParams(prompt="p")
|
||||
assert params.num_inference_steps == 5
|
||||
assert params.guidance_scale == 1.0
|
||||
|
||||
with pytest.raises(ValueError, match="exactly five sigma grid points"):
|
||||
FastH3SamplingParams(prompt="p", num_inference_steps=50)
|
||||
|
||||
with pytest.raises(ValueError, match="distilled for t2va only"):
|
||||
FastH3SamplingParams(
|
||||
prompt="p",
|
||||
task="fl2va",
|
||||
conditions=[{"type": "image", "uri": "x.png", "role": "first_frame"}],
|
||||
target={
|
||||
"short_edge": 768,
|
||||
"aspect_ratio": "16:9",
|
||||
"duration_seconds": 5.0,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def test_fasth3_pipeline_config_gates_and_rejections() -> None:
|
||||
config = FastH3PipelineConfig()
|
||||
assert config.dit_config.arch_config.has_gate_compress
|
||||
assert not MiniMaxH3PipelineConfig().dit_config.arch_config.has_gate_compress
|
||||
mapping = config.dit_config.arch_config.param_names_mapping
|
||||
source = "transformer_blocks.7.attn.to_gate_compress.weight"
|
||||
targets = [
|
||||
re.sub(pattern, target if isinstance(target, str) else target[0], source)
|
||||
for pattern, target in mapping.items()
|
||||
if re.match(pattern, source)
|
||||
]
|
||||
assert targets == ["blocks.7.attn.to_gate_compress.weight"]
|
||||
|
||||
with pytest.raises(ValueError, match="--model-variant does not apply"):
|
||||
config.validate_server_args(SimpleNamespace(model_variant="ref2va"))
|
||||
with pytest.raises(ValueError, match="no.*audited high-quality deployment"):
|
||||
config.validate_quality_deployment(server_args=None)
|
||||
|
||||
|
||||
def test_fasth3_lora_bundle_is_rejected_loudly() -> None:
|
||||
model = SimpleNamespace(arch=SimpleNamespace(adaln_affine_input_dim=None))
|
||||
plain = {
|
||||
"blocks.0.attn.qkv_proj.lora_A": torch.zeros(3, 64, 8),
|
||||
"blocks.0.attn.qkv_proj.lora_B": torch.zeros(3, 8, 64),
|
||||
}
|
||||
assert MiniMaxH3DiTModel.prepare_lora_adapter(model, dict(plain)) == plain
|
||||
|
||||
bundle = dict(plain)
|
||||
bundle["blocks.0.attn.qkv_proj.diff"] = torch.zeros(3, 64, 64)
|
||||
bundle["audio_patch_proj.diff_b"] = torch.zeros(64)
|
||||
bundle["blocks.0.attn.to_gate_compress.set_weight"] = torch.zeros(64, 64)
|
||||
with pytest.raises(ValueError, match="3 non-LoRA tensors.*set_weight"):
|
||||
MiniMaxH3DiTModel.prepare_lora_adapter(model, bundle)
|
||||
|
||||
|
||||
def test_fasth3_gates_stay_bf16_under_runtime_quantization() -> None:
|
||||
_ensure_single_process_parallel_runtime()
|
||||
with torch.device("meta"):
|
||||
model = MiniMaxH3DiTModel(
|
||||
config=FastH3PipelineConfig().dit_config,
|
||||
hf_config={},
|
||||
quant_config=Fp8Config(),
|
||||
)
|
||||
|
||||
attn = model.blocks[0].attn
|
||||
assert not isinstance(attn.qkv_proj.quant_method, UnquantizedLinearMethod)
|
||||
assert isinstance(attn.to_gate_compress.quant_method, UnquantizedLinearMethod)
|
||||
assert attn.to_gate_compress.weight.dtype == torch.bfloat16
|
||||
assert attn.to_gate_compress.weight.missing_param_init == "error"
|
||||
assert model.token_refiner.blocks[0].attn.to_gate_compress is None
|
||||
@@ -0,0 +1,241 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""VSA-H3 backend contracts.
|
||||
|
||||
The load-bearing check is sparsity -> 0: every tile is inside the budget, so
|
||||
the block-sparse kernel must reproduce dense attention over the packed rows to
|
||||
bf16 rounding. That single assertion pins the tile routing indices, ragged
|
||||
tile masking, the untile permutation, and the softmax scale.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from sglang.multimodal_gen.runtime.layers.attention.backends.video_sparse_attn_h3 import (
|
||||
VSA_H3_TILE_ELEMS,
|
||||
VideoSparseAttentionH3Impl,
|
||||
VideoSparseAttentionH3MetadataBuilder,
|
||||
_topk_tile_lists,
|
||||
)
|
||||
|
||||
requires_cuda = pytest.mark.skipif(
|
||||
not torch.cuda.is_available(), reason="VSA-H3 kernels need CUDA"
|
||||
)
|
||||
|
||||
# Ragged on purpose: text 70 and audio 100 are not tile multiples, and the
|
||||
# video canvas (5, 6, 10) is ragged in every tile dimension.
|
||||
PREFIX_SEGMENTS = (70, 0, 100)
|
||||
VIDEO_SHAPE = (5, 6, 10)
|
||||
HEADS = 4
|
||||
HEAD_DIM = 128
|
||||
|
||||
|
||||
def _build_metadata(sparsity: float, device, **kwargs):
|
||||
return VideoSparseAttentionH3MetadataBuilder().build(
|
||||
current_timestep=0,
|
||||
raw_latent_shape=VIDEO_SHAPE,
|
||||
patch_size=(1, 1, 1),
|
||||
VSA_sparsity=sparsity,
|
||||
prefix_segments=PREFIX_SEGMENTS,
|
||||
device=device,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
||||
def _packed_qkv(device, seed: int = 7):
|
||||
used = sum(PREFIX_SEGMENTS) + math.prod(VIDEO_SHAPE)
|
||||
total = (used + 63) // 64 * 64
|
||||
generator = torch.Generator(device="cpu").manual_seed(seed)
|
||||
tensors = [
|
||||
torch.randn(
|
||||
(total, HEADS, HEAD_DIM), generator=generator, dtype=torch.float32
|
||||
).to(device=device, dtype=torch.bfloat16)
|
||||
for _ in range(3)
|
||||
]
|
||||
return used, total, tensors
|
||||
|
||||
|
||||
def _dense_reference(q, k, v, used: int) -> torch.Tensor:
|
||||
qf = q[:used].float().permute(1, 0, 2)
|
||||
kf = k[:used].float().permute(1, 0, 2)
|
||||
vf = v[:used].float().permute(1, 0, 2)
|
||||
scores = qf @ kf.transpose(-2, -1) / math.sqrt(HEAD_DIM)
|
||||
return (torch.softmax(scores, dim=-1) @ vf).permute(1, 0, 2)
|
||||
|
||||
|
||||
def _impl():
|
||||
impl = VideoSparseAttentionH3Impl(
|
||||
num_heads=HEADS,
|
||||
head_size=HEAD_DIM,
|
||||
causal=False,
|
||||
softmax_scale=HEAD_DIM**-0.5,
|
||||
num_kv_heads=HEADS,
|
||||
prefix="blocks.3.attn",
|
||||
)
|
||||
assert impl.layer_idx == 3
|
||||
return impl
|
||||
|
||||
|
||||
def _run(impl, meta, used, total, q, k, v, gate=None):
|
||||
cu = torch.tensor([0, used, total], dtype=torch.int32, device=q.device)
|
||||
return impl.forward_varlen(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
cu_seqlens=cu,
|
||||
max_seqlen=used,
|
||||
cu_seqlens_host=(0, used, total),
|
||||
attn_metadata=meta,
|
||||
gate_compress=gate,
|
||||
)
|
||||
|
||||
|
||||
@requires_cuda
|
||||
def test_zero_sparsity_matches_dense() -> None:
|
||||
device = torch.device("cuda")
|
||||
meta = _build_metadata(0.0, device)
|
||||
used, total, (q, k, v) = _packed_qkv(device)
|
||||
assert meta.total_seq_length == used
|
||||
|
||||
out = _run(_impl(), meta, used, total, q, k, v)
|
||||
reference = _dense_reference(q, k, v, used)
|
||||
diff = (out[:used].float() - reference).abs().max().item()
|
||||
assert diff < 2e-2, f"sparse(0) vs dense max diff {diff}"
|
||||
assert torch.all(out[used:] == 0)
|
||||
|
||||
|
||||
@requires_cuda
|
||||
def test_zero_gate_is_noop_and_trained_gate_activates() -> None:
|
||||
device = torch.device("cuda")
|
||||
meta = _build_metadata(0.5, device)
|
||||
used, total, (q, k, v) = _packed_qkv(device)
|
||||
impl = _impl()
|
||||
|
||||
base = _run(impl, meta, used, total, q, k, v).clone()
|
||||
zero_gate = torch.zeros_like(q)
|
||||
gated_zero = _run(impl, meta, used, total, q, k, v, gate=zero_gate)
|
||||
assert torch.equal(base, gated_zero)
|
||||
|
||||
gate = torch.randn_like(q) * 0.1
|
||||
gated = _run(impl, meta, used, total, q, k, v, gate=gate)
|
||||
assert not torch.equal(base, gated)
|
||||
|
||||
|
||||
@requires_cuda
|
||||
def test_dense_overrides() -> None:
|
||||
device = torch.device("cuda")
|
||||
meta = _build_metadata(0.9, device, dense_layers=(3,))
|
||||
used, total, (q, k, v) = _packed_qkv(device)
|
||||
out = _run(_impl(), meta, used, total, q, k, v)
|
||||
reference = _dense_reference(q, k, v, used)
|
||||
diff = (out[:used].float() - reference).abs().max().item()
|
||||
assert diff < 2e-2, f"dense-layer opt-out vs dense max diff {diff}"
|
||||
|
||||
assert _build_metadata(0.9, device, dense_first_n_steps=0).VSA_sparsity == 0.9
|
||||
assert _build_metadata(0.9, device, dense_first_n_steps=1).VSA_sparsity == 0.0
|
||||
|
||||
|
||||
def _lists_to_mask(lists: torch.Tensor, n_tiles: int) -> torch.Tensor:
|
||||
mask = torch.zeros(*lists.shape[:-1], n_tiles, dtype=torch.bool)
|
||||
mask.scatter_(-1, lists.long(), True)
|
||||
return mask
|
||||
|
||||
|
||||
def test_topk_tile_list_semantics() -> None:
|
||||
num_prefix, num_video = 3, 10
|
||||
n_tiles = num_prefix + num_video
|
||||
scores = torch.randn(2, n_tiles, n_tiles)
|
||||
keep = math.ceil(0.5 * num_video)
|
||||
|
||||
exempt = _topk_tile_lists(scores, num_prefix, num_video, 0.5, True)
|
||||
assert exempt.shape == (2, num_video, num_prefix + keep)
|
||||
assert exempt.dtype == torch.int32
|
||||
assert torch.equal(exempt, exempt.sort(dim=-1).values)
|
||||
mask = _lists_to_mask(exempt, n_tiles)
|
||||
assert mask[..., :num_prefix].all()
|
||||
assert (mask[..., num_prefix:].sum(dim=-1) == keep).all()
|
||||
|
||||
compete = _topk_tile_lists(scores, num_prefix, num_video, 0.5, False)
|
||||
assert compete.shape == (2, num_video, min(keep + num_prefix, n_tiles))
|
||||
assert (_lists_to_mask(compete, n_tiles).sum(dim=-1) == keep + num_prefix).all()
|
||||
|
||||
|
||||
def _masked_dense_reference(meta, used, q, k, v, gate, sparsity):
|
||||
"""fp32 reference over the padded tile layout with the top-k tile mask."""
|
||||
n_tiles = meta.num_tiles
|
||||
seq_pad = n_tiles * VSA_H3_TILE_ELEMS
|
||||
valid = meta.pack_index >= 0
|
||||
src = meta.pack_index.clamp(min=0).long()
|
||||
|
||||
def tile(x):
|
||||
t = x[src].float().permute(1, 0, 2) # [H, S_pad, D]
|
||||
return t * valid.to(t.dtype)[None, :, None]
|
||||
|
||||
qt, kt, vt, gt = tile(q), tile(k), tile(v), tile(gate)
|
||||
sizes = meta.variable_block_sizes.float()
|
||||
pooled = [
|
||||
t.view(HEADS, n_tiles, VSA_H3_TILE_ELEMS, HEAD_DIM).sum(2)
|
||||
/ sizes[None, :, None]
|
||||
for t in (qt, kt, vt)
|
||||
]
|
||||
scores = pooled[0] @ pooled[1].transpose(-1, -2) / math.sqrt(HEAD_DIM)
|
||||
lists = _topk_tile_lists(
|
||||
scores, meta.num_prefix_tiles, meta.num_video_tiles, sparsity, meta.exempt
|
||||
)
|
||||
tile_mask = torch.ones(HEADS, n_tiles, n_tiles, dtype=torch.bool, device=q.device)
|
||||
tile_mask[:, meta.num_prefix_tiles :] = _lists_to_mask(lists.cpu(), n_tiles).to(
|
||||
q.device
|
||||
)
|
||||
row_mask = tile_mask.repeat_interleave(VSA_H3_TILE_ELEMS, 1).repeat_interleave(
|
||||
VSA_H3_TILE_ELEMS, 2
|
||||
)
|
||||
row_mask &= valid[None, None, :]
|
||||
logits = qt @ kt.transpose(-1, -2) / math.sqrt(HEAD_DIM)
|
||||
logits = logits.masked_fill(~row_mask, float("-inf"))
|
||||
out = torch.softmax(logits, dim=-1) @ vt
|
||||
compress = torch.softmax(scores, dim=-1) @ pooled[2] # [H, n_tiles, D]
|
||||
out = out + compress.repeat_interleave(VSA_H3_TILE_ELEMS, 1) * gt
|
||||
return out[:, meta.unpack_index.long()].permute(1, 0, 2) # [used, H, D]
|
||||
|
||||
|
||||
@requires_cuda
|
||||
def test_sparse_gated_matches_masked_dense_reference() -> None:
|
||||
device = torch.device("cuda")
|
||||
for exempt in (True, False):
|
||||
meta = _build_metadata(0.5, device, exempt=exempt)
|
||||
used, total, (q, k, v) = _packed_qkv(device)
|
||||
gate = (torch.randn_like(q) * 0.1).to(torch.bfloat16)
|
||||
out = _run(_impl(), meta, used, total, q, k, v, gate=gate)
|
||||
reference = _masked_dense_reference(meta, used, q, k, v, gate, 0.5)
|
||||
diff = (out[:used].float() - reference).abs().max().item()
|
||||
assert diff < 2e-2, f"exempt={exempt}: sparse+gate vs reference {diff}"
|
||||
assert torch.all(out[used:] == 0)
|
||||
|
||||
|
||||
def test_metadata_tile_geometry_accounts_every_row() -> None:
|
||||
device = torch.device("cpu")
|
||||
meta = _build_metadata(0.9, device)
|
||||
used = sum(PREFIX_SEGMENTS) + math.prod(VIDEO_SHAPE)
|
||||
assert int(meta.variable_block_sizes.sum()) == used
|
||||
assert meta.unpack_index.numel() == used
|
||||
assert int((meta.pack_index >= 0).sum()) == used
|
||||
assert meta.pack_index.numel() == meta.num_tiles * VSA_H3_TILE_ELEMS
|
||||
# Prefix chunks never straddle segment boundaries: 70 -> 64+6, 100 -> 64+36.
|
||||
assert meta.num_prefix_tiles == 4
|
||||
assert meta.variable_block_sizes[: meta.num_prefix_tiles].tolist() == [
|
||||
64,
|
||||
6,
|
||||
64,
|
||||
36,
|
||||
]
|
||||
video_tiles = (
|
||||
math.ceil(VIDEO_SHAPE[0] / 4)
|
||||
* math.ceil(VIDEO_SHAPE[1] / 4)
|
||||
* math.ceil(VIDEO_SHAPE[2] / 4)
|
||||
)
|
||||
assert meta.num_video_tiles == video_tiles
|
||||
assert meta.variable_block_sizes.numel() == meta.num_prefix_tiles + video_tiles
|
||||
assert VSA_H3_TILE_ELEMS == 64
|
||||
Reference in New Issue
Block a user