Compare commits
39
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
0a861031c7 | ||
|
|
f08c5ee8af | ||
|
|
755f4a7967 | ||
|
|
d543a67b10 | ||
|
|
741aa8d289 | ||
|
|
ad9cd63122 | ||
|
|
68e6ffca9e | ||
|
|
64cdcf6be4 | ||
|
|
aa0d98a6b8 | ||
|
|
93b03bc14d | ||
|
|
160f0c9ccf | ||
|
|
e5d1110a0f | ||
|
|
622217ff2a | ||
|
|
cbab605eff | ||
|
|
ac98869aa1 | ||
|
|
56d4a6074f | ||
|
|
9713ea1275 | ||
|
|
37aa382cce | ||
|
|
3f00983287 | ||
|
|
2dc57f4070 | ||
|
|
9df19be719 | ||
|
|
089eea3970 | ||
|
|
dd8447ecc5 | ||
|
|
dca423fd31 | ||
|
|
1b43af8e8e | ||
|
|
628591b620 | ||
|
|
0980ca563f | ||
|
|
b158388733 | ||
|
|
c4ad4227c0 | ||
|
|
a63ccce73d | ||
|
|
ac56806aff | ||
|
|
0462e1b0e7 | ||
|
|
907f2100ec | ||
|
|
e0a3db5651 | ||
|
|
fca45bc8e1 | ||
|
|
942f7db3db | ||
|
|
aadb23f409 | ||
|
|
74b409d7cf | ||
|
|
528cef02c4 |
@@ -6,6 +6,7 @@ on:
|
||||
paths:
|
||||
- 'docs/**'
|
||||
- 'examples/**'
|
||||
- 'scripts/inference/**'
|
||||
- 'mkdocs.yml'
|
||||
- 'requirements-mkdocs.in'
|
||||
- 'requirements-mkdocs.txt'
|
||||
@@ -16,6 +17,7 @@ on:
|
||||
paths:
|
||||
- 'docs/**'
|
||||
- 'examples/**'
|
||||
- 'scripts/inference/**'
|
||||
- 'mkdocs.yml'
|
||||
- 'requirements-mkdocs.in'
|
||||
- 'requirements-mkdocs.txt'
|
||||
|
||||
@@ -0,0 +1,52 @@
|
||||
{
|
||||
"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"
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -0,0 +1,60 @@
|
||||
(() => {
|
||||
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);
|
||||
})();
|
||||
@@ -42,6 +42,46 @@ 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;
|
||||
|
||||
@@ -0,0 +1,44 @@
|
||||
# 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,6 +76,10 @@ 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."
|
||||
|
||||
@@ -2,6 +2,7 @@
|
||||
# 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
|
||||
@@ -19,6 +20,40 @@ 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:
|
||||
@@ -536,6 +571,7 @@ 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!")
|
||||
@@ -549,6 +585,7 @@ 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!")
|
||||
|
||||
@@ -65,6 +65,7 @@ uv pip install flash-attn --no-build-isolation -v
|
||||
|
||||
## Next Steps
|
||||
|
||||
- [Quick Start Guide](quick_start.md) - Get started with your first video generation
|
||||
- [Quick Start](quick_start.md) - Generate your first video
|
||||
- [Inference Cookbook](../cookbook/index.md) - Choose a maintained recipe
|
||||
- [Configuration](../inference/configuration.md) - Learn about configuration options
|
||||
- [Examples](../inference/examples/examples_inference_index.md) - Explore example scripts and notebooks
|
||||
- [Examples](../inference/examples/examples_inference_index.md) - Explore scripts and notebooks
|
||||
|
||||
@@ -23,61 +23,21 @@ Also optionally install flash-attn:
|
||||
uv pip install flash-attn --no-build-isolation -v
|
||||
```
|
||||
|
||||
## Basic Usage
|
||||
## Choose a maintained recipe
|
||||
|
||||
### Text-to-Video Generation
|
||||
The cookbook selects complete, checked-in recipes instead of mixing model,
|
||||
parallelism, offload, and attention settings independently.
|
||||
|
||||
```python
|
||||
from fastvideo import VideoGenerator
|
||||
[Open the inference cookbook](../cookbook/index.md){ .md-button .md-button--primary }
|
||||
|
||||
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()
|
||||
```
|
||||
!!! 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.
|
||||
|
||||
## 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
|
||||
|
||||
@@ -33,6 +33,14 @@ 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!
|
||||
|
||||
@@ -0,0 +1,180 @@
|
||||
# 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,7 +27,8 @@ 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 converted with tools/minimax_h3/fit_adaln_basis.py).
|
||||
# (or a local dir produced by
|
||||
# scripts/checkpoint_conversion/convert_minimax_h3_adaln_rank.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,7 +27,8 @@ 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 converted with tools/minimax_h3/fit_adaln_basis.py).
|
||||
# (or a local dir produced by
|
||||
# scripts/checkpoint_conversion/convert_minimax_h3_adaln_rank.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,7 +24,8 @@ 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 converted with tools/minimax_h3/fit_adaln_basis.py).
|
||||
# (or a local dir produced by
|
||||
# scripts/checkpoint_conversion/convert_minimax_h3_adaln_rank.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.
|
||||
|
||||
@@ -318,6 +318,38 @@ 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}
|
||||
)
|
||||
@@ -333,10 +365,14 @@ 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
|
||||
|
||||
@@ -0,0 +1,7 @@
|
||||
// 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
@@ -0,0 +1,201 @@
|
||||
#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
|
||||
@@ -0,0 +1,114 @@
|
||||
// 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};
|
||||
}
|
||||
@@ -0,0 +1,877 @@
|
||||
// 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,10 +28,31 @@ 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
|
||||
|
||||
|
||||
@@ -0,0 +1,104 @@
|
||||
# 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)
|
||||
+19
-8
@@ -237,7 +237,12 @@ 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)
|
||||
qkT = tl.dot(k, qT)
|
||||
# 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)
|
||||
pT = tl.math.exp2(qkT - m[None, :])
|
||||
mask = tl.arange(0, BLOCK_N1) < block_size
|
||||
pT = tl.where(mask[:, None], pT, 0.0)
|
||||
@@ -268,6 +273,7 @@ def _attn_bwd_dq(
|
||||
do,
|
||||
m,
|
||||
D,
|
||||
sm_scale,
|
||||
# shared by Q/K/V/DO.
|
||||
q2k_index,
|
||||
q2k_num,
|
||||
@@ -315,7 +321,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)
|
||||
qk = tl.dot(q, kT) * (sm_scale * 1.4426950408889634)
|
||||
p = tl.math.exp2(qk - m)
|
||||
offs_in_block = half * step_n + tl.arange(0, BLOCK_N2)
|
||||
mask = offs_in_block < block_size
|
||||
@@ -324,8 +330,7 @@ def _attn_bwd_dq(
|
||||
dp = tl.dot(do, vT).to(tl.float32)
|
||||
ds = p * (dp - Di[:, None])
|
||||
ds = ds.to(tl.bfloat16)
|
||||
# Compute dQ.
|
||||
# NOTE: We need to de-scale dq in the end, because kT was pre-scaled.
|
||||
# Compute dQ (kT is raw; the caller applies sm_scale once at the end).
|
||||
dq += tl.dot(ds, tl.trans(kT))
|
||||
# Increment pointers.
|
||||
return dq
|
||||
@@ -453,6 +458,7 @@ def _attn_bwd(
|
||||
do,
|
||||
m,
|
||||
D, #
|
||||
sm_scale,
|
||||
q2k_index,
|
||||
q2k_num,
|
||||
max_kv_blks,
|
||||
@@ -470,7 +476,7 @@ def _attn_bwd(
|
||||
)
|
||||
# Write back dQ.
|
||||
dq_ptrs = DQ + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d
|
||||
dq *= LN2
|
||||
dq *= sm_scale
|
||||
tl.store(dq_ptrs, dq)
|
||||
|
||||
|
||||
@@ -591,6 +597,7 @@ def _attn_bwd_dq_kernel(
|
||||
Q,
|
||||
K,
|
||||
V,
|
||||
sm_scale,
|
||||
DO, #
|
||||
DQ,
|
||||
M,
|
||||
@@ -663,6 +670,7 @@ def _attn_bwd_dq_kernel(
|
||||
do,
|
||||
m,
|
||||
D,
|
||||
sm_scale,
|
||||
q2k_index,
|
||||
q2k_num,
|
||||
max_kv_blks,
|
||||
@@ -680,7 +688,7 @@ def _attn_bwd_dq_kernel(
|
||||
)
|
||||
|
||||
dq_ptrs = DQ + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d
|
||||
dq_acc *= LN2
|
||||
dq_acc *= sm_scale
|
||||
tl.store(dq_ptrs, dq_acc)
|
||||
|
||||
|
||||
@@ -748,9 +756,11 @@ 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
|
||||
RCP_LN2 = 1.4426950408889634 # = 1.0 / ln(2)
|
||||
# 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.)
|
||||
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)
|
||||
@@ -813,6 +823,7 @@ 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,7 +14,9 @@ import math
|
||||
import torch
|
||||
|
||||
VSA_TILE_SIZE = (4, 4, 4)
|
||||
_SUPPORTED_VSA_BLOCK_VOLUMES = (64, 256)
|
||||
# 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)
|
||||
|
||||
|
||||
def _canonicalize_device(device: torch.device | str) -> torch.device:
|
||||
|
||||
@@ -0,0 +1,104 @@
|
||||
"""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}"
|
||||
@@ -19,7 +19,12 @@ from fastvideo.attention.backends.abstract import (
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
logger.info("Using FlashAttention-%s backend", fa_version)
|
||||
# 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)
|
||||
|
||||
# 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,13 +5,17 @@ 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 (4,8,8) video tiles]``;
|
||||
prefix tiles never straddle segment boundaries.
|
||||
- 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``).
|
||||
- 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 H3
|
||||
checkpoint does not carry: the loader zero-initializes it, so untrained
|
||||
inference is exactly pure sparse and finetuning can learn the gate.
|
||||
- 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.
|
||||
- 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,
|
||||
@@ -20,22 +24,46 @@ 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.
|
||||
|
||||
Targets sm10.x through the FA4 CuTe 256-tile path
|
||||
At tile 256 this 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.
|
||||
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.
|
||||
"""
|
||||
|
||||
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)
|
||||
@@ -43,51 +71,115 @@ 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
|
||||
|
||||
VSA_H3_TILE_SIZE = (4, 8, 8) # 256 elements -> FA4 CuTe fastpath on sm10.x
|
||||
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)
|
||||
_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) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
def token_tile_and_valid(variable_block_sizes: torch.Tensor,
|
||||
tile_elems: int = _TILE_ELEMS) -> 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 = VSA_H3_TILE_SIZE
|
||||
ts_t, ts_h, ts_w = tile_shape
|
||||
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, VSA_H3_TILE_SIZE)
|
||||
video_sizes = construct_variable_block_sizes(dit_seq_shape, num_tiles, device, tile_shape)
|
||||
num_video_tiles = int(video_sizes.numel())
|
||||
|
||||
video_indices = get_tile_partition_indices(dit_seq_shape, VSA_H3_TILE_SIZE, device) + prefix_len
|
||||
video_indices = get_tile_partition_indices(dit_seq_shape, tile_shape, device) + prefix_len
|
||||
tile_partition_indices = torch.cat([
|
||||
torch.arange(prefix_len, device=device, dtype=torch.long),
|
||||
video_indices,
|
||||
@@ -100,9 +192,11 @@ 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)
|
||||
|
||||
|
||||
@@ -139,6 +233,9 @@ 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
|
||||
@@ -158,24 +255,28 @@ 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, ...] = (),
|
||||
**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, ...] = (),
|
||||
tile_size: int = _TILE_ELEMS,
|
||||
**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)
|
||||
num_video_tiles) = _h3_tile_geometry(prefix_segments, dit_seq_shape, device, tile_shape)
|
||||
|
||||
return MiniMaxH3VSAMetadata(
|
||||
current_timestep=current_timestep,
|
||||
@@ -186,13 +287,14 @@ 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) -> torch.Tensor:
|
||||
"""fp32 mean over each 256-token tile. x: [B, S_pad, H, D] -> [B, H, n_tiles, D].
|
||||
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].
|
||||
|
||||
Pad positions in the tile buffer are guaranteed zero (zeros-init, never
|
||||
written), so a plain sum with fp32 accumulation needs no validity mask
|
||||
@@ -200,8 +302,8 @@ def _pool_tiles(x: torch.Tensor, variable_block_sizes: torch.Tensor) -> torch.Te
|
||||
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)
|
||||
|
||||
@@ -232,6 +334,24 @@ 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__(
|
||||
@@ -259,7 +379,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 * _TILE_ELEMS, x.shape[-2], x.shape[-1])
|
||||
target_shape = (x.shape[0], n_tiles * attn_metadata.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
|
||||
@@ -281,7 +401,11 @@ class MiniMaxH3VSAImpl(AttentionImpl):
|
||||
gate_compress: torch.Tensor | None,
|
||||
attn_metadata: MiniMaxH3VSAMetadata,
|
||||
) -> torch.Tensor:
|
||||
if block_sparse_attn_256_bshd is None:
|
||||
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:
|
||||
raise NotImplementedError("fastvideo_kernel.block_sparse_attn_256 is not installed")
|
||||
|
||||
# probe-guided per-layer opt-out: diffuse layers run dense (all-True
|
||||
@@ -291,8 +415,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)
|
||||
k_pooled = _pool_tiles(key, attn_metadata.variable_block_sizes)
|
||||
q_pooled = _pool_tiles(query, attn_metadata.variable_block_sizes, tile_elems)
|
||||
k_pooled = _pool_tiles(key, attn_metadata.variable_block_sizes, tile_elems)
|
||||
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)
|
||||
@@ -309,14 +433,66 @@ class MiniMaxH3VSAImpl(AttentionImpl):
|
||||
attn_metadata.exempt,
|
||||
)
|
||||
|
||||
out, _ = block_sparse_attn_256_bshd(query, key, value, mask, attn_metadata.variable_block_sizes)
|
||||
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)
|
||||
|
||||
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)
|
||||
v_pooled = _pool_tiles(value, attn_metadata.variable_block_sizes, tile_elems)
|
||||
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
|
||||
@@ -325,7 +501,7 @@ class MiniMaxH3VSAImpl(AttentionImpl):
|
||||
# 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_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)
|
||||
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)
|
||||
token_tile, token_valid = token_tile_and_valid(attn_metadata.variable_block_sizes, attn_metadata.tile_elems)
|
||||
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,6 +62,8 @@ 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
|
||||
@@ -107,7 +109,7 @@ class MiniMaxH3Qwen3VLArchConfig(TextEncoderArchConfig):
|
||||
vision_initializer_range: float = 0.02
|
||||
vision_deepstack_visual_indexes: tuple[int, ...] = (8, 16, 24)
|
||||
|
||||
output_hidden_states: bool = True
|
||||
output_hidden_states: bool = False
|
||||
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,
|
||||
@@ -118,6 +120,17 @@ 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:
|
||||
|
||||
@@ -21,12 +21,17 @@ 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
|
||||
@@ -217,10 +222,34 @@ 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":
|
||||
|
||||
@@ -146,6 +146,19 @@ 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).
|
||||
@@ -169,6 +182,7 @@ 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
|
||||
@@ -286,8 +300,27 @@ 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``.
|
||||
|
||||
@@ -631,6 +664,18 @@ 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,
|
||||
@@ -644,6 +689,12 @@ 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(
|
||||
|
||||
+4
-1
@@ -114,7 +114,10 @@ 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):
|
||||
logger.log(logging.INFO, msg, *args, stacklevel=2, **kwargs)
|
||||
# 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)
|
||||
|
||||
global _warned_local_main_process, _warned_main_process
|
||||
|
||||
|
||||
@@ -10,6 +10,7 @@ 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
|
||||
@@ -23,12 +24,50 @@ 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):
|
||||
@@ -62,6 +101,7 @@ 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(
|
||||
@@ -78,11 +118,15 @@ 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)
|
||||
hidden_states, gate = hidden_states.chunk(2, dim=-1)
|
||||
hidden_states = hidden_states * F.silu(gate)
|
||||
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, _ = self.fc_out(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
@@ -99,6 +143,7 @@ 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
|
||||
@@ -134,6 +179,7 @@ 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
|
||||
@@ -211,11 +257,18 @@ 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))
|
||||
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)
|
||||
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)
|
||||
|
||||
# H3 rotates only 96/128 channels, which the generic `freqs_cis`
|
||||
# branch cannot express. Apply it above, then pass no RoPE here.
|
||||
@@ -397,6 +450,9 @@ 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)
|
||||
@@ -408,6 +464,7 @@ 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(
|
||||
@@ -415,6 +472,7 @@ 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,
|
||||
@@ -423,6 +481,7 @@ class MiniMaxH3TransformerBlock(nn.Module):
|
||||
prefix=f"{prefix}.adaln_proj",
|
||||
apply_silu=adaln_apply_silu,
|
||||
)
|
||||
self.fuse_modulate = fuse_modulate
|
||||
|
||||
def forward(
|
||||
self,
|
||||
@@ -435,19 +494,39 @@ 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))
|
||||
|
||||
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)
|
||||
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)
|
||||
attention_output = self.attn(norm_hidden_states, rotary_emb, original_seq_len)
|
||||
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)
|
||||
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)
|
||||
feed_forward_output = self.ff(norm_hidden_states)
|
||||
return residual + gate_mlp.index_select(0, adaln_indices) * feed_forward_output
|
||||
return hidden_states + gate_mlp.index_select(0, adaln_indices) * feed_forward_output
|
||||
|
||||
|
||||
class MiniMaxH3Transformer3DModel(BaseDiT):
|
||||
@@ -493,6 +572,17 @@ 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 "
|
||||
@@ -546,7 +636,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 "
|
||||
"tools/minimax_h3/fit_adaln_basis.py.")
|
||||
"scripts/checkpoint_conversion/convert_minimax_h3_adaln_rank.py.")
|
||||
adaln_dim = self.adaln_rank or arch.time_embed_dim
|
||||
self.adaln_basis = ReplicatedLinear(
|
||||
arch.time_embed_dim,
|
||||
@@ -590,6 +680,9 @@ 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(
|
||||
@@ -616,6 +709,20 @@ 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,
|
||||
@@ -734,14 +841,17 @@ class MiniMaxH3Transformer3DModel(BaseDiT):
|
||||
local_timestep_indices, _ = sequence_model_parallel_shard(local_timestep_indices, dim=0)
|
||||
rotary_emb = (rotary_cos, rotary_sin)
|
||||
|
||||
for block in self.transformer_blocks:
|
||||
packed_hidden_states = block(
|
||||
packed_hidden_states,
|
||||
temb,
|
||||
adaln_indices,
|
||||
rotary_emb,
|
||||
original_seq_len,
|
||||
)
|
||||
# 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,
|
||||
)
|
||||
|
||||
packed_hidden_states = self.norm_out(
|
||||
packed_hidden_states,
|
||||
|
||||
@@ -0,0 +1,20 @@
|
||||
# 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",
|
||||
]
|
||||
@@ -0,0 +1,302 @@
|
||||
# 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)
|
||||
@@ -0,0 +1,174 @@
|
||||
# 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"]
|
||||
@@ -0,0 +1,104 @@
|
||||
# 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"]
|
||||
@@ -1,6 +1,7 @@
|
||||
# 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
|
||||
@@ -8,11 +9,16 @@ 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__()
|
||||
@@ -23,13 +29,7 @@ class TextEncoder(nn.Module, ABC):
|
||||
raise ValueError(f"Subclass {self.__class__.__name__} must define _supported_attention_backends")
|
||||
|
||||
@abstractmethod
|
||||
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:
|
||||
def forward(self, *args: Any, **kwargs: Any) -> TextEncoderOutputT:
|
||||
pass
|
||||
|
||||
@property
|
||||
|
||||
@@ -0,0 +1,453 @@
|
||||
# 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,10 +227,15 @@ 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(config.num_hidden_layers))
|
||||
self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
|
||||
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)
|
||||
self.rotary_emb = MiniMaxH3Qwen3VLTextRotaryEmbedding(config)
|
||||
|
||||
def forward(
|
||||
@@ -238,18 +243,14 @@ 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,
|
||||
) -> BaseEncoderOutput:
|
||||
) -> torch.Tensor:
|
||||
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:
|
||||
@@ -258,10 +259,9 @@ 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
|
||||
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)
|
||||
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}]")
|
||||
|
||||
|
||||
class MiniMaxH3Qwen3VLVisionPatchEmbed(nn.Module):
|
||||
@@ -499,10 +499,18 @@ class MiniMaxH3Qwen3VLVisionModel(nn.Module):
|
||||
return self.merger(hidden_states), deepstack_features
|
||||
|
||||
|
||||
class MiniMaxH3Qwen3VLConditioner(TextEncoder):
|
||||
"""FastVideo-native Qwen3-VL body without the unused language-model head."""
|
||||
class MiniMaxH3Qwen3VLConditioner(TextEncoder[torch.Tensor]):
|
||||
"""H3 conditioner returning the unnormalized layer-50 hidden tensor."""
|
||||
|
||||
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)
|
||||
@@ -518,6 +526,10 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder):
|
||||
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,
|
||||
@@ -610,35 +622,39 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder):
|
||||
f"tokens={int(mask.sum())}, features={features.shape[0]}")
|
||||
return mask
|
||||
|
||||
def forward(
|
||||
# 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(
|
||||
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,
|
||||
input_ids: torch.Tensor,
|
||||
*,
|
||||
pixel_values: torch.Tensor | None = None,
|
||||
pixel_values_videos: 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,
|
||||
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")
|
||||
) -> 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)
|
||||
|
||||
image_mask = None
|
||||
video_mask = None
|
||||
image_deepstack = None
|
||||
video_deepstack = None
|
||||
if pixel_values is not None:
|
||||
if input_ids is None or image_grid_thw is None:
|
||||
if 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)
|
||||
@@ -646,7 +662,7 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder):
|
||||
"image")
|
||||
inputs_embeds = inputs_embeds.masked_scatter(image_mask.unsqueeze(-1), image_features)
|
||||
if pixel_values_videos is not None:
|
||||
if input_ids is None or video_grid_thw is None:
|
||||
if 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)
|
||||
@@ -674,25 +690,34 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder):
|
||||
visual_mask = video_mask
|
||||
deepstack_features = video_deepstack
|
||||
|
||||
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(
|
||||
position_ids = self._get_rope_index(input_ids, image_grid_thw, video_grid_thw, None)
|
||||
hidden_states = self.language_model(
|
||||
inputs_embeds,
|
||||
position_ids,
|
||||
attention_mask,
|
||||
output_hidden_states,
|
||||
None,
|
||||
visual_mask,
|
||||
deepstack_features,
|
||||
)
|
||||
outputs.attention_mask = attention_mask
|
||||
return outputs
|
||||
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,
|
||||
)
|
||||
|
||||
def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:
|
||||
parameters = dict(self.named_parameters())
|
||||
@@ -702,6 +727,8 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder):
|
||||
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]
|
||||
@@ -710,7 +737,23 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder):
|
||||
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"]
|
||||
__all__ = [
|
||||
"MiniMaxH3Qwen3VLConditioner",
|
||||
"MiniMaxH3SerializedFP8Config",
|
||||
]
|
||||
|
||||
@@ -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 cast
|
||||
from typing import Any, cast
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
@@ -30,9 +30,13 @@ 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,
|
||||
@@ -347,22 +351,46 @@ 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,
|
||||
@@ -381,11 +409,20 @@ 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):
|
||||
loaded_weights: set[str] = model.load_weights(
|
||||
safetensors_weights_iterator(
|
||||
[fastvideo_args.override_text_encoder_safetensors],
|
||||
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],
|
||||
to_cpu=use_cpu_offload,
|
||||
)) # type: ignore
|
||||
)
|
||||
loaded_weights: set[str] = model.load_weights(override_weights) # type: ignore
|
||||
else:
|
||||
loaded_weights: set[str] = model.load_weights(
|
||||
self._get_all_weights(
|
||||
@@ -400,6 +437,10 @@ 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)
|
||||
|
||||
@@ -442,7 +483,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:
|
||||
if weights_not_loaded and (model_config.quant_config is None or checkpoint_quant_config is not None):
|
||||
raise ValueError("Following weights were not initialized from "
|
||||
f"checkpoint: {weights_not_loaded}")
|
||||
|
||||
@@ -1057,7 +1098,12 @@ 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)
|
||||
logger.info("transformer attention backend: %s", resolved.name if resolved else "automatic selection")
|
||||
# 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)
|
||||
model = maybe_load_fsdp_model(
|
||||
model_cls=model_cls,
|
||||
init_params={
|
||||
|
||||
@@ -0,0 +1,127 @@
|
||||
# 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,9 +392,16 @@ 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
|
||||
|
||||
@@ -0,0 +1,378 @@
|
||||
# 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",
|
||||
]
|
||||
@@ -7,6 +7,7 @@ This module intentionally uses only PyTorch and FastVideo configuration types.
|
||||
"""
|
||||
|
||||
import math
|
||||
from collections.abc import Iterator
|
||||
from dataclasses import dataclass
|
||||
|
||||
import torch
|
||||
@@ -14,7 +15,10 @@ 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:
|
||||
@@ -291,6 +295,7 @@ 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
|
||||
@@ -302,12 +307,34 @@ 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))
|
||||
@@ -328,9 +355,17 @@ 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)
|
||||
|
||||
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)
|
||||
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)
|
||||
return self.to_out[0](hidden_states)
|
||||
|
||||
|
||||
@@ -433,6 +468,7 @@ 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,
|
||||
@@ -482,6 +518,11 @@ 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."""
|
||||
|
||||
@@ -489,6 +530,7 @@ 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__()
|
||||
@@ -654,12 +696,15 @@ 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 = []
|
||||
@@ -676,6 +721,12 @@ 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))
|
||||
@@ -699,36 +750,64 @@ 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]
|
||||
return self._stitch_tiles(rows, latent_y_overlaps, latent_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()
|
||||
|
||||
def _decode_clip(self, z: torch.Tensor) -> torch.Tensor:
|
||||
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)
|
||||
"""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()
|
||||
|
||||
def _encode(self, x: torch.Tensor) -> torch.Tensor:
|
||||
clip_length = self.config.clip_length
|
||||
@@ -747,43 +826,157 @@ class AutoencoderKLMiniMaxH3(nn.Module):
|
||||
moments = moments[:, :, :-self.config.token_drop]
|
||||
return moments
|
||||
|
||||
def _decode(self, z: torch.Tensor) -> torch.Tensor:
|
||||
tokens_chunk_size = self.tokens_chunk_size
|
||||
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
|
||||
tokens_chunk_size = self.tokens_chunk_size
|
||||
temporal_ratio = self.temporal_compression_ratio
|
||||
chunk_num_frames = tokens_chunk_size * temporal_ratio
|
||||
num_tokens = z.shape[2] + token_drop
|
||||
num_tokens = latent_num_frames + 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)
|
||||
|
||||
decoded_chunks = []
|
||||
output_frame_start = 0
|
||||
overlap = None
|
||||
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:
|
||||
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]
|
||||
if overlap is not None:
|
||||
chunk = self._blend(overlap, chunk, self.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)
|
||||
num_frames = min(chunk.shape[2], output_num_frames - output_frame_start)
|
||||
chunk = chunk[:, :, :num_frames]
|
||||
|
||||
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
|
||||
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)
|
||||
|
||||
def encode(
|
||||
self,
|
||||
@@ -799,6 +992,34 @@ 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,
|
||||
@@ -825,6 +1046,26 @@ 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,25 +42,6 @@ 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,
|
||||
@@ -155,20 +136,10 @@ class MiniMaxH3ConditioningStage(PipelineStage):
|
||||
device: torch.device,
|
||||
**vision_inputs: torch.Tensor | None,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
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,
|
||||
)
|
||||
input_ids = torch.tensor(token_ids, dtype=torch.long, device=device)
|
||||
dtype = self.conditioner.dtype
|
||||
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,
|
||||
prompt_embeds = self.conditioner(
|
||||
input_ids,
|
||||
**{
|
||||
name:
|
||||
None if value is None else value.to(
|
||||
@@ -178,10 +149,10 @@ class MiniMaxH3ConditioningStage(PipelineStage):
|
||||
for name, value in vision_inputs.items()
|
||||
},
|
||||
)
|
||||
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}]`.")
|
||||
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)}")
|
||||
return (
|
||||
outputs.hidden_states[hidden_state_index].to(device=device, dtype=dtype),
|
||||
prompt_embeds.unsqueeze(0).to(device=device, dtype=dtype),
|
||||
torch.tensor(token_tags, dtype=torch.long),
|
||||
)
|
||||
|
||||
@@ -286,6 +257,7 @@ 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
|
||||
@@ -293,10 +265,13 @@ class MiniMaxH3ConditioningStage(PipelineStage):
|
||||
if moved_for_forward:
|
||||
self.conditioner.to(device)
|
||||
try:
|
||||
if self.ref2va:
|
||||
prompt_embeds, text_token_tags = self._encode_ref2va(batch, device)
|
||||
else:
|
||||
prompt_embeds, text_token_tags = self._encode_fl2va(batch, device)
|
||||
# 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)
|
||||
finally:
|
||||
if moved_for_forward:
|
||||
self.conditioner.to("cpu")
|
||||
|
||||
@@ -7,10 +7,13 @@ from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.distributed import get_local_torch_device, get_sp_group, model_parallel_is_initialized
|
||||
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,
|
||||
@@ -21,6 +24,9 @@ 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:
|
||||
@@ -30,6 +36,23 @@ 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."""
|
||||
|
||||
@@ -54,6 +77,16 @@ 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.")
|
||||
@@ -71,13 +104,33 @@ class MiniMaxH3VideoDecodingStage(PipelineStage):
|
||||
try:
|
||||
latents = self.vae.denormalize_latents(latents.to(device=device, dtype=torch.float32))
|
||||
if fastvideo_args.output_type == "latent":
|
||||
batch.output = latents.detach().float().cpu()
|
||||
# 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
|
||||
return batch
|
||||
|
||||
# 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()
|
||||
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
|
||||
return batch
|
||||
finally:
|
||||
if fastvideo_args.vae_cpu_offload:
|
||||
@@ -107,6 +160,15 @@ 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.")
|
||||
@@ -124,7 +186,10 @@ class MiniMaxH3AudioDecodingStage(PipelineStage):
|
||||
self._clear_runtime(batch)
|
||||
return batch
|
||||
|
||||
decoded = self.audio_vae.decode(latents).sample.float()
|
||||
# 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()
|
||||
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,6 +89,7 @@ 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.")
|
||||
@@ -145,9 +146,15 @@ 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:
|
||||
with profiler_region("inference_denoising"):
|
||||
# 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"):
|
||||
for index, (video_timestep,
|
||||
audio_timestep) in enumerate(zip(video_timesteps, audio_timesteps, strict=True)):
|
||||
unique_timesteps, timestep_indices = row_timestep_plan[index]
|
||||
@@ -167,6 +174,7 @@ 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,8 +9,10 @@ import numpy as np
|
||||
import torch
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.distributed import get_local_torch_device, get_sp_group, model_parallel_is_initialized
|
||||
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,
|
||||
@@ -36,6 +38,8 @@ 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"
|
||||
|
||||
|
||||
@@ -105,8 +109,20 @@ 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":
|
||||
@@ -119,9 +135,11 @@ 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(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
|
||||
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
|
||||
latents = self.vae.normalize_latents(_sample_visual_posterior(posterior).to(
|
||||
torch.float16).float()).cpu()
|
||||
reference.num_latent_frames = int(latents.shape[2])
|
||||
@@ -202,7 +220,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)
|
||||
video_rows = self._encode_visual_rows(references, vae_device, fastvideo_args)
|
||||
finally:
|
||||
if fastvideo_args.vae_cpu_offload:
|
||||
self.vae.to("cpu")
|
||||
|
||||
@@ -38,6 +38,25 @@ 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."""
|
||||
|
||||
@@ -5,20 +5,29 @@ 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, token_tile_and_valid)
|
||||
_pool_tiles, _validate_h3_tile_geometry,
|
||||
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):
|
||||
def _build(spec, sparsity=0.0, device=_CPU, tile_size=_TILE_ELEMS):
|
||||
return MiniMaxH3VSAMetadataBuilder().build(
|
||||
current_timestep=0,
|
||||
raw_latent_shape=spec["raw_latent_shape"],
|
||||
@@ -26,6 +35,7 @@ def _build(spec, sparsity=0.0, device=_CPU):
|
||||
VSA_sparsity=sparsity,
|
||||
prefix_segments=spec["prefix_segments"],
|
||||
device=device,
|
||||
tile_size=tile_size,
|
||||
)
|
||||
|
||||
|
||||
@@ -36,7 +46,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)
|
||||
token_tile, token_valid = token_tile_and_valid(meta.variable_block_sizes, meta.tile_elems)
|
||||
out = torch.empty_like(query)
|
||||
for b in range(query.shape[0]):
|
||||
for h in range(query.shape[2]):
|
||||
@@ -133,9 +143,117 @@ 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")
|
||||
|
||||
@@ -0,0 +1,181 @@
|
||||
# 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,6 +15,12 @@ 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
|
||||
@@ -105,3 +111,73 @@ 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),
|
||||
]
|
||||
|
||||
@@ -0,0 +1,281 @@
|
||||
# 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
|
||||
@@ -0,0 +1,187 @@
|
||||
# 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))])
|
||||
@@ -0,0 +1,185 @@
|
||||
# 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)
|
||||
@@ -0,0 +1,231 @@
|
||||
# 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()
|
||||
@@ -0,0 +1,173 @@
|
||||
# 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,
|
||||
)
|
||||
@@ -0,0 +1,192 @@
|
||||
# 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)
|
||||
@@ -0,0 +1,73 @@
|
||||
# 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)
|
||||
@@ -0,0 +1,314 @@
|
||||
# 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")
|
||||
@@ -0,0 +1,110 @@
|
||||
# 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)
|
||||
@@ -0,0 +1,428 @@
|
||||
# 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)
|
||||
@@ -0,0 +1,168 @@
|
||||
# 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,11 +1,40 @@
|
||||
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
|
||||
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,
|
||||
)
|
||||
|
||||
|
||||
def _worker_returning(output_batch: ForwardBatch) -> Worker:
|
||||
|
||||
@@ -4,6 +4,7 @@ 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
|
||||
@@ -13,6 +14,14 @@ 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):
|
||||
@@ -61,6 +70,8 @@ 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
|
||||
|
||||
+3
-1
@@ -78,7 +78,7 @@ plugins:
|
||||
minify_js: true
|
||||
minify_css: true
|
||||
cache_safe: true
|
||||
js_files: [assets/copy-page.js]
|
||||
js_files: [assets/copy-page.js, assets/cookbook.js]
|
||||
css_files: [assets/custom.css]
|
||||
- api-autonav:
|
||||
modules: ["fastvideo"]
|
||||
@@ -154,6 +154,7 @@ 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
|
||||
@@ -238,3 +239,4 @@ extra_css:
|
||||
# Custom JavaScript
|
||||
extra_javascript:
|
||||
- assets/copy-page.js
|
||||
- assets/cookbook.js
|
||||
|
||||
@@ -14,6 +14,7 @@ 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
|
||||
|
||||
+1
-1
@@ -24,7 +24,7 @@ directly: ``adaln_rank`` in ``config.json`` flows onto the arch config through
|
||||
|
||||
Usage::
|
||||
|
||||
python tools/minimax_h3/fit_adaln_basis.py \
|
||||
python scripts/checkpoint_conversion/convert_minimax_h3_adaln_rank.py \
|
||||
--src /path/to/MiniMax-H3/transformer \
|
||||
--dst /path/to/MiniMax-H3-r16/transformer --rank 16
|
||||
|
||||
@@ -0,0 +1,83 @@
|
||||
# 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,6 +5,9 @@ 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
|
||||
@@ -146,13 +149,13 @@ def _make_cases(root: Path) -> dict[str, dict[str, torch.Tensor]]:
|
||||
return cases
|
||||
|
||||
|
||||
def _run_cases(
|
||||
def _run_reference_cases(
|
||||
model: torch.nn.Module,
|
||||
cases: dict[str, dict[str, torch.Tensor]],
|
||||
device: torch.device,
|
||||
) -> dict[str, tuple[torch.Tensor, ...]]:
|
||||
) -> dict[str, torch.Tensor]:
|
||||
dtype = next(model.parameters()).dtype
|
||||
outputs: dict[str, tuple[torch.Tensor, ...]] = {}
|
||||
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)
|
||||
@@ -166,7 +169,28 @@ def _run_cases(
|
||||
)
|
||||
assert result.hidden_states is not None
|
||||
assert len(result.hidden_states) > MINIMAX_H3_TEXT_ENCODER_LAYER
|
||||
outputs[name] = tuple(hidden_state.detach().cpu() for hidden_state in result.hidden_states)
|
||||
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()
|
||||
return outputs
|
||||
|
||||
|
||||
@@ -211,21 +235,25 @@ 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_cases(official, cases, device)
|
||||
expected = _run_reference_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_cases(production, cases, device)
|
||||
actual = _run_production_cases(production, cases, device)
|
||||
|
||||
assert actual.keys() == expected.keys()
|
||||
for name in expected:
|
||||
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]
|
||||
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}",
|
||||
)
|
||||
drift = (result.float() - reference.float()).abs()
|
||||
print(
|
||||
f"{name}: max_abs={drift.max().item():.8f} mean_abs={drift.mean().item():.8f}",
|
||||
|
||||
@@ -14,6 +14,7 @@
|
||||
| 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 |
|
||||
@@ -34,6 +35,7 @@ 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
|
||||
|
||||
|
||||
@@ -12,6 +12,14 @@ FastVideo-owned unit contracts belong under `fastvideo/tests/`.
|
||||
The reference helper verifies the pinned source and import origin. A missing checkout may skip a source-parity module;
|
||||
that skip is not parity evidence.
|
||||
|
||||
## FastVideo unit contracts
|
||||
|
||||
```bash
|
||||
pytest \
|
||||
fastvideo/tests/vaes/test_minimax_h3_video_vae_streaming.py \
|
||||
fastvideo/tests/stages/test_minimax_h3_vae_streaming.py -q
|
||||
```
|
||||
|
||||
## Registry smoke
|
||||
|
||||
```bash
|
||||
@@ -49,7 +57,34 @@ pytest \
|
||||
```
|
||||
|
||||
With a gate enabled, missing CUDA, source, or weights is a failure. Recorded component evidence is exact for both DiT
|
||||
partitions, the video VAE, and all Qwen3-VL hidden states; audio decode has maximum absolute drift `2.4e-7`.
|
||||
partitions and the video VAE; audio decode has maximum absolute drift `2.4e-7`. The encoder gate compares the slim
|
||||
forward's selected layer-50 hidden state bit-exactly against the same state from the official full stack across text,
|
||||
image, and video inputs.
|
||||
|
||||
The video VAE test verifies the reference checkout at commit
|
||||
`abc5e9bf71fd38f53cd471bc3acaa84bc5ecbfdc` and compares the production CPU `uint8` `encode_pixels()` path against
|
||||
the official posterior element by element.
|
||||
|
||||
## Video VAE memory benchmark
|
||||
|
||||
The benchmark uses one warmup and three measured runs with `vae_cpu_offload=True`. It reports absolute and
|
||||
stage-incremental allocated/reserved CUDA peaks for every rank. For SP runs, the reported aggregate is explicitly the
|
||||
sum of rank-local maxima, not a simultaneous node peak.
|
||||
|
||||
```bash
|
||||
python tests/local_tests/vaes/benchmark_minimax_h3_video_vae_memory.py \
|
||||
--source-root "$PWD" --model-root "$MINIMAX_H3_MODEL_ROOT" \
|
||||
--revision-label candidate --operation encode
|
||||
|
||||
python -m torch.distributed.run --nproc_per_node=4 \
|
||||
tests/local_tests/vaes/benchmark_minimax_h3_video_vae_memory.py \
|
||||
--source-root "$PWD" --model-root "$MINIMAX_H3_MODEL_ROOT" \
|
||||
--revision-label candidate-sp4 --operation decode
|
||||
```
|
||||
|
||||
Run the same script with `--source-root` pointed at the base checkout for a comparable baseline. The default workload
|
||||
is deterministic `124 x 768 x 1344` video geometry with seed `20260803`; the JSON record includes source/model
|
||||
revisions, software/allocator metadata, exact measurement boundaries, per-repetition values, and output shapes.
|
||||
|
||||
FastVideo joint audio/video generation and SP=1/SP=4 latent consistency have been validated. T2VA, FL2VA, and
|
||||
Ref2VA video/audio latents match the pinned Diffusers pipeline exactly.
|
||||
|
||||
@@ -3,6 +3,7 @@
|
||||
import hashlib
|
||||
import inspect
|
||||
import os
|
||||
import subprocess
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
@@ -13,6 +14,28 @@ REFERENCE_SRC = REFERENCE_ROOT / "src"
|
||||
PINNED_COMMIT = "abc5e9bf71fd38f53cd471bc3acaa84bc5ecbfdc"
|
||||
|
||||
|
||||
def _run_git(*args: str) -> str:
|
||||
try:
|
||||
result = subprocess.run(
|
||||
["git", "-C", str(REFERENCE_ROOT), *args],
|
||||
check=True,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
)
|
||||
except (FileNotFoundError, subprocess.CalledProcessError) as error:
|
||||
raise RuntimeError(f"Could not verify the MiniMax-H3 reference checkout at {REFERENCE_ROOT}.") from error
|
||||
return result.stdout.strip()
|
||||
|
||||
|
||||
def assert_reference_revision() -> None:
|
||||
actual_commit = _run_git("rev-parse", "HEAD")
|
||||
if actual_commit != PINNED_COMMIT:
|
||||
raise RuntimeError(f"MiniMax-H3 parity requires Diffusers commit {PINNED_COMMIT}, got {actual_commit}.")
|
||||
dirty_source = _run_git("status", "--short", "--untracked-files=no", "--", "src/diffusers")
|
||||
if dirty_source:
|
||||
raise RuntimeError(f"MiniMax-H3 reference source has tracked changes:\n{dirty_source}")
|
||||
|
||||
|
||||
def assert_pinned_reference(relative_path: str, sha256: str) -> Path:
|
||||
path = REFERENCE_ROOT / relative_path
|
||||
if not path.is_file():
|
||||
|
||||
@@ -0,0 +1,300 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Measure MiniMax-H3 production VAE stage memory with CPU offload enabled."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import gc
|
||||
import hashlib
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
import statistics
|
||||
import subprocess
|
||||
import sys
|
||||
import time
|
||||
from types import SimpleNamespace
|
||||
|
||||
|
||||
def _parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--source-root", type=Path, required=True)
|
||||
parser.add_argument("--model-root", type=Path, required=True)
|
||||
parser.add_argument("--revision-label", required=True)
|
||||
parser.add_argument("--operation", choices=("encode", "decode"), required=True)
|
||||
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)
|
||||
parser.add_argument("--seed", type=int, default=20260803)
|
||||
parser.add_argument("--warmups", type=int, default=1)
|
||||
parser.add_argument("--repetitions", type=int, default=3)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def _git(source_root: Path, *args: str) -> str:
|
||||
result = subprocess.run(
|
||||
["git", "-C", str(source_root), *args],
|
||||
check=True,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
)
|
||||
return result.stdout.strip()
|
||||
|
||||
|
||||
def _sha256(path: Path) -> str:
|
||||
digest = hashlib.sha256()
|
||||
with path.open("rb") as file:
|
||||
for block in iter(lambda: file.read(1024 * 1024), b""):
|
||||
digest.update(block)
|
||||
return digest.hexdigest()
|
||||
|
||||
|
||||
def _model_snapshot(model_root: Path) -> str | None:
|
||||
parts = model_root.resolve().parts
|
||||
try:
|
||||
return parts[parts.index("snapshots") + 1]
|
||||
except (ValueError, IndexError):
|
||||
return None
|
||||
|
||||
|
||||
def _make_layout(rows, latent_shape):
|
||||
import torch
|
||||
|
||||
from fastvideo.pipelines.basic.minimax_h3.packing import MiniMaxH3PackedLayout
|
||||
|
||||
empty = torch.empty(0, dtype=torch.long, device=rows.device)
|
||||
return MiniMaxH3PackedLayout(
|
||||
sequence_length=rows.shape[0],
|
||||
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 _build_operation(args, vae, device):
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.models.dits.minimax_h3 import MiniMaxH3Config
|
||||
from fastvideo.pipelines.basic.minimax_h3.packing import patchify_video_latents, video_latent_num_frames
|
||||
from fastvideo.pipelines.basic.minimax_h3.reference import MiniMaxH3PreparedReference
|
||||
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
|
||||
|
||||
patch_size = MiniMaxH3Config().arch_config.patch_size
|
||||
transformer = SimpleNamespace(patch_size=patch_size)
|
||||
runtime_args = SimpleNamespace(output_type="pil", pin_cpu_memory=False, vae_cpu_offload=True)
|
||||
|
||||
if args.operation == "encode":
|
||||
if int(os.environ.get("WORLD_SIZE", "1")) != 1:
|
||||
raise ValueError("The encode benchmark is single-rank; use one process.")
|
||||
frames = np.random.default_rng(args.seed).integers(
|
||||
0,
|
||||
256,
|
||||
size=(args.num_frames, args.height, args.width, 3),
|
||||
dtype=np.uint8,
|
||||
)
|
||||
stage = MiniMaxH3LatentPreparationStage(
|
||||
transformer=transformer,
|
||||
vae=vae,
|
||||
audio_vae=None,
|
||||
scheduler=None,
|
||||
ref2va=True,
|
||||
)
|
||||
|
||||
def run_once():
|
||||
reference = MiniMaxH3PreparedReference(media_type="video", frames=frames)
|
||||
vae.to(device)
|
||||
try:
|
||||
return stage._encode_visual_rows([reference], device)[0]
|
||||
finally:
|
||||
vae.to("cpu")
|
||||
|
||||
return run_once, {
|
||||
"input": "NumPy PCG64 CPU uint8 RGB pixels",
|
||||
"input_shape": [1, 3, args.num_frames, args.height, args.width],
|
||||
"boundary": "before VAE CPU-to-GPU transfer through post-encode VAE CPU offload",
|
||||
}
|
||||
|
||||
latent_shape = (
|
||||
1,
|
||||
vae.latent_channels,
|
||||
video_latent_num_frames(args.num_frames),
|
||||
args.height // vae.spatial_compression_ratio,
|
||||
args.width // vae.spatial_compression_ratio,
|
||||
)
|
||||
generator = torch.Generator(device=device).manual_seed(args.seed)
|
||||
latents = torch.randn(latent_shape, generator=generator, device=device, dtype=torch.float32)
|
||||
rows = patchify_video_latents(latents, patch_size)
|
||||
layout = _make_layout(rows, latent_shape)
|
||||
stage = MiniMaxH3VideoDecodingStage(vae, transformer)
|
||||
|
||||
def run_once():
|
||||
batch = ForwardBatch(data_type="video", latents=rows, raw_latent_shape=latent_shape)
|
||||
batch.extra[MINIMAX_H3_LAYOUT_KEY] = layout
|
||||
return stage.forward(batch, runtime_args).output
|
||||
|
||||
return run_once, {
|
||||
"input": "PyTorch Philox normalized FP32 latents",
|
||||
"input_shape": list(latent_shape),
|
||||
"boundary": "MiniMaxH3VideoDecodingStage.forward including VAE CPU-to-GPU transfer and CPU offload",
|
||||
}
|
||||
|
||||
|
||||
def _measure(run_once, device, warmups: int, repetitions: int) -> dict:
|
||||
import torch
|
||||
|
||||
for _ in range(warmups):
|
||||
with torch.inference_mode():
|
||||
result = run_once()
|
||||
torch.cuda.synchronize(device)
|
||||
del result
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
records = []
|
||||
for repetition in range(repetitions):
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.synchronize(device)
|
||||
start_allocated = torch.cuda.memory_allocated(device)
|
||||
start_reserved = torch.cuda.memory_reserved(device)
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
started = time.perf_counter()
|
||||
with torch.inference_mode():
|
||||
result = run_once()
|
||||
torch.cuda.synchronize(device)
|
||||
elapsed = time.perf_counter() - started
|
||||
records.append({
|
||||
"repetition": repetition,
|
||||
"elapsed_seconds": elapsed,
|
||||
"start_allocated_bytes": start_allocated,
|
||||
"peak_allocated_bytes": torch.cuda.max_memory_allocated(device),
|
||||
"incremental_peak_allocated_bytes": torch.cuda.max_memory_allocated(device) - start_allocated,
|
||||
"start_reserved_bytes": start_reserved,
|
||||
"peak_reserved_bytes": torch.cuda.max_memory_reserved(device),
|
||||
"incremental_peak_reserved_bytes": torch.cuda.max_memory_reserved(device) - start_reserved,
|
||||
"output_shape": list(result.shape),
|
||||
})
|
||||
del result
|
||||
return {
|
||||
"repetitions": records,
|
||||
"median_elapsed_seconds": statistics.median(record["elapsed_seconds"] for record in records),
|
||||
"max_peak_allocated_bytes": max(record["peak_allocated_bytes"] for record in records),
|
||||
"max_incremental_peak_allocated_bytes": max(
|
||||
record["incremental_peak_allocated_bytes"] for record in records),
|
||||
"max_peak_reserved_bytes": max(record["peak_reserved_bytes"] for record in records),
|
||||
"max_incremental_peak_reserved_bytes": max(
|
||||
record["incremental_peak_reserved_bytes"] for record in records),
|
||||
}
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = _parse_args()
|
||||
if args.warmups < 0 or args.repetitions < 1:
|
||||
raise ValueError("warmups must be non-negative and repetitions must be positive.")
|
||||
source_root = args.source_root.resolve()
|
||||
sys.path.insert(0, str(source_root))
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
|
||||
import fastvideo
|
||||
from fastvideo.configs.pipelines.minimax_h3 import MiniMaxH3PipelineConfig
|
||||
from fastvideo.distributed import cleanup_dist_env_and_memory, maybe_init_distributed_environment_and_model_parallel
|
||||
from fastvideo.models.loader.component_loader import VAELoader
|
||||
|
||||
imported_root = Path(fastvideo.__file__).resolve().parents[1]
|
||||
if imported_root != source_root:
|
||||
raise RuntimeError(f"Imported FastVideo from {imported_root}, expected {source_root}.")
|
||||
|
||||
os.environ.setdefault("MASTER_ADDR", "127.0.0.1")
|
||||
os.environ.setdefault("MASTER_PORT", "29673")
|
||||
os.environ.setdefault("RANK", "0")
|
||||
os.environ.setdefault("WORLD_SIZE", "1")
|
||||
os.environ.setdefault("LOCAL_RANK", "0")
|
||||
world_size = int(os.environ["WORLD_SIZE"])
|
||||
maybe_init_distributed_environment_and_model_parallel(1, world_size)
|
||||
rank = dist.get_rank()
|
||||
device = torch.device(f"cuda:{int(os.environ['LOCAL_RANK'])}")
|
||||
|
||||
component_dir = args.model_root.resolve() / "vae"
|
||||
config_path = component_dir / "config.json"
|
||||
index_path = component_dir / "diffusion_pytorch_model.safetensors.index.json"
|
||||
for path in (config_path, index_path):
|
||||
if not path.is_file():
|
||||
raise FileNotFoundError(path)
|
||||
|
||||
loader_args = SimpleNamespace(
|
||||
pipeline_config=MiniMaxH3PipelineConfig(),
|
||||
model_paths={},
|
||||
vae_cpu_offload=True,
|
||||
)
|
||||
vae = VAELoader().load(str(component_dir), loader_args)
|
||||
parameter_bytes = sum(parameter.numel() * parameter.element_size() for parameter in vae.parameters())
|
||||
run_once, workload = _build_operation(args, vae, device)
|
||||
rank_result = _measure(run_once, device, args.warmups, args.repetitions)
|
||||
rank_result.update({
|
||||
"rank": rank,
|
||||
"device_index": device.index,
|
||||
"device_name": torch.cuda.get_device_name(device),
|
||||
"device_total_memory_bytes": torch.cuda.get_device_properties(device).total_memory,
|
||||
})
|
||||
rank_results = [None] * world_size
|
||||
dist.all_gather_object(rank_results, rank_result)
|
||||
|
||||
if rank == 0:
|
||||
result = {
|
||||
"schema_version": 1,
|
||||
"exit_status": 0,
|
||||
"revision_label": args.revision_label,
|
||||
"source_root": str(source_root),
|
||||
"source_git_head": _git(source_root, "rev-parse", "HEAD"),
|
||||
"source_tracked_status": _git(source_root, "status", "--short", "--untracked-files=no"),
|
||||
"benchmark_script_sha256": _sha256(Path(__file__).resolve()),
|
||||
"model_root": str(args.model_root.resolve()),
|
||||
"model_snapshot": _model_snapshot(args.model_root),
|
||||
"model_config_sha256": _sha256(config_path),
|
||||
"model_index_sha256": _sha256(index_path),
|
||||
"vae_parameter_bytes": parameter_bytes,
|
||||
"operation": args.operation,
|
||||
"vae_cpu_offload": True,
|
||||
"pin_cpu_memory": False,
|
||||
"warmups": args.warmups,
|
||||
"measurement_repetitions": args.repetitions,
|
||||
"seed": args.seed,
|
||||
"workload": workload,
|
||||
"world_size": world_size,
|
||||
"torch_version": torch.__version__,
|
||||
"cuda_version": torch.version.cuda,
|
||||
"allocator_config": os.environ.get("PYTORCH_ALLOC_CONF")
|
||||
or os.environ.get("PYTORCH_CUDA_ALLOC_CONF"),
|
||||
"rank_results": rank_results,
|
||||
"sum_rank_local_max_peak_allocated_bytes": sum(
|
||||
item["max_peak_allocated_bytes"] for item in rank_results),
|
||||
"sum_rank_local_max_incremental_peak_allocated_bytes": sum(
|
||||
item["max_incremental_peak_allocated_bytes"] for item in rank_results),
|
||||
"sum_rank_local_max_peak_reserved_bytes": sum(
|
||||
item["max_peak_reserved_bytes"] for item in rank_results),
|
||||
"sum_rank_local_max_incremental_peak_reserved_bytes": sum(
|
||||
item["max_incremental_peak_reserved_bytes"] for item in rank_results),
|
||||
}
|
||||
print("MINIMAX_H3_VAE_MEMORY=" + json.dumps(result, sort_keys=True), flush=True)
|
||||
|
||||
cleanup_dist_env_and_memory()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -62,8 +62,9 @@ def _require_assets() -> tuple[torch.device, Path]:
|
||||
|
||||
def _load_official(component_dir: Path, device: torch.device) -> torch.nn.Module:
|
||||
from diffusers.models.autoencoders.autoencoder_kl_minimax_h3 import AutoencoderKLMiniMaxH3
|
||||
from tests.local_tests.minimax_h3._reference import assert_reference_source
|
||||
from tests.local_tests.minimax_h3._reference import assert_reference_revision, assert_reference_source
|
||||
|
||||
assert_reference_revision()
|
||||
assert_reference_source(
|
||||
AutoencoderKLMiniMaxH3,
|
||||
"src/diffusers/models/autoencoders/autoencoder_kl_minimax_h3.py",
|
||||
@@ -101,14 +102,35 @@ def _load_fastvideo(component_dir: Path) -> torch.nn.Module:
|
||||
return model
|
||||
|
||||
|
||||
def _make_video() -> torch.Tensor:
|
||||
def _make_pixels() -> torch.Tensor:
|
||||
generator = torch.Generator(device="cpu").manual_seed(20260803)
|
||||
# Two logical 17-frame clips after padding are required for H3's three-token
|
||||
# temporal overlap contract; 32 px remains safely above reflect-pad minima.
|
||||
return torch.randn(1, 3, 22, 32, 32, generator=generator, dtype=torch.float32)
|
||||
return torch.randint(0, 256, (1, 3, 22, 32, 32), generator=generator, dtype=torch.uint8)
|
||||
|
||||
|
||||
def _run(model: torch.nn.Module, video: torch.Tensor) -> dict[str, torch.Tensor]:
|
||||
def _normalize_pixels(pixels: torch.Tensor, device: torch.device) -> torch.Tensor:
|
||||
video = pixels.to(device=device, dtype=torch.float32).div_(255.0)
|
||||
pixel_mean = torch.tensor((0.485, 0.456, 0.406), device=device).view(1, -1, 1, 1, 1)
|
||||
pixel_std = torch.tensor((0.229, 0.224, 0.225), device=device).view(1, -1, 1, 1, 1)
|
||||
return (video - pixel_mean) / pixel_std
|
||||
|
||||
|
||||
def _run_encode_pixels(model: torch.nn.Module, pixels: torch.Tensor) -> dict[str, torch.Tensor]:
|
||||
with torch.inference_mode():
|
||||
posterior = model.encode_pixels(pixels, return_dict=False)[0]
|
||||
return {
|
||||
"mean": posterior.mean.detach().cpu(),
|
||||
"logvar": posterior.logvar.detach().cpu(),
|
||||
}
|
||||
|
||||
|
||||
def _run(
|
||||
model: torch.nn.Module,
|
||||
video: torch.Tensor,
|
||||
*,
|
||||
stream_output: bool = False,
|
||||
) -> dict[str, torch.Tensor]:
|
||||
with torch.inference_mode():
|
||||
posterior = model.encode(video, return_dict=False)[0]
|
||||
latents = posterior.mode()
|
||||
@@ -121,12 +143,25 @@ def _run(model: torch.nn.Module, video: torch.Tensor) -> dict[str, torch.Tensor]
|
||||
dtype=latents.dtype).view(1, -1, 1, 1, 1)
|
||||
normalized = (latents - mean) / std
|
||||
decoded = model.decode(latents, return_dict=False)[0]
|
||||
return {
|
||||
pixel_mean = torch.tensor((0.485, 0.456, 0.406), device=decoded.device,
|
||||
dtype=torch.float32).view(1, -1, 1, 1, 1)
|
||||
pixel_std = torch.tensor((0.229, 0.224, 0.225), device=decoded.device,
|
||||
dtype=torch.float32).view(1, -1, 1, 1, 1)
|
||||
pixels = (decoded.float() * pixel_std + pixel_mean).clamp_(0, 1)
|
||||
streamed = None
|
||||
if stream_output:
|
||||
streamed = torch.empty(model.decoded_pixel_shape(latents.shape), dtype=torch.float32, device="cpu")
|
||||
model.decode_to_pixels(latents, streamed)
|
||||
result = {
|
||||
"mean": posterior.mean.detach().cpu(),
|
||||
"logvar": posterior.logvar.detach().cpu(),
|
||||
"normalized": normalized.detach().cpu(),
|
||||
"decoded": decoded.detach().cpu(),
|
||||
"pixels": pixels.detach().cpu(),
|
||||
}
|
||||
if streamed is not None:
|
||||
result["streamed"] = streamed
|
||||
return result
|
||||
|
||||
|
||||
def _reclaim_vram() -> None:
|
||||
@@ -151,15 +186,16 @@ def _assert_tensor_parity(name: str, actual: torch.Tensor, expected: torch.Tenso
|
||||
def test_minimax_h3_video_vae_parity() -> None:
|
||||
"""Match posterior, normalization, geometry, and deterministic decode."""
|
||||
device, component_dir = _require_assets()
|
||||
video = _make_video()
|
||||
pixels = _make_pixels()
|
||||
|
||||
official = _load_official(component_dir, device)
|
||||
expected = _run(official, video.to(device))
|
||||
expected = _run(official, _normalize_pixels(pixels, device))
|
||||
del official
|
||||
_reclaim_vram()
|
||||
|
||||
fastvideo = _load_fastvideo(component_dir)
|
||||
actual = _run(fastvideo, video.to(device))
|
||||
actual = _run(fastvideo, _normalize_pixels(pixels, device), stream_output=True)
|
||||
streaming_encode = _run_encode_pixels(fastvideo, pixels)
|
||||
assert fastvideo.temporal_compression_ratio == 4
|
||||
assert fastvideo.spatial_compression_ratio == 16
|
||||
del fastvideo
|
||||
@@ -167,5 +203,8 @@ def test_minimax_h3_video_vae_parity() -> None:
|
||||
|
||||
_assert_tensor_parity("video_vae.mean", actual["mean"], expected["mean"], 0.0)
|
||||
_assert_tensor_parity("video_vae.logvar", actual["logvar"], expected["logvar"], 0.0)
|
||||
_assert_tensor_parity("video_vae.streaming_encode.mean", streaming_encode["mean"], expected["mean"], 0.0)
|
||||
_assert_tensor_parity("video_vae.streaming_encode.logvar", streaming_encode["logvar"], expected["logvar"], 0.0)
|
||||
_assert_tensor_parity("video_vae.normalized", actual["normalized"], expected["normalized"], 0.0)
|
||||
_assert_tensor_parity("video_vae.decode", actual["decoded"], expected["decoded"], 0.0)
|
||||
_assert_tensor_parity("video_vae.streaming_decode", actual["streamed"], expected["pixels"], 0.0)
|
||||
|
||||
@@ -0,0 +1,254 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Correctness tests for the sm_100a CUDA block-sparse VSA forward.
|
||||
|
||||
Compared against an explicit PyTorch reference rather than the Triton kernel: Triton's
|
||||
block-sparse forward is hardcoded to 64-token blocks (BLOCK_M = BLOCK_N = 64) while this
|
||||
extension also carries a 128-token build, so a direct comparison would be comparing two
|
||||
different sparsity granularities. The reference below applies exactly the semantics the
|
||||
kernel is supposed to implement -- selected blocks only, keys past variable_block_sizes
|
||||
masked. Every case runs at both block sizes.
|
||||
|
||||
Run with: python -m pytest tests/test_block_sparse_sm100a.py -v
|
||||
"""
|
||||
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from fastvideo_kernel import block_sparse_attn_sm100a as vsa
|
||||
|
||||
HEAD_DIM = 128
|
||||
BLOCK_SIZES = [64, 128]
|
||||
|
||||
pytestmark = pytest.mark.skipif(
|
||||
not torch.cuda.is_available() or torch.cuda.get_device_capability() != (10, 0)
|
||||
or not vsa._HAS_VSA_SM100A,
|
||||
reason="requires Blackwell (sm_100a) and a built fastvideo_kernel extension",
|
||||
)
|
||||
|
||||
|
||||
def make_case(block, num_blocks=8, topk=4, heads=4, batch=1, ragged=False, seed=0):
|
||||
torch.manual_seed(seed)
|
||||
S = num_blocks * block
|
||||
shape = (batch, heads, S, HEAD_DIM) if vsa.BHSD else (batch, S, heads, HEAD_DIM)
|
||||
q, k, v = (torch.randn(shape, device="cuda", dtype=torch.bfloat16) for _ in range(3))
|
||||
|
||||
idx = torch.empty((batch * heads * num_blocks, topk), dtype=torch.int32, device="cuda")
|
||||
for r in range(idx.shape[0]):
|
||||
idx[r] = torch.randperm(num_blocks, device="cuda")[:topk].to(torch.int32).sort().values
|
||||
num = torch.full((batch * heads * num_blocks, ), topk, dtype=torch.int32, device="cuda")
|
||||
|
||||
if ragged:
|
||||
vbs = torch.randint(block // 2, block + 1, (num_blocks, ), dtype=torch.int32,
|
||||
device="cuda")
|
||||
else:
|
||||
vbs = torch.full((num_blocks, ), block, dtype=torch.int32, device="cuda")
|
||||
return q, k, v, idx, num, vbs
|
||||
|
||||
|
||||
def reference(q, k, v, idx, num, vbs, block):
|
||||
"""Dense attention restricted to the selected blocks, with padded keys masked."""
|
||||
if not vsa.BHSD:
|
||||
q, k, v = (t.transpose(1, 2) for t in (q, k, v)) # -> [B, H, S, D]
|
||||
B, H, S, D = q.shape
|
||||
num_blocks = vbs.numel()
|
||||
scale = 1.0 / (D**0.5)
|
||||
|
||||
keep = torch.zeros((B, H, S, S), dtype=torch.bool, device=q.device)
|
||||
for b in range(B):
|
||||
for h in range(H):
|
||||
for qb in range(num_blocks):
|
||||
row = (b * H + h) * num_blocks + qb
|
||||
for j in range(int(num[row])):
|
||||
kb = int(idx[row, j])
|
||||
valid = int(vbs[kb])
|
||||
keep[b, h, qb * block:(qb + 1) * block,
|
||||
kb * block:kb * block + valid] = True
|
||||
|
||||
scores = (q.float() @ k.float().transpose(-1, -2)) * scale
|
||||
scores = scores.masked_fill(~keep, float("-inf"))
|
||||
p = torch.softmax(scores, dim=-1)
|
||||
out = p @ v.float()
|
||||
lse = torch.logsumexp(scores, dim=-1) * 1.4426950408889634
|
||||
return out, lse
|
||||
|
||||
|
||||
def run_and_compare(block, ragged, num_blocks=8, topk=4, heads=4, atol=0.02):
|
||||
q, k, v, idx, num, vbs = make_case(block, num_blocks=num_blocks, topk=topk, heads=heads,
|
||||
ragged=ragged)
|
||||
assert vsa.is_supported(q, vbs)
|
||||
got, got_lse = vsa.block_sparse_attn_sm100a(q, k, v, idx, num, vbs)
|
||||
ref, ref_lse = reference(q, k, v, idx, num, vbs, block)
|
||||
|
||||
got_o = got if vsa.BHSD else got.transpose(1, 2)
|
||||
diff = (got_o.float() - ref).abs().max().item()
|
||||
assert diff < atol, f"out: max |diff| = {diff:.5f}"
|
||||
lse_diff = (got_lse.float() - ref_lse).abs().max().item()
|
||||
assert lse_diff < 0.05, f"lse: max |diff| = {lse_diff:.5f}"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("block", BLOCK_SIZES)
|
||||
def test_forward_matches_reference(block):
|
||||
run_and_compare(block, ragged=False)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("block", BLOCK_SIZES)
|
||||
def test_forward_matches_reference_ragged(block):
|
||||
"""variable_block_sizes is what FastVideo always passes; padded keys must be masked."""
|
||||
run_and_compare(block, ragged=True)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("block", BLOCK_SIZES)
|
||||
@pytest.mark.parametrize("topk", [1, 2, 3, 5, 7])
|
||||
def test_topk_not_a_multiple_of_the_group(block, topk):
|
||||
"""The kernel groups selected blocks; a ragged final group must still be correct."""
|
||||
run_and_compare(block, ragged=True, num_blocks=8, topk=topk)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("block", BLOCK_SIZES)
|
||||
@pytest.mark.parametrize("num_blocks", [4, 8, 16])
|
||||
def test_sequence_lengths(block, num_blocks):
|
||||
run_and_compare(block, ragged=True, num_blocks=num_blocks, topk=3)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("block", BLOCK_SIZES)
|
||||
def test_lse_is_not_vacuous(block):
|
||||
"""Guards the lse assertion: a wrong lse must actually fail the comparison."""
|
||||
q, k, v, idx, num, vbs = make_case(block, ragged=True)
|
||||
_, got_lse = vsa.block_sparse_attn_sm100a(q, k, v, idx, num, vbs)
|
||||
_, ref_lse = reference(q, k, v, idx, num, vbs, block)
|
||||
assert (got_lse.float() - (ref_lse + 1.0)).abs().max().item() > 0.5
|
||||
|
||||
|
||||
def test_unsupported_is_rejected():
|
||||
q, _, _, _, _, vbs = make_case(64)
|
||||
assert not vsa.is_supported(q.float(), vbs) # wrong dtype
|
||||
assert not vsa.is_supported(q[..., :64].contiguous(), vbs) # wrong head_dim
|
||||
odd = torch.full((7, ), 64, dtype=torch.int32, device="cuda")
|
||||
assert not vsa.is_supported(q, odd) # seqlen/blocks mismatch
|
||||
thirty_two = torch.full((16, ), 32, dtype=torch.int32, device="cuda")
|
||||
assert not vsa.is_supported(q, thirty_two) # block size with no build
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Per-q-tile q2k_num regressions. A CTA owns an ADJACENT PAIR of query blocks;
|
||||
# the kernel must honor each row's own count -- not the even row's -- in both
|
||||
# the q2k_idx window clamp and the vbs-threshold masking. These cases pin the
|
||||
# two failure modes of a pair-shared count: silent corruption of the odd tile
|
||||
# when the pair's rows differ, and a hang when an even row's count is zero.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def make_two_class_case(block, num_blocks=16, prefix=3, topk=4, heads=4, seed=0, zero_rows=()):
|
||||
"""Two row classes, like a packed multimodal layout: `prefix` DENSE rows
|
||||
(count = num_blocks) followed by rows at a uniform smaller count
|
||||
(prefix + topk). With an ODD `prefix`, pair (prefix-1, prefix) straddles
|
||||
the classes, so the two rows of one CTA carry different counts."""
|
||||
torch.manual_seed(seed)
|
||||
S = num_blocks * block
|
||||
shape = (1, heads, S, HEAD_DIM) if vsa.BHSD else (1, S, heads, HEAD_DIM)
|
||||
q, k, v = (torch.randn(shape, device="cuda", dtype=torch.bfloat16) for _ in range(3))
|
||||
vbs = torch.randint(block // 2, block + 1, (num_blocks, ), dtype=torch.int32, device="cuda")
|
||||
|
||||
rows = heads * num_blocks
|
||||
# Pad past each row's count with 0 -- a VALID block id, like real metadata
|
||||
# padding -- so a regression to the pair-shared count reads plausible
|
||||
# in-bounds garbage and fails by WRONG VALUES rather than by luck.
|
||||
idx = torch.zeros((rows, num_blocks), dtype=torch.int32, device="cuda")
|
||||
num = torch.zeros((rows, ), dtype=torch.int32, device="cuda")
|
||||
g = torch.Generator().manual_seed(seed + 1)
|
||||
for h in range(heads):
|
||||
for t in range(num_blocks):
|
||||
r = h * num_blocks + t
|
||||
if t < prefix:
|
||||
sel = torch.arange(num_blocks, dtype=torch.int32)
|
||||
else:
|
||||
vid = torch.randperm(num_blocks - prefix, generator=g)[:topk] + prefix
|
||||
sel = torch.cat([torch.arange(prefix), vid.sort().values]).to(torch.int32)
|
||||
if t in zero_rows:
|
||||
sel = sel[:0]
|
||||
idx[r, :sel.numel()] = sel.cuda()
|
||||
num[r] = sel.numel()
|
||||
return q, k, v, idx, num, vbs
|
||||
|
||||
|
||||
@pytest.mark.parametrize("block", BLOCK_SIZES)
|
||||
def test_nonuniform_counts_straddling_a_pair(block):
|
||||
"""A dense row paired with a top-k row must both be exact. Regression: the
|
||||
pair-shared count computed the odd tile with the even row's count (max
|
||||
|out diff| ~0.6 on this exact case), leaving every other tile correct."""
|
||||
q, k, v, idx, num, vbs = make_two_class_case(block)
|
||||
assert vsa.is_supported(q, vbs)
|
||||
got, got_lse = vsa.block_sparse_attn_sm100a(q, k, v, idx, num, vbs)
|
||||
ref, ref_lse = reference(q, k, v, idx, num, vbs, block)
|
||||
got_o = got if vsa.BHSD else got.transpose(1, 2)
|
||||
diff = (got_o.float() - ref).abs().max().item()
|
||||
assert diff < 0.02, f"out: max |diff| = {diff:.5f}"
|
||||
lse_diff = (got_lse.float() - ref_lse).abs().max().item()
|
||||
assert lse_diff < 0.05, f"lse: max |diff| = {lse_diff:.5f}"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("block", BLOCK_SIZES)
|
||||
def test_same_seed_determinism(block):
|
||||
"""Two same-seed runs must be bitwise identical. Also documents that the
|
||||
uniform-count path is unchanged by the per-tile fix: with cnt0 == cnt1 >= 1
|
||||
the pair trip count max(cnt0, cnt1, 1) and the per-tile clamps reduce to
|
||||
the original expressions."""
|
||||
runs = []
|
||||
for _ in range(2):
|
||||
q, k, v, idx, num, vbs = make_case(block, ragged=True, seed=3)
|
||||
o, m = vsa.block_sparse_attn_sm100a(q, k, v, idx, num, vbs)
|
||||
torch.cuda.synchronize()
|
||||
runs.append((o, m))
|
||||
assert torch.equal(runs[0][0], runs[1][0]), "out differs across same-seed runs"
|
||||
assert torch.equal(runs[0][1], runs[1][1]), "lse differs across same-seed runs"
|
||||
|
||||
|
||||
def _zero_count_case_main(block):
|
||||
"""Body of test_zero_count_rows, run in a subprocess (see the test)."""
|
||||
zero_rows = (5, 6, 8, 9) # odd member; even member with nonzero sibling 7; a whole pair
|
||||
q, k, v, idx, num, vbs = make_two_class_case(block, zero_rows=zero_rows)
|
||||
assert vsa.is_supported(q, vbs)
|
||||
got, got_lse = vsa.block_sparse_attn_sm100a(q, k, v, idx, num, vbs)
|
||||
torch.cuda.synchronize()
|
||||
ref, ref_lse = reference(q, k, v, idx, num, vbs, block)
|
||||
got_o = (got if vsa.BHSD else got.transpose(1, 2)).float()
|
||||
|
||||
empty = torch.zeros(ref.shape[:-1], dtype=torch.bool, device="cuda") # [B, H, S]
|
||||
for t in zero_rows:
|
||||
rows = slice(t * block, (t + 1) * block)
|
||||
empty[:, :, rows] = True
|
||||
zmax = got_o[:, :, rows].abs().max().item()
|
||||
assert zmax == 0.0, f"zero-count tile {t}: expected exact zeros, got max |out| = {zmax}"
|
||||
assert torch.isfinite(got_lse).all(), "lse must stay finite on zero-count rows"
|
||||
# every non-empty row -- including tile 7, whose pair sibling is empty -- is exact
|
||||
diff = (got_o - ref).abs().amax(dim=-1)
|
||||
assert diff[~empty].max().item() < 0.02
|
||||
lse_diff = (got_lse.float() - ref_lse).abs()
|
||||
assert lse_diff[~empty].max().item() < 0.05
|
||||
print("ZERO_COUNT_CASE_OK")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("block", BLOCK_SIZES)
|
||||
def test_zero_count_rows(block):
|
||||
"""q2k_num == 0 rows must produce exactly-zero output rows (finite lse) and
|
||||
leave every other row intact. Runs in a subprocess with a timeout because
|
||||
the failure mode this pins is a HANG (a zero count on an even row starved
|
||||
the softmax/correction mbarrier handshake); a regression must fail the
|
||||
suite, not wedge it."""
|
||||
proc = subprocess.run(
|
||||
[sys.executable, os.path.abspath(__file__), "--zero-count-case", str(block)],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=240,
|
||||
)
|
||||
assert proc.returncode == 0 and "ZERO_COUNT_CASE_OK" in proc.stdout, (
|
||||
f"rc={proc.returncode}\nstdout: {proc.stdout[-2000:]}\nstderr: {proc.stderr[-2000:]}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
if len(sys.argv) == 3 and sys.argv[1] == "--zero-count-case":
|
||||
_zero_count_case_main(int(sys.argv[2]))
|
||||
Reference in New Issue
Block a user