Compare commits

..
106 changed files with 1498 additions and 10507 deletions
-2
View File
@@ -6,7 +6,6 @@ on:
paths:
- 'docs/**'
- 'examples/**'
- 'scripts/inference/**'
- 'mkdocs.yml'
- 'requirements-mkdocs.in'
- 'requirements-mkdocs.txt'
@@ -17,7 +16,6 @@ on:
paths:
- 'docs/**'
- 'examples/**'
- 'scripts/inference/**'
- 'mkdocs.yml'
- 'requirements-mkdocs.in'
- 'requirements-mkdocs.txt'
-1
View File
@@ -9,7 +9,6 @@
**FastVideo is a unified post-training and real-time inference framework for accelerated video generation.**
## NEWS
- `2026/08/19`: FastVideo now supports MLX on Apple Silicon with [FastMetal-QAD](https://huggingface.co/collections/FastVideo/fastmetal), a family of 1.3B, 5B, and 14B models optimized for Mac—follow the [Apple Silicon guide](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/mps/) and read the [Blog](https://haoailab.com/blogs/fastmetal/).
- `2026/06/23`: Release FastWan-QAD: 5s of Video generated in 1.8s E2E. See the [FastWan-QAD models](https://huggingface.co/FastVideo/FastWan-QAD-FP8-1.3B), [Attn-QAT training guide](https://haoailab.com/FastVideo/training/attn_qat/), and [blog](https://haoailab.com/blogs/fastwan-qad/).
- `2026/03/17`: Release demo: Into the Dreamverse: Vibe Directing in FastVideo, check out the [Blog](https://haoailab.com/blogs/dreamverse/).
- `2026/03/13`: Release demo: Create a 5s 1080p Video in 4.5s with FastVideo on a Single GPU, check out the [Blog](https://haoailab.com/blogs/fastvideo_realtime_1080p/).
+1 -1
View File
@@ -68,7 +68,7 @@ ARG FLASH_ATTN_WHEEL_RELEASE_ARM64=https://github.com/mjun0812/flash-attention-p
# cutlass-4.4 `cute.core.ThrMma` API, which crashes on the cutlass-dsl 4.5 that
# flashinfer/quack pull in. After the wheel install we overlay this cutlass-4.5-safe
# upstream cute (flash-attn-4) so the image runs FA4 instead of the FA2 fallback.
ARG FA4_CUTE_REF=14c377950125c70b7a9dabf9c561fca53715ac7d
ARG FA4_CUTE_REF=82d6441eec5d4dfec120153db2c0145ae855a083
# Provided automatically by BuildKit/buildx (e.g. "amd64" / "arm64") and used to
# select the prebuilt flash-attn wheel. Empty under a plain `docker build` without
-52
View File
@@ -1,52 +0,0 @@
{
"recipes": [
{
"id": "fastwan21-t2v",
"task": "Text to video",
"label": "FastWan2.1 1.3B (distilled + VSA)",
"model": "FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
"source": "scripts/inference/inference_wan_VSA_DMD_1_3B.yaml",
"command": "FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN fastvideo generate --config scripts/inference/inference_wan_VSA_DMD_1_3B.yaml"
},
{
"id": "wan22-t2v",
"task": "Text to video",
"label": "Wan2.2 A14B",
"model": "Wan-AI/Wan2.2-T2V-A14B-Diffusers",
"source": "examples/inference/basic/basic_wan2_2.py",
"command": "python examples/inference/basic/basic_wan2_2.py"
},
{
"id": "wan21-i2v",
"task": "Image to video",
"label": "Wan2.1 14B 480P",
"model": "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers",
"source": "scripts/inference/inference_wan_i2v.yaml",
"command": "fastvideo generate --config scripts/inference/inference_wan_i2v.yaml"
},
{
"id": "turbowan22-i2v",
"task": "Image to video",
"label": "TurboWan2.2 A14B",
"model": "loayrashid/TurboWan2.2-I2V-A14B-Diffusers",
"source": "examples/inference/basic/basic_turbodiffusion_i2v.py",
"command": "python examples/inference/basic/basic_turbodiffusion_i2v.py"
},
{
"id": "wan22-ti2v",
"task": "Text or image to video",
"label": "Wan2.2 TI2V 5B",
"model": "Wan-AI/Wan2.2-TI2V-5B-Diffusers",
"source": "examples/inference/basic/basic_wan2_2_ti2v.py",
"command": "python examples/inference/basic/basic_wan2_2_ti2v.py"
},
{
"id": "matrix-game-2",
"task": "Interactive world",
"label": "Matrix Game 2.0",
"model": "FastVideo/Matrix-Game-2.0-Base-Distilled-Diffusers",
"source": "examples/inference/basic/basic_matrixgame2.py",
"command": "python examples/inference/basic/basic_matrixgame2.py"
}
]
}
-60
View File
@@ -1,60 +0,0 @@
(() => {
let recipesPromise;
const loadRecipes = (url) => {
recipesPromise ||= fetch(url).then((response) => {
if (!response.ok) throw new Error(`HTTP ${response.status}`);
return response.json();
});
return recipesPromise;
};
const init = () => {
document.querySelectorAll("[data-cookbook]").forEach(async (root) => {
if (root.dataset.initialized) return;
root.dataset.initialized = "true";
const select = root.querySelector("[data-cookbook-recipe]");
const model = root.querySelector("[data-cookbook-model]");
const source = root.querySelector("[data-cookbook-source]");
const command = root.querySelector("[data-cookbook-command]");
const status = root.querySelector("[data-cookbook-status]");
try {
const { recipes } = await loadRecipes(root.dataset.recipes);
const byId = new Map(recipes.map((recipe) => [recipe.id, recipe]));
const groups = new Map();
select.replaceChildren();
recipes.forEach((recipe) => {
if (!groups.has(recipe.task)) {
const group = document.createElement("optgroup");
group.label = recipe.task;
groups.set(recipe.task, group);
select.append(group);
}
groups.get(recipe.task).append(new Option(recipe.label, recipe.id));
});
const render = () => {
const recipe = byId.get(select.value);
model.textContent = recipe.model;
source.textContent = recipe.source;
source.href = `https://github.com/hao-ai-lab/FastVideo/blob/main/${recipe.source}`;
command.textContent = recipe.command;
status.textContent = `${recipe.label} selected.`;
};
select.addEventListener("change", render);
select.disabled = false;
render();
} catch (error) {
status.textContent = "Recipes could not be loaded. Use the examples link below.";
console.error("Failed to load FastVideo cookbook recipes", error);
}
});
};
if (window.document$) window.document$.subscribe(init);
else document.addEventListener("DOMContentLoaded", init);
})();
-40
View File
@@ -42,46 +42,6 @@ img {
margin: 0 auto;
}
.cookbook-picker {
padding: 1rem;
border: 0.05rem solid var(--md-default-fg-color--lightest);
border-radius: 0.2rem;
}
.cookbook-picker select {
width: 100%;
padding: 0.6rem;
color: var(--md-default-fg-color);
background: var(--md-default-bg-color);
border: 0.05rem solid var(--md-default-fg-color--lighter);
border-radius: 0.2rem;
}
.cookbook-picker dl {
display: grid;
grid-template-columns: max-content 1fr;
gap: 0.25rem 1rem;
}
.cookbook-picker dt {
font-weight: 700;
}
.cookbook-picker dd {
margin: 0;
min-width: 0;
overflow-wrap: anywhere;
}
.cookbook-picker__status {
position: absolute;
width: 1px;
height: 1px;
overflow: hidden;
clip: rect(0, 0, 0, 0);
white-space: nowrap;
}
.md-typeset .copy-page-button.md-button {
float: right;
margin: 0 0 1rem 1rem;
+68 -14
View File
@@ -9,33 +9,82 @@ FastVideo exposes a process-wide torch profiler that you can enable via environm
```bash
FASTVIDEO_TORCH_PROFILER_DIR=/mnt/traces/fastvideo \
FASTVIDEO_TORCH_PROFILE_REGIONS="profiler_region_model_loading,profiler_region_training_step"
FASTVIDEO_TORCH_PROFILE_REGIONS="model_loading,training_train_one_step" \
bash examples/train/run.sh /path/to/config.yaml
```
All profiled regions must be registered in `fastvideo.profiler`; the current list includes:
The `profiler_region_` prefix is optional. All profiled regions must be
registered in `fastvideo.profiler`; the current list includes:
- `profiler_region_model_loading` — pipeline/module loading
- `profiler_region_inference_pre_denoising`
- `profiler_region_inference_denoising`
- `profiler_region_inference_post_denoising`
- `profiler_region_training_checkpoint_saving`
- `profiler_region_training_dit`
- `profiler_region_training_train` — the complete training run
- `profiler_region_training_train_one_step` — one complete optimizer step
- `profiler_region_training_dataloader` — fetch the next batch in the trainer process
- `profiler_region_training_forward` — method forward and loss computation
- `profiler_region_training_backward` — backward pass
- `profiler_region_training_optimizer` — gradient clipping, optimizer/scheduler step, and zeroing gradients
- `profiler_region_training_callbacks` — end-of-step callbacks such as EMA updates
- `profiler_region_training_validation`
- `profiler_region_training_epoch`
- `profiler_region_training_step`
- `profiler_region_training_backward`
- `profiler_region_training_optimizer`
- `profiler_region_training_save_checkpoint`
- `profiler_region_distillation_teacher_forward`
- `profiler_region_distillation_student_forward`
- `profiler_region_distillation_loss`
- `profiler_region_distillation_update`
- `profiler_region_dmd2_student_rollout` — one DMD2 student rollout, including simulated prefix steps
- `profiler_region_dmd2_generator_loss` — teacher/critic scoring for the generator loss
- `profiler_region_dmd2_critic_loss` — critic flow-matching loss, including its student rollout
### Profiling modular training
The YAML-driven trainer under `fastvideo/train/` initializes and flushes the
profiler automatically. For example, this captures six DMD2 steps: one warmup
step, four steady-state critic updates, and the configured 1-in-5 generator
update, without changing the method's update cadence:
```bash
TRACE_DIR="$(pwd)/profiler_traces/wan_dmd2"
mkdir -p "$TRACE_DIR"
NUM_GPUS=4 \
WANDB_MODE=offline \
FASTVIDEO_TORCH_PROFILER_DIR="$TRACE_DIR" \
FASTVIDEO_TORCH_PROFILE_REGIONS="training_train,training_train_one_step,training_dataloader,training_forward,training_backward,training_optimizer,training_callbacks,dmd2_student_rollout,dmd2_generator_loss,dmd2_critic_loss" \
bash examples/train/run.sh \
examples/train/configs/distribution_matching/wan/dmd2_t2v.yaml \
--training.loop.max_train_steps 6 \
--training.checkpoint.training_state_checkpointing_steps 0 \
--callbacks.validation.every_steps 0
```
The two overrides disable validation and checkpoint writes so their I/O does
not contaminate training-step measurements. Remove them when profiling those
regions. Methods that manage optimization internally emit the enclosing
`training_train_one_step` region but do not emit the generic
dataloader/forward/backward/optimizer child regions because those boundaries
belong to the method. With `dataloader_num_workers > 0`, the
`training_dataloader` region measures the training process waiting for and
receiving the next batch; work performed inside dataloader worker subprocesses
does not appear in that process's torch-profiler trace.
The modular trainer also logs `dataloader_time_sec` on every ordinary training
step (including DMD2), so routine runs can monitor rank 0's
`next(dataloader)` wait without enabling the torch profiler. It is
intentionally not reduced across ranks; cross-rank synchronization would
perturb the hot path being measured.
While profiling is enabled, FastVideo records additional annotations:
- `fastvideo.region::<name>` spans are emitted when entering a region.
- `fastvideo.profiler.enable_collection` / `fastvideo.profiler.disable_collection` events mark when torch profiler collection is toggled on or off.
Only one profiler instance is created per process; subsequent pipelines reuse the same controller. If you set `FASTVIDEO_TORCH_PROFILE_REGIONS` incorrectly (e.g. misspelled name), FastVideo logs a warning and ignores that entry.
Only one profiler controller is created per process; subsequent pipelines
reuse it. Each outermost enabled region invocation produces one complete
CPU/CUDA trace segment. Enabled regions nested inside it are annotations in
that segment. Include an enclosing region such as `training_train_one_step`
when selecting its child regions so they share one trace and profiler startup
does not perturb each phase independently. If you set
`FASTVIDEO_TORCH_PROFILE_REGIONS` incorrectly (e.g. misspelled name), FastVideo
logs a warning and ignores that entry.
Additional knobs:
@@ -44,13 +93,18 @@ Additional knobs:
- `FASTVIDEO_TORCH_PROFILER_WITH_STACK`
- `FASTVIDEO_TORCH_PROFILER_WITH_FLOPS`
Traces can be visualized using <https://ui.perfetto.dev/>.
Traces can be visualized using <https://ui.perfetto.dev/>. Each rank also
writes `summary_rank<N>_segment<MMMM>_<region>.txt` and the corresponding JSON
beside every trace segment for a quick operator-level view without loading the
full timeline. The JSON contains CPU and device totals; input-shape grouping is
only enabled with `FASTVIDEO_TORCH_PROFILER_RECORD_SHAPES=1` because it makes
summary generation materially more expensive on large traces.
### Best Practices
- Keep the profiled step count small; traces can be large and slow down job shutdown while the profiler flushes data.
- After profiling, clean up trace directories to avoid filling disk storage.
- When adding new regions, register them in `fastvideo.profiler` and wrap the corresponding code block with `with self.profiler_controller.region("your_region"):` or the `@profile_region` decorator.
- When adding new regions, register them in `fastvideo.profiler` and wrap the corresponding code block with `profiler_region("your_region")` or the `@profile_region` decorator.
## Related: Activation Trace Mode
-44
View File
@@ -1,44 +0,0 @@
# Inference Cookbook
Choose a complete recipe maintained in the FastVideo repository. Each command
runs its checked-in source directly, so coupled model, GPU, offload, and
attention settings do not drift into unsupported combinations.
The commands expect a local clone:
```bash
git clone https://github.com/hao-ai-lab/FastVideo.git
cd FastVideo
```
<div class="cookbook-picker" data-cookbook data-recipes="../assets/cookbook-recipes.json">
<label for="cookbook-recipe"><strong>Recipe</strong></label>
<select id="cookbook-recipe" data-cookbook-recipe disabled>
<option>Loading recipes…</option>
</select>
<dl>
<dt>Model</dt>
<dd data-cookbook-model>Loading…</dd>
<dt>Source</dt>
<dd><a data-cookbook-source href="../inference/examples/basic/">Browse maintained examples</a></dd>
</dl>
<pre><code class="language-bash" data-cookbook-command>Loading…</code></pre>
<p class="cookbook-picker__status" role="status" aria-live="polite" data-cookbook-status></p>
<noscript>
JavaScript is needed for the recipe picker. Browse the
<a href="../inference/examples/examples_inference_index/">inference examples</a>
instead.
</noscript>
</div>
## Customize a recipe
Start from the checked-in source, then change only the settings your model
supports:
- [Configuration](../inference/configuration.md) covers the Python and CLI
config surfaces.
- [Optimizations](../inference/optimizations.md) covers attention backends,
compilation, and memory tradeoffs.
- [Support matrix](../inference/support_matrix.md) lists supported models and
optimizations.
@@ -76,10 +76,6 @@ surfaces:
lora_target_modules: "Legacy LoRA configuration surface pending dedicated component API."
output_type: "Legacy output formatting surface pending GenerationResult cleanup."
VSA_sparsity: "Model-specific inference optimization not yet represented in the typed public schema."
VSA_tile_size: "VSA-H3 tile geometry request; model-specific optimization not yet represented in the typed public schema."
vae_parallel_decode: "MiniMax-H3 sequence-parallel VAE decode opt-in; model-specific optimization not yet represented in the typed public schema."
vae_parallel_encode: "MiniMax-H3 sequence-parallel reference VAE encode opt-in; model-specific optimization not yet represented in the typed public schema."
vae_parallel_decode_strategy: "Chunk-transport collective for vae_parallel_decode; model-specific optimization not yet represented in the typed public schema."
attention_backend: "Process-wide default attention-backend request applied per component at load time; kernel-selection knob not yet represented in the typed public schema."
moba_config_path: "Model-specific MoBA optimization surface not yet represented in the typed public schema."
master_port: "Executor/bootstrap compatibility field; not part of the canonical inference schema."
-37
View File
@@ -2,7 +2,6 @@
# adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/docs/source/generate_examples.py
import itertools
import json
import os
import re
from dataclasses import dataclass, field
@@ -20,40 +19,6 @@ GENERATED_DOC_PREFIXES = (
"training/examples/",
"distillation/examples/",
)
COOKBOOK_DATA = ROOT_DIR / "docs/assets/cookbook-recipes.json"
COOKBOOK_SOURCE_ROOTS = (
ROOT_DIR / "examples/inference",
ROOT_DIR / "scripts/inference",
)
def validate_cookbook() -> None:
"""Keep cookbook entries tied to checked-in runnable sources."""
recipes = json.loads(COOKBOOK_DATA.read_text(encoding="utf-8")).get("recipes")
if not isinstance(recipes, list) or not recipes:
raise ValueError(f"{COOKBOOK_DATA}: recipes must be a non-empty list")
seen: set[str] = set()
for recipe in recipes:
required = ("id", "task", "label", "model", "source", "command")
missing = {key for key in required if not recipe.get(key)}
if missing:
raise ValueError(f"Cookbook recipe is missing: {', '.join(sorted(missing))}")
if recipe["id"] in seen:
raise ValueError(f"Duplicate cookbook recipe id: {recipe['id']}")
seen.add(recipe["id"])
source = (ROOT_DIR / recipe["source"]).resolve()
if not any(source.is_relative_to(root.resolve()) for root in COOKBOOK_SOURCE_ROOTS):
raise ValueError(f"Cookbook source is outside an approved directory: {recipe['source']}")
if not source.is_file():
raise ValueError(f"Cookbook source does not exist: {recipe['source']}")
source_text = source.read_text(encoding="utf-8")
if recipe["model"] not in source_text:
raise ValueError(f"Cookbook model is not present in {recipe['source']}: {recipe['model']}")
if recipe["source"] not in recipe["command"]:
raise ValueError(f"Cookbook command does not invoke its source: {recipe['id']}")
def fix_case(text: str) -> str:
@@ -571,7 +536,6 @@ def on_pre_build(config, **kwargs):
MkDocs hook to generate examples before building the documentation.
This function is called automatically by MkDocs' native hook system.
"""
validate_cookbook()
print("Generating example documentation...")
generate_examples(generate_main_index=True)
print("Example documentation generated successfully!")
@@ -585,7 +549,6 @@ def on_page_context(context, page, **kwargs):
if __name__ == "__main__":
validate_cookbook()
print("Generating example documentation...")
generate_examples(generate_main_index=True)
print("Example documentation generated successfully!")
+2 -3
View File
@@ -65,7 +65,6 @@ uv pip install flash-attn --no-build-isolation -v
## Next Steps
- [Quick Start](quick_start.md) - Generate your first video
- [Inference Cookbook](../cookbook/index.md) - Choose a maintained recipe
- [Quick Start Guide](quick_start.md) - Get started with your first video generation
- [Configuration](../inference/configuration.md) - Learn about configuration options
- [Examples](../inference/examples/examples_inference_index.md) - Explore scripts and notebooks
- [Examples](../inference/examples/examples_inference_index.md) - Explore example scripts and notebooks
+49 -9
View File
@@ -23,21 +23,61 @@ Also optionally install flash-attn:
uv pip install flash-attn --no-build-isolation -v
```
## Choose a maintained recipe
## Basic Usage
The cookbook selects complete, checked-in recipes instead of mixing model,
parallelism, offload, and attention settings independently.
### Text-to-Video Generation
[Open the inference cookbook](../cookbook/index.md){ .md-button .md-button--primary }
```python
from fastvideo import VideoGenerator
!!! tip "Need more control?"
Start from a maintained recipe, then use the
[configuration](../inference/configuration.md) and
[optimization](../inference/optimizations.md) guides for supported changes.
def main():
# Create a video generator with a pre-trained model
generator = VideoGenerator.from_pretrained(
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
num_gpus=1, # Adjust based on your hardware
)
# Define a prompt for your video
prompt = "A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes wide with interest."
# Generate the video
video = generator.generate_video(
prompt,
output_path="my_videos/", # Controls where videos are saved
save_video=True
)
if __name__ == '__main__':
main()
```
### Image-to-Video Generation
```python
from fastvideo import VideoGenerator, SamplingParam
def main():
# Create the generator
model_name = "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
generator = VideoGenerator.from_pretrained(model_name, num_gpus=1)
# Set up parameters with an initial image
sampling_param = SamplingParam.from_pretrained(model_name)
sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
sampling_param.num_frames = 107
# Generate video based on the image
prompt = "A photograph coming to life with gentle movement"
generator.generate_video(prompt, sampling_param=sampling_param,
output_path="my_videos/",
save_video=True)
if __name__ == '__main__':
main()
```
## Next Steps
- [Inference Cookbook](../cookbook/index.md) - Choose a maintained recipe
- [Installation Guide](installation.md) - Detailed installation instructions
- [Configuration](../inference/configuration.md) - Learn about configuration options
- [Examples](../inference/examples/examples_inference_index.md) - Explore more
+3 -1
View File
@@ -431,7 +431,9 @@ callbacks:
```
The EMA callback owns its own state and checkpoints independently — EMA weights
are saved and restored automatically on resume.
are saved and restored automatically on resume. It advances only after the
training method reports that the student optimizer updated; alternating methods
such as DMD2 therefore skip EMA work on critic-only steps.
### ValidationCallback
-8
View File
@@ -33,14 +33,6 @@ For the typed config/request path added during the inference API refactor:
python examples/inference/basic/basic_dmd_new_api.py
```
For the few-step (4-step, DMD2-distilled) MiniMax-H3 preview, generating synchronized video and audio, optionally with block-sparse VSA attention:
```
python examples/inference/basic/basic_fasth3.py --prompt "your prompt" [--vsa-sparsity 0.9]
```
The default checkpoint `FastVideo/FastVideo-Minimax-FastH3-Preview-v0.1` is private on the Hub while its license review completes; until it flips public, pass `--model-path` with a local snapshot of the release.
On Blackwell (sm_100) GPUs with a `fastvideo-kernel` build that carries the sm_100a block-sparse extension, `--vsa-kernel sm100a` routes the tile-64 attention forwards through the CUDA kernel instead of Triton (it sets `FASTVIDEO_VSA_SM100A=1` before the pipeline boots); if the extension or the arch is missing, the run warns once and falls back to Triton.
## Basic Walkthrough
All you need to generate videos using multi-gpus from state-of-the-art diffusion pipelines is the following few lines!
-180
View File
@@ -1,180 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Few-step video+audio generation with the DMD2-distilled MiniMax H3 preview.
FastVideo/FastVideo-Minimax-FastH3-Preview-v0.1 is a 4-step distillation of
MiniMaxAI/MiniMax-H3 (data-free DMD2): it walks a 4-step grid on the release
sampler's shift-12 schedule instead of the base model's 50 steps, generating
synchronized video and audio in one pipeline call.
The student was trained with block-sparse video attention (VSA, 64-token
tiles) and its checkpoint carries the trained sparse-gate parameters
(``attn.to_gate_compress``), so this script always runs the VSA-H3 attention
backend. At the default ``--vsa-sparsity 0.0`` the attention math is exactly
dense (every tile is selected); raise the sparsity for additional speedup.
"""
from __future__ import annotations
import argparse
import os
from pathlib import Path
from fastvideo import VideoGenerator
from fastvideo.api import (
CompileConfig,
EngineConfig,
GenerationRequest,
GeneratorConfig,
OffloadConfig,
OutputConfig,
ParallelismConfig,
PipelineSelection,
SamplingConfig,
)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--model-path", default="FastVideo/FastVideo-Minimax-FastH3-Preview-v0.1")
# The HF repo is private while the MiniMax H3 Community License review
# completes; until it flips public, pass --model-path with a local
# snapshot of the release instead (e.g. the team export at
# /mnt/lustre/vlm-wlsaidhi/fastvideo/exports/FastVideo-Minimax-FastH3-Preview-v0.1).
parser.add_argument("--prompt", required=True)
parser.add_argument("--output", default="outputs/fasth3")
parser.add_argument("--height", type=int, default=768)
parser.add_argument("--width", type=int, default=1344)
parser.add_argument("--num-frames", type=int, default=124)
# num_inference_steps counts sigma-GRID POINTS, matching the base model's
# convention ("50 steps" = a 50-point grid = 49 transformer forwards). The
# student's distilled 4-step grid is 4 FORWARDS, i.e. a 5-point grid
# (t = 1000, 750, 500, 250 -> 0 on the shift-12 schedule) — so the correct
# default here is 5. Other grids are off-distribution.
parser.add_argument("--steps",
type=int,
default=5,
help="num_inference_steps = sigma-grid points; N points run N-1 denoising "
"forwards. 5 (default) is the distilled 4-forward grid")
parser.add_argument("--seed", type=int, default=0)
parser.add_argument("--num-gpus", type=int, default=4)
parser.add_argument("--vsa-sparsity",
type=float,
default=0.0,
help="Run-level VSA sparsity in [0, 1). 0.0 (default) selects every tile, which is "
"exactly dense attention; the student was trained at 0.9")
# 64 is the trained contract: the student was TRAINED with 64-token
# (4,4,4) tiles, and its to_gate_compress gates were learned against
# pooling at that granularity — keep 64 unless you are ablating.
parser.add_argument("--vsa-tile-size",
type=int,
choices=(64, 256),
default=64,
help="VSA-H3 tile size in tokens; 64 (default) is what the student was trained "
"with and runs the native Triton block-sparse path, 256 is the FA4-CuTe-capable "
"geometry for ablations")
parser.add_argument("--vsa-kernel",
choices=("triton", "sm100a"),
default="triton",
help="Block-sparse kernel for the tile-64 attention forward: triton (default, "
"portable fwd+bwd) or sm100a — the opt-in Blackwell CUDA forward "
"(fastvideo_kernel.block_sparse_attn_sm100a). sm100a needs an sm_100 GPU and a "
"fastvideo-kernel build that carries the extension; if a precondition fails at "
"run time the attention layer logs one warning and falls back to Triton. Only "
"meaningful with --vsa-tile-size 64")
parser.add_argument("--torch-compile", action="store_true", help="torch.compile the DiT transformer path")
parser.add_argument("--compile-mode",
default=None,
help='torch.compile mode, e.g. "reduce-overhead" for CUDA graphs')
parser.add_argument("--repeats",
type=int,
default=1,
help="generate N times; with --torch-compile the first run pays "
"compilation, so steady-state is the last repeat")
return parser.parse_args()
def main() -> None:
args = parse_args()
output_dir = Path(args.output)
output_dir.mkdir(parents=True, exist_ok=True)
if args.vsa_kernel == "sm100a":
# The attention backend reads FASTVIDEO_VSA_SM100A per forward; set it
# before the pipeline boots so spawned GPU workers inherit it. The
# kernel is forward-only and inference runs under no-grad, so every
# denoising forward qualifies for the CUDA route.
os.environ["FASTVIDEO_VSA_SM100A"] = "1"
# Boot-time run configuration, folded into FastVideoArgs (the same route
# examples/inference/basic/basic_minimax_h3_t2v.py uses for sparsity):
# - attention_backend: the checkpoint carries trained to_gate_compress
# gates, which only exist under the VSA-H3 backend — a dense-backend
# load would reject them as unexpected weights. Layers that do not
# support VSA-H3 (e.g. the token refiner) fall back to flash attention.
# - VSA_tile_size: forwarded even at sparsity 0.0 because the gate-compress
# branch pools per tile, and the gates were trained at 64 tokens/tile.
experimental: dict[str, object] = {
"attention_backend": "VIDEO_SPARSE_ATTN_H3",
"VSA_tile_size": args.vsa_tile_size,
}
if args.vsa_sparsity > 0.0:
experimental["VSA_sparsity"] = args.vsa_sparsity
generator = VideoGenerator.from_config(
GeneratorConfig(
model_path=args.model_path,
pipeline=PipelineSelection(experimental=experimental),
engine=EngineConfig(
num_gpus=args.num_gpus,
use_fsdp_inference=args.num_gpus > 1,
parallelism=ParallelismConfig(tp_size=1, sp_size=args.num_gpus),
offload=OffloadConfig(
dit=False,
dit_layerwise=False,
text_encoder=True,
vae=True,
pin_cpu_memory=False,
),
compile=CompileConfig(
enabled=args.torch_compile,
mode=args.compile_mode,
),
),
))
try:
request = GenerationRequest(
prompt=args.prompt,
negative_prompt="",
sampling=SamplingConfig(
height=args.height,
width=args.width,
num_frames=args.num_frames,
fps=24,
num_inference_steps=args.steps,
# the base model is guidance-distilled; the student inherits it
guidance_scale=1.0,
batch_cfg=False,
seed=args.seed,
),
output=OutputConfig(
output_path=str(output_dir / "fasth3.mp4"),
save_video=True,
return_frames=False,
),
)
result = generator.generate(request)
print(f"Output written to: {result.video_path}")
if result.generation_time is not None:
# machine-readable: benchmark harnesses parse this line to separate
# generation from model-load time (last occurrence = steady state)
print(f"Generation time: {result.generation_time:.2f}s")
for _ in range(args.repeats - 1):
result = generator.generate(request)
if result.generation_time is not None:
print(f"Generation time: {result.generation_time:.2f}s")
finally:
generator.shutdown()
if __name__ == "__main__":
main()
@@ -27,8 +27,7 @@ def parse_args() -> argparse.Namespace:
parser.add_argument("--model-path", default="MiniMaxAI/MiniMax-H3")
# Rank-reduced AdaLN checkpoint (-39% params, -23 GiB VRAM): pass
# --model-path noctuashap/MiniMax-H3-pruned-r16
# (or a local dir produced by
# scripts/checkpoint_conversion/convert_minimax_h3_adaln_rank.py).
# (or a local dir converted with tools/minimax_h3/fit_adaln_basis.py).
# adaln_rank is read from the checkpoint config; no other flags needed.
# Rank-reduced checkpoints are inference-only: training needs the
# full-rank release.
@@ -27,8 +27,7 @@ def parse_args() -> argparse.Namespace:
parser.add_argument("--model-path", default="MiniMaxAI/MiniMax-H3")
# Rank-reduced AdaLN checkpoint (-39% params, -23 GiB VRAM): pass
# --model-path noctuashap/MiniMax-H3-pruned-r16
# (or a local dir produced by
# scripts/checkpoint_conversion/convert_minimax_h3_adaln_rank.py).
# (or a local dir converted with tools/minimax_h3/fit_adaln_basis.py).
# adaln_rank is read from the checkpoint config; no other flags needed.
# Rank-reduced checkpoints are inference-only: training needs the
# full-rank release.
@@ -24,8 +24,7 @@ def parse_args() -> argparse.Namespace:
parser.add_argument("--model-path", default="MiniMaxAI/MiniMax-H3")
# Rank-reduced AdaLN checkpoint (-39% params, -23 GiB VRAM): pass
# --model-path noctuashap/MiniMax-H3-pruned-r16
# (or a local dir produced by
# scripts/checkpoint_conversion/convert_minimax_h3_adaln_rank.py).
# (or a local dir converted with tools/minimax_h3/fit_adaln_basis.py).
# adaln_rank is read from the checkpoint config; no other flags needed.
# Rank-reduced checkpoints are inference-only: training needs the
# full-rank release.
-36
View File
@@ -318,38 +318,6 @@ if(BUILD_CXX_KERNELS)
# Combined FastVideo Extension
# Using name 'fastvideo_kernel_ops' to distinguish from the python package namespace
# ---------------------------------------------------------------------------
# VSA block-sparse attention forward, Blackwell (sm_100a) only.
#
# NOTE the "a" suffix: -arch=sm_100a is NOT enough -- it emits a plain sm_100 target and
# ptxas rejects every tcgen05 / setmaxnreg instruction. The explicit gencode spelling
# below is required, and matches the 10.0a entry in TORCH_CUDA_ARCH_LIST.
#
# Built for 64-token sparse blocks -- FastVideo's default (4,4,4) tiling, so no
# tile-size change and no top-k granularity change is needed. VSA_BHSD
# selects [B, H, S, D]. Both are compile-time; the Python is_supported() checks incoming
# tensors against them so callers fall back to Triton rather than getting a wrong answer.
# Read the ENVIRONMENT as well as the cache variable. When TORCH_CUDA_ARCH_LIST is
# exported (build.sh, and `pip install` with it set) the branch above only prints it --
# the cmake variable stays empty, so testing that alone silently skips the kernel and
# leaves a build that succeeds with the op missing.
set(ENABLE_VSA_SM100A OFF)
set(_VSA_ARCH_LIST "${TORCH_CUDA_ARCH_LIST}")
if(NOT _VSA_ARCH_LIST AND DEFINED ENV{TORCH_CUDA_ARCH_LIST})
set(_VSA_ARCH_LIST "$ENV{TORCH_CUDA_ARCH_LIST}")
endif()
if(_VSA_ARCH_LIST MATCHES "(^|[; ,])(10\\.0a|100a|sm_100a)([; ,]|$)")
set(ENABLE_VSA_SM100A ON)
endif()
if(ENABLE_VSA_SM100A)
message(STATUS "fastvideo-kernel: building block_sparse_sm100a (Blackwell, 64- and 128-token blocks)")
list(APPEND EXTENSION_SOURCES csrc/attention/block_sparse_sm100a.cu
csrc/attention/block_sparse_blk128_sm100a.cu)
set_source_files_properties(csrc/attention/block_sparse_sm100a.cu
csrc/attention/block_sparse_blk128_sm100a.cu PROPERTIES
COMPILE_OPTIONS "-gencode;arch=compute_100a,code=sm_100a;-DVSA_BHSD=true")
endif()
Python_add_library(fastvideo_kernel_ops MODULE USE_SABI ${SKBUILD_SABI_VERSION} WITH_SOABI
${EXTENSION_SOURCES}
)
@@ -365,14 +333,10 @@ if(BUILD_CXX_KERNELS)
# Build compile definitions list
set(COMPILE_DEFS TORCH_EXTENSION_NAME=fastvideo_kernel_ops)
if(ENABLE_VSA_SM100A)
list(APPEND COMPILE_DEFS TK_COMPILE_BLOCK_SPARSE_VSA_SM100A)
endif()
if(ENABLE_TK_KERNELS)
list(APPEND COMPILE_DEFS TK_COMPILE_ST_ATTN TK_COMPILE_BLOCK_SPARSE)
endif()
target_compile_definitions(fastvideo_kernel_ops PRIVATE ${COMPILE_DEFS})
target_compile_options(fastvideo_kernel_ops PRIVATE
+11 -40
View File
@@ -17,7 +17,7 @@ Runtime-JIT kernels (no build step, ship in every wheel/image):
| Kernels | Where | Used when |
|---|---|---|
| Triton: STA, VSA block-sparse, SLA, fused compress+topk, FP4 QAT training, quant/norm utils | `python/fastvideo_kernel/triton_kernels/` | automatic fallback when the matching C++ op is absent (`ops.py`, `turbodiffusion_ops.py`) |
| FA4 CuTe-DSL block-sparse forward/backward (VSA-128/256 fastpath on `sm_100`) | `block_sparse_attn_cute_fwd.py` | optional `flash_attn.cute` dependency, see below |
| FA4 CuTe-DSL block-sparse (VSA-256 fastpath on `sm_100`) | `block_sparse_attn_cute_fwd.py` | optional `flash_attn.cute` dependency, see below |
| VMoBA `moba_attn_varlen` | `vmoba.py` | wraps flash-attn varlen |
## What gets built where, and when
@@ -64,49 +64,28 @@ cd fastvideo-kernel
./build.sh --rocm
```
### Optional: FA4 CuTe block-sparse backend (VSA-128/256 fastpath)
### Optional: FA4 CuTe block-sparse backend (VSA-256 fastpath)
The VSA-128/256 fastpaths (tile volume 128 or 256, on NVIDIA Blackwell / sm_100) route to the
The VSA-256 fastpath (tile volume 256, on NVIDIA Blackwell / sm_100) routes to the
FlashAttention-4 CuTe-DSL block-sparse kernel exposed as `flash_attn.cute`. This is
an **optional** dependency: it is imported lazily, and `video_sparse_attn`
transparently falls back to the Triton backend when it is absent (so the package is
fully usable without it).
The symbols the fastpath needs (`flash_attn.cute.block_sparsity.BlockSparseTensorsTorch`
and the public/private forward-backward bridges in `flash_attn.cute.interface`) are provided upstream by
The symbols the fastpath needs (`flash_attn.cute.block_sparsity.BlockSparseTensorsTorch`,
`flash_attn.cute.interface._flash_attn_fwd`) are provided upstream by
[Dao-AILab/flash-attention](https://github.com/Dao-AILab/flash-attention). Pin to
commit `14c377950125c70b7a9dabf9c561fca53715ac7d`, the revision FastVideo pins as
commit `940cd9680f3315f2f06b43ab5bea2c2cf2d96806`, the revision FastVideo pins as
the `flash-attn-4` source in the repo-root `pyproject.toml`; other revisions may
have incompatible block-sparse forward/backward interfaces.
Install it under its distribution name so its own runtime stack resolves with it.
Do **not** pre-install `nvidia-cutlass-dsl` by hand: this revision pins
`nvidia-cutlass-dsl==4.6.0.dev0` exactly, and a hand-installed 4.5.x floor either
gets silently upgraded or, if something else holds it back, leaves the CuTe
kernels broken.
have an incompatible `_flash_attn_fwd` signature.
```bash
pip install torchvision
pip install "flash-attn-4 @ git+https://github.com/Dao-AILab/flash-attention.git@14c377950125c70b7a9dabf9c561fca53715ac7d#subdirectory=flash_attn/cute"
pip install "nvidia-cutlass-dsl>=4.5.0" torchvision
pip install "git+https://github.com/Dao-AILab/flash-attention.git@940cd9680f3315f2f06b43ab5bea2c2cf2d96806#subdirectory=flash_attn/cute"
```
That resolves `nvidia-cutlass-dsl` to 4.6.0.dev0 and `quack-kernels` to 0.5.3, a
combination this revision works with. A mismatched CuTe DSL only surfaces when the
kernel JIT-compiles, so the error points at CuTe internals rather than at the
install:
| Error on first VSA-128/256 CuTe call | Cause |
|---|---|
| `TypeError: fmax() missing 1 required positional argument: 'b'` | `nvidia-cutlass-dsl` 4.5.x |
| `AttributeError: module 'cutlass.cute.core' has no attribute 'ThrMma'` | `quack-kernels` older than 0.5.1 |
| `ImportError: cannot import name 'alloc_reserved_mbarrier'` | `quack-kernels` 0.6.2 or newer |
An environment whose `flash_attn.cute` came from a prebuilt flash-attn wheel rather
than from this pin hits the first row; that is what the overlay step in
`docker/Dockerfile` works around.
The CuTe kernels JIT-compile on first use. Forward and backward are verified on
Blackwell (sm_100) against `tests/test_vsa128_*.py` and `tests/test_vsa256_*.py`.
The CuTe kernel JIT-compiles on first use. Verified on Blackwell (sm_100) against
`tests/test_vsa256_forward*.py`.
## Usage
@@ -163,14 +142,6 @@ After building/installing `fastvideo-kernel`, run:
```bash
cd fastvideo-kernel
python benchmarks/bench_vsa.py --batch_size 1 --num_heads 16 --head_dim 128 --q_seq_lens 49152 --topk 64
# VSA-256 FA4 CuTe forward/backward on Blackwell
python benchmarks/bench_vsa.py --block_size 256 --use_cute \
--batch_size 1 --num_heads 12 --head_dim 128 --q_seq_lens 39936 --topk 20
# VSA-128 FA4 CuTe forward/backward on Blackwell
python benchmarks/bench_vsa.py --block_size 128 --use_cute \
--batch_size 1 --num_heads 12 --head_dim 128 --q_seq_lens 39936 --topk 40
```
### TurboDiffusion Kernels
+20 -48
View File
@@ -2,9 +2,8 @@
"""
Benchmark VSA *wrapper* performance (forward + backward) and report TFLOPs.
This script benchmarks the autograd-enabled wrappers:
- 64-token TK/Triton: fastvideo_kernel.block_sparse_attn.block_sparse_attn
- 128/256-token Triton/CuTe: fastvideo_kernel.block_sparse_attn_256
This script benchmarks the autograd-enabled wrapper:
- fastvideo_kernel.block_sparse_attn.block_sparse_attn
So measured time includes wrapper overhead (map->index conversion, dispatch) plus kernel time.
"""
@@ -24,6 +23,9 @@ try:
except Exception as e: # pragma: no cover
raise ImportError("This benchmark requires triton (for triton.testing.do_bench).") from e
BLOCK_M = 64
BLOCK_N = 64
def set_seed(seed: int = 42) -> None:
random.seed(seed)
@@ -39,11 +41,7 @@ def parse_arguments() -> argparse.Namespace:
p.add_argument("--num_heads", type=int, default=12)
p.add_argument("--head_dim", type=int, default=128, choices=[64, 128])
p.add_argument("--topk", type=int, default=None, help="KV blocks per Q block (default: ~90%% sparsity)")
p.add_argument("--q_seq_lens",
type=int,
nargs="+",
default=[49152],
help="Q sequence lengths (must be divisible by --block_size)")
p.add_argument("--q_seq_lens", type=int, nargs="+", default=[49152], help="Q sequence lengths (must be /64)")
p.add_argument("--kv_seq_lens",
type=int,
nargs="+",
@@ -53,13 +51,9 @@ def parse_arguments() -> argparse.Namespace:
p.add_argument("--rep", type=int, default=20)
p.add_argument("--seed", type=int, default=42)
p.add_argument("--dtype", type=str, default="bf16", choices=["bf16", "fp16"])
p.add_argument("--block_size", type=int, default=64, choices=[64, 128, 256])
p.add_argument("--force_triton",
action="store_true",
help="Force wrapper to use Triton path (if supported by shapes).")
p.add_argument("--use_cute",
action="store_true",
help="Use the optional FA4 CuTe forward/backward path (requires --block_size 128 or 256).")
return p.parse_args()
@@ -90,38 +84,18 @@ def bench_ms(fn: Callable[[], object], warmup: int, rep: int) -> float:
return do_bench(fn, warmup=warmup, rep=rep, quantiles=None)
def _configure_backend(args: argparse.Namespace) -> None:
if args.use_cute and args.block_size not in (128, 256):
raise ValueError("--use_cute requires --block_size 128 or 256")
if args.use_cute and args.force_triton:
raise ValueError("--use_cute and --force_triton are mutually exclusive")
if args.force_triton:
os.environ.pop("FASTVIDEO_VSA_CUTEDSL", None)
os.environ["FASTVIDEO_KERNEL_VSA_FORCE_TRITON"] = "1"
elif args.use_cute:
os.environ.pop("FASTVIDEO_VSA_TRITON", None)
os.environ.pop("FASTVIDEO_KERNEL_VSA_FORCE_TRITON", None)
os.environ["FASTVIDEO_VSA_CUTEDSL"] = "1"
def main() -> None:
args = parse_arguments()
set_seed(args.seed)
_configure_backend(args)
dtype = torch.bfloat16 if args.dtype == "bf16" else torch.float16
if args.force_triton:
os.environ["FASTVIDEO_KERNEL_VSA_FORCE_TRITON"] = "1"
from fastvideo_kernel.block_sparse_attn import block_sparse_attn
from fastvideo_kernel.block_sparse_attn_256 import block_sparse_attn_128, block_sparse_attn_256
bs, h, d = args.batch_size, args.num_heads, args.head_dim
block_size = args.block_size
attention = {
64: block_sparse_attn,
128: block_sparse_attn_128,
256: block_sparse_attn_256,
}[block_size]
kv_seq_lens = args.kv_seq_lens
if kv_seq_lens is None:
kv_seq_lens = args.q_seq_lens
@@ -131,22 +105,20 @@ def main() -> None:
print("VSA Block-Sparse Attention Benchmark (WRAPPER)")
print(f"device: {torch.cuda.get_device_name(0)}")
print(f"batch={bs}, heads={h}, head_dim={d}, dtype={args.dtype}")
print(f"block_size={block_size}")
print(f"BLOCK_M={BLOCK_M}, BLOCK_N={BLOCK_N}")
print("NOTE: timings include wrapper overhead (map->index + dispatch).")
if args.use_cute:
print("dispatch: FA4 CuTe")
elif args.force_triton:
if args.force_triton:
print("dispatch: forced Triton (FASTVIDEO_KERNEL_VSA_FORCE_TRITON=1)")
else:
print("dispatch: SM90 if available, else Triton")
for q_len, kv_len in zip(args.q_seq_lens, kv_seq_lens):
if q_len % block_size != 0 or kv_len % block_size != 0:
print(f"[skip] q_len={q_len}, kv_len={kv_len} must be divisible by {block_size}")
if q_len % BLOCK_M != 0 or kv_len % BLOCK_N != 0:
print(f"[skip] q_len={q_len}, kv_len={kv_len} must be divisible by 64")
continue
num_q_blocks = q_len // block_size
num_kv_blocks = kv_len // block_size
num_q_blocks = q_len // BLOCK_M
num_kv_blocks = kv_len // BLOCK_N
topk = args.topk if args.topk is not None else max(1, num_kv_blocks // 10)
topk = min(topk, num_kv_blocks)
@@ -157,11 +129,11 @@ def main() -> None:
q, k, v = create_qkv(bs, h, q_len, kv_len, d, dtype)
block_map = make_block_map(bs, h, num_q_blocks, num_kv_blocks, topk)
# Variable block sizes: default full logical blocks.
variable_block_sizes = torch.full((num_kv_blocks, ), block_size, dtype=torch.int32, device="cuda")
# Variable block sizes: default full blocks (64 tokens per KV block)
variable_block_sizes = torch.full((num_kv_blocks, ), BLOCK_N, dtype=torch.int32, device="cuda")
def _fwd():
return attention(q, k, v, block_map, variable_block_sizes)
return block_sparse_attn(q, k, v, block_map, variable_block_sizes)
fwd_ms = bench_ms(_fwd, warmup=args.warmup, rep=args.rep)
@@ -170,7 +142,7 @@ def main() -> None:
q_ = q.detach().requires_grad_(True)
k_ = k.detach().requires_grad_(True)
v_ = v.detach().requires_grad_(True)
o_, _aux_ = attention(q_, k_, v_, block_map, variable_block_sizes)
o_, _aux_ = block_sparse_attn(q_, k_, v_, block_map, variable_block_sizes)
og = torch.randn_like(o_)
loss = (o_ * og).sum()
@@ -184,7 +156,7 @@ def main() -> None:
rep=max(5, args.rep // 2),
)
flops = flops_sparse_attention(bs, h, d, q_len, topk, block_size)
flops = flops_sparse_attention(bs, h, d, q_len, topk, BLOCK_N)
fwd_tflops = flops / fwd_ms * 1e-12 * 1e3
# Rough backward multiplier (attention backward typically ~2-3x forward)
bwd_tflops = (2.5 * flops) / bwd_ms * 1e-12 * 1e3
@@ -1,7 +0,0 @@
// block_sparse_blk128_sm100a.cu -- the 128-token-block instantiation of the torch binding.
//
// Same source as block_sparse_sm100a.cu with VSA_BLK128 set: the kernel and launch land in
// namespace vsa_blk128 (distinct symbols, no ODR clash with the blk64 objects) and the
// exported entry point becomes block_sparse_sm100a_blk128_fwd.
#define VSA_BLK128 true
#include "block_sparse_sm100a.cu"
File diff suppressed because it is too large Load Diff
@@ -1,201 +0,0 @@
#ifndef BLOCK_SPARSE_VSA_LAUNCH_SM100A_CUH
#define BLOCK_SPARSE_VSA_LAUNCH_SM100A_CUH
// Launch surface for the sm_100a VSA block-sparse FMHA forward.
//
// Everything a caller needs: a POD argument struct, a predicate saying whether this build can
// run those arguments, and one launch entry point. The benchmark in
// block_sparse_bench_sm100a.cu and the torch binding both go through here, so there is
// one tensormap construction and one launch configuration rather than two that can drift.
//
// Two compile-time knobs select the four builds:
// VSA_BLK128 false -> 64-token sparse blocks, true -> 128-token
// VSA_BHSD false -> [token][head][dim] (BSHD), true -> [batch][head][token][dim] (BHSD)
#include "block_sparse_kernel_sm100a.cuh"
namespace VSA_NAMESPACE {
struct BlockSparseVsaArgs {
const __nv_bfloat16* q;
const __nv_bfloat16* k;
const __nv_bfloat16* v; // natural layout; only blk128 reads it (blk64 still needs v_t)
const __nv_bfloat16* v_t; // unused: kept so the bench's V_T buffer still binds
__nv_bfloat16* o;
float* lse; // [batch, num_heads, seqlen] fp32, or nullptr
const int* q2k_idx; // [batch*num_heads*num_blocks, max_kv] int32
const int* q2k_num; // [batch*num_heads*num_blocks] int32
const int* variable_block_sizes; // [num_blocks] int32, valid tokens per block
int batch;
int num_heads;
int seqlen;
int head_dim;
int num_blocks;
int max_kv;
float sm_scale;
};
// cudaSuccess iff this build can run `a`. Deliberately conservative: the caller is expected
// to fall back to its own implementation rather than get a wrong answer.
__host__ inline cudaError_t block_sparse_supported(const BlockSparseVsaArgs& a) {
if (a.head_dim != HEAD_DIM) return cudaErrorInvalidValue; // compile-time in the kernel
if (a.num_blocks % 2 != 0) return cudaErrorInvalidValue; // a CTA owns an adjacent pair
if (a.seqlen != a.num_blocks * BLOCK) return cudaErrorInvalidValue;
if (a.max_kv < 1 || a.num_blocks < 1) return cudaErrorInvalidValue;
if (a.q == nullptr || a.k == nullptr || a.o == nullptr) return cudaErrorInvalidValue;
if (a.q2k_idx == nullptr || a.q2k_num == nullptr) return cudaErrorInvalidValue;
// FastVideo always supplies this; without it padded keys would be attended as real zeros.
if (a.variable_block_sizes == nullptr) return cudaErrorInvalidValue;
// V is read MN-major at BOTH block sizes now, so no pre-transposed V_T is ever needed.
if (a.v == nullptr) return cudaErrorInvalidValue;
return cudaSuccess;
}
__host__ inline cudaError_t launch_block_sparse_sm100a(const BlockSparseVsaArgs& a,
cudaStream_t stream) {
const cudaError_t sup = block_sparse_supported(a);
if (sup != cudaSuccess) return sup;
const int B = a.batch, H = a.num_heads, S = a.seqlen, hd = a.head_dim;
const int num_blocks = a.num_blocks, max_kv = a.max_kv;
const long tq = (long)B * S;
const int packed_mtiles_per_seq = num_blocks / 2;
const int total_work = B * H * packed_mtiles_per_seq;
constexpr bool BHSD = VSA_BHSD;
CUtensorMap tq_, tk_, tvt_, tv_, to_;
{
uint64_t gd[4] = { (uint64_t)SUB_COLS_BF16, BHSD ? (uint64_t)((long)B * H) : (uint64_t)H,
BHSD ? (uint64_t)S : (uint64_t)tq, (uint64_t)Q_SUBTILES };
uint64_t gs[3] = { BHSD ? (uint64_t)((long)S * hd) * 2u : (uint64_t)hd * 2u,
BHSD ? (uint64_t)hd * 2u : (uint64_t)((long)H * hd) * 2u,
(uint64_t)SUB_COLS_BF16 * 2u };
uint32_t bd[4] = { (uint32_t)SUB_COLS_BF16, 1u, (uint32_t)M_TILE, (uint32_t)Q_SUBTILES };
uint32_t es[4] = { 1u, 1u, 1u, 1u };
if (cuTensorMapEncodeTiled(&tq_, CU_TENSOR_MAP_DATA_TYPE_BFLOAT16, 4,
const_cast<__nv_bfloat16*>(a.q), gd, gs, bd, es,
CU_TENSOR_MAP_INTERLEAVE_NONE, CU_TENSOR_MAP_SWIZZLE_128B,
CU_TENSOR_MAP_L2_PROMOTION_L2_128B,
CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE) != CUDA_SUCCESS)
return cudaErrorInvalidValue;
if (cuTensorMapEncodeTiled(&to_, CU_TENSOR_MAP_DATA_TYPE_BFLOAT16, 4, a.o, gd, gs, bd, es,
CU_TENSOR_MAP_INTERLEAVE_NONE, CU_TENSOR_MAP_SWIZZLE_128B,
CU_TENSOR_MAP_L2_PROMOTION_L2_128B,
CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE) != CUDA_SUCCESS)
return cudaErrorInvalidValue;
}
{
uint64_t gd[4] = { (uint64_t)SUB_COLS_BF16,
BHSD ? (uint64_t)S : (uint64_t)tq,
BHSD ? (uint64_t)(hd / SUB_COLS_BF16)
: (uint64_t)((long)H * hd / SUB_COLS_BF16),
(uint64_t)((long)B * H) };
uint64_t gs[3] = { BHSD ? (uint64_t)hd * 2u : (uint64_t)((long)H * hd) * 2u,
(uint64_t)SUB_COLS_BF16 * 2u,
(uint64_t)((long)S * hd) * 2u };
uint32_t bd[4] = { (uint32_t)SUB_COLS_BF16, (uint32_t)BLOCK,
BLK128 ? (uint32_t)K_SUBTILES : 1u, 1u };
uint32_t es[4] = { 1u, 1u, 1u, 1u };
if (cuTensorMapEncodeTiled(&tk_, CU_TENSOR_MAP_DATA_TYPE_BFLOAT16, BHSD ? 4 : 3,
const_cast<__nv_bfloat16*>(a.k), gd, gs, bd, es,
CU_TENSOR_MAP_INTERLEAVE_NONE, CU_TENSOR_MAP_SWIZZLE_128B,
CU_TENSOR_MAP_L2_PROMOTION_L2_128B,
CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE) != CUDA_SUCCESS)
return cudaErrorInvalidValue;
// V map is byte-for-byte the K map over a.v: MN-major V needs no transpose (blk128).
const __nv_bfloat16* vbase = a.v ? a.v : a.k;
if (cuTensorMapEncodeTiled(&tv_, CU_TENSOR_MAP_DATA_TYPE_BFLOAT16, BHSD ? 4 : 3,
const_cast<__nv_bfloat16*>(vbase), gd, gs, bd, es,
CU_TENSOR_MAP_INTERLEAVE_NONE, CU_TENSOR_MAP_SWIZZLE_128B,
CU_TENSOR_MAP_L2_PROMOTION_L2_128B,
CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE) != CUDA_SUCCESS)
return cudaErrorInvalidValue;
}
// V_T map: blk64 only. Unused at blk128 but must still be a valid tensormap to pass by value.
{
const __nv_bfloat16* vt = a.v_t ? a.v_t : a.k;
if constexpr (BLK128) {
uint64_t gd[3] = { (uint64_t)SUB_COLS_BF16, (uint64_t)((long)H * hd),
(uint64_t)((long)tq / SUB_COLS_BF16) };
uint64_t gs[2] = { (uint64_t)tq * 2u, (uint64_t)SUB_COLS_BF16 * 2u };
uint32_t bd[3] = { (uint32_t)SUB_COLS_BF16, (uint32_t)hd, (uint32_t)V_SUBTILES };
uint32_t es[3] = { 1u, 1u, 1u };
if (cuTensorMapEncodeTiled(&tvt_, CU_TENSOR_MAP_DATA_TYPE_BFLOAT16, 3,
const_cast<__nv_bfloat16*>(vt), gd, gs, bd, es,
CU_TENSOR_MAP_INTERLEAVE_NONE, CU_TENSOR_MAP_SWIZZLE_128B,
CU_TENSOR_MAP_L2_PROMOTION_L2_128B,
CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE) != CUDA_SUCCESS)
return cudaErrorInvalidValue;
} else {
if (make_tma_2d_tiled(&tvt_, const_cast<__nv_bfloat16*>(vt), (long)H * hd, (int)tq, hd,
SUB_COLS_BF16, 2, CU_TENSOR_MAP_DATA_TYPE_BFLOAT16,
CU_TENSOR_MAP_SWIZZLE_128B) != cudaSuccess)
return cudaErrorInvalidValue;
}
}
const size_t smem =
(size_t)2 * Q_TILE_BYTES + NUM_KV_STAGES * KV_RING_SLOT_BYTES
+ (size_t)2 * M_TILE * HEAD_DIM * sizeof(__nv_bfloat16)
+ (2 * NUM_KV_STAGES + 22) * 8
+ (size_t)CLC_STAGES * (2 * 8 + 16) + 16
+ 8
+ (size_t)2 * STAT_REGIONS * STATS * sizeof(float)
+ 256;
#ifndef VSA_NAMED_BAR
#define VSA_NAMED_BAR false
#endif
#ifndef VSA_THROTTLE
#define VSA_THROTTLE false
#endif
#ifndef VSA_USE_CLC
#define VSA_USE_CLC true
#endif
constexpr bool FULL_NAMED_BAR = VSA_NAMED_BAR, EX2_EMU = true, SPLIT_P = true,
SOFTMAX_THROTTLE = VSA_THROTTLE, USE_CLC = VSA_USE_CLC,
Q_RASTER = true, MHA = true;
auto kfn = &fmha_context_bf16_gen_kernel<32, FULL_NAMED_BAR, EX2_EMU, SPLIT_P,
SOFTMAX_THROTTLE, USE_CLC, Q_RASTER, MHA,
/*RESCALE_THRESHOLD=*/8, /*BHSD=*/VSA_BHSD>;
cudaError_t e = cudaFuncSetAttribute(kfn, cudaFuncAttributeMaxDynamicSharedMemorySize, (int)smem);
if (e != cudaSuccess) return e;
const unsigned long long magic0 = make_magic((unsigned)(H * packed_mtiles_per_seq));
const unsigned long long magic1 = make_magic((unsigned)H);
const unsigned long long magic2 = make_magic((unsigned)packed_mtiles_per_seq);
const float scale_log2 = a.sm_scale * (float)M_LOG2E;
int numSM = 0;
e = cudaDeviceGetAttribute(&numSM, cudaDevAttrMultiProcessorCount, 0);
if (e != cudaSuccess) return e;
const int num_ctas = USE_CLC ? total_work : (total_work < numSM ? total_work : numSM);
dim3 grid(num_ctas, 1, 1), block(N_WARPS * 32, 1, 1);
if (USE_CLC) {
cudaLaunchConfig_t cfg = {};
cfg.gridDim = grid; cfg.blockDim = block; cfg.dynamicSmemBytes = smem; cfg.stream = stream;
cudaLaunchAttribute cfgAttr[1];
cfgAttr[0].id = cudaLaunchAttributeClusterDimension;
cfgAttr[0].val.clusterDim.x = 1; cfgAttr[0].val.clusterDim.y = 1;
cfgAttr[0].val.clusterDim.z = 1;
cfg.attrs = cfgAttr; cfg.numAttrs = 1;
return cudaLaunchKernelEx(&cfg, kfn, tq_, tk_, tvt_, tv_, to_, S, H, scale_log2, B,
num_blocks, packed_mtiles_per_seq, max_kv, magic0, magic1, magic2,
a.q2k_idx, a.q2k_num, a.variable_block_sizes, a.lse);
}
kfn<<<grid, block, smem, stream>>>(tq_, tk_, tvt_, tv_, to_, S, H, scale_log2, B, num_blocks,
packed_mtiles_per_seq, max_kv, magic0, magic1, magic2,
a.q2k_idx, a.q2k_num, a.variable_block_sizes, a.lse);
return cudaGetLastError();
}
} // namespace VSA_NAMESPACE
// Callers (the bench, the torch binding) keep using unqualified names; each translation unit
// only ever sees the one configuration its VSA_BLK128 selected.
using namespace VSA_NAMESPACE;
#endif // BLOCK_SPARSE_VSA_LAUNCH_SM100A_CUH
@@ -1,114 +0,0 @@
// block_sparse_sm100a.cu -- torch binding for the sm_100a VSA block-sparse FMHA forward.
//
// Forward only: returns (out, lse) so FastVideo's existing Triton backward keeps working
// unchanged. lse is exactly the M tensor triton_block_sparse_attn_forward writes --
// max(qk * qk_scale) + log2(l), [B, H, S] fp32 -- which is what lets
// block_sparse_attn_backward_triton run against our forward untouched.
//
// The build is fixed at compile time by two flags, so one extension carries one configuration:
// VSA_BLK128 false -> 64-token sparse blocks, true -> 128-token
// VSA_BHSD false -> [B, S, H, D], true -> [B, H, S, D]
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAGuard.h>
#include "block_sparse_launch_sm100a.cuh"
namespace {
void check_qkv(const torch::Tensor& t, const char* name, int64_t B, int64_t H, int64_t S,
int64_t D) {
TORCH_CHECK(t.is_cuda(), name, " must be a CUDA tensor");
TORCH_CHECK(t.scalar_type() == at::kBFloat16, name, " must be bfloat16, got ", t.scalar_type());
TORCH_CHECK(t.dim() == 4, name, " must be 4-D, got ", t.dim(), " dims");
TORCH_CHECK(t.is_contiguous(), name, " must be contiguous");
if (VSA_BHSD) {
TORCH_CHECK(t.size(0) == B && t.size(1) == H && t.size(2) == S && t.size(3) == D, name,
" has shape ", t.sizes(), ", expected [", B, ",", H, ",", S, ",", D, "]");
} else {
TORCH_CHECK(t.size(0) == B && t.size(1) == S && t.size(2) == H && t.size(3) == D, name,
" has shape ", t.sizes(), ", expected [", B, ",", S, ",", H, ",", D, "]");
}
}
void check_index(const torch::Tensor& t, const char* name) {
TORCH_CHECK(t.is_cuda(), name, " must be a CUDA tensor");
TORCH_CHECK(t.scalar_type() == at::kInt, name, " must be int32, got ", t.scalar_type());
TORCH_CHECK(t.is_contiguous(), name, " must be contiguous");
}
} // namespace
// The exported symbol carries the block size: block_sparse_sm100a_fwd is the 64-token build,
// block_sparse_sm100a_blk128_fwd the 128-token one (block_sparse_blk128_sm100a.cu re-includes
// this file with VSA_BLK128 set). The python backend picks by the metadata's block size.
#if VSA_BLK128
#define BLOCK_SPARSE_SM100A_FWD block_sparse_sm100a_blk128_fwd
#else
#define BLOCK_SPARSE_SM100A_FWD block_sparse_sm100a_fwd
#endif
// Returns {out} or {out, lse}. Layout of out matches the inputs.
std::vector<torch::Tensor> BLOCK_SPARSE_SM100A_FWD(torch::Tensor q, torch::Tensor k,
torch::Tensor v,
c10::optional<torch::Tensor> v_t,
torch::Tensor q2k_idx,
torch::Tensor q2k_num,
torch::Tensor variable_block_sizes,
double sm_scale, bool need_lse) {
const at::cuda::OptionalCUDAGuard guard(device_of(q));
const int64_t B = q.size(0);
const int64_t H = VSA_BHSD ? q.size(1) : q.size(2);
const int64_t S = VSA_BHSD ? q.size(2) : q.size(1);
const int64_t D = q.size(3);
check_qkv(q, "q", B, H, S, D);
check_qkv(k, "k", B, H, S, D);
check_qkv(v, "v", B, H, S, D);
check_index(q2k_idx, "q2k_idx");
check_index(q2k_num, "q2k_num");
check_index(variable_block_sizes, "variable_block_sizes");
const int64_t num_blocks = variable_block_sizes.numel();
const int64_t max_kv = q2k_idx.size(-1);
TORCH_CHECK(S == num_blocks * BLOCK, "seqlen ", S, " must equal num_blocks (", num_blocks,
") * ", BLOCK, "; FastVideo pads the sequence up to whole blocks");
auto out = torch::empty_like(q);
torch::Tensor lse;
if (need_lse) lse = torch::empty({B, H, S}, q.options().dtype(torch::kFloat32));
BlockSparseVsaArgs a{};
a.q = reinterpret_cast<const __nv_bfloat16*>(q.data_ptr());
a.k = reinterpret_cast<const __nv_bfloat16*>(k.data_ptr());
a.v = reinterpret_cast<const __nv_bfloat16*>(v.data_ptr());
a.v_t = v_t.has_value() ? reinterpret_cast<const __nv_bfloat16*>(v_t->data_ptr()) : nullptr;
a.o = reinterpret_cast<__nv_bfloat16*>(out.data_ptr());
a.lse = need_lse ? lse.data_ptr<float>() : nullptr;
a.q2k_idx = q2k_idx.data_ptr<int>();
a.q2k_num = q2k_num.data_ptr<int>();
a.variable_block_sizes = variable_block_sizes.data_ptr<int>();
a.batch = (int)B;
a.num_heads = (int)H;
a.seqlen = (int)S;
a.head_dim = (int)D;
a.num_blocks = (int)num_blocks;
a.max_kv = (int)max_kv;
a.sm_scale = (float)sm_scale;
// Report an unsupported regime loudly rather than returning plausible-looking wrong values.
TORCH_CHECK(block_sparse_supported(a) == cudaSuccess,
"block_sparse_sm100a: unsupported configuration -- requires head_dim==",
HEAD_DIM, ", an even num_blocks, seqlen == num_blocks*", BLOCK,
", and a variable_block_sizes tensor. Got head_dim=", D, " num_blocks=",
num_blocks, " seqlen=", S);
const cudaError_t err = launch_block_sparse_sm100a(a, at::cuda::getCurrentCUDAStream());
TORCH_CHECK(err == cudaSuccess,
"block_sparse_sm100a launch failed: ", cudaGetErrorString(err));
if (need_lse) return {out, lse};
return {out};
}
@@ -1,877 +0,0 @@
// primitives.cuh -- device primitives for the sm_100a VSA block-sparse attention
// forward: tcgen05 (alloc / mma / ld / st / commit / wait / fence), TMA load / store /
// tensormap, mbarrier, cluster launch control, setmaxnreg, fast math, and the FMHA helpers.
//
// Generated and pruned to what the kernel reaches -- do not edit by hand.
#pragma once
#include <cstdint>
#include <cstdio>
#include <cstdlib>
#include <cuda.h>
#include <cuda_bf16.h>
#include <cuda_fp16.h>
#include <cuda_runtime.h>
#include <cmath>
#include <cassert>
#include <cstring>
#include <vector_types.h>
#ifndef CUDA_CHECK
#define CUDA_CHECK(stmt) do { \
cudaError_t _e = (stmt); \
if (_e != cudaSuccess) { \
fprintf(stderr, "CUDA error %s:%d: %s -> %s\n", \
__FILE__, __LINE__, #stmt, cudaGetErrorString(_e)); \
std::exit(1); \
} \
} while (0)
#endif
__device__ __forceinline__
uint64_t mbarrier_arrive(uint32_t mbar_smem) {
uint64_t state;
asm volatile("mbarrier.arrive.shared::cta.b64 %0, [%1];\n"
: "=l"(state) : "r"(mbar_smem) : "memory");
return state;
}
__device__ __forceinline__
void mbarrier_arrive_nostate(uint32_t mbar_smem) {
asm volatile("mbarrier.arrive.shared::cta.b64 _, [%0];\n"
:: "r"(mbar_smem) : "memory");
}
__device__ __forceinline__
void mbarrier_arrive_cluster_default(uint32_t cluster_smem_addr) {
asm volatile("mbarrier.arrive.shared::cluster.b64 _, [%0];\n"
:: "r"(cluster_smem_addr) : "memory");
}
__device__ __forceinline__
void mbarrier_arrive_expect_tx(uint32_t mbar_smem, uint32_t expected_bytes) {
asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;\n"
:: "r"(mbar_smem), "r"(expected_bytes) : "memory");
}
__device__ __forceinline__
void mbarrier_wait_parity_suspend(uint32_t mbar_smem, uint32_t phase_parity) {
asm volatile(
"{\n"
".reg .pred P1;\n"
"LAB_WAIT:\n"
"mbarrier.try_wait.parity.shared::cta.b64 P1, [%0], %1, 10000000;\n"
"@!P1 bra.uni LAB_WAIT;\n"
"}\n"
:: "r"(mbar_smem), "r"(phase_parity) : "memory");
}
__device__ __forceinline__
void mbarrier_wait_parity(uint32_t mbar_smem, uint32_t phase_parity) {
asm volatile(
"{\n"
".reg .pred P1;\n"
"LAB_WAIT_HOT:\n"
"mbarrier.try_wait.parity.shared::cta.b64 P1, [%0], %1;\n"
"@!P1 bra.uni LAB_WAIT_HOT;\n"
"}\n"
:: "r"(mbar_smem), "r"(phase_parity) : "memory");
}
__device__ __forceinline__
void fence_proxy_async_shared_cta() {
asm volatile("fence.proxy.async.shared::cta;\n" ::: "memory");
}
__device__ __forceinline__
void fence_proxy_async_shared() {
asm volatile("fence.proxy.async.shared::cta;\n" ::: "memory");
}
__device__ __forceinline__ void clc_try_cancel_async(
uint32_t smem_dst, uint32_t mbar_smem) {
asm volatile(
"clusterlaunchcontrol.try_cancel.async.shared::cta.mbarrier::complete_tx::bytes.b128"
" [%0], [%1];\n"
:: "r"(smem_dst), "r"(mbar_smem) : "memory");
}
__device__ __forceinline__ void clc_load_response(
uint32_t smem_slot, uint32_t& r0, uint32_t& r1,
uint32_t& r2, uint32_t& r3) {
asm volatile("ld.shared::cta.v4.b32 {%0, %1, %2, %3}, [%4];\n"
: "=r"(r0), "=r"(r1), "=r"(r2), "=r"(r3)
: "r"(smem_slot));
}
template <int NUM_STAGES>
struct MbarrierPhaseTracker {
uint32_t phase[NUM_STAGES];
int idx;
__device__ __forceinline__
void init() {
for (int i = 0; i < NUM_STAGES; ++i) phase[i] = 0;
idx = 0;
}
__device__ __forceinline__
uint32_t current_phase() const { return phase[idx]; }
__device__ __forceinline__
void advance() {
phase[idx] ^= 1u;
idx = (idx + 1) % NUM_STAGES;
}
__device__ __forceinline__
int stage() const { return idx; }
};
template <int NUM_STAGES>
struct PhaseTracker {
int stage;
uint32_t phase;
__device__ __forceinline__
PhaseTracker() : stage(0), phase(0) {}
__device__ __forceinline__
void advance() {
stage++;
if (stage == NUM_STAGES) {
stage = 0;
phase ^= 1;
}
}
__device__ __forceinline__
int get_stage() const { return stage; }
__device__ __forceinline__
uint32_t get_phase() const { return phase; }
};
template <int NUM_STAGES>
struct EmptyPhaseTracker {
int stage;
uint32_t phase;
__device__ __forceinline__
EmptyPhaseTracker() : stage(0), phase(1) {}
__device__ __forceinline__
void advance() {
stage++;
if (stage == NUM_STAGES) {
stage = 0;
phase ^= 1;
}
}
__device__ __forceinline__
int get_stage() const { return stage; }
__device__ __forceinline__
uint32_t get_phase() const { return phase; }
};
template <int STAGES>
__device__ __forceinline__
void advance_stage_phase(int& stage, uint32_t& phase) {
++stage;
if (stage == STAGES) {
stage = 0;
phase ^= 1u;
}
}
static constexpr uint32_t SM100_CLC_PEER_MASK = 0xFEFFFFFF;
struct ClcTileInfo {
int m_tile;
int n_tile;
bool valid;
};
enum class ClcRasterOrder { AlongN, AlongM };
__device__ __forceinline__
void clc_arrive_expect_tx_cta(uint32_t clc_full_local_addr, uint32_t tx_bytes) {
if ((threadIdx.x & 31) == 0) {
mbarrier_arrive_expect_tx(clc_full_local_addr, tx_bytes);
}
}
__device__ __forceinline__
void clc_consumer_release(uint32_t clc_empty_local_addr) {
uint32_t peer0_addr = clc_empty_local_addr & SM100_CLC_PEER_MASK;
mbarrier_arrive_cluster_default(peer0_addr);
}
__device__ __forceinline__
void clc_consumer_release_cta(uint32_t clc_empty_local_addr) {
mbarrier_arrive_nostate(clc_empty_local_addr);
}
template <int CLUSTER_SHAPE_M, int CLUSTER_SHAPE_N, ClcRasterOrder ORDER>
__device__ __forceinline__
ClcTileInfo clc_parse_response(uint32_t resp_smem_addr) {
uint32_t d0, d1, d2, d3;
fence_proxy_async_shared_cta();
clc_load_response(resp_smem_addr, d0, d1, d2, d3);
const int ctaid_x = static_cast<int>(d0);
const int ctaid_y = static_cast<int>(d1 & 0xFFFFu);
const bool valid = (d2 & 1u) != 0u;
(void)d3;
ClcTileInfo info;
info.valid = valid;
if constexpr (ORDER == ClcRasterOrder::AlongN) {
info.m_tile = ctaid_y / CLUSTER_SHAPE_M;
info.n_tile = ctaid_x / CLUSTER_SHAPE_N;
} else {
info.m_tile = ctaid_x / CLUSTER_SHAPE_M;
info.n_tile = ctaid_y / CLUSTER_SHAPE_N;
}
return info;
}
template <int CLUSTER_SHAPE_M, int CLUSTER_SHAPE_N, ClcRasterOrder ORDER,
int CTA_GROUP = 2, bool SUSPEND = false>
__device__ __forceinline__
ClcTileInfo clc_fetch_next_tile(
uint64_t* clc_full_bar, uint64_t* clc_empty_bar, uint32_t* clc_response,
int clc_cons_stage, uint32_t clc_cons_phase, bool do_release) {
uint32_t full_addr = static_cast<uint32_t>(
__cvta_generic_to_shared(&clc_full_bar[clc_cons_stage]));
if constexpr (SUSPEND) mbarrier_wait_parity_suspend(full_addr, clc_cons_phase);
else mbarrier_wait_parity(full_addr, clc_cons_phase);
uint32_t resp_addr = static_cast<uint32_t>(
__cvta_generic_to_shared(&clc_response[clc_cons_stage * 4]));
ClcTileInfo t = clc_parse_response<
CLUSTER_SHAPE_M, CLUSTER_SHAPE_N, ORDER>(resp_addr);
if (do_release) {
uint32_t empty_local = static_cast<uint32_t>(
__cvta_generic_to_shared(&clc_empty_bar[clc_cons_stage]));
if constexpr (CTA_GROUP == 1) {
clc_consumer_release_cta(empty_local);
} else {
clc_consumer_release(empty_local);
}
}
return t;
}
template <int STAGES = 2>
__device__ __forceinline__
void clc_fetch_next_tile_advance(int& clc_cons_stage,
uint32_t& clc_cons_phase) {
advance_stage_phase<STAGES>(clc_cons_stage, clc_cons_phase);
}
__device__ __forceinline__ unsigned fdiv(unsigned n, unsigned long long pk) {
unsigned M = (unsigned)pk;
if (M == 0u) return n;
return __umulhi(n, M) >> (unsigned)(pk >> 32);
}
__host__ inline unsigned long long make_magic(unsigned d) {
if (d <= 1u) return 0ULL;
unsigned l = 0; while ((1u << (l + 1)) <= d) ++l;
unsigned p = 31u + l;
unsigned long long m = ((1ull << p) + (unsigned long long)d - 1ull) / d;
return (m & 0xffffffffULL) | ((unsigned long long)(p - 32u) << 32);
}
template <int BARRIER_ID>
__device__ __forceinline__
void bar_sync(uint32_t thread_count) {
static_assert(BARRIER_ID >= 0 && BARRIER_ID <= 15,
"bar.sync: BARRIER_ID must be in [0, 15]");
asm volatile("bar.sync %0, %1;\n"
:: "n"(BARRIER_ID), "r"(thread_count) : "memory");
}
template <int BARRIER_ID>
__device__ __forceinline__
void bar_arrive(uint32_t thread_count) {
static_assert(BARRIER_ID >= 0 && BARRIER_ID <= 15,
"bar.arrive: BARRIER_ID must be in [0, 15]");
asm volatile("bar.arrive %0, %1;\n"
:: "n"(BARRIER_ID), "r"(thread_count) : "memory");
}
__device__ __forceinline__ void full_bar_arrive(int m_tile, int band) {
switch (1 + m_tile * 4 + band) {
case 1: bar_arrive<1>(64); break;
case 2: bar_arrive<2>(64); break;
case 3: bar_arrive<3>(64); break;
case 4: bar_arrive<4>(64); break;
case 5: bar_arrive<5>(64); break;
case 6: bar_arrive<6>(64); break;
case 7: bar_arrive<7>(64); break;
case 8: bar_arrive<8>(64); break;
}
}
__device__ __forceinline__ void full_bar_wait(int m_tile, int band) {
switch (1 + m_tile * 4 + band) {
case 1: bar_sync<1>(64); break;
case 2: bar_sync<2>(64); break;
case 3: bar_sync<3>(64); break;
case 4: bar_sync<4>(64); break;
case 5: bar_sync<5>(64); break;
case 6: bar_sync<6>(64); break;
case 7: bar_sync<7>(64); break;
case 8: bar_sync<8>(64); break;
}
}
template <bool IS_CAUSAL, int K_TILE>
__device__ __forceinline__ void mask_s_row_r2p(float* scores, int k_offset, int q_pos, int seqlen_k) {
int n_keep = seqlen_k - k_offset;
if constexpr (IS_CAUSAL) {
const int causal = q_pos - k_offset + 1;
n_keep = n_keep < causal ? n_keep : causal;
}
#pragma unroll
for (int s = 0; s < K_TILE / 32; ++s) {
int m = (s + 1) * 32 - n_keep;
m = m < 0 ? 0 : (m > 32 ? 32 : m);
const uint32_t keep = (m >= 32) ? 0u : (0xFFFFFFFFu >> m);
#pragma unroll
for (int i = 0; i < 32; ++i)
if (!(keep & (1u << i))) scores[s * 32 + i] = -INFINITY;
}
}
template <int CTA_GROUP>
__device__ __forceinline__ void tcgen05_alloc(uint32_t smem_dst_ptr,
uint32_t n_cols) {
static_assert(CTA_GROUP == 1 || CTA_GROUP == 2,
"tcgen05_alloc: CTA_GROUP must be 1 or 2");
if constexpr (CTA_GROUP == 1) {
asm volatile(
"tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;\n"
:: "r"(smem_dst_ptr), "r"(n_cols));
} else {
asm volatile(
"tcgen05.alloc.cta_group::2.sync.aligned.shared::cta.b32 [%0], %1;\n"
:: "r"(smem_dst_ptr), "r"(n_cols));
}
}
__device__ __forceinline__ void tcgen05_st_32x32b_x16(
uint32_t tmem_addr, const uint32_t (&r)[16]) {
asm volatile("tcgen05.st.sync.aligned.32x32b.x16.b32 "
"[%0], {%1,%2,%3,%4,%5,%6,%7,%8,%9,%10,"
"%11,%12,%13,%14,%15,%16};\n"
:: "r"(tmem_addr),
"r"(r[0]),"r"(r[1]),"r"(r[2]),"r"(r[3]),
"r"(r[4]),"r"(r[5]),"r"(r[6]),"r"(r[7]),
"r"(r[8]),"r"(r[9]),"r"(r[10]),"r"(r[11]),
"r"(r[12]),"r"(r[13]),"r"(r[14]),"r"(r[15]));
}
__device__ __forceinline__ void tcgen05_st_32x32b_x32(
uint32_t tmem_addr, const uint32_t (&r)[32]) {
asm volatile("tcgen05.st.sync.aligned.32x32b.x32.b32 "
"[%0], {%1,%2,%3,%4,%5,%6,%7,%8,%9,%10,"
"%11,%12,%13,%14,%15,%16,%17,%18,%19,%20,"
"%21,%22,%23,%24,%25,%26,%27,%28,%29,%30,"
"%31,%32};\n"
:: "r"(tmem_addr),
"r"(r[0]),"r"(r[1]),"r"(r[2]),"r"(r[3]),
"r"(r[4]),"r"(r[5]),"r"(r[6]),"r"(r[7]),
"r"(r[8]),"r"(r[9]),"r"(r[10]),"r"(r[11]),
"r"(r[12]),"r"(r[13]),"r"(r[14]),"r"(r[15]),
"r"(r[16]),"r"(r[17]),"r"(r[18]),"r"(r[19]),
"r"(r[20]),"r"(r[21]),"r"(r[22]),"r"(r[23]),
"r"(r[24]),"r"(r[25]),"r"(r[26]),"r"(r[27]),
"r"(r[28]),"r"(r[29]),"r"(r[30]),"r"(r[31]));
}
__device__ __forceinline__ void tcgen05_commit1_lead(uint32_t lead, uint32_t mbar_smem_addr) {
asm volatile(
"{\n\t"
".reg .pred q;\n\t"
"setp.ne.b32 q, %0, 0;\n\t"
"@q tcgen05.commit.cta_group::1.mbarrier::arrive::one.b64 [%1];\n\t"
"}\n"
:: "r"(lead), "r"(mbar_smem_addr));
}
__device__ __forceinline__ void tcgen05_wait_st() {
asm volatile("tcgen05.wait::st.sync.aligned;\n" ::: "memory");
}
__device__ __forceinline__ void tcgen05_fence_before_thread_sync() {
asm volatile("tcgen05.fence::before_thread_sync;\n" ::: "memory");
}
__device__ __forceinline__
void tma_load_2d(uint32_t smem_dst, const void* tensormap_ptr,
uint32_t mbar_smem, int coord_x, int coord_y) {
asm volatile(
"cp.async.bulk.tensor.2d.shared::cluster.global.tile"
".mbarrier::complete_tx::bytes"
" [%0], [%1, {%3, %4}], [%2];\n"
:: "r"(smem_dst), "l"(tensormap_ptr),
"r"(mbar_smem), "r"(coord_x), "r"(coord_y)
: "memory");
}
__device__ __forceinline__
void tma_load_3d(uint32_t smem_dst, const void* tensormap_ptr,
uint32_t mbar_smem, int c0, int c1, int c2) {
asm volatile(
"cp.async.bulk.tensor.3d.shared::cluster.global.tile"
".mbarrier::complete_tx::bytes"
" [%0], [%1, {%3, %4, %5}], [%2];\n"
:: "r"(smem_dst), "l"(tensormap_ptr),
"r"(mbar_smem), "r"(c0), "r"(c1), "r"(c2)
: "memory");
}
__device__ __forceinline__
void tma_load_4d(uint32_t smem_dst, const void* tensormap_ptr,
uint32_t mbar_smem, int c0, int c1, int c2, int c3) {
asm volatile(
"cp.async.bulk.tensor.4d.shared::cluster.global.tile"
".mbarrier::complete_tx::bytes"
" [%0], [%1, {%3, %4, %5, %6}], [%2];\n"
:: "r"(smem_dst), "l"(tensormap_ptr),
"r"(mbar_smem), "r"(c0), "r"(c1), "r"(c2), "r"(c3)
: "memory");
}
template <int CTA_GROUP>
__device__ __forceinline__ void tcgen05_dealloc(uint32_t tmem_addr,
uint32_t n_cols) {
static_assert(CTA_GROUP == 1 || CTA_GROUP == 2,
"tcgen05_dealloc: CTA_GROUP must be 1 or 2");
if constexpr (CTA_GROUP == 1) {
asm volatile(
"tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;\n"
:: "r"(tmem_addr), "r"(n_cols));
} else {
asm volatile(
"tcgen05.dealloc.cta_group::2.sync.aligned.b32 %0, %1;\n"
:: "r"(tmem_addr), "r"(n_cols));
}
}
__device__ __forceinline__
void tma_store_2d(const void* tensormap_ptr, int coord_x, int coord_y,
uint32_t smem_src) {
asm volatile(
"cp.async.bulk.tensor.2d.global.shared::cta.tile.bulk_group"
" [%0, {%1, %2}], [%3];\n"
:: "l"(tensormap_ptr), "r"(coord_x), "r"(coord_y),
"r"(smem_src)
: "memory");
}
__device__ __forceinline__
void tma_store_3d(const void* tensormap_ptr, int c0, int c1, int c2,
uint32_t smem_src) {
asm volatile(
"cp.async.bulk.tensor.3d.global.shared::cta.tile.bulk_group"
" [%0, {%1, %2, %3}], [%4];\n"
:: "l"(tensormap_ptr), "r"(c0), "r"(c1), "r"(c2), "r"(smem_src)
: "memory");
}
__device__ __forceinline__
void tma_store_4d(const void* tensormap_ptr, int c0, int c1, int c2, int c3,
uint32_t smem_src) {
asm volatile(
"cp.async.bulk.tensor.4d.global.shared::cta.tile.bulk_group"
" [%0, {%1, %2, %3, %4}], [%5];\n"
:: "l"(tensormap_ptr), "r"(c0), "r"(c1), "r"(c2), "r"(c3), "r"(smem_src)
: "memory");
}
inline cudaError_t make_tma_2d_tiled(
CUtensorMap* out,
const void* ptr, int rows, int cols, int box_rows, int box_cols,
int elem_bytes, CUtensorMapDataType dtype,
CUtensorMapSwizzle swizzle = CU_TENSOR_MAP_SWIZZLE_128B,
CUtensorMapL2promotion l2 = CU_TENSOR_MAP_L2_PROMOTION_L2_128B,
CUtensorMapFloatOOBfill oob = CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE) {
uint64_t globalDim[2] = { (uint64_t)cols, (uint64_t)rows };
uint64_t globalStrides[1] = { (uint64_t)cols * (uint64_t)elem_bytes };
uint32_t boxDim[2] = { (uint32_t)box_cols, (uint32_t)box_rows };
uint32_t elemStrides[2] = { 1u, 1u };
CUresult r = cuTensorMapEncodeTiled(
out, dtype, 2,
const_cast<void*>(ptr), globalDim, globalStrides,
boxDim, elemStrides,
CU_TENSOR_MAP_INTERLEAVE_NONE,
swizzle, l2, oob);
return (r == CUDA_SUCCESS) ? cudaSuccess : cudaErrorInvalidValue;
}
__device__ __forceinline__
void cp_async_bulk_commit_group() {
asm volatile("cp.async.bulk.commit_group;\n" ::: "memory");
}
template <int N>
__device__ __forceinline__
void cp_async_bulk_wait_group_read() {
asm volatile("cp.async.bulk.wait_group.read %0;\n" :: "n"(N) : "memory");
}
__device__ __forceinline__
void mbarrier_init(uint32_t mbar_smem, uint32_t arrive_count) {
asm volatile("mbarrier.init.shared::cta.b64 [%0], %1;\n"
:: "r"(mbar_smem), "r"(arrive_count) : "memory");
}
template <int CTA_GROUP>
__device__ __forceinline__ void tcgen05_relinquish_alloc_permit() {
static_assert(CTA_GROUP == 1 || CTA_GROUP == 2,
"tcgen05_relinquish_alloc_permit: CTA_GROUP must be 1 or 2");
if constexpr (CTA_GROUP == 1) {
asm volatile("tcgen05.relinquish_alloc_permit.cta_group::1.sync.aligned;\n" ::);
} else {
asm volatile("tcgen05.relinquish_alloc_permit.cta_group::2.sync.aligned;\n" ::);
}
}
__device__ __forceinline__
void fence_mbarrier_init_release_cluster() {
asm volatile("fence.mbarrier_init.release.cluster;\n" ::: "memory");
}
__device__ __forceinline__ void tcgen05_mma_f16_ss_lead(uint32_t lead,
uint32_t tmem_c, uint64_t desc_a, uint64_t desc_b, uint32_t idesc,
bool enable_input_d) {
asm volatile(
"{\n\t"
".reg .pred p, q;\n\t"
"setp.ne.b32 q, %0, 0;\n\t"
"setp.ne.b32 p, %5, 0;\n\t"
"@q tcgen05.mma.cta_group::1.kind::f16 [%1], %2, %3, %4, {%6, %7, %8, %9}, p;\n\t"
"}\n"
:: "r"(lead), "r"(tmem_c), "l"(desc_a), "l"(desc_b), "r"(idesc),
"r"(enable_input_d ? 1u : 0u), "r"(0u), "r"(0u), "r"(0u), "r"(0u));
}
__device__ __forceinline__ void tcgen05_mma_f16_ts_1sm_lead(uint32_t lead,
uint32_t tmem_c, uint32_t tmem_a, uint64_t desc_b, uint32_t idesc,
bool enable_input_d) {
asm volatile(
"{\n\t"
".reg .pred p, q;\n\t"
"setp.ne.b32 q, %0, 0;\n\t"
"setp.ne.b32 p, %5, 0;\n\t"
"@q tcgen05.mma.cta_group::1.kind::f16 [%1], [%2], %3, %4, {%6, %7, %8, %9}, p;\n\t"
"}\n"
:: "r"(lead), "r"(tmem_c), "r"(tmem_a), "l"(desc_b), "r"(idesc),
"r"(enable_input_d ? 1u : 0u), "r"(0u), "r"(0u), "r"(0u), "r"(0u));
}
enum class SmemSwizzleBlackwell : uint32_t {
None = 0,
B128_32atom = 1,
B128 = 2,
B64 = 4,
B32 = 6,
};
__device__ __host__ __forceinline__ uint64_t build_smem_desc_blackwell(
uint32_t smem_addr,
uint32_t stride_byte_offset,
uint32_t leading_byte_offset,
SmemSwizzleBlackwell swizzle = SmemSwizzleBlackwell::B128,
uint32_t base_offset = 0) {
uint64_t d = 0;
d |= static_cast<uint64_t>((smem_addr >> 4) & 0x3FFF);
d |= static_cast<uint64_t>((leading_byte_offset >> 4) & 0x3FFF) << 16;
d |= static_cast<uint64_t>((stride_byte_offset >> 4) & 0x3FFF) << 32;
d |= static_cast<uint64_t>(1) << 46;
d |= (static_cast<uint64_t>(base_offset) & 0x7) << 49;
d |= static_cast<uint64_t>(static_cast<uint32_t>(swizzle) & 0x7) << 61;
return d;
}
__device__ __forceinline__
uint32_t elect_one_sync() {
uint32_t elected;
asm volatile(
"{\n\t"
".reg .pred p;\n\t"
"elect.sync %0|p, 0xffffffff;\n\t"
"selp.b32 %0, 1, 0, p;\n\t"
"}\n"
: "=r"(elected));
return elected;
}
__device__ __forceinline__
uint32_t elect_one_sync(uint32_t membermask) {
uint32_t elected;
asm volatile(
"{\n\t"
".reg .pred p;\n\t"
"elect.sync %0|p, %1;\n\t"
"selp.b32 %0, 1, 0, p;\n\t"
"}\n"
: "=r"(elected) : "r"(membermask));
return elected;
}
template <int N>
__device__ __forceinline__
void setmaxnreg_dec() {
static_assert(N >= 24 && N <= 256 && N % 8 == 0,
"setmaxnreg_dec: N must be in [24, 256] and a multiple of 8");
asm volatile("setmaxnreg.dec.sync.aligned.u32 %0;\n"
:: "n"(N) : "memory");
}
template <int N>
__device__ __forceinline__
void setmaxnreg_inc() {
static_assert(N >= 24 && N <= 256 && N % 8 == 0,
"setmaxnreg_inc: N must be in [24, 256] and a multiple of 8");
asm volatile("setmaxnreg.inc.sync.aligned.u32 %0;\n"
:: "n"(N) : "memory");
}
__device__ __forceinline__
uint32_t cvt_f32x2_to_bf16x2(float a, float b) {
uint32_t r;
asm volatile("cvt.rn.bf16x2.f32 %0, %2, %1;\n"
: "=r"(r) : "f"(a), "f"(b));
return r;
}
namespace {
__device__ __forceinline__ uint64_t f32x2_bits(float2 v) {
uint64_t b; __builtin_memcpy(&b, &v, 8); return b;
}
__device__ __forceinline__ float2 f32x2_make(uint64_t b) {
float2 v; __builtin_memcpy(&v, &b, 8); return v;
}
}
__device__ __forceinline__ float2 fmul2(float2 a, float2 b) {
uint64_t d;
asm volatile("mul.f32x2 %0, %1, %2;\n" : "=l"(d) : "l"(f32x2_bits(a)), "l"(f32x2_bits(b)));
return f32x2_make(d);
}
__device__ __forceinline__ float2 fadd2(float2 a, float2 b) {
uint64_t d;
asm volatile("add.f32x2 %0, %1, %2;\n" : "=l"(d) : "l"(f32x2_bits(a)), "l"(f32x2_bits(b)));
return f32x2_make(d);
}
__device__ __forceinline__ float2 ffma2(float2 a, float2 b, float2 c) {
uint64_t d;
asm volatile("fma.rn.f32x2 %0, %1, %2, %3;\n"
: "=l"(d) : "l"(f32x2_bits(a)), "l"(f32x2_bits(b)), "l"(f32x2_bits(c)));
return f32x2_make(d);
}
__device__ __forceinline__ float2 f32x2_splat(float s) { return make_float2(s, s); }
__device__ __forceinline__ float ex2_approx_f32(float z) {
float d;
asm volatile("ex2.approx.ftz.f32 %0, %1;\n" : "=f"(d) : "f"(z));
return d;
}
__device__ __forceinline__ float2 ex2_emu_f32x2(float x, float y) {
uint32_t ox, oy;
asm volatile(
"{\n\t"
".reg .f32 f1,f2,f3,f4,f5,f6,f7;\n\t"
".reg .b64 l1,l2,l3,l4,l5,l6,l7,l8,l9,l10;\n\t"
".reg .s32 r1,r2,r3,r4,r5,r6,r7,r8;\n\t"
"max.f32 f1, %2, 0fC2FE0000;\n\t"
"max.f32 f2, %3, 0fC2FE0000;\n\t"
"mov.b64 l1, {f1, f2};\n\t"
"mov.f32 f3, 0f4B400000;\n\t"
"mov.b64 l2, {f3, f3};\n\t"
"add.rm.f32x2 l7, l1, l2;\n\t"
"sub.rn.f32x2 l8, l7, l2;\n\t"
"sub.rn.f32x2 l9, l1, l8;\n\t"
"mov.f32 f7, 0f3D9DF09D;\n\t"
"mov.b64 l6, {f7, f7};\n\t"
"mov.f32 f6, 0f3E6906A4;\n\t"
"mov.b64 l5, {f6, f6};\n\t"
"mov.f32 f5, 0f3F31F519;\n\t"
"mov.b64 l4, {f5, f5};\n\t"
"mov.f32 f4, 0f3F800000;\n\t"
"mov.b64 l3, {f4, f4};\n\t"
"fma.rn.f32x2 l10, l9, l6, l5;\n\t"
"fma.rn.f32x2 l10, l10, l9, l4;\n\t"
"fma.rn.f32x2 l10, l10, l9, l3;\n\t"
"mov.b64 {r1, r2}, l7;\n\t"
"mov.b64 {r3, r4}, l10;\n\t"
"shl.b32 r5, r1, 23;\n\t"
"add.s32 r7, r5, r3;\n\t"
"shl.b32 r6, r2, 23;\n\t"
"add.s32 r8, r6, r4;\n\t"
"mov.b32 %0, r7;\n\t"
"mov.b32 %1, r8;\n\t"
"}\n"
: "=r"(ox), "=r"(oy) : "f"(x), "f"(y));
float2 r; __builtin_memcpy(&r.x, &ox, 4); __builtin_memcpy(&r.y, &oy, 4);
return r;
}
__device__ __forceinline__ float rcp_approx_ftz_f32(float x) {
float d;
asm("rcp.approx.ftz.f32 %0, %1;\n" : "=f"(d) : "f"(x));
return d;
}
__device__ __forceinline__ uint32_t make_idesc_table44(
int M, int N,
uint32_t dtype, uint32_t atype, uint32_t btype,
bool transpose_a = false, bool transpose_b = false,
bool negate_a = false, bool negate_b = false) {
uint32_t idesc = 0;
idesc |= (dtype & 0x3) << 4;
idesc |= (atype & 0x7) << 7;
idesc |= (btype & 0x7) << 10;
idesc |= (negate_a ? 1u : 0u) << 13;
idesc |= (negate_b ? 1u : 0u) << 14;
idesc |= (transpose_a ? 1u : 0u) << 15;
idesc |= (transpose_b ? 1u : 0u) << 16;
idesc |= ((static_cast<uint32_t>(N) >> 3) & 0x3F) << 17;
idesc |= ((static_cast<uint32_t>(M) >> 4) & 0x1F) << 24;
return idesc;
}
__device__ __forceinline__ uint32_t make_idesc_bf16_f32(
int M, int N, bool ta = false, bool tb = false) {
return make_idesc_table44(M, N, 1,
1, 1, ta, tb);
}
__device__ __forceinline__ void tcgen05_ld_32x32b_x16(
uint32_t tmem_addr, uint32_t (&r)[16]) {
asm volatile("tcgen05.ld.sync.aligned.32x32b.x16.b32 "
"{%0,%1,%2,%3,%4,%5,%6,%7,%8,%9,"
"%10,%11,%12,%13,%14,%15}, [%16];\n"
: "=r"(r[0]),"=r"(r[1]),"=r"(r[2]),"=r"(r[3]),
"=r"(r[4]),"=r"(r[5]),"=r"(r[6]),"=r"(r[7]),
"=r"(r[8]),"=r"(r[9]),"=r"(r[10]),"=r"(r[11]),
"=r"(r[12]),"=r"(r[13]),"=r"(r[14]),"=r"(r[15])
: "r"(tmem_addr));
}
__device__ __forceinline__ void tcgen05_ld_32x32b_x32(
uint32_t tmem_addr, uint32_t (&r)[32]) {
asm volatile("tcgen05.ld.sync.aligned.32x32b.x32.b32 "
"{%0,%1,%2,%3,%4,%5,%6,%7,%8,%9,"
"%10,%11,%12,%13,%14,%15,%16,%17,%18,%19,"
"%20,%21,%22,%23,%24,%25,%26,%27,%28,%29,"
"%30,%31}, [%32];\n"
: "=r"(r[0]),"=r"(r[1]),"=r"(r[2]),"=r"(r[3]),
"=r"(r[4]),"=r"(r[5]),"=r"(r[6]),"=r"(r[7]),
"=r"(r[8]),"=r"(r[9]),"=r"(r[10]),"=r"(r[11]),
"=r"(r[12]),"=r"(r[13]),"=r"(r[14]),"=r"(r[15]),
"=r"(r[16]),"=r"(r[17]),"=r"(r[18]),"=r"(r[19]),
"=r"(r[20]),"=r"(r[21]),"=r"(r[22]),"=r"(r[23]),
"=r"(r[24]),"=r"(r[25]),"=r"(r[26]),"=r"(r[27]),
"=r"(r[28]),"=r"(r[29]),"=r"(r[30]),"=r"(r[31])
: "r"(tmem_addr));
}
__device__ __forceinline__ void tcgen05_ld_32x32b_x64(
uint32_t tmem_addr, uint32_t (&r)[64]) {
asm volatile("tcgen05.ld.sync.aligned.32x32b.x64.b32 "
"{%0,%1,%2,%3,%4,%5,%6,%7,%8,%9,"
"%10,%11,%12,%13,%14,%15,%16,%17,%18,%19,"
"%20,%21,%22,%23,%24,%25,%26,%27,%28,%29,"
"%30,%31,%32,%33,%34,%35,%36,%37,%38,%39,"
"%40,%41,%42,%43,%44,%45,%46,%47,%48,%49,"
"%50,%51,%52,%53,%54,%55,%56,%57,%58,%59,"
"%60,%61,%62,%63}, [%64];\n"
: "=r"(r[0]),"=r"(r[1]),"=r"(r[2]),"=r"(r[3]),
"=r"(r[4]),"=r"(r[5]),"=r"(r[6]),"=r"(r[7]),
"=r"(r[8]),"=r"(r[9]),"=r"(r[10]),"=r"(r[11]),
"=r"(r[12]),"=r"(r[13]),"=r"(r[14]),"=r"(r[15]),
"=r"(r[16]),"=r"(r[17]),"=r"(r[18]),"=r"(r[19]),
"=r"(r[20]),"=r"(r[21]),"=r"(r[22]),"=r"(r[23]),
"=r"(r[24]),"=r"(r[25]),"=r"(r[26]),"=r"(r[27]),
"=r"(r[28]),"=r"(r[29]),"=r"(r[30]),"=r"(r[31]),
"=r"(r[32]),"=r"(r[33]),"=r"(r[34]),"=r"(r[35]),
"=r"(r[36]),"=r"(r[37]),"=r"(r[38]),"=r"(r[39]),
"=r"(r[40]),"=r"(r[41]),"=r"(r[42]),"=r"(r[43]),
"=r"(r[44]),"=r"(r[45]),"=r"(r[46]),"=r"(r[47]),
"=r"(r[48]),"=r"(r[49]),"=r"(r[50]),"=r"(r[51]),
"=r"(r[52]),"=r"(r[53]),"=r"(r[54]),"=r"(r[55]),
"=r"(r[56]),"=r"(r[57]),"=r"(r[58]),"=r"(r[59]),
"=r"(r[60]),"=r"(r[61]),"=r"(r[62]),"=r"(r[63])
: "r"(tmem_addr));
}
__device__ __forceinline__ void tcgen05_ld_32x32b_x128(
uint32_t tmem_addr, uint32_t (&r)[128]) {
asm volatile("tcgen05.ld.sync.aligned.32x32b.x128.b32 "
"{%0,%1,%2,%3,%4,%5,%6,%7,%8,%9,"
"%10,%11,%12,%13,%14,%15,%16,%17,%18,%19,"
"%20,%21,%22,%23,%24,%25,%26,%27,%28,%29,"
"%30,%31,%32,%33,%34,%35,%36,%37,%38,%39,"
"%40,%41,%42,%43,%44,%45,%46,%47,%48,%49,"
"%50,%51,%52,%53,%54,%55,%56,%57,%58,%59,"
"%60,%61,%62,%63,%64,%65,%66,%67,%68,%69,"
"%70,%71,%72,%73,%74,%75,%76,%77,%78,%79,"
"%80,%81,%82,%83,%84,%85,%86,%87,%88,%89,"
"%90,%91,%92,%93,%94,%95,%96,%97,%98,%99,"
"%100,%101,%102,%103,%104,%105,%106,%107,%108,%109,"
"%110,%111,%112,%113,%114,%115,%116,%117,%118,%119,"
"%120,%121,%122,%123,%124,%125,%126,%127}, [%128];\n"
: "=r"(r[ 0]),"=r"(r[ 1]),"=r"(r[ 2]),"=r"(r[ 3]),
"=r"(r[ 4]),"=r"(r[ 5]),"=r"(r[ 6]),"=r"(r[ 7]),
"=r"(r[ 8]),"=r"(r[ 9]),"=r"(r[ 10]),"=r"(r[ 11]),
"=r"(r[ 12]),"=r"(r[ 13]),"=r"(r[ 14]),"=r"(r[ 15]),
"=r"(r[ 16]),"=r"(r[ 17]),"=r"(r[ 18]),"=r"(r[ 19]),
"=r"(r[ 20]),"=r"(r[ 21]),"=r"(r[ 22]),"=r"(r[ 23]),
"=r"(r[ 24]),"=r"(r[ 25]),"=r"(r[ 26]),"=r"(r[ 27]),
"=r"(r[ 28]),"=r"(r[ 29]),"=r"(r[ 30]),"=r"(r[ 31]),
"=r"(r[ 32]),"=r"(r[ 33]),"=r"(r[ 34]),"=r"(r[ 35]),
"=r"(r[ 36]),"=r"(r[ 37]),"=r"(r[ 38]),"=r"(r[ 39]),
"=r"(r[ 40]),"=r"(r[ 41]),"=r"(r[ 42]),"=r"(r[ 43]),
"=r"(r[ 44]),"=r"(r[ 45]),"=r"(r[ 46]),"=r"(r[ 47]),
"=r"(r[ 48]),"=r"(r[ 49]),"=r"(r[ 50]),"=r"(r[ 51]),
"=r"(r[ 52]),"=r"(r[ 53]),"=r"(r[ 54]),"=r"(r[ 55]),
"=r"(r[ 56]),"=r"(r[ 57]),"=r"(r[ 58]),"=r"(r[ 59]),
"=r"(r[ 60]),"=r"(r[ 61]),"=r"(r[ 62]),"=r"(r[ 63]),
"=r"(r[ 64]),"=r"(r[ 65]),"=r"(r[ 66]),"=r"(r[ 67]),
"=r"(r[ 68]),"=r"(r[ 69]),"=r"(r[ 70]),"=r"(r[ 71]),
"=r"(r[ 72]),"=r"(r[ 73]),"=r"(r[ 74]),"=r"(r[ 75]),
"=r"(r[ 76]),"=r"(r[ 77]),"=r"(r[ 78]),"=r"(r[ 79]),
"=r"(r[ 80]),"=r"(r[ 81]),"=r"(r[ 82]),"=r"(r[ 83]),
"=r"(r[ 84]),"=r"(r[ 85]),"=r"(r[ 86]),"=r"(r[ 87]),
"=r"(r[ 88]),"=r"(r[ 89]),"=r"(r[ 90]),"=r"(r[ 91]),
"=r"(r[ 92]),"=r"(r[ 93]),"=r"(r[ 94]),"=r"(r[ 95]),
"=r"(r[ 96]),"=r"(r[ 97]),"=r"(r[ 98]),"=r"(r[ 99]),
"=r"(r[100]),"=r"(r[101]),"=r"(r[102]),"=r"(r[103]),
"=r"(r[104]),"=r"(r[105]),"=r"(r[106]),"=r"(r[107]),
"=r"(r[108]),"=r"(r[109]),"=r"(r[110]),"=r"(r[111]),
"=r"(r[112]),"=r"(r[113]),"=r"(r[114]),"=r"(r[115]),
"=r"(r[116]),"=r"(r[117]),"=r"(r[118]),"=r"(r[119]),
"=r"(r[120]),"=r"(r[121]),"=r"(r[122]),"=r"(r[123]),
"=r"(r[124]),"=r"(r[125]),"=r"(r[126]),"=r"(r[127])
: "r"(tmem_addr));
}
__device__ __forceinline__
uint32_t smem_ptr_u32(const void* ptr) {
return static_cast<uint32_t>(__cvta_generic_to_shared(ptr));
}
__device__ __forceinline__
void sts_f32(uint32_t smem_addr, float val) {
asm volatile("st.shared.f32 [%0], %1;" :: "r"(smem_addr), "f"(val) : "memory");
}
@@ -28,31 +28,10 @@ void register_rms_norm(pybind11::module_ &);
void register_layer_norm(pybind11::module_ &);
void register_gemm(pybind11::module_ &);
#ifdef TK_COMPILE_BLOCK_SPARSE_VSA_SM100A
extern std::vector<torch::Tensor> block_sparse_sm100a_fwd(
torch::Tensor q, torch::Tensor k, torch::Tensor v, c10::optional<torch::Tensor> v_t,
torch::Tensor q2k_idx, torch::Tensor q2k_num, torch::Tensor variable_block_sizes,
double sm_scale, bool need_lse);
extern std::vector<torch::Tensor> block_sparse_sm100a_blk128_fwd(
torch::Tensor q, torch::Tensor k, torch::Tensor v, c10::optional<torch::Tensor> v_t,
torch::Tensor q2k_idx, torch::Tensor q2k_num, torch::Tensor variable_block_sizes,
double sm_scale, bool need_lse);
#endif
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.doc() = "FastVideo CUDA Kernels";
#ifdef TK_COMPILE_BLOCK_SPARSE_VSA_SM100A
m.def("block_sparse_sm100a_fwd",
torch::wrap_pybind_function(block_sparse_sm100a_fwd),
"VSA block-sparse attention forward, 64-token blocks (Blackwell sm100a)");
m.def("block_sparse_sm100a_blk128_fwd",
torch::wrap_pybind_function(block_sparse_sm100a_blk128_fwd),
"VSA block-sparse attention forward, 128-token blocks (Blackwell sm100a)");
#endif
#ifdef TK_COMPILE_ST_ATTN
m.def("sta_fwd", torch::wrap_pybind_function(sta_forward), "sliding tile attention (Hopper)");
#endif
@@ -1,4 +1,4 @@
"""VSA-128/256 block-sparse attention wrappers.
"""VSA-256 block-sparse attention wrapper.
The default 256-block path is Triton: it expands the logical 256-block map
to the existing 64-block Triton kernel via a dense 4x4 expansion per logical
@@ -7,8 +7,8 @@ edge ("route A"), and requires no optional dependencies.
The FA4 CuTe block-sparse fastpath (intended for Blackwell sm_100+) is
*opt-in* via ``FASTVIDEO_VSA_CUTEDSL=1``. It routes to
:mod:`fastvideo_kernel.block_sparse_attn_cute_fwd`, which natively operates
on 128-token Q/KV blocks (the 256 wrapper expands its logical KV map and
sizes into that physical representation). The CuTe kernel
on 128-token KV blocks (this wrapper expands the logical 256-block map /
sizes into that physical 128-block representation). The CuTe kernel
(``flash_attn.cute`` with block-sparsity) is an optional dependency,
imported lazily only when this fastpath is selected.
@@ -35,7 +35,7 @@ _KV_BLOCK_TRITON = 64 # Existing Triton path uses 64-token KV blocks.
def _resolve_backend() -> str:
"""Pick the backend for the 128/256-block VSA paths.
"""Pick the backend for the 256-block VSA path.
Default is Triton (no optional deps). The FA4 CuTe fastpath is opt-in
via ``FASTVIDEO_VSA_CUTEDSL=1`` and requires the optional FA4 CuTe
@@ -49,26 +49,6 @@ def _resolve_backend() -> str:
return "triton"
def _expand_mask_and_sizes_128_to_64(
logical_mask_128: torch.Tensor,
logical_kv_sizes_128: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Expand a [B, H, Qb128, KVb128] map to 64-token Triton tiles."""
expanded_mask = logical_mask_128.repeat_interleave(2, dim=2).repeat_interleave(2, dim=3)
sizes_i32 = logical_kv_sizes_128.to(torch.int32)
offsets = torch.tensor(
[0, _KV_BLOCK_TRITON],
dtype=torch.int32,
device=sizes_i32.device,
)
expanded_sizes = torch.clamp(
sizes_i32[:, None] - offsets[None, :],
min=0,
max=_KV_BLOCK_TRITON,
).reshape(-1)
return expanded_mask, expanded_sizes
def _expand_mask_and_sizes_256_to_128(
logical_mask_256: torch.Tensor,
logical_kv_sizes_256: torch.Tensor,
@@ -132,63 +112,6 @@ def _triton_via_route_a(
return block_sparse_attn_triton(q, k, v, q2k_idx, q2k_num, sizes_64)
def _triton_via_route_a_128(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
logical_mask_128: torch.Tensor,
logical_kv_sizes_128: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
from .triton_kernels.index import map_to_index as triton_map_to_index
mask_64, sizes_64 = _expand_mask_and_sizes_128_to_64(logical_mask_128, logical_kv_sizes_128)
q2k_idx, q2k_num = triton_map_to_index(mask_64.to(torch.bool))
return block_sparse_attn_triton(q, k, v, q2k_idx, q2k_num, sizes_64)
def block_sparse_attn_128(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
logical_block_map_128: torch.Tensor,
logical_variable_block_sizes_128: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""VSA-128 sparse-branch entrypoint for [B, H, S, D] inputs."""
if logical_block_map_128.dim() == 3:
logical_block_map_128 = logical_block_map_128.unsqueeze(0)
if _resolve_backend() == "triton":
return _triton_via_route_a_128(q, k, v, logical_block_map_128, logical_variable_block_sizes_128)
from .block_sparse_attn_cute_fwd import block_sparse_attn_cute_fwd
return block_sparse_attn_cute_fwd(q, k, v, logical_block_map_128, logical_variable_block_sizes_128)
def block_sparse_attn_128_bshd(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
logical_block_map_128: torch.Tensor,
logical_variable_block_sizes_128: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""VSA-128 sparse-branch entrypoint for [B, S, H, D] inputs."""
if logical_block_map_128.dim() == 3:
logical_block_map_128 = logical_block_map_128.unsqueeze(0)
if _resolve_backend() == "triton":
out_bhsd, aux = _triton_via_route_a_128(
q.transpose(1, 2).contiguous(),
k.transpose(1, 2).contiguous(),
v.transpose(1, 2).contiguous(),
logical_block_map_128,
logical_variable_block_sizes_128,
)
return out_bhsd.transpose(1, 2).contiguous(), aux
from .block_sparse_attn_cute_fwd import block_sparse_attn_cute_fwd_bshd
return block_sparse_attn_cute_fwd_bshd(q, k, v, logical_block_map_128, logical_variable_block_sizes_128)
def block_sparse_attn_256(
q: torch.Tensor,
k: torch.Tensor,
@@ -1,18 +1,18 @@
"""FA4 CuTe-DSL block-sparse attention adapter.
"""CuTe-DSL block-sparse attention forward kernel.
This module adapts VSA's ``(block_map, variable_block_sizes)`` inputs into
FA4's forward and backward ``BlockSparseTensorsTorch`` representations.
FA4's public ``flash_attn_func`` owns the forward/backward autograd bridge.
Thin wrapper around `flash_attn.cute.interface._flash_attn_fwd` that adapts
VSA's `(block_map, variable_block_sizes)` inputs into FA4's
`BlockSparseTensorsTorch` representation and the per-KV-block validity mask.
Both [B, H, S, D] (BHSD) and [B, S, H, D] (BSHD) entrypoints are provided.
The BSHD variant is preferred from VSA-128/256 callers to avoid layout
The BSHD variant is preferred from VSA-256 callers to avoid layout
round-trips on the hot path.
The FA4 CuTe block-sparse kernel (``flash_attn.cute`` with
``block_sparsity``) is an *optional* dependency: it is imported lazily and
only exercised when the VSA-128/256 CuTe fastpath is explicitly selected
(``FASTVIDEO_VSA_CUTEDSL=1``). The default path is Triton and does not require
it. Also needs ``nvidia-cutlass-dsl`` and ``quack-kernels``.
only exercised when the VSA-256 CuTe fastpath is explicitly selected
(``FASTVIDEO_VSA_CUTEDSL=1``). The default VSA-256 path is Triton and does
not require it. Also needs ``nvidia-cutlass-dsl`` and ``quack-kernels``.
"""
from __future__ import annotations
@@ -22,14 +22,13 @@ from typing import Tuple
import torch
_FA4_IMPORT_HINT = ("VSA-128/256 CuTe fastpath requires a FlashAttention-4 CuTe build that "
_FA4_IMPORT_HINT = ("VSA-256 CuTe fastpath requires a FlashAttention-4 CuTe build that "
"provides `flash_attn.cute` with block-sparsity support (plus "
"`nvidia-cutlass-dsl` and `quack-kernels`). This is an optional "
"dependency; the default path is Triton. Install the FA4 CuTe "
"dependency; the default VSA-256 path is Triton. Install the FA4 CuTe "
"build and set FASTVIDEO_VSA_CUTEDSL=1 to enable the CuTe fastpath.")
@functools.lru_cache(maxsize=1)
def _load_fa4_cute():
"""Lazily import the optional FA4 CuTe block-sparse symbols.
@@ -39,39 +38,14 @@ def _load_fa4_cute():
"""
try:
from flash_attn.cute.block_sparsity import BlockSparseTensorsTorch
from flash_attn.cute.interface import (
_flash_attn_bwd,
_flash_attn_fwd,
flash_attn_func,
)
from flash_attn.cute.interface import _flash_attn_fwd
except ImportError as exc: # pragma: no cover - optional dependency
raise ImportError(_FA4_IMPORT_HINT) from exc
return BlockSparseTensorsTorch, flash_attn_func, _flash_attn_fwd, _flash_attn_bwd
return BlockSparseTensorsTorch, _flash_attn_fwd
# FA4's physical Q tile size; KV block size comes from the VSA caller.
_FA4_Q_BLOCK_SIZE = 128
class _SingleQStageLength(int):
"""Keep the real length while selecting FA4's one-stage Q128 path.
On sm_100 FA4 derives ``q_stage`` from ``max_seqlen_q > tile_m``. Its
kernel supports one 128-token Q stage, but the fixed-length public wrapper
does not expose that choice. VSA-128 must select it explicitly; otherwise
adjacent logical Q blocks are merged into a 256-token sparse block.
"""
def __mul__(self, other):
return type(self)(int(self) * int(other))
def __rmul__(self, other):
return type(self)(int(other) * int(self))
def __gt__(self, other):
if int(other) == _FA4_Q_BLOCK_SIZE:
return False
return int(self) > int(other)
# Q-side tile size; kv_block_size comes from the caller's VSA logical KV block.
_M_BLOCK_SIZE_DEFAULT = 128
def _map_to_index(block_map: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
@@ -90,12 +64,12 @@ def _map_to_index(block_map: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
return triton_map_to_index(block_map)
def _choose_q_sparse_block_size(q_len: int, q_tile_size: int = _FA4_Q_BLOCK_SIZE) -> int:
# FA4 supports a doubled Q sparsity granularity on sm_100+ when q_len > q_tile_size.
def _choose_q_sparse_block_size(q_len: int, m_block_size: int = _M_BLOCK_SIZE_DEFAULT) -> int:
# FA4 supports a doubled Q sparsity granularity on sm_100+ when q_len > m_block_size.
major, _ = torch.cuda.get_device_capability()
if major >= 10 and q_len > q_tile_size:
return 2 * q_tile_size
return q_tile_size
if major >= 10 and q_len > m_block_size:
return 2 * m_block_size
return m_block_size
def _aggregate_q_block_map(
@@ -160,35 +134,23 @@ def _build_vbs_mask_mod(kv_block_size: int):
return _vbs_mask_mod
def _build_sparse_tensors(
def _cute_forward(
q_bshd: torch.Tensor,
k_bshd: torch.Tensor,
v_bshd: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
*,
q_len: int,
q_block_size: int,
kv_block_size: int,
need_backward: bool,
force_q_sparse_block_size: int | None = None,
) -> Tuple[object, object | None]:
"""Build the Q-owned forward and KV-owned backward sparse metadata.
``need_backward`` is False on inference-only calls: the backward metadata
is a pair of dense ``[B, H, kv_blocks, q_blocks]`` int32 index tensors that
FA4 keeps alive on its autograd ctx until backward runs, so building it
when nothing requires grad is pure overhead (~80 MiB per call at Wan-14B
720p shape).
"""
BlockSparseTensorsTorch, _, _, _ = _load_fa4_cute()
if force_q_sparse_block_size is None:
q_sparse_candidate = _choose_q_sparse_block_size(q_len)
q_sparse_block_size = max(
q_block_size,
((q_sparse_candidate + q_block_size - 1) // q_block_size) * q_block_size,
)
else:
q_sparse_block_size = force_q_sparse_block_size
if q_sparse_block_size < q_block_size or q_sparse_block_size % q_block_size != 0:
raise ValueError("force_q_sparse_block_size must be a positive multiple of q_block_size")
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Internal: FA4 CuTe BSA fwd with BSHD inputs."""
BlockSparseTensorsTorch, _flash_attn_fwd = _load_fa4_cute()
q_sparse_candidate = _choose_q_sparse_block_size(q_bshd.shape[1])
q_sparse_block_size = max(
q_block_size,
((q_sparse_candidate + q_block_size - 1) // q_block_size) * q_block_size,
)
sparse_map = _aggregate_q_block_map(
block_map,
q_sparse_block_size=q_sparse_block_size,
@@ -196,166 +158,35 @@ def _build_sparse_tensors(
)
kv_full = (variable_block_sizes == kv_block_size).view(1, 1, 1, -1)
kv_partial = ((variable_block_sizes > 0) & (variable_block_sizes < kv_block_size)).view(1, 1, 1, -1)
full_map = sparse_map & kv_full
mask_map = sparse_map & kv_partial
def from_maps(full_map: torch.Tensor, mask_map: torch.Tensor) -> object:
full_block_idx, full_block_cnt = _map_to_index(full_map.contiguous())
mask_block_idx, mask_block_cnt = _map_to_index(mask_map.contiguous())
return BlockSparseTensorsTorch(
full_block_cnt=full_block_cnt.to(torch.int32).contiguous(),
full_block_idx=full_block_idx.to(torch.int32).contiguous(),
mask_block_cnt=mask_block_cnt.to(torch.int32).contiguous(),
mask_block_idx=mask_block_idx.to(torch.int32).contiguous(),
block_size=(q_sparse_block_size, kv_block_size),
)
full_block_idx, full_block_cnt = _map_to_index(full_map)
mask_block_idx, mask_block_cnt = _map_to_index(mask_map)
forward_sparse_tensors = from_maps(
sparse_map & kv_full,
sparse_map & kv_partial,
sparse_tensors = BlockSparseTensorsTorch(
full_block_cnt=full_block_cnt.to(torch.int32).contiguous(),
full_block_idx=full_block_idx.to(torch.int32).contiguous(),
mask_block_cnt=mask_block_cnt.to(torch.int32).contiguous(),
mask_block_idx=mask_block_idx.to(torch.int32).contiguous(),
block_size=(q_sparse_block_size, kv_block_size),
)
if not need_backward:
return forward_sparse_tensors, None
# FA4 backward is KV-owned: for each physical KV tile, list the sparse
# query tiles that selected it. Full and partial KV tiles stay separate
# so the token-level validity mask only runs for padded tiles.
backward_sparse_tensors = from_maps(
(sparse_map & kv_full).transpose(2, 3),
(sparse_map & kv_partial).transpose(2, 3),
)
return forward_sparse_tensors, backward_sparse_tensors
def _cute_attention_q128_forward(
q_bshd: torch.Tensor,
k_bshd: torch.Tensor,
v_bshd: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
*,
need_backward: bool,
) -> Tuple[torch.Tensor, torch.Tensor, object | None]:
"""Run FA4 with one physical Q stage per logical VSA-128 block."""
_, _, flash_attn_fwd, _ = _load_fa4_cute()
forward_sparse_tensors, backward_sparse_tensors = _build_sparse_tensors(
block_map,
variable_block_sizes,
q_len=q_bshd.shape[1],
q_block_size=_FA4_Q_BLOCK_SIZE,
kv_block_size=_FA4_Q_BLOCK_SIZE,
need_backward=need_backward,
force_q_sparse_block_size=_FA4_Q_BLOCK_SIZE,
)
out, lse = flash_attn_fwd(
# _flash_attn_fwd returns (out, lse, p, row_max); keep the first two.
out, lse = _flash_attn_fwd(
q_bshd,
k_bshd,
v_bshd,
tile_mn=(_FA4_Q_BLOCK_SIZE, _FA4_Q_BLOCK_SIZE),
max_seqlen_q=_SingleQStageLength(q_bshd.shape[1]),
mask_mod=_build_vbs_mask_mod(_FA4_Q_BLOCK_SIZE),
block_sparse_tensors=forward_sparse_tensors,
tile_mn=(_M_BLOCK_SIZE_DEFAULT, kv_block_size),
mask_mod=_build_vbs_mask_mod(kv_block_size),
block_sparse_tensors=sparse_tensors,
aux_tensors=[variable_block_sizes],
causal=False,
return_lse=True,
)[:2]
return out, lse, backward_sparse_tensors
class _CuteAttentionQ128(torch.autograd.Function):
@staticmethod
def forward(ctx, q_bshd, k_bshd, v_bshd, block_map, variable_block_sizes):
out, lse, backward_sparse_tensors = _cute_attention_q128_forward(
q_bshd,
k_bshd,
v_bshd,
block_map,
variable_block_sizes,
need_backward=True,
)
ctx.save_for_backward(q_bshd, k_bshd, v_bshd, out, lse, variable_block_sizes)
ctx.backward_sparse_tensors = backward_sparse_tensors
ctx.mark_non_differentiable(lse)
ctx.set_materialize_grads(False)
return out, lse
@staticmethod
def backward(ctx, grad_out, grad_lse):
del grad_lse
q_bshd, k_bshd, v_bshd, out, lse, variable_block_sizes = ctx.saved_tensors
if grad_out is None:
grad_out = torch.zeros_like(out)
_, _, _, flash_attn_bwd = _load_fa4_cute()
dq, dk, dv = flash_attn_bwd(
q_bshd,
k_bshd,
v_bshd,
out,
grad_out.contiguous(),
lse,
softmax_scale=q_bshd.shape[-1]**-0.5,
mask_mod=_build_vbs_mask_mod(_FA4_Q_BLOCK_SIZE),
aux_tensors=[variable_block_sizes],
block_sparse_tensors=ctx.backward_sparse_tensors,
)
return dq, dk, dv, None, None
def _cute_attention_q128(
q_bshd: torch.Tensor,
k_bshd: torch.Tensor,
v_bshd: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
need_backward = torch.is_grad_enabled() and any(t.requires_grad for t in (q_bshd, k_bshd, v_bshd))
if need_backward:
return _CuteAttentionQ128.apply(q_bshd, k_bshd, v_bshd, block_map, variable_block_sizes)
out, lse, _ = _cute_attention_q128_forward(
q_bshd,
k_bshd,
v_bshd,
block_map,
variable_block_sizes,
need_backward=False,
)
return out, lse
def _cute_attention(
q_bshd: torch.Tensor,
k_bshd: torch.Tensor,
v_bshd: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Run FA4's autograd-enabled block-sparse attention with BSHD inputs."""
_, flash_attn_func, _, _ = _load_fa4_cute()
q_block_size = q_bshd.shape[1] // block_map.shape[2]
kv_block_size = k_bshd.shape[1] // block_map.shape[3]
if q_block_size == kv_block_size == _FA4_Q_BLOCK_SIZE:
return _cute_attention_q128(q_bshd, k_bshd, v_bshd, block_map, variable_block_sizes)
need_backward = torch.is_grad_enabled() and any(t.requires_grad for t in (q_bshd, k_bshd, v_bshd))
forward_sparse_tensors, backward_sparse_tensors = _build_sparse_tensors(
block_map,
variable_block_sizes,
q_len=q_bshd.shape[1],
q_block_size=q_block_size,
kv_block_size=kv_block_size,
need_backward=need_backward,
)
return flash_attn_func(
q_bshd,
k_bshd,
v_bshd,
mask_mod=_build_vbs_mask_mod(kv_block_size),
aux_tensors=[variable_block_sizes],
block_sparse_tensors=forward_sparse_tensors,
block_sparse_tensors_bwd=backward_sparse_tensors,
return_lse=True,
)
def block_sparse_attn_cute_fwd(
q: torch.Tensor,
k: torch.Tensor,
@@ -363,25 +194,34 @@ def block_sparse_attn_cute_fwd(
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Autograd-enabled CuTe block-sparse attention for [B, H, S, D]."""
"""CuTe forward-only block-sparse attention with [B, H, S, D] inputs."""
if block_map.dim() == 3:
block_map = block_map.unsqueeze(0)
q_block_size = q.shape[2] // block_map.shape[2]
kv_block_size = k.shape[2] // block_map.shape[3]
q_bshd = q.transpose(1, 2).contiguous()
k_bshd = k.transpose(1, 2).contiguous()
v_bshd = v.transpose(1, 2).contiguous()
out_bshd, lse = _cute_attention(
out_bshd, lse_bshd = _cute_forward(
q_bshd,
k_bshd,
v_bshd,
block_map,
variable_block_sizes,
q_block_size=q_block_size,
kv_block_size=kv_block_size,
)
out = out_bshd.transpose(1, 2).contiguous()
# FA4 already returns lse as [B, H, S], matching the Triton path's aux
# contract, so it needs no transpose. Detach before any further op: the
# value is informational and callers never backprop through it.
return out, lse.detach()
if lse_bshd is None:
lse = torch.empty(
(q.shape[0], q.shape[1], q.shape[2]),
dtype=torch.float32,
device=q.device,
)
else:
lse = lse_bshd.transpose(1, 2).contiguous()
return out, lse
def block_sparse_attn_cute_fwd_bshd(
@@ -391,16 +231,27 @@ def block_sparse_attn_cute_fwd_bshd(
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Autograd-enabled CuTe block-sparse attention for [B, S, H, D]."""
"""CuTe forward-only block-sparse attention with [B, S, H, D] inputs."""
if block_map.dim() == 3:
block_map = block_map.unsqueeze(0)
q_block_size = q.shape[1] // block_map.shape[2]
kv_block_size = k.shape[1] // block_map.shape[3]
out, lse = _cute_attention(
out, lse_bshd = _cute_forward(
q,
k,
v,
block_map,
variable_block_sizes,
q_block_size=q_block_size,
kv_block_size=kv_block_size,
)
# lse is [B, H, S] regardless of the q/k/v layout; see above.
return out, lse.detach()
if lse_bshd is None:
lse = torch.empty(
(q.shape[0], q.shape[2], q.shape[1]),
dtype=torch.float32,
device=q.device,
)
else:
lse = lse_bshd.transpose(1, 2).contiguous()
return out, lse
@@ -1,104 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""sm_100a (Blackwell) CUDA block-sparse VSA forward.
A third backend behind the same VSA op as the Triton and CuTe-DSL paths. Forward only: it
returns ``(out, lse)`` with ``lse`` in exactly the form ``triton_block_sparse_attn_forward``
writes -- ``max(qk * qk_scale) + log2(l)``, ``[B, H, S]`` fp32 -- so
``block_sparse_attn_backward_triton`` runs against it unchanged.
The extension carries TWO instantiations of the kernel, for 64- and 128-token sparse blocks
(tile volumes 64 and 128 in ``build_vsa_metadata``); the block size is inferred from the
tensors and picks the op. Anything else falls back to Triton via ``is_supported``.
"""
from typing import Tuple
import torch
try:
# The pybind symbols live on fastvideo_kernel_ops, NOT on the _C package that contains it.
# `import fastvideo_kernel._C as _C` resolves to the namespace package, whose __init__ is
# empty, so hasattr() fails on a wheel install and the caller silently falls back with the
# kernel built and present.
from fastvideo_kernel._C import fastvideo_kernel_ops as _C
_FWD_BY_BLOCK = {
64: getattr(_C, "block_sparse_sm100a_fwd", None),
128: getattr(_C, "block_sparse_sm100a_blk128_fwd", None),
}
_HAS_VSA_SM100A = any(_FWD_BY_BLOCK.values())
except ImportError: # pragma: no cover - extension not built
_C = None
_FWD_BY_BLOCK = {}
_HAS_VSA_SM100A = False
_SM100 = (10, 0)
HEAD_DIM = 128
# Must match the -DVSA_BHSD the extension was compiled with (see CMakeLists).
BHSD = True
def _block_size(q: torch.Tensor, variable_block_sizes: torch.Tensor) -> int:
num_blocks = variable_block_sizes.numel()
seqlen = q.shape[2] if BHSD else q.shape[1]
return 0 if num_blocks == 0 or seqlen % num_blocks else seqlen // num_blocks
def is_supported(q: torch.Tensor, variable_block_sizes: torch.Tensor) -> bool:
"""True iff this build can run these tensors; otherwise the caller uses Triton.
Static facts only -- shapes, dtypes, arch, layout. Deliberately NO reads of tensor
contents: the previous ``int(variable_block_sizes.min())`` was a GPU->CPU sync on every
call, and the kernel no longer needs it (see below). This predicate must stay cheap
enough to sit on a per-layer dispatch path.
What the kernel accepts (and is tested to handle):
* q/k/v: contiguous 4-D bf16 CUDA tensors on an sm_100 device, head_dim 128, laid out
as compiled (BHSD here); seqlen == num_blocks * block with an EVEN num_blocks (a CTA
owns an adjacent pair of query blocks) and a 64- or 128-token build present.
* q2k_num: any per-row counts in [0, max_kv], NON-uniform across rows included. Rows
with count 0 produce exactly-zero output rows (and a finite LSE sentinel) rather
than attending anywhere -- so no ``.min()`` floor is required of the caller.
* q2k_idx: rows only need valid entries (in [0, num_blocks)) BELOW that row's count;
padding past the count (e.g. map_to_index's -1 fill) is never dereferenced. max_kv
(= q2k_idx.shape[-1]) must be >= 1, which the host launcher re-checks.
* variable_block_sizes: per-KV-block valid-token counts in [0, block]; keys at or past
a block's count are masked. Integer metadata is converted to int32/contiguous by
``block_sparse_attn_sm100a`` itself, so int64 inputs merely cost a cast.
"""
if not _HAS_VSA_SM100A or not q.is_cuda:
return False
if torch.cuda.get_device_capability(q.device) != _SM100:
return False
if q.dtype != torch.bfloat16 or q.dim() != 4 or q.shape[-1] != HEAD_DIM:
return False
if not q.is_contiguous():
return False
if _FWD_BY_BLOCK.get(_block_size(q, variable_block_sizes)) is None:
return False
# A CTA owns an adjacent pair of query blocks.
if variable_block_sizes.numel() % 2 != 0:
return False
# Metadata must be integer-typed so the wrapper's int32 conversion is value-preserving.
if not variable_block_sizes.is_cuda or variable_block_sizes.dtype not in (torch.int32, torch.int64):
return False
return True
def block_sparse_attn_sm100a(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
variable_block_sizes: torch.Tensor,
need_lse: bool = True,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Forward pass. Returns ``(out, lse)``; ``out`` has q's layout."""
fwd = _FWD_BY_BLOCK[_block_size(q, variable_block_sizes)]
idx = q2k_idx.to(torch.int32).contiguous()
num = q2k_num.to(torch.int32).contiguous()
vbs = variable_block_sizes.to(torch.int32).contiguous()
sm_scale = 1.0 / (q.shape[-1]**0.5)
res = fwd(q.contiguous(), k.contiguous(), v.contiguous(), None,
idx, num, vbs, sm_scale, need_lse)
return (res[0], res[1]) if need_lse else (res[0], None)
+22 -24
View File
@@ -2,8 +2,6 @@ import math
import torch
from .block_sparse_attn import block_sparse_attn
from .block_sparse_attn_256 import (
block_sparse_attn_128,
block_sparse_attn_128_bshd,
block_sparse_attn_256,
block_sparse_attn_256_bshd,
)
@@ -76,13 +74,12 @@ def video_sparse_attn(
Dispatches the sparse branch by ``block_elements = prod(block_size)``:
- 64 -> existing TK/Triton path (see ``block_sparse_attn_from_indices``).
- 128 -> Triton fallback or CuTe FA4 block-sparse attention.
- 256 -> CuTe FA4 block-sparse attention (see ``block_sparse_attn_256``).
Backend overrides:
- ``FASTVIDEO_VSA_TRITON=1`` forces Triton in either path.
- ``FASTVIDEO_VSA_TK=1`` prefers the sm_90 TK kernel in the 64-block path.
- ``FASTVIDEO_VSA_CUTEDSL=1`` prefers CuTe in the 128/256-block paths.
- ``FASTVIDEO_VSA_CUTEDSL=1`` prefers CuTe in the 256-block path.
"""
if isinstance(block_size, int):
block_size = (block_size, block_size, block_size)
@@ -122,9 +119,8 @@ def video_sparse_attn(
# Sparse branch (fused Triton topk mask)
mask = fused_topk_mask(scores, topk)
if block_elements in (128, 256):
attention = block_sparse_attn_128 if block_elements == 128 else block_sparse_attn_256
out_s = attention(q, k, v, mask, variable_block_sizes)[0]
if block_elements == 256:
out_s = block_sparse_attn_256(q, k, v, mask, variable_block_sizes)[0]
else:
out_s = block_sparse_attn(q, k, v, mask, variable_block_sizes)[0]
@@ -146,14 +142,14 @@ def video_sparse_attn_bshd(
"""VSA entrypoint for [B, S, H, D] tensors.
Avoids the BHSD<->BSHD round-trip that ``video_sparse_attn`` performs on
the CuTe 128/256-block paths; the 64-block path still expects BHSD and is not
the CuTe 256-block path; the 64-block path still expects BHSD and is not
supported here.
"""
if isinstance(block_size, int):
block_size = (block_size, block_size, block_size)
block_elements = block_size[0] * block_size[1] * block_size[2]
if block_elements not in (128, 256):
raise ValueError("video_sparse_attn_bshd is only defined for block_elements=128 or 256 "
if block_elements != 256:
raise ValueError("video_sparse_attn_bshd is only defined for block_elements=256 "
f"(got {block_elements}); use video_sparse_attn for the 64-block path.")
batch, q_seq_len, heads, dim = q.shape
@@ -175,15 +171,19 @@ def video_sparse_attn_bshd(
raise ValueError(f"q_variable_block_sizes must have length q_num_blocks={q_num_blocks}, "
f"got {q_variable_block_sizes.numel()}")
# Compression branch (BSHD-native: match fused_block_mean's semantics).
# Padding values are expected to be zero; gradients are broadcast across
# the full padded block, just like the BHSD fused common path.
# Compression branch (BSHD-native: mean over the 256-token axis).
token_idx = torch.arange(block_elements, device=q.device, dtype=torch.int32)
q_token_valid = (token_idx.view(1, -1) < q_variable_block_sizes.view(-1,
1)).view(1, q_num_blocks, block_elements, 1, 1)
kv_token_valid = (token_idx.view(1, -1) < variable_block_sizes.view(-1,
1)).view(1, kv_num_blocks, block_elements, 1, 1)
q_c = q.view(batch, q_num_blocks, block_elements, heads, dim)
k_c = k.view(batch, kv_num_blocks, block_elements, heads, dim)
v_c = v.view(batch, kv_num_blocks, block_elements, heads, dim)
q_c = (q_c.float().sum(dim=2) / q_variable_block_sizes.view(1, -1, 1, 1)).to(q.dtype)
k_c = (k_c.float().sum(dim=2) / variable_block_sizes.view(1, -1, 1, 1)).to(k.dtype)
v_c = (v_c.float().sum(dim=2) / variable_block_sizes.view(1, -1, 1, 1)).to(v.dtype)
q_c = ((q_c.float() * q_token_valid).sum(dim=2) / q_variable_block_sizes.view(1, -1, 1, 1)).to(q.dtype)
k_c = ((k_c.float() * kv_token_valid).sum(dim=2) / variable_block_sizes.view(1, -1, 1, 1)).to(k.dtype)
v_c = ((v_c.float() * kv_token_valid).sum(dim=2) / variable_block_sizes.view(1, -1, 1, 1)).to(v.dtype)
q_ch = q_c.permute(0, 2, 1, 3).contiguous()
k_ch = k_c.permute(0, 2, 1, 3).contiguous()
v_ch = v_c.permute(0, 2, 1, 3).contiguous()
@@ -195,15 +195,13 @@ def video_sparse_attn_bshd(
# Sparse branch (fused Triton topk mask + CuTe BSHD).
mask = fused_topk_mask(scores, topk)
attention = block_sparse_attn_128_bshd if block_elements == 128 else block_sparse_attn_256_bshd
out_s, _ = attention(q, k, v, mask, variable_block_sizes)
out_s, _ = block_sparse_attn_256_bshd(q, k, v, mask, variable_block_sizes)
# Out-of-place: ``out_s`` is the tensor FA4's autograd node saved for its
# backward, so mutating it in place invalidates the graph.
out_view = out_s.view(batch, q_num_blocks, block_elements, heads, dim)
out = out_s
out_view = out.view(batch, q_num_blocks, block_elements, heads, dim)
if compress_attn_weight is not None:
gate_view = compress_attn_weight.view(batch, q_num_blocks, block_elements, heads, dim)
out = out_view + out_c_blk.unsqueeze(2) * gate_view
out_view.add_(out_c_blk.unsqueeze(2) * gate_view)
else:
out = out_view + out_c_blk.unsqueeze(2)
return out.view(batch, q_seq_len, heads, dim)
out_view.add_(out_c_blk.unsqueeze(2))
return out
@@ -237,12 +237,7 @@ def _attn_bwd_dkdv(
# Load m before computing qk to reduce pipeline stall.
offs_m = start_m + block_sparse_offset + tl.arange(0, BLOCK_M1)
m = tl.load(M + offs_m)
# Recompute logits exactly as the forward does: raw bf16 operands into
# the dot, fp32 scale after accumulation. A bf16 pre-scaled K perturbs
# the recomputed logits relative to the saved M by an error
# proportional to |logit|, which exp2 amplifies into arbitrarily wrong
# probabilities at large activations.
qkT = tl.dot(k, qT) * (sm_scale * 1.4426950408889634)
qkT = tl.dot(k, qT)
pT = tl.math.exp2(qkT - m[None, :])
mask = tl.arange(0, BLOCK_N1) < block_size
pT = tl.where(mask[:, None], pT, 0.0)
@@ -273,7 +268,6 @@ def _attn_bwd_dq(
do,
m,
D,
sm_scale,
# shared by Q/K/V/DO.
q2k_index,
q2k_num,
@@ -321,7 +315,7 @@ def _attn_bwd_dq(
block_sparse_offset = (kv_idx * 2 + half) * step_n * stride_tok
kT = tl.load(kT_ptrs + block_sparse_offset)
vT = tl.load(vT_ptrs + block_sparse_offset)
qk = tl.dot(q, kT) * (sm_scale * 1.4426950408889634)
qk = tl.dot(q, kT)
p = tl.math.exp2(qk - m)
offs_in_block = half * step_n + tl.arange(0, BLOCK_N2)
mask = offs_in_block < block_size
@@ -330,7 +324,8 @@ def _attn_bwd_dq(
dp = tl.dot(do, vT).to(tl.float32)
ds = p * (dp - Di[:, None])
ds = ds.to(tl.bfloat16)
# Compute dQ (kT is raw; the caller applies sm_scale once at the end).
# Compute dQ.
# NOTE: We need to de-scale dq in the end, because kT was pre-scaled.
dq += tl.dot(ds, tl.trans(kT))
# Increment pointers.
return dq
@@ -458,7 +453,6 @@ def _attn_bwd(
do,
m,
D, #
sm_scale,
q2k_index,
q2k_num,
max_kv_blks,
@@ -476,7 +470,7 @@ def _attn_bwd(
)
# Write back dQ.
dq_ptrs = DQ + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d
dq *= sm_scale
dq *= LN2
tl.store(dq_ptrs, dq)
@@ -597,7 +591,6 @@ def _attn_bwd_dq_kernel(
Q,
K,
V,
sm_scale,
DO, #
DQ,
M,
@@ -670,7 +663,6 @@ def _attn_bwd_dq_kernel(
do,
m,
D,
sm_scale,
q2k_index,
q2k_num,
max_kv_blks,
@@ -688,7 +680,7 @@ def _attn_bwd_dq_kernel(
)
dq_ptrs = DQ + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d
dq_acc *= sm_scale
dq_acc *= LN2
tl.store(dq_ptrs, dq_acc)
@@ -756,11 +748,9 @@ def triton_block_sparse_attn_backward(do, q, k, v, o, M, q2k_index, q2k_num, k2q
dv = torch.empty_like(v)
BATCH, N_HEAD = q.shape[:2]
BLOCK_M1, BLOCK_N1, BLOCK_M2, BLOCK_N2 = 32, 64, 64, 32
# K stays raw: the backward kernels apply sm_scale in fp32 after the dot,
# matching the forward's rounding exactly. (A bf16 pre-scaled K perturbs
# the recomputed logits vs the saved M; exp2 turns that into unboundedly
# wrong probabilities at large activations.)
RCP_LN2 = 1.4426950408889634 # = 1.0 / ln(2)
arg_k = k
arg_k = arg_k * (sm_scale * RCP_LN2)
PRE_BLOCK = 64
assert Tq % PRE_BLOCK == 0
pre_grid = (Tq // PRE_BLOCK, BATCH * N_HEAD)
@@ -823,7 +813,6 @@ def triton_block_sparse_attn_backward(do, q, k, v, o, M, q2k_index, q2k_num, k2q
q,
arg_k,
v,
sm_scale,
do,
dq,
M,
@@ -14,9 +14,7 @@ import math
import torch
VSA_TILE_SIZE = (4, 4, 4)
# 128 is served by the sm_100a CUDA backend (block_sparse_attn_sm100a); 64 and 256 by
# Triton and the CuTe-DSL path. A volume here only needs a backend that accepts it.
_SUPPORTED_VSA_BLOCK_VOLUMES = (64, 128, 256)
_SUPPORTED_VSA_BLOCK_VOLUMES = (64, 256)
def _canonicalize_device(device: torch.device | str) -> torch.device:
@@ -1,146 +0,0 @@
"""VSA-128 CuTe/Triton forward and backward parity on Blackwell."""
from __future__ import annotations
import math
import pytest
import torch
from fastvideo_kernel import video_sparse_attn, video_sparse_attn_bshd
from fastvideo_kernel.block_sparse_attn_256 import block_sparse_attn_128
from .test_vsa256_triton import _metrics, _torch_vsa256_reference
_BLOCK = 128
_BLOCK_SIZE_3D = (2, 8, 8)
def _select_backend(monkeypatch, backend: str) -> None:
if not torch.cuda.is_available():
pytest.skip("CUDA is required")
if backend == "cute":
pytest.importorskip(
"flash_attn.cute.block_sparsity",
reason="optional FA4 CuTe build (flash_attn.cute) not installed",
)
monkeypatch.setenv("FASTVIDEO_VSA_CUTEDSL", "1")
monkeypatch.delenv("FASTVIDEO_VSA_TRITON", raising=False)
monkeypatch.delenv("FASTVIDEO_KERNEL_VSA_FORCE_TRITON", raising=False)
else:
monkeypatch.setenv("FASTVIDEO_VSA_TRITON", "1")
monkeypatch.delenv("FASTVIDEO_VSA_CUTEDSL", raising=False)
def _dense_sparse_reference(q, k, v, block_map, variable_block_sizes):
token_mask = block_map.repeat_interleave(_BLOCK, dim=2).repeat_interleave(_BLOCK, dim=3)
kv_valid = torch.arange(_BLOCK, device=k.device) < variable_block_sizes[:, None]
token_mask = token_mask & kv_valid.reshape(1, 1, 1, -1)
logits = torch.matmul(q.float(), k.float().transpose(-2, -1)) / math.sqrt(q.shape[-1])
probabilities = torch.softmax(logits.masked_fill(~token_mask, float("-inf")), dim=-1)
return torch.matmul(probabilities, v.float()).to(q.dtype)
def _check(tag: str, expected: torch.Tensor, actual: torch.Tensor, avg_tol: float, rel_tol: float) -> None:
assert torch.isfinite(actual).all().item(), f"{tag}: non-finite values"
avg_abs, max_rel = _metrics(expected, actual)
print(f" {tag}: avg_abs={avg_abs:.6e}, max_rel={max_rel:.6e}")
assert avg_abs < avg_tol
assert max_rel < rel_tol
@pytest.mark.cuda
@pytest.mark.parametrize("backend", ["cute", "triton"])
def test_vsa128_explicit_routes_forward_backward(backend: str, monkeypatch) -> None:
"""Adjacent Q128 blocks must keep independent routes instead of merging."""
_select_backend(monkeypatch, backend)
torch.manual_seed(53)
shape = (1, 1, 3 * _BLOCK, 128)
base = [torch.randn(shape, device="cuda", dtype=torch.bfloat16) for _ in range(3)]
grad_output = torch.randn(shape, device="cuda", dtype=torch.bfloat16)
variable_block_sizes = torch.tensor([128, 91, 37], device="cuda", dtype=torch.int32)
block_map = torch.eye(3, device="cuda", dtype=torch.bool).view(1, 1, 3, 3)
actual_inputs = [tensor.detach().clone().requires_grad_() for tensor in base]
actual, _ = block_sparse_attn_128(*actual_inputs, block_map, variable_block_sizes)
(actual * grad_output).sum().backward()
reference_inputs = [tensor.detach().clone().requires_grad_() for tensor in base]
expected = _dense_sparse_reference(*reference_inputs, block_map, variable_block_sizes)
(expected * grad_output).sum().backward()
print(f"[vsa128-explicit-{backend}]")
_check("out", expected, actual, 1e-3, 0.2)
for name, reference, candidate in zip(("dq", "dk", "dv"), reference_inputs, actual_inputs, strict=True):
_check(name, reference.grad, candidate.grad, 2e-2, 0.5)
def _zero_kv_tail(x: torch.Tensor, variable_block_sizes: torch.Tensor) -> torch.Tensor:
valid = torch.arange(_BLOCK, device=x.device) < variable_block_sizes[:, None]
valid = valid.view(1, 1, -1, _BLOCK, 1).expand_as(x.view(1, x.shape[1], -1, _BLOCK, x.shape[-1]))
return x * valid.reshape_as(x).to(x.dtype)
@pytest.mark.cuda
@pytest.mark.parametrize("backend", ["cute", "triton"])
@pytest.mark.parametrize("layout", ["bhsd", "bshd"])
def test_vsa128_wrapper_forward_backward(backend: str, layout: str, monkeypatch) -> None:
_select_backend(monkeypatch, backend)
torch.manual_seed(59)
batch, heads, dim = 1, 2, 128
q_blocks, kv_blocks, topk = 3, 4, 2
q_shape = (batch, heads, q_blocks * _BLOCK, dim)
kv_shape = (batch, heads, kv_blocks * _BLOCK, dim)
q_base = torch.randn(q_shape, device="cuda", dtype=torch.bfloat16)
kv_sizes = torch.tensor([128, 91, 37, 128], device="cuda", dtype=torch.int32)
q_sizes = torch.full((q_blocks, ), _BLOCK, device="cuda", dtype=torch.int32)
k_base = _zero_kv_tail(torch.randn(kv_shape, device="cuda", dtype=torch.bfloat16), kv_sizes)
v_base = _zero_kv_tail(torch.randn(kv_shape, device="cuda", dtype=torch.bfloat16), kv_sizes)
gate_base = torch.randn(q_shape, device="cuda", dtype=torch.bfloat16) * 0.1
grad_output = torch.randn(q_shape, device="cuda", dtype=torch.bfloat16)
if layout == "bhsd":
actual_inputs = [tensor.detach().clone().requires_grad_() for tensor in (q_base, k_base, v_base)]
actual_gate = gate_base.detach().clone().requires_grad_()
actual = video_sparse_attn(
*actual_inputs,
kv_sizes,
q_sizes,
topk,
block_size=_BLOCK_SIZE_3D,
compress_attn_weight=actual_gate,
)
actual_grads = actual_inputs
else:
bshd_inputs = [tensor.transpose(1, 2).contiguous().detach().requires_grad_()
for tensor in (q_base, k_base, v_base)]
bshd_gate = gate_base.transpose(1, 2).contiguous().detach().requires_grad_()
actual = video_sparse_attn_bshd(
*bshd_inputs,
kv_sizes,
q_sizes,
topk,
block_size=_BLOCK_SIZE_3D,
compress_attn_weight=bshd_gate,
).transpose(1, 2)
actual_grads = bshd_inputs
(actual * grad_output).sum().backward()
reference_inputs = [tensor.detach().clone().requires_grad_() for tensor in (q_base, k_base, v_base)]
reference_gate = gate_base.detach().clone().requires_grad_()
expected = _torch_vsa256_reference(
*reference_inputs,
q_sizes,
kv_sizes,
topk,
compress_attn_weight=reference_gate,
)
(expected * grad_output).sum().backward()
print(f"[vsa128-wrapper-{backend}-{layout}]")
_check("out", expected, actual, 1e-3, 0.2)
for name, reference, candidate in zip(("dq", "dk", "dv"), reference_inputs, actual_grads, strict=True):
candidate_grad = candidate.grad if layout == "bhsd" else candidate.grad.transpose(1, 2)
_check(name, reference.grad, candidate_grad, 2e-2, 0.5)
actual_gate_grad = actual_gate.grad if layout == "bhsd" else bshd_gate.grad.transpose(1, 2)
_check("dgate", reference_gate.grad, actual_gate_grad, 1e-3, 0.2)
@@ -1,224 +0,0 @@
"""VSA-256 FA4 CuTe forward/backward parity for BHSD and BSHD APIs.
Covers the shapes the CuTe backward actually sees in production: the gated
compression branch (`compress_attn_weight`), partially filled Q tiles,
and q_len != kv_len. Also pins the inference fast path, which must skip the
KV-owned backward metadata without changing the forward result.
"""
from __future__ import annotations
from typing import Tuple
import pytest
import torch
from fastvideo_kernel import video_sparse_attn, video_sparse_attn_bshd
from .test_vsa256_triton import _metrics, _torch_vsa256_reference
_BLOCK = 256
_BLOCK_SIZE_3D = (4, 8, 8) # prod == 256
# Measured on GB200 (sm_100) with bf16 inputs: grads land around 1e-4 avg_abs
# and <=0.11 max_rel across every case below, so these leave ~10x headroom
# without being loose enough to hide a real regression.
_OUT_TOL = (1e-3, 0.2)
_GRAD_TOL = (1e-3, 0.25)
@pytest.fixture(autouse=True)
def _require_cute_backend(monkeypatch):
pytest.importorskip(
"flash_attn.cute.block_sparsity",
reason="optional FA4 CuTe build (flash_attn.cute) not installed",
)
if not torch.cuda.is_available():
pytest.skip("CUDA is required")
monkeypatch.setenv("FASTVIDEO_VSA_CUTEDSL", "1")
monkeypatch.delenv("FASTVIDEO_VSA_TRITON", raising=False)
monkeypatch.delenv("FASTVIDEO_KERNEL_VSA_FORCE_TRITON", raising=False)
def _zero_pad_tail(x: torch.Tensor, var: torch.Tensor) -> torch.Tensor:
"""Zero the padded tail of every 256-token tile of a [B, H, S, D] tensor.
VSA callers scatter into a zeroed tile buffer, so padded slots are zero;
both the kernel and the reference rely on that.
"""
bsz, heads, _, dim = x.shape
blocks = var.numel()
token_idx = torch.arange(_BLOCK, device=x.device, dtype=torch.int32)
valid = (token_idx.view(1, -1) < var.view(-1, 1)).view(1, 1, blocks, _BLOCK, 1)
valid = valid.expand(bsz, heads, blocks, _BLOCK, dim).reshape_as(x)
return x * valid.to(x.dtype)
def _make_inputs(
q_blocks: int,
kv_blocks: int,
kv_var: torch.Tensor,
q_var: torch.Tensor,
heads: int = 2,
dim: int = 128,
seed: int = 42,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
torch.manual_seed(seed)
device = torch.device("cuda")
dtype = torch.bfloat16
sq, skv = q_blocks * _BLOCK, kv_blocks * _BLOCK
q = torch.randn(1, heads, sq, dim, device=device, dtype=dtype)
k = torch.randn(1, heads, skv, dim, device=device, dtype=dtype)
v = torch.randn(1, heads, skv, dim, device=device, dtype=dtype)
grad_out = torch.randn_like(q)
return _zero_pad_tail(q, q_var), _zero_pad_tail(k, kv_var), _zero_pad_tail(v, kv_var), grad_out
def _check(tag: str, ref: torch.Tensor, got: torch.Tensor, tol: Tuple[float, float]) -> None:
assert torch.isfinite(got).all().item(), f"{tag}: non-finite values"
avg_abs, max_rel = _metrics(ref, got)
print(f" {tag}: avg_abs={avg_abs:.6e}, max_rel={max_rel:.6e}")
assert avg_abs < tol[0], f"{tag}: avg_abs {avg_abs:.3e} >= {tol[0]:.3e}"
assert max_rel < tol[1], f"{tag}: max_rel {max_rel:.3e} >= {tol[1]:.3e}"
def _run_bhsd(q, k, v, kv_var, q_var, topk, gate=None):
qg, kg, vg = (t.detach().clone().requires_grad_(True) for t in (q, k, v))
out = video_sparse_attn(qg, kg, vg, kv_var, q_var, topk, block_size=_BLOCK_SIZE_3D, compress_attn_weight=gate)
return out, (qg, kg, vg)
def _run_bshd(q, k, v, kv_var, q_var, topk, gate=None):
qg, kg, vg = (t.transpose(1, 2).contiguous().requires_grad_(True) for t in (q, k, v))
gate_bshd = None if gate is None else gate.transpose(1, 2).contiguous()
out = video_sparse_attn_bshd(qg,
kg,
vg,
kv_var,
q_var,
topk,
block_size=_BLOCK_SIZE_3D,
compress_attn_weight=gate_bshd)
return out.transpose(1, 2), (qg, kg, vg)
def _reference(q, k, v, q_var, kv_var, topk, gate=None):
qr, kr, vr = (t.detach().clone().requires_grad_(True) for t in (q, k, v))
out = _torch_vsa256_reference(qr, kr, vr, q_var, kv_var, topk, compress_attn_weight=gate)
return out, (qr, kr, vr)
def _compare(tag, layout, q, k, v, kv_var, q_var, topk, grad_out, gate=None):
runner = _run_bhsd if layout == "bhsd" else _run_bshd
out, (qg, kg, vg) = runner(q, k, v, kv_var, q_var, topk, gate=gate)
(out * grad_out).sum().backward()
grads = [g.grad if g.grad.dim() == 4 and layout == "bhsd" else g.grad for g in (qg, kg, vg)]
if layout == "bshd":
grads = [g.transpose(1, 2) for g in grads]
out_ref, refs = _reference(q, k, v, q_var, kv_var, topk, gate=gate)
(out_ref * grad_out).sum().backward()
print(f"[{tag}-{layout}]")
_check("out", out_ref, out, _OUT_TOL)
for name, ref, got in zip(("dq", "dk", "dv"), refs, grads):
_check(name, ref.grad, got, _GRAD_TOL)
@pytest.mark.cuda
@pytest.mark.parametrize("layout", ["bhsd", "bshd"])
def test_vsa256_cute_forward_backward_vs_torch_ref(layout: str) -> None:
kv_var = torch.tensor([256, 173, 79, 256], dtype=torch.int32, device="cuda")
q_var = torch.full((3, ), _BLOCK, dtype=torch.int32, device="cuda")
q, k, v, grad_out = _make_inputs(3, 4, kv_var, q_var)
_compare("vsa256-cute", layout, q, k, v, kv_var, q_var, 2, grad_out)
@pytest.mark.cuda
@pytest.mark.parametrize("layout", ["bhsd", "bshd"])
def test_vsa256_cute_backward_with_compress_gate(layout: str) -> None:
"""The gated compression branch is what Wan and MiniMax-H3 actually run.
It is also the branch that composes the sparse output with the compression
output, so it is the one that breaks if that composition mutates FA4's
saved output in place.
"""
kv_var = torch.tensor([256, 200, 256, 91], dtype=torch.int32, device="cuda")
q_var = torch.full((3, ), _BLOCK, dtype=torch.int32, device="cuda")
q, k, v, grad_out = _make_inputs(3, 4, kv_var, q_var, seed=7)
gate = torch.randn_like(q) * 0.1
_compare("vsa256-cute-gated", layout, q, k, v, kv_var, q_var, 2, grad_out, gate=gate)
@pytest.mark.cuda
@pytest.mark.parametrize("layout", ["bhsd", "bshd"])
def test_vsa256_cute_backward_partial_q_blocks(layout: str) -> None:
"""Q tiles that are not full: only the compression divisor depends on it,
but it is the one axis the existing coverage held constant."""
kv_var = torch.tensor([256, 128, 256], dtype=torch.int32, device="cuda")
q_var = torch.tensor([256, 61, 199], dtype=torch.int32, device="cuda")
q, k, v, grad_out = _make_inputs(3, 3, kv_var, q_var, seed=11)
_compare("vsa256-cute-partial-q", layout, q, k, v, kv_var, q_var, 2, grad_out)
@pytest.mark.cuda
@pytest.mark.parametrize("layout", ["bhsd", "bshd"])
def test_vsa256_cute_backward_cross_q_kv(layout: str) -> None:
"""q_len != kv_len: forward has coverage, backward did not."""
kv_var = torch.tensor([256, 143, 256, 256, 88], dtype=torch.int32, device="cuda")
q_var = torch.full((2, ), _BLOCK, dtype=torch.int32, device="cuda")
q, k, v, grad_out = _make_inputs(2, 5, kv_var, q_var, seed=13)
_compare("vsa256-cute-cross", layout, q, k, v, kv_var, q_var, 3, grad_out)
@pytest.mark.cuda
def test_vsa256_cute_inference_matches_training_forward() -> None:
"""The KV-owned backward metadata is only built when something requires
grad. Skipping it must not perturb the forward result."""
kv_var = torch.tensor([256, 173, 79, 256], dtype=torch.int32, device="cuda")
q_var = torch.full((3, ), _BLOCK, dtype=torch.int32, device="cuda")
q, k, v, _ = _make_inputs(3, 4, kv_var, q_var, seed=5)
with torch.no_grad():
out_infer = video_sparse_attn_bshd(
q.transpose(1, 2).contiguous(),
k.transpose(1, 2).contiguous(),
v.transpose(1, 2).contiguous(),
kv_var,
q_var,
2,
block_size=_BLOCK_SIZE_3D,
compress_attn_weight=None,
)
out_train, _ = _run_bshd(q, k, v, kv_var, q_var, 2)
torch.testing.assert_close(out_infer, out_train.transpose(1, 2).detach(), rtol=0, atol=0)
@pytest.mark.cuda
def test_vsa256_cute_lse_is_bhs() -> None:
"""The aux return is [B, H, S] on both entrypoints, matching the Triton
path's contract."""
from fastvideo_kernel.block_sparse_attn_256 import (block_sparse_attn_256, block_sparse_attn_256_bshd)
device = torch.device("cuda")
heads, dim, q_blocks, kv_blocks = 2, 128, 3, 4
sq, skv = q_blocks * _BLOCK, kv_blocks * _BLOCK
q = torch.randn(1, heads, sq, dim, device=device, dtype=torch.bfloat16)
k = torch.randn(1, heads, skv, dim, device=device, dtype=torch.bfloat16)
v = torch.randn(1, heads, skv, dim, device=device, dtype=torch.bfloat16)
vbs = torch.full((kv_blocks, ), _BLOCK, dtype=torch.int32, device=device)
mask = torch.zeros(1, heads, q_blocks, kv_blocks, dtype=torch.bool, device=device)
mask[..., :2] = True
_, lse_bhsd = block_sparse_attn_256(q, k, v, mask, vbs)
assert lse_bhsd.shape == (1, heads, sq), lse_bhsd.shape
_, lse_bshd = block_sparse_attn_256_bshd(
q.transpose(1, 2).contiguous(),
k.transpose(1, 2).contiguous(),
v.transpose(1, 2).contiguous(),
mask,
vbs,
)
assert lse_bshd.shape == (1, heads, sq), lse_bshd.shape
@@ -22,7 +22,6 @@ def _torch_vsa256_reference(
q_var: torch.Tensor,
kv_var: torch.Tensor,
topk_logical: int,
compress_attn_weight: torch.Tensor | None = None,
) -> torch.Tensor:
bsz, heads, _sq, dim = q.shape
q_blocks = q_var.numel()
@@ -56,8 +55,6 @@ def _torch_vsa256_reference(
logits = logits.masked_fill(~token_mask, float("-inf"))
prob = torch.softmax(logits, dim=-1)
out_s = torch.matmul(prob, vf).to(q.dtype)
if compress_attn_weight is not None:
return out_c * compress_attn_weight + out_s
return out_c + out_s
@@ -1,104 +0,0 @@
"""Regression: Triton block-sparse backward gradient parity at realistic activation scale.
The backward used to fold ``sm_scale / ln(2)`` into K in bf16 before the
exp2-based logit recompute. The bf16 rounding error on the pre-scaled K grows
proportionally to |logit| and exp2 amplifies it into exponentially wrong
probabilities, so dQ/dK/dV were correct at unit scale (every pre-existing test)
but off by orders of magnitude at real activation magnitudes.
This test sweeps the input scale and checks the Triton kernel's gradients
against an fp32 masked-dense SDPA reference. The unit-scale case is the
control (it passed even with the broken kernel); the large-scale cases are
the regression.
"""
import pytest
import torch
from fastvideo_kernel.block_sparse_attn import _map_to_index, block_sparse_attn_triton
from .utils import generate_block_sparse_mask_for_function
BLOCK = 64
@pytest.fixture(autouse=True)
def _seed_rng():
"""Pin the RNG so these cases do not depend on what ran before them.
Same convention as test_vsa_varlen.py: every tensor here comes from the
global torch RNG and the checks use tight thresholds, so an unseeded run
would shift inputs whenever an earlier test file draws a different number
of randoms.
"""
torch.manual_seed(0)
torch.cuda.manual_seed_all(0)
def _dense_reference(q, k, v, block_mask):
"""fp32 masked-dense SDPA over the token-expanded block mask.
q/k/v: [B, H, S, D]; block_mask: [B, H, S // BLOCK, S // BLOCK] bool.
"""
qf, kf, vf = q.float(), k.float(), v.float()
token_mask = block_mask.repeat_interleave(BLOCK, dim=-2).repeat_interleave(BLOCK, dim=-1)
logits = torch.matmul(qf, kf.transpose(-2, -1)) * (q.shape[-1]**-0.5)
logits = logits.masked_fill(~token_mask, float("-inf"))
return torch.matmul(logits.softmax(dim=-1), vf)
@pytest.mark.cuda
@pytest.mark.parametrize("scale", [1.0, 4.0, 16.0])
def test_triton_backward_grad_parity_across_input_scales(scale: float) -> None:
"""Kernel dQ/dK/dV must stay within a few percent of the fp32 reference
regardless of input magnitude.
With the bf16 K pre-scaling bug, scale<=4.0 passes at this geometry while
scale=16.0 fails (measured on GB200: dq relative L2 error 5.9e-1 vs 6.9e-3
fixed); at larger geometries and real activation magnitudes the broken
kernel is off by orders of magnitude. The passing unit-scale case is
exactly how the bug survived the original test suite.
"""
if not torch.cuda.is_available():
pytest.skip("CUDA is required")
device = torch.device("cuda")
dtype = torch.bfloat16
batch, heads, dim = 1, 4, 128
num_blocks = 8
seq = num_blocks * BLOCK
q = torch.randn(batch, heads, seq, dim, device=device, dtype=dtype) * scale
k = torch.randn(batch, heads, seq, dim, device=device, dtype=dtype) * scale
v = torch.randn(batch, heads, seq, dim, device=device, dtype=dtype)
grad_out = torch.randn_like(q)
block_mask = generate_block_sparse_mask_for_function(heads, num_blocks, num_blocks, k=3,
device=device).unsqueeze(0)
q2k_idx, q2k_num = _map_to_index(block_mask)
variable_block_sizes = torch.full((num_blocks, ), BLOCK, dtype=torch.int32, device=device)
q_ker, k_ker, v_ker = (t.detach().clone().requires_grad_(True) for t in (q, k, v))
out_ker, _ = block_sparse_attn_triton(q_ker, k_ker, v_ker, q2k_idx, q2k_num, variable_block_sizes)
(out_ker.float() * grad_out.float()).sum().backward()
q_ref, k_ref, v_ref = (t.detach().clone().requires_grad_(True) for t in (q, k, v))
out_ref = _dense_reference(q_ref, k_ref, v_ref, block_mask)
(out_ref * grad_out.float()).sum().backward()
# Forward is exact at any scale; this pins the harness itself.
fwd_rel = ((out_ker.float() - out_ref).norm() / out_ref.norm()).item()
assert fwd_rel < 2e-2, f"scale={scale}: forward rel err {fwd_rel:.3e}"
for name, g_ker, g_ref in (
("dq", q_ker.grad, q_ref.grad),
("dk", k_ker.grad, k_ref.grad),
("dv", v_ker.grad, v_ref.grad),
):
assert torch.isfinite(g_ker).all().item(), f"scale={scale}: non-finite {name}"
ref_norm = g_ref.float().norm()
rel = ((g_ker.float() - g_ref.float()).norm() / ref_norm.clamp_min(1e-12)).item()
ratio = (g_ker.float().norm() / ref_norm.clamp_min(1e-12)).item()
print(f"scale={scale} {name}: rel_l2={rel:.4e} norm_ratio={ratio:.4f}")
assert rel < 5e-2, f"scale={scale}: {name} rel l2 err {rel:.3e} >= 5e-2"
assert 0.98 < ratio < 1.02, f"scale={scale}: {name} grad-norm ratio {ratio:.4f}"
-14
View File
@@ -23,20 +23,6 @@ from fastvideo_kernel.block_sparse_attn import (
from fastvideo_kernel.block_sparse_attn_varlen import block_sparse_attn_varlen
@pytest.fixture(autouse=True)
def _seed_rng():
"""Pin the RNG so these cases do not depend on what ran before them.
Every tensor and every variable block size here comes from the global
torch RNG, and the gradient checks use a tight max_rel threshold. Without
a seed the inputs shift whenever an earlier test file draws a different
number of randoms, which surfaces as an unrelated-looking failure in
whichever case happens to land on unlucky data.
"""
torch.manual_seed(0)
torch.cuda.manual_seed_all(0)
def _reference_per_sequence(
q_list,
k_list,
+1 -6
View File
@@ -19,12 +19,7 @@ from fastvideo.attention.backends.abstract import (
from fastvideo.logger import init_logger
logger = init_logger(__name__)
# Every worker records the loaded FlashAttention implementation so a
# distributed profiling log contains one backend receipt per rank.
logger.info("Worker %s Using FlashAttention-%s backend",
os.environ.get("RANK", "0"),
fa_version,
local_main_process_only=False)
logger.info("Using FlashAttention-%s backend", fa_version)
# FP4 FA4 support: quantize Q/K to NVFP4 E2M1 for block-scaled MMA on Blackwell.
# Requires: flash-attention-fp4, flashinfer, cutlass-dsl. Enable via nvfp4_fa4=True kwarg.
@@ -5,17 +5,13 @@ H3 runs one joint bidirectional attention over
``[text | condition keyframes | audio | generated video]``, so this
backend differs from the Wan-tuned ``video_sparse_attn``:
- Tiles are ``[segment-pure prefix chunks] + [3D video tiles]``; prefix
tiles never straddle segment boundaries. The tile size is selectable at
metadata build time: 256 tokens ``(4,8,8)`` (default) or 64 tokens
``(4,4,4)`` (see ``VSA_H3_TILE_SHAPES``).
- Tiles are ``[segment-pure prefix chunks] + [3D (4,8,8) video tiles]``;
prefix tiles never straddle segment boundaries.
- Selection is pure Python on pooled tile scores; the block-sparse kernel
consumes an explicit bool mask, so no kernel changes are needed.
- The compression branch is gated by ``to_gate_compress``, which the base
H3 checkpoint does not carry: the loader zero-initializes it, so
untrained inference is exactly pure sparse and finetuning can learn the
gate. VSA-distilled students (e.g. FastVideo-Minimax-H3-Preview) ship
trained gates, which load and activate the branch.
- The compression branch is gated by ``to_gate_compress``, which the H3
checkpoint does not carry: the loader zero-initializes it, so untrained
inference is exactly pure sparse and finetuning can learn the gate.
- 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 ablation axis,
@@ -24,46 +20,22 @@ backend differs from the Wan-tuned ``video_sparse_attn``:
(``vsa_dense_first_n_steps``, ``vsa_dense_layers``) let mixed schedules
run the diffuse steps/layers dense while pushing the rest harder.
At tile 256 this targets sm10.x through the FA4 CuTe 256-tile path
Targets sm10.x through the FA4 CuTe 256-tile path
(``FASTVIDEO_VSA_CUTEDSL=1``); the Triton 256→64 expansion is the
fallback and keeps identical mask semantics. At tile 64 the block map is
already at the kernels' native 64-token granularity, so both forward and
backward run the Triton block-sparse kernels directly (no expansion,
``FASTVIDEO_VSA_CUTEDSL`` does not apply). A third, opt-in route exists
for the tile-64 FORWARD only: ``FASTVIDEO_VSA_SM100A=1`` sends no-grad
forwards through the sm_100a CUDA block-sparse kernel
(``fastvideo_kernel.block_sparse_attn_sm100a``, upstream PR #1719 plus
our per-q-tile ``q2k_num`` fix) when the extension is built, the device
is sm_100, and the geometry qualifies; grad-tracking forwards and every
backward stay on Triton unchanged. If the env is set but a precondition
fails, the route logs one warning and falls back.
fallback and keeps identical mask semantics.
"""
import functools
import math
import os
from dataclasses import dataclass
from typing import Any
import torch
try:
from fastvideo_kernel.block_sparse_attn import block_sparse_attn as block_sparse_attn_64_bhsd
from fastvideo_kernel.block_sparse_attn_256 import block_sparse_attn_256_bshd
from fastvideo_kernel.triton_kernels.index import map_to_index
except ImportError:
block_sparse_attn_64_bhsd = None
block_sparse_attn_256_bshd = None
map_to_index = None
try:
# Optional: only present in fastvideo_kernel builds that carry the sm_100a
# CUDA block-sparse forward (upstream PR #1719). The module itself imports
# fine without the compiled symbols (`_HAS_VSA_SM100A` is then False and
# `is_supported` says no), so this only guards *module* availability.
from fastvideo_kernel import block_sparse_attn_sm100a as _sm100a
except ImportError:
_sm100a = None
from fastvideo.attention.backends.abstract import (AttentionBackend, AttentionImpl, AttentionMetadata,
AttentionMetadataBuilder, layer_idx_from_prefix)
@@ -71,115 +43,51 @@ from fastvideo.attention.backends.video_sparse_attn import (compute_topk, constr
get_non_pad_index, get_tile_partition_indices,
scatter_into_tile_buf)
from fastvideo.attention.backends.video_sparse_attn_h3_probe import probe_enabled, record_probe
from fastvideo.logger import init_logger
logger = init_logger(__name__)
# Opt-in switch for the sm_100a CUDA forward on the tile-64 no-grad path.
VSA_SM100A_ENV = "FASTVIDEO_VSA_SM100A"
VSA_H3_TILE_SIZE = (4, 8, 8) # 256 elements -> FA4 CuTe fastpath on sm10.x (default)
VSA_H3_TILE_SIZE = (4, 8, 8) # 256 elements -> FA4 CuTe fastpath on sm10.x
_TILE_ELEMS = math.prod(VSA_H3_TILE_SIZE)
# Selectable tile geometries, keyed by element count (= the build-time
# ``tile_size``). 64 runs the native 64-token Triton block-sparse kernels for
# forward AND backward — the block map is already at kernel granularity, so no
# 256->64 mask expansion is involved and FASTVIDEO_VSA_CUTEDSL does not apply.
VSA_H3_TILE_SHAPES: dict[int, tuple[int, int, int]] = {
_TILE_ELEMS: VSA_H3_TILE_SIZE,
64: (4, 4, 4),
}
def token_tile_and_valid(variable_block_sizes: torch.Tensor,
tile_elems: int = _TILE_ELEMS) -> tuple[torch.Tensor, torch.Tensor]:
def token_tile_and_valid(variable_block_sizes: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
"""Per padded-token tile id and pad-validity mask.
The single encoding of the padding contract, shared by the probe and the
test oracle so they cannot drift from the backend's tile geometry.
``tile_elems`` must match the metadata the sizes came from
(``MiniMaxH3VSAMetadata.tile_elems``).
"""
device = variable_block_sizes.device
token_tile = torch.arange(variable_block_sizes.numel(), device=device).repeat_interleave(tile_elems)
token_valid = (torch.arange(tile_elems, device=device)[None, :] < variable_block_sizes[:, None]).reshape(-1)
token_tile = torch.arange(variable_block_sizes.numel(), device=device).repeat_interleave(_TILE_ELEMS)
token_valid = (torch.arange(_TILE_ELEMS, device=device)[None, :] < variable_block_sizes[:, None]).reshape(-1)
return token_tile, token_valid
def _validate_h3_tile_geometry(
prefix_segments: tuple[int, ...],
dit_seq_shape: tuple[int, int, int],
variable_block_sizes: torch.Tensor,
untile_combined_index: torch.Tensor,
tile_elems: int = _TILE_ELEMS,
) -> None:
"""Fail synchronously on out-of-bounds tile geometry.
Invariants the block-sparse kernel trusts without checking:
every tile's valid size is in (0, tile_elems]; the sizes sum to the
packed sequence length; and ``untile_combined_index`` maps each packed
row to exactly one non-pad slot of the padded tile buffer. A violation
would surface only as an async device fault at some later kernel or
collective (e.g. an FSDP all-gather), which is unattributable — so raise
here, once per cached geometry, with the numbers in hand.
"""
total = sum(prefix_segments) + math.prod(dit_seq_shape)
n_pad = variable_block_sizes.numel() * tile_elems
sizes_min = int(variable_block_sizes.min())
sizes_max = int(variable_block_sizes.max())
sizes_sum = int(variable_block_sizes.sum())
if sizes_min < 1 or sizes_max > tile_elems or sizes_sum != total:
raise ValueError(f"VSA-H3 tile sizes out of bounds for prefix={prefix_segments}, video={dit_seq_shape}, "
f"tile_elems={tile_elems}: min={sizes_min}, max={sizes_max}, sum={sizes_sum}, "
f"expected sum={total}.")
if untile_combined_index.numel() != total:
raise ValueError(f"VSA-H3 untile index has {untile_combined_index.numel()} entries for a packed "
f"sequence of {total} rows (prefix={prefix_segments}, video={dit_seq_shape}).")
idx_min = int(untile_combined_index.min())
idx_max = int(untile_combined_index.max())
if idx_min < 0 or idx_max >= n_pad:
# Range first: the pad-slot gather below would itself index out of
# bounds (the very async fault this guard exists to preempt).
raise ValueError(f"VSA-H3 untile index is not an injective map into non-pad slots: range "
f"[{idx_min}, {idx_max}] vs padded length {n_pad} "
f"(prefix={prefix_segments}, video={dit_seq_shape}).")
in_tile_offset = untile_combined_index % tile_elems
maps_into_pad = bool((in_tile_offset >= variable_block_sizes[untile_combined_index // tile_elems]).any())
if maps_into_pad or int(torch.unique(untile_combined_index).numel()) != total:
raise ValueError(f"VSA-H3 untile index is not an injective map into non-pad slots: "
f"pad-slot hit={maps_into_pad} "
f"(prefix={prefix_segments}, video={dit_seq_shape}).")
@functools.lru_cache(maxsize=10)
def _h3_tile_geometry(
prefix_segments: tuple[int, ...],
dit_seq_shape: tuple[int, int, int],
device: torch.device,
tile_shape: tuple[int, int, int] = VSA_H3_TILE_SIZE,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, int, int]:
"""Tile the packed sequence: segment-pure prefix chunks, then video tiles.
Returns (tile_partition_indices, variable_block_sizes,
untile_combined_index, num_prefix_tiles, num_video_tiles).
"""
tile_elems = math.prod(tile_shape)
prefix_len = sum(prefix_segments)
prefix_sizes: list[int] = []
for segment in prefix_segments:
full, rem = divmod(segment, tile_elems)
prefix_sizes.extend([tile_elems] * full)
full, rem = divmod(segment, _TILE_ELEMS)
prefix_sizes.extend([_TILE_ELEMS] * full)
if rem:
prefix_sizes.append(rem)
num_prefix_tiles = len(prefix_sizes)
ts_t, ts_h, ts_w = tile_shape
ts_t, ts_h, ts_w = VSA_H3_TILE_SIZE
t, h, w = dit_seq_shape
num_tiles = (math.ceil(t / ts_t), math.ceil(h / ts_h), math.ceil(w / ts_w))
video_sizes = construct_variable_block_sizes(dit_seq_shape, num_tiles, device, tile_shape)
video_sizes = construct_variable_block_sizes(dit_seq_shape, num_tiles, device, VSA_H3_TILE_SIZE)
num_video_tiles = int(video_sizes.numel())
video_indices = get_tile_partition_indices(dit_seq_shape, tile_shape, device) + prefix_len
video_indices = get_tile_partition_indices(dit_seq_shape, VSA_H3_TILE_SIZE, device) + prefix_len
tile_partition_indices = torch.cat([
torch.arange(prefix_len, device=device, dtype=torch.long),
video_indices,
@@ -192,11 +100,9 @@ def _h3_tile_geometry(
# get_non_pad_index is lru-cached on tensor identity; variable_block_sizes
# is itself cached by this function, so the identity stays stable.
non_pad_index = get_non_pad_index(variable_block_sizes, tile_elems)
non_pad_index = get_non_pad_index(variable_block_sizes, _TILE_ELEMS)
untile_combined_index = non_pad_index[torch.argsort(tile_partition_indices)]
# One-time (lru-cached) synchronous bounds check; see _validate_h3_tile_geometry.
_validate_h3_tile_geometry(prefix_segments, dit_seq_shape, variable_block_sizes, untile_combined_index, tile_elems)
return (tile_partition_indices, variable_block_sizes, untile_combined_index, num_prefix_tiles, num_video_tiles)
@@ -233,9 +139,6 @@ class MiniMaxH3VSAMetadata(AttentionMetadata):
exempt: bool
variable_block_sizes: torch.Tensor
untile_combined_index: torch.Tensor
# tokens per tile (256 or 64); selects the tile geometry AND the kernel
# route in forward() (256 -> VSA-256 CuTe/Triton, 64 -> native Triton)
tile_elems: int = _TILE_ELEMS
# layers forced dense regardless of sparsity (probe-guided opt-outs)
dense_layers: tuple[int, ...] = ()
# Single-slot holder for the padded tile buffer, owned by the BUILDER so
@@ -255,28 +158,24 @@ class MiniMaxH3VSAMetadataBuilder(AttentionMetadataBuilder):
pass
def build( # type: ignore
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, ...] = (),
tile_size: int = _TILE_ELEMS,
**kwargs: dict[str, Any],
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, ...] = (),
**kwargs: dict[str, Any],
) -> MiniMaxH3VSAMetadata:
tile_shape = VSA_H3_TILE_SHAPES.get(int(tile_size))
if tile_shape is None:
raise ValueError(f"VSA-H3 tile_size must be one of {sorted(VSA_H3_TILE_SHAPES)}, got {tile_size!r}")
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)
total_seq_length = sum(prefix_segments) + math.prod(dit_seq_shape)
(_tile_partition_indices, variable_block_sizes, untile_combined_index, num_prefix_tiles,
num_video_tiles) = _h3_tile_geometry(prefix_segments, dit_seq_shape, device, tile_shape)
num_video_tiles) = _h3_tile_geometry(prefix_segments, dit_seq_shape, device)
return MiniMaxH3VSAMetadata(
current_timestep=current_timestep,
@@ -287,14 +186,13 @@ class MiniMaxH3VSAMetadataBuilder(AttentionMetadataBuilder):
exempt=exempt,
variable_block_sizes=variable_block_sizes,
untile_combined_index=untile_combined_index,
tile_elems=int(tile_size),
dense_layers=tuple(int(layer) for layer in dense_layers),
tile_buf_holder=self._tile_buf_holder,
)
def _pool_tiles(x: torch.Tensor, variable_block_sizes: torch.Tensor, tile_elems: int = _TILE_ELEMS) -> torch.Tensor:
"""fp32 mean over each tile_elems-token tile. x: [B, S_pad, H, D] -> [B, H, n_tiles, D].
def _pool_tiles(x: torch.Tensor, variable_block_sizes: torch.Tensor) -> torch.Tensor:
"""fp32 mean over each 256-token tile. x: [B, S_pad, H, D] -> [B, H, n_tiles, D].
Pad positions in the tile buffer are guaranteed zero (zeros-init, never
written), so a plain sum with fp32 accumulation needs no validity mask
@@ -302,8 +200,8 @@ def _pool_tiles(x: torch.Tensor, variable_block_sizes: torch.Tensor, tile_elems:
the masked mean exactly.
"""
batch, seq_len, heads, dim = x.shape
n_tiles = seq_len // tile_elems
pooled = x.view(batch, n_tiles, tile_elems, heads, dim).sum(dim=2, dtype=torch.float32)
n_tiles = seq_len // _TILE_ELEMS
pooled = x.view(batch, n_tiles, _TILE_ELEMS, heads, dim).sum(dim=2, dtype=torch.float32)
pooled = pooled / variable_block_sizes.view(1, -1, 1, 1)
return pooled.permute(0, 2, 1, 3)
@@ -334,24 +232,6 @@ def _build_block_mask(
return mask
def _sm100a_unavailable_reason(sm100a_mod: Any, query_bhsd: torch.Tensor, variable_block_sizes: torch.Tensor,
grad_mode: bool) -> str | None:
"""Why the opt-in sm_100a forward route cannot run here, or None if it can.
Pure decision logic, split out so the routing is unit-testable without a
GPU or the compiled extension (tests substitute ``sm100a_mod``). Order
matters only for the message: the cheapest, most actionable reason first.
"""
if sm100a_mod is None:
return "fastvideo_kernel.block_sparse_attn_sm100a is not installed"
if grad_mode:
return "inputs require grad and the sm_100a kernel is forward-only; grad paths keep Triton"
if not sm100a_mod.is_supported(query_bhsd, variable_block_sizes):
return ("block_sparse_attn_sm100a.is_supported returned False (needs an sm_100 device, a built "
"extension, bf16, head_dim 128, an even tile count, and integer tile sizes)")
return None
class MiniMaxH3VSAImpl(AttentionImpl):
def __init__(
@@ -379,7 +259,7 @@ class MiniMaxH3VSAImpl(AttentionImpl):
f"got {x.shape[1]}. A non-packed sequence (e.g. the token refiner) is "
"routed to the VSA-H3 backend; exclude it from the supported backends.")
n_tiles = attn_metadata.variable_block_sizes.numel()
target_shape = (x.shape[0], n_tiles * attn_metadata.tile_elems, x.shape[-2], x.shape[-1])
target_shape = (x.shape[0], n_tiles * _TILE_ELEMS, x.shape[-2], x.shape[-1])
# single scatter: untile_combined_index maps original row i to its
# padded slot, so this is exactly the inverse of postprocess_output
@@ -401,11 +281,7 @@ class MiniMaxH3VSAImpl(AttentionImpl):
gate_compress: torch.Tensor | None,
attn_metadata: MiniMaxH3VSAMetadata,
) -> torch.Tensor:
tile_elems = attn_metadata.tile_elems
if tile_elems == 64:
if block_sparse_attn_64_bhsd is None:
raise NotImplementedError("fastvideo_kernel.block_sparse_attn is not installed")
elif block_sparse_attn_256_bshd is None:
if block_sparse_attn_256_bshd is None:
raise NotImplementedError("fastvideo_kernel.block_sparse_attn_256 is not installed")
# probe-guided per-layer opt-out: diffuse layers run dense (all-True
@@ -415,8 +291,8 @@ class MiniMaxH3VSAImpl(AttentionImpl):
scores = None
if layer_sparsity > 0.0 or gate_compress is not None or probe_dir is not None:
q_pooled = _pool_tiles(query, attn_metadata.variable_block_sizes, tile_elems)
k_pooled = _pool_tiles(key, attn_metadata.variable_block_sizes, tile_elems)
q_pooled = _pool_tiles(query, attn_metadata.variable_block_sizes)
k_pooled = _pool_tiles(key, attn_metadata.variable_block_sizes)
scores = torch.matmul(q_pooled, k_pooled.transpose(-2, -1)) / (query.shape[-1]**0.5)
if probe_dir is not None:
record_probe(probe_dir, self.layer_idx, query, key, scores, attn_metadata)
@@ -433,75 +309,18 @@ class MiniMaxH3VSAImpl(AttentionImpl):
attn_metadata.exempt,
)
if tile_elems == 64:
# Native 64-token path: the block map is already at the kernels'
# granularity. Both 64-token entries take BHSD ([B, H, S_pad, D]);
# mirror block_sparse_attn_256_bshd's Triton branch and transpose
# around the call.
q_bhsd = query.transpose(1, 2).contiguous()
k_bhsd = key.transpose(1, 2).contiguous()
v_bhsd = value.transpose(1, 2).contiguous()
# Opt-in sm_100a CUDA forward (upstream PR #1719 + per-q-tile
# q2k_num fix). Forward-only: grad-tracking calls stay on Triton
# so autograd keeps the Triton fwd+bwd pairing untouched. The
# kernel does return an LSE in Triton's M format, so a future
# fwd/bwd pairing is possible, but it is not built here.
use_sm100a = False
if os.environ.get(VSA_SM100A_ENV, "0") == "1":
grad_mode = torch.is_grad_enabled() and (query.requires_grad or key.requires_grad
or value.requires_grad)
reason = _sm100a_unavailable_reason(_sm100a, q_bhsd, attn_metadata.variable_block_sizes, grad_mode)
if reason is None and map_to_index is None:
reason = "fastvideo_kernel.triton_kernels.index (map_to_index) is not importable"
if reason is None:
use_sm100a = True
elif not torch.compiler.is_compiling():
logger.warning_once(f"{VSA_SM100A_ENV}=1 but falling back to the Triton-64 kernels: {reason}")
if use_sm100a:
# The sm_100a entry is index-native; compact the bool map the
# same way the Triton bool entry does internally. Per-row
# counts are NON-uniform here (prefix query tiles are dense,
# video tiles run prefix+top-k) -- legal for the fixed kernel,
# silently wrong on the pre-fix upstream one.
q2k_idx, q2k_num = map_to_index(mask)
out_bhsd, _ = _sm100a.block_sparse_attn_sm100a(
q_bhsd,
k_bhsd,
v_bhsd,
q2k_idx,
q2k_num,
attn_metadata.variable_block_sizes.to(torch.int32),
need_lse=False,
)
else:
out_bhsd, _ = block_sparse_attn_64_bhsd(
q_bhsd,
k_bhsd,
v_bhsd,
mask,
attn_metadata.variable_block_sizes,
)
out = out_bhsd.transpose(1, 2).contiguous()
else:
out, _ = block_sparse_attn_256_bshd(query, key, value, mask, attn_metadata.variable_block_sizes)
out, _ = block_sparse_attn_256_bshd(query, key, value, mask, attn_metadata.variable_block_sizes)
if gate_compress is not None:
# Wan-style compression branch: dense attention over pooled tiles,
# broadcast to each tile's rows, scaled by the learned gate
# (zero-initialized for H3 => branch contributes nothing until
# finetuned; the model layer skips it entirely for all-zero gates).
v_pooled = _pool_tiles(value, attn_metadata.variable_block_sizes, tile_elems)
v_pooled = _pool_tiles(value, attn_metadata.variable_block_sizes)
out_c = torch.matmul(torch.softmax(scores, dim=-1), v_pooled) # [B, H, n_tiles, D]
out_c = out_c.permute(0, 2, 1, 3).to(out.dtype) # [B, n_tiles, H, D]
batch, seq_len, heads, dim = out.shape
batch, _, heads, dim = out.shape
n_tiles = attn_metadata.variable_block_sizes.numel()
# Out-of-place: on the CuTe backend ``out`` is the tensor FA4's
# autograd node saved for its backward, so an in-place add here
# bumps its version counter and backward dies with "one of the
# variables needed for gradient computation has been modified".
out_tiled = out.view(batch, n_tiles, tile_elems, heads, dim)
gate_tiled = gate_compress.view(batch, n_tiles, tile_elems, heads, dim)
out = (out_tiled + out_c.unsqueeze(2) * gate_tiled).view(batch, seq_len, heads, dim)
out.view(batch, n_tiles, _TILE_ELEMS, heads,
dim).addcmul_(out_c.unsqueeze(2), gate_compress.view(batch, n_tiles, _TILE_ELEMS, heads, dim))
return out
@@ -60,7 +60,7 @@ def record_probe(
gen = torch.Generator(device="cpu").manual_seed(step * 1000 + layer)
# sample among video rows in the PADDED/tiled domain that are non-pad
from fastvideo.attention.backends.video_sparse_attn_h3 import token_tile_and_valid
token_tile, token_valid = token_tile_and_valid(attn_metadata.variable_block_sizes, attn_metadata.tile_elems)
token_tile, token_valid = token_tile_and_valid(attn_metadata.variable_block_sizes)
video_rows = torch.nonzero((token_tile >= P) & token_valid, as_tuple=False).flatten()
idx = video_rows[torch.randint(0, video_rows.numel(), (_TRUE_ROWS, ), generator=gen).to(query.device)]
@@ -62,8 +62,6 @@ class MiniMaxH3Qwen3VLArchConfig(TextEncoderArchConfig):
hidden_size: int = 5120
intermediate_size: int = 25600
num_hidden_layers: int = 64
output_hidden_state_index: int = 50
num_hidden_layers_override: int | None = 50
num_attention_heads: int = 64
num_key_value_heads: int = 8
head_dim: int = 128
@@ -109,7 +107,7 @@ class MiniMaxH3Qwen3VLArchConfig(TextEncoderArchConfig):
vision_initializer_range: float = 0.02
vision_deepstack_visual_indexes: tuple[int, ...] = (8, 16, 24)
output_hidden_states: bool = False
output_hidden_states: bool = True
stacked_params_mapping: list[tuple[str, str, str | int]] = field(default_factory=list)
_fsdp_shard_conditions: list = field(default_factory=lambda: [
_is_language_transformer_layer,
@@ -120,17 +118,6 @@ class MiniMaxH3Qwen3VLArchConfig(TextEncoderArchConfig):
])
def __post_init__(self) -> None:
if self.output_hidden_state_index <= 0 or self.output_hidden_state_index > self.num_hidden_layers:
raise ValueError("MiniMax H3 Qwen3-VL output_hidden_state_index must be in "
f"[1, {self.num_hidden_layers}], got {self.output_hidden_state_index}.")
if self.num_hidden_layers_override is not None:
if self.num_hidden_layers_override <= 0:
raise ValueError("MiniMax H3 Qwen3-VL num_hidden_layers_override must be positive or None.")
if self.num_hidden_layers_override < self.output_hidden_state_index:
raise ValueError("MiniMax H3 Qwen3-VL num_hidden_layers_override must build through "
f"hidden_states[{self.output_hidden_state_index}], got "
f"{self.num_hidden_layers_override}.")
rope_scaling = dict(self.rope_scaling or {})
self.mrope_interleaved = bool(rope_scaling.get("mrope_interleaved", self.mrope_interleaved))
if not self.mrope_interleaved:
+18 -43
View File
@@ -808,19 +808,14 @@ class VideoGenerator:
latent_batch_size = _infer_latent_batch_size(batch)
is_latent_output = fastvideo_args.output_type == "latent"
needs_frame_output = batch.return_frames or (batch.save_video and not is_latent_output)
# A populated ``samples`` has exactly one consumer — the result
# dict (``"samples": samples if batch.return_frames else None``).
# Post-decode frame building reads ``output_batch.output``
# directly (the GPU ``vid_u8`` path), not ``samples``. So when
# ``return_frames=False`` the pinned fp32 alloc + D->H copy are
# dead weight — the CLI generate flow (``save_video=True``,
# ``return_frames=False``) hits this on every call.
# ``output_type == "latent"`` keeps its existing branch (shape
# mismatch falls through to ``.cpu()`` below) for callers that
# *do* ask for the latent samples via ``return_frames=True``.
needs_samples_buffer = batch.return_frames or needs_frame_output
# When ``output_type == "latent"`` the forward output has latent
# shape (e.g. ``[B, C_latent, T_latent, H_latent, W_latent]``)
# rather than the pre-allocation's pixel shape. Skip the pinned
# ~50 MB buffer entirely. Also skip it for metadata-only calls;
# neither the result nor save path will consume the decoded tensor.
# ``skip_pixel_prealloc`` also gates the slow-path warning.
needs_samples_out = batch.return_frames
skip_pixel_prealloc = is_latent_output or not needs_samples_out
skip_pixel_prealloc = is_latent_output or not needs_samples_buffer
if skip_pixel_prealloc:
samples = torch.empty(0, device='cpu')
else:
@@ -840,11 +835,9 @@ class VideoGenerator:
"This usually means the executor/pipeline failed earlier.")
audio_only = bool(output_batch.extra.get("audio_only"))
if not needs_samples_out:
# Nothing downstream reads ``samples`` (the result dict
# returns None when ``return_frames=False``); keep the empty
# placeholder allocated above and skip the fp32 D->H copy
# entirely.
if not needs_samples_buffer or (audio_only and not batch.return_frames):
# Metadata-only/audio-only request: keep the empty placeholder and
# avoid the decoded tensor D->H copy.
pass
elif audio_only:
# Audio-only return-frames requests expose the small placeholder
@@ -876,13 +869,8 @@ class VideoGenerator:
# `GenerationResult.size` describes the produced media, not only the
# base-stage request. Refiner pipelines can change the final pixel
# dimensions, so derive this result metadata from the decoded output.
# Read the geometry from `output_batch.output` (a shape-only access,
# no D->H copy): when `return_frames=False` the `samples` mirror
# stays an empty placeholder and no longer carries the decoded
# shape. Metadata-only calls keep the request fallback and never
# inspect the (possibly dropped) worker output.
output_size = _resolve_output_size(
output_batch.output if needs_frame_output else samples,
samples,
(target_height, target_width, batch.num_frames),
pixel_output=not is_latent_output and not audio_only,
)
@@ -894,26 +882,13 @@ class VideoGenerator:
elif not needs_frame_output:
frames = None
else:
# Quantize on the source device (typically CUDA) BEFORE the
# device->host copy. `samples` above is just the pinned-CPU
# mirror of `output_batch.output` (`samples.copy_(output)` or
# `output.cpu()`) with no intervening preprocessing, so reading
# `output_batch.output` here is the same data. The old path
# paid a full fp32 video D->H copy (which scales with
# resolution x frames x batch) and then a single-threaded
# per-frame CPU *255/cast loop. Casting to uint8 on-device
# makes the transfer 4x smaller, ships it in a single copy,
# and moves the elementwise work onto the GPU. clamp_() also
# fixes a latent overflow bug: VAE output slightly outside
# [0, 1] wrapped mod 256 in the old unclamped cast.
# (Equivalence is SSIM-gated, not bit-exact: float->uint8
# differs <=1 LSB CPU vs GPU.)
src = output_batch.output
vid_u8 = (src * 255).clamp_(0, 255).to(torch.uint8)
vid_u8 = rearrange(vid_u8, "b c t h w -> t b c h w").cpu()
frames = [
torchvision.utils.make_grid(x, nrow=6).permute(1, 2, 0).squeeze(-1).contiguous().numpy() for x in vid_u8
]
videos = rearrange(samples, "b c t h w -> t b c h w")
frames = []
for x in videos:
x = torchvision.utils.make_grid(x, nrow=6)
x = x.permute(1, 2, 0).squeeze(-1)
x = (x * 255).to(torch.uint8)
frames.append(x.contiguous().cpu().numpy())
postprocess_time = time.perf_counter() - postprocess_start
logger.info("PostDecodeFrameProcessStage completed in %.3f s", postprocess_time)
if logging_info is not None:
+1 -35
View File
@@ -21,17 +21,12 @@ if TYPE_CHECKING:
FASTVIDEO_TRACE_FUNCTION: int = 0
FASTVIDEO_ATTENTION_BACKEND: str | None = None
FASTVIDEO_FA4: bool = False
FASTVIDEO_MINIMAX_H3_FUSIONS: str = ""
FASTVIDEO_VAE_PARALLEL_DECODE: bool = False
FASTVIDEO_VAE_PARALLEL_ENCODE: bool = False
FASTVIDEO_VAE_PARALLEL_DECODE_STRATEGY: str | None = None
FASTVIDEO_WORKER_MULTIPROC_METHOD: str = "spawn"
FASTVIDEO_TARGET_DEVICE: str = "cuda"
MAX_JOBS: str | None = None
NVCC_THREADS: str | None = None
CMAKE_BUILD_TYPE: str | None = None
VERBOSE: bool = False
FASTVIDEO_NVTX_PROFILE: bool = False
FASTVIDEO_TORCH_PROFILER_DIR: str | None = None
FASTVIDEO_TORCH_PROFILER_RECORD_SHAPES: bool = False
FASTVIDEO_TORCH_PROFILER_WITH_PROFILE_MEMORY: bool = False
@@ -222,34 +217,10 @@ environment_variables: dict[str, Callable[[], Any]] = {
"FASTVIDEO_FA4":
lambda: os.getenv("FASTVIDEO_FA4", "0") != "0",
# If set (=1), MiniMax-H3 VAE decode (and, with the ENCODE variant,
# reference-video encode) round-robins its temporal chunks across the
# sequence-parallel ranks instead of running serially on the output rank.
# Folded into FastVideoArgs.vae_parallel_decode / vae_parallel_encode at
# construction (parse-once). The STRATEGY variant picks the chunk
# transport collective: "gather" (default) or "all_gather".
"FASTVIDEO_VAE_PARALLEL_DECODE":
lambda: os.getenv("FASTVIDEO_VAE_PARALLEL_DECODE", "0") != "0",
"FASTVIDEO_VAE_PARALLEL_ENCODE":
lambda: os.getenv("FASTVIDEO_VAE_PARALLEL_ENCODE", "0") != "0",
"FASTVIDEO_VAE_PARALLEL_DECODE_STRATEGY":
lambda: os.getenv("FASTVIDEO_VAE_PARALLEL_DECODE_STRATEGY", None),
# Opt-in MiniMax-H3 inference-only Triton fusions adapted from the
# NVlabs/Sana Sol-Engine implementation. Accepts `all`, `1`, or a
# comma-separated subset of `modulate,qknorm_rope,swiglu`. An empty value
# (the default), `0`, or `none` keeps the eager implementation.
"FASTVIDEO_MINIMAX_H3_FUSIONS":
lambda: os.getenv("FASTVIDEO_MINIMAX_H3_FUSIONS", ""),
# Use dedicated multiprocess context for workers.
"FASTVIDEO_WORKER_MULTIPROC_METHOD":
lambda: os.getenv("FASTVIDEO_WORKER_MULTIPROC_METHOD", "spawn"),
# Emit lightweight NVTX ranges for external profilers such as Nsight Systems.
"FASTVIDEO_NVTX_PROFILE":
lambda: os.getenv("FASTVIDEO_NVTX_PROFILE", "0") != "0",
# Enables torch profiler if set. Path to the directory where torch profiler
# traces are saved. Note that it must be an absolute path.
"FASTVIDEO_TORCH_PROFILER_DIR":
@@ -279,12 +250,7 @@ environment_variables: dict[str, Callable[[], Any]] = {
# not profile flops.
"FASTVIDEO_TORCH_PROFILER_WITH_FLOPS":
lambda: bool(os.getenv("FASTVIDEO_TORCH_PROFILER_WITH_FLOPS", "0") != "0"),
# Wait steps per profiling cycle (torch.profiler.schedule wait parameter)
# Defaults to 2 if not set.
# Warmup steps per profiling cycle (torch.profiler.schedule warmup parameter)
# Defaults to 1 if not set.
# Active steps per profiling cycle (torch.profiler.schedule active parameter)
# Defaults to 2 if not set.
# Comma-separated names of registered profiler regions to capture.
"FASTVIDEO_TORCH_PROFILE_REGIONS":
lambda: os.getenv("FASTVIDEO_TORCH_PROFILE_REGIONS", ""),
-51
View File
@@ -146,19 +146,6 @@ class FastVideoArgs:
vae_cpu_offload: bool = True
pin_cpu_memory: bool = True
# Sequence-parallel MiniMax-H3 VAE (opt-in, default off). With SP > 1 the
# video VAE's temporal chunks (decode) and clips (reference encode) are
# round-robined across the sequence-parallel ranks and reassembled
# bit-exactly on the group's first rank instead of running serially on
# one rank while the others idle. ``__post_init__`` folds the
# FASTVIDEO_VAE_PARALLEL_DECODE / FASTVIDEO_VAE_PARALLEL_ENCODE env vars
# into these fields (parse-once, like attention_backend), and
# FASTVIDEO_VAE_PARALLEL_DECODE_STRATEGY overrides the chunk-transport
# collective ("gather" or "all_gather").
vae_parallel_decode: bool = False
vae_parallel_encode: bool = False
vae_parallel_decode_strategy: str | None = None
# Compilation
# ``enable_torch_compile`` covers the DiT path (transformer,
# transformer_2, and the LTX-2 stage-2 transformer_refine).
@@ -182,7 +169,6 @@ class FastVideoArgs:
# VSA parameters
VSA_sparsity: float = 0.0 # inference/validation sparsity
VSA_tile_size: int = 256 # VSA-H3 tile size (256 or 64); 64 = native Triton path
# V-MoBA parameters
moba_config_path: str | None = None
@@ -300,27 +286,8 @@ class FastVideoArgs:
env_backend = envs.FASTVIDEO_ATTENTION_BACKEND
if env_backend is not None and backend_name_to_enum(env_backend) is not None:
self.attention_backend = env_backend
self._fold_vae_parallel_env()
self.check_fastvideo_args()
def _fold_vae_parallel_env(self) -> None:
"""Parse-once adapters for the sequence-parallel VAE env vars."""
import fastvideo.envs as envs
# Mirrors fastvideo.models.vaes.minimax_h3_parallel.DECODE_GATHER_STRATEGIES /
# DEFAULT_DECODE_GATHER_STRATEGY (kept literal here so constructing args
# never imports model modules; a unit test pins the two in sync).
strategies = ("gather", "all_gather")
if not self.vae_parallel_decode and envs.FASTVIDEO_VAE_PARALLEL_DECODE:
self.vae_parallel_decode = True
if not self.vae_parallel_encode and envs.FASTVIDEO_VAE_PARALLEL_ENCODE:
self.vae_parallel_encode = True
if self.vae_parallel_decode_strategy is None:
self.vae_parallel_decode_strategy = envs.FASTVIDEO_VAE_PARALLEL_DECODE_STRATEGY or "gather"
if self.vae_parallel_decode_strategy not in strategies:
raise ValueError(f"vae_parallel_decode_strategy must be one of {strategies}, "
f"got {self.vae_parallel_decode_strategy!r}.")
def _apply_transformer_quant(self) -> None:
"""Pin the typed ``transformer_quant`` instance onto ``dit_config``.
@@ -664,18 +631,6 @@ class FastVideoArgs:
"Pin memory for CPU offload. Only added as a temp workaround if it throws \"CUDA error: invalid argument\". "
"Should be enabled in almost all cases",
)
parser.add_argument(
"--vae-parallel-decode",
action=StoreBoolean,
help="With sequence parallelism, round-robin MiniMax-H3 VAE decode chunks across the SP ranks "
"and reassemble bit-exactly on the output rank (default: serial decode on the output rank)",
)
parser.add_argument(
"--vae-parallel-encode",
action=StoreBoolean,
help="With sequence parallelism, round-robin MiniMax-H3 reference-video VAE encode clips across "
"the SP ranks; every rank keeps the identical full encoding (default: serial encode on every rank)",
)
parser.add_argument(
"--disable-autocast",
action=StoreBoolean,
@@ -689,12 +644,6 @@ class FastVideoArgs:
default=FastVideoArgs.VSA_sparsity,
help="Validation sparsity for VSA",
)
parser.add_argument(
"--VSA-tile-size",
type=int,
default=FastVideoArgs.VSA_tile_size,
help="VSA-H3 tile size in tokens (256 or 64); 64 runs the native Triton block-sparse path",
)
# Master port for distributed training/inference
parser.add_argument(
+1 -4
View File
@@ -114,10 +114,7 @@ def _info(logger: Logger,
is_local_main_process = local_rank == 0
if (main_process_only and is_main_process) or (local_main_process_only and is_local_main_process):
# Honor an explicit stacklevel (info_once routes through here with
# stacklevel already set) instead of passing the keyword twice.
stacklevel = kwargs.pop("stacklevel", 2)
logger.log(logging.INFO, msg, *args, stacklevel=stacklevel, **kwargs)
logger.log(logging.INFO, msg, *args, stacklevel=2, **kwargs)
global _warned_local_main_process, _warned_main_process
+27 -137
View File
@@ -10,7 +10,6 @@ import torch
import torch.nn as nn
import torch.nn.functional as F
from fastvideo import envs
from fastvideo.attention import DistributedAttention
from fastvideo.attention.layer import DistributedAttention_VSA
from fastvideo.attention.selector import get_attn_backend
@@ -24,50 +23,12 @@ from fastvideo.layers.linear import ReplicatedLinear
from fastvideo.layers.mlp import MLP
from fastvideo.layers.quantization import QuantizationConfig
from fastvideo.layers.visual_embedding import Timesteps
from fastvideo.logger import init_logger
from fastvideo.models.dits.base import BaseDiT
from fastvideo.models.dits.minimax_h3_fusions import (
HAVE_TRITON,
fused_qknorm_rope,
fused_residual_gate_rmsnorm_modulate,
fused_rmsnorm_modulate,
minimax_h3_swiglu,
)
from fastvideo.platforms import AttentionBackendEnum
from fastvideo.profiler import nvtx_range
from fastvideo.utils import get_compute_dtype
logger = init_logger(__name__)
MINIMAX_H3_MODALITY_NUM = 3
_CFG = MiniMaxH3Config()
_MINIMAX_H3_FUSION_NAMES = frozenset({"modulate", "qknorm_rope", "swiglu"})
def _enabled_minimax_h3_fusions(value: str | None = None) -> frozenset[str]:
"""Parse the independently switchable inference fusion set."""
raw = envs.FASTVIDEO_MINIMAX_H3_FUSIONS if value is None else value
normalized = raw.strip().lower()
if normalized in {"", "0", "none"}:
return frozenset()
if normalized in {"1", "all"}:
return _MINIMAX_H3_FUSION_NAMES
enabled = frozenset(item.strip() for item in normalized.split(",") if item.strip())
unknown = enabled - _MINIMAX_H3_FUSION_NAMES
if unknown:
supported = ",".join(sorted(_MINIMAX_H3_FUSION_NAMES))
raise ValueError(f"Unknown MiniMax H3 fusion(s) {sorted(unknown)}; expected a subset of {supported}.")
return enabled
def _can_run_minimax_h3_fusion(tensor: torch.Tensor) -> bool:
"""Triton kernels are inference-only and stay outside Dynamo capture.
The ``HAVE_TRITON`` check makes the eager fallback exact: on a CUDA build
whose Triton failed to import, an enabled fusion falls back instead of
hitting the strict wrappers' hard RuntimeError mid-forward.
"""
return (HAVE_TRITON and tensor.is_cuda and not torch.is_grad_enabled() and not torch.compiler.is_compiling())
class MiniMaxH3RotaryPosEmbed(nn.Module):
@@ -101,7 +62,6 @@ class MiniMaxH3FeedForward(nn.Module):
ffn_dim: int,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
fuse_swiglu: bool = False,
) -> None:
super().__init__()
self.fc_in = ReplicatedLinear(
@@ -118,15 +78,11 @@ class MiniMaxH3FeedForward(nn.Module):
quant_config=quant_config,
prefix=f"{prefix}.fc_out",
)
self.fuse_swiglu = fuse_swiglu
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
hidden_states, _ = self.fc_in(hidden_states)
if self.fuse_swiglu and _can_run_minimax_h3_fusion(hidden_states):
hidden_states = minimax_h3_swiglu(hidden_states)
else:
hidden_states, gate = hidden_states.chunk(2, dim=-1)
hidden_states = hidden_states * F.silu(gate)
hidden_states, gate = hidden_states.chunk(2, dim=-1)
hidden_states = hidden_states * F.silu(gate)
hidden_states, _ = self.fc_out(hidden_states)
return hidden_states
@@ -143,7 +99,6 @@ class MiniMaxH3Attention(nn.Module):
supported_attention_backends: tuple[AttentionBackendEnum, ...],
quant_config: QuantizationConfig | None,
prefix: str,
fuse_qknorm_rope: bool = False,
) -> None:
super().__init__()
self.num_attention_heads = num_attention_heads
@@ -179,7 +134,6 @@ class MiniMaxH3Attention(nn.Module):
quant_config=quant_config,
prefix=f"{prefix}.to_out",
)
self.fuse_qknorm_rope = fuse_qknorm_rope
# VSA carries a learned gate on its pooled-compression branch. The H3
# checkpoint has no such weight, so the loader zero-initializes it
# (ALLOWED_NEW_PARAM_PATTERNS) and the branch is exactly disabled
@@ -257,18 +211,11 @@ class MiniMaxH3Attention(nn.Module):
query = query.unflatten(-1, (self.num_attention_heads, self.attention_head_dim))
key = key.unflatten(-1, (self.num_attention_heads, self.attention_head_dim))
value = value.unflatten(-1, (self.num_attention_heads, self.attention_head_dim))
if (self.fuse_qknorm_rope and rotary_emb is not None and _can_run_minimax_h3_fusion(query)):
cos, sin = rotary_emb
cos = cos.to(query.dtype)
sin = sin.to(query.dtype)
query = fused_qknorm_rope(query, self.norm_q.weight, cos, sin, self.norm_q.eps)
key = fused_qknorm_rope(key, self.norm_k.weight, cos, sin, self.norm_k.eps)
else:
query = self.norm_q(query)
key = self.norm_k(key)
if rotary_emb is not None:
query = self._apply_rotary_emb(query, rotary_emb)
key = self._apply_rotary_emb(key, rotary_emb)
query = self.norm_q(query)
key = self.norm_k(key)
if rotary_emb is not None:
query = self._apply_rotary_emb(query, rotary_emb)
key = self._apply_rotary_emb(key, rotary_emb)
# H3 rotates only 96/128 channels, which the generic `freqs_cis`
# branch cannot express. Apply it above, then pass no RoPE here.
@@ -450,9 +397,6 @@ class MiniMaxH3TransformerBlock(nn.Module):
quant_config: QuantizationConfig | None,
prefix: str,
adaln_apply_silu: bool = True,
fuse_modulate: bool = False,
fuse_qknorm_rope: bool = False,
fuse_swiglu: bool = False,
) -> None:
super().__init__()
self.norm1 = nn.RMSNorm(hidden_size, eps=norm_eps)
@@ -464,7 +408,6 @@ class MiniMaxH3TransformerBlock(nn.Module):
supported_attention_backends,
quant_config,
prefix=f"{prefix}.attn",
fuse_qknorm_rope=fuse_qknorm_rope,
)
self.norm2 = nn.RMSNorm(hidden_size, eps=norm_eps)
self.ff = MiniMaxH3FeedForward(
@@ -472,7 +415,6 @@ class MiniMaxH3TransformerBlock(nn.Module):
ffn_dim,
quant_config=quant_config,
prefix=f"{prefix}.ff",
fuse_swiglu=fuse_swiglu,
)
self.adaln_proj = MiniMaxH3AdaLayerNormModulation(
time_embed_dim,
@@ -481,7 +423,6 @@ class MiniMaxH3TransformerBlock(nn.Module):
prefix=f"{prefix}.adaln_proj",
apply_silu=adaln_apply_silu,
)
self.fuse_modulate = fuse_modulate
def forward(
self,
@@ -494,39 +435,19 @@ class MiniMaxH3TransformerBlock(nn.Module):
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = (
t.to(hidden_states.dtype) for t in self.adaln_proj(temb))
use_modulate_fusion = self.fuse_modulate and _can_run_minimax_h3_fusion(hidden_states)
if use_modulate_fusion:
norm_hidden_states = fused_rmsnorm_modulate(
hidden_states,
self.norm1.weight,
scale_msa,
shift_msa,
adaln_indices,
self.norm1.eps,
)
else:
norm_hidden_states = self.norm1(hidden_states)
norm_hidden_states = norm_hidden_states * (
1.0 + scale_msa.index_select(0, adaln_indices)) + shift_msa.index_select(0, adaln_indices)
residual = hidden_states
norm_hidden_states = self.norm1(hidden_states)
norm_hidden_states = norm_hidden_states * (
1.0 + scale_msa.index_select(0, adaln_indices)) + shift_msa.index_select(0, adaln_indices)
attention_output = self.attn(norm_hidden_states, rotary_emb, original_seq_len)
if use_modulate_fusion:
hidden_states, norm_hidden_states = fused_residual_gate_rmsnorm_modulate(
hidden_states,
attention_output,
gate_msa,
self.norm2.weight,
scale_mlp,
shift_mlp,
adaln_indices,
self.norm2.eps,
)
else:
hidden_states = hidden_states + gate_msa.index_select(0, adaln_indices) * attention_output
norm_hidden_states = self.norm2(hidden_states)
norm_hidden_states = norm_hidden_states * (
1.0 + scale_mlp.index_select(0, adaln_indices)) + shift_mlp.index_select(0, adaln_indices)
hidden_states = residual + gate_msa.index_select(0, adaln_indices) * attention_output
residual = hidden_states
norm_hidden_states = self.norm2(hidden_states)
norm_hidden_states = norm_hidden_states * (
1.0 + scale_mlp.index_select(0, adaln_indices)) + shift_mlp.index_select(0, adaln_indices)
feed_forward_output = self.ff(norm_hidden_states)
return hidden_states + gate_mlp.index_select(0, adaln_indices) * feed_forward_output
return residual + gate_mlp.index_select(0, adaln_indices) * feed_forward_output
class MiniMaxH3Transformer3DModel(BaseDiT):
@@ -572,17 +493,6 @@ class MiniMaxH3Transformer3DModel(BaseDiT):
def __init__(self, config: MiniMaxH3Config, hf_config: dict[str, Any]) -> None:
super().__init__(config, hf_config)
arch = config.arch_config
self.enabled_fusions = _enabled_minimax_h3_fusions()
if self.enabled_fusions:
if HAVE_TRITON:
logger.info(
"MiniMax H3 inference fusions enabled: %s (CUDA inference-only; grad-enabled and "
"torch.compile-captured forwards fall back to eager).",
",".join(sorted(self.enabled_fusions)))
else:
logger.warning(
"FASTVIDEO_MINIMAX_H3_FUSIONS requested %s but Triton is unavailable; "
"every forward stays on the eager path.", ",".join(sorted(self.enabled_fusions)))
sp_world_size = get_sp_world_size() if model_parallel_is_initialized() else 1
if arch.num_attention_heads % sp_world_size:
raise ValueError(f"MiniMax H3 attention heads ({arch.num_attention_heads}) must be divisible by "
@@ -636,7 +546,7 @@ class MiniMaxH3Transformer3DModel(BaseDiT):
"parameter, but factorized AdaLN weights are pinned to FP16 "
"(BF16 reconstructs them ~1.7x worse). Fine-tune the full-rank "
"checkpoint instead, then re-fit the basis with "
"scripts/checkpoint_conversion/convert_minimax_h3_adaln_rank.py.")
"tools/minimax_h3/fit_adaln_basis.py.")
adaln_dim = self.adaln_rank or arch.time_embed_dim
self.adaln_basis = ReplicatedLinear(
arch.time_embed_dim,
@@ -680,9 +590,6 @@ class MiniMaxH3Transformer3DModel(BaseDiT):
config.quant_config,
prefix=f"{config.prefix}.transformer_blocks.{index}",
adaln_apply_silu=self.adaln_rank is None,
fuse_modulate="modulate" in self.enabled_fusions,
fuse_qknorm_rope="qknorm_rope" in self.enabled_fusions,
fuse_swiglu="swiglu" in self.enabled_fusions,
) for index in range(arch.num_layers)
])
self.norm_out = MiniMaxH3AdaLayerNormOut(
@@ -709,20 +616,6 @@ class MiniMaxH3Transformer3DModel(BaseDiT):
)
self.__post_init__()
def prepare_for_compile(self) -> None:
"""Pipeline hook, called once right before torch.compile wraps the blocks.
Dynamo capture traces the eager branch of every fusion guard, so an
enabled ``FASTVIDEO_MINIMAX_H3_FUSIONS`` set is silently inert inside
compiled block forwards (H3 compiles per-block by default). Say so
once instead of leaving the flag looking active.
"""
if self.enabled_fusions:
logger.warning(
"torch.compile is enabled for MiniMax H3, so the requested inference fusions (%s) are "
"inert inside compiled block forwards; the compiled eager path runs instead.",
",".join(sorted(self.enabled_fusions)))
def materialize_non_persistent_buffers(
self,
device: torch.device,
@@ -841,17 +734,14 @@ class MiniMaxH3Transformer3DModel(BaseDiT):
local_timestep_indices, _ = sequence_model_parallel_shard(local_timestep_indices, dim=0)
rotary_emb = (rotary_cos, rotary_sin)
# The eager driver owns profiling markers while each block's compiled
# forward owns the graph that the marker surrounds.
for block_index, block in enumerate(self.transformer_blocks):
with nvtx_range(f"minimax_h3.transformer_block.{block_index}"):
packed_hidden_states = block(
packed_hidden_states,
temb,
adaln_indices,
rotary_emb,
original_seq_len,
)
for block in self.transformer_blocks:
packed_hidden_states = block(
packed_hidden_states,
temb,
adaln_indices,
rotary_emb,
original_seq_len,
)
packed_hidden_states = self.norm_out(
packed_hidden_states,
@@ -1,20 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Inference-only MiniMax H3 fusions adapted from NVlabs/Sana Sol-Engine.
Source: https://github.com/NVlabs/Sana/tree/sol-engine/models/minimax_h3/GB200
"""
from .modulation import (
fused_residual_gate_rmsnorm_modulate,
fused_rmsnorm_modulate,
)
from .qknorm_rope import HAVE_TRITON, fused_qknorm_rope
from .swiglu import minimax_h3_swiglu
__all__ = [
"HAVE_TRITON",
"fused_qknorm_rope",
"fused_residual_gate_rmsnorm_modulate",
"fused_rmsnorm_modulate",
"minimax_h3_swiglu",
]
@@ -1,302 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""MiniMax H3 RMSNorm and row-indexed modulation fusions."""
from __future__ import annotations
import math
import torch
try:
import triton
import triton.language as tl
except ImportError as exc: # pragma: no cover - depends on the runtime image
triton = None
tl = None
_TRITON_IMPORT_ERROR: ImportError | None = exc
else:
_TRITON_IMPORT_ERROR = None
__all__ = [
"fused_residual_gate_rmsnorm_modulate",
"fused_rmsnorm_modulate",
]
_SUPPORTED_DTYPES = (torch.float16, torch.bfloat16, torch.float32)
_rmsnorm_modulate_kernel = None
_residual_gate_rmsnorm_modulate_kernel = None
if triton is not None:
@triton.jit
def _rmsnorm_modulate_kernel(
out_ptr,
x_ptr,
weight_ptr,
scale_ptr,
shift_ptr,
index_ptr,
n_cols,
n_index,
eps,
stride_x_row,
stride_scale_row,
stride_shift_row,
BLOCK: tl.constexpr,
):
row = tl.program_id(0).to(tl.int64)
cols = tl.arange(0, BLOCK)
mask = cols < n_cols
x_offsets = row * stride_x_row + cols
table_row = tl.load(index_ptr + row % n_index).to(tl.int64)
x = tl.load(x_ptr + x_offsets, mask=mask, other=0.0).to(tl.float32)
variance = tl.sum(x * x, axis=0) / n_cols
normed = x * tl.math.rsqrt(variance + eps)
weight = tl.load(weight_ptr + cols, mask=mask, other=0.0).to(tl.float32)
scale = tl.load(
scale_ptr + table_row * stride_scale_row + cols,
mask=mask,
other=0.0,
).to(tl.float32)
shift = tl.load(
shift_ptr + table_row * stride_shift_row + cols,
mask=mask,
other=0.0,
).to(tl.float32)
output = normed * weight * (1.0 + scale) + shift
tl.store(out_ptr + x_offsets, output.to(out_ptr.dtype.element_ty), mask=mask)
@triton.jit
def _residual_gate_rmsnorm_modulate_kernel(
hidden_out_ptr,
normed_out_ptr,
residual_ptr,
branch_ptr,
gate_ptr,
weight_ptr,
scale_ptr,
shift_ptr,
index_ptr,
n_cols,
n_index,
eps,
stride_input_row,
stride_gate_row,
stride_scale_row,
stride_shift_row,
BLOCK: tl.constexpr,
):
row = tl.program_id(0).to(tl.int64)
cols = tl.arange(0, BLOCK)
mask = cols < n_cols
input_offsets = row * stride_input_row + cols
table_row = tl.load(index_ptr + row % n_index).to(tl.int64)
residual = tl.load(residual_ptr + input_offsets, mask=mask, other=0.0).to(tl.float32)
branch = tl.load(branch_ptr + input_offsets, mask=mask, other=0.0).to(tl.float32)
gate = tl.load(
gate_ptr + table_row * stride_gate_row + cols,
mask=mask,
other=0.0,
).to(tl.float32)
hidden = residual + gate * branch
tl.store(hidden_out_ptr + input_offsets, hidden.to(hidden_out_ptr.dtype.element_ty), mask=mask)
variance = tl.sum(hidden * hidden, axis=0) / n_cols
normed = hidden * tl.math.rsqrt(variance + eps)
weight = tl.load(weight_ptr + cols, mask=mask, other=0.0).to(tl.float32)
scale = tl.load(
scale_ptr + table_row * stride_scale_row + cols,
mask=mask,
other=0.0,
).to(tl.float32)
shift = tl.load(
shift_ptr + table_row * stride_shift_row + cols,
mask=mask,
other=0.0,
).to(tl.float32)
output = normed * weight * (1.0 + scale) + shift
tl.store(normed_out_ptr + input_offsets, output.to(normed_out_ptr.dtype.element_ty), mask=mask)
def _validate_contract(
x: torch.Tensor,
weight: torch.Tensor,
tables: tuple[torch.Tensor, ...],
index: torch.Tensor,
eps: float,
) -> None:
if x.ndim < 2:
raise ValueError(f"x must have shape (..., sequence_length, hidden_size), got {tuple(x.shape)}.")
if x.numel() == 0 or x.shape[-1] == 0:
raise ValueError("x must not be empty.")
if x.dtype not in _SUPPORTED_DTYPES:
raise TypeError(f"x must use float16, bfloat16, or float32, got {x.dtype}.")
hidden_size = x.shape[-1]
sequence_length = x.shape[-2]
if weight.shape != (hidden_size, ):
raise ValueError(f"weight must have shape ({hidden_size},), got {tuple(weight.shape)}.")
if weight.dtype not in _SUPPORTED_DTYPES:
raise TypeError(f"weight must use float16, bfloat16, or float32, got {weight.dtype}.")
if index.ndim != 1 or index.numel() != sequence_length:
raise ValueError(
f"index must have shape ({sequence_length},) so it can wrap over batch rows, got {tuple(index.shape)}."
)
if index.dtype not in (torch.int32, torch.int64):
raise TypeError(f"index must use int32 or int64, got {index.dtype}.")
if not isinstance(eps, (float, int)) or isinstance(eps, bool) or not math.isfinite(eps) or eps <= 0:
raise ValueError(f"eps must be a positive finite number, got {eps!r}.")
table_rows = tables[0].shape[0] if tables and tables[0].ndim == 2 else None
for name, table in zip(("gate", "scale", "shift")[-len(tables):], tables, strict=True):
if table.ndim != 2 or table.shape[1] != hidden_size:
raise ValueError(f"{name} must have shape (table_rows, {hidden_size}), got {tuple(table.shape)}.")
if table.shape[0] == 0 or table.shape[0] != table_rows:
raise ValueError("all modulation tables must have the same non-zero row count.")
if table.dtype not in _SUPPORTED_DTYPES:
raise TypeError(f"{name} must use float16, bfloat16, or float32, got {table.dtype}.")
tensors = (x, weight, *tables, index)
if any(tensor.device != x.device for tensor in tensors[1:]):
raise ValueError("x, weight, modulation tables, and index must be on the same device.")
def _validate_residual_branch(residual: torch.Tensor, branch: torch.Tensor) -> None:
if branch.shape != residual.shape:
raise ValueError(f"branch must match residual shape {tuple(residual.shape)}, got {tuple(branch.shape)}.")
if branch.dtype != residual.dtype:
raise TypeError(f"branch dtype must match residual dtype {residual.dtype}, got {branch.dtype}.")
if branch.device != residual.device:
raise ValueError("branch and residual must be on the same device.")
def _require_triton_cuda(x: torch.Tensor) -> None:
if triton is None:
detail = f": {_TRITON_IMPORT_ERROR}" if _TRITON_IMPORT_ERROR is not None else ""
raise RuntimeError(f"MiniMax H3 modulation fusion requires Triton{detail}.")
if x.device.type != "cuda":
raise RuntimeError(f"MiniMax H3 modulation fusion requires CUDA tensors, got device {x.device}.")
def _require_forward_only(*tensors: torch.Tensor) -> None:
if torch.is_grad_enabled() and any(tensor.requires_grad for tensor in tensors):
raise RuntimeError("MiniMax H3 modulation fusion is forward-only and does not support autograd.")
def _next_power_of_two(value: int) -> int:
return 1 << (value - 1).bit_length()
def _num_warps(block_size: int) -> int:
if block_size >= 8192:
return 16
if block_size >= 2048:
return 8
return 4
def _row_addressable(table: torch.Tensor) -> torch.Tensor:
return table if table.stride(-1) == 1 else table.contiguous()
def fused_rmsnorm_modulate(
x: torch.Tensor,
weight: torch.Tensor,
scale: torch.Tensor,
shift: torch.Tensor,
index: torch.Tensor,
eps: float,
) -> torch.Tensor:
"""Run RMSNorm and row-indexed modulation in one strict Triton kernel.
``index`` values must lie in ``[0, table_rows)``. Unlike eager
``index_select``, the kernel does not raise on out-of-range values (a
device-side bounds check would synchronize); callers are safe by
construction (``timestep_indices * 3 + token_tags``, SP pads with 0).
"""
_validate_contract(x, weight, (scale, shift), index, eps)
_require_forward_only(x, weight, scale, shift)
_require_triton_cuda(x)
hidden_size = x.shape[-1]
flat_x = x.reshape(-1, hidden_size).contiguous()
weight = weight.contiguous()
scale = _row_addressable(scale)
shift = _row_addressable(shift)
index = index.contiguous()
output = torch.empty_like(flat_x)
block_size = _next_power_of_two(hidden_size)
_rmsnorm_modulate_kernel[(flat_x.shape[0], )](
output,
flat_x,
weight,
scale,
shift,
index,
hidden_size,
index.numel(),
eps,
flat_x.stride(0),
scale.stride(0),
shift.stride(0),
BLOCK=block_size,
num_warps=_num_warps(block_size),
)
return output.view_as(x)
def fused_residual_gate_rmsnorm_modulate(
residual: torch.Tensor,
branch: torch.Tensor,
gate: torch.Tensor,
weight: torch.Tensor,
scale: torch.Tensor,
shift: torch.Tensor,
index: torch.Tensor,
eps: float,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Fuse residual update, row-indexed gate, RMSNorm, and modulation.
``index`` values must lie in ``[0, table_rows)``; see
:func:`fused_rmsnorm_modulate` for why the wrapper does not check them.
"""
_validate_residual_branch(residual, branch)
_validate_contract(residual, weight, (gate, scale, shift), index, eps)
_require_forward_only(residual, branch, gate, weight, scale, shift)
_require_triton_cuda(residual)
hidden_size = residual.shape[-1]
flat_residual = residual.reshape(-1, hidden_size).contiguous()
flat_branch = branch.reshape(-1, hidden_size).contiguous()
weight = weight.contiguous()
gate = _row_addressable(gate)
scale = _row_addressable(scale)
shift = _row_addressable(shift)
index = index.contiguous()
hidden = torch.empty_like(flat_residual)
modulated = torch.empty_like(flat_residual)
block_size = _next_power_of_two(hidden_size)
_residual_gate_rmsnorm_modulate_kernel[(flat_residual.shape[0], )](
hidden,
modulated,
flat_residual,
flat_branch,
gate,
weight,
scale,
shift,
index,
hidden_size,
index.numel(),
eps,
flat_residual.stride(0),
gate.stride(0),
scale.stride(0),
shift.stride(0),
BLOCK=block_size,
num_warps=_num_warps(block_size),
)
return hidden.view_as(residual), modulated.view_as(residual)
@@ -1,174 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Fused per-head RMSNorm and partial rotary embedding for MiniMax H3."""
from __future__ import annotations
import math
import torch
try:
import triton
import triton.language as tl
HAVE_TRITON = True
except ImportError: # pragma: no cover - exercised only in environments without Triton
triton = None
tl = None
HAVE_TRITON = False
if HAVE_TRITON:
@triton.jit
def _qknorm_partial_rope_kernel(
out_ptr,
x_ptr,
weight_ptr,
cos_ptr,
sin_ptr,
head_dim,
rotary_dim,
half_rotary_dim,
num_heads,
seq_len,
eps,
BLOCK_SIZE: tl.constexpr,
):
# int64, like the sibling kernels: with int32 program ids,
# ``row * head_dim`` wraps once the flattened input reaches 2**31
# elements (H3's 56 heads x 128 head_dim crosses that at
# batch*seq >= 299_593 tokens per rank) and the loads/stores below
# become out-of-bounds. ``seq_index`` inherits int64 from ``row``.
row = tl.program_id(0).to(tl.int64)
seq_index = (row // num_heads) % seq_len
cols = tl.arange(0, BLOCK_SIZE)
head_mask = cols < head_dim
row_offset = row * head_dim
x = tl.load(x_ptr + row_offset + cols, mask=head_mask, other=0.0).to(tl.float32)
variance = tl.sum(x * x, axis=0) / head_dim
inv_rms = tl.math.rsqrt(variance + eps)
weight = tl.load(weight_ptr + cols, mask=head_mask, other=0.0).to(tl.float32)
normalized = x * inv_rms * weight
rotary_mask = cols < rotary_dim
first_half = cols < half_rotary_dim
partner_col = tl.where(first_half, cols + half_rotary_dim, cols - half_rotary_dim)
partner_x = tl.load(
x_ptr + row_offset + partner_col,
mask=rotary_mask,
other=0.0,
).to(tl.float32)
partner_weight = tl.load(weight_ptr + partner_col, mask=rotary_mask, other=0.0).to(tl.float32)
partner_normalized = partner_x * inv_rms * partner_weight
rotated = tl.where(first_half, -partner_normalized, partner_normalized)
table_offset = seq_index * rotary_dim + cols
cos = tl.load(cos_ptr + table_offset, mask=rotary_mask, other=1.0).to(tl.float32)
sin = tl.load(sin_ptr + table_offset, mask=rotary_mask, other=0.0).to(tl.float32)
rotary_output = normalized * cos + rotated * sin
output = tl.where(rotary_mask, rotary_output, normalized)
tl.store(out_ptr + row_offset + cols, output.to(out_ptr.dtype.element_ty), mask=head_mask)
_SUPPORTED_DTYPES = (torch.float16, torch.bfloat16, torch.float32)
def _validate_inputs(
x: torch.Tensor,
weight: torch.Tensor,
cos: torch.Tensor,
sin: torch.Tensor,
eps: float,
) -> tuple[int, int, int, int, int]:
for name, tensor in (("x", x), ("weight", weight), ("cos", cos), ("sin", sin)):
if not isinstance(tensor, torch.Tensor):
raise TypeError(f"{name} must be a torch.Tensor, got {type(tensor).__name__}")
if x.ndim != 4:
raise ValueError(f"x must have shape (batch, seq, heads, head_dim), got {tuple(x.shape)}")
batch, seq_len, num_heads, head_dim = x.shape
if min(batch, seq_len, num_heads, head_dim) <= 0:
raise ValueError(f"x dimensions must all be positive, got {tuple(x.shape)}")
if weight.shape != (head_dim, ):
raise ValueError(f"weight must have shape ({head_dim},), got {tuple(weight.shape)}")
if cos.ndim != 2:
raise ValueError(f"cos must have shape (seq, rotary_dim), got {tuple(cos.shape)}")
if sin.shape != cos.shape:
raise ValueError(f"sin must match cos shape {tuple(cos.shape)}, got {tuple(sin.shape)}")
if cos.shape[0] != seq_len:
raise ValueError(f"cos/sin sequence length must be {seq_len}, got {cos.shape[0]}")
rotary_dim = cos.shape[1]
if rotary_dim <= 0:
raise ValueError(f"rotary_dim must be positive, got {rotary_dim}")
if rotary_dim > head_dim:
raise ValueError(f"rotary_dim must not exceed head_dim, got rotary_dim={rotary_dim}, head_dim={head_dim}")
if rotary_dim % 2:
raise ValueError(f"rotary_dim must be even, got {rotary_dim}")
if x.dtype not in _SUPPORTED_DTYPES:
raise TypeError(f"x dtype must be float16, bfloat16, or float32, got {x.dtype}")
for name, tensor in (("weight", weight), ("cos", cos), ("sin", sin)):
if tensor.dtype != x.dtype:
raise TypeError(f"{name} dtype must match x dtype {x.dtype}, got {tensor.dtype}")
if tensor.device != x.device:
raise ValueError(f"{name} device must match x device {x.device}, got {tensor.device}")
if not isinstance(eps, (float, int)) or not math.isfinite(float(eps)) or eps <= 0:
raise ValueError(f"eps must be a positive finite number, got {eps!r}")
return batch, seq_len, num_heads, head_dim, rotary_dim
def fused_qknorm_rope(
x: torch.Tensor,
weight: torch.Tensor,
cos: torch.Tensor,
sin: torch.Tensor,
eps: float,
) -> torch.Tensor:
"""Run per-head RMSNorm and partial RoPE in one Sol-Engine-style kernel.
RMSNorm reduction and RoPE arithmetic stay in FP32 registers until the
final store. Triton's reduction order and the absence of eager's BF16
intermediate materializations can produce small, expected rounding drift.
Row offsets are computed in int64, so inputs beyond 2**31 total elements
(about 300k tokens per rank at H3's 56 heads x 128 head_dim) address
correctly.
"""
batch, seq_len, num_heads, head_dim, rotary_dim = _validate_inputs(x, weight, cos, sin, eps)
if not weight.is_contiguous():
raise ValueError("weight must be contiguous")
if not cos.is_contiguous() or not sin.is_contiguous():
raise ValueError("cos and sin must be contiguous (seq, rotary_dim) tables")
if not x.is_cuda:
raise RuntimeError("fused_qknorm_rope requires CUDA tensors")
if not HAVE_TRITON:
raise RuntimeError("fused_qknorm_rope requires Triton")
if torch.is_grad_enabled() and any(tensor.requires_grad for tensor in (x, weight, cos, sin)):
raise RuntimeError("fused_qknorm_rope is inference-only and does not implement autograd")
flat_x = x.reshape(-1, head_dim).contiguous()
flat_out = torch.empty_like(flat_x)
block_size = 1 << (head_dim - 1).bit_length()
_qknorm_partial_rope_kernel[(flat_x.shape[0], )](
flat_out,
flat_x,
weight,
cos,
sin,
head_dim,
rotary_dim,
rotary_dim // 2,
num_heads,
seq_len,
eps,
BLOCK_SIZE=block_size,
num_warps=4,
)
return flat_out.view(batch, seq_len, num_heads, head_dim)
__all__ = ["HAVE_TRITON", "fused_qknorm_rope"]
@@ -1,104 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""MiniMax H3's value-first packed SwiGLU fusion."""
from __future__ import annotations
import torch
try:
import triton
import triton.language as tl
HAVE_TRITON = True
except ImportError: # pragma: no cover - exercised only in environments without Triton
triton = None
tl = None
HAVE_TRITON = False
def _validate_input(x: torch.Tensor) -> int:
if x.ndim == 0:
raise ValueError("MiniMax H3 SwiGLU expects at least one dimension")
packed_width = x.shape[-1]
if packed_width == 0 or packed_width % 2 != 0:
raise ValueError(
"MiniMax H3 SwiGLU expects a positive even last dimension containing packed (value, gate) halves, "
f"got {packed_width}"
)
if not x.is_floating_point():
raise TypeError(f"MiniMax H3 SwiGLU expects a floating-point tensor, got {x.dtype}")
return packed_width // 2
if HAVE_TRITON:
@triton.jit
def _minimax_h3_swiglu_kernel(
out_ptr,
x_ptr,
ffn_dim,
stride_in_row,
stride_out_row,
BLOCK_SIZE: tl.constexpr,
):
row = tl.program_id(0).to(tl.int64)
cols = tl.arange(0, BLOCK_SIZE)
mask = cols < ffn_dim
value = tl.load(x_ptr + row * stride_in_row + cols, mask=mask, other=0.0).to(tl.float32)
gate = tl.load(x_ptr + row * stride_in_row + ffn_dim + cols, mask=mask, other=0.0).to(tl.float32)
# Match Sol-Engine: keep the complete SwiGLU expression in FP32 and
# convert only the final output store.
out = value * (gate * tl.sigmoid(gate))
tl.store(out_ptr + row * stride_out_row + cols, out.to(out_ptr.dtype.element_ty), mask=mask)
else:
_minimax_h3_swiglu_kernel = None
def _num_warps(block_size: int) -> int:
if block_size >= 8192:
return 16
if block_size >= 2048:
return 8
return 4
def minimax_h3_swiglu(x: torch.Tensor) -> torch.Tensor:
"""Run the forward-only Triton fusion over an H3 ``(..., 2 * ffn_dim)`` input.
This is intentionally a strict kernel wrapper: callers own fallback policy and
must only invoke it for a supported CUDA inference path.
"""
ffn_dim = _validate_input(x)
if not x.is_cuda:
raise ValueError("MiniMax H3 fused SwiGLU requires a CUDA tensor")
if x.dtype not in (torch.float16, torch.bfloat16, torch.float32):
raise TypeError(f"MiniMax H3 fused SwiGLU supports float16, bfloat16, and float32, got {x.dtype}")
if torch.is_grad_enabled() and x.requires_grad:
raise RuntimeError("MiniMax H3 fused SwiGLU is forward-only and does not implement autograd")
if _minimax_h3_swiglu_kernel is None:
raise RuntimeError("MiniMax H3 fused SwiGLU requires Triton")
packed_width = x.shape[-1]
flat = x.reshape(-1, packed_width).contiguous()
output_shape = (*x.shape[:-1], ffn_dim)
if flat.shape[0] == 0:
return torch.empty(output_shape, dtype=x.dtype, device=x.device)
out = torch.empty((flat.shape[0], ffn_dim), dtype=x.dtype, device=x.device)
block_size = triton.next_power_of_2(ffn_dim)
_minimax_h3_swiglu_kernel[(flat.shape[0],)](
out,
flat,
ffn_dim,
flat.stride(0),
out.stride(0),
BLOCK_SIZE=block_size,
num_warps=_num_warps(block_size),
)
return out.view(output_shape)
__all__ = ["HAVE_TRITON", "minimax_h3_swiglu"]
+8 -8
View File
@@ -1,7 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
from abc import ABC, abstractmethod
from dataclasses import field
from typing import Any, Generic, TypeVar
import torch
from torch import nn
@@ -9,16 +8,11 @@ from torch import nn
from fastvideo.configs.models.encoders import (BaseEncoderOutput, ImageEncoderConfig, TextEncoderConfig)
from fastvideo.platforms import AttentionBackendEnum
TextEncoderOutputT = TypeVar("TextEncoderOutputT")
class TextEncoder(nn.Module, ABC, Generic[TextEncoderOutputT]):
"""Base for native encoders with a model-specific forward output contract."""
class TextEncoder(nn.Module, ABC):
_fsdp_shard_conditions: list = field(default_factory=lambda: [])
_stacked_params_mapping: list[tuple[str, str, str]] = field(default_factory=list)
_supported_attention_backends: tuple[AttentionBackendEnum, ...] = TextEncoderConfig()._supported_attention_backends
supported_checkpoint_quantization_methods: frozenset[str] = frozenset()
def __init__(self, config: TextEncoderConfig) -> None:
super().__init__()
@@ -29,7 +23,13 @@ class TextEncoder(nn.Module, ABC, Generic[TextEncoderOutputT]):
raise ValueError(f"Subclass {self.__class__.__name__} must define _supported_attention_backends")
@abstractmethod
def forward(self, *args: Any, **kwargs: Any) -> TextEncoderOutputT:
def forward(self,
input_ids: torch.Tensor | None,
position_ids: torch.Tensor | None = None,
attention_mask: torch.Tensor | None = None,
inputs_embeds: torch.Tensor | None = None,
output_hidden_states: bool | None = None,
**kwargs) -> BaseEncoderOutput:
pass
@property
@@ -1,453 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Serialized block-FP8 execution for the MiniMax-H3 Qwen3-VL encoder."""
from typing import Any
import torch
from torch import nn
from torch.nn.parameter import Parameter
try:
import triton
import triton.language as tl
except ImportError:
triton = None
tl = None
from fastvideo.distributed import get_tp_world_size
from fastvideo.layers.linear import LinearBase, LinearMethodBase
from fastvideo.layers.quantization.base_config import QuantizationConfig
from fastvideo.layers.quantization.fp8_config import FP8_DTYPE
from fastvideo.models.utils import set_weight_attrs
class MiniMaxH3SerializedFP8Config(QuantizationConfig):
"""Serialized 128x128 block-FP8 contract for the H3 text encoder."""
def __init__(self, weight_block_size: tuple[int, int]) -> None:
super().__init__()
if weight_block_size != (128, 128):
raise ValueError("MiniMax-H3 serialized FP8 requires weight_block_size=[128, 128], "
f"got {list(weight_block_size)}")
self.weight_block_size = weight_block_size
self.is_checkpoint_fp8_serialized = True
self.activation_scheme = "dynamic"
@classmethod
def get_name(cls) -> str:
return "fp8"
@classmethod
def get_supported_act_dtypes(cls) -> list[torch.dtype]:
return [torch.bfloat16]
@classmethod
def get_min_capability(cls) -> int:
return 100
@staticmethod
def get_config_filenames() -> list[str]:
return []
@classmethod
def from_config(cls, config: dict[str, Any]) -> "MiniMaxH3SerializedFP8Config":
quant_method = str(config.get("quant_method", "")).lower()
if quant_method != "fp8":
raise ValueError(f"MiniMax-H3 only supports serialized FP8 text-encoder checkpoints, got {quant_method!r}")
if str(config.get("activation_scheme", "")).lower() != "dynamic":
raise ValueError("MiniMax-H3 serialized FP8 requires dynamic activation quantization")
if str(config.get("fmt", "e4m3")).lower() not in ("e4m3", "float8_e4m3fn"):
raise ValueError(f"MiniMax-H3 serialized FP8 requires E4M3 weights, got {config.get('fmt')!r}")
block_size = config.get("weight_block_size")
if not isinstance(block_size, list | tuple) or len(block_size) != 2:
raise ValueError("MiniMax-H3 serialized FP8 requires a two-dimensional weight_block_size")
ignored_layers = config.get("modules_to_not_convert", config.get("ignored_layers", []))
if not isinstance(ignored_layers, list | tuple):
raise ValueError("MiniMax-H3 serialized FP8 modules_to_not_convert must be a sequence")
language_exclusions = [
name for name in ignored_layers
if isinstance(name, str) and (name.startswith("language_model.") or ".language_model." in name)
]
if language_exclusions:
raise ValueError("MiniMax-H3 does not support partially quantized language stacks; "
f"ignored language layers: {language_exclusions[:3]}")
if not any(isinstance(name, str) and "visual" in name for name in ignored_layers):
raise ValueError("MiniMax-H3 serialized FP8 requires the vision stack to be listed in "
"modules_to_not_convert")
return cls((int(block_size[0]), int(block_size[1])))
def validate_runtime(self, device: torch.device) -> None:
if device.type != "cuda":
raise RuntimeError("MiniMax-H3 serialized blockwise FP8 requires a CUDA device; "
f"got {device.type!r}")
capability = torch.cuda.get_device_capability(device)
capability_number = capability[0] * 10 + capability[1]
if capability_number < self.get_min_capability():
raise RuntimeError("MiniMax-H3 serialized blockwise FP8 requires GPU capability "
f"sm{self.get_min_capability()} or newer, got sm{capability_number}")
if capability[0] not in (10, 12):
raise RuntimeError("MiniMax-H3 serialized blockwise FP8 currently adapts SGLang's Blackwell "
f"FlashInfer path; got unsupported sm{capability_number}")
_require_sglang_per_token_group_fp8_quantization()
_get_flashinfer_groupwise_fp8_gemm()
def get_quant_method(self, layer: torch.nn.Module, prefix: str):
if isinstance(layer, LinearBase) and ".language_model.layers." in prefix:
return MiniMaxH3SerializedFP8LinearMethod(self.weight_block_size)
return None
# Copyright 2024 SGLang Team
# Licensed under the Apache License, Version 2.0.
# Adapted from SGLang's per-token-group quantization kernels and Blackwell
# FlashInfer dispatch at commit f99c62063c7dcfcd06784b885dc08cb52cf23865:
# https://github.com/sgl-project/sglang/blob/f99c62063c7dcfcd06784b885dc08cb52cf23865/python/sglang/kernels/ops/quantization/fp8_kernel.py
# https://github.com/sgl-project/sglang/blob/f99c62063c7dcfcd06784b885dc08cb52cf23865/python/sglang/srt/layers/quantization/fp8_utils.py
if triton is not None:
@triton.jit
def _h3_per_token_group_quant_fp8_row_major(
input_ptr,
output_ptr,
scale_ptr,
group_size,
eps,
fp8_min,
fp8_max,
BLOCK: tl.constexpr,
):
group_id = tl.program_id(0)
input_ptr += group_id.to(tl.int64) * group_size
output_ptr += group_id.to(tl.int64) * group_size
scale_ptr += group_id
offsets = tl.arange(0, BLOCK)
mask = offsets < group_size
values = tl.load(input_ptr + offsets, mask=mask, other=0.0).to(tl.float32)
absmax = tl.maximum(tl.max(tl.abs(values)), eps)
scale = absmax / fp8_max
quantized = tl.clamp(values / scale, fp8_min, fp8_max).to(output_ptr.dtype.element_ty)
tl.store(output_ptr + offsets, quantized, mask=mask)
tl.store(scale_ptr, scale)
@triton.jit
def _h3_per_token_group_quant_fp8_column_major(
input_ptr,
output_ptr,
scale_ptr,
group_size,
input_columns,
scale_column_stride,
eps,
fp8_min,
fp8_max,
BLOCK: tl.constexpr,
):
group_id = tl.program_id(0)
input_ptr += group_id.to(tl.int64) * group_size
output_ptr += group_id.to(tl.int64) * group_size
groups_per_row = input_columns // group_size
scale_column = group_id % groups_per_row
scale_row = group_id // groups_per_row
scale_ptr += scale_column * scale_column_stride + scale_row
offsets = tl.arange(0, BLOCK)
mask = offsets < group_size
values = tl.load(input_ptr + offsets, mask=mask, other=0.0).to(tl.float32)
absmax = tl.maximum(tl.max(tl.abs(values)), eps)
scale = absmax / fp8_max
quantized = tl.clamp(values / scale, fp8_min, fp8_max).to(output_ptr.dtype.element_ty)
tl.store(output_ptr + offsets, quantized, mask=mask)
tl.store(scale_ptr, scale)
else:
_h3_per_token_group_quant_fp8_row_major = None
_h3_per_token_group_quant_fp8_column_major = None
def _require_sglang_per_token_group_fp8_quantization() -> None:
if (triton is None or _h3_per_token_group_quant_fp8_row_major is None
or _h3_per_token_group_quant_fp8_column_major is None):
raise RuntimeError(
"MiniMax-H3 serialized blockwise FP8 requires Triton for SGLang-compatible "
"per-token-group activation quantization")
def _sglang_per_token_group_quant_fp8(
input_tensor: torch.Tensor,
group_size: int,
*,
column_major_scales: bool,
) -> tuple[torch.Tensor, torch.Tensor]:
"""SGLang-compatible dynamic FP8 quantization for contiguous 2-D activations."""
_require_sglang_per_token_group_fp8_quantization()
if input_tensor.ndim != 2:
raise ValueError(f"per-token-group FP8 quantization expects 2-D input, got {input_tensor.ndim}-D")
if not input_tensor.is_contiguous():
raise ValueError("per-token-group FP8 quantization requires contiguous input")
if input_tensor.shape[-1] % group_size:
raise ValueError(f"activation width {input_tensor.shape[-1]} is not divisible by group_size={group_size}")
quantized = torch.empty_like(input_tensor, dtype=FP8_DTYPE)
rows, columns = input_tensor.shape
groups_per_row = columns // group_size
if column_major_scales:
scales = torch.empty(
(groups_per_row, rows),
device=input_tensor.device,
dtype=torch.float32,
).permute(1, 0)
else:
scales = torch.empty(
(rows, groups_per_row),
device=input_tensor.device,
dtype=torch.float32,
)
if rows:
num_groups = input_tensor.numel() // group_size
block = triton.next_power_of_2(group_size)
num_warps = min(max(block // 256, 1), 8)
if column_major_scales:
_h3_per_token_group_quant_fp8_column_major[(num_groups,)](
input_tensor,
quantized,
scales,
group_size,
columns,
scales.stride(1),
1e-10,
-448.0,
448.0,
BLOCK=block,
num_warps=num_warps,
num_stages=1,
)
else:
_h3_per_token_group_quant_fp8_row_major[(num_groups,)](
input_tensor,
quantized,
scales,
group_size,
1e-10,
-448.0,
448.0,
BLOCK=block,
num_warps=num_warps,
num_stages=1,
)
return quantized, scales
def _get_flashinfer_groupwise_fp8_gemm():
try:
from flashinfer.gemm import gemm_fp8_nt_groupwise
except (AttributeError, ImportError) as error:
raise RuntimeError(
"MiniMax-H3 serialized blockwise FP8 requires "
"flashinfer.gemm.gemm_fp8_nt_groupwise (validated with flashinfer-python==0.6.8). "
"FastVideo will not re-quantize this checkpoint to tensorwise FP8.") from error
return gemm_fp8_nt_groupwise
def _get_flashinfer_groupwise_backend(device: torch.device) -> str:
capability = torch.cuda.get_device_capability(device)
if capability[0] >= 12:
return "cutlass"
if capability[0] == 10:
return "trtllm"
capability_number = capability[0] * 10 + capability[1]
raise RuntimeError(f"FlashInfer groupwise FP8 requires a Blackwell GPU, got sm{capability_number}")
def _flashinfer_gemm_w8a8_block_fp8_linear_with_fallback(
input_tensor: torch.Tensor,
weight: torch.Tensor,
block_size: tuple[int, int],
weight_scale: torch.Tensor,
bias: torch.Tensor | None = None,
) -> torch.Tensor:
input_2d = input_tensor.view(-1, input_tensor.shape[-1])
output_shape = [*input_tensor.shape[:-1], weight.shape[0]]
backend = _get_flashinfer_groupwise_backend(input_tensor.device)
if input_2d.dtype != torch.bfloat16:
raise RuntimeError("MiniMax-H3 FlashInfer groupwise FP8 requires BF16 activations; "
f"got {input_2d.dtype}. The SGLang FP16 Triton GEMM fallback is not enabled for H3.")
if backend == "trtllm" and input_2d.shape[1] < 256:
raise RuntimeError("MiniMax-H3 FlashInfer TRTLLM groupwise FP8 requires K >= 256; "
f"got K={input_2d.shape[1]}. The SGLang Triton GEMM fallback is not enabled for H3.")
gemm_fp8_nt_groupwise = _get_flashinfer_groupwise_fp8_gemm()
block_n, block_k = block_size
q_input, x_scale = _sglang_per_token_group_quant_fp8(
input_2d,
block_k,
column_major_scales=(backend == "trtllm"),
)
if backend == "cutlass":
m, k = input_2d.shape
n = weight.shape[0]
expected_x_scale_shape = (k // block_k, m)
expected_weight_scale_shape = (k // block_k, n // block_n)
if x_scale.shape == (m, k // block_k):
x_scale = x_scale.transpose(-1, -2).contiguous()
if weight_scale.shape == (n // block_n, k // block_k):
weight_scale = weight_scale.transpose(-1, -2).contiguous()
if x_scale.shape != expected_x_scale_shape or weight_scale.shape != expected_weight_scale_shape:
raise RuntimeError("FlashInfer CUTLASS block-FP8 scale layout mismatch: "
f"x_scale={tuple(x_scale.shape)}, weight_scale={tuple(weight_scale.shape)}, "
f"expected={expected_x_scale_shape}/{expected_weight_scale_shape}")
if x_scale.dtype != torch.float32 or weight_scale.dtype != torch.float32:
raise RuntimeError("FlashInfer CUTLASS block-FP8 scales must be float32")
output = gemm_fp8_nt_groupwise(
q_input,
weight,
x_scale.contiguous(),
weight_scale.contiguous(),
out_dtype=input_2d.dtype,
backend="cutlass",
scale_major_mode="MN",
)
else:
expected_x_scale_shape = (input_2d.shape[0], input_2d.shape[1] // block_k)
expected_weight_scale_shape = (weight.shape[0] // block_n, weight.shape[1] // block_k)
if x_scale.shape != expected_x_scale_shape or x_scale.stride(0) != 1:
raise RuntimeError("FlashInfer TRTLLM block-FP8 activation scale layout mismatch: "
f"shape={tuple(x_scale.shape)}, stride={x_scale.stride()}, "
f"expected column-major {expected_x_scale_shape}")
if weight_scale.shape != expected_weight_scale_shape:
raise RuntimeError("FlashInfer TRTLLM block-FP8 weight scale layout mismatch: "
f"shape={tuple(weight_scale.shape)}, expected={expected_weight_scale_shape}")
output = gemm_fp8_nt_groupwise(
q_input,
weight,
x_scale,
weight_scale,
out_dtype=input_2d.dtype,
backend="trtllm",
)
if bias is not None:
output += bias
return output.to(dtype=input_2d.dtype).view(*output_shape)
class MiniMaxH3SerializedFP8LinearMethod(LinearMethodBase):
"""Execute serialized 128x128 block-FP8 weights without re-quantizing them."""
def __init__(self, weight_block_size: tuple[int, int]) -> None:
super().__init__()
self.weight_block_size = weight_block_size
def create_weights(
self,
layer: torch.nn.Module,
input_size_per_partition: int,
output_partition_sizes: list[int],
input_size: int,
output_size: int,
params_dtype: torch.dtype,
**extra_weight_attrs,
) -> None:
output_size_per_partition = sum(output_partition_sizes)
block_n, block_k = self.weight_block_size
tp_size = get_tp_world_size()
if tp_size > 1 and input_size // input_size_per_partition == tp_size:
if input_size_per_partition % block_k:
raise ValueError(f"Weight input_size_per_partition={input_size_per_partition} is not divisible "
f"by block_k={block_k}")
if tp_size > 1 and output_size // output_size_per_partition == tp_size:
for output_partition_size in output_partition_sizes:
if output_partition_size % block_n:
raise ValueError(f"Weight output_partition_size={output_partition_size} is not divisible "
f"by block_n={block_n}")
layer.logical_widths = output_partition_sizes
layer.input_size_per_partition = input_size_per_partition
layer.output_size_per_partition = output_size_per_partition
layer.orig_dtype = params_dtype
weight_loader = extra_weight_attrs.get("weight_loader")
weight = Parameter(
torch.empty(output_size_per_partition, input_size_per_partition, dtype=FP8_DTYPE),
requires_grad=False,
)
set_weight_attrs(weight, {
"input_dim": 1,
"output_dim": 0,
"weight_loader": weight_loader,
})
layer.register_parameter("weight", weight)
scale = Parameter(
torch.empty((output_size_per_partition + block_n - 1) // block_n,
(input_size_per_partition + block_k - 1) // block_k,
dtype=torch.float32),
requires_grad=False,
)
set_weight_attrs(scale, {
"input_dim": 1,
"output_dim": 0,
"weight_loader": weight_loader,
})
scale.data.fill_(torch.finfo(torch.float32).min)
layer.register_parameter("weight_scale_inv", scale)
layer.register_parameter("input_scale", None)
def process_weights_after_loading(self, layer: nn.Module) -> None:
weight = getattr(layer, "weight", None)
block_scales = getattr(layer, "weight_scale_inv", None)
if weight is None or block_scales is None:
raise ValueError("Serialized MiniMax-H3 FP8 linear is missing weight or weight_scale_inv")
if weight.dtype != FP8_DTYPE:
raise ValueError(f"Serialized MiniMax-H3 FP8 weight must be {FP8_DTYPE}, got {weight.dtype}")
if block_scales.dtype != torch.float32:
raise ValueError("Serialized MiniMax-H3 FP8 weight_scale_inv must be float32, "
f"got {block_scales.dtype}")
block_n, block_k = self.weight_block_size
output_size, input_size = weight.shape
if output_size % block_n or input_size % block_k:
raise ValueError("Serialized MiniMax-H3 FP8 weight dimensions must be divisible by the 128x128 block size; "
f"got {tuple(weight.shape)}")
expected_scale_shape = (output_size // block_n, input_size // block_k)
if tuple(block_scales.shape) != expected_scale_shape:
raise ValueError("Serialized MiniMax-H3 FP8 scale shape mismatch: "
f"expected {expected_scale_shape}, got {tuple(block_scales.shape)}")
if not bool(torch.isfinite(block_scales).all()) or bool((block_scales <= 0).any()):
raise ValueError("Serialized MiniMax-H3 FP8 weight_scale_inv must contain finite positive values")
layer.weight.data = weight.data
layer.weight_scale_inv.data = block_scales.data
def apply(
self,
layer: torch.nn.Module,
x: torch.Tensor,
bias: torch.Tensor | None = None,
) -> torch.Tensor:
if x.device.type != "cuda":
raise RuntimeError("MiniMax-H3 serialized blockwise FP8 execution requires CUDA")
capability = torch.cuda.get_device_capability(x.device)
capability_number = capability[0] * 10 + capability[1]
if capability_number < MiniMaxH3SerializedFP8Config.get_min_capability():
raise RuntimeError("MiniMax-H3 serialized blockwise FP8 requires GPU capability "
f"sm{MiniMaxH3SerializedFP8Config.get_min_capability()} or newer, "
f"got sm{capability_number}")
if not x.is_contiguous():
x = x.contiguous()
return _flashinfer_gemm_w8a8_block_fp8_linear_with_fallback(
x,
layer.weight,
self.weight_block_size,
layer.weight_scale_inv,
bias,
)
__all__ = [
"MiniMaxH3SerializedFP8Config",
"MiniMaxH3SerializedFP8LinearMethod",
]
@@ -8,13 +8,13 @@ import torch
import torch.nn.functional as F
from torch import nn
from fastvideo.configs.models.encoders import BaseEncoderOutput
from fastvideo.configs.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConfig
from fastvideo.distributed import get_tp_world_size
from fastvideo.layers.layernorm import RMSNorm
from fastvideo.layers.linear import ColumnParallelLinear, RowParallelLinear
from fastvideo.layers.vocab_parallel_embedding import VocabParallelEmbedding
from fastvideo.models.encoders.base import TextEncoder
from fastvideo.models.encoders.minimax_h3_checkpoint_fp8 import MiniMaxH3SerializedFP8Config
from fastvideo.models.loader.weight_utils import default_weight_loader
@@ -227,15 +227,10 @@ class MiniMaxH3Qwen3VLLanguageModel(nn.Module):
org_num_embeddings=config.vocab_size,
quant_config=quant_config,
)
override = config.num_hidden_layers_override
self.num_layers = (config.num_hidden_layers
if override is None else min(config.num_hidden_layers, override))
self.output_hidden_state_index = config.output_hidden_state_index
self.layers = nn.ModuleList(
MiniMaxH3Qwen3VLTextDecoderLayer(config, prefix=f"{config.prefix}.language_model.layers.{index}")
for index in range(self.num_layers))
self.norm = (RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
if self.num_layers == config.num_hidden_layers else None)
for index in range(config.num_hidden_layers))
self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
self.rotary_emb = MiniMaxH3Qwen3VLTextRotaryEmbedding(config)
def forward(
@@ -243,14 +238,18 @@ class MiniMaxH3Qwen3VLLanguageModel(nn.Module):
inputs_embeds: torch.Tensor,
position_ids: torch.Tensor,
attention_mask: torch.Tensor | None,
output_hidden_states: bool,
visual_pos_masks: torch.Tensor | None,
deepstack_visual_embeds: list[torch.Tensor] | None,
) -> torch.Tensor:
) -> BaseEncoderOutput:
if attention_mask is not None and bool(attention_mask.to(torch.bool).all()):
attention_mask = None
position_embeddings = self.rotary_emb(inputs_embeds, position_ids)
hidden_states = inputs_embeds
all_hidden_states: tuple[torch.Tensor, ...] | None = () if output_hidden_states else None
for layer_index, layer in enumerate(self.layers):
if all_hidden_states is not None:
all_hidden_states += (hidden_states, )
hidden_states = layer(hidden_states, position_embeddings, attention_mask)
if deepstack_visual_embeds is not None and layer_index < len(deepstack_visual_embeds):
if visual_pos_masks is None:
@@ -259,9 +258,10 @@ class MiniMaxH3Qwen3VLLanguageModel(nn.Module):
visual = deepstack_visual_embeds[layer_index].to(hidden_states.device, hidden_states.dtype)
updated = hidden_states[mask].clone() + visual
hidden_states[mask] = updated
if layer_index + 1 == self.output_hidden_state_index:
return hidden_states
raise RuntimeError(f"MiniMax-H3 text stack did not reach hidden_states[{self.output_hidden_state_index}]")
hidden_states = self.norm(hidden_states)
if all_hidden_states is not None:
all_hidden_states += (hidden_states, )
return BaseEncoderOutput(last_hidden_state=hidden_states, hidden_states=all_hidden_states)
class MiniMaxH3Qwen3VLVisionPatchEmbed(nn.Module):
@@ -499,18 +499,10 @@ class MiniMaxH3Qwen3VLVisionModel(nn.Module):
return self.merger(hidden_states), deepstack_features
class MiniMaxH3Qwen3VLConditioner(TextEncoder[torch.Tensor]):
"""H3 conditioner returning the unnormalized layer-50 hidden tensor."""
class MiniMaxH3Qwen3VLConditioner(TextEncoder):
"""FastVideo-native Qwen3-VL body without the unused language-model head."""
supports_hf_from_pretrained = False
supported_checkpoint_quantization_methods = frozenset({"fp8"})
@classmethod
def checkpoint_quantization_config_from_metadata(
cls,
metadata: dict[str, Any],
) -> MiniMaxH3SerializedFP8Config:
return MiniMaxH3SerializedFP8Config.from_config(metadata)
def __init__(self, config: MiniMaxH3Qwen3VLConfig) -> None:
super().__init__(config)
@@ -526,10 +518,6 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder[torch.Tensor]):
def num_hidden_layers(self) -> int:
return self.config.num_hidden_layers
@property
def num_built_hidden_layers(self) -> int:
return self.language_model.num_layers
def _get_rope_index(
self,
input_ids: torch.Tensor,
@@ -622,39 +610,35 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder[torch.Tensor]):
f"tokens={int(mask.sum())}, features={features.shape[0]}")
return mask
# no_grad, NOT inference_mode: with text_encoder_cpu_offload=True (the
# FastVideoArgs default) the loader FSDP2-shards this conditioner, and
# FSDP2's wait_for_unshard reads tensor._version via
# _unsafe_preserve_version_counter - inference tensors do not track
# version counters, so inference_mode crashes the first encode. no_grad
# frees the same activation memory and keeps prompt_embeds ordinary
# tensors (safe for any future backward through the conditioning).
@torch.no_grad()
def encode_ids(
def forward(
self,
input_ids: torch.Tensor,
*,
input_ids: torch.Tensor | None,
position_ids: torch.Tensor | None = None,
attention_mask: torch.Tensor | None = None,
inputs_embeds: torch.Tensor | None = None,
output_hidden_states: bool | None = None,
pixel_values: torch.Tensor | None = None,
image_grid_thw: torch.Tensor | None = None,
pixel_values_videos: torch.Tensor | None = None,
image_grid_thw: torch.Tensor | None = None,
video_grid_thw: torch.Tensor | None = None,
) -> torch.Tensor:
if input_ids.ndim != 1:
raise ValueError(f"MiniMax-H3 slim forward expects 1-D input_ids, got shape={tuple(input_ids.shape)}")
if (pixel_values is None) != (image_grid_thw is None):
raise ValueError("pixel_values and image_grid_thw must be provided together")
if (pixel_values_videos is None) != (video_grid_thw is None):
raise ValueError("pixel_values_videos and video_grid_thw must be provided together")
input_ids = input_ids.unsqueeze(0)
inputs_embeds = self.language_model.embed_tokens(input_ids)
mm_token_type_ids: torch.Tensor | None = None,
**kwargs: Any,
) -> BaseEncoderOutput:
del mm_token_type_ids, kwargs
if (input_ids is None) == (inputs_embeds is None):
raise ValueError("Exactly one of input_ids or inputs_embeds is required")
if inputs_embeds is None:
assert input_ids is not None
inputs_embeds = self.language_model.embed_tokens(input_ids)
if input_ids is None and (pixel_values is not None or pixel_values_videos is not None):
raise ValueError("Multimodal Qwen3-VL inputs require input_ids for placeholder matching")
image_mask = None
video_mask = None
image_deepstack = None
video_deepstack = None
if pixel_values is not None:
if image_grid_thw is None:
if input_ids is None or image_grid_thw is None:
raise ValueError("pixel_values require input_ids and image_grid_thw")
image_features, image_deepstack = self._visual_features(pixel_values, image_grid_thw)
image_features = image_features.to(inputs_embeds.device, inputs_embeds.dtype)
@@ -662,7 +646,7 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder[torch.Tensor]):
"image")
inputs_embeds = inputs_embeds.masked_scatter(image_mask.unsqueeze(-1), image_features)
if pixel_values_videos is not None:
if video_grid_thw is None:
if input_ids is None or video_grid_thw is None:
raise ValueError("pixel_values_videos require input_ids and video_grid_thw")
video_features, video_deepstack = self._visual_features(pixel_values_videos, video_grid_thw)
video_features = video_features.to(inputs_embeds.device, inputs_embeds.dtype)
@@ -690,34 +674,25 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder[torch.Tensor]):
visual_mask = video_mask
deepstack_features = video_deepstack
position_ids = self._get_rope_index(input_ids, image_grid_thw, video_grid_thw, None)
hidden_states = self.language_model(
if position_ids is None:
if input_ids is None:
sequence_length = inputs_embeds.shape[1]
position_ids = torch.arange(sequence_length,
device=inputs_embeds.device).view(1, 1,
-1).expand(3, inputs_embeds.shape[0], -1)
else:
position_ids = self._get_rope_index(input_ids, image_grid_thw, video_grid_thw, attention_mask)
output_hidden_states = self.config.output_hidden_states if output_hidden_states is None else output_hidden_states
outputs = self.language_model(
inputs_embeds,
position_ids,
None,
attention_mask,
output_hidden_states,
visual_mask,
deepstack_features,
)
if hidden_states.ndim != 3 or hidden_states.shape[0] != 1:
raise RuntimeError(f"MiniMax-H3 language model returned unexpected shape={tuple(hidden_states.shape)}")
return hidden_states[0]
def forward(
self,
input_ids: torch.Tensor,
*,
pixel_values: torch.Tensor | None = None,
image_grid_thw: torch.Tensor | None = None,
pixel_values_videos: torch.Tensor | None = None,
video_grid_thw: torch.Tensor | None = None,
) -> torch.Tensor:
return self.encode_ids(
input_ids,
pixel_values=pixel_values,
image_grid_thw=image_grid_thw,
pixel_values_videos=pixel_values_videos,
video_grid_thw=video_grid_thw,
)
outputs.attention_mask = attention_mask
return outputs
def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:
parameters = dict(self.named_parameters())
@@ -727,8 +702,6 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder[torch.Tensor]):
if source_name == "lm_head.weight":
continue
name = source_name[6:] if source_name.startswith("model.") else source_name
if self._is_omitted_checkpoint_key(name):
continue
if name not in parameters:
raise ValueError(f"Unexpected MiniMax-H3 Qwen3-VL checkpoint key: {source_name}")
parameter = parameters[name]
@@ -737,23 +710,7 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder[torch.Tensor]):
loaded.add(name)
return loaded
def _is_omitted_checkpoint_key(self, name: str) -> bool:
"""Return whether a valid checkpoint key belongs to an unbuilt layer."""
language_model = self.language_model
if language_model.norm is not None:
return False
if name == "language_model.norm.weight":
return True
prefix = "language_model.layers."
if not name.startswith(prefix):
return False
index = name[len(prefix):].split(".", 1)[0]
return (index.isdigit() and language_model.num_layers <= int(index) < self.config.num_hidden_layers)
EntryClass = MiniMaxH3Qwen3VLConditioner
__all__ = [
"MiniMaxH3Qwen3VLConditioner",
"MiniMaxH3SerializedFP8Config",
]
__all__ = ["MiniMaxH3Qwen3VLConditioner"]
+15 -61
View File
@@ -9,7 +9,7 @@ from abc import ABC, abstractmethod
from collections.abc import Generator, Iterable
from contextlib import nullcontext
from copy import deepcopy
from typing import Any, cast
from typing import cast
import torch
import torch.distributed as dist
@@ -30,13 +30,9 @@ from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.layers.quantization import get_quantization_config
from fastvideo.logger import init_logger
from fastvideo.models.encoders.base import TextEncoder
from fastvideo.models.hf_transformer_utils import get_diffusers_config
from fastvideo.models.loader.fsdp_load import maybe_load_fsdp_model, shard_model
from fastvideo.models.loader.text_encoder_quantization import (
_configure_text_encoder_quantization,
_process_quantized_text_encoder_weights,
_resolve_text_encoder_checkpoint_path,
)
from fastvideo.models.loader.utils import set_default_torch_dtype
from fastvideo.models.loader.weight_utils import (
filter_duplicate_safetensors_files,
@@ -351,46 +347,22 @@ class TextEncoderLoader(ComponentLoader):
if cpu_offload is None:
cpu_offload = fastvideo_args.text_encoder_cpu_offload
use_cpu_offload = (cpu_offload and len(getattr(model_config, "_fsdp_shard_conditions", [])) > 0)
runtime_device = get_local_torch_device()
from fastvideo.platforms import current_platform
if cpu_offload:
target_device = (torch.device("mps") if current_platform.is_mps() else torch.device("cpu"))
# Set quantization config if specified
if (use_text_encoder_override and fastvideo_args.override_text_encoder_quant is not None):
if fastvideo_args.override_text_encoder_safetensors is None:
raise ValueError("override_text_encoder_quant is set but override_text_encoder_safetensors is None")
quant_cls = get_quantization_config(fastvideo_args.override_text_encoder_quant)
model_config.quant_config = quant_cls()
with set_default_torch_dtype(PRECISION_TO_TYPE[dtype]):
architectures = getattr(model_config, "architectures", [])
model_cls, _ = ModelRegistry.resolve_model_cls(architectures)
checkpoint_path = _resolve_text_encoder_checkpoint_path(
model_path,
fastvideo_args,
use_text_encoder_override,
)
checkpoint_quant_config = _configure_text_encoder_quantization(
model_config,
model_cls,
checkpoint_path,
)
if checkpoint_quant_config is not None:
if fastvideo_args.override_text_encoder_quant is not None:
raise ValueError("Serialized checkpoint quantization is selected from checkpoint metadata; "
"override_text_encoder_quant is an online conversion option and must be unset")
requested_dtype = PRECISION_TO_TYPE[dtype]
if requested_dtype not in checkpoint_quant_config.get_supported_act_dtypes():
raise ValueError(f"Serialized {checkpoint_quant_config.get_name()} text encoder does not support "
f"activation dtype {requested_dtype}")
checkpoint_quant_config.validate_runtime(runtime_device)
logger.info(
"Selected serialized %s text-encoder checkpoint execution from %s",
checkpoint_quant_config.get_name(),
checkpoint_path,
)
elif use_text_encoder_override and fastvideo_args.override_text_encoder_quant is not None:
if fastvideo_args.override_text_encoder_safetensors is None:
raise ValueError("override_text_encoder_quant is set but override_text_encoder_safetensors is None")
quant_cls = get_quantization_config(fastvideo_args.override_text_encoder_quant)
model_config.quant_config = quant_cls()
if getattr(model_cls, "supports_hf_from_pretrained", False):
model = model_cls.from_pretrained_local( # type: ignore[attr-defined]
model_path,
@@ -409,20 +381,11 @@ class TextEncoderLoader(ComponentLoader):
weights_to_load = {name for name, _ in model.named_parameters()}
if (use_text_encoder_override and fastvideo_args.override_text_encoder_safetensors is not None):
if os.path.isdir(checkpoint_path):
override_weights = self._get_all_weights(
model,
checkpoint_path,
to_cpu=bool(cpu_offload),
)
else:
if self.counter_before_loading_weights == 0.0:
self.counter_before_loading_weights = time.perf_counter()
override_weights = safetensors_weights_iterator(
[checkpoint_path],
loaded_weights: set[str] = model.load_weights(
safetensors_weights_iterator(
[fastvideo_args.override_text_encoder_safetensors],
to_cpu=use_cpu_offload,
)
loaded_weights: set[str] = model.load_weights(override_weights) # type: ignore
)) # type: ignore
else:
loaded_weights: set[str] = model.load_weights(
self._get_all_weights(
@@ -437,10 +400,6 @@ class TextEncoderLoader(ComponentLoader):
self.counter_after_loading_weights - self.counter_before_loading_weights,
)
if checkpoint_quant_config is not None:
processed_linears = _process_quantized_text_encoder_weights(model, runtime_device)
logger.info("Validated %d serialized blockwise FP8 text-encoder linears", processed_linears)
# Explicitly move model to target device after loading weights
model = model.to(target_device)
@@ -483,7 +442,7 @@ class TextEncoderLoader(ComponentLoader):
# that have loaded weights tracking currently.
# if loaded_weights is not None:
weights_not_loaded = weights_to_load - loaded_weights
if weights_not_loaded and (model_config.quant_config is None or checkpoint_quant_config is not None):
if weights_not_loaded and model_config.quant_config is None:
raise ValueError("Following weights were not initialized from "
f"checkpoint: {weights_not_loaded}")
@@ -1098,12 +1057,7 @@ class TransformerLoader(ComponentLoader):
# so recording here makes the decision readable from the loaded
# transformer — and records the narrowed one for teacher/critic.
resolved = record_resolved_attention_backend(dit_config)
# Every worker records its resolved backend so distributed profile
# snapshots can prove that all ranks use the requested kernels.
logger.info("Worker %s transformer attention backend: %s",
os.environ.get("RANK", "0"),
resolved.name if resolved else "automatic selection",
local_main_process_only=False)
logger.info("transformer attention backend: %s", resolved.name if resolved else "automatic selection")
model = maybe_load_fsdp_model(
model_cls=model_cls,
init_params={
@@ -1,127 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Checkpoint-serialized quantization lifecycle for native text encoders."""
import json
import os
from itertools import chain
from typing import Any
import torch
import torch.nn as nn
from safetensors.torch import safe_open
from fastvideo.configs.models import EncoderConfig
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.layers.linear import LinearBase, UnquantizedLinearMethod
from fastvideo.layers.quantization.base_config import QuantizationConfig
from fastvideo.models.encoders.base import TextEncoder
def _resolve_text_encoder_checkpoint_path(
model_path: str,
fastvideo_args: FastVideoArgs,
use_text_encoder_override: bool,
) -> str:
override = fastvideo_args.override_text_encoder_safetensors if use_text_encoder_override else None
checkpoint_path = override or model_path
if not os.path.exists(checkpoint_path):
raise FileNotFoundError(f"Text-encoder checkpoint does not exist: {checkpoint_path}")
if not os.path.isdir(checkpoint_path) and not os.path.isfile(checkpoint_path):
raise ValueError(f"Text-encoder checkpoint must be a file or directory: {checkpoint_path}")
return checkpoint_path
def _read_text_encoder_checkpoint_quantization_config(checkpoint_path: str) -> dict[str, Any] | None:
checkpoint_dir = checkpoint_path if os.path.isdir(checkpoint_path) else os.path.dirname(checkpoint_path)
config_path = os.path.join(checkpoint_dir, "config.json")
if os.path.isfile(config_path):
try:
with open(config_path, encoding="utf-8") as config_file:
checkpoint_config = json.load(config_file)
except json.JSONDecodeError as error:
raise ValueError(f"Invalid text-encoder checkpoint config: {config_path}") from error
quantization_config = checkpoint_config.get("quantization_config")
if quantization_config is not None:
if not isinstance(quantization_config, dict):
raise ValueError(f"quantization_config in {config_path} must be an object")
return quantization_config
if not os.path.isfile(checkpoint_path) or not checkpoint_path.endswith(".safetensors"):
return None
with safe_open(checkpoint_path, framework="pt", device="cpu") as checkpoint_file:
metadata = checkpoint_file.metadata() or {}
for key in ("quantization_config", "_quantization_metadata"):
serialized = metadata.get(key)
if serialized is None:
continue
try:
quantization_config = json.loads(serialized)
except json.JSONDecodeError as error:
raise ValueError(f"Invalid {key} metadata in {checkpoint_path}") from error
if not isinstance(quantization_config, dict):
raise ValueError(f"{key} metadata in {checkpoint_path} must decode to an object")
return quantization_config
return None
def _configure_text_encoder_quantization(
model_config: EncoderConfig,
model_cls: type[nn.Module],
checkpoint_path: str,
) -> QuantizationConfig | None:
if not issubclass(model_cls, TextEncoder):
return None
checkpoint_quantization = _read_text_encoder_checkpoint_quantization_config(checkpoint_path)
if checkpoint_quantization is None:
return None
quant_method = str(checkpoint_quantization.get("quant_method", "")).lower()
if not quant_method:
raise ValueError(f"Quantized text-encoder checkpoint {checkpoint_path} does not declare quant_method")
supported_methods = getattr(model_cls, "supported_checkpoint_quantization_methods", frozenset())
if quant_method not in supported_methods:
supported = ", ".join(sorted(supported_methods)) or "none"
raise ValueError(f"Text encoder {model_cls.__name__} does not support serialized {quant_method!r} "
f"checkpoints (supported: {supported})")
factory = getattr(model_cls, "checkpoint_quantization_config_from_metadata", None)
if not callable(factory):
raise ValueError(f"Text encoder {model_cls.__name__} advertises serialized {quant_method!r} support "
"without a checkpoint quantization factory")
quant_config = factory(checkpoint_quantization)
model_config.quant_config = quant_config
return quant_config
def _module_tensor_device(module: nn.Module) -> torch.device | None:
devices = {
tensor.device
for tensor in chain(
module.parameters(recurse=False),
module.buffers(recurse=False),
)
}
if len(devices) > 1:
raise ValueError(f"Quantized text-encoder module {type(module).__name__} spans multiple devices: {devices}")
return next(iter(devices), None)
def _process_quantized_text_encoder_weights(model: nn.Module, process_device: torch.device) -> int:
"""Run quantized post-load hooks one linear at a time on ``process_device``."""
processed = 0
for module in model.modules():
if not isinstance(module, LinearBase) or isinstance(module.quant_method, UnquantizedLinearMethod):
continue
if module.quant_method is None:
continue
original_device = _module_tensor_device(module)
try:
module.to(process_device)
module.quant_method.process_weights_after_loading(module)
finally:
if original_device is not None:
module.to(original_device)
processed += 1
if processed == 0:
raise ValueError("Serialized quantized text-encoder checkpoint selected, but no quantized linear layers exist")
return processed
@@ -392,16 +392,9 @@ class MiniMaxH3AudioBigVGANDecoder(nn.Module):
return torch.clamp(hidden_states, min=-1.0, max=1.0)
def _is_minimax_h3_audio_vae_decoder(name: str, submodule: nn.Module) -> bool:
"""Select the audio decoder that serves the H3 VAE ``decode`` path."""
return name == "decoder" and isinstance(submodule, MiniMaxH3AudioBigVGANDecoder)
class MiniMaxH3AudioVAE(nn.Module):
"""DAC encoder plus BigVGAN decoder for mono 32 kHz waveforms."""
_compile_conditions = [_is_minimax_h3_audio_vae_decoder]
def __init__(self, config: MiniMaxH3AudioVAEConfig):
super().__init__()
self.config = config
@@ -1,378 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Sequence-parallel chunk scheduling for the MiniMax-H3 video VAE.
The H3 video VAE decodes a video as a series of temporal-chunk decoder
forwards whose outputs are joined by a short deterministic frame blend
(``AutoencoderKLMiniMaxH3._decode_chunks``), and encodes videos as fully
independent ``clip_length``-frame encoder forwards. Neither the chunk decode
nor the clip encode has any cross-chunk data dependency — only the *joining*
of decoded chunks (overlap blending, frame trimming) is sequential. This
module round-robins the chunk/clip forwards across the ranks of a
sequence-parallel group and replays the serial joining logic on the
assembling rank, reproducing the serial result bit for bit.
Bit-exactness contract:
- every rank holds an identical copy of the inputs (the H3 DiT all-gathers
its outputs, and reference pixels are prepared identically on all ranks);
- a chunk decoded on any rank is bitwise the tensor the serial loop would
produce (identical weights, inputs, and deterministic kernels on identical
GPUs), and NCCL transports it bitwise;
- every serialization point of the serial algorithm (overlap blending, frame
trimming, pixel denormalization, output-buffer copies, moment
concatenation and token-drop trimming) runs on the assembling rank in
serial order via the same VAE methods the serial path uses.
Collective safety: all group ranks must call these functions together with
identically shaped inputs. Work proceeds in rounds of one collective each;
ranks without a chunk in the final round contribute a placeholder tensor, so
participation is uniform by construction and no rank-dependent branch guards
a collective.
Caveat — compiled decoders (``enable_torch_compile_vae``): inductor autotunes
kernel configs per process at first call, so a compiled decoder is only
deterministic WITHIN a process, not across processes. Chunks decoded on other
ranks then differ from the serial rank's decode of the same chunk exactly as
two serial runs in different processes would (measured on GB200 at 124f:
max 63/255 on <0.5% of pixels, mean ~1e-2/255, first chunk bit-identical).
With the eager decoder — the pipeline default — parallel output is bitwise
equal to serial ``decode_to_pixels``.
"""
from __future__ import annotations
from typing import TYPE_CHECKING
import torch
from fastvideo.models.vaes.minimax_h3_video import (
AutoencoderKLMiniMaxH3,
AutoencoderKLOutput,
DiagonalGaussianDistribution,
)
from fastvideo.profiler import nvtx_range
if TYPE_CHECKING:
from fastvideo.distributed.parallel_state import GroupCoordinator
# Collective used to move decoded chunk segments to the assembling rank.
# "gather" moves each segment once (destination-only); "all_gather" also
# leaves every rank with every segment. Both are exact; the default is the
# faster one measured on GB200 NVL72 (see the PR notes).
DECODE_GATHER_STRATEGIES = ("gather", "all_gather")
DEFAULT_DECODE_GATHER_STRATEGY = "gather"
def parallel_chunk_indices(num_chunks: int, world_size: int, rank_in_group: int) -> list[int]:
"""Round-robin chunk ownership: chunk ``i`` belongs to rank ``i % world_size``."""
if num_chunks < 0:
raise ValueError(f"num_chunks must be non-negative, got {num_chunks}.")
if world_size < 1:
raise ValueError(f"world_size must be positive, got {world_size}.")
if not 0 <= rank_in_group < world_size:
raise ValueError(f"rank_in_group {rank_in_group} out of range for world_size {world_size}.")
return list(range(rank_in_group, num_chunks, world_size))
def _num_rounds(num_chunks: int, world_size: int) -> int:
return -(-num_chunks // world_size)
def _decode_segment(vae: AutoencoderKLMiniMaxH3, z_padded: torch.Tensor, chunk_index: int) -> torch.Tensor:
"""Decode one temporal chunk's clip and keep the frames the join consumes.
The serial loop uses two spans of each decoded clip: the chunk body
``clip[:, :, frame_pre_padding:chunk_num_frames]`` and (when
``token_drop > 0``) the blend tail
``clip[:, :, chunk_num_frames + frame_pre_padding:]``. Everything from
``frame_pre_padding`` on covers both, so one contiguous slice per chunk
travels over the wire. ``.contiguous()`` also detaches the segment from
any decoder-owned storage (e.g. a compiled decoder's reuse pools) before
the next chunk decode can overwrite it.
"""
start = chunk_index * vae.tokens_chunk_size
with nvtx_range(f"minimax_h3.vae.parallel_chunk.{chunk_index}"):
clip = vae._decode_clip(z_padded[:, :, start:start + vae.tokens_chunk_size + vae.token_overlap])
return clip[:, :, vae.frame_pre_padding:].contiguous()
class _ChunkAssembler:
"""Replay the serial chunk-joining semantics of ``_decode_chunks`` +
``_decode_to_pixels`` on gathered chunk segments, in chunk order.
On CUDA the joining kernels and output copies run on a dedicated side
stream: they depend only on already-gathered segments, so running them
off the main stream keeps the assembling rank's next chunk decode (and
therefore every other rank's next collective) off the assembly's tail.
Stream placement cannot change values — the ops and their order are
identical — so bit-exactness with the serial path is unaffected.
"""
def __init__(self, vae: AutoencoderKLMiniMaxH3, output: torch.Tensor, output_num_frames: int,
non_blocking: bool, device: torch.device) -> None:
self._vae = vae
self._output = output
self._output_num_frames = output_num_frames
self._non_blocking = non_blocking
self._body_frames = vae.tokens_chunk_size * vae.temporal_compression_ratio - vae.frame_pre_padding
self._overlap: torch.Tensor | None = None
self._frame_start = 0
self._stream = torch.cuda.Stream(device) if device.type == "cuda" else None
def push(self, segment: torch.Tensor) -> None:
"""Consume the next chunk's segment (``clip[:, :, frame_pre_padding:]``)."""
if self._stream is None:
self._push(segment)
return
# The segment is produced on the current (collective) stream; hand it
# to the assembly stream and pin its storage until assembly reads it.
self._stream.wait_stream(torch.cuda.current_stream(segment.device))
segment.record_stream(self._stream)
with torch.cuda.stream(self._stream):
self._push(segment)
def _push(self, segment: torch.Tensor) -> None:
vae = self._vae
chunk = segment[:, :, :self._body_frames]
if self._overlap is not None:
chunk = vae._blend(self._overlap, chunk, vae.frame_overlap, dim=-3)
num_frames = min(chunk.shape[2], self._output_num_frames - self._frame_start)
chunk = chunk[:, :, :num_frames]
# The tail past the body (and its pre-padding gap) is the next
# chunk's blend overlap — the serial loop's ``next_overlap``.
self._overlap = segment[:, :, self._body_frames + vae.frame_pre_padding:] if vae.config.token_drop > 0 else None
if num_frames > 0:
self._emit(chunk)
def finalize(self) -> None:
"""Emit the final overlap tail exactly as the serial generator does."""
if self._overlap is not None and self._frame_start < self._output_num_frames:
tail = self._overlap[:, :, :self._output_num_frames - self._frame_start]
if self._stream is None:
self._emit(tail)
else:
with torch.cuda.stream(self._stream):
self._emit(tail)
if self._frame_start != self._output.shape[2]:
raise RuntimeError(
f"MiniMax-H3 decode wrote {self._frame_start} frames into an output buffer expecting "
f"{self._output.shape[2]}.")
def synchronize(self) -> None:
"""Drain assembly kernels and output copies before the buffer is read."""
if self._stream is not None:
self._stream.synchronize()
def _emit(self, chunk: torch.Tensor) -> None:
pixels = self._vae.denormalize_pixels(chunk.float()).clamp_(0, 1)
self._vae._copy_chunk_pixels(pixels, self._output, self._frame_start, self._non_blocking)
self._frame_start += pixels.shape[2]
def _broadcast_segment_meta(group: "GroupCoordinator",
segment: torch.Tensor | None) -> tuple[torch.dtype, tuple[int, ...]]:
"""Share the leader's real segment dtype/shape so placeholder tensors match.
The decoder's output dtype depends on the surrounding autocast context;
deriving it on the leader from an actually decoded segment (instead of
predicting it) keeps collective dtypes correct by construction.
"""
meta = (segment.dtype, tuple(segment.shape)) if segment is not None else None
meta = group.broadcast_object(meta, src=0)
if meta is None:
raise RuntimeError("MiniMax-H3 parallel VAE meta broadcast returned no leader metadata.")
return meta
def decode_to_pixels_parallel(
vae: AutoencoderKLMiniMaxH3,
z: torch.Tensor,
output: torch.Tensor | None,
group: "GroupCoordinator",
strategy: str = DEFAULT_DECODE_GATHER_STRATEGY,
) -> torch.Tensor | None:
"""Chunk-parallel ``decode_to_pixels`` across a sequence-parallel group.
All group ranks call this together with identical ``z``. Temporal chunks
are decoded round-robin across the group and their segments move to the
group's first rank, which assembles bitwise the serial
``decode_to_pixels`` result into ``output``. Only the first rank passes
``output`` (validated exactly like the serial API); other ranks pass
``None`` and receive ``None``.
"""
if strategy not in DECODE_GATHER_STRATEGIES:
raise ValueError(f"Unknown parallel-decode strategy {strategy!r}; expected one of {DECODE_GATHER_STRATEGIES}.")
is_leader = group.rank_in_group == 0
if is_leader:
if output is None:
raise ValueError("The first sequence-parallel rank must provide the CPU output buffer.")
expected_shape = vae.decoded_pixel_shape(z.shape)
if output.device.type != "cpu" or output.dtype != torch.float32 or tuple(output.shape) != expected_shape:
raise ValueError(
"`output` must be a CPU float32 tensor with shape "
f"{expected_shape}, got device={output.device}, dtype={output.dtype}, shape={tuple(output.shape)}.")
elif output is not None:
raise ValueError("Only the first sequence-parallel rank may provide an output buffer.")
if group.world_size == 1:
return vae.decode_to_pixels(z, output)
try:
if vae.use_slicing and z.shape[0] > 1:
for batch_index, z_slice in enumerate(z.split(1)):
slice_output = output[batch_index:batch_index + 1] if output is not None else None
_decode_single_parallel(vae, z_slice, slice_output, group, strategy)
else:
_decode_single_parallel(vae, z, output, group, strategy)
finally:
# Drain the leader's async chunk copies before the caller (or an
# exception handler) can read or release the pinned buffer.
if output is not None and vae._streams_chunk_copies(z, output):
torch.cuda.current_stream(z.device).synchronize()
return output
def _decode_single_parallel(
vae: AutoencoderKLMiniMaxH3,
z: torch.Tensor,
output: torch.Tensor | None,
group: "GroupCoordinator",
strategy: str,
) -> None:
pad_tokens, num_chunks, output_num_frames = vae._temporal_decode_plan(z.shape[2])
if pad_tokens > 0:
z = torch.cat([z, z[:, :, -1:].repeat(1, 1, pad_tokens, 1, 1)], dim=2)
world_size = group.world_size
rank = group.rank_in_group
# Every rank decodes its round-0 chunk BEFORE the metadata rendezvous so
# the first decodes run concurrently (a rank that waited on the broadcast
# first would idle a full chunk-decode behind the leader). The leader
# owns chunk 0 under round-robin assignment, so its segment supplies real
# dtype/shape for placeholder rounds instead of guessing autocast state.
first_segment = _decode_segment(vae, z, rank) if rank < num_chunks else None
segment_dtype, segment_shape = _broadcast_segment_meta(group, first_segment if rank == 0 else None)
assembler = None
if output is not None:
non_blocking = vae._streams_chunk_copies(z, output)
assembler = _ChunkAssembler(vae, output, output_num_frames, non_blocking, z.device)
try:
segment_frames = segment_shape[2]
for round_index in range(_num_rounds(num_chunks, world_size)):
chunk_index = round_index * world_size + rank
if chunk_index >= num_chunks:
segment = torch.zeros(segment_shape, dtype=segment_dtype, device=z.device)
elif round_index == 0 and first_segment is not None:
segment = first_segment
else:
segment = _decode_segment(vae, z, chunk_index)
with nvtx_range(f"minimax_h3.vae.parallel_{strategy}.{round_index}"):
if strategy == "gather":
gathered = group.gather(segment, dst=0, dim=2)
else:
gathered = group.all_gather(segment, dim=2)
if assembler is None or gathered is None:
continue
for slot in range(world_size):
if round_index * world_size + slot >= num_chunks:
break
assembler.push(gathered.narrow(2, slot * segment_frames, segment_frames))
if assembler is not None:
assembler.finalize()
finally:
# Drain assembly-stream copies into ``output`` even on the error path
# so an exception cannot leave an in-flight DMA into a buffer the
# caller may release.
if assembler is not None:
assembler.synchronize()
def _encode_clip_moments(vae: AutoencoderKLMiniMaxH3, pixels: torch.Tensor, clip_index: int) -> torch.Tensor:
"""Encode one ``clip_length``-frame clip exactly as ``_encode_pixels`` does."""
clip_length = vae.config.clip_length
frame_start = clip_index * clip_length
with nvtx_range(f"minimax_h3.vae.parallel_encode_clip.{clip_index}"):
clip = pixels[:, :, frame_start:frame_start + clip_length].to(
device=vae.pixel_mean.device,
dtype=torch.float32,
)
if pixels.dtype == torch.uint8:
clip = clip / 255.0
if clip.shape[2] < clip_length:
pad_frames = clip[:, :, -1:].repeat(1, 1, clip_length - clip.shape[2], 1, 1)
clip = torch.cat([clip, pad_frames], dim=2)
clip = vae.normalize_pixels(clip)
return vae._encode_clip(clip).contiguous()
def encode_pixels_parallel(
vae: AutoencoderKLMiniMaxH3,
pixels: torch.Tensor,
group: "GroupCoordinator",
) -> AutoencoderKLOutput:
"""Clip-parallel ``encode_pixels`` across a sequence-parallel group.
Encoder clips have no cross-clip dependency (no overlap, no blending), so
ranks encode disjoint clips and all-gather the per-clip moment tensors.
Every rank returns the identical full posterior — preserving the serial
contract that all ranks hold the same encoded latents — bitwise equal to
``vae.encode_pixels(pixels)``. Moments are latent-sized (a few MB per
clip), so the all-gather is negligible next to the clip forwards.
"""
if pixels.ndim != 5 or pixels.shape[1] != vae.config.in_channels or pixels.shape[2] <= 0:
raise ValueError(
f"`pixels` must have shape [B, {vae.config.in_channels}, T, H, W] with T > 0, "
f"got {tuple(pixels.shape)}.")
if pixels.device.type != "cpu":
raise ValueError(f"`pixels` must remain on CPU, got device={pixels.device}.")
if pixels.dtype != torch.uint8 and not pixels.is_floating_point():
raise TypeError(f"`pixels` must use uint8 or a floating-point dtype, got {pixels.dtype}.")
if group.world_size == 1:
return vae.encode_pixels(pixels)
if vae.use_slicing and pixels.shape[0] > 1:
moments = torch.cat([_encode_single_parallel(vae, pixel_slice, group) for pixel_slice in pixels.split(1)])
else:
moments = _encode_single_parallel(vae, pixels, group)
return AutoencoderKLOutput(latent_dist=DiagonalGaussianDistribution(moments))
def _encode_single_parallel(vae: AutoencoderKLMiniMaxH3, pixels: torch.Tensor,
group: "GroupCoordinator") -> torch.Tensor:
clip_length = vae.config.clip_length
num_clips = -(-pixels.shape[2] // clip_length)
world_size = group.world_size
rank = group.rank_in_group
# Same first-work-then-rendezvous ordering as the decode path: encode the
# round-0 clip before the metadata broadcast so first encodes overlap.
first_moments = _encode_clip_moments(vae, pixels, rank) if rank < num_clips else None
moment_dtype, moment_shape = _broadcast_segment_meta(group, first_moments if rank == 0 else None)
moment_tokens = moment_shape[2]
parts: list[torch.Tensor] = []
for round_index in range(_num_rounds(num_clips, world_size)):
clip_index = round_index * world_size + rank
if clip_index >= num_clips:
moments = torch.zeros(moment_shape, dtype=moment_dtype, device=vae.pixel_mean.device)
elif round_index == 0 and first_moments is not None:
moments = first_moments
else:
moments = _encode_clip_moments(vae, pixels, clip_index)
gathered = group.all_gather(moments, dim=2)
for slot in range(world_size):
if round_index * world_size + slot >= num_clips:
break
parts.append(gathered.narrow(2, slot * moment_tokens, moment_tokens))
encoded = torch.cat(parts, dim=2)
if vae.config.token_drop > 0:
encoded = encoded[:, :, :-vae.config.token_drop]
return encoded
__all__ = [
"DECODE_GATHER_STRATEGIES",
"DEFAULT_DECODE_GATHER_STRATEGY",
"decode_to_pixels_parallel",
"encode_pixels_parallel",
"parallel_chunk_indices",
]
+57 -298
View File
@@ -7,7 +7,6 @@ This module intentionally uses only PyTorch and FastVideo configuration types.
"""
import math
from collections.abc import Iterator
from dataclasses import dataclass
import torch
@@ -15,10 +14,7 @@ import torch.nn as nn
import torch.nn.functional as F
from torch.utils.checkpoint import checkpoint
from fastvideo.attention import get_attn_backend
from fastvideo.configs.models.vaes.minimax_h3_video import MiniMaxH3VideoVAEConfig
from fastvideo.platforms import AttentionBackendEnum
from fastvideo.profiler import nvtx_range
class DiagonalGaussianDistribution:
@@ -295,7 +291,6 @@ class MiniMaxH3VideoRotaryPosEmbed(nn.Module):
class MiniMaxH3VideoAttention(nn.Module):
def __init__(self, dim: int, heads: int, dim_head: int, eps: float = 1e-5, bias: bool = True) -> None:
"""Build projections and the selected dense FastVideo attention implementation."""
super().__init__()
self.heads = heads
self.dim_head = dim_head
@@ -307,34 +302,12 @@ class MiniMaxH3VideoAttention(nn.Module):
self.to_k = nn.Linear(dim, inner_dim, bias=bias)
self.to_v = nn.Linear(dim, inner_dim, bias=bias)
self.to_out = nn.ModuleList([nn.Linear(inner_dim, dim, bias=bias), nn.Dropout(0.0)])
self.attn_impl = None
from fastvideo.platforms import current_platform
if current_platform.is_cuda_alike():
attention_backend = get_attn_backend(
dim_head,
# FlashAttention executes the FP32 VAE activations in BF16 and
# restores FP32 output, so resolve against the kernel dtype.
torch.bfloat16,
supported_attention_backends=(
AttentionBackendEnum.TORCH_SDPA,
AttentionBackendEnum.FLASH_ATTN,
),
)
self.attn_impl = attention_backend.get_impl_cls()(
num_heads=heads,
head_size=dim_head,
softmax_scale=dim_head**-0.5,
num_kv_heads=heads,
causal=False,
)
def forward(
self,
hidden_states: torch.Tensor,
rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None,
) -> torch.Tensor:
"""Apply dense self-attention to one spatial VAE token sequence."""
query = self.to_q(hidden_states).unflatten(2, (self.heads, -1))
key = self.to_k(hidden_states).unflatten(2, (self.heads, -1))
value = self.to_v(hidden_states).unflatten(2, (self.heads, -1))
@@ -355,17 +328,9 @@ class MiniMaxH3VideoAttention(nn.Module):
query = torch.cat([query_rotary * cos + query_rotated * sin, query_pass], dim=-1)
key = torch.cat([key_rotary * cos + key_rotated * sin, key_pass], dim=-1)
if self.attn_impl is not None and query.device.type != "cpu":
# VAE decoding has no diffusion-step metadata, so call the selected
# backend implementation directly with dense BSHD tensors.
hidden_states = self.attn_impl.forward(query, key, value, None)
hidden_states = hidden_states.flatten(2, 3)
else:
# Keep CPU construction and execution available without requiring
# an accelerator attention backend.
query, key, value = (tensor.permute(0, 2, 1, 3) for tensor in (query, key, value))
hidden_states = F.scaled_dot_product_attention(query, key, value)
hidden_states = hidden_states.permute(0, 2, 1, 3).flatten(2, 3)
query, key, value = (tensor.permute(0, 2, 1, 3) for tensor in (query, key, value))
hidden_states = F.scaled_dot_product_attention(query, key, value)
hidden_states = hidden_states.permute(0, 2, 1, 3).flatten(2, 3)
return self.to_out[0](hidden_states)
@@ -468,7 +433,6 @@ class MiniMaxH3VideoViTDecoder3d(nn.Module):
self.gradient_checkpointing = False
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
"""Decode one latent spatial input through the H3 video transformer."""
batch_size, num_channels, num_frames, height, width = hidden_states.shape
hidden_states = hidden_states.permute(0, 2, 3, 4, 1).reshape(
batch_size,
@@ -518,11 +482,6 @@ class MiniMaxH3VideoViTDecoder3d(nn.Module):
)
def _is_minimax_h3_video_vae_decoder(name: str, submodule: nn.Module) -> bool:
"""Select the video decoder that serves the H3 VAE ``decode`` path."""
return name == "decoder" and isinstance(submodule, MiniMaxH3VideoViTDecoder3d)
class AutoencoderKLMiniMaxH3(nn.Module):
"""MiniMax-H3 causal encoder and ViT decoder with exact release geometry."""
@@ -530,7 +489,6 @@ class AutoencoderKLMiniMaxH3(nn.Module):
_no_split_modules = ["MiniMaxH3VideoResnetBlock3d", "MiniMaxH3VideoTransformerBlock"]
_repeated_blocks = ["MiniMaxH3VideoTransformerBlock"]
_keep_in_fp32_modules = ["encoder", "decoder", "quant_conv", "post_quant_conv"]
_compile_conditions = [_is_minimax_h3_video_vae_decoder]
def __init__(self, config: MiniMaxH3VideoVAEConfig) -> None:
super().__init__()
@@ -696,15 +654,12 @@ class AutoencoderKLMiniMaxH3(nn.Module):
slice_rest[dim] = slice(blend_extent, None)
return torch.cat([blended, b[tuple(slice_rest)]], dim=dim)
# The fixed spatial tile grid reuses one compiled blend-and-concatenate graph.
@torch.compile(backend="inductor", mode="reduce-overhead", dynamic=False)
def _stitch_tiles(
self,
tiles: list[list[torch.Tensor]],
height_overlaps: list[int],
width_overlaps: list[int],
) -> torch.Tensor:
"""Blend decoded tile overlaps and concatenate the spatial canvas."""
result_rows = []
for row_index, row in enumerate(tiles):
result_row = []
@@ -721,12 +676,6 @@ class AutoencoderKLMiniMaxH3(nn.Module):
result_rows.append(torch.cat(result_row, dim=-1))
return torch.cat(result_rows, dim=-2)
# Each fixed-shape latent tile reuses one compiled decoder-input projection.
@torch.compile(backend="inductor", mode="reduce-overhead", dynamic=False)
def _project_decoder_tile(self, tile: torch.Tensor) -> torch.Tensor:
"""Project one spatial latent tile into the decoder input channels."""
return self.post_quant_conv(tile)
def _encode_clip(self, x: torch.Tensor) -> torch.Tensor:
if not self.use_tiling:
return self.quant_conv(self.encoder(x))
@@ -750,64 +699,36 @@ class AutoencoderKLMiniMaxH3(nn.Module):
rows.append(row)
latent_y_overlaps = [overlap // self.spatial_compression_ratio for overlap in y_overlaps]
latent_x_overlaps = [overlap // self.spatial_compression_ratio for overlap in x_overlaps]
# Under mode="reduce-overhead" the stitched canvas is a CUDA-graph
# static buffer that the next _stitch_tiles replay overwrites. Callers
# (_encode/_encode_pixels/encode_keyframe) collect per-clip results
# across replays before concatenating, so hand them a caller-owned
# tensor instead of cudagraph-pooled storage.
return self._stitch_tiles(rows, latent_y_overlaps, latent_x_overlaps).clone()
return self._stitch_tiles(rows, latent_y_overlaps, latent_x_overlaps)
def _decode_clip(self, z: torch.Tensor) -> torch.Tensor:
"""Decode one temporal clip, with optional overlapping spatial tiles."""
with nvtx_range("minimax_h3.vae.decode_clip"):
if not self.use_tiling:
with nvtx_range("minimax_h3.vae.decode_clip.no_s_tile.post_quant_conv"):
projected_clip = self.post_quant_conv(z)
with nvtx_range("minimax_h3.vae.decode_clip.no_s_tile.decoder_forward"):
return self.decoder(projected_clip)
height = z.shape[-2] * self.spatial_compression_ratio
width = z.shape[-1] * self.spatial_compression_ratio
with nvtx_range("minimax_h3.vae.decode_clip.split_tiles"):
y_indices, y_lengths, y_overlaps = self._split_tiles(
height,
self.tile_sample_min_height,
self.tile_sample_min_overlap_height,
)
x_indices, x_lengths, x_overlaps = self._split_tiles(
width,
self.tile_sample_min_width,
self.tile_sample_min_overlap_width,
)
ratio = self.spatial_compression_ratio
rows = []
# The eager tile driver owns NVTX so each marker remains outside
# the compiled decoder graph.
with nvtx_range("minimax_h3.vae.decode_clip.decode_tiles"):
for row_index, (y_position, y_length) in enumerate(zip(y_indices, y_lengths)):
row = []
for column_index, (x_position, x_length) in enumerate(zip(x_indices, x_lengths)):
with nvtx_range(f"minimax_h3.vae.decode_clip.tile.{row_index}.{column_index}"):
tile = z[
...,
y_position // ratio:y_position // ratio + y_length // ratio,
x_position // ratio:x_position // ratio + x_length // ratio,
]
projected_tile = self._project_decoder_tile(tile)
with nvtx_range("minimax_h3.vae.decode_clip.tile.decoder_forward"):
decoded_tile = self.decoder(projected_tile)
row.append(decoded_tile)
rows.append(row)
with nvtx_range("minimax_h3.vae.decode_clip.stitch_tiles"):
# Same CUDA-graph output-ownership contract as _encode_clip:
# _decode collects chunks across _stitch_tiles replays before
# torch.cat, so the pooled canvas must not escape this driver.
# (The streaming _decode_to_pixels path copies each chunk out
# before the next decode and never held stale storage; the
# clone keeps that path correct too at one D2D copy per chunk.)
return self._stitch_tiles(rows, y_overlaps, x_overlaps).clone()
if not self.use_tiling:
return self.decoder(self.post_quant_conv(z))
height = z.shape[-2] * self.spatial_compression_ratio
width = z.shape[-1] * self.spatial_compression_ratio
y_indices, y_lengths, y_overlaps = self._split_tiles(
height,
self.tile_sample_min_height,
self.tile_sample_min_overlap_height,
)
x_indices, x_lengths, x_overlaps = self._split_tiles(
width,
self.tile_sample_min_width,
self.tile_sample_min_overlap_width,
)
ratio = self.spatial_compression_ratio
rows = []
for y_position, y_length in zip(y_indices, y_lengths):
row = []
for x_position, x_length in zip(x_indices, x_lengths):
tile = z[
...,
y_position // ratio:y_position // ratio + y_length // ratio,
x_position // ratio:x_position // ratio + x_length // ratio,
]
row.append(self.decoder(self.post_quant_conv(tile)))
rows.append(row)
return self._stitch_tiles(rows, y_overlaps, x_overlaps)
def _encode(self, x: torch.Tensor) -> torch.Tensor:
clip_length = self.config.clip_length
@@ -826,157 +747,43 @@ class AutoencoderKLMiniMaxH3(nn.Module):
moments = moments[:, :, :-self.config.token_drop]
return moments
def _encode_pixels(self, pixels: torch.Tensor) -> torch.Tensor:
"""Encode unnormalized pixels while keeping full videos off the accelerator."""
clip_length = self.config.clip_length
moments = []
for frame_start in range(0, pixels.shape[2], clip_length):
clip = pixels[:, :, frame_start:frame_start + clip_length].to(
device=self.pixel_mean.device,
dtype=torch.float32,
)
if pixels.dtype == torch.uint8:
clip = clip / 255.0
if clip.shape[2] < clip_length:
pad_frames = clip[:, :, -1:].repeat(1, 1, clip_length - clip.shape[2], 1, 1)
clip = torch.cat([clip, pad_frames], dim=2)
clip = self.normalize_pixels(clip)
moments.append(self._encode_clip(clip))
del clip
encoded = torch.cat(moments, dim=2)
if self.config.token_drop > 0:
encoded = encoded[:, :, :-self.config.token_drop]
return encoded
def _temporal_decode_plan(self, latent_num_frames: int) -> tuple[int, int, int]:
"""Return pad tokens, chunk count, and exact decoded frame count."""
if latent_num_frames <= 0:
raise ValueError(f"MiniMax-H3 latent frame count must be positive, got {latent_num_frames}.")
token_drop = self.config.token_drop
def _decode(self, z: torch.Tensor) -> torch.Tensor:
tokens_chunk_size = self.tokens_chunk_size
token_drop = self.config.token_drop
temporal_ratio = self.temporal_compression_ratio
num_tokens = latent_num_frames + token_drop
chunk_num_frames = tokens_chunk_size * temporal_ratio
num_tokens = z.shape[2] + token_drop
pad_tokens = (-num_tokens) % tokens_chunk_size
num_chunks = (num_tokens + pad_tokens) // tokens_chunk_size - int(token_drop > 0)
if num_chunks < 1:
pad_tokens += tokens_chunk_size
num_chunks = 1
decoded_num_frames = num_chunks * (tokens_chunk_size * temporal_ratio - self.frame_pre_padding)
if token_drop > 0:
decoded_num_frames += self.frame_overlap
if pad_tokens > 0:
intra_tail = self.config.clip_length % temporal_ratio
pad_frames = sum(intra_tail if intra_tail and (latent_num_frames + offset) %
tokens_chunk_size == 0 else temporal_ratio for offset in range(pad_tokens))
decoded_num_frames -= pad_frames
if decoded_num_frames <= 0:
raise RuntimeError(
f"MiniMax-H3 decode plan produced {decoded_num_frames} frames for {latent_num_frames} latent "
"frames; the clip_length/token_drop configuration is inconsistent.")
return pad_tokens, num_chunks, decoded_num_frames
def _decode_chunks(self, z: torch.Tensor) -> Iterator[torch.Tensor]:
"""Yield finalized temporal chunks in decode order."""
tokens_chunk_size = self.tokens_chunk_size
chunk_num_frames = tokens_chunk_size * self.temporal_compression_ratio
pad_tokens, num_chunks, output_num_frames = self._temporal_decode_plan(z.shape[2])
if pad_tokens > 0:
z = torch.cat([z, z[:, :, -1:].repeat(1, 1, pad_tokens, 1, 1)], dim=2)
output_frame_start = 0
decoded_chunks = []
overlap = None
for chunk_index in range(num_chunks):
with nvtx_range(f"minimax_h3.vae.temporal_chunk.{chunk_index}"):
start = chunk_index * tokens_chunk_size
clip = self._decode_clip(z[:, :, start:start + tokens_chunk_size + self.token_overlap])
with nvtx_range(f"minimax_h3.vae.temporal_chunk.{chunk_index}.frame_segment.0"):
chunk = clip[:, :, self.frame_pre_padding:chunk_num_frames]
for index in range(num_chunks):
start = index * tokens_chunk_size
clip = self._decode_clip(z[:, :, start:start + tokens_chunk_size + self.token_overlap])
for overlap_index in range(int(token_drop > 0) + 1):
frame_start = overlap_index * chunk_num_frames
chunk = clip[:, :, frame_start:frame_start + chunk_num_frames]
chunk = chunk[:, :, self.frame_pre_padding:]
if overlap_index == 0:
if overlap is not None:
chunk = self._blend(overlap, chunk, self.frame_overlap, dim=-3)
num_frames = min(chunk.shape[2], output_num_frames - output_frame_start)
chunk = chunk[:, :, :num_frames]
decoded_chunks.append(chunk)
else:
overlap = chunk
if overlap is not None:
decoded_chunks.append(overlap)
decoded = torch.cat(decoded_chunks, dim=2)
next_overlap = None
if self.config.token_drop > 0:
with nvtx_range(f"minimax_h3.vae.temporal_chunk.{chunk_index}.frame_segment.1"):
next_overlap = clip[:, :, chunk_num_frames + self.frame_pre_padding:].clone()
# Yield after the ranges close so consumer-side CPU copies do not inflate decoder timing.
overlap = next_overlap
if num_frames > 0:
output_frame_start += num_frames
yield chunk
if overlap is not None and output_frame_start < output_num_frames:
yield overlap[:, :, :output_num_frames - output_frame_start]
def decoded_pixel_shape(self, latent_shape: torch.Size | tuple[int, ...]) -> tuple[int, int, int, int, int]:
"""Return the exact CPU pixel-buffer shape for a latent tensor shape."""
if len(latent_shape) != 5:
raise ValueError(f"MiniMax-H3 latents must be five-dimensional, got shape {tuple(latent_shape)}.")
batch_size, channels, latent_num_frames, latent_height, latent_width = map(int, latent_shape)
if channels != self.latent_channels:
raise ValueError(f"MiniMax-H3 latents must have {self.latent_channels} channels, got {channels}.")
_, _, decoded_num_frames = self._temporal_decode_plan(latent_num_frames)
return (
batch_size,
int(self.config.out_channels),
decoded_num_frames,
latent_height * self.spatial_compression_ratio,
latent_width * self.spatial_compression_ratio,
)
@staticmethod
def _streams_chunk_copies(z: torch.Tensor, output: torch.Tensor) -> bool:
"""Whether finalized chunks copy to ``output`` asynchronously on the current CUDA stream."""
return z.device.type == "cuda" and output.is_pinned()
@staticmethod
def _copy_chunk_pixels(pixels: torch.Tensor, output: torch.Tensor, frame_start: int, non_blocking: bool) -> None:
"""Copy one finalized fp32 pixel chunk into the CPU ``output`` buffer.
Device-to-host copies run per (batch, channel) plane: the temporal
slice of ``output`` is strided across channels, but each plane is
contiguous on both sides, so every transfer stays a direct memcpy
instead of staging through a pageable CPU temporary. With a pinned
``output`` and ``non_blocking=True`` the copies are additionally
asynchronous on the current CUDA stream; callers synchronize once
before releasing the buffer.
"""
target = output[:, :, frame_start:frame_start + pixels.shape[2]]
if pixels.device.type == "cuda":
pixels = pixels.contiguous()
for batch_index in range(pixels.shape[0]):
for channel_index in range(pixels.shape[1]):
target[batch_index, channel_index].copy_(pixels[batch_index, channel_index],
non_blocking=non_blocking)
else:
target.copy_(pixels)
def _decode_to_pixels(self, z: torch.Tensor, output: torch.Tensor) -> None:
"""Decode temporal chunks and immediately copy finalized pixels to CPU.
Each finalized chunk streams through ``_copy_chunk_pixels`` (direct
per-plane memcpys; asynchronous with a pinned ``output``) so the
copies overlap the next chunk's decode; ``decode_to_pixels``
synchronizes once before returning.
"""
non_blocking = self._streams_chunk_copies(z, output)
output_frame_start = 0
for chunk in self._decode_chunks(z):
num_frames = chunk.shape[2]
pixels = self.denormalize_pixels(chunk.float()).clamp_(0, 1)
self._copy_chunk_pixels(pixels, output, output_frame_start, non_blocking)
output_frame_start += num_frames
if output_frame_start != output.shape[2]:
raise RuntimeError(
f"MiniMax-H3 decode wrote {output_frame_start} frames into an output buffer expecting "
f"{output.shape[2]}.")
def _decode(self, z: torch.Tensor) -> torch.Tensor:
return torch.cat(list(self._decode_chunks(z)), dim=2)
if pad_tokens > 0:
intra_tail = self.config.clip_length % temporal_ratio
num_tokens_before_pad = z.shape[2] - pad_tokens
pad_frames = sum(intra_tail if intra_tail and (num_tokens_before_pad + offset) %
tokens_chunk_size == 0 else temporal_ratio for offset in range(pad_tokens))
decoded = decoded[:, :, :-pad_frames]
return decoded
def encode(
self,
@@ -992,34 +799,6 @@ class AutoencoderKLMiniMaxH3(nn.Module):
return (posterior, )
return AutoencoderKLOutput(latent_dist=posterior)
def encode_pixels(
self,
pixels: torch.Tensor,
return_dict: bool = True,
) -> AutoencoderKLOutput | tuple[DiagonalGaussianDistribution]:
"""Encode CPU-resident pixels one VAE clip at a time.
``pixels`` stays on CPU as ``uint8`` in ``[0, 255]`` or floating point
in ``[0, 1]``; each clip is moved to the VAE device, normalized, and
encoded so only one clip of pixels is resident on the accelerator.
"""
if pixels.ndim != 5 or pixels.shape[1] != self.config.in_channels or pixels.shape[2] <= 0:
raise ValueError(
f"`pixels` must have shape [B, {self.config.in_channels}, T, H, W] with T > 0, "
f"got {tuple(pixels.shape)}.")
if pixels.device.type != "cpu":
raise ValueError(f"`pixels` must remain on CPU, got device={pixels.device}.")
if pixels.dtype != torch.uint8 and not pixels.is_floating_point():
raise TypeError(f"`pixels` must use uint8 or a floating-point dtype, got {pixels.dtype}.")
if self.use_slicing and pixels.shape[0] > 1:
moments = torch.cat([self._encode_pixels(pixel_slice) for pixel_slice in pixels.split(1)])
else:
moments = self._encode_pixels(pixels)
posterior = DiagonalGaussianDistribution(moments)
if not return_dict:
return (posterior, )
return AutoencoderKLOutput(latent_dist=posterior)
def encode_keyframe(
self,
x: torch.Tensor,
@@ -1046,26 +825,6 @@ class AutoencoderKLMiniMaxH3(nn.Module):
return (decoded, )
return DecoderOutput(sample=decoded)
def decode_to_pixels(self, z: torch.Tensor, output: torch.Tensor) -> torch.Tensor:
"""Stream decoded ``[0, 1]`` FP32 pixels into a caller-owned CPU buffer."""
expected_shape = self.decoded_pixel_shape(z.shape)
if output.device.type != "cpu" or output.dtype != torch.float32 or tuple(output.shape) != expected_shape:
raise ValueError(
"`output` must be a CPU float32 tensor with shape "
f"{expected_shape}, got device={output.device}, dtype={output.dtype}, shape={tuple(output.shape)}.")
try:
if self.use_slicing and z.shape[0] > 1:
for batch_index, z_slice in enumerate(z.split(1)):
self._decode_to_pixels(z_slice, output[batch_index:batch_index + 1])
else:
self._decode_to_pixels(z, output)
finally:
# Drain async chunk copies before the caller (or an exception
# handler) can read or release the pinned buffer.
if self._streams_chunk_copies(z, output):
torch.cuda.current_stream(z.device).synchronize()
return output
def forward(
self,
sample: torch.Tensor,
@@ -12,9 +12,9 @@ from torch.distributed.tensor import DTensor
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConditioner
from fastvideo.profiler import nvtx_range
from fastvideo.pipelines.basic.minimax_h3.packing import (
MINIMAX_H3_IMAGE_PAD_TOKEN,
MINIMAX_H3_TEXT_ENCODER_LAYER,
MINIMAX_H3_TEXT_TAG,
MINIMAX_H3_VIDEO_PAD_TOKEN,
MINIMAX_H3_VIDEO_TAG,
@@ -42,6 +42,25 @@ def _token_ids(tokenized: Any) -> list[int]:
return [int(token_id) for token_id in input_ids]
def _create_mm_token_type_ids(processor: Any, token_ids: list[int]) -> list[list[int]]:
"""Build Qwen3-VL modality IDs across old and new Transformers releases."""
create_ids = getattr(processor, "create_mm_token_type_ids", None)
if callable(create_ids):
return create_ids([token_ids])
modality_ids = [0] * len(token_ids)
for modality, modality_type in (("image", 1), ("video", 2), ("audio", 3)):
special_ids = getattr(processor, f"{modality}_token_ids", None)
if special_ids is None:
special_id = getattr(processor, f"{modality}_token_id", None)
special_ids = [] if special_id is None else [special_id]
resolved_ids = {int(special_id) for special_id in special_ids if special_id is not None}
for index, token_id in enumerate(token_ids):
if token_id in resolved_ids:
modality_ids[index] = modality_type
return [modality_ids]
def build_ref2va_presentation(
tokenizer: Any,
prompt: str,
@@ -136,10 +155,20 @@ class MiniMaxH3ConditioningStage(PipelineStage):
device: torch.device,
**vision_inputs: torch.Tensor | None,
) -> tuple[torch.Tensor, torch.Tensor]:
input_ids = torch.tensor(token_ids, dtype=torch.long, device=device)
hidden_state_index = MINIMAX_H3_TEXT_ENCODER_LAYER
input_ids = torch.tensor([token_ids], dtype=torch.long, device=device)
mm_token_type_ids = torch.as_tensor(
_create_mm_token_type_ids(self.processor, token_ids),
dtype=torch.long,
device=device,
)
dtype = self.conditioner.dtype
prompt_embeds = self.conditioner(
input_ids,
outputs = self.conditioner(
input_ids=input_ids,
attention_mask=torch.ones_like(input_ids),
mm_token_type_ids=mm_token_type_ids,
use_cache=False,
output_hidden_states=True,
**{
name:
None if value is None else value.to(
@@ -149,10 +178,10 @@ class MiniMaxH3ConditioningStage(PipelineStage):
for name, value in vision_inputs.items()
},
)
if prompt_embeds.ndim != 2 or prompt_embeds.shape[0] != len(token_ids):
raise ValueError(f"MiniMax-H3 slim text encoder returned unexpected shape={tuple(prompt_embeds.shape)}")
if outputs.hidden_states is None or len(outputs.hidden_states) <= hidden_state_index:
raise ValueError(f"Qwen3-VL did not return `hidden_states[{hidden_state_index}]`.")
return (
prompt_embeds.unsqueeze(0).to(device=device, dtype=dtype),
outputs.hidden_states[hidden_state_index].to(device=device, dtype=dtype),
torch.tensor(token_tags, dtype=torch.long),
)
@@ -257,7 +286,6 @@ class MiniMaxH3ConditioningStage(PipelineStage):
@torch.no_grad()
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
"""Encode one H3 prompt presentation and attach its packed text features."""
device = get_local_torch_device()
first_param = next(self.conditioner.parameters(), None)
moved_for_forward = (fastvideo_args.text_encoder_cpu_offload and first_param is not None
@@ -265,13 +293,10 @@ class MiniMaxH3ConditioningStage(PipelineStage):
if moved_for_forward:
self.conditioner.to(device)
try:
# Keep both H3 prompt-presentation modes under one text-encoding
# range so Nsight Systems exposes their complete conditioning cost.
with nvtx_range("minimax_h3.text_encoding"):
if self.ref2va:
prompt_embeds, text_token_tags = self._encode_ref2va(batch, device)
else:
prompt_embeds, text_token_tags = self._encode_fl2va(batch, device)
if self.ref2va:
prompt_embeds, text_token_tags = self._encode_ref2va(batch, device)
else:
prompt_embeds, text_token_tags = self._encode_fl2va(batch, device)
finally:
if moved_for_forward:
self.conditioner.to("cpu")
@@ -7,13 +7,10 @@ from typing import Any
import torch
from fastvideo.distributed import get_local_torch_device, get_sp_group, model_parallel_is_initialized
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.models.vaes.minimax_h3_audio import MiniMaxH3AudioVAE
from fastvideo.models.vaes.minimax_h3_parallel import DEFAULT_DECODE_GATHER_STRATEGY, decode_to_pixels_parallel
from fastvideo.models.vaes.minimax_h3_video import AutoencoderKLMiniMaxH3
from fastvideo.profiler import nvtx_range
from fastvideo.pipelines.basic.minimax_h3.packing import (
MiniMaxH3PackedLayout,
unpack_audio_tokens,
@@ -24,9 +21,6 @@ from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.base import PipelineStage
from fastvideo.pipelines.stages.validators import StageValidators as V
from fastvideo.pipelines.stages.validators import VerificationResult
from fastvideo.utils import is_pin_memory_available
logger = init_logger(__name__)
def _layout(batch: ForwardBatch) -> MiniMaxH3PackedLayout:
@@ -36,23 +30,6 @@ def _layout(batch: ForwardBatch) -> MiniMaxH3PackedLayout:
return layout
def _decode_participation(fastvideo_args: FastVideoArgs, want_parallel: bool) -> tuple[Any, bool, bool]:
"""Resolve (sp_group, is_output_rank, parallel) for the VAE decode stages.
The executors consume rank 0's ForwardBatch and the training validation
callback consumes each sequence-parallel group leader's, so the output
rank is the SP group's first rank (identical to world rank 0 in the
single-group e2e case). ``parallel`` is only true when every group rank
will run the decode body — the collectives inside require uniform
participation, so no rank-dependent branch may guard them.
"""
if not model_parallel_is_initialized():
return None, True, False
sp_group = get_sp_group()
parallel = bool(want_parallel) and sp_group.world_size > 1
return sp_group, sp_group.is_first_rank, parallel
class MiniMaxH3VideoDecodingStage(PipelineStage):
"""Drop visual condition rows, unpatchify, and decode the target video."""
@@ -77,16 +54,6 @@ class MiniMaxH3VideoDecodingStage(PipelineStage):
@torch.no_grad()
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
"""Decode H3 video latents into normalized CPU pixels."""
placeholder = torch.empty((0, 3, 0, 0, 0), device="cpu", dtype=torch.float32)
sp_group, is_output_rank, parallel = _decode_participation(fastvideo_args, fastvideo_args.vae_parallel_decode)
if not is_output_rank and not parallel:
# Consumers read the output rank's ForwardBatch. Keep a
# verifier-compatible placeholder on other ranks and avoid
# duplicating the full VAE decode and CPU output buffer.
batch.output = placeholder
return batch
layout = _layout(batch)
if batch.latents is None or batch.raw_latent_shape is None or len(batch.raw_latent_shape) != 5:
raise ValueError("MiniMax-H3 video latents or raw geometry are missing at decode.")
@@ -104,33 +71,13 @@ class MiniMaxH3VideoDecodingStage(PipelineStage):
try:
latents = self.vae.denormalize_latents(latents.to(device=device, dtype=torch.float32))
if fastvideo_args.output_type == "latent":
# No collectives on this path, so uniform participation is
# trivial: every rank returns here.
batch.output = latents.detach().float().cpu() if is_output_rank else placeholder
batch.output = latents.detach().float().cpu()
return batch
output = None
if is_output_rank:
output = torch.empty(
self.vae.decoded_pixel_shape(latents.shape),
device="cpu",
dtype=torch.float32,
pin_memory=fastvideo_args.pin_cpu_memory and is_pin_memory_available(),
)
# Attribute the streamed decoder computation while retaining
# per-chunk device-to-host transfer and pinned-buffer reuse.
with (
nvtx_range("minimax_h3.vae"),
torch.autocast(device_type=device.type, dtype=torch.float16, enabled=device.type == "cuda"),
):
if parallel:
strategy = fastvideo_args.vae_parallel_decode_strategy or DEFAULT_DECODE_GATHER_STRATEGY
logger.info_once(f"MiniMax-H3 VAE decode: sequence-parallel chunks across "
f"{sp_group.world_size} ranks ({strategy})")
decode_to_pixels_parallel(self.vae, latents, output, sp_group, strategy=strategy)
else:
self.vae.decode_to_pixels(latents, output)
batch.output = output if is_output_rank else placeholder
# The published decode recipe uses FP16 autocast over FP32 weights.
with torch.autocast(device_type=device.type, dtype=torch.float16, enabled=device.type == "cuda"):
video = self.vae.decode(latents).sample
batch.output = self.vae.denormalize_pixels(video.float()).clamp_(0, 1).cpu()
return batch
finally:
if fastvideo_args.vae_cpu_offload:
@@ -160,15 +107,6 @@ class MiniMaxH3AudioDecodingStage(PipelineStage):
@torch.no_grad()
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
"""Decode H3 audio latents into a stereo CPU waveform."""
# Audio decode is sub-second, so it always runs serially on the SP
# group's first rank (the rank whose ForwardBatch consumers read).
if model_parallel_is_initialized() and not get_sp_group().is_first_rank:
batch.extra["audio"] = torch.empty((0, 2), device="cpu", dtype=torch.float32)
batch.extra["audio_sample_rate"] = self.audio_vae.sampling_rate
self._clear_runtime(batch)
return batch
layout = _layout(batch)
if batch.audio_latents is None:
raise ValueError("MiniMax-H3 audio latents are missing at decode.")
@@ -186,10 +124,7 @@ class MiniMaxH3AudioDecodingStage(PipelineStage):
self._clear_runtime(batch)
return batch
# The range isolates waveform synthesis from packing and runtime
# cleanup so the audio decoder has one stable timeline boundary.
with nvtx_range("minimax_h3.audio_vae"):
decoded = self.audio_vae.decode(latents).sample.float()
decoded = self.audio_vae.decode(latents).sample.float()
if decoded.ndim != 3 or decoded.shape[0] != 2 or decoded.shape[1] != 1:
raise ValueError("MiniMax-H3 audio VAE must decode stereo channels as two mono batch items; "
f"got {tuple(decoded.shape)}.")
@@ -11,8 +11,8 @@ from fastvideo.attention.selector import component_attention_backend, get_attn_b
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.profiler import profiler_region
from fastvideo.hooks.activation_trace import trace_step
from fastvideo.profiler import nvtx_range, profiler_region
from fastvideo.pipelines.basic.minimax_h3.packing import (
MINIMAX_H3_KEYFRAME_NOISE_AUG,
MiniMaxH3PackedLayout,
@@ -89,7 +89,6 @@ class MiniMaxH3DenoisingStage(PipelineStage):
@torch.no_grad()
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
"""Denoise the packed H3 video and audio streams over one shared schedule."""
layout = batch.extra.get(MINIMAX_H3_LAYOUT_KEY)
if not isinstance(layout, MiniMaxH3PackedLayout):
raise ValueError("MiniMax-H3 packed layout is missing before denoising.")
@@ -146,15 +145,9 @@ class MiniMaxH3DenoisingStage(PipelineStage):
vsa_exempt = vsa_mode == "exempt"
vsa_dense_layers = tuple(batch.extra.get("vsa_dense_layers", ()))
vsa_dense_first_n = int(batch.extra.get("vsa_dense_first_n_steps", 0))
# Run-level tile geometry (256 default, 64 = native Triton path),
# plumbed like the run-level sparsity; the builder validates the
# value against VSA_H3_TILE_SHAPES.
vsa_tile_size = int(fastvideo_args.VSA_tile_size)
try:
# The stage range groups the complete denoising loop while the
# indexed model ranges retain timing detail for every H3 block.
with profiler_region("inference_denoising"), nvtx_range("minimax_h3.dit"):
with profiler_region("inference_denoising"):
for index, (video_timestep,
audio_timestep) in enumerate(zip(video_timesteps, audio_timesteps, strict=True)):
unique_timesteps, timestep_indices = row_timestep_plan[index]
@@ -174,7 +167,6 @@ class MiniMaxH3DenoisingStage(PipelineStage):
device=device,
exempt=vsa_exempt,
dense_layers=vsa_dense_layers,
tile_size=vsa_tile_size,
)
# Under torch.compile(mode="reduce-overhead") each denoising
# step must be marked, or cudagraph trees flag cross-step
@@ -9,10 +9,8 @@ import numpy as np
import torch
from diffusers.utils.torch_utils import randn_tensor
from fastvideo.distributed import get_local_torch_device, get_sp_group, model_parallel_is_initialized
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.models.vaes.minimax_h3_parallel import encode_pixels_parallel
from fastvideo.pipelines.basic.minimax_h3.packing import (
MINIMAX_H3_AUDIO_CHANNELS,
MINIMAX_H3_KEYFRAME_ENCODE_SEED,
@@ -38,8 +36,6 @@ from fastvideo.pipelines.stages.base import PipelineStage
from fastvideo.pipelines.stages.validators import StageValidators as V
from fastvideo.pipelines.stages.validators import VerificationResult
logger = init_logger(__name__)
MINIMAX_H3_LAYOUT_KEY = "minimax_h3_layout"
@@ -109,20 +105,8 @@ class MiniMaxH3LatentPreparationStage(PipelineStage):
self,
references: list[MiniMaxH3PreparedReference],
device: torch.device,
fastvideo_args: FastVideoArgs,
) -> list[torch.Tensor]:
patch_size = self.transformer.patch_size
# Reference encode runs on every rank (all ranks hold identical
# prepared references), so clip-parallel encode keeps participation
# uniform by construction: each rank encodes a clip subset and the
# all-gather leaves the identical full posterior everywhere.
parallel_group = None
if fastvideo_args.vae_parallel_encode and model_parallel_is_initialized():
sp_group = get_sp_group()
if sp_group.world_size > 1:
parallel_group = sp_group
logger.info_once(f"MiniMax-H3 reference VAE encode: sequence-parallel clips across "
f"{sp_group.world_size} ranks")
rows: list[torch.Tensor] = []
for reference in references:
if reference.media_type == "audio":
@@ -135,11 +119,9 @@ class MiniMaxH3LatentPreparationStage(PipelineStage):
if reference.frames is None:
raise ValueError("MiniMax-H3 reference video frames are missing.")
frames = reference.frames[:trim_reference_num_frames(reference.frames.shape[0])]
pixels = torch.from_numpy(np.ascontiguousarray(frames)).permute(3, 0, 1, 2)[None]
if parallel_group is not None:
posterior = encode_pixels_parallel(self.vae, pixels, parallel_group).latent_dist
else:
posterior = self.vae.encode_pixels(pixels).latent_dist
pixels = torch.from_numpy(frames.copy()).permute(3, 0, 1, 2)[None]
pixels = pixels.to(device=device, dtype=torch.float32).div_(255.0)
posterior = self.vae.encode(self.vae.normalize_pixels(pixels)).latent_dist
latents = self.vae.normalize_latents(_sample_visual_posterior(posterior).to(
torch.float16).float()).cpu()
reference.num_latent_frames = int(latents.shape[2])
@@ -220,7 +202,7 @@ class MiniMaxH3LatentPreparationStage(PipelineStage):
vae_device = get_local_torch_device()
self.vae.to(vae_device)
try:
video_rows = self._encode_visual_rows(references, vae_device, fastvideo_args)
video_rows = self._encode_visual_rows(references, vae_device)
finally:
if fastvideo_args.vae_cpu_offload:
self.vae.to("cpu")
+185 -103
View File
@@ -2,18 +2,17 @@
"""Utilities for managing the PyTorch profiler within FastVideo.
The profiler is shared across the process; this module adds a light-weight
controller that gates collection based on named *regions*. Regions may be
enabled through dedicated environment variables (e.g.
``FASTVIDEO_TORCH_PROFILE_MODEL_LOADING=1``) or via the consolidated
``FASTVIDEO_TORCH_PROFILE_REGIONS`` comma-separated list. Short names work
controller that gates collection based on named *regions*. Regions are enabled
through the ``FASTVIDEO_TORCH_PROFILE_REGIONS`` comma-separated list. Short names work
(``FASTVIDEO_TORCH_PROFILE_REGIONS=model_loading,training_train`` resolves the
``profiler_region_`` prefix automatically).
Typical usage from client code::
controller = TorchProfilerController(profiler, activities)
with controller.region("training_dit"):
controller = get_or_create_profiler("/tmp/fastvideo-traces")
with controller.region("training_train"):
run_training_step()
controller.stop()
To introduce a new region, register it via :func:`register_profiler_region`
and wrap the corresponding code in :meth:`TorchProfilerController.region`.
@@ -22,11 +21,11 @@ and wrap the corresponding code in :meth:`TorchProfilerController.region`.
from __future__ import annotations
import contextlib
import functools
import os
from collections.abc import Callable, Iterable
from dataclasses import dataclass
from typing import Any
from collections.abc import Callable
import functools
from collections.abc import Iterable
import torch
@@ -38,25 +37,6 @@ logger = init_logger(__name__)
_GLOBAL_CONTROLLER: TorchProfilerController | None = None
@contextlib.contextmanager
def nvtx_range(name: str):
"""Emit one optional NVTX range for an external CUDA profiler.
``FASTVIDEO_NVTX_PROFILE=1`` enables the marker. The context manager stays
a no-op without CUDA so call sites can remain shared with CPU tests.
"""
enabled = envs.FASTVIDEO_NVTX_PROFILE and torch.cuda.is_available()
if not enabled:
yield
return
torch.cuda.nvtx.range_push(name)
try:
yield
finally:
torch.cuda.nvtx.range_pop()
@dataclass(frozen=True)
class ProfilerRegion:
"""Metadata describing a profiler region."""
@@ -161,6 +141,26 @@ register_profiler_region(
name="profiler_region_training_train_one_step",
description="Single optimizer step including forward/backward passes.",
)
register_profiler_region(
name="profiler_region_training_forward",
description="Training method forward pass and loss computation.",
)
register_profiler_region(
name="profiler_region_training_dataloader",
description="Fetch the next training batch in the trainer process.",
)
register_profiler_region(
name="profiler_region_training_backward",
description="Training backward pass.",
)
register_profiler_region(
name="profiler_region_training_optimizer",
description="Gradient clipping, optimizer/scheduler steps, and zero_grad.",
)
register_profiler_region(
name="profiler_region_training_callbacks",
description="End-of-step training callbacks such as EMA updates.",
)
register_profiler_region(
name="profiler_region_training_train",
description="High-level step orchestration in the training loop.",
@@ -184,6 +184,21 @@ register_profiler_region(
description="Parameter updates specific to distillation workflows.",
)
# DMD2 method regions. These sit inside ``training_forward`` and make the
# method's multi-model forward path distinguishable in a single trace.
register_profiler_region(
name="profiler_region_dmd2_student_rollout",
description="DMD2 student rollout, including its simulated prefix steps.",
)
register_profiler_region(
name="profiler_region_dmd2_generator_loss",
description="DMD2 generator loss, including teacher and critic scoring.",
)
register_profiler_region(
name="profiler_region_dmd2_critic_loss",
description="DMD2 critic flow-matching loss, including its student rollout.",
)
def get_or_create_profiler(trace_dir: str | None) -> TorchProfilerController:
"""Create or reuse the process-wide torch profiler controller."""
@@ -208,25 +223,28 @@ def get_or_create_profiler(trace_dir: str | None) -> TorchProfilerController:
)
logger.info("FASTVIDEO_TORCH_PROFILE_REGIONS=%s", envs.FASTVIDEO_TORCH_PROFILE_REGIONS)
profiler = torch.profiler.profile(
activities=_DEFAULT_ACTIVITIES,
record_shapes=envs.FASTVIDEO_TORCH_PROFILER_RECORD_SHAPES,
profile_memory=envs.FASTVIDEO_TORCH_PROFILER_WITH_PROFILE_MEMORY,
with_stack=envs.FASTVIDEO_TORCH_PROFILER_WITH_STACK,
with_flops=envs.FASTVIDEO_TORCH_PROFILER_WITH_FLOPS,
# No schedule: nothing in the codebase calls profiler.step(), so a
# wait/warmup schedule never advances and the profiler records nothing.
# Region toggling gates collection; the single trace exports at stop().
on_trace_ready=torch.profiler.tensorboard_trace_handler(trace_dir, use_gzip=True),
def profiler_factory() -> Any:
return torch.profiler.profile(
activities=_DEFAULT_ACTIVITIES,
record_shapes=envs.FASTVIDEO_TORCH_PROFILER_RECORD_SHAPES,
profile_memory=envs.FASTVIDEO_TORCH_PROFILER_WITH_PROFILE_MEMORY,
with_stack=envs.FASTVIDEO_TORCH_PROFILER_WITH_STACK,
with_flops=envs.FASTVIDEO_TORCH_PROFILER_WITH_FLOPS,
on_trace_ready=torch.profiler.tensorboard_trace_handler(trace_dir, use_gzip=True),
)
controller = TorchProfilerController(
None,
_DEFAULT_ACTIVITIES,
profiler_factory=profiler_factory,
trace_dir=trace_dir,
)
controller = TorchProfilerController(profiler, _DEFAULT_ACTIVITIES)
controller._trace_dir = trace_dir
controller.start()
# The trace only exports at stop(); inference paths have no shutdown hook
# that calls it, so register one. stop() is idempotent.
# Region exit normally exports each trace segment. Keep an atexit hook for
# exceptions or process shutdown while a region is still active.
import atexit
atexit.register(controller.stop)
logger.info("Torch profiler started")
logger.info("Torch profiler armed; collection starts at the first enabled region")
return controller
@@ -279,7 +297,14 @@ class TorchProfilerConfig:
class TorchProfilerController:
"""Helper that toggles torch profiler collection for named regions.
"""Create complete torch-profiler trace segments for named regions.
PyTorch's dynamic CUDA collection toggle can fail to re-enable CUPTI on
some supported stacks. In that failure mode it emits CPU operators while
silently dropping every CUDA kernel. This controller therefore starts a
fresh profiler at each outermost enabled region and stops it when that
region exits. Nested enabled regions become annotations in the same
CPU/CUDA trace segment.
Parameters
----------
@@ -292,12 +317,17 @@ class TorchProfilerController:
config:
Optional :class:`TorchProfilerConfig`. If omitted, :meth:`from_env`
constructs one during initialization.
profiler_factory:
Factory for fresh profiler instances. Required to profile more than
one outermost region invocation.
trace_dir:
Directory for per-segment summaries.
Examples
--------
Enabling an existing region from the command line::
FASTVIDEO_TORCH_PROFILE_REGIONS=model_loading,training_dit \
FASTVIDEO_TORCH_PROFILE_REGIONS=model_loading,training_train \
python fastvideo/training/wan_training_pipeline.py ...
Wrapping a code block in a registered region::
@@ -318,24 +348,29 @@ class TorchProfilerController:
activities: Iterable[torch.profiler.ProfilerActivity],
config: TorchProfilerConfig | None = None,
disabled: bool = False,
profiler_factory: Callable[[], Any] | None = None,
trace_dir: str | None = None,
) -> None:
activities_tuple = tuple(activities)
existing = get_global_controller()
if existing is not None and not disabled:
raise RuntimeError("TorchProfilerController already initialized globally. Use get_global_controller().")
self._profiler = profiler
self._profiler_factory = profiler_factory
self._activities = activities_tuple
self._trace_dir = trace_dir
self._segment_index = 0
self._segment_region: str | None = None
self._active_region_depth = 0
self._collection_enabled = False
if disabled:
self._profiler = None
self._configured = False
self._armed = False
return
self._profiler = profiler
self._activities = activities_tuple
self._config = config or TorchProfilerConfig.from_env()
# torch.profiler collects from start(); reflect that so the initial
# _set_collection(False) in start() actually toggles it off instead of
# short-circuiting (which captured everything before the first region).
self._collection_enabled = True
self._active_region_depth = 0
self._trace_dir: str | None = None
self._configured = True
self._armed = False
logger.info("PROFILER: TorchProfilerController initialized with config: %s", self._config)
set_global_controller(self)
@@ -343,29 +378,61 @@ class TorchProfilerController:
def is_enabled(self) -> bool:
"""Return ``True`` when the underlying profiler is collecting."""
if self._profiler is None:
return False
return self._collection_enabled
def is_region_enabled(self, region: str) -> bool:
"""Return ``True`` if ``region`` should be collected."""
if self._profiler is None:
if not self.has_profiler:
return False
resolved = resolve_profiler_region(region)
if resolved is None:
return False
return self._config.regions.get(resolved.name, False)
def _set_collection(self, enabled: bool) -> None:
if self._profiler is None:
def _new_profiler(self) -> Any:
if self._profiler is not None:
profiler = self._profiler
self._profiler = None
return profiler
if self._profiler_factory is None:
raise RuntimeError("Torch profiler cannot start another trace segment without a profiler_factory")
return self._profiler_factory()
def _start_segment(self, region: str) -> None:
if self._collection_enabled:
return
if self._collection_enabled == enabled:
self._profiler = self._new_profiler()
logger.info(
"PROFILER: Starting segment %d for region %s",
self._segment_index,
region,
)
self._profiler.start()
self._segment_region = region
self._collection_enabled = True
def _finish_segment(self) -> None:
if self._profiler is None or not self._collection_enabled:
return
event = ("fastvideo.profiler.enable_collection" if enabled else "fastvideo.profiler.disable_collection")
with torch.profiler.record_function(event):
self._profiler.toggle_collection_dynamic(enabled, self._activities)
self._collection_enabled = enabled
profiler = self._profiler
segment_index = self._segment_index
segment_region = self._segment_region or "unknown"
logger.info(
"PROFILER: Stopping segment %d for region %s",
segment_index,
segment_region,
)
profiler.stop()
self._write_summary(
profiler,
segment_index=segment_index,
segment_region=segment_region,
)
self._profiler = None
self._collection_enabled = False
self._segment_region = None
self._segment_index += 1
_warned_unregistered: set[str] = set()
@@ -373,7 +440,7 @@ class TorchProfilerController:
def region(self, region: str):
"""Context manager that enables profiling for ``region`` if configured."""
if self._profiler is None:
if not self.has_profiler:
yield
return
@@ -390,64 +457,70 @@ class TorchProfilerController:
yield
return
# NVTX range so the same region names are visible in nsys timelines
if self._active_region_depth == 0:
self._start_segment(region)
# NVTX range so the same region names are visible in nsys timelines.
# Push after profiler startup so Kineto also records the annotation.
nvtx = torch.cuda.is_available()
if nvtx:
torch.cuda.nvtx.range_push(f"fastvideo.region::{region}")
self._active_region_depth += 1
if self._active_region_depth == 1:
logger.info("PROFILER: Setting collection to True (depth=%s) for region %s", self._active_region_depth,
region)
self._set_collection(True)
try:
# record_function opens after collection is enabled so the region
# marker itself lands in the trace.
with torch.profiler.record_function(f"fastvideo.region::{region}"):
yield
finally:
self._active_region_depth -= 1
logger.info("PROFILER: Decreasing active region depth to %s", self._active_region_depth)
if self._active_region_depth == 0:
logger.info("PROFILER: Setting collection to False upon exiting region %s", region)
self._set_collection(False)
if nvtx:
torch.cuda.nvtx.range_pop()
try:
if nvtx:
torch.cuda.nvtx.range_pop()
finally:
# Close NVTX before stopping Kineto so both profilers see a
# balanced outermost range in the exported segment.
if self._active_region_depth == 0:
self._finish_segment()
def start(self) -> None:
"""Start the profiler and pause collection until a region is entered."""
"""Arm the controller; collection begins at an enabled region."""
logger.info("PROFILER: Starting profiler...")
if self._profiler is None:
if not self._configured:
return
self._profiler.start()
logger.info("PROFILER: Profiler started")
# Profiler starts with collection disabled by default.
logger.info("PROFILER: Setting collection to False")
self._set_collection(False)
logger.info("PROFILER: Profiler started with collection disabled")
self._armed = True
logger.info("PROFILER: Controller armed")
def _write_summary(self) -> None:
def _write_summary(
self,
profiler: Any,
*,
segment_index: int,
segment_region: str,
) -> None:
"""Compact per-rank op summary next to the trace: a key_averages table
and a JSON with input shapes, so operator-split analysis does not
require parsing multi-GB chrome traces."""
if self._profiler is None or not self._trace_dir:
if not self._trace_dir:
return
try:
import json as _json
import os as _os
rank = _os.environ.get("RANK", "0")
averages = self._profiler.key_averages(group_by_input_shape=True)
stem = _os.path.join(self._trace_dir, f"summary_rank{rank}")
with open(f"{stem}.txt", "w") as fh:
rank = os.environ.get("RANK", "0")
short_region = segment_region.removeprefix("profiler_region_")
stem = os.path.join(
self._trace_dir,
f"summary_rank{rank}_segment{segment_index:04d}_{short_region}",
)
averages = profiler.key_averages(group_by_input_shape=envs.FASTVIDEO_TORCH_PROFILER_RECORD_SHAPES, )
with open(f"{stem}.txt", "w", encoding="utf-8") as fh:
fh.write(averages.table(sort_by="self_device_time_total", row_limit=60))
rows = [{
"name": e.key,
"shapes": str(e.input_shapes),
"self_cpu_us": e.self_cpu_time_total,
"cpu_us": e.cpu_time_total,
"self_device_us": e.self_device_time_total,
"device_us": e.device_time_total,
"count": e.count,
} for e in averages]
with open(f"{stem}.json", "w") as fh:
with open(f"{stem}.json", "w", encoding="utf-8") as fh:
_json.dump(rows, fh)
if rank == "0":
logger.info("PROFILER: summary written to %s.txt", stem)
@@ -455,24 +528,25 @@ class TorchProfilerController:
logger.exception("PROFILER: summary generation failed")
def stop(self) -> None:
"""Stop the profiler after disabling collection and clearing state."""
"""Flush any active segment and disable this controller."""
if self._profiler is None:
if not self._configured:
return
logger.info("PROFILER: Stopping profiler...")
self._profiler.stop()
self._write_summary()
self._profiler = None # makes stop() idempotent (atexit may re-enter)
self._finish_segment()
self._profiler = None
self._configured = False
self._armed = False
logger.info("PROFILER: Profiler stopped")
self._active_region_depth = 0
set_global_controller(None)
@property
def has_profiler(self) -> bool:
"""Return ``True`` when a profiler instance is available."""
"""Return ``True`` when this controller is configured and armed."""
return self._profiler is not None
return self._configured and self._armed
@property
def activities(self) -> tuple[torch.profiler.ProfilerActivity, ...]:
@@ -495,13 +569,21 @@ def profiler_region(region: str):
def profile_region(region: str) -> Callable[[Callable[..., Any]], Callable[..., Any]]:
"""Wrap a bound method so it runs inside a profiler region if available."""
"""Wrap a bound method so it runs inside a profiler region if available.
Prefer a controller attached to the instance, then fall back to the
process-wide controller. The fallback lets lightweight owners such as the
modular trainer and its callbacks add regions without threading profiler
plumbing through their public constructors.
"""
def decorator(fn: Callable[..., Any]) -> Callable[..., Any]:
@functools.wraps(fn)
def wrapped(self, *args, **kwargs):
controller = getattr(self, "profiler_controller", None)
if controller is None:
controller = get_global_controller()
if controller is None or not controller.has_profiler:
return fn(self, *args, **kwargs)
with controller.region(region):
@@ -1,110 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""GPU backward checks for the VSA-H3 backend.
The CuTe backend returns FA4's own output tensor, which FA4's autograd node
saved for its backward. Composing the compression branch onto it in place
therefore poisons the graph, and the failure only appears once the VSA-256
CuTe path has a backward at all. These tests pin the composition.
"""
import pytest
import torch
from fastvideo.attention.backends.video_sparse_attn_h3 import (MiniMaxH3VSAImpl, MiniMaxH3VSAMetadataBuilder)
_SPEC = dict(raw_latent_shape=(16, 16, 24), patch_size=(1, 2, 2), prefix_segments=(64, 32, 16))
_HEADS = 2
_DIM = 128
def _build_meta(device, sparsity=0.5):
return MiniMaxH3VSAMetadataBuilder().build(
current_timestep=0,
raw_latent_shape=_SPEC["raw_latent_shape"],
patch_size=_SPEC["patch_size"],
VSA_sparsity=sparsity,
prefix_segments=_SPEC["prefix_segments"],
device=device,
)
def _select_backend(monkeypatch, backend):
if backend == "cute":
pytest.importorskip(
"flash_attn.cute.block_sparsity",
reason="optional FA4 CuTe build (flash_attn.cute) not installed",
)
monkeypatch.setenv("FASTVIDEO_VSA_CUTEDSL", "1")
monkeypatch.delenv("FASTVIDEO_VSA_TRITON", raising=False)
monkeypatch.delenv("FASTVIDEO_KERNEL_VSA_FORCE_TRITON", raising=False)
else:
monkeypatch.setenv("FASTVIDEO_VSA_TRITON", "1")
monkeypatch.delenv("FASTVIDEO_VSA_CUTEDSL", raising=False)
def _forward_backward(impl, meta, gate_compress, device):
seq = meta.total_seq_length
torch.manual_seed(0)
q, k, v = (torch.randn(1, seq, _HEADS, _DIM, device=device, dtype=torch.bfloat16, requires_grad=True)
for _ in range(3))
tq, tk, tv = (impl.tile(t, meta).clone() for t in (q, k, v))
gate = None
if gate_compress:
gate = torch.randn(1, tq.shape[1], _HEADS, _DIM, device=device, dtype=torch.bfloat16) * 0.1
out = impl.forward(tq, tk, tv, gate, meta)
out = impl.postprocess_output(out, meta)
out.float().pow(2).sum().backward()
return out, (q, k, v)
@pytest.mark.parametrize("backend", ["triton", "cute"])
@pytest.mark.parametrize("gate_compress", [False, True])
def test_h3_vsa_backward_runs(monkeypatch, backend: str, gate_compress: bool) -> None:
"""Regression: with the CuTe backend and a non-zero gate this used to die
with "one of the variables needed for gradient computation has been
modified by an inplace operation ... output 0 of FlashAttnFuncBackward".
"""
if not torch.cuda.is_available():
pytest.skip("CUDA is required")
_select_backend(monkeypatch, backend)
device = torch.device("cuda")
meta = _build_meta(device)
impl = MiniMaxH3VSAImpl(num_heads=_HEADS, head_size=_DIM, causal=False, softmax_scale=_DIM**-0.5)
out, leaves = _forward_backward(impl, meta, gate_compress, device)
assert torch.isfinite(out).all().item()
for name, leaf in zip(("q", "k", "v"), leaves):
assert leaf.grad is not None, f"{name} received no gradient"
assert torch.isfinite(leaf.grad).all().item(), f"{name}.grad has non-finite values"
assert leaf.grad.abs().sum().item() > 0, f"{name}.grad is all zero"
@pytest.mark.parametrize("gate_compress", [False, True])
def test_h3_vsa_backward_cute_matches_triton(monkeypatch, gate_compress: bool) -> None:
"""CuTe and Triton take different routes to the same math; their gradients
should agree to bf16 tolerance."""
if not torch.cuda.is_available():
pytest.skip("CUDA is required")
device = torch.device("cuda")
impl = MiniMaxH3VSAImpl(num_heads=_HEADS, head_size=_DIM, causal=False, softmax_scale=_DIM**-0.5)
grads = {}
for backend in ("triton", "cute"):
with monkeypatch.context() as m:
_select_backend(m, backend)
meta = _build_meta(device)
_, leaves = _forward_backward(impl, meta, gate_compress, device)
grads[backend] = [leaf.grad.detach().float() for leaf in leaves]
for name, ref, got in zip(("dq", "dk", "dv"), grads["triton"], grads["cute"]):
diff = (ref - got).abs()
avg_abs = diff.mean().item()
max_rel = (diff.max() / (ref.abs().mean() + 1e-6)).item()
print(f"[h3-vsa gate={gate_compress}] {name}: avg_abs={avg_abs:.6e}, max_rel={max_rel:.6e}")
assert avg_abs < 1e-2, f"{name}: avg_abs {avg_abs:.3e}"
assert max_rel < 0.5, f"{name}: max_rel {max_rel:.3e}"
@@ -5,29 +5,20 @@ reference. The same reference doubles as the GPU kernel parity oracle."""
import math
import pytest
import torch
import torch.nn.functional as F
from fastvideo.attention.backends.video_sparse_attn_h3 import (_TILE_ELEMS, MiniMaxH3VSAImpl,
MiniMaxH3VSAMetadataBuilder, _build_block_mask,
_pool_tiles, _validate_h3_tile_geometry,
token_tile_and_valid)
_pool_tiles, token_tile_and_valid)
_720P = dict(raw_latent_shape=(30, 44, 80), patch_size=(1, 2, 2), prefix_segments=(512, 1760, 400))
_TINY = dict(raw_latent_shape=(8, 8, 12), patch_size=(1, 2, 2), prefix_segments=(7, 5, 3))
# (4,4,4) coverage: dit grid (9, 10, 13) is ragged in all three dims
# (t: 4+4+1, h: 4+4+2, w: 4+4+4+1) and every prefix segment leaves a
# partial tail tile at 64 (70 -> 64+6, 5 -> 5, 130 -> 64+64+2).
_TINY64 = dict(raw_latent_shape=(9, 20, 26), patch_size=(1, 2, 2), prefix_segments=(70, 5, 130))
# production-shape request: 768x1344, 124 frames -> latents (37, 48, 84),
# patch (1,2,2) -> token grid (37, 24, 42); text 300 + audio 414 rows.
_PROD = dict(raw_latent_shape=(37, 48, 84), patch_size=(1, 2, 2), prefix_segments=(300, 0, 414))
_CPU = torch.device("cpu")
def _build(spec, sparsity=0.0, device=_CPU, tile_size=_TILE_ELEMS):
def _build(spec, sparsity=0.0, device=_CPU):
return MiniMaxH3VSAMetadataBuilder().build(
current_timestep=0,
raw_latent_shape=spec["raw_latent_shape"],
@@ -35,7 +26,6 @@ def _build(spec, sparsity=0.0, device=_CPU, tile_size=_TILE_ELEMS):
VSA_sparsity=sparsity,
prefix_segments=spec["prefix_segments"],
device=device,
tile_size=tile_size,
)
@@ -46,7 +36,7 @@ def _impl():
def reference_sparse_attention(query, key, value, mask, meta):
"""Token-level oracle: SDPA over the padded tile buffer with the block
mask expanded to tokens. query/key/value: tiled [B, S_pad, H, D]."""
token_tile, token_valid = token_tile_and_valid(meta.variable_block_sizes, meta.tile_elems)
token_tile, token_valid = token_tile_and_valid(meta.variable_block_sizes)
out = torch.empty_like(query)
for b in range(query.shape[0]):
for h in range(query.shape[2]):
@@ -143,117 +133,9 @@ def test_prefix_queries_stay_dense_at_high_sparsity():
"video rows should actually be sparse at 75%"
# ---------------------------------------------------------------------------
# 64-token (4,4,4) tile geometry
# ---------------------------------------------------------------------------
def test_geometry_tile64_ragged_tails():
"""Hand-computed (4,4,4) oracle on a grid ragged in all three dims."""
meta = _build(_TINY64, tile_size=64)
assert meta.tile_elems == 64
t, h, w = 9, 10, 13 # raw latents (9, 20, 26) under patch (1, 2, 2)
n_t, n_h, n_w = 3, 3, 4
prefix_len = sum(_TINY64["prefix_segments"])
seq = prefix_len + t * h * w
assert meta.total_seq_length == seq
assert meta.num_prefix_tiles == 2 + 1 + 3
assert meta.num_video_tiles == n_t * n_h * n_w
assert int(meta.variable_block_sizes.sum()) == seq
assert int(meta.variable_block_sizes.max()) <= 64
assert meta.variable_block_sizes[:meta.num_prefix_tiles].tolist() == [64, 6, 5, 64, 64, 2]
# per-tile valid sizes: product of the per-dim clamped tails
expected = torch.tensor([
min(4, t - 4 * tt) * min(4, h - 4 * hh) * min(4, w - 4 * ww) for tt in range(n_t) for hh in range(n_h)
for ww in range(n_w)
],
dtype=torch.long)
assert torch.equal(meta.variable_block_sizes[meta.num_prefix_tiles:], expected)
assert int(expected.min()) == 1 * 2 * 1 # the (t,h,w) ragged corner
# every packed video row lands in the 3D tile its (t,h,w) coordinate says
idx = meta.untile_combined_index
row = torch.arange(t * h * w)
row_t, row_h, row_w = row // (h * w), (row // w) % h, row % w
expected_tile = meta.num_prefix_tiles + ((row_t // 4) * n_h + row_h // 4) * n_w + row_w // 4
assert torch.equal(idx[prefix_len:] // 64, expected_tile)
# and in a non-pad slot of that tile
assert bool((idx % 64 < meta.variable_block_sizes[idx // 64]).all())
# untile(tile(x)) == x on the 64-wide padded buffer
x = torch.randn(1, seq, 2, 4)
buf = _impl().tile(x, meta)
assert buf.shape[1] == meta.variable_block_sizes.numel() * 64
assert torch.equal(buf[:, idx], x)
def test_geometry_tile64_production_shape():
"""Production latents (37, 48, 84): ragged t and w tails at (4,4,4)."""
meta64 = _build(_PROD, tile_size=64)
assert meta64.num_prefix_tiles == 5 + 7 # 300 -> 4x64+44, 414 -> 6x64+30
assert meta64.num_video_tiles == 10 * 6 * 11 # (37, 24, 42) / (4, 4, 4)
assert meta64.total_seq_length == 300 + 414 + 37 * 24 * 42
assert int(meta64.variable_block_sizes.sum()) == meta64.total_seq_length
sizes_vid = meta64.variable_block_sizes[meta64.num_prefix_tiles:]
assert int(sizes_vid.max()) == 64 and int(sizes_vid.min()) == 1 * 4 * 2 # (t, w) ragged corner
# same packed sequence under the default 256 geometry, fewer tiles
meta256 = _build(_PROD)
assert meta256.tile_elems == _TILE_ELEMS
assert meta256.num_prefix_tiles == 2 + 2
assert meta256.num_video_tiles == 10 * 3 * 6
assert meta256.total_seq_length == meta64.total_seq_length
x = torch.randn(1, meta64.total_seq_length, 2, 4)
buf = _impl().tile(x, meta64)
assert torch.equal(buf[:, meta64.untile_combined_index], x)
def test_sparsity_zero_matches_dense_sdpa_tile64():
torch.manual_seed(2)
meta = _build(_TINY64, tile_size=64)
seq = meta.total_seq_length
q, k, v = (torch.randn(1, seq, 2, 8) for _ in range(3))
impl = _impl()
tq, tk, tv = (impl.tile(t, meta).clone() for t in (q, k, v))
scores = torch.matmul(_pool_tiles(tq, meta.variable_block_sizes, meta.tile_elems),
_pool_tiles(tk, meta.variable_block_sizes, meta.tile_elems).transpose(-2, -1))
mask = _build_block_mask(scores, meta.num_prefix_tiles, meta.num_video_tiles, 0.0, exempt=True)
sparse_out = impl.postprocess_output(reference_sparse_attention(tq, tk, tv, mask, meta), meta)
dense_out = F.scaled_dot_product_attention(q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2)).transpose(1, 2)
assert torch.allclose(sparse_out, dense_out, atol=1e-5), (sparse_out - dense_out).abs().max()
def test_geometry_guard_enforces_tile64_bound():
"""A 65-token tile passes the 256 bound but must fail the 64 one."""
meta = _build(_TINY64, tile_size=64)
prefix = tuple(s for s in _TINY64["prefix_segments"] if s > 0)
dit_shape = (9, 10, 13)
sizes = meta.variable_block_sizes.clone()
sizes[0] = 65
with pytest.raises(ValueError, match="tile sizes out of bounds"):
_validate_h3_tile_geometry(prefix, dit_shape, sizes, meta.untile_combined_index, 64)
# the untampered tile-64 geometry passes its own bound
_validate_h3_tile_geometry(prefix, dit_shape, meta.variable_block_sizes, meta.untile_combined_index, 64)
def test_builder_rejects_unknown_tile_size():
for bad in (0, 128, 512):
with pytest.raises(ValueError, match="tile_size"):
_build(_TINY, tile_size=bad)
if __name__ == "__main__":
test_geometry_720p()
test_mask_policy()
test_sparsity_zero_matches_dense_sdpa()
test_prefix_queries_stay_dense_at_high_sparsity()
test_geometry_tile64_ragged_tails()
test_geometry_tile64_production_shape()
test_sparsity_zero_matches_dense_sdpa_tile64()
test_geometry_guard_enforces_tile64_bound()
test_builder_rejects_unknown_tile_size()
print("all VSA-H3 CPU checks passed")
@@ -1,181 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""CPU checks for the VSA-H3 tile-64 sm_100a route selection.
The opt-in third kernel route (``FASTVIDEO_VSA_SM100A=1``) must (a) stay off by
default, (b) engage only when the extension is present, the device qualifies,
and the forward carries no grad, and (c) fall back to the Triton-64 entry with
one warning when the env is set but a precondition fails. All device/extension
probes are monkeypatched; no GPU or kernel install needed.
"""
import pytest
import torch
import fastvideo.attention.backends.video_sparse_attn_h3 as vsa_h3
from fastvideo.attention.backends.video_sparse_attn_h3 import (VSA_SM100A_ENV, MiniMaxH3VSAImpl,
MiniMaxH3VSAMetadataBuilder, _sm100a_unavailable_reason)
# Small tile-64 geometry: 2 prefix segments + a (4,4,8)-token video grid.
_SPEC = dict(raw_latent_shape=(4, 8, 16), patch_size=(1, 2, 2), prefix_segments=(70, 30))
_HEADS, _DIM = 2, 128
def _build_meta():
return MiniMaxH3VSAMetadataBuilder().build(
current_timestep=0,
raw_latent_shape=_SPEC["raw_latent_shape"],
patch_size=_SPEC["patch_size"],
VSA_sparsity=0.0,
prefix_segments=_SPEC["prefix_segments"],
device=torch.device("cpu"),
tile_size=64,
)
def _tiled_qkv(meta, requires_grad=False):
# bf16 like the real tiled buffers, so forward()'s dtype-cast warning
# stays out of the warning assertions below.
s_pad = meta.variable_block_sizes.numel() * 64
return tuple(
torch.randn(1, s_pad, _HEADS, _DIM, dtype=torch.bfloat16, requires_grad=requires_grad) for _ in range(3))
class _FakeSm100a:
"""Stands in for fastvideo_kernel.block_sparse_attn_sm100a."""
def __init__(self, supported=True):
self.supported = supported
self.calls = []
def is_supported(self, q, variable_block_sizes):
return self.supported
def block_sparse_attn_sm100a(self, q, k, v, q2k_idx, q2k_num, variable_block_sizes, need_lse=True):
self.calls.append(dict(q=q, q2k_idx=q2k_idx, q2k_num=q2k_num, vbs=variable_block_sizes,
need_lse=need_lse))
return q.clone(), None
def _fake_map_to_index(block_map):
"""Pure-torch stand-in for the Triton map_to_index (same contract)."""
b, h, t, n = block_map.shape
idx = torch.full((b, h, t, n), -1, dtype=torch.int32)
num = block_map.sum(dim=-1, dtype=torch.int32)
for bi in range(b):
for hi in range(h):
for ti in range(t):
cols = torch.nonzero(block_map[bi, hi, ti], as_tuple=False).flatten()
idx[bi, hi, ti, :cols.numel()] = cols.to(torch.int32)
return idx, num
class _FakeTriton:
def __init__(self):
self.calls = 0
def __call__(self, q, k, v, mask, variable_block_sizes):
self.calls += 1
return q.clone(), None
@pytest.fixture()
def routed(monkeypatch):
"""Backend with both kernel entries faked; returns (fakes, run)."""
fake_sm = _FakeSm100a()
fake_triton = _FakeTriton()
monkeypatch.setattr(vsa_h3, "_sm100a", fake_sm)
monkeypatch.setattr(vsa_h3, "block_sparse_attn_64_bhsd", fake_triton)
monkeypatch.setattr(vsa_h3, "map_to_index", _fake_map_to_index)
meta = _build_meta()
impl = MiniMaxH3VSAImpl(num_heads=_HEADS, head_size=_DIM, causal=False, softmax_scale=_DIM**-0.5)
def run(requires_grad=False):
q, k, v = _tiled_qkv(meta, requires_grad=requires_grad)
return impl.forward(q, k, v, None, meta)
return fake_sm, fake_triton, run, meta
def test_reason_covers_every_precondition():
q = torch.randn(1, _HEADS, 128, _DIM)
vbs = torch.full((2, ), 64, dtype=torch.long)
assert "not installed" in _sm100a_unavailable_reason(None, q, vbs, grad_mode=False)
ok = _FakeSm100a(supported=True)
assert "forward-only" in _sm100a_unavailable_reason(ok, q, vbs, grad_mode=True)
bad = _FakeSm100a(supported=False)
assert "is_supported" in _sm100a_unavailable_reason(bad, q, vbs, grad_mode=False)
assert _sm100a_unavailable_reason(ok, q, vbs, grad_mode=False) is None
def test_default_off_routes_triton(routed, monkeypatch):
fake_sm, fake_triton, run, _ = routed
monkeypatch.delenv(VSA_SM100A_ENV, raising=False)
run()
assert fake_triton.calls == 1
assert fake_sm.calls == []
def test_env_on_routes_sm100a_with_index_metadata(routed, monkeypatch):
fake_sm, fake_triton, run, meta = routed
monkeypatch.setenv(VSA_SM100A_ENV, "1")
out = run()
assert fake_triton.calls == 0
assert len(fake_sm.calls) == 1
call = fake_sm.calls[0]
n_tiles = meta.variable_block_sizes.numel()
# sparsity 0 -> all-True mask -> every row's count is n_tiles
assert call["q2k_num"].dtype == torch.int32 and (call["q2k_num"] == n_tiles).all()
assert call["q2k_idx"].shape[-1] == n_tiles and call["q2k_idx"].dtype == torch.int32
assert call["vbs"].dtype == torch.int32
assert call["need_lse"] is False
# BHSD kernel result comes back in the backend's BSHD layout
assert out.shape == (1, n_tiles * 64, _HEADS, _DIM)
def test_env_on_grad_inputs_fall_back_to_triton(routed, monkeypatch):
fake_sm, fake_triton, run, _ = routed
monkeypatch.setenv(VSA_SM100A_ENV, "1")
run(requires_grad=True)
assert fake_triton.calls == 1
assert fake_sm.calls == []
# ...but the same process still routes no-grad forwards to sm_100a
run(requires_grad=False)
assert len(fake_sm.calls) == 1
def test_env_on_unsupported_warns_once_and_falls_back(routed, monkeypatch):
fake_sm, fake_triton, run, _ = routed
fake_sm.supported = False
monkeypatch.setenv(VSA_SM100A_ENV, "1")
warnings = []
monkeypatch.setattr(vsa_h3.logger, "warning_once", warnings.append)
run()
run()
assert fake_triton.calls == 2
assert fake_sm.calls == []
assert len(warnings) == 2 # warning_once dedups by message; both carry the same one line
assert warnings[0] == warnings[1]
assert VSA_SM100A_ENV in warnings[0] and "is_supported" in warnings[0]
def test_env_on_missing_module_warns_and_falls_back(routed, monkeypatch):
fake_sm, fake_triton, run, _ = routed
monkeypatch.setattr(vsa_h3, "_sm100a", None)
monkeypatch.setenv(VSA_SM100A_ENV, "1")
warnings = []
monkeypatch.setattr(vsa_h3.logger, "warning_once", warnings.append)
run()
assert fake_triton.calls == 1
assert warnings and "not installed" in warnings[0]
def test_env_on_no_grad_context_detaches_route_from_leaf_flags(routed, monkeypatch):
"""A requires_grad leaf under torch.no_grad() is still a no-grad forward."""
fake_sm, fake_triton, run, meta = routed
monkeypatch.setenv(VSA_SM100A_ENV, "1")
impl = MiniMaxH3VSAImpl(num_heads=_HEADS, head_size=_DIM, causal=False, softmax_scale=_DIM**-0.5)
q, k, v = _tiled_qkv(meta, requires_grad=True)
with torch.no_grad():
impl.forward(q, k, v, None, meta)
assert len(fake_sm.calls) == 1
assert fake_triton.calls == 0
@@ -15,13 +15,10 @@ import json
import os
import subprocess
import sys
from unittest.mock import Mock
import pytest
import torch
from fastvideo.profiler import nvtx_range
# Five-window child: ops before any region, inside a region, between regions,
# inside a second (short-named) region, after the last region. Exits without
# calling stop() — export must happen via the atexit hook. Each window is
@@ -67,19 +64,21 @@ def _run_child(tmp_path):
return trace_dir, proc.stdout + proc.stderr
def _trace_event_names(trace_dir):
def _trace_events(trace_dir):
traces = glob.glob(os.path.join(trace_dir, "**", "*.json*"), recursive=True)
traces = [t for t in traces if "summary" not in os.path.basename(t)]
assert traces, f"no trace exported to {trace_dir} (atexit hook missing?)"
opener = gzip.open if traces[0].endswith(".gz") else open
with opener(traces[0], "rt") as fh:
events = json.load(fh).get("traceEvents", [])
return {e.get("name", "") for e in events}
events = []
for trace in traces:
opener = gzip.open if trace.endswith(".gz") else open
with opener(trace, "rt") as fh:
events.extend(json.load(fh).get("traceEvents", []))
return events
def test_regions_gate_collection_and_atexit_exports(tmp_path):
trace_dir, output = _run_child(tmp_path)
names = _trace_event_names(trace_dir)
names = {event.get("name", "") for event in _trace_events(trace_dir)}
assert "win_region1" in names, "op inside an enabled region was not captured"
assert "win_region2" in names, "short region name did not resolve/capture"
@@ -93,8 +92,61 @@ def test_regions_gate_collection_and_atexit_exports(tmp_path):
assert output.count("is not registered") == 1
# per-rank op summary written next to the trace
summaries = glob.glob(os.path.join(trace_dir, "summary_rank0.*"))
assert sorted(os.path.splitext(s)[1] for s in summaries) == [".json", ".txt"]
summaries = glob.glob(os.path.join(trace_dir, "summary_rank0_segment*.*"))
assert sorted(os.path.splitext(s)[1] for s in summaries) == [
".json",
".json",
".txt",
".txt",
]
with open(next(s for s in summaries if s.endswith(".json")), encoding="utf-8") as fh:
summary_rows = json.load(fh)
assert summary_rows
assert {
"name",
"shapes",
"self_cpu_us",
"cpu_us",
"self_device_us",
"device_us",
"count",
} <= summary_rows[0].keys()
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required")
def test_cuda_region_exports_kernel_events(tmp_path):
trace_dir = str(tmp_path / "cuda_traces")
child = r"""
import torch
from fastvideo.profiler import get_or_create_profiler, profiler_region
controller = get_or_create_profiler({trace_dir!r})
x = torch.randn(1024, 1024, device="cuda")
y = torch.randn(1024, 1024, device="cuda")
torch.cuda.synchronize()
with profiler_region("training_forward"):
torch.mm(x, y)
torch.cuda.synchronize()
controller.stop()
""".format(trace_dir=trace_dir)
env = os.environ.copy()
env["FASTVIDEO_TORCH_PROFILER_DIR"] = trace_dir
env["FASTVIDEO_TORCH_PROFILE_REGIONS"] = "training_forward"
proc = subprocess.run(
[sys.executable, "-c", child],
env=env,
capture_output=True,
text=True,
timeout=300,
)
assert proc.returncode == 0, proc.stderr
events = _trace_events(trace_dir)
categories = {event.get("cat", "") for event in events}
names = {event.get("name", "") for event in events}
assert "kernel" in categories
assert "cuda_runtime" in categories
assert "fastvideo.region::training_forward" in names
def test_noop_without_profiler_dir(tmp_path):
@@ -111,73 +163,3 @@ def test_noop_without_profiler_dir(tmp_path):
proc = subprocess.run([sys.executable, "-c", child], env=env,
capture_output=True, text=True, timeout=300)
assert proc.returncode == 0, proc.stderr
def test_nvtx_range_disabled_is_noop(monkeypatch):
"""Keep CUDA NVTX untouched when external profiling is disabled."""
range_push = Mock()
range_pop = Mock()
monkeypatch.setenv("FASTVIDEO_NVTX_PROFILE", "0")
monkeypatch.setattr(torch.cuda, "is_available", lambda: True)
monkeypatch.setattr(torch.cuda.nvtx, "range_push", range_push)
monkeypatch.setattr(torch.cuda.nvtx, "range_pop", range_pop)
with nvtx_range("disabled"):
body_executed = True
assert body_executed is True
range_push.assert_not_called()
range_pop.assert_not_called()
def test_nvtx_range_without_cuda_is_noop(monkeypatch):
"""Keep NVTX untouched when profiling is enabled on a CPU-only process."""
range_push = Mock()
range_pop = Mock()
monkeypatch.setenv("FASTVIDEO_NVTX_PROFILE", "1")
monkeypatch.setattr(torch.cuda, "is_available", lambda: False)
monkeypatch.setattr(torch.cuda.nvtx, "range_push", range_push)
monkeypatch.setattr(torch.cuda.nvtx, "range_pop", range_pop)
with nvtx_range("cpu-only"):
body_executed = True
assert body_executed is True
range_push.assert_not_called()
range_pop.assert_not_called()
def test_nvtx_range_enabled_orders_push_body_pop(monkeypatch):
"""Place the profiled body between one matching NVTX push and pop."""
events = []
monkeypatch.setenv("FASTVIDEO_NVTX_PROFILE", "1")
monkeypatch.setattr(torch.cuda, "is_available", lambda: True)
monkeypatch.setattr(torch.cuda.nvtx, "range_push", lambda name: events.append(("push", name)))
monkeypatch.setattr(torch.cuda.nvtx, "range_pop", lambda: events.append(("pop", None)))
with nvtx_range("minimax_h3.test"):
events.append(("body", None))
assert events == [
("push", "minimax_h3.test"),
("body", None),
("pop", None),
]
def test_nvtx_range_body_exception_pops_and_propagates(monkeypatch):
"""Balance the NVTX stack while preserving a body exception."""
events = []
monkeypatch.setenv("FASTVIDEO_NVTX_PROFILE", "1")
monkeypatch.setattr(torch.cuda, "is_available", lambda: True)
monkeypatch.setattr(torch.cuda.nvtx, "range_push", lambda name: events.append(("push", name)))
monkeypatch.setattr(torch.cuda.nvtx, "range_pop", lambda: events.append(("pop", None)))
with pytest.raises(RuntimeError, match="profile body failed"):
with nvtx_range("minimax_h3.failure"):
raise RuntimeError("profile body failed")
assert events == [
("push", "minimax_h3.failure"),
("pop", None),
]
@@ -1,281 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import json
import os
import pytest
import torch
os.environ.setdefault("MASTER_ADDR", "localhost")
os.environ.setdefault("MASTER_PORT", "29514")
import fastvideo.models.encoders.minimax_h3_checkpoint_fp8 as h3_fp8
from fastvideo.configs.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConfig
from fastvideo.layers.linear import ColumnParallelLinear, UnquantizedLinearMethod
from fastvideo.layers.vocab_parallel_embedding import UnquantizedEmbeddingMethod, VocabParallelEmbedding
from fastvideo.models.encoders.base import TextEncoder
from fastvideo.models.encoders.minimax_h3_checkpoint_fp8 import (
MiniMaxH3SerializedFP8Config,
MiniMaxH3SerializedFP8LinearMethod,
)
from fastvideo.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConditioner
from fastvideo.models.loader.text_encoder_quantization import (
_configure_text_encoder_quantization,
_process_quantized_text_encoder_weights,
_read_text_encoder_checkpoint_quantization_config,
)
def _checkpoint_quantization_config(**overrides) -> dict:
config = {
"quant_method": "fp8",
"activation_scheme": "dynamic",
"fmt": "e4m3",
"weight_block_size": [128, 128],
"modules_to_not_convert": ["model.visual", "lm_head"],
}
config.update(overrides)
return config
def test_h3_accepts_only_the_serialized_blockwise_checkpoint_contract() -> None:
config = MiniMaxH3SerializedFP8Config.from_config(_checkpoint_quantization_config())
assert config.weight_block_size == (128, 128)
assert config.get_supported_act_dtypes() == [torch.bfloat16]
with pytest.raises(ValueError, match=r"weight_block_size=\[128, 128\]"):
MiniMaxH3SerializedFP8Config.from_config(
_checkpoint_quantization_config(weight_block_size=[1, 128]))
with pytest.raises(ValueError, match="dynamic activation"):
MiniMaxH3SerializedFP8Config.from_config(
_checkpoint_quantization_config(activation_scheme="static"))
with pytest.raises(ValueError, match="vision stack"):
MiniMaxH3SerializedFP8Config.from_config(
_checkpoint_quantization_config(modules_to_not_convert=["lm_head"]))
with pytest.raises(ValueError, match="partially quantized language"):
MiniMaxH3SerializedFP8Config.from_config(
_checkpoint_quantization_config(modules_to_not_convert=["model.visual", "language_model.layers.3"]))
def test_serialized_fp8_allocates_checkpoint_weight_and_scale_without_requantization(distributed_setup) -> None:
config = MiniMaxH3SerializedFP8Config.from_config(_checkpoint_quantization_config())
layer = ColumnParallelLinear(
input_size=128,
output_size=256,
bias=False,
quant_config=config,
prefix="minimax_h3_qwen3_vl.language_model.layers.0.self_attn.q_proj",
)
assert isinstance(layer.quant_method, MiniMaxH3SerializedFP8LinearMethod)
assert layer.weight.dtype == torch.float8_e4m3fn
assert layer.weight.shape == (256, 128)
assert layer.weight_scale_inv.dtype == torch.float32
assert layer.weight_scale_inv.shape == (2, 1)
layer.weight.data.zero_()
layer.weight_scale_inv.data.fill_(0.25)
weight_pointer = layer.weight.data_ptr()
scale_pointer = layer.weight_scale_inv.data_ptr()
layer.quant_method.process_weights_after_loading(layer)
assert layer.weight.data_ptr() == weight_pointer
assert layer.weight_scale_inv.data_ptr() == scale_pointer
assert not hasattr(layer, "_fp8_weight")
def test_serialized_fp8_quantizes_only_language_linears(distributed_setup) -> None:
config = MiniMaxH3SerializedFP8Config.from_config(_checkpoint_quantization_config())
visual_linear = ColumnParallelLinear(
input_size=128,
output_size=128,
bias=False,
quant_config=config,
prefix="minimax_h3_qwen3_vl.visual.blocks.0.attn.proj",
)
embedding = VocabParallelEmbedding(
num_embeddings=128,
embedding_dim=128,
org_num_embeddings=128,
quant_config=config,
prefix="minimax_h3_qwen3_vl.language_model.embed_tokens",
)
assert isinstance(visual_linear.quant_method, UnquantizedLinearMethod)
assert visual_linear.weight.dtype == torch.get_default_dtype()
assert isinstance(embedding.quant_method, UnquantizedEmbeddingMethod)
assert embedding.weight.dtype == torch.get_default_dtype()
def test_serialized_fp8_cpu_execution_fails_closed(distributed_setup) -> None:
config = MiniMaxH3SerializedFP8Config.from_config(_checkpoint_quantization_config())
layer = ColumnParallelLinear(
input_size=128,
output_size=128,
bias=False,
quant_config=config,
prefix="minimax_h3_qwen3_vl.language_model.layers.0.mlp.up_proj",
)
layer.weight.data.zero_()
layer.weight_scale_inv.data.fill_(1.0)
assert isinstance(layer.quant_method, MiniMaxH3SerializedFP8LinearMethod)
layer.quant_method.process_weights_after_loading(layer)
with pytest.raises(RuntimeError, match="requires CUDA"):
layer(torch.zeros(2, 128, dtype=torch.bfloat16))
def test_runtime_preflight_reports_capability_and_missing_dependencies(monkeypatch) -> None:
config = MiniMaxH3SerializedFP8Config.from_config(_checkpoint_quantization_config())
monkeypatch.setattr(torch.cuda, "get_device_capability", lambda device: (8, 0))
with pytest.raises(RuntimeError, match="sm100 or newer"):
config.validate_runtime(torch.device("cuda"))
monkeypatch.setattr(torch.cuda, "get_device_capability", lambda device: (10, 0))
def missing_quantizer() -> None:
raise RuntimeError("SGLang-compatible Triton quantizer is missing")
monkeypatch.setattr(h3_fp8, "_require_sglang_per_token_group_fp8_quantization", missing_quantizer)
with pytest.raises(RuntimeError, match="Triton quantizer is missing"):
config.validate_runtime(torch.device("cuda"))
monkeypatch.setattr(h3_fp8, "_require_sglang_per_token_group_fp8_quantization", lambda: None)
def missing_flashinfer():
raise RuntimeError("FlashInfer groupwise GEMM is missing")
monkeypatch.setattr(h3_fp8, "_get_flashinfer_groupwise_fp8_gemm", missing_flashinfer)
with pytest.raises(RuntimeError, match="FlashInfer groupwise GEMM is missing"):
config.validate_runtime(torch.device("cuda"))
def test_loader_detects_and_capability_gates_checkpoint_metadata(tmp_path) -> None:
checkpoint_config = _checkpoint_quantization_config()
(tmp_path / "config.json").write_text(
json.dumps({"quantization_config": checkpoint_config}),
encoding="utf-8",
)
assert _read_text_encoder_checkpoint_quantization_config(str(tmp_path)) == checkpoint_config
model_config = MiniMaxH3Qwen3VLConfig()
quant_config = _configure_text_encoder_quantization(
model_config,
MiniMaxH3Qwen3VLConditioner,
str(tmp_path),
)
assert isinstance(quant_config, MiniMaxH3SerializedFP8Config)
assert model_config.quant_config is quant_config
unsupported_config = MiniMaxH3Qwen3VLConfig()
with pytest.raises(ValueError, match="does not support serialized 'fp8'"):
_configure_text_encoder_quantization(
unsupported_config,
TextEncoder,
str(tmp_path),
)
def test_loader_leaves_bf16_checkpoint_path_unchanged(tmp_path) -> None:
(tmp_path / "config.json").write_text(json.dumps({"architectures": ["Qwen3VLModel"]}), encoding="utf-8")
model_config = MiniMaxH3Qwen3VLConfig()
quant_config = _configure_text_encoder_quantization(
model_config,
MiniMaxH3Qwen3VLConditioner,
str(tmp_path),
)
assert quant_config is None
assert model_config.quant_config is None
def test_post_load_processing_visits_only_serialized_fp8_linears(distributed_setup) -> None:
config = MiniMaxH3SerializedFP8Config.from_config(_checkpoint_quantization_config())
quantized = ColumnParallelLinear(
input_size=128,
output_size=128,
bias=False,
quant_config=config,
prefix="minimax_h3_qwen3_vl.language_model.layers.0.self_attn.q_proj",
)
plain = ColumnParallelLinear(
input_size=128,
output_size=128,
bias=False,
prefix="plain",
)
quantized.weight.data.zero_()
quantized.weight_scale_inv.data.fill_(1.0)
model = torch.nn.ModuleList([quantized, plain])
assert _process_quantized_text_encoder_weights(model, torch.device("cpu")) == 1
assert quantized.weight.device.type == "cpu"
assert plain.weight.device.type == "cpu"
def test_flashinfer_groupwise_path_pins_output_dtype_and_trtllm_scale_layout(monkeypatch) -> None:
input_tensor = torch.zeros(2, 256, dtype=torch.bfloat16)
weight = torch.zeros(128, 256, dtype=torch.float8_e4m3fn)
weight_scale = torch.ones(1, 2, dtype=torch.float32)
quantized_input = torch.zeros_like(input_tensor, dtype=torch.float8_e4m3fn)
input_scale = torch.empty(2, 2, dtype=torch.float32).t()
input_scale.fill_(1.0)
receipt: dict[str, object] = {}
def fake_quantize(
value: torch.Tensor,
group_size: int,
*,
column_major_scales: bool,
) -> tuple[torch.Tensor, torch.Tensor]:
assert value.data_ptr() == input_tensor.data_ptr()
assert value.shape == input_tensor.shape
assert group_size == 128
assert column_major_scales is True
return quantized_input, input_scale
def fake_gemm(
activation: torch.Tensor,
checkpoint_weight: torch.Tensor,
activation_scale: torch.Tensor,
checkpoint_scale: torch.Tensor,
*,
out_dtype: torch.dtype,
backend: str,
) -> torch.Tensor:
receipt.update(
activation=activation,
checkpoint_weight=checkpoint_weight,
activation_scale=activation_scale,
checkpoint_scale=checkpoint_scale,
out_dtype=out_dtype,
backend=backend,
)
return torch.zeros(activation.shape[0], checkpoint_weight.shape[0], dtype=out_dtype)
monkeypatch.setattr(h3_fp8, "_get_flashinfer_groupwise_backend", lambda device: "trtllm")
monkeypatch.setattr(h3_fp8, "_sglang_per_token_group_quant_fp8", fake_quantize)
monkeypatch.setattr(h3_fp8, "_get_flashinfer_groupwise_fp8_gemm", lambda: fake_gemm)
previous_default_dtype = torch.get_default_dtype()
torch.set_default_dtype(torch.float32)
try:
output = h3_fp8._flashinfer_gemm_w8a8_block_fp8_linear_with_fallback(
input_tensor,
weight,
(128, 128),
weight_scale,
)
assert torch.get_default_dtype() == torch.float32
finally:
torch.set_default_dtype(previous_default_dtype)
assert output.dtype == torch.bfloat16
assert receipt["out_dtype"] == torch.bfloat16
assert receipt["backend"] == "trtllm"
assert receipt["activation"] is quantized_input
assert receipt["checkpoint_weight"] is weight
assert receipt["checkpoint_scale"] is weight_scale
assert receipt["activation_scale"] is input_scale
@@ -1,187 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""MiniMax-H3 Qwen3-VL layer truncation and slim-forward tests."""
from __future__ import annotations
import inspect
import os
import pytest
import torch
os.environ.setdefault("MASTER_ADDR", "localhost")
os.environ.setdefault("MASTER_PORT", "29513")
from fastvideo.configs.models.encoders.minimax_h3_qwen3_vl import (
MiniMaxH3Qwen3VLArchConfig,
MiniMaxH3Qwen3VLConfig,
)
from fastvideo.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLLanguageModel
from fastvideo.pipelines.basic.minimax_h3.packing import MINIMAX_H3_TEXT_ENCODER_LAYER
def _small_arch(**overrides) -> MiniMaxH3Qwen3VLArchConfig:
kwargs: dict = dict(
vocab_size=64,
hidden_size=16,
intermediate_size=32,
num_hidden_layers=8,
output_hidden_state_index=5,
num_attention_heads=2,
num_key_value_heads=1,
head_dim=8,
rope_scaling={
"mrope_interleaved": True,
"mrope_section": [2, 1, 1],
"rope_type": "default",
},
vision_out_hidden_size=16,
)
kwargs.update(overrides)
return MiniMaxH3Qwen3VLArchConfig(**kwargs)
def _small_config(**overrides) -> MiniMaxH3Qwen3VLConfig:
config = MiniMaxH3Qwen3VLConfig()
config.arch_config = _small_arch(**overrides)
return config
def test_default_matches_the_index_the_pipeline_reads() -> None:
config = MiniMaxH3Qwen3VLArchConfig()
assert config.output_hidden_state_index == MINIMAX_H3_TEXT_ENCODER_LAYER
assert config.num_hidden_layers_override == MINIMAX_H3_TEXT_ENCODER_LAYER
def test_rejects_build_depth_that_cannot_reach_the_output() -> None:
for override in (0, 4):
with pytest.raises(ValueError, match="num_hidden_layers_override"):
_small_arch(num_hidden_layers_override=override)
def test_rejects_output_index_above_the_checkpoint_depth() -> None:
with pytest.raises(ValueError, match="output_hidden_state_index"):
_small_arch(output_hidden_state_index=9, num_hidden_layers_override=None)
def test_builds_only_up_to_the_override(distributed_setup) -> None:
model = MiniMaxH3Qwen3VLLanguageModel(_small_config(num_hidden_layers_override=5))
assert model.num_layers == 5
assert len(model.layers) == 5
assert model.norm is None
def test_override_none_keeps_the_full_stack(distributed_setup) -> None:
model = MiniMaxH3Qwen3VLLanguageModel(_small_config(num_hidden_layers_override=None))
assert model.num_layers == 8
assert model.norm is not None
def test_nominal_and_built_depths_remain_distinct(distributed_setup) -> None:
from fastvideo.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConditioner
conditioner = MiniMaxH3Qwen3VLConditioner(_small_config(num_hidden_layers_override=5))
assert conditioner.num_hidden_layers == 8
assert conditioner.num_built_hidden_layers == 5
def test_override_above_the_stack_does_not_over_build(distributed_setup) -> None:
model = MiniMaxH3Qwen3VLLanguageModel(_small_config(num_hidden_layers_override=99))
assert model.num_layers == 8
assert model.norm is not None
def test_tapped_hidden_state_is_unchanged_by_truncation(distributed_setup) -> None:
"""The slim model returns the raw output at the selected layer."""
tap = 5
full = MiniMaxH3Qwen3VLLanguageModel(_small_config(num_hidden_layers_override=None))
cut = MiniMaxH3Qwen3VLLanguageModel(_small_config(num_hidden_layers_override=tap))
torch.manual_seed(0)
for parameter in full.parameters():
parameter.data.normal_(std=0.02)
for (_, a), (_, b) in zip(full.layers[:tap].named_parameters(),
cut.layers[:tap].named_parameters(),
strict=True):
b.data.copy_(a.data)
torch.manual_seed(1)
inputs_embeds = torch.randn(1, 6, 16)
position_ids = torch.arange(6).view(1, 1, 6).expand(3, 1, 6)
with torch.no_grad():
expected = inputs_embeds
position_embeddings = full.rotary_emb(inputs_embeds, position_ids)
for layer in full.layers[:tap]:
expected = layer(expected, position_embeddings, None)
full_out = full(inputs_embeds, position_ids, None, None, None)
cut_out = cut(inputs_embeds, position_ids, None, None, None)
assert torch.equal(expected, full_out)
assert torch.equal(expected, cut_out)
def test_conditioning_stage_adapts_slim_sequence_output() -> None:
from fastvideo.pipelines.basic.minimax_h3.stages.minimax_h3_conditioning import MiniMaxH3ConditioningStage
class FakeConditioner:
dtype = torch.float32
def __call__(self, input_ids: torch.Tensor, **kwargs) -> torch.Tensor:
assert input_ids.ndim == 1
assert not kwargs
return torch.ones(input_ids.shape[0], 4)
stage = MiniMaxH3ConditioningStage(conditioner=FakeConditioner(), tokenizer=None, processor=None, ref2va=False)
embeddings, tags = stage._encode_tokens([1, 2, 3], [0, 0, 0], torch.device("cpu"))
assert embeddings.shape == (1, 3, 4)
assert tags.shape == (3, )
def test_conditioner_exposes_only_the_slim_forward_contract() -> None:
from fastvideo.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConditioner
assert tuple(inspect.signature(MiniMaxH3Qwen3VLConditioner.forward).parameters) == (
"self",
"input_ids",
"pixel_values",
"image_grid_thw",
"pixel_values_videos",
"video_grid_thw",
)
def test_truncated_model_drops_the_surplus_checkpoint_keys(distributed_setup) -> None:
from fastvideo.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConditioner
conditioner = MiniMaxH3Qwen3VLConditioner(_small_config(num_hidden_layers_override=5))
assert conditioner._is_omitted_checkpoint_key("language_model.layers.5.mlp.gate_proj.weight")
assert conditioner._is_omitted_checkpoint_key("language_model.layers.7.self_attn.q_proj.weight")
assert conditioner._is_omitted_checkpoint_key("language_model.norm.weight")
assert not conditioner._is_omitted_checkpoint_key("language_model.layers.4.mlp.gate_proj.weight")
assert not conditioner._is_omitted_checkpoint_key("language_model.embed_tokens.weight")
assert not conditioner._is_omitted_checkpoint_key("visual.blocks.0.attn.qkv.weight")
assert not conditioner._is_omitted_checkpoint_key("language_model.layers.8.mlp.gate_proj.weight")
def test_full_stack_filters_nothing(distributed_setup) -> None:
from fastvideo.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConditioner
conditioner = MiniMaxH3Qwen3VLConditioner(_small_config(num_hidden_layers_override=None))
assert not conditioner._is_omitted_checkpoint_key("language_model.layers.7.mlp.gate_proj.weight")
assert not conditioner._is_omitted_checkpoint_key("language_model.norm.weight")
def test_corrupt_layer_above_checkpoint_depth_remains_unexpected(distributed_setup) -> None:
from fastvideo.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConditioner
conditioner = MiniMaxH3Qwen3VLConditioner(_small_config(num_hidden_layers_override=5))
with pytest.raises(ValueError, match="Unexpected"):
conditioner.load_weights([("language_model.layers.8.mlp.gate_proj.weight", torch.empty(1))])
@@ -2,10 +2,8 @@ import os
from types import SimpleNamespace
import warnings
import numpy as np
import pytest
import torch
from einops import rearrange
import fastvideo.entrypoints.video_generator as video_generator_module
from fastvideo.api import (
@@ -273,70 +271,6 @@ def test_generate_single_video_return_frames_still_materializes_output(tmp_path)
assert result["video_path"] is None
def test_generate_single_video_frames_match_legacy_cpu_loop(tmp_path):
"""The on-device quantize path (#1362) must reproduce the legacy
per-frame CPU loop (make_grid -> permute -> *255 -> uint8) bit-exactly
for in-range fp32 pixels: same uint8 dtype, same HWC grid layout with
nrow=6 (batch>1), odd frame count. CPU-only: on CUDA the float->uint8
cast may differ by <=1 LSB, but on CPU both orderings run identical
fp32 ops, so exact equality is required."""
torch.manual_seed(0)
output = torch.rand((2, 3, 3, 16, 16), dtype=torch.float32)
output_batch = _single_video_output_batch(output)
fastvideo_args = _single_video_args()
generator = _single_video_generator(output_batch, fastvideo_args)
sampling_param = _small_sampling_param(save_video=False, return_frames=True)
sampling_param.num_frames = 3
sampling_param.num_videos_per_prompt = 2
result = generator._generate_single_video(
prompt="grid parity",
sampling_param=sampling_param,
fastvideo_args=fastvideo_args,
output_path=str(tmp_path / "unused.mp4"),
)
legacy_frames = []
for x in rearrange(output, "b c t h w -> t b c h w"):
grid = video_generator_module.torchvision.utils.make_grid(x, nrow=6)
grid = grid.permute(1, 2, 0).squeeze(-1)
legacy_frames.append((grid * 255).to(torch.uint8).contiguous().cpu().numpy())
torch.testing.assert_close(result["samples"], output)
assert len(result["frames"]) == 3
for got, want in zip(result["frames"], legacy_frames, strict=True):
assert got.dtype == np.uint8
assert got.shape == want.shape
np.testing.assert_array_equal(got, want)
def test_generate_single_video_frames_clamp_out_of_range_pixels(tmp_path):
"""VAE output slightly outside [0, 1] must saturate at 0/255 in the
uint8 frames. The pre-#1362 unclamped cast wrapped mod 256 (e.g.
1.5 -> 126). CPU-only."""
output = torch.full((1, 3, 2, 16, 16), 1.5, dtype=torch.float32)
output[:, :, 1] = -0.5
output_batch = _single_video_output_batch(output)
fastvideo_args = _single_video_args()
generator = _single_video_generator(output_batch, fastvideo_args)
result = generator._generate_single_video(
prompt="clamp",
sampling_param=_small_sampling_param(save_video=False, return_frames=True),
fastvideo_args=fastvideo_args,
output_path=str(tmp_path / "unused.mp4"),
)
frames = result["frames"]
assert len(frames) == 2
# make_grid passes a single image through without grid padding, so
# every pixel comes from the (clamped) output tensor.
assert frames[0].dtype == np.uint8
assert frames[0].shape == (16, 16, 3)
assert (frames[0] == 255).all()
assert (frames[1] == 0).all()
def test_generate_single_video_save_video_still_builds_frames(monkeypatch, tmp_path):
output = torch.ones((1, 3, 2, 16, 16), dtype=torch.float32) * 0.5
output_batch = _single_video_output_batch(output)
@@ -371,37 +305,6 @@ def test_generate_single_video_save_video_still_builds_frames(monkeypatch, tmp_p
}
def test_generate_single_video_save_only_reports_refined_output_size(monkeypatch, tmp_path):
"""`GenerationResult.size` must describe the decoded media even when the
fp32 `samples` mirror is skipped (`return_frames=False`, the CLI save
flow). Refiner pipelines can change the final pixel geometry, so the size
has to come from `output_batch.output`, not the base request. CPU-only."""
# Refiner-style output: request asks for 2 frames of 16x16, pipeline
# produces 5 frames of 32x48.
output = torch.full((1, 3, 5, 32, 48), 0.5, dtype=torch.float32)
output_batch = _single_video_output_batch(output)
fastvideo_args = _single_video_args()
generator = _single_video_generator(output_batch, fastvideo_args)
saved = {}
def fake_mimsave(path, frames, *, fps, format):
saved["frame_count"] = len(frames)
monkeypatch.setattr(video_generator_module.imageio, "mimsave", fake_mimsave)
result = generator._generate_single_video(
prompt="refined save",
sampling_param=_small_sampling_param(save_video=True, return_frames=False),
fastvideo_args=fastvideo_args,
output_path=str(tmp_path / "refined.mp4"),
)
assert result["samples"] is None
assert result["frames"] is None
assert result["size"] == (32, 48, 5)
assert saved["frame_count"] == 5
def test_generate_single_video_audio_only_metadata_returns_audio_without_frames(tmp_path):
audio = torch.zeros((16, ), dtype=torch.float32)
output_batch = _single_video_output_batch(
@@ -1,185 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from types import SimpleNamespace
import numpy as np
import torch
from fastvideo.pipelines.basic.minimax_h3.packing import (
MiniMaxH3PackedLayout,
patchify_video_latents,
)
from fastvideo.pipelines.basic.minimax_h3.reference import MiniMaxH3PreparedReference
from fastvideo.pipelines.basic.minimax_h3.stages import minimax_h3_decoding
from fastvideo.pipelines.basic.minimax_h3.stages.minimax_h3_decoding import MiniMaxH3VideoDecodingStage
from fastvideo.pipelines.basic.minimax_h3.stages.minimax_h3_latent_preparation import (
MINIMAX_H3_LAYOUT_KEY,
MiniMaxH3LatentPreparationStage,
)
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
def _layout(rows: int, latent_shape: tuple[int, ...]) -> MiniMaxH3PackedLayout:
empty = torch.empty(0, dtype=torch.long)
return MiniMaxH3PackedLayout(
sequence_length=rows,
position_ids=empty,
token_tags=empty,
video_indices=empty,
audio_indices=empty,
text_indices=empty,
num_condition_video_rows=0,
num_condition_audio_rows=0,
num_video_latent_frames=latent_shape[2],
latent_height=latent_shape[3],
latent_width=latent_shape[4],
num_audio_latents=0,
)
def test_reference_video_encode_keeps_pixels_on_cpu() -> None:
observed = {}
class VAE:
def encode_pixels(self, pixels):
observed["pixels"] = pixels
posterior = SimpleNamespace(sample=lambda generator=None: torch.zeros(1, 4, 7, 4, 4))
return SimpleNamespace(latent_dist=posterior)
def normalize_latents(self, latents):
return latents
stage = MiniMaxH3LatentPreparationStage(
transformer=SimpleNamespace(patch_size=(1, 1, 1)),
vae=VAE(),
audio_vae=None,
scheduler=None,
ref2va=True,
)
reference = MiniMaxH3PreparedReference(
media_type="video",
frames=np.zeros((22, 16, 16, 3), dtype=np.uint8),
)
args = SimpleNamespace(vae_parallel_encode=False)
rows = stage._encode_visual_rows([reference], torch.device("cpu"), args)
assert observed["pixels"].dtype == torch.uint8
assert observed["pixels"].device.type == "cpu"
assert rows[0].shape == (7 * 4 * 4, 4)
def test_decode_stage_uses_cpu_output_buffer(monkeypatch) -> None:
latent_shape = (1, 4, 2, 4, 4)
latents = torch.randn(latent_shape)
rows = patchify_video_latents(latents, (1, 1, 1))
batch = ForwardBatch(data_type="video", latents=rows, raw_latent_shape=latent_shape)
batch.extra[MINIMAX_H3_LAYOUT_KEY] = _layout(rows.shape[0], latent_shape)
observed = {}
class VAE:
def to(self, device):
return self
def denormalize_latents(self, decoded_latents):
return decoded_latents
def decoded_pixel_shape(self, shape):
assert tuple(shape) == latent_shape
return (1, 3, 5, 16, 16)
def decode_to_pixels(self, decoded_latents, output):
observed["latents"] = decoded_latents
observed["output"] = output
output.fill_(0.25)
monkeypatch.setattr(minimax_h3_decoding, "get_local_torch_device", lambda: torch.device("cpu"))
result = MiniMaxH3VideoDecodingStage(VAE(), SimpleNamespace(patch_size=(1, 1, 1))).forward(
batch,
SimpleNamespace(output_type="pil", pin_cpu_memory=False, vae_cpu_offload=False, vae_parallel_decode=False),
)
torch.testing.assert_close(observed["latents"], latents)
assert observed["output"] is result.output
assert result.output.device.type == "cpu"
assert torch.all(result.output == 0.25)
def test_decode_stages_skip_vae_on_non_output_rank(monkeypatch) -> None:
class VAE:
sampling_rate = 32000
def to(self, device):
raise AssertionError("non-output ranks must not execute a VAE")
monkeypatch.setattr(minimax_h3_decoding, "model_parallel_is_initialized", lambda: True)
monkeypatch.setattr(minimax_h3_decoding, "get_sp_group",
lambda: SimpleNamespace(is_first_rank=False, world_size=4, rank_in_group=1))
args = SimpleNamespace(output_type="pil", pin_cpu_memory=False, vae_cpu_offload=True, vae_parallel_decode=False)
video = MiniMaxH3VideoDecodingStage(VAE(), SimpleNamespace()).forward(ForwardBatch(data_type="video"), args)
assert video.output.shape == (0, 3, 0, 0, 0)
audio_batch = ForwardBatch(data_type="audio", latents=torch.zeros(1), audio_latents=torch.zeros(1))
audio_batch.extra[MINIMAX_H3_LAYOUT_KEY] = object()
audio = minimax_h3_decoding.MiniMaxH3AudioDecodingStage(VAE()).forward(audio_batch, args)
assert audio.extra["audio"].shape == (0, 2)
assert audio.extra["audio_sample_rate"] == 32000
assert audio.latents is None
assert audio.audio_latents is None
assert MINIMAX_H3_LAYOUT_KEY not in audio.extra
def test_parallel_decode_runs_on_every_rank(monkeypatch) -> None:
"""With vae_parallel_decode, non-leader ranks must enter the decode body
(the collectives inside require uniform participation) and only the
leader owns the CPU output buffer."""
latent_shape = (1, 4, 2, 4, 4)
rows = patchify_video_latents(torch.randn(latent_shape), (1, 1, 1))
calls = []
class VAE:
def to(self, device):
return self
def denormalize_latents(self, decoded_latents):
return decoded_latents
def decoded_pixel_shape(self, shape):
return (1, 3, 5, 16, 16)
def fake_parallel(vae, latents, output, group, strategy):
calls.append((group.rank_in_group, output, strategy))
if output is not None:
output.fill_(0.5)
return output
monkeypatch.setattr(minimax_h3_decoding, "get_local_torch_device", lambda: torch.device("cpu"))
monkeypatch.setattr(minimax_h3_decoding, "model_parallel_is_initialized", lambda: True)
monkeypatch.setattr(minimax_h3_decoding, "decode_to_pixels_parallel", fake_parallel)
args = SimpleNamespace(output_type="pil",
pin_cpu_memory=False,
vae_cpu_offload=False,
vae_parallel_decode=True,
vae_parallel_decode_strategy="gather")
for rank, is_first in ((0, True), (2, False)):
monkeypatch.setattr(
minimax_h3_decoding, "get_sp_group",
lambda rank=rank, is_first=is_first: SimpleNamespace(is_first_rank=is_first,
world_size=4,
rank_in_group=rank))
batch = ForwardBatch(data_type="video", latents=rows.clone(), raw_latent_shape=latent_shape)
batch.extra[MINIMAX_H3_LAYOUT_KEY] = _layout(rows.shape[0], latent_shape)
result = MiniMaxH3VideoDecodingStage(VAE(), SimpleNamespace(patch_size=(1, 1, 1))).forward(batch, args)
if is_first:
assert result.output.shape == (1, 3, 5, 16, 16)
assert torch.all(result.output == 0.5)
else:
assert result.output.shape == (0, 3, 0, 0, 0)
assert [(rank, output is not None) for rank, output, _ in calls] == [(0, True), (2, False)]
assert all(strategy == "gather" for _, _, strategy in calls)
@@ -42,9 +42,15 @@ class _Method:
self,
transformer: torch.nn.Module | None,
tracker: Any | None = None,
updated_iterations: set[int] | None = None,
) -> None:
self.student = _Student(transformer)
self.tracker = tracker
self.updated_iterations = updated_iterations
def did_update_role(self, role: str, iteration: int) -> bool:
assert role == "student"
return (self.updated_iterations is None or iteration in self.updated_iterations)
def _tiny_transformer(*, fill: float = 0.0) -> torch.nn.Module:
@@ -155,6 +161,44 @@ class TestOnTrainingStepEnd:
assert any(payload.get("ema/decay") == 0.99 and step == 0 for payload, step in tracker.entries)
def test_only_updates_after_student_optimizer_step(self) -> None:
transformer = _tiny_transformer(fill=1.0)
tracker = _RecordingTracker()
method = _Method(
transformer,
tracker=tracker,
updated_iterations={5},
)
cb = EMACallback(decay=0.9, start_iter=0)
cb.on_train_start(method, iteration=0)
with torch.no_grad():
transformer.weight.fill_(7.0)
cb.on_training_step_end(method, loss_dict={}, iteration=4)
assert not cb._ema_started
assert torch.allclose(
cb.student_ema.shadow["weight"],
torch.full((2, 4), 1.0),
)
assert tracker.entries == []
cb.on_training_step_end(method, loss_dict={}, iteration=5)
assert cb._ema_started
assert torch.allclose(
cb.student_ema.shadow["weight"],
torch.full((2, 4), 7.0),
)
assert tracker.entries == [({"ema/decay": 0.9}, 5)]
with torch.no_grad():
transformer.weight.fill_(11.0)
cb.on_training_step_end(method, loss_dict={}, iteration=6)
assert torch.allclose(
cb.student_ema.shadow["weight"],
torch.full((2, 4), 7.0),
)
assert tracker.entries == [({"ema/decay": 0.9}, 5)]
# ---------------------------------------------------------------------------
# C. ema_context
@@ -0,0 +1,189 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
from types import SimpleNamespace
from typing import Any
import pytest
import torch
from fastvideo.train.methods.distribution_matching.dmd2 import DMD2Method
class _RecordingStudent:
def __init__(self) -> None:
self.predict_calls: list[torch.Tensor] = []
self.add_noise_calls: list[torch.Tensor] = []
def predict_x0(
self,
noisy_latents: torch.Tensor,
timestep: torch.Tensor,
batch: Any,
**kwargs: Any,
) -> torch.Tensor:
del batch, kwargs
self.predict_calls.append(timestep.detach().clone())
return noisy_latents + timestep.to(noisy_latents.dtype)
def add_noise(
self,
clean_latents: torch.Tensor,
noise: torch.Tensor,
timestep: torch.Tensor,
) -> torch.Tensor:
self.add_noise_calls.append(timestep.detach().clone())
return clean_latents + noise
def _rollout_method(seed: int = 1234) -> tuple[DMD2Method, _RecordingStudent]:
method = object.__new__(DMD2Method)
torch.nn.Module.__init__(method)
student = _RecordingStudent()
method.student = student
method._rollout_mode = "simulate"
method._cfg_uncond = None
method._denoising_step_list = torch.tensor([1000, 750, 500, 250])
method.cuda_generator = torch.Generator(device="cpu").manual_seed(seed)
return method, student
def _legacy_full_rollout_reference(
*,
seed: int,
target_idx: int,
shape: tuple[int, ...],
) -> tuple[torch.Tensor, torch.Tensor]:
"""Reproduce the pre-optimization simulate rollout and RNG state."""
step_list = torch.tensor([1000, 750, 500, 250])
generator = torch.Generator(device="cpu").manual_seed(seed)
current = torch.randn(shape, generator=generator)
initial = current.clone()
noise_latents: list[torch.Tensor] = []
for step_idx in range(len(step_list) - 1):
pred_clean = current + step_list[step_idx].to(current.dtype)
noise = torch.randn(shape, generator=generator)
current = pred_clean + noise
noise_latents.append(current.clone())
noisy_input: torch.Tensor
if target_idx == 0:
noisy_input = initial
else:
noisy_input = noise_latents[target_idx - 1]
output = noisy_input + step_list[target_idx].to(noisy_input.dtype)
return output, generator.get_state()
@pytest.mark.parametrize("target_idx", range(4))
def test_simulate_rollout_only_runs_required_prefix_forwards(
monkeypatch: pytest.MonkeyPatch,
target_idx: int,
) -> None:
method, student = _rollout_method()
batch = SimpleNamespace(
latents=torch.zeros((1, 2)),
dmd_latent_vis_dict={},
)
def _fixed_target(*args: Any, **kwargs: Any) -> torch.Tensor:
del args, kwargs
return torch.tensor([target_idx], dtype=torch.long)
monkeypatch.setattr(torch, "randint", _fixed_target)
method._student_rollout(batch, with_grad=True)
# One prediction per required prefix step, plus the differentiable target
# prediction. Prefix noising only happens for the required prefix steps.
assert len(student.predict_calls) == target_idx + 1
assert len(student.add_noise_calls) == target_idx
@pytest.mark.parametrize("target_idx", range(4))
def test_simulate_rollout_preserves_method_generator_progress(
monkeypatch: pytest.MonkeyPatch,
target_idx: int,
) -> None:
seed = 4321
method, _ = _rollout_method(seed)
shape = (1, 2)
batch = SimpleNamespace(
latents=torch.zeros(shape),
dmd_latent_vis_dict={},
)
def _fixed_target(*args: Any, **kwargs: Any) -> torch.Tensor:
del args, kwargs
return torch.tensor([target_idx], dtype=torch.long)
monkeypatch.setattr(torch, "randint", _fixed_target)
output = method._student_rollout(batch, with_grad=False)
# The previous implementation drew the initial latent and one noise tensor
# for each of the three possible prefix transitions. Keep consuming those
# draws so subsequent DMD2 randomness remains aligned across the change.
reference_output, reference_state = _legacy_full_rollout_reference(
seed=seed,
target_idx=target_idx,
shape=shape,
)
assert torch.equal(output, reference_output)
assert torch.equal(
method.cuda_generator.get_state(),
reference_state,
)
@pytest.mark.parametrize("target_idx", range(4))
def test_simulate_rollout_uses_global_max_prefix_without_changing_local_output(
monkeypatch: pytest.MonkeyPatch,
target_idx: int,
) -> None:
seed = 9876
method, student = _rollout_method(seed)
shape = (1, 2)
batch = SimpleNamespace(
latents=torch.zeros(shape),
dmd_latent_vis_dict={},
)
monkeypatch.setattr(
torch,
"randint",
lambda *args, **kwargs: torch.tensor([target_idx], dtype=torch.long),
)
monkeypatch.setattr(
method,
"_max_rollout_target_idx_across_ranks",
lambda sampled_idx: 3,
)
output = method._student_rollout(batch, with_grad=False)
reference_output, reference_state = _legacy_full_rollout_reference(
seed=seed,
target_idx=target_idx,
shape=shape,
)
# Every rank participates in the globally required three prefix forwards,
# then evaluates its own target. The local result and method-owned RNG
# sequence still match the legacy full rollout exactly.
assert len(student.predict_calls) == 4
assert len(student.add_noise_calls) == 3
assert torch.equal(output, reference_output)
assert torch.equal(method.cuda_generator.get_state(), reference_state)
def test_dmd2_role_update_cadence_follows_selected_optimizers() -> None:
method = object.__new__(DMD2Method)
torch.nn.Module.__init__(method)
method.method_config = {"generator_update_interval": 5}
method._student_optimizer = object()
method._critic_optimizer = object()
assert not method.did_update_role("student", iteration=4)
assert method.did_update_role("student", iteration=5)
assert method.did_update_role("critic", iteration=4)
assert not method.did_update_role("teacher", iteration=5)
@@ -0,0 +1,61 @@
# SPDX-License-Identifier: Apache-2.0
"""CPU-only tests for modular training profiler lifecycle ownership."""
from __future__ import annotations
from contextlib import contextmanager
from types import SimpleNamespace
from fastvideo.train.entrypoint.train import run_training_from_config
class _RecordingProfiler:
def __init__(self) -> None:
self.events: list[tuple[str, str] | tuple[str]] = []
@contextmanager
def region(self, name: str):
self.events.append(("enter", name))
try:
yield
finally:
self.events.append(("exit", name))
def stop(self) -> None:
self.events.append(("stop", ))
def test_dry_run_profiles_model_build_and_flushes(monkeypatch) -> None:
profiler = _RecordingProfiler()
training = SimpleNamespace(
distributed=SimpleNamespace(tp_size=1, sp_size=1),
vsa_sparsity=0.0,
model_path="model",
)
cfg = SimpleNamespace(training=training)
monkeypatch.setattr(
"fastvideo.train.utils.config.load_run_config",
lambda *args, **kwargs: cfg,
)
monkeypatch.setattr(
"fastvideo.distributed.maybe_init_distributed_environment_and_model_parallel",
lambda *args, **kwargs: None,
)
monkeypatch.setattr(
"fastvideo.train.utils.builder.build_from_config",
lambda loaded_cfg: (loaded_cfg.training, object(), object(), 0),
)
monkeypatch.setattr(
"fastvideo.train.entrypoint.train.get_or_create_profiler",
lambda trace_dir: profiler,
)
run_training_from_config("unused.yaml", dry_run=True)
assert profiler.events == [
("enter", "profiler_region_model_loading"),
("exit", "profiler_region_model_loading"),
("stop", ),
]
@@ -0,0 +1,137 @@
# SPDX-License-Identifier: Apache-2.0
"""CPU-only tests for modular Trainer profiler boundaries."""
from __future__ import annotations
from contextlib import contextmanager
from types import SimpleNamespace
from typing import Any
import torch
from fastvideo.profiler import list_profiler_regions
from fastvideo.train.trainer import Trainer
from fastvideo.train.utils.training_config import TrainingConfig
class _RecordingProfiler:
def __init__(self) -> None:
self.events: list[tuple[str, str]] = []
@property
def has_profiler(self) -> bool:
return True
@contextmanager
def region(self, name: str):
self.events.append(("enter", name))
try:
yield
finally:
self.events.append(("exit", name))
class _DummyTracker:
def log(self, metrics: dict[str, float], step: int) -> None:
del metrics, step
def finish(self) -> None:
pass
class _DummyMethod:
def __init__(self) -> None:
self.weight = torch.nn.Parameter(torch.tensor(1.0))
def set_tracker(self, tracker: Any) -> None:
del tracker
def on_train_start(self) -> None:
pass
def manages_optimization(self) -> bool:
return False
def single_train_step(
self,
batch: dict[str, Any],
iteration: int,
) -> tuple[dict[str, torch.Tensor], dict[str, Any], dict[str, float]]:
del batch, iteration
return {"total_loss": self.weight.square()}, {}, {}
def backward(
self,
loss_map: dict[str, torch.Tensor],
outputs: dict[str, Any],
*,
grad_accum_rounds: int,
) -> None:
del outputs
(loss_map["total_loss"] / grad_accum_rounds).backward()
def optimizers_schedulers_step(self, iteration: int) -> None:
del iteration
def optimizers_zero_grad(self, iteration: int) -> None:
del iteration
self.weight.grad = None
def test_modular_trainer_emits_nested_step_regions(monkeypatch) -> None:
group = SimpleNamespace(rank=0, local_rank=0, rank_in_group=0, world_size=1)
profiler = _RecordingProfiler()
monkeypatch.setattr("fastvideo.train.trainer.get_world_group", lambda: group)
monkeypatch.setattr("fastvideo.train.trainer.get_sp_group", lambda: group)
monkeypatch.setattr(
"fastvideo.train.trainer.build_tracker",
lambda *args, **kwargs: _DummyTracker(),
)
monkeypatch.setattr("fastvideo.profiler._GLOBAL_CONTROLLER", profiler)
trainer = Trainer(TrainingConfig())
trainer.run(
_DummyMethod(),
dataloader=[{}],
max_steps=1,
)
assert profiler.events == [
("enter", "profiler_region_training_train"),
("enter", "profiler_region_training_train_one_step"),
("enter", "profiler_region_training_dataloader"),
("exit", "profiler_region_training_dataloader"),
("enter", "profiler_region_training_forward"),
("exit", "profiler_region_training_forward"),
("enter", "profiler_region_training_backward"),
("exit", "profiler_region_training_backward"),
("enter", "profiler_region_training_optimizer"),
("exit", "profiler_region_training_optimizer"),
("enter", "profiler_region_training_callbacks"),
("exit", "profiler_region_training_callbacks"),
("exit", "profiler_region_training_train_one_step"),
("exit", "profiler_region_training_train"),
]
def test_modular_training_regions_are_registered() -> None:
names = {region.name for region in list_profiler_regions()}
assert {
"profiler_region_model_loading",
"profiler_region_training_train",
"profiler_region_training_train_one_step",
"profiler_region_training_dataloader",
"profiler_region_training_forward",
"profiler_region_training_backward",
"profiler_region_training_optimizer",
"profiler_region_training_callbacks",
"profiler_region_training_save_checkpoint",
"profiler_region_training_validation",
"profiler_region_dmd2_student_rollout",
"profiler_region_dmd2_generator_loss",
"profiler_region_dmd2_critic_loss",
} <= names
@@ -136,4 +136,5 @@ def test_trainer_runs_validation_callback_during_training(monkeypatch, ) -> None
assert method.zero_grad_steps == [0, 1, 2, 3]
assert method.optimizer_steps == [1, 2, 3]
assert [step for _, step in tracker.logs] == [1, 2, 3]
assert all("dataloader_time_sec" in metrics for metrics, _ in tracker.logs)
assert tracker.finished is True
@@ -1,231 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Focused routing and FA4 integration checks for MiniMax-H3 fusions."""
from __future__ import annotations
import pytest
import torch
from fastvideo.platforms import AttentionBackendEnum
@pytest.mark.parametrize(
("raw", "expected"),
[
("", frozenset()),
("0", frozenset()),
("none", frozenset()),
("all", frozenset({"modulate", "qknorm_rope", "swiglu"})),
("1", frozenset({"modulate", "qknorm_rope", "swiglu"})),
("swiglu, modulate", frozenset({"swiglu", "modulate"})),
],
)
def test_minimax_h3_fusion_selector(raw: str, expected: frozenset[str]) -> None:
from fastvideo.models.dits.minimax_h3 import _enabled_minimax_h3_fusions
assert _enabled_minimax_h3_fusions(raw) == expected
def test_minimax_h3_fusion_selector_rejects_unknown_name() -> None:
from fastvideo.models.dits.minimax_h3 import _enabled_minimax_h3_fusions
with pytest.raises(ValueError, match="Unknown MiniMax H3 fusion"):
_enabled_minimax_h3_fusions("swiglu,unknown")
def test_swiglu_fusion_stays_on_eager_path_with_grad(monkeypatch: pytest.MonkeyPatch) -> None:
import fastvideo.models.dits.minimax_h3 as h3
def unexpected_kernel(_: torch.Tensor) -> torch.Tensor:
raise AssertionError("inference-only fusion ran with grad enabled")
monkeypatch.setattr(h3, "minimax_h3_swiglu", unexpected_kernel)
layer = h3.MiniMaxH3FeedForward(8, 16, fuse_swiglu=True)
inputs = torch.randn(2, 3, 8, requires_grad=True)
layer(inputs).sum().backward()
assert inputs.grad is not None
def test_all_minimax_h3_fusions_match_one_eager_block_under_fa4(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""Exercise the real block wiring without loading any H3 checkpoint."""
if not torch.cuda.is_available() or not torch.cuda.is_bf16_supported():
pytest.skip("BF16 CUDA is required")
pytest.importorskip("triton")
flash_attn = pytest.importorskip("flash_attn")
if "fa4" not in getattr(flash_attn, "__version__", "").lower():
pytest.skip("the focused integration test requires the FA4 environment")
monkeypatch.setenv("FASTVIDEO_ATTENTION_BACKEND", "FLASH_ATTN")
monkeypatch.setenv("FASTVIDEO_FA4", "1")
monkeypatch.setenv("MASTER_ADDR", "127.0.0.1")
monkeypatch.setenv("MASTER_PORT", "29573")
monkeypatch.setenv("RANK", "0")
monkeypatch.setenv("WORLD_SIZE", "1")
monkeypatch.setenv("LOCAL_RANK", "0")
from fastvideo.distributed import cleanup_dist_env_and_memory, maybe_init_distributed_environment_and_model_parallel
from fastvideo.forward_context import set_forward_context
from fastvideo.models.dits.minimax_h3 import MiniMaxH3RotaryPosEmbed, MiniMaxH3TransformerBlock
maybe_init_distributed_environment_and_model_parallel(1, 1)
try:
kwargs = dict(
hidden_size=128,
num_attention_heads=1,
attention_head_dim=128,
ffn_dim=256,
time_embed_dim=64,
norm_eps=1e-5,
qk_norm_eps=1e-5,
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN, ),
quant_config=None,
prefix="minimax_h3.test_block",
)
previous_default_dtype = torch.get_default_dtype()
torch.set_default_dtype(torch.bfloat16)
try:
eager = MiniMaxH3TransformerBlock(**kwargs)
fused = MiniMaxH3TransformerBlock(
**kwargs,
fuse_modulate=True,
fuse_qknorm_rope=True,
fuse_swiglu=True,
)
finally:
torch.set_default_dtype(previous_default_dtype)
with torch.no_grad():
for name, parameter in eager.named_parameters():
if "norm" in name and name.endswith("weight"):
parameter.fill_(1.0)
elif parameter.ndim > 1:
torch.nn.init.normal_(parameter, mean=0.0, std=0.02)
else:
parameter.zero_()
fused.load_state_dict(eager.state_dict(), strict=True)
device = torch.device("cuda")
eager = eager.to(device=device, dtype=torch.bfloat16).eval()
fused = fused.to(device=device, dtype=torch.bfloat16).eval()
generator = torch.Generator(device=device).manual_seed(2026)
hidden_states = torch.randn(2, 12, 128, generator=generator, device=device, dtype=torch.bfloat16)
temb = torch.randn(2, 64, generator=generator, device=device, dtype=torch.bfloat16)
adaln_indices = torch.arange(12, device=device, dtype=torch.long).remainder(6)
position_ids = torch.zeros(12, 3, device=device, dtype=torch.float32)
position_ids[:, 0] = torch.arange(12, device=device)
rotary_emb = MiniMaxH3RotaryPosEmbed(rope_freq_dim=16, rope_theta=10000.0).to(device)(position_ids)
inputs = dict(
hidden_states=hidden_states,
temb=temb,
adaln_indices=adaln_indices,
rotary_emb=tuple(value.to(torch.bfloat16) for value in rotary_emb),
original_seq_len=12,
)
with torch.inference_mode(), set_forward_context(current_timestep=0, attn_metadata=None):
eager_output = eager(**inputs)
fused_output = fused(**inputs)
# Sol-Engine keeps fused intermediates in FP32 registers until their
# final BF16 stores, so the opt-in path is close but not bit-identical.
torch.testing.assert_close(fused_output, eager_output, atol=3e-2, rtol=3e-2)
finally:
cleanup_dist_env_and_memory()
def test_minimax_h3_fusions_engage_on_cuda_inference(monkeypatch: pytest.MonkeyPatch) -> None:
"""Pin the positive side of the routing guard.
The parity test above still passes if ``_can_run_minimax_h3_fusion``
silently degrades to always-False (both blocks then run the identical
eager path), so count the fused-kernel calls: one CUDA inference forward
through a fully fused block must hit ``fused_rmsnorm_modulate`` once,
``fused_residual_gate_rmsnorm_modulate`` once, ``fused_qknorm_rope``
twice (q and k), and ``minimax_h3_swiglu`` once -- and a grad-enabled
forward must leave every counter unchanged.
"""
if not torch.cuda.is_available() or not torch.cuda.is_bf16_supported():
pytest.skip("BF16 CUDA is required")
pytest.importorskip("triton")
monkeypatch.setenv("FASTVIDEO_ATTENTION_BACKEND", "TORCH_SDPA")
monkeypatch.setenv("MASTER_ADDR", "127.0.0.1")
monkeypatch.setenv("MASTER_PORT", "29574")
monkeypatch.setenv("RANK", "0")
monkeypatch.setenv("WORLD_SIZE", "1")
monkeypatch.setenv("LOCAL_RANK", "0")
import fastvideo.models.dits.minimax_h3 as h3
from fastvideo.distributed import cleanup_dist_env_and_memory, maybe_init_distributed_environment_and_model_parallel
from fastvideo.forward_context import set_forward_context
calls = dict.fromkeys(("rmsnorm_modulate", "residual_gate_rmsnorm_modulate", "qknorm_rope", "swiglu"), 0)
def _counting(name: str, real):
def wrapper(*args, **kwargs):
calls[name] += 1
return real(*args, **kwargs)
return wrapper
monkeypatch.setattr(h3, "fused_rmsnorm_modulate", _counting("rmsnorm_modulate", h3.fused_rmsnorm_modulate))
monkeypatch.setattr(h3, "fused_residual_gate_rmsnorm_modulate",
_counting("residual_gate_rmsnorm_modulate", h3.fused_residual_gate_rmsnorm_modulate))
monkeypatch.setattr(h3, "fused_qknorm_rope", _counting("qknorm_rope", h3.fused_qknorm_rope))
monkeypatch.setattr(h3, "minimax_h3_swiglu", _counting("swiglu", h3.minimax_h3_swiglu))
maybe_init_distributed_environment_and_model_parallel(1, 1)
try:
block = h3.MiniMaxH3TransformerBlock(
hidden_size=128,
num_attention_heads=1,
attention_head_dim=128,
ffn_dim=256,
time_embed_dim=64,
norm_eps=1e-5,
qk_norm_eps=1e-5,
supported_attention_backends=(AttentionBackendEnum.TORCH_SDPA, ),
quant_config=None,
prefix="minimax_h3.engagement_block",
fuse_modulate=True,
fuse_qknorm_rope=True,
fuse_swiglu=True,
)
device = torch.device("cuda")
block = block.to(device=device, dtype=torch.bfloat16).eval()
generator = torch.Generator(device=device).manual_seed(2026)
hidden_states = torch.randn(2, 12, 128, generator=generator, device=device, dtype=torch.bfloat16)
temb = torch.randn(2, 64, generator=generator, device=device, dtype=torch.bfloat16)
adaln_indices = torch.arange(12, device=device, dtype=torch.long).remainder(6)
position_ids = torch.zeros(12, 3, device=device, dtype=torch.float32)
position_ids[:, 0] = torch.arange(12, device=device)
rotary_emb = h3.MiniMaxH3RotaryPosEmbed(rope_freq_dim=16, rope_theta=10000.0).to(device)(position_ids)
inputs = dict(
hidden_states=hidden_states,
temb=temb,
adaln_indices=adaln_indices,
rotary_emb=tuple(value.to(torch.bfloat16) for value in rotary_emb),
original_seq_len=12,
)
with torch.inference_mode(), set_forward_context(current_timestep=0, attn_metadata=None):
block(**inputs)
engaged = dict(calls)
assert engaged == {
"rmsnorm_modulate": 1,
"residual_gate_rmsnorm_modulate": 1,
"qknorm_rope": 2,
"swiglu": 1,
}, engaged
grad_inputs = {**inputs, "hidden_states": hidden_states.clone().requires_grad_(True)}
with set_forward_context(current_timestep=0, attn_metadata=None):
block(**grad_inputs)
assert dict(calls) == engaged, f"a fusion ran under grad: {calls} vs {engaged}"
finally:
cleanup_dist_env_and_memory()
@@ -1,173 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
import pytest
import torch
import torch.nn.functional as F
from fastvideo.models.dits.minimax_h3_fusions.modulation import (
fused_residual_gate_rmsnorm_modulate,
fused_rmsnorm_modulate,
)
EPS = 1e-6
SOL_ENGINE_BF16_TOLERANCE = 3e-2
def _chunk_tables(rows: int, hidden_size: int, *, device: torch.device | str = "cpu", dtype=torch.float32):
wide = torch.randn(rows, 6 * hidden_size, device=device, dtype=dtype)
tables = wide.chunk(6, dim=-1)
assert all(table.stride() == (6 * hidden_size, 1) for table in tables)
assert all(not table.is_contiguous() for table in tables)
return tables
def _eager_rmsnorm_modulate(x, weight, scale, shift, index):
normed = F.rms_norm(x, (x.shape[-1], ), weight, EPS)
return normed * (1.0 + scale.index_select(0, index)) + shift.index_select(0, index)
def _eager_residual_gate_rmsnorm_modulate(residual, branch, gate, weight, scale, shift, index):
hidden = residual + gate.index_select(0, index) * branch
normed = F.rms_norm(hidden, (hidden.shape[-1], ), weight, EPS)
modulated = normed * (1.0 + scale.index_select(0, index)) + shift.index_select(0, index)
return hidden, modulated
@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires a CUDA GPU")
def test_bf16_sol_engine_fusions_match_production_eager_within_tolerance():
"""Sol-Engine keeps fused intermediates in FP32 until its output stores."""
pytest.importorskip("triton")
torch.manual_seed(2)
device = torch.device("cuda")
batch, sequence_length, hidden_size, table_rows = 2, 9, 5376, 6
residual = torch.randn(batch, sequence_length, hidden_size, device=device, dtype=torch.bfloat16)
branch = torch.randn_like(residual)
weight = torch.randn(hidden_size, device=device, dtype=torch.bfloat16)
shift, scale, gate, shift_mlp, scale_mlp, _ = _chunk_tables(
table_rows,
hidden_size,
device=device,
dtype=torch.bfloat16,
)
index = torch.tensor([5, 0, 4, 1, 3, 2, 5, 1, 0], device=device, dtype=torch.int64)
expected_norm1 = _eager_rmsnorm_modulate(residual, weight, scale, shift, index)
actual_norm1 = fused_rmsnorm_modulate(residual, weight, scale, shift, index, EPS)
expected_hidden, expected_norm2 = _eager_residual_gate_rmsnorm_modulate(
residual,
branch,
gate,
weight,
scale_mlp,
shift_mlp,
index,
)
actual_hidden, actual_norm2 = fused_residual_gate_rmsnorm_modulate(
residual,
branch,
gate,
weight,
scale_mlp,
shift_mlp,
index,
EPS,
)
# Production eager materializes BF16 after each PyTorch operator. The
# single-kernel Sol-Engine path deliberately removes those round points.
torch.testing.assert_close(
actual_norm1,
expected_norm1,
rtol=SOL_ENGINE_BF16_TOLERANCE,
atol=SOL_ENGINE_BF16_TOLERANCE,
)
torch.testing.assert_close(
actual_hidden,
expected_hidden,
rtol=SOL_ENGINE_BF16_TOLERANCE,
atol=SOL_ENGINE_BF16_TOLERANCE,
)
torch.testing.assert_close(
actual_norm2,
expected_norm2,
rtol=SOL_ENGINE_BF16_TOLERANCE,
atol=SOL_ENGINE_BF16_TOLERANCE,
)
@pytest.mark.parametrize(
("mutate", "error", "match"),
[
(lambda args: args | {"x": args["x"][0, 0]}, ValueError, "shape"),
(lambda args: args | {"weight": args["weight"][:-1]}, ValueError, "weight"),
(lambda args: args | {"index": args["index"][:-1]}, ValueError, "index"),
(lambda args: args | {"index": args["index"].float()}, TypeError, "index"),
(lambda args: args | {"eps": 0.0}, ValueError, "eps"),
],
)
def test_fusion_rejects_invalid_contracts(mutate, error, match):
args = {
"x": torch.randn(2, 3, 8),
"weight": torch.randn(8),
"scale": torch.randn(4, 8),
"shift": torch.randn(4, 8),
"index": torch.tensor([0, 3, 1]),
"eps": EPS,
}
with pytest.raises(error, match=match):
fused_rmsnorm_modulate(**mutate(args))
def test_residual_fusion_rejects_mismatched_branch():
residual = torch.randn(2, 3, 8)
with pytest.raises(ValueError, match="branch"):
fused_residual_gate_rmsnorm_modulate(
residual,
torch.randn(2, 2, 8),
torch.randn(4, 8),
torch.randn(8),
torch.randn(4, 8),
torch.randn(4, 8),
torch.tensor([0, 1, 2]),
EPS,
)
def test_triton_wrappers_fail_explicitly_on_cpu():
x = torch.randn(1, 2, 8)
branch = torch.randn_like(x)
weight = torch.randn(8)
gate = torch.randn(3, 8)
scale = torch.randn(3, 8)
shift = torch.randn(3, 8)
index = torch.tensor([0, 2])
with pytest.raises(RuntimeError, match="Triton|CUDA"):
fused_rmsnorm_modulate(x, weight, scale, shift, index, EPS)
with pytest.raises(RuntimeError, match="Triton|CUDA"):
fused_residual_gate_rmsnorm_modulate(x, branch, gate, weight, scale, shift, index, EPS)
def test_triton_wrappers_reject_autograd_before_backend_check():
x = torch.randn(1, 2, 8)
branch = torch.randn_like(x, requires_grad=True)
weight = torch.randn(8, requires_grad=True)
gate = torch.randn(3, 8)
scale = torch.randn(3, 8)
shift = torch.randn(3, 8)
index = torch.tensor([0, 2])
with pytest.raises(RuntimeError, match="forward-only"):
fused_rmsnorm_modulate(x, weight, scale, shift, index, EPS)
with pytest.raises(RuntimeError, match="forward-only"):
fused_residual_gate_rmsnorm_modulate(
x,
branch,
gate,
weight.detach(),
scale,
shift,
index,
EPS,
)
@@ -1,192 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
import pytest
import torch
import torch.nn.functional as F
from fastvideo.models.dits.minimax_h3 import MiniMaxH3Attention
from fastvideo.models.dits.minimax_h3_fusions.qknorm_rope import (
HAVE_TRITON,
fused_qknorm_rope,
)
def _rotary_tables(
seq_len: int,
rotary_dim: int,
*,
dtype: torch.dtype,
device: torch.device | str,
) -> tuple[torch.Tensor, torch.Tensor]:
angles = torch.randn(seq_len, rotary_dim // 2, dtype=torch.float32, device=device)
angles = torch.cat((angles, angles), dim=-1)
return angles.cos().to(dtype), angles.sin().to(dtype)
def _eager_qknorm_rope(
x: torch.Tensor,
weight: torch.Tensor,
cos: torch.Tensor,
sin: torch.Tensor,
eps: float,
) -> torch.Tensor:
normalized = F.rms_norm(x, (x.shape[-1], ), weight, eps)
return MiniMaxH3Attention._apply_rotary_emb(normalized, (cos, sin))
def test_fused_qknorm_rope_rejects_invalid_rotary_dim() -> None:
x = torch.randn(2, 3, 4, 128)
weight = torch.ones(128)
cos, sin = _rotary_tables(3, 130, dtype=x.dtype, device=x.device)
with pytest.raises(ValueError, match="rotary_dim must not exceed head_dim"):
fused_qknorm_rope(x, weight, cos, sin, 1e-6)
cos = torch.randn(3, 95)
sin = torch.randn_like(cos)
with pytest.raises(ValueError, match="rotary_dim must be even"):
fused_qknorm_rope(x, weight, cos, sin, 1e-6)
def test_fused_qknorm_rope_rejects_shape_mismatches() -> None:
x = torch.randn(2, 3, 4, 128)
weight = torch.ones(128)
cos, sin = _rotary_tables(3, 96, dtype=x.dtype, device=x.device)
with pytest.raises(ValueError, match="x must have shape"):
fused_qknorm_rope(x[0], weight, cos, sin, 1e-6)
with pytest.raises(ValueError, match="weight must have shape"):
fused_qknorm_rope(x, weight[:-1], cos, sin, 1e-6)
with pytest.raises(ValueError, match="sequence length"):
fused_qknorm_rope(x, weight, cos[:-1], sin[:-1], 1e-6)
with pytest.raises(ValueError, match="sin must match cos shape"):
fused_qknorm_rope(x, weight, cos, sin[:, :-2], 1e-6)
@pytest.mark.parametrize("noncontiguous_input", ["weight", "cos", "sin"])
def test_fused_qknorm_rope_rejects_noncontiguous_linear_inputs(noncontiguous_input: str) -> None:
x = torch.randn(2, 3, 4, 128)
weight = torch.ones(256)[::2]
cos = torch.randn(3, 192)[:, ::2]
sin = torch.randn(3, 192)[:, ::2]
assert not weight.is_contiguous()
assert not cos.is_contiguous()
assert not sin.is_contiguous()
if noncontiguous_input != "weight":
weight = weight.contiguous()
if noncontiguous_input != "cos":
cos = cos.contiguous()
if noncontiguous_input != "sin":
sin = sin.contiguous()
with pytest.raises(ValueError, match="must be contiguous"):
fused_qknorm_rope(x, weight, cos, sin, 1e-6)
def test_fused_qknorm_rope_requires_matching_precast_dtype() -> None:
x = torch.randn(2, 3, 4, 128, dtype=torch.float32)
weight = torch.ones(128, dtype=torch.float32)
cos, sin = _rotary_tables(3, 96, dtype=torch.bfloat16, device=x.device)
with pytest.raises(TypeError, match="cos dtype must match x dtype"):
fused_qknorm_rope(x, weight, cos, sin, 1e-6)
def test_fused_qknorm_rope_does_not_accept_missing_rotary_tables() -> None:
x = torch.randn(2, 3, 4, 128)
weight = torch.ones(128)
with pytest.raises(TypeError, match="cos must be a torch.Tensor"):
fused_qknorm_rope(x, weight, None, None, 1e-6) # type: ignore[arg-type]
def test_fused_qknorm_rope_requires_cuda() -> None:
x = torch.randn(2, 3, 4, 128)
weight = torch.ones(128)
cos, sin = _rotary_tables(3, 96, dtype=x.dtype, device=x.device)
with pytest.raises(RuntimeError, match="requires CUDA"):
fused_qknorm_rope(x, weight, cos, sin, 1e-6)
@pytest.mark.parametrize(
"rotary_dim,shape,use_input_view",
[
pytest.param(96, (2, 11, 5, 128), False, id="partial-96-batch2-seq11-heads5"),
pytest.param(128, (2, 7, 3, 128), True, id="full-128-noncontiguous-input-view"),
],
)
def test_fused_qknorm_rope_matches_eager_bf16_cuda(
rotary_dim: int,
shape: tuple[int, ...],
use_input_view: bool,
) -> None:
if not torch.cuda.is_available():
pytest.skip("CUDA is required for the Triton fusion")
if not HAVE_TRITON:
pytest.skip("Triton is required for the fusion")
torch.manual_seed(1)
device = torch.device("cuda")
if use_input_view:
x = torch.randn(*shape[:-1], shape[-1] * 2, dtype=torch.bfloat16, device=device)[..., ::2]
assert not x.is_contiguous()
else:
x = torch.randn(shape, dtype=torch.bfloat16, device=device)
weight = (1.0 + 0.05 * torch.randn(shape[-1], dtype=torch.bfloat16, device=device)).contiguous()
cos, sin = _rotary_tables(shape[1], rotary_dim, dtype=x.dtype, device=device)
with torch.inference_mode():
actual = fused_qknorm_rope(x, weight, cos, sin, 1e-6)
expected = _eager_qknorm_rope(x, weight, cos, sin, 1e-6)
# The fused kernel keeps RMSNorm and both RoPE products in FP32 registers
# until its final BF16 store. Eager materializes BF16 intermediates, and
# PyTorch/Triton reductions need not use the same summation order.
torch.testing.assert_close(actual, expected, atol=2e-2, rtol=2e-2)
assert actual.shape == x.shape
assert actual.dtype == x.dtype
assert actual.is_contiguous()
def test_fused_qknorm_rope_matches_eager_beyond_int32_element_count() -> None:
"""Regression: kernel row offsets must be int64.
With int32 offsets, ``row * head_dim`` wraps once the flattened input
crosses 2**31 elements and the kernel reads/writes out of bounds (CUDA
illegal memory access). ``(1, 8_500_000, 2, 128)`` is 2.176e9 elements,
just past the boundary; for H3's 56 heads x 128 head_dim the equivalent
is ``batch*seq >= 299_593`` tokens per rank, reachable at SP=1.
GPU assumption: needs ~16 GiB free CUDA memory (input + output at
bf16 plus the fp32 rotary-table construction); skips below 20 GiB.
"""
if not torch.cuda.is_available():
pytest.skip("CUDA is required for the Triton fusion")
if not HAVE_TRITON:
pytest.skip("Triton is required for the fusion")
free_bytes, _ = torch.cuda.mem_get_info()
if free_bytes < 20 * 1024**3:
pytest.skip("needs ~20 GiB free GPU memory for a >2**31-element input")
heads, head_dim, rotary_dim = 2, 128, 96
seq_len = 8_500_000
assert seq_len * heads * head_dim > 2**31
torch.manual_seed(9)
device = torch.device("cuda")
x = torch.randn(1, seq_len, heads, head_dim, dtype=torch.bfloat16, device=device)
weight = (1.0 + 0.05 * torch.randn(head_dim, dtype=torch.bfloat16, device=device)).contiguous()
cos, sin = _rotary_tables(seq_len, rotary_dim, dtype=x.dtype, device=device)
with torch.inference_mode():
fused = fused_qknorm_rope(x, weight, cos, sin, 1e-6)
torch.cuda.synchronize()
# Compare only head/tail slices against eager: a full-tensor eager
# reference would double peak memory for no extra coverage, and the
# tail rows are exactly the ones an int32 wrap corrupts first.
expected_head = _eager_qknorm_rope(x[:, :8], weight, cos[:8], sin[:8], 1e-6)
expected_tail = _eager_qknorm_rope(x[:, -8:], weight, cos[-8:], sin[-8:], 1e-6)
torch.testing.assert_close(fused[:, :8], expected_head, atol=2e-2, rtol=2e-2)
torch.testing.assert_close(fused[:, -8:], expected_tail, atol=2e-2, rtol=2e-2)
@@ -1,73 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Focused tests for MiniMax H3's value-first packed SwiGLU fusion."""
from __future__ import annotations
import pytest
import torch
import torch.nn.functional as F
from fastvideo.models.dits.minimax_h3_fusions.swiglu import minimax_h3_swiglu
def _require_bf16_triton_cuda() -> None:
if not torch.cuda.is_available():
pytest.skip("MiniMax H3 fused SwiGLU requires CUDA")
if not torch.cuda.is_bf16_supported():
pytest.skip("MiniMax H3 fused SwiGLU parity requires BF16 support")
pytest.importorskip("triton", reason="MiniMax H3 fused SwiGLU requires Triton")
def _assert_bf16_parity(x: torch.Tensor) -> None:
value, gate = x.chunk(2, dim=-1)
expected = value * F.silu(gate)
actual = minimax_h3_swiglu(x)
assert actual.shape == (*x.shape[:-1], x.shape[-1] // 2)
assert actual.dtype == x.dtype
assert actual.device == x.device
# Sol-Engine keeps the full SwiGLU expression in FP32 until the output
# store, while eager F.silu materializes a BF16 intermediate.
torch.testing.assert_close(actual, expected, rtol=2e-2, atol=2e-2)
@pytest.mark.parametrize("last_dim", [0, 7])
def test_minimax_h3_swiglu_rejects_nonpositive_or_odd_last_dimension(last_dim: int) -> None:
x = torch.empty((2, last_dim), dtype=torch.float32)
with pytest.raises(ValueError, match="positive even last dimension"):
minimax_h3_swiglu(x)
def test_minimax_h3_swiglu_strict_wrapper_rejects_cpu() -> None:
with pytest.raises(ValueError, match="requires a CUDA tensor"):
minimax_h3_swiglu(torch.randn(2, 8))
@pytest.mark.gpu
def test_minimax_h3_swiglu_multidimensional_bf16_gpu_parity() -> None:
_require_bf16_triton_cuda()
torch.manual_seed(0)
x = torch.randn((2, 3, 5, 66), device="cuda", dtype=torch.bfloat16)
_assert_bf16_parity(x)
@pytest.mark.gpu
def test_minimax_h3_swiglu_noncontiguous_bf16_gpu_parity() -> None:
_require_bf16_triton_cuda()
torch.manual_seed(1)
storage = torch.randn((2, 3, 148), device="cuda", dtype=torch.bfloat16)
x = storage[..., ::2]
assert not x.is_contiguous()
_assert_bf16_parity(x)
@pytest.mark.gpu
def test_minimax_h3_swiglu_real_ffn_dim_bf16_gpu() -> None:
_require_bf16_triton_cuda()
torch.manual_seed(2)
ffn_dim = 14336
x = torch.randn((1, 2 * ffn_dim), device="cuda", dtype=torch.bfloat16)
_assert_bf16_parity(x)
@@ -1,314 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""CPU tests for sequence-parallel MiniMax-H3 VAE chunk decode / clip encode.
The collective transport is simulated with a threaded fake group (one thread
per simulated rank, barrier-synchronized slots), so the REAL drivers in
``fastvideo.models.vaes.minimax_h3_parallel`` — chunk assignment, placeholder
rounds, metadata broadcast, gathered-segment assembly, halo/blend math — run
end to end on CPU and are checked bit-exactly against the serial APIs.
"""
import threading
import pytest
import torch
from torch.testing import assert_close
from fastvideo.configs.models.vaes.minimax_h3_video import (
MiniMaxH3VideoVAEArchConfig,
MiniMaxH3VideoVAEConfig,
)
from fastvideo.models.vaes.minimax_h3_parallel import (
DECODE_GATHER_STRATEGIES,
DEFAULT_DECODE_GATHER_STRATEGY,
decode_to_pixels_parallel,
encode_pixels_parallel,
parallel_chunk_indices,
)
from fastvideo.models.vaes.minimax_h3_video import AutoencoderKLMiniMaxH3
def _tiny_vae(token_drop: int = 3) -> AutoencoderKLMiniMaxH3:
arch = MiniMaxH3VideoVAEArchConfig(
latent_channels=4,
block_out_channels=(32, 32),
layers_per_block=1,
spatial_downsample_factors=(2, 2),
temporal_downsample_factors=(2, 2),
decoder_num_layers=1,
decoder_num_attention_heads=1,
decoder_attention_head_dim=8,
decoder_num_register_tokens=2,
decoder_ffn_mult=1,
token_drop=token_drop,
latents_mean=(0.0, ) * 4,
latents_std=(1.0, ) * 4,
)
return AutoencoderKLMiniMaxH3(
MiniMaxH3VideoVAEConfig(
arch_config=arch,
use_tiling=False,
use_temporal_tiling=False,
use_parallel_tiling=False,
)).eval()
class _ThreadedFakeGroup:
"""Barrier-synchronized in-process stand-in for a GroupCoordinator.
One thread per simulated rank runs the SPMD driver; ``gather`` /
``all_gather`` / ``broadcast_object`` rendezvous through shared slots
with a double barrier (all writes land, everyone reads, then slots are
reusable). Matches the GroupCoordinator call signatures the drivers use.
"""
def __init__(self, world_size: int) -> None:
self.world_size = world_size
self._local = threading.local()
self._barrier = threading.Barrier(world_size)
self._slots: list = [None] * world_size
self._object = None
@property
def rank_in_group(self) -> int:
return self._local.rank
def broadcast_object(self, obj=None, src: int = 0):
if self.world_size == 1:
return obj
if self.rank_in_group == src:
self._object = obj
self._barrier.wait()
received = self._object
self._barrier.wait()
return received
def all_gather(self, input_: torch.Tensor, dim: int = -1) -> torch.Tensor:
if self.world_size == 1:
return input_
self._slots[self.rank_in_group] = input_
self._barrier.wait()
gathered = torch.cat([slot for slot in self._slots], dim=dim)
self._barrier.wait()
return gathered
def gather(self, input_: torch.Tensor, dst: int = 0, dim: int = -1):
if self.world_size == 1:
return input_
self._slots[self.rank_in_group] = input_
self._barrier.wait()
gathered = torch.cat([slot for slot in self._slots], dim=dim) if self.rank_in_group == dst else None
self._barrier.wait()
return gathered
def run(self, fn) -> list:
"""Run ``fn(rank)`` on one thread per rank; re-raise the first error."""
results: list = [None] * self.world_size
errors: list = [None] * self.world_size
def _target(rank: int) -> None:
self._local.rank = rank
try:
# inference_mode is thread-local; the drivers run inference-only.
with torch.inference_mode():
results[rank] = fn(rank)
except BaseException as error: # noqa: BLE001 - propagate to the test
errors[rank] = error
self._barrier.abort()
threads = [threading.Thread(target=_target, args=(rank, )) for rank in range(self.world_size)]
for thread in threads:
thread.start()
for thread in threads:
thread.join()
for error in errors:
if error is not None and not isinstance(error, threading.BrokenBarrierError):
raise error
for error in errors:
if error is not None:
raise error
return results
@pytest.mark.parametrize("num_chunks,world_size", ((0, 4), (1, 4), (7, 4), (8, 4), (20, 4), (5, 3), (2, 5)))
def test_parallel_chunk_indices_partition(num_chunks: int, world_size: int) -> None:
"""Round-robin ownership covers every chunk exactly once, in order."""
owned = [parallel_chunk_indices(num_chunks, world_size, rank) for rank in range(world_size)]
flattened = sorted(index for indices in owned for index in indices)
assert flattened == list(range(num_chunks))
for rank, indices in enumerate(owned):
assert indices == sorted(indices)
assert all(index % world_size == rank for index in indices)
# Round-robin balance: no rank holds more than one extra chunk.
assert len(indices) in (num_chunks // world_size, -(-num_chunks // world_size))
def test_parallel_chunk_indices_validates() -> None:
with pytest.raises(ValueError, match="world_size"):
parallel_chunk_indices(4, 0, 0)
with pytest.raises(ValueError, match="rank_in_group"):
parallel_chunk_indices(4, 2, 2)
with pytest.raises(ValueError, match="num_chunks"):
parallel_chunk_indices(-1, 2, 0)
# Latent frames cover: one padded chunk (3), pad on the intra-clip tail (6),
# two blended chunks (12), three chunks plus pad trim (13). World sizes cover
# fewer chunks than ranks, uneven rounds, and the exact-multiple case.
@pytest.mark.parametrize("world_size", (2, 3, 4, 5))
@pytest.mark.parametrize("latent_frames", (3, 6, 12, 13))
@pytest.mark.parametrize("strategy", DECODE_GATHER_STRATEGIES)
@torch.inference_mode()
def test_parallel_decode_matches_serial(world_size: int, latent_frames: int, strategy: str) -> None:
torch.manual_seed(20260821 + latent_frames)
vae = _tiny_vae()
latents = torch.randn(1, 4, latent_frames, 4, 4)
expected = torch.empty(vae.decoded_pixel_shape(latents.shape), dtype=torch.float32)
vae.decode_to_pixels(latents, expected)
group = _ThreadedFakeGroup(world_size)
def _rank_main(rank: int):
output = torch.empty(vae.decoded_pixel_shape(latents.shape), dtype=torch.float32) if rank == 0 else None
return decode_to_pixels_parallel(vae, latents.clone(), output, group, strategy=strategy)
results = group.run(_rank_main)
assert all(result is None for result in results[1:])
assert_close(results[0], expected, atol=0.0, rtol=0.0)
@torch.inference_mode()
def test_parallel_decode_without_token_drop() -> None:
"""token_drop == 0 has no overlap halo; the assembler must skip blending."""
torch.manual_seed(20260822)
vae = _tiny_vae(token_drop=0)
assert vae.frame_overlap == 0
latents = torch.randn(1, 4, 10, 4, 4)
expected = torch.empty(vae.decoded_pixel_shape(latents.shape), dtype=torch.float32)
vae.decode_to_pixels(latents, expected)
group = _ThreadedFakeGroup(3)
def _rank_main(rank: int):
output = torch.empty(vae.decoded_pixel_shape(latents.shape), dtype=torch.float32) if rank == 0 else None
return decode_to_pixels_parallel(vae, latents.clone(), output, group)
results = group.run(_rank_main)
assert_close(results[0], expected, atol=0.0, rtol=0.0)
@torch.inference_mode()
def test_parallel_decode_batched_slicing_matches_serial() -> None:
torch.manual_seed(20260823)
vae = _tiny_vae()
vae.enable_slicing()
latents = torch.randn(2, 4, 7, 4, 4)
expected = torch.empty(vae.decoded_pixel_shape(latents.shape), dtype=torch.float32)
vae.decode_to_pixels(latents, expected)
group = _ThreadedFakeGroup(2)
def _rank_main(rank: int):
output = torch.empty(vae.decoded_pixel_shape(latents.shape), dtype=torch.float32) if rank == 0 else None
return decode_to_pixels_parallel(vae, latents.clone(), output, group)
results = group.run(_rank_main)
assert_close(results[0], expected, atol=0.0, rtol=0.0)
@torch.inference_mode()
def test_parallel_decode_world_size_one_is_serial() -> None:
torch.manual_seed(20260824)
vae = _tiny_vae()
latents = torch.randn(1, 4, 7, 4, 4)
expected = torch.empty(vae.decoded_pixel_shape(latents.shape), dtype=torch.float32)
vae.decode_to_pixels(latents, expected)
group = _ThreadedFakeGroup(1)
def _rank_main(rank: int):
output = torch.empty(vae.decoded_pixel_shape(latents.shape), dtype=torch.float32)
return decode_to_pixels_parallel(vae, latents, output, group)
results = group.run(_rank_main)
assert_close(results[0], expected, atol=0.0, rtol=0.0)
def test_parallel_decode_validates_buffers_and_strategy() -> None:
vae = _tiny_vae()
latents = torch.randn(1, 4, 7, 4, 4)
output = torch.empty(vae.decoded_pixel_shape(latents.shape), dtype=torch.float32)
group = _ThreadedFakeGroup(1)
group._local.rank = 0
with pytest.raises(ValueError, match="strategy"):
decode_to_pixels_parallel(vae, latents, output, group, strategy="scatter")
with pytest.raises(ValueError, match="must provide the CPU output buffer"):
decode_to_pixels_parallel(vae, latents, None, group)
with pytest.raises(ValueError, match="CPU float32 tensor"):
decode_to_pixels_parallel(vae, latents, output[:, :, :-1], group)
group._local.rank = 1 # simulate a non-leader passing a buffer
group.world_size = 2
with pytest.raises(ValueError, match="Only the first sequence-parallel rank"):
decode_to_pixels_parallel(vae, latents, output, group)
@pytest.mark.parametrize("world_size", (2, 4))
@pytest.mark.parametrize("num_frames", (16, 22, 40))
@torch.inference_mode()
def test_parallel_encode_matches_serial(world_size: int, num_frames: int) -> None:
"""Every rank must hold the full serial moments, bit for bit."""
torch.manual_seed(20260825 + num_frames)
vae = _tiny_vae()
pixels = torch.randint(0, 256, (1, 3, num_frames, 16, 16), dtype=torch.uint8)
expected = vae.encode_pixels(pixels).latent_dist.parameters
group = _ThreadedFakeGroup(world_size)
results = group.run(lambda rank: encode_pixels_parallel(vae, pixels, group).latent_dist.parameters)
for moments in results:
assert_close(moments, expected, atol=0.0, rtol=0.0)
@torch.inference_mode()
def test_parallel_encode_float_and_batched_slicing() -> None:
torch.manual_seed(20260826)
vae = _tiny_vae()
vae.enable_slicing()
pixels = torch.rand(2, 3, 22, 16, 16)
expected = vae.encode_pixels(pixels).latent_dist.parameters
group = _ThreadedFakeGroup(3)
results = group.run(lambda rank: encode_pixels_parallel(vae, pixels, group).latent_dist.parameters)
for moments in results:
assert_close(moments, expected, atol=0.0, rtol=0.0)
def test_parallel_encode_validates_input() -> None:
vae = _tiny_vae()
group = _ThreadedFakeGroup(1)
group._local.rank = 0
with pytest.raises(ValueError, match="must remain on CPU"):
encode_pixels_parallel(vae, torch.empty(1, 3, 4, 16, 16, device="meta"), group)
with pytest.raises(TypeError, match="uint8 or a floating-point"):
encode_pixels_parallel(vae, torch.zeros(1, 3, 4, 16, 16, dtype=torch.int32), group)
with pytest.raises(ValueError, match="must have shape"):
encode_pixels_parallel(vae, torch.zeros(1, 4, 4, 16, 16), group)
def test_fastvideo_args_strategy_literals_match_module() -> None:
"""fastvideo_args mirrors the strategy literals to avoid importing model
modules at args construction; keep the two in sync."""
from fastvideo.fastvideo_args import FastVideoArgs
args = FastVideoArgs(model_path="test/parallel-vae")
assert args.vae_parallel_decode is False
assert args.vae_parallel_encode is False
assert args.vae_parallel_decode_strategy == DEFAULT_DECODE_GATHER_STRATEGY
assert args.vae_parallel_decode_strategy in DECODE_GATHER_STRATEGIES
for strategy in DECODE_GATHER_STRATEGIES:
assert FastVideoArgs(model_path="test/parallel-vae",
vae_parallel_decode_strategy=strategy).vae_parallel_decode_strategy == strategy
with pytest.raises(ValueError, match="vae_parallel_decode_strategy"):
FastVideoArgs(model_path="test/parallel-vae", vae_parallel_decode_strategy="scatter")
@@ -1,110 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""GPU regression test for sequence-parallel MiniMax-H3 VAE decode/encode.
Requires a multi-GPU torchrun launch (real NCCL collectives across an SP
group); skipped otherwise:
torchrun --nproc-per-node=4 -m pytest \
fastvideo/tests/vaes/test_minimax_h3_parallel_vae_gpu.py -q
Asserts the parallel drivers are bitwise equal to the serial rank-local
decode/encode under the pipeline's fp16 autocast, for both transport
strategies, and that repeated parallel runs are deterministic.
"""
import os
import pytest
import torch
from torch.testing import assert_close
from fastvideo.configs.models.vaes.minimax_h3_video import (
MiniMaxH3VideoVAEArchConfig,
MiniMaxH3VideoVAEConfig,
)
from fastvideo.models.vaes.minimax_h3_video import AutoencoderKLMiniMaxH3
_WORLD_SIZE = int(os.environ.get("WORLD_SIZE", "1"))
def _tiny_vae() -> AutoencoderKLMiniMaxH3:
"""Same tiny geometry as test_minimax_h3_parallel_vae (test dirs are not packages)."""
arch = MiniMaxH3VideoVAEArchConfig(
latent_channels=4,
block_out_channels=(32, 32),
layers_per_block=1,
spatial_downsample_factors=(2, 2),
temporal_downsample_factors=(2, 2),
decoder_num_layers=1,
decoder_num_attention_heads=1,
decoder_attention_head_dim=8,
decoder_num_register_tokens=2,
decoder_ffn_mult=1,
latents_mean=(0.0, ) * 4,
latents_std=(1.0, ) * 4,
)
return AutoencoderKLMiniMaxH3(
MiniMaxH3VideoVAEConfig(
arch_config=arch,
use_tiling=False,
use_temporal_tiling=False,
use_parallel_tiling=False,
)).eval()
pytestmark = [
pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA"),
pytest.mark.skipif(_WORLD_SIZE < 2, reason="requires a torchrun launch with WORLD_SIZE > 1"),
]
@pytest.fixture(scope="module")
def sp_group():
from fastvideo.distributed import get_sp_group, maybe_init_distributed_environment_and_model_parallel
maybe_init_distributed_environment_and_model_parallel(1, _WORLD_SIZE)
return get_sp_group()
@pytest.mark.parametrize("strategy", ("gather", "all_gather"))
@pytest.mark.parametrize("latent_frames", (3, 13))
@torch.no_grad()
def test_parallel_decode_bitwise_matches_serial_on_gpu(sp_group, strategy: str, latent_frames: int) -> None:
from fastvideo.models.vaes.minimax_h3_parallel import decode_to_pixels_parallel
device = torch.device("cuda", torch.cuda.current_device())
torch.manual_seed(20260821) # identical weights on every rank
vae = _tiny_vae().to(device)
latents = torch.randn(1, 4, latent_frames, 4, 4, generator=torch.Generator().manual_seed(7)).to(device)
with torch.autocast(device_type="cuda", dtype=torch.float16):
expected = torch.empty(vae.decoded_pixel_shape(latents.shape), dtype=torch.float32, pin_memory=True)
vae.decode_to_pixels(latents, expected)
outputs = []
for _ in range(3): # repeat-determinism
output = None
if sp_group.is_first_rank:
output = torch.empty(vae.decoded_pixel_shape(latents.shape), dtype=torch.float32, pin_memory=True)
result = decode_to_pixels_parallel(vae, latents, output, sp_group, strategy=strategy)
outputs.append(result.clone() if result is not None else None)
if sp_group.is_first_rank:
for output in outputs:
assert_close(output, expected, atol=0.0, rtol=0.0)
else:
assert all(output is None for output in outputs)
@torch.no_grad()
def test_parallel_encode_bitwise_matches_serial_on_gpu(sp_group) -> None:
from fastvideo.models.vaes.minimax_h3_parallel import encode_pixels_parallel
device = torch.device("cuda", torch.cuda.current_device())
torch.manual_seed(20260821)
vae = _tiny_vae().to(device)
pixels = torch.randint(0, 256, (1, 3, 40, 16, 16), dtype=torch.uint8,
generator=torch.Generator().manual_seed(9))
expected = vae.encode_pixels(pixels).latent_dist.parameters
for _ in range(3):
moments = encode_pixels_parallel(vae, pixels, sp_group).latent_dist.parameters
assert_close(moments, expected, atol=0.0, rtol=0.0)
@@ -1,428 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""CPU contract tests for MiniMax H3 VAE compilation and profiling ranges.
The final test is a CUDA regression gate for the reduce-overhead tile path
(real tiled decode/encode with an unmocked ``_stitch_tiles``).
"""
from contextlib import contextmanager
from types import MethodType, SimpleNamespace
from typing import Any
from unittest.mock import Mock, patch
import pytest
import torch
import torch.nn as nn
import torch.nn.functional as F
from fastvideo.models.vaes.minimax_h3_audio import (
MiniMaxH3AudioBigVGANDecoder,
MiniMaxH3AudioVAE,
)
from fastvideo.models.vaes.minimax_h3_video import (
AutoencoderKLMiniMaxH3,
MiniMaxH3VideoAttention,
MiniMaxH3VideoViTDecoder3d,
)
from fastvideo.platforms import AttentionBackendEnum
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
def _empty_typed_module(module_type: type[nn.Module]) -> nn.Module:
"""Create a weightless instance that retains its production module type."""
module = object.__new__(module_type)
nn.Module.__init__(module)
return module
def _assert_dynamic_compile_selects_decoder(
vae_type: type[nn.Module],
decoder_type: type[nn.Module],
) -> None:
"""Verify one H3 VAE compiles only its top-level decoder in place."""
vae = _empty_typed_module(vae_type)
decoder = _empty_typed_module(decoder_type)
same_type_under_another_name = _empty_typed_module(decoder_type)
unrelated_submodule = nn.Identity()
vae.decoder = decoder
vae.same_type_under_another_name = same_type_under_another_name
vae.unrelated_submodule = unrelated_submodule
compiled_forward = Mock(name="compiled_forward")
compile_kwargs = {"backend": "inductor", "dynamic": False}
with patch(
"fastvideo.pipelines.composed_pipeline_base.torch.compile",
return_value=compiled_forward,
) as compile_mock:
compiled_count = ComposedPipelineBase._compile_with_conditions(vae, compile_kwargs)
assert compiled_count == 1
compile_mock.assert_called_once()
selected_forward = compile_mock.call_args.args[0]
assert selected_forward.__self__ is decoder
assert selected_forward.__func__ is decoder_type.forward
assert compile_mock.call_args.kwargs == compile_kwargs
assert decoder.forward is compiled_forward
assert "forward" not in same_type_under_another_name.__dict__
assert "forward" not in unrelated_submodule.__dict__
wrong_type_vae = _empty_typed_module(vae_type)
wrong_type_vae.decoder = nn.Identity()
with patch("fastvideo.pipelines.composed_pipeline_base.torch.compile") as wrong_type_compile:
wrong_type_count = ComposedPipelineBase._compile_with_conditions(wrong_type_vae, compile_kwargs)
assert wrong_type_count == 0
wrong_type_compile.assert_not_called()
def _assert_reduce_overhead_compile(compiled_function: Any) -> None:
"""Verify a class-owned compile boundary enables CUDA Graph replay."""
assert hasattr(compiled_function, "get_compiler_config")
assert compiled_function.get_compiler_config()["triton.cudagraphs"] is True
def test_video_attention_uses_selected_fastvideo_backend() -> None:
"""Pass BSHD tensors to the selected dense backend without forward metadata."""
backend_call: dict[str, Any] = {}
class RecordingAttentionImpl:
"""Record the backend construction and forward contracts."""
def __init__(self, **kwargs: Any) -> None:
backend_call["init"] = kwargs
def forward(
self,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
attention_metadata: Any,
) -> torch.Tensor:
"""Return values unchanged after recording the backend inputs."""
backend_call["shapes"] = (query.shape, key.shape, value.shape)
backend_call["metadata"] = attention_metadata
return value
class RecordingAttentionBackend:
"""Supply the recording implementation through the backend API."""
@staticmethod
def get_impl_cls() -> type[RecordingAttentionImpl]:
return RecordingAttentionImpl
with (
patch("fastvideo.platforms.current_platform") as current_platform,
patch(
"fastvideo.models.vaes.minimax_h3_video.get_attn_backend",
return_value=RecordingAttentionBackend,
) as get_backend,
):
current_platform.is_cuda_alike.return_value = True
attention = MiniMaxH3VideoAttention(dim=8, heads=2, dim_head=4)
attention.to_q = nn.Identity()
attention.to_k = nn.Identity()
attention.to_v = nn.Identity()
attention.norm_q = nn.Identity()
attention.norm_k = nn.Identity()
attention.to_out[0] = nn.Identity()
output = attention(torch.empty((1, 3, 8), device="meta"))
get_backend.assert_called_once_with(
4,
torch.bfloat16,
supported_attention_backends=(
AttentionBackendEnum.TORCH_SDPA,
AttentionBackendEnum.FLASH_ATTN,
),
)
assert backend_call["init"] == {
"num_heads": 2,
"head_size": 4,
"softmax_scale": 0.5,
"num_kv_heads": 2,
"causal": False,
}
assert backend_call["shapes"] == ((1, 3, 2, 4), ) * 3
assert backend_call["metadata"] is None
assert output.shape == (1, 3, 8)
def test_video_attention_cpu_uses_torch_sdpa() -> None:
"""Use PyTorch SDPA when H3 VAE attention receives CPU tensors."""
with (
patch("fastvideo.platforms.current_platform") as current_platform,
patch("fastvideo.models.vaes.minimax_h3_video.get_attn_backend") as get_backend,
):
current_platform.is_cuda_alike.return_value = False
attention = MiniMaxH3VideoAttention(dim=8, heads=2, dim_head=4)
get_backend.assert_not_called()
assert attention.attn_impl is None
attention.to_q = nn.Identity()
attention.to_k = nn.Identity()
attention.to_v = nn.Identity()
attention.norm_q = nn.Identity()
attention.norm_k = nn.Identity()
attention.to_out[0] = nn.Identity()
hidden_states = torch.randn(1, 3, 8)
query = hidden_states.unflatten(2, (2, 4)).permute(0, 2, 1, 3)
expected = F.scaled_dot_product_attention(query, query, query).permute(0, 2, 1, 3).flatten(2, 3)
torch.testing.assert_close(attention(hidden_states), expected)
def test_compile_with_conditions_selects_minimax_h3_video_decoder() -> None:
"""Compile the registered video decoder with the VAE runtime kwargs."""
assert not hasattr(MiniMaxH3VideoViTDecoder3d.forward, "get_compiler_config")
_assert_dynamic_compile_selects_decoder(AutoencoderKLMiniMaxH3, MiniMaxH3VideoViTDecoder3d)
def test_project_decoder_tile_uses_reduce_overhead_compile() -> None:
"""Compile the per-tile decoder-input projection with CUDA Graph replay."""
_assert_reduce_overhead_compile(AutoencoderKLMiniMaxH3._project_decoder_tile)
def test_stitch_tiles_uses_reduce_overhead_compile() -> None:
"""Compile spatial tile blending and concatenation with CUDA Graph replay."""
_assert_reduce_overhead_compile(AutoencoderKLMiniMaxH3._stitch_tiles)
def test_compile_with_conditions_selects_minimax_h3_audio_decoder() -> None:
"""Compile the audio VAE decoder that the H3 waveform decode path calls."""
_assert_dynamic_compile_selects_decoder(MiniMaxH3AudioVAE, MiniMaxH3AudioBigVGANDecoder)
def test_decode_emits_indexed_temporal_chunk_ranges() -> None:
"""Nest frame-segment ranges under each temporal decoder chunk range."""
vae = _empty_typed_module(AutoencoderKLMiniMaxH3)
vae.tokens_chunk_size = 1
vae.token_overlap = 1
vae.temporal_compression_ratio = 1
vae.frame_pre_padding = 0
vae.frame_overlap = 1
vae.config = SimpleNamespace(token_drop=1)
vae._decode_clip = Mock(return_value=torch.zeros((1, 1, 2, 1, 1)))
range_events = []
@contextmanager
def record_range(name: str):
range_events.append(("enter", name))
try:
yield
finally:
range_events.append(("exit", name))
with patch("fastvideo.models.vaes.minimax_h3_video.nvtx_range", record_range):
decoded = vae._decode(torch.zeros((1, 1, 2, 1, 1)))
assert decoded.shape == (1, 1, 3, 1, 1)
assert vae._decode_clip.call_count == 2
assert range_events == [
("enter", "minimax_h3.vae.temporal_chunk.0"),
("enter", "minimax_h3.vae.temporal_chunk.0.frame_segment.0"),
("exit", "minimax_h3.vae.temporal_chunk.0.frame_segment.0"),
("enter", "minimax_h3.vae.temporal_chunk.0.frame_segment.1"),
("exit", "minimax_h3.vae.temporal_chunk.0.frame_segment.1"),
("exit", "minimax_h3.vae.temporal_chunk.0"),
("enter", "minimax_h3.vae.temporal_chunk.1"),
("enter", "minimax_h3.vae.temporal_chunk.1.frame_segment.0"),
("exit", "minimax_h3.vae.temporal_chunk.1.frame_segment.0"),
("enter", "minimax_h3.vae.temporal_chunk.1.frame_segment.1"),
("exit", "minimax_h3.vae.temporal_chunk.1.frame_segment.1"),
("exit", "minimax_h3.vae.temporal_chunk.1"),
]
def test_decode_clip_no_spatial_tiling_stage_ranges() -> None:
"""Separate untiled latent projection and decoder ranges."""
vae = _empty_typed_module(AutoencoderKLMiniMaxH3)
vae.use_tiling = False
range_events = []
vae.post_quant_conv = nn.Identity()
vae.post_quant_conv.register_forward_hook(
lambda _module, _args, _output: range_events.append(("call", "post_quant_conv")))
vae.decoder = nn.Identity()
vae.decoder.register_forward_hook(
lambda _module, _args, _output: range_events.append(("call", "decoder_forward")))
latent_clip = torch.zeros((1, 1, 1, 2, 2))
@contextmanager
def record_range(name: str):
range_events.append(("enter", name))
try:
yield
finally:
range_events.append(("exit", name))
with patch("fastvideo.models.vaes.minimax_h3_video.nvtx_range", record_range):
decoded_clip = vae._decode_clip(latent_clip)
assert decoded_clip is latent_clip
assert range_events == [
("enter", "minimax_h3.vae.decode_clip"),
("enter", "minimax_h3.vae.decode_clip.no_s_tile.post_quant_conv"),
("call", "post_quant_conv"),
("exit", "minimax_h3.vae.decode_clip.no_s_tile.post_quant_conv"),
("enter", "minimax_h3.vae.decode_clip.no_s_tile.decoder_forward"),
("call", "decoder_forward"),
("exit", "minimax_h3.vae.decode_clip.no_s_tile.decoder_forward"),
("exit", "minimax_h3.vae.decode_clip"),
]
def test_decode_clip_emits_tiled_stage_ranges() -> None:
"""Nest indexed decoder tiles between tile-splitting and stitching ranges."""
vae = _empty_typed_module(AutoencoderKLMiniMaxH3)
vae.use_tiling = True
vae.spatial_compression_ratio = 1
vae.tile_sample_min_height = 1
vae.tile_sample_min_width = 1
vae.tile_sample_min_overlap_height = 0
vae.tile_sample_min_overlap_width = 0
vae._split_tiles = Mock(side_effect=[
([0, 1], [1, 1], [0]),
([0, 1], [1, 1], [0]),
])
vae.post_quant_conv = nn.Identity()
vae._project_decoder_tile = Mock(side_effect=vae.post_quant_conv)
vae.decoder = nn.Identity()
stitched_clip = torch.zeros((1, 1, 1, 2, 2))
vae._stitch_tiles = Mock(return_value=stitched_clip)
range_events = []
@contextmanager
def record_range(name: str):
range_events.append(("enter", name))
try:
yield
finally:
range_events.append(("exit", name))
with patch("fastvideo.models.vaes.minimax_h3_video.nvtx_range", record_range):
decoded_clip = vae._decode_clip(torch.zeros((1, 1, 1, 2, 2)))
# The tile driver must hand back a caller-owned copy: under
# mode="reduce-overhead" the stitched canvas is CUDA-graph pooled storage
# that the next replay overwrites, so returning it by identity is a bug.
assert decoded_clip is not stitched_clip
assert torch.equal(decoded_clip, stitched_clip)
assert vae._split_tiles.call_count == 2
assert vae._project_decoder_tile.call_count == 4
assert vae._stitch_tiles.call_count == 1
assert range_events == [
("enter", "minimax_h3.vae.decode_clip"),
("enter", "minimax_h3.vae.decode_clip.split_tiles"),
("exit", "minimax_h3.vae.decode_clip.split_tiles"),
("enter", "minimax_h3.vae.decode_clip.decode_tiles"),
("enter", "minimax_h3.vae.decode_clip.tile.0.0"),
("enter", "minimax_h3.vae.decode_clip.tile.decoder_forward"),
("exit", "minimax_h3.vae.decode_clip.tile.decoder_forward"),
("exit", "minimax_h3.vae.decode_clip.tile.0.0"),
("enter", "minimax_h3.vae.decode_clip.tile.0.1"),
("enter", "minimax_h3.vae.decode_clip.tile.decoder_forward"),
("exit", "minimax_h3.vae.decode_clip.tile.decoder_forward"),
("exit", "minimax_h3.vae.decode_clip.tile.0.1"),
("enter", "minimax_h3.vae.decode_clip.tile.1.0"),
("enter", "minimax_h3.vae.decode_clip.tile.decoder_forward"),
("exit", "minimax_h3.vae.decode_clip.tile.decoder_forward"),
("exit", "minimax_h3.vae.decode_clip.tile.1.0"),
("enter", "minimax_h3.vae.decode_clip.tile.1.1"),
("enter", "minimax_h3.vae.decode_clip.tile.decoder_forward"),
("exit", "minimax_h3.vae.decode_clip.tile.decoder_forward"),
("exit", "minimax_h3.vae.decode_clip.tile.1.1"),
("exit", "minimax_h3.vae.decode_clip.decode_tiles"),
("enter", "minimax_h3.vae.decode_clip.stitch_tiles"),
("exit", "minimax_h3.vae.decode_clip.stitch_tiles"),
("exit", "minimax_h3.vae.decode_clip"),
]
def _tiny_real_vae() -> AutoencoderKLMiniMaxH3:
"""Random-weight VAE small enough for a real tiled decode/encode on GPU."""
from fastvideo.configs.models.vaes.minimax_h3_video import (
MiniMaxH3VideoVAEArchConfig,
MiniMaxH3VideoVAEConfig,
)
arch = MiniMaxH3VideoVAEArchConfig(
latent_channels=4,
block_out_channels=(32, 32),
layers_per_block=1,
spatial_downsample_factors=(2, 2),
temporal_downsample_factors=(2, 2),
decoder_num_layers=1,
decoder_num_attention_heads=1,
decoder_attention_head_dim=8,
decoder_num_register_tokens=2,
decoder_ffn_mult=1,
latents_mean=(0.0, ) * 4,
latents_std=(1.0, ) * 4,
)
return AutoencoderKLMiniMaxH3(
MiniMaxH3VideoVAEConfig(
arch_config=arch,
use_tiling=False,
use_temporal_tiling=False,
use_parallel_tiling=False,
)).eval()
def _dynamo_original(compiled_function: Any) -> Any:
"""Return the eager callable behind a ``torch.compile``-decorated function."""
original = getattr(compiled_function, "_torchdynamo_orig_callable", None)
if original is None:
original = getattr(compiled_function, "__wrapped__", None)
assert original is not None, "cannot recover the eager tile helpers"
return original
@pytest.mark.skipif(not torch.cuda.is_available(), reason="reduce-overhead tile compile requires CUDA graphs")
@torch.inference_mode()
def test_tiled_decode_and_encode_survive_cudagraph_buffer_reuse_on_cuda() -> None:
"""Real tiled decode()/encode() with an unmocked reduce-overhead ``_stitch_tiles``.
Regression gate for the CUDA-graph output-clobbering bug: the stitched
canvas is a cudagraph static buffer, and the collect-then-``torch.cat``
consumers (``_decode``/``_encode``/``_encode_pixels``) hold chunk/clip
results across subsequent ``_stitch_tiles`` replays. Without the eager
``.clone()`` at the tile-driver returns, the first tiled ``decode()`` with
>=2 temporal chunks raises ``accessing tensor output of CUDAGraphs that
has been overwritten by a subsequent run``. This test needs >=2 chunks
(decode), >=2 clips (encode), and a >=2x2 spatial tile grid.
"""
torch.manual_seed(20260821)
vae = _tiny_real_vae().to("cuda")
vae.enable_tiling(16, 16, 4, 4)
# 8 latent tokens = 2 temporal chunks (tokens_chunk_size 5); 8x8 latents =
# 32x32 pixels = a 2x2 grid of 16px tiles.
z = torch.randn(1, 4, 8, 8, 8, device="cuda")
pad_tokens, num_chunks, _ = vae._temporal_decode_plan(z.shape[2])
assert num_chunks >= 2, "decode workload must span multiple stitch replays"
decoded_first = vae.decode(z).sample
decoded_second = vae.decode(z).sample
assert torch.equal(decoded_first, decoded_second)
# 34 frames = 2 encode clips of clip_length 17 -> 2 stitch replays.
pixels = torch.rand(1, 3, 34, 32, 32, device="cuda")
encoded_first = vae.encode(pixels).latent_dist.parameters
encoded_second = vae.encode(pixels).latent_dist.parameters
assert torch.equal(encoded_first, encoded_second)
uint8_pixels = torch.randint(0, 256, (1, 3, 34, 32, 32), dtype=torch.uint8)
streamed_first = vae.encode_pixels(uint8_pixels).latent_dist.parameters
streamed_second = vae.encode_pixels(uint8_pixels).latent_dist.parameters
assert torch.equal(streamed_first, streamed_second)
# Output parity vs the fully eager tile helpers (same weights, same math;
# the tolerance absorbs inductor fusion reassociation only).
eager_vae = _tiny_real_vae().to("cuda")
eager_vae.load_state_dict(vae.state_dict())
eager_vae.enable_tiling(16, 16, 4, 4)
eager_vae._stitch_tiles = MethodType(_dynamo_original(AutoencoderKLMiniMaxH3._stitch_tiles), eager_vae)
eager_vae._project_decoder_tile = MethodType(_dynamo_original(AutoencoderKLMiniMaxH3._project_decoder_tile),
eager_vae)
torch.testing.assert_close(decoded_first, eager_vae.decode(z).sample, atol=2e-4, rtol=2e-4)
torch.testing.assert_close(encoded_first, eager_vae.encode(pixels).latent_dist.parameters, atol=2e-4, rtol=2e-4)
@@ -1,168 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
import pytest
import torch
from torch.testing import assert_close
from fastvideo.configs.models.vaes.minimax_h3_video import (
MiniMaxH3VideoVAEArchConfig,
MiniMaxH3VideoVAEConfig,
)
from fastvideo.models.vaes.minimax_h3_video import AutoencoderKLMiniMaxH3
def _tiny_vae() -> AutoencoderKLMiniMaxH3:
arch = MiniMaxH3VideoVAEArchConfig(
latent_channels=4,
block_out_channels=(32, 32),
layers_per_block=1,
spatial_downsample_factors=(2, 2),
temporal_downsample_factors=(2, 2),
decoder_num_layers=1,
decoder_num_attention_heads=1,
decoder_attention_head_dim=8,
decoder_num_register_tokens=2,
decoder_ffn_mult=1,
latents_mean=(0.0, ) * 4,
latents_std=(1.0, ) * 4,
)
return AutoencoderKLMiniMaxH3(
MiniMaxH3VideoVAEConfig(
arch_config=arch,
use_tiling=False,
use_temporal_tiling=False,
use_parallel_tiling=False,
)).eval()
@torch.inference_mode()
def test_encode_pixels_matches_encode() -> None:
torch.manual_seed(20260810)
vae = _tiny_vae()
pixels = torch.randint(0, 256, (1, 3, 22, 16, 16), dtype=torch.uint8)
expected = vae.encode(vae.normalize_pixels(pixels.float().div(255))).latent_dist.parameters
assert_close(vae.encode_pixels(pixels).latent_dist.parameters, expected, atol=0.0, rtol=0.0)
float_pixels = torch.rand(1, 3, 22, 16, 16)
original_pixels = float_pixels.clone()
expected = vae.encode(vae.normalize_pixels(float_pixels)).latent_dist.parameters
actual = vae.encode_pixels(float_pixels).latent_dist.parameters
assert_close(float_pixels, original_pixels, atol=0.0, rtol=0.0)
assert_close(actual, expected, atol=0.0, rtol=0.0)
with pytest.raises(ValueError, match="must remain on CPU"):
vae.encode_pixels(torch.empty(1, 3, 1, 16, 16, device="meta"))
def _legacy_decode(vae: AutoencoderKLMiniMaxH3, z: torch.Tensor) -> torch.Tensor:
"""Verbatim pre-streaming ``_decode`` (main @ 8208536cd) as a reference oracle."""
tokens_chunk_size = vae.tokens_chunk_size
token_drop = vae.config.token_drop
temporal_ratio = vae.temporal_compression_ratio
chunk_num_frames = tokens_chunk_size * temporal_ratio
num_tokens = z.shape[2] + token_drop
pad_tokens = (-num_tokens) % tokens_chunk_size
num_chunks = (num_tokens + pad_tokens) // tokens_chunk_size - int(token_drop > 0)
if pad_tokens > 0:
z = torch.cat([z, z[:, :, -1:].repeat(1, 1, pad_tokens, 1, 1)], dim=2)
decoded_chunks = []
overlap = None
for index in range(num_chunks):
start = index * tokens_chunk_size
clip = vae._decode_clip(z[:, :, start:start + tokens_chunk_size + vae.token_overlap])
for overlap_index in range(int(token_drop > 0) + 1):
frame_start = overlap_index * chunk_num_frames
chunk = clip[:, :, frame_start:frame_start + chunk_num_frames]
chunk = chunk[:, :, vae.frame_pre_padding:]
if overlap_index == 0:
if overlap is not None:
chunk = vae._blend(overlap, chunk, vae.frame_overlap, dim=-3)
decoded_chunks.append(chunk)
else:
overlap = chunk
if overlap is not None:
decoded_chunks.append(overlap)
decoded = torch.cat(decoded_chunks, dim=2)
if pad_tokens > 0:
intra_tail = vae.config.clip_length % temporal_ratio
num_tokens_before_pad = z.shape[2] - pad_tokens
pad_frames = sum(intra_tail if intra_tail and (num_tokens_before_pad + offset) %
tokens_chunk_size == 0 else temporal_ratio for offset in range(pad_tokens))
decoded = decoded[:, :, :-pad_frames]
return decoded
@pytest.mark.parametrize("latent_frames", (2, 12))
@torch.inference_mode()
def test_decode_to_pixels_matches_decode(latent_frames: int) -> None:
torch.manual_seed(20260810)
vae = _tiny_vae()
latents = torch.randn(1, 4, latent_frames, 4, 4)
expected = vae.denormalize_pixels(vae.decode(latents).sample.float()).clamp_(0, 1)
actual = torch.empty(vae.decoded_pixel_shape(latents.shape), dtype=torch.float32)
vae.decode_to_pixels(latents, actual)
assert_close(actual, expected, atol=0.0, rtol=0.0)
# 3: one chunk with pad tokens; 6: pad hitting the intra-clip tail;
# 12: two blended chunks without padding; 13: three chunks plus pad trim.
@pytest.mark.parametrize("latent_frames", (3, 6, 12, 13))
@torch.inference_mode()
def test_decode_matches_legacy_algorithm(latent_frames: int) -> None:
"""The chunk iterator must stay bit-exact with the pre-streaming decode."""
torch.manual_seed(20260811 + latent_frames)
vae = _tiny_vae()
latents = torch.randn(1, 4, latent_frames, 4, 4)
expected = _legacy_decode(vae, latents)
decoded = vae.decode(latents).sample
assert_close(decoded, expected, atol=0.0, rtol=0.0)
streamed = torch.empty(vae.decoded_pixel_shape(latents.shape), dtype=torch.float32)
vae.decode_to_pixels(latents, streamed)
assert_close(streamed, vae.denormalize_pixels(expected.float()).clamp_(0, 1), atol=0.0, rtol=0.0)
@torch.inference_mode()
def test_streaming_slicing_matches_unbatched() -> None:
torch.manual_seed(20260812)
vae = _tiny_vae()
vae.enable_slicing()
latents = torch.randn(2, 4, 7, 4, 4)
expected = vae.denormalize_pixels(vae.decode(latents).sample.float()).clamp_(0, 1)
actual = torch.empty(vae.decoded_pixel_shape(latents.shape), dtype=torch.float32)
vae.decode_to_pixels(latents, actual)
assert_close(actual, expected, atol=0.0, rtol=0.0)
pixels = torch.randint(0, 256, (2, 3, 22, 16, 16), dtype=torch.uint8)
expected_moments = vae.encode(vae.normalize_pixels(pixels.float().div(255))).latent_dist.parameters
assert_close(vae.encode_pixels(pixels).latent_dist.parameters, expected_moments, atol=0.0, rtol=0.0)
@pytest.mark.skipif(not torch.cuda.is_available(), reason="pinned-buffer streaming requires CUDA")
@torch.inference_mode()
def test_decode_to_pixels_pinned_buffer_matches_dense_on_cuda() -> None:
"""Async chunk copies into a pinned buffer must equal the dense decode."""
torch.manual_seed(20260813)
vae = _tiny_vae().to("cuda")
latents = torch.randn(1, 4, 12, 4, 4, device="cuda")
expected = vae.denormalize_pixels(vae.decode(latents).sample.float()).clamp_(0, 1).cpu()
pinned = torch.empty(vae.decoded_pixel_shape(latents.shape), dtype=torch.float32, pin_memory=True)
vae.decode_to_pixels(latents, pinned)
assert_close(pinned, expected, atol=0.0, rtol=0.0)
def test_decode_to_pixels_rejects_incomplete_output(monkeypatch) -> None:
vae = _tiny_vae()
latents = torch.randn(1, 4, 2, 4, 4)
output = torch.empty(vae.decoded_pixel_shape(latents.shape), dtype=torch.float32)
monkeypatch.setattr(vae, "_decode_chunks", lambda _: iter(()))
with pytest.raises(RuntimeError, match="wrote 0 frames"):
vae.decode_to_pixels(latents, output)
+1 -30
View File
@@ -1,40 +1,11 @@
from types import SimpleNamespace
from unittest.mock import Mock
import pytest
import torch
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.pipelines import ForwardBatch
from fastvideo.worker.gpu_worker import Worker, _log_cuda_device_uuid
def test_cuda_device_uuid_receipt_is_disabled_without_nvtx_profiling(monkeypatch) -> None:
"""Avoid NVIDIA property access during ordinary worker initialization."""
get_device_properties = Mock()
monkeypatch.setenv("FASTVIDEO_NVTX_PROFILE", "0")
monkeypatch.setattr(torch.cuda, "get_device_properties", get_device_properties)
_log_cuda_device_uuid(0, torch.device("cuda:0"))
get_device_properties.assert_not_called()
def test_cuda_device_uuid_receipt_identifies_profiled_worker(monkeypatch) -> None:
"""Bind one profiled worker rank to its NVIDIA device UUID in logs."""
log_info = Mock()
monkeypatch.setenv("FASTVIDEO_NVTX_PROFILE", "1")
monkeypatch.setattr(torch.cuda, "get_device_properties", lambda device: SimpleNamespace(uuid="device-uuid"))
monkeypatch.setattr("fastvideo.worker.gpu_worker.logger.info", log_info)
_log_cuda_device_uuid(2, torch.device("cuda:0"))
log_info.assert_called_once_with(
"Worker %d CUDA device UUID: GPU-%s",
2,
"device-uuid",
local_main_process_only=False,
)
from fastvideo.worker.gpu_worker import Worker
def _worker_returning(output_batch: ForwardBatch) -> Worker:
+2
View File
@@ -94,6 +94,8 @@ class EMACallback(Callback):
if iteration < self._start_iter:
return
if not method.did_update_role("student", iteration):
return
if not self._ema_started:
logger.info(
"Starting EMA updates at iteration %d "
+2
View File
@@ -31,6 +31,7 @@ from fastvideo.distributed import (
get_world_group,
)
from fastvideo.logger import init_logger
from fastvideo.profiler import profile_region
from fastvideo.pipelines import ForwardBatch
from fastvideo.train.callbacks.callback import Callback
from fastvideo.train.utils.instantiate import resolve_target
@@ -285,6 +286,7 @@ class ValidationCallback(Callback):
# Core validation logic
# ----------------------------------------------------------
@profile_region("profiler_region_training_validation")
def _run_validation(
self,
method: TrainingMethod,
+45 -35
View File
@@ -24,7 +24,9 @@ from typing import Any
import torch
import fastvideo.envs as envs
from fastvideo.logger import init_logger
from fastvideo.profiler import get_or_create_profiler
logger = init_logger(__name__)
@@ -74,47 +76,55 @@ def run_training_from_config(
tc.distributed.sp_size,
)
_, method, dataloader, start_step = build_from_config(cfg)
profiler_controller = get_or_create_profiler(envs.FASTVIDEO_TORCH_PROFILER_DIR, )
try:
with profiler_controller.region("profiler_region_model_loading"):
_, method, dataloader, start_step = build_from_config(cfg)
if dry_run:
logger.info("Dry-run: config parsed and "
"build_from_config succeeded.")
return
if dry_run:
logger.info("Dry-run: config parsed and "
"build_from_config succeeded.")
return
trainer = Trainer(
tc,
config=cfg.resolved_config(),
callback_configs=cfg.callbacks,
)
trainer = Trainer(
tc,
config=cfg.resolved_config(),
callback_configs=cfg.callbacks,
)
# Attach the exact YAML used for this run to the
# tracker (e.g., W&B Files).
trainer.tracker.log_file(
os.path.abspath(os.path.expanduser(config_path)),
name="run.yaml",
)
# Attach the exact YAML used for this run to the
# tracker (e.g., W&B Files).
trainer.tracker.log_file(
os.path.abspath(os.path.expanduser(config_path)),
name="run.yaml",
)
ckpt_config = CheckpointConfig(
save_steps=int(tc.checkpoint.training_state_checkpointing_steps or 0),
keep_last=int(tc.checkpoint.checkpoints_total_limit or 0),
)
ckpt_config = CheckpointConfig(
save_steps=int(tc.checkpoint.training_state_checkpointing_steps or 0),
keep_last=int(tc.checkpoint.checkpoints_total_limit or 0),
)
checkpoint_manager = CheckpointManager(
method=method,
dataloader=dataloader,
output_dir=tc.checkpoint.output_dir,
config=ckpt_config,
callbacks=trainer.callbacks,
raw_config=cfg.raw,
)
checkpoint_manager = CheckpointManager(
method=method,
dataloader=dataloader,
output_dir=tc.checkpoint.output_dir,
config=ckpt_config,
callbacks=trainer.callbacks,
raw_config=cfg.raw,
)
trainer.run(
method,
dataloader=dataloader,
max_steps=tc.loop.max_train_steps,
start_step=start_step,
checkpoint_manager=checkpoint_manager,
)
trainer.run(
method,
dataloader=dataloader,
max_steps=tc.loop.max_train_steps,
start_step=start_step,
checkpoint_manager=checkpoint_manager,
)
finally:
# torch.profiler exports its trace at stop(). Keep atexit as a fallback
# for abrupt exits, but flush here so normal training returns only after
# all per-rank traces and summaries are complete.
profiler_controller.stop()
def main(
+17
View File
@@ -208,6 +208,23 @@ class TrainingMethod(torch.nn.Module, ABC):
"""
return False
def did_update_role(
self,
role: str,
iteration: int,
) -> bool:
"""Whether ``role`` completed an optimizer update this iteration.
Callbacks run after optimization and use this hook to follow the
method's actual update cadence. The default implementation matches
the role's optimizer against the optimizers selected for this
iteration; methods that step optimizers internally can override it.
"""
role_optimizer = self._optimizer_dict.get(role)
if role_optimizer is None:
return False
return any(optimizer is role_optimizer for optimizer in self.get_optimizers(iteration))
def managed_train_step(
self,
data_stream: Any,
@@ -6,8 +6,10 @@ from __future__ import annotations
from typing import Any, Literal
import torch
import torch.distributed as dist
import torch.nn.functional as F
from fastvideo.profiler import profile_region
from fastvideo.train.methods.base import TrainingMethod, LogScalar
from fastvideo.train.models.base import ModelBase
from fastvideo.train.utils.optimizer import (
@@ -431,6 +433,23 @@ class DMD2Method(TrainingMethod):
)
return step_list[index]
@staticmethod
def _max_rollout_target_idx_across_ranks(target_timestep_idx: torch.Tensor, ) -> int:
"""Return the largest sampled rollout index across all data ranks.
Each rank intentionally samples its own DMD2 timestep, but FSDP ranks
must execute the same number of transformer forwards so their
collectives stay ordered. The largest local target is therefore the
minimum prefix length every rank must execute this iteration.
"""
max_target_timestep_idx = target_timestep_idx.detach().clone()
if dist.is_available() and dist.is_initialized():
dist.all_reduce(
max_target_timestep_idx,
op=dist.ReduceOp.MAX,
)
return int(max_target_timestep_idx.item())
def _parse_score_timestep_bounds(self) -> tuple[int, int]:
"""Resolve the score-model timestep window used by legacy DMD.
@@ -476,6 +495,7 @@ class DMD2Method(TrainingMethod):
self._score_max_timestep,
)
@profile_region("profiler_region_dmd2_student_rollout")
def _student_rollout(
self,
batch: Any,
@@ -516,6 +536,7 @@ class DMD2Method(TrainingMethod):
generator=self.cuda_generator,
)
target_timestep_idx_int = int(target_timestep_idx.item())
synchronized_target_idx = self._max_rollout_target_idx_across_ranks(target_timestep_idx, )
target_timestep = step_list[target_timestep_idx]
current_noise_latents = torch.randn(
@@ -524,56 +545,62 @@ class DMD2Method(TrainingMethod):
dtype=dtype,
generator=self.cuda_generator,
)
current_noise_latents_copy = (current_noise_latents.clone())
max_target_idx = len(step_list) - 1
noise_latents: list[torch.Tensor] = []
noise_latent_index = target_timestep_idx_int - 1
noisy_input = current_noise_latents
if max_target_idx > 0:
with torch.no_grad():
for step_idx in range(max_target_idx):
current_timestep = step_list[step_idx]
current_timestep_tensor = (current_timestep * torch.ones(
1,
device=device,
dtype=torch.long,
))
# FSDP ranks must run the same number of forwards. Ranks
# whose local target is earlier keep advancing a throwaway
# trajectory until the largest target sampled globally.
needs_prefix_step = step_idx < synchronized_target_idx
pred_clean: torch.Tensor | None = None
noise_dtype = dtype
if needs_prefix_step:
current_timestep = step_list[step_idx]
current_timestep_tensor = (current_timestep * torch.ones(
1,
device=device,
dtype=torch.long,
))
pred_clean = self.student.predict_x0(
current_noise_latents,
current_timestep_tensor,
batch,
conditional=True,
cfg_uncond=self._cfg_uncond,
attn_kind="vsa",
)
pred_clean = self.student.predict_x0(
current_noise_latents,
current_timestep_tensor,
batch,
conditional=True,
cfg_uncond=self._cfg_uncond,
attn_kind="vsa",
)
noise_dtype = pred_clean.dtype
next_timestep = step_list[step_idx + 1]
next_timestep_tensor = (next_timestep * torch.ones(
1,
device=device,
dtype=torch.long,
))
# Preserve the method-owned generator sequence even when
# the corresponding prefix forward is unnecessary. This
# keeps all later DMD2 random draws aligned with the old
# full-rollout implementation.
noise = torch.randn(
latents.shape,
device=device,
dtype=pred_clean.dtype,
dtype=noise_dtype,
generator=self.cuda_generator,
)
current_noise_latents = (self.student.add_noise(
pred_clean,
noise,
next_timestep_tensor,
))
noise_latents.append(current_noise_latents.clone())
if noise_latent_index >= 0:
if noise_latent_index >= len(noise_latents):
raise RuntimeError("noise_latent_index is out of bounds")
noisy_input = noise_latents[noise_latent_index]
else:
noisy_input = current_noise_latents_copy
if needs_prefix_step:
assert pred_clean is not None
next_timestep = step_list[step_idx + 1]
next_timestep_tensor = (next_timestep * torch.ones(
1,
device=device,
dtype=torch.long,
))
current_noise_latents = (self.student.add_noise(
pred_clean,
noise,
next_timestep_tensor,
))
if step_idx + 1 == target_timestep_idx_int:
noisy_input = current_noise_latents
if with_grad:
pred_x0 = self.student.predict_x0(
@@ -598,6 +625,7 @@ class DMD2Method(TrainingMethod):
batch.dmd_latent_vis_dict["generator_timestep"] = target_timestep.float().detach()
return pred_x0
@profile_region("profiler_region_dmd2_critic_loss")
def _critic_flow_matching_loss(
self,
batch: Any,
@@ -638,6 +666,7 @@ class DMD2Method(TrainingMethod):
outputs,
)
@profile_region("profiler_region_dmd2_generator_loss")
def _dmd_loss(
self,
generator_pred_x0: torch.Tensor,
+113 -77
View File
@@ -11,6 +11,7 @@ import torch
from tqdm.auto import tqdm
from fastvideo.distributed import get_sp_group, get_world_group
from fastvideo.profiler import profile_region, profiler_region
from fastvideo.train.callbacks.callback import CallbackDict
from fastvideo.train.methods.base import LogScalar, TrainingMethod
from fastvideo.train.utils.tracking import build_tracker
@@ -98,6 +99,113 @@ class Trainer:
if self.global_rank == 0 and validation_metrics:
self.tracker.log(validation_metrics, iteration)
@profile_region("profiler_region_training_train_one_step")
def _run_train_step(
self,
method: TrainingMethod,
*,
data_stream: Iterator[dict[str, Any]],
step: int,
grad_accum: int,
method_manages_optimization: bool,
) -> None:
t0 = time.perf_counter()
# Accumulate on GPU during grad-accum; materialise to CPU once per
# step right before logging.
loss_sums: dict[str, float | torch.Tensor] = {}
metric_sums: dict[str, float | torch.Tensor] = {}
dataloader_time_sec = 0.0
if method_manages_optimization:
# Managed methods own their forward/backward/optimizer boundaries,
# so the enclosing training_train_one_step region is the truthful
# granularity available to the trainer.
loss_map, outputs, step_metrics = method.managed_train_step(
data_stream,
step,
)
for k, v in loss_map.items():
if isinstance(v, torch.Tensor):
loss_sums[k] = v.detach()
for k, v in step_metrics.items():
if k in loss_sums:
raise ValueError(f"Metric key {k!r} collides "
"with loss key. Use a "
"different name (e.g. prefix "
"with 'train/').")
metric_sums[k] = _coerce_log_scalar(
v,
where=("method.managed_train_step()"
f".metrics[{k!r}]"),
)
else:
for _ in range(grad_accum):
dataloader_t0 = time.perf_counter()
with profiler_region("profiler_region_training_dataloader"):
batch = next(data_stream)
dataloader_time_sec += (time.perf_counter() - dataloader_t0)
with profiler_region("profiler_region_training_forward"):
loss_map, outputs, step_metrics = (method.single_train_step(
batch,
step,
))
with profiler_region("profiler_region_training_backward"):
method.backward(
loss_map,
outputs,
grad_accum_rounds=grad_accum,
)
for k, v in loss_map.items():
if isinstance(v, torch.Tensor):
prev = loss_sums.get(k, 0.0)
loss_sums[k] = prev + v.detach()
for k, v in step_metrics.items():
if k in loss_sums:
raise ValueError(f"Metric key {k!r} collides "
"with loss key. Use a "
"different name (e.g. prefix "
"with 'train/').")
prev = metric_sums.get(k, 0.0)
metric_sums[k] = (prev + _coerce_log_scalar(
v,
where=("method.single_train_step()"
f".metrics[{k!r}]"),
))
if not method_manages_optimization:
with profiler_region("profiler_region_training_optimizer"):
self.callbacks.on_before_optimizer_step(
method,
iteration=step,
)
method.optimizers_schedulers_step(step)
method.optimizers_zero_grad(step)
# Single CPU sync point: materialise GPU tensors to float right before
# logging.
divisor = 1 if method_manages_optimization else grad_accum
metrics = {k: float(v) / divisor for k, v in loss_sums.items()}
metrics.update({k: float(v) / divisor for k, v in metric_sums.items()})
metrics["step_time_sec"] = (time.perf_counter() - t0)
if not method_manages_optimization:
# This is the local training process's wait for next(data_stream).
# Track it without a cross-rank reduction to avoid adding a
# synchronization to every training step.
metrics["dataloader_time_sec"] = dataloader_time_sec
metrics["vsa_sparsity"] = float(self.training_config.vsa_sparsity)
if self.global_rank == 0 and metrics:
self.tracker.log(metrics, step)
with profiler_region("profiler_region_training_callbacks"):
self.callbacks.on_training_step_end(
method,
metrics,
iteration=step,
)
@profile_region("profiler_region_training_train")
def run(
self,
method: TrainingMethod,
@@ -150,84 +258,12 @@ class Trainer:
# Allow method-specific optimization flow (e.g. DiffusionNFT).
method_manages_optimization = bool(method.manages_optimization())
for step in progress:
t0 = time.perf_counter()
# Accumulate on GPU during grad-accum; materialise
# to CPU once per step right before logging.
loss_sums: dict[str, float | torch.Tensor] = {}
metric_sums: dict[str, float | torch.Tensor] = {}
if method_manages_optimization:
loss_map, outputs, step_metrics = method.managed_train_step(
data_stream,
step,
)
for k, v in loss_map.items():
if isinstance(v, torch.Tensor):
loss_sums[k] = v.detach()
for k, v in step_metrics.items():
if k in loss_sums:
raise ValueError(f"Metric key {k!r} collides "
"with loss key. Use a "
"different name (e.g. prefix "
"with 'train/').")
metric_sums[k] = _coerce_log_scalar(
v,
where=("method.managed_train_step()"
f".metrics[{k!r}]"),
)
else:
for accum_iter in range(grad_accum):
batch = next(data_stream)
loss_map, outputs, step_metrics = (method.single_train_step(
batch,
step,
))
method.backward(
loss_map,
outputs,
grad_accum_rounds=grad_accum,
)
for k, v in loss_map.items():
if isinstance(v, torch.Tensor):
prev = loss_sums.get(k, 0.0)
loss_sums[k] = prev + v.detach()
for k, v in step_metrics.items():
if k in loss_sums:
raise ValueError(f"Metric key {k!r} collides "
"with loss key. Use a "
"different name (e.g. prefix "
"with 'train/').")
prev = metric_sums.get(k, 0.0)
metric_sums[k] = (prev + _coerce_log_scalar(
v,
where=("method.single_train_step()"
f".metrics[{k!r}]"),
))
if not method_manages_optimization:
self.callbacks.on_before_optimizer_step(
method,
iteration=step,
)
method.optimizers_schedulers_step(step)
method.optimizers_zero_grad(step)
# Single CPU sync point: materialise GPU tensors
# to float right before logging.
divisor = 1 if method_manages_optimization else grad_accum
metrics = {k: float(v) / divisor for k, v in loss_sums.items()}
metrics.update({k: float(v) / divisor for k, v in metric_sums.items()})
metrics["step_time_sec"] = (time.perf_counter() - t0)
metrics["vsa_sparsity"] = float(tc.vsa_sparsity)
if self.global_rank == 0 and metrics:
self.tracker.log(metrics, step)
self.callbacks.on_training_step_end(
self._run_train_step(
method,
metrics,
iteration=step,
data_stream=data_stream,
step=step,
grad_accum=grad_accum,
method_manages_optimization=method_manages_optimization,
)
if checkpoint_manager is not None:
+2
View File
@@ -23,6 +23,7 @@ from torch.distributed.checkpoint.state_dict import (
from torch.distributed.checkpoint.stateful import Stateful
from fastvideo.logger import init_logger
from fastvideo.profiler import profile_region
logger = init_logger(__name__)
@@ -246,6 +247,7 @@ class CheckpointManager:
return
self.save(step)
@profile_region("profiler_region_training_save_checkpoint")
def save(self, step: int) -> None:
checkpoint_dir = self._checkpoint_dir(step)
dcp_dir = self._dcp_dir(step)
-11
View File
@@ -4,7 +4,6 @@ from typing import Any, cast
import torch
import fastvideo.envs as envs
from fastvideo.distributed import (cleanup_dist_env_and_memory, maybe_init_distributed_environment_and_model_parallel)
from fastvideo.distributed.parallel_state import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
@@ -14,14 +13,6 @@ from fastvideo.pipelines import ForwardBatch, LoRAPipeline, build_pipeline
logger = init_logger(__name__)
def _log_cuda_device_uuid(rank: int, device: torch.device) -> None:
"""Record an NVIDIA worker UUID when external NVTX profiling is enabled."""
if not envs.FASTVIDEO_NVTX_PROFILE:
return
device_uuid = torch.cuda.get_device_properties(device).uuid
logger.info("Worker %d CUDA device UUID: GPU-%s", rank, device_uuid, local_main_process_only=False)
class Worker:
def __init__(self, fastvideo_args: FastVideoArgs, local_rank: int, rank: int, distributed_init_method: str):
@@ -70,8 +61,6 @@ class Worker:
if current_platform.is_cuda_alike():
torch.cuda.set_device(self.device)
self.init_gpu_memory = torch.cuda.mem_get_info(self.device)[0]
if current_platform.is_cuda():
_log_cuda_device_uuid(self.rank, self.device)
else:
# For MPS, we can't get memory info the same way
self.init_gpu_memory = 0
+1 -3
View File
@@ -78,7 +78,7 @@ plugins:
minify_js: true
minify_css: true
cache_safe: true
js_files: [assets/copy-page.js, assets/cookbook.js]
js_files: [assets/copy-page.js]
css_files: [assets/custom.css]
- api-autonav:
modules: ["fastvideo"]
@@ -154,7 +154,6 @@ nav:
- Apple Silicon FastWan: getting_started/installation/mps.md
- Quick Start: getting_started/quick_start.md
- V1 API: getting_started/v1_api.md
- Cookbook: cookbook/index.md
- Inference:
- Quick Start: inference/inference_quick_start.md
- Configuration: inference/configuration.md
@@ -239,4 +238,3 @@ extra_css:
# Custom JavaScript
extra_javascript:
- assets/copy-page.js
- assets/cookbook.js
+3 -5
View File
@@ -122,11 +122,9 @@ fastvideo-kernel = [
# uv reads UV_TORCH_BACKEND / --torch-backend (not a pyproject setting). The
# supported backends (cu126, cu130) both provide torch 2.12.0.
imagebind = { git = "https://github.com/facebookresearch/ImageBind.git", rev = "53680b02d7e37b19b124fa37bae4b6c98c38f5be" }
# FA4 cute. This revision pins nvidia-cutlass-dsl==4.6.0.dev0; holding cutlass-dsl
# back at 4.5.x breaks the CuTe kernels at JIT time, not at import
# (see fastvideo-kernel/README.md).
# torch.compile support comes from FastVideo's own custom_op wrappers.
flash-attn-4 = { git = "https://github.com/Dao-AILab/flash-attention.git", rev = "14c377950125c70b7a9dabf9c561fca53715ac7d", subdirectory = "flash_attn/cute" }
# FA4 cute, pinned to a cutlass-4.5-compatible revision. torch.compile support
# comes from FastVideo's own custom_op wrappers.
flash-attn-4 = { git = "https://github.com/Dao-AILab/flash-attention.git", rev = "82d6441eec5d4dfec120153db2c0145ae855a083", subdirectory = "flash_attn/cute" }
[project.optional-dependencies]
-1
View File
@@ -14,7 +14,6 @@ checkpoint_conversion/
├── convert_gamecraft_weights.py # DiT only
├── convert_gen3c_to_fastvideo.py
├── convert_ltx2_weights.py
├── convert_minimax_h3_adaln_rank.py # Rank-reduces AdaLN in place; same family
├── convert_turbodiffusion_to_diffusers.py
├── convert_turbodiffusion_i2v_to_diffusers.py
├── extract_llava_text_encoder.py # Encoder extraction from a multimodal repo
-83
View File
@@ -1,83 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""In-framework benchmark of the sm_100a VSA backend at real Wan shapes.
Not the powers-of-two bench grid: this builds the metadata FastVideo actually builds, at the
deployed latent shapes, and drives the same block_sparse_attn_from_indices entry point the
model calls -- so it measures the path that will run, including the dispatch and the index
tensors as constructed rather than synthesised.
PYTHONPATH=fastvideo-kernel/python python tests/bench_block_sparse_sm100a.py
"""
import time
import torch
from fastvideo_kernel import block_sparse_attn_sm100a as vsa
from fastvideo_kernel.block_sparse_attn import block_sparse_attn_triton
from fastvideo_kernel.triton_kernels.index import map_to_index
from fastvideo_kernel.vsa_utils import build_vsa_metadata
HEAD_DIM = 128
# (label, latent (T,H,W), tile, heads). The latents are Wan's; heads is per-rank.
CASES = [
("480P", (21, 30, 52), (4, 4, 4), 40),
("720P", (21, 45, 80), (4, 4, 4), 40),
]
def timed(fn, iters=20, warmup=5):
for _ in range(warmup):
fn()
torch.cuda.synchronize()
t0 = time.perf_counter()
for _ in range(iters):
fn()
torch.cuda.synchronize()
return (time.perf_counter() - t0) / iters * 1e3
def main():
if torch.cuda.get_device_capability() != (10, 0):
print("not Blackwell; skipping")
return
print(f"{'case':<8} {'S_pad':>7} {'blocks':>7} {'topk':>5} {'sm100a ms':>10} "
f"{'triton ms':>10} {'speedup':>8} selected")
for label, latent, tile, heads in CASES:
meta = build_vsa_metadata(latent, tile_size=tile, device="cuda")
vbs = meta["variable_block_sizes"].to(torch.int32)
nb = vbs.numel()
block = int(meta["max_block_size"])
S = nb * block
topk = max(1, int(0.1 * nb)) # sparsity 0.9, as deployed
q, k, v = (torch.randn(1, heads, S, HEAD_DIM, device="cuda", dtype=torch.bfloat16)
for _ in range(3))
# index tensors exactly as FastVideo builds them, from a top-k bool map
scores = torch.randn(1, heads, nb, nb, device="cuda")
keep = torch.zeros_like(scores, dtype=torch.bool)
keep.scatter_(-1, scores.topk(topk, dim=-1).indices, True)
idx, num = map_to_index(keep)
idx, num = idx.to(torch.int32).contiguous(), num.to(torch.int32).contiguous()
ok = vsa.is_supported(q, vbs)
ours = timed(lambda: vsa.block_sparse_attn_sm100a(q, k, v, idx, num, vbs)) if ok else float("nan")
# Triton needs 64-token blocks: expand each 128-block into its two halves.
if block == 128: # Triton is 64-granular; expand only when we run 128
keep64 = keep.repeat_interleave(2, dim=-1).repeat_interleave(2, dim=-2)
vbs64 = torch.stack([vbs.clamp(max=64), (vbs - 64).clamp(min=0)], dim=-1).flatten()
i64, n64 = map_to_index(keep64)
i64, n64 = i64.to(torch.int32).contiguous(), n64.to(torch.int32).contiguous()
vbs64 = vbs64.to(torch.int32).contiguous()
else:
keep64, i64, n64, vbs64 = keep, idx, num, vbs
tri = timed(lambda: block_sparse_attn_triton(q, k, v, i64, n64, vbs64))
print(f"{label:<8} {S:>7} {nb:>7} {topk:>5} {ours:>10.3f} {tri:>10.3f} "
f"{tri / ours:>7.2f}x {'sm100a' if ok else 'FALLBACK'}")
if __name__ == "__main__":
main()
@@ -5,9 +5,6 @@ The test compares the exact Transformers base model used by the official
pipeline with FastVideo's production ``TextEncoderLoader`` path. It covers
the three numerical branches the H3 pipelines exercise: text-only tokens,
image features, and video features.
The production encoder returns only the selected layer-50 hidden state, which
is compared bit-exactly with the same state from the official full stack.
"""
from __future__ import annotations
@@ -149,13 +146,13 @@ def _make_cases(root: Path) -> dict[str, dict[str, torch.Tensor]]:
return cases
def _run_reference_cases(
def _run_cases(
model: torch.nn.Module,
cases: dict[str, dict[str, torch.Tensor]],
device: torch.device,
) -> dict[str, torch.Tensor]:
) -> dict[str, tuple[torch.Tensor, ...]]:
dtype = next(model.parameters()).dtype
outputs: dict[str, torch.Tensor] = {}
outputs: dict[str, tuple[torch.Tensor, ...]] = {}
for name, case in cases.items():
inputs = {
key: value.to(device=device, dtype=dtype if key.startswith("pixel_values") else value.dtype)
@@ -169,28 +166,7 @@ def _run_reference_cases(
)
assert result.hidden_states is not None
assert len(result.hidden_states) > MINIMAX_H3_TEXT_ENCODER_LAYER
outputs[name] = result.hidden_states[MINIMAX_H3_TEXT_ENCODER_LAYER][0].detach().cpu()
return outputs
def _run_production_cases(
model: torch.nn.Module,
cases: dict[str, dict[str, torch.Tensor]],
device: torch.device,
) -> dict[str, torch.Tensor]:
dtype = next(model.parameters()).dtype
outputs: dict[str, torch.Tensor] = {}
for name, case in cases.items():
inputs = {
key: value.to(device=device, dtype=dtype if key.startswith("pixel_values") else value.dtype)
for key, value in case.items()
if key not in {"attention_mask", "mm_token_type_ids"}
}
inputs["input_ids"] = inputs["input_ids"][0]
with torch.inference_mode():
result = model(**inputs)
assert result.ndim == 2
outputs[name] = result.detach().cpu()
outputs[name] = tuple(hidden_state.detach().cpu() for hidden_state in result.hidden_states)
return outputs
@@ -235,25 +211,21 @@ def test_minimax_h3_qwen3_vl_parity() -> None:
assert not load_errors, f"Official Qwen3-VL checkpoint did not load strictly: {load_errors}"
official = official_full.model.eval().to(device)
del official_full
expected = _run_reference_cases(official, cases, device)
expected = _run_cases(official, cases, device)
del official
_reclaim_vram()
production = TextEncoderLoader().load(str(root / "text_encoder"), _production_loader_args())
assert getattr(production, "_fastvideo_input_device", device) == device
actual = _run_production_cases(production, cases, device)
actual = _run_cases(production, cases, device)
assert actual.keys() == expected.keys()
for name in expected:
result = actual[name]
reference = expected[name]
assert_close(
result,
reference,
atol=0.0,
rtol=0.0,
msg=lambda message: f"{name} layer {MINIMAX_H3_TEXT_ENCODER_LAYER}: {message}",
)
assert len(actual[name]) == len(expected[name])
for layer, (result, reference) in enumerate(zip(actual[name], expected[name], strict=True)):
assert_close(result, reference, atol=0.0, rtol=0.0, msg=lambda message: f"{name} layer {layer}: {message}")
result = actual[name][MINIMAX_H3_TEXT_ENCODER_LAYER]
reference = expected[name][MINIMAX_H3_TEXT_ENCODER_LAYER]
drift = (result.float() - reference.float()).abs()
print(
f"{name}: max_abs={drift.max().item():.8f} mean_abs={drift.mean().item():.8f}",
@@ -14,7 +14,6 @@
| Qwen3-VL encoder | exact text/image/video hidden states through the production loader | complete |
| FL2VA and Ref2VA DiTs | exact video/audio heads for both model partitions | complete |
| Video VAE | exact encode, normalization, and decode through the production loader | complete |
| Video VAE streaming | exact chunked encode/decode and output-rank-only distributed decode | complete |
| Audio VAE | exact encode and normalization; decode maximum absolute drift `2.4e-7` | complete |
| Video/audio schedulers | pinned `12/3` schedule parity | complete |
| FL2VA packing | pinned row, position, tag, timestep, and RNG parity | complete |
@@ -35,7 +34,6 @@ T2VA, FL2VA, and Ref2VA match the official video/audio latents exactly.
- Load `transformer/` for T2VA/FL2VA and `transformer_ref/` for Ref2VA.
- Keep `last_image`, `references`, and `audio_latents` on the typed request path.
- Treat the published component folders as the loading boundary.
- Keep reference videos on CPU between VAE clips and decode final pixels only on the executor's output rank.
## Evidence boundary

Some files were not shown because too many files have changed in this diff Show More