Compare commits

..
Author SHA1 Message Date
SolitaryThinker 3164e57e1f [ci]: add preprocessing integration coverage 2026-07-16 18:13:49 -07:00
277 changed files with 859 additions and 34029 deletions
+34 -2
View File
@@ -1,8 +1,6 @@
env:
IMAGE_VERSION: "py3.12-latest"
BUILDKITE_CLEAN_CHECKOUT: true
# Buildkite only launches Modal; remote jobs initialize their own submodules.
BUILDKITE_GIT_SUBMODULES: false
notify:
- github_commit_status:
@@ -79,6 +77,17 @@ steps:
limit: 2
agents:
queue: "default"
- label: ":microscope: Preprocessing Tests"
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "preprocessing"
command: "timeout 90m .buildkite/scripts/pr_test.sh"
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
agents:
queue: "default"
- label: ":microscope: DreamVerse App Tests"
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "dreamverse_app"
command: "timeout 90m .buildkite/scripts/pr_test.sh"
@@ -324,6 +333,29 @@ steps:
- TEST_TYPE=unit_test
agents:
queue: "default"
- path:
- "fastvideo/pipelines/preprocess/**"
- "fastvideo/workflow/preprocess/**"
- "fastvideo/dataset/dataloader/parquet_io.py"
- "fastvideo/dataset/dataloader/record_schema.py"
- "fastvideo/dataset/dataloader/schema.py"
- "fastvideo/configs/configs.py"
- "fastvideo/fastvideo_args.py"
- "fastvideo/tests/workflow/test_t2v_preprocessing_e2e.py"
- "fastvideo/tests/nightly/reference_video_1_sample_v0.mp4"
- "fastvideo/tests/modal/pr_test.py"
- ".buildkite/pipeline.yml"
- ".buildkite/scripts/pr_test.sh"
- ".github/workflows/ci-slash-commands.yml"
- "pyproject.toml"
- "docker/Dockerfile"
config:
command: "timeout 20m .buildkite/scripts/pr_test.sh"
label: ":microscope: Preprocessing Tests"
env:
- TEST_TYPE=preprocessing
agents:
queue: "default"
- path:
- "apps/dreamverse/**"
- "pyproject.toml"
+4
View File
@@ -233,6 +233,10 @@ case "$TEST_TYPE" in
log "Running unit tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_unit_test"
;;
"preprocessing")
log "Running preprocessing integration test..."
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_preprocessing_tests"
;;
"dreamverse_app")
log "Running DreamVerse app tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_dreamverse_app_tests"
+4 -17
View File
@@ -27,25 +27,14 @@ jobs:
ref: ${{ inputs.ref || '' }}
# For PR events, lint the PR head — but keep the hook definitions from
# the base branch so an untrusted PR cannot alter what gets executed.
# The gate scripts are saved too: the self-test step below executes them,
# so it must run the base-branch copies, not the PR head's.
- name: Save trusted hook config and gate scripts
- name: Save trusted hook config
if: github.event_name == 'pull_request_target'
run: |
cp .pre-commit-config.yaml "$RUNNER_TEMP/trusted-pre-commit-config.yaml"
cp -a .github/scripts "$RUNNER_TEMP/trusted-scripts"
echo "GATE_SCRIPTS_DIR=$RUNNER_TEMP/trusted-scripts" >> "$GITHUB_ENV"
# allow-unsafe-pr-checkout acknowledges checkout's pull_request_target
# guard: the head is data for the trusted hooks to lint; nothing from it
# is executed (config and gate scripts are pinned to the base branch
# above) and credentials are not persisted. SHA-pinned to v4.4.0 because
# actionlint's action schema does not know the new input yet.
- uses: actions/checkout@11d5960a326750d5838078e36cf38b85af677262 # v4.4.0
run: cp .pre-commit-config.yaml "$RUNNER_TEMP/trusted-pre-commit-config.yaml"
- uses: actions/checkout@v4
if: github.event_name == 'pull_request_target'
with:
ref: ${{ github.event.pull_request.head.sha }}
persist-credentials: false
allow-unsafe-pr-checkout: true
- name: Restore trusted hook config
if: github.event_name == 'pull_request_target'
run: cp "$RUNNER_TEMP/trusted-pre-commit-config.yaml" .pre-commit-config.yaml
@@ -59,7 +48,5 @@ jobs:
with:
extra_args: --all-files --hook-stage manual
# After pre-commit so a self-test failure cannot mask lint failures.
# GATE_SCRIPTS_DIR points at the base-branch copy on fork PRs (set above);
# push / workflow_call runs use the checked-out tree directly.
- name: Full-suite gate self-test
run: bash "${GATE_SCRIPTS_DIR:-.github/scripts}/test_gate_full_suite.sh"
run: bash .github/scripts/test_gate_full_suite.sh
+3 -2
View File
@@ -129,7 +129,7 @@ jobs:
set -euo pipefail
TEST_NAME=$(echo "$COMMENT" | grep -oP '(?<=/test\s)\S+' | head -1 || true)
VALID="encoder vae transformer kernel unit dreamverse ssim training lora-inference lora-training lora-extraction distillation self-forcing vsa vmoba performance api train-framework eval full fastcheck pre-commit"
VALID="encoder vae transformer kernel unit preprocessing dreamverse ssim training lora-inference lora-training lora-extraction distillation self-forcing vsa vmoba performance api train-framework eval full fastcheck pre-commit"
if [ -z "$TEST_NAME" ] || ! echo "$VALID" | grep -qw "$TEST_NAME"; then
echo "Unknown test: '$TEST_NAME'. Valid: $VALID"
exit 1
@@ -137,7 +137,8 @@ jobs:
declare -A MAP=(
[encoder]=encoder [vae]=vae [transformer]=transformer
[kernel]=kernel_tests [unit]=unit_test [dreamverse]=dreamverse_app
[kernel]=kernel_tests [unit]=unit_test [preprocessing]=preprocessing
[dreamverse]=dreamverse_app
[ssim]=ssim [training]=training
[lora-inference]=inference_lora [lora-training]=training_lora
[lora-extraction]=lora_extraction
+14 -8
View File
@@ -11,7 +11,9 @@ on:
- 'requirements-mkdocs.txt'
- 'scripts/check_docs_links.py'
- '.github/workflows/infra-docs.yml'
pull_request:
# Run the trusted base-branch workflow so fork PRs can be skipped without
# waiting for maintainer approval.
pull_request_target:
branches: [ main ]
paths:
- 'docs/**'
@@ -24,21 +26,19 @@ on:
permissions:
contents: read
pages: write
id-token: write
concurrency:
group: "pages"
cancel-in-progress: false
jobs:
build:
# MkDocs executes repository code; only trusted same-repository PRs run it.
if: github.event_name == 'push' || github.event.pull_request.head.repo.full_name == github.repository
runs-on: ubuntu-latest
steps:
- name: Checkout
uses: actions/checkout@v4
with:
ref: ${{ github.event.pull_request.head.sha || github.sha }}
fetch-depth: 0
persist-credentials: false
- name: Setup Python
uses: actions/setup-python@v5
@@ -52,6 +52,7 @@ jobs:
run: uv pip install --system -r requirements-mkdocs.txt
- name: Setup Pages
if: github.event_name == 'push'
uses: actions/configure-pages@v4
- name: Build documentation
@@ -61,17 +62,22 @@ jobs:
run: python scripts/check_docs_links.py
- name: Upload artifact
if: github.event_name == 'push'
uses: actions/upload-pages-artifact@v3
with:
path: ./site
deploy:
permissions:
pages: write
id-token: write
concurrency: pages
environment:
name: github-pages
url: ${{ steps.deployment.outputs.page_url }}
runs-on: ubuntu-latest
needs: build
if: github.ref == 'refs/heads/main'
if: github.event_name == 'push'
steps:
- name: Deploy to GitHub Pages
id: deployment
-2
View File
@@ -6,7 +6,6 @@ results/
wandb/
*.ipynb
*.jpg
!examples/dataset/lingbotworld2/image.jpg
*.safetensors
*.mp4
*.png
@@ -35,7 +34,6 @@ env
*.log
weights/
logs/
/Z-Image/
official_weights/
converted_weights/
+1 -1
View File
@@ -9,7 +9,7 @@
**FastVideo is a unified post-training and real-time inference framework for accelerated video generation.**
## NEWS
- `2026/06/23`: Release FastWan-QAD: 5s of Video generated in 1.8s E2E. See the [FastWan-QAD models](https://huggingface.co/FastVideo/FastWan-QAD-FP8-1.3B), [Attn-QAT training guide](https://haoailab.com/FastVideo/training/attn_qat/), and [blog](https://haoailab.com/blogs/fastwan-qad/).
- `2026/06/23`: Release FastWan-QAD: 5s of Video generated in 1.8s E2E. [FastWan-QAD models](https://huggingface.co/FastVideo/FastWan-QAD-FP8-1.3B), check out the [Blog](https://haoailab.com/blogs/fastwan-qad/).
- `2026/03/17`: Release demo: Into the Dreamverse: Vibe Directing in FastVideo, check out the [Blog](https://haoailab.com/blogs/dreamverse/).
- `2026/03/13`: Release demo: Create a 5s 1080p Video in 4.5s with FastVideo on a Single GPU, check out the [Blog](https://haoailab.com/blogs/fastvideo_realtime_1080p/).
- `2025/11/19`: Release [CausalWan2.2 I2V A14B Preview](https://huggingface.co/FastVideo/CausalWan2.2-I2V-A14B-Preview-Diffusers) models, [Blog](https://hao-ai-lab.github.io/blogs/fastvideo_causalwan_preview/) and [Inference Code!](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_self_forcing_causal_wan2_2_i2v.py).
@@ -39,17 +39,6 @@ export FASTVIDEO_GENERATION_SEGMENT_CAP="${FASTVIDEO_GENERATION_SEGMENT_CAP:-6}"
export FASTVIDEO_PROMPT_AUTO_SLEEP_MS="${FASTVIDEO_PROMPT_AUTO_SLEEP_MS:-120}"
export FASTVIDEO_PROMPT_AUTO_TIMEOUT_MS="${FASTVIDEO_PROMPT_AUTO_TIMEOUT_MS:-1800}"
if [[ "${ENABLE_TORCH_COMPILE}" == "1" ]]; then
# Persist Inductor, AOTAutograd, and Triton artifacts across launches.
export DREAMVERSE_TORCH_COMPILE_CACHE_ROOT="${DREAMVERSE_TORCH_COMPILE_CACHE_ROOT:-${HOME}/.cache/dreamverse/torch_compile}"
export TORCHINDUCTOR_CACHE_DIR="${TORCHINDUCTOR_CACHE_DIR:-${DREAMVERSE_TORCH_COMPILE_CACHE_ROOT}/inductor}"
export TRITON_CACHE_DIR="${TRITON_CACHE_DIR:-${DREAMVERSE_TORCH_COMPILE_CACHE_ROOT}/triton}"
export TORCHINDUCTOR_FX_GRAPH_CACHE="${TORCHINDUCTOR_FX_GRAPH_CACHE:-1}"
export TORCHINDUCTOR_AUTOGRAD_CACHE="${TORCHINDUCTOR_AUTOGRAD_CACHE:-1}"
mkdir -p "${TORCHINDUCTOR_CACHE_DIR}" "${TRITON_CACHE_DIR}"
echo "[launch-demo] torch.compile cache: ${DREAMVERSE_TORCH_COMPILE_CACHE_ROOT}"
fi
cd "${DREAMVERSE_ROOT}"
if ! command -v dreamverse-server >/dev/null 2>&1; then
-35
View File
@@ -1,35 +0,0 @@
# Causal WanTrack Control
This standalone prototype prepares one image and continuously generates causal
WanTrack blocks. Handle updates received during block N are committed only at
the next block boundary, so generated frames and control history are immutable.
One GPU session may generate at a time. The SF checkpoint uses its fixed
DMD four-step schedule (`method.dmd_denoising_steps`, default
`[1000, 750, 500, 250]` with warp) without classifier-free guidance; the
server does not accept client overrides for steps or guidance.
Set a Diffusers-format causal WanTrack export and its Self-Forcing training
YAML, then launch:
```bash
export WANTRACK_MODEL_DIR=/path/to/wantrack-causal-export
export WANTRACK_YAML_PATH=/path/to/sf/config/run.yaml
export WANTRACK_TAEHV_CHECKPOINT=/path/to/taew2_1.pth
python -m apps.wantrack_control
```
Open `http://127.0.0.1:8010`. FFmpeg with `libx264` is required. Completed
blocks are preserved under `WANTRACK_OUTPUT_DIR` (or the system temporary
directory) and concatenated into a downloadable MP4 on Stop, disconnect, or a
recoverable failure.
The interactive preview uses the official TAEHV `StreamingTAEHV` decoder with
the Wan 2.1 `taew2_1.pth` weights. The full Wan VAE remains loaded only because
input preparation still needs its encoder.
The `/ws` endpoint accepts `prepare`, `start`, `control_update`, and `stop`
JSON messages. It emits phase-specific `progress`, `prepared`,
`session_started`, `block_started`, `block_encoding`, `control_applied`,
`media_init`, `media_segment_complete`, `stream_complete`, and terminal
`error` events. Binary frames carry one fMP4 initialization section followed
by ordered block fragments.
-5
View File
@@ -1,5 +0,0 @@
"""Standalone causal WanTrack control prototype."""
from apps.wantrack_control.server import create_app
__all__ = ["create_app"]
-18
View File
@@ -1,18 +0,0 @@
"""Launch the standalone WanTrack control server."""
import os
import uvicorn
def main() -> None:
uvicorn.run(
"apps.wantrack_control.server:app",
host=os.getenv("WANTRACK_HOST", "127.0.0.1"),
port=int(os.getenv("WANTRACK_PORT", "8010")),
reload=False,
)
if __name__ == "__main__":
main()
-236
View File
@@ -1,236 +0,0 @@
"""Video-only fragmented-MP4 encoding and completed-prefix finalization."""
from __future__ import annotations
from dataclasses import dataclass
import os
from pathlib import Path
import shutil
import subprocess
from collections.abc import Iterable
import uuid
import numpy as np
MEDIA_MIME = 'video/mp4; codecs="avc1.42E01E"'
@dataclass(frozen=True, slots=True)
class EncodedMediaSegment:
block_index: int
init_bytes: bytes
media_bytes: bytes
path: Path
mime: str = MEDIA_MIME
def split_fmp4(data: bytes) -> tuple[bytes, bytes]:
"""Split top-level fMP4 initialization boxes from moof/mdat media."""
offset = 0
first_fragment: int | None = None
while offset + 8 <= len(data):
size = int.from_bytes(data[offset:offset + 4], "big")
box_type = data[offset + 4:offset + 8]
header = 8
if size == 1:
if offset + 16 > len(data):
break
size = int.from_bytes(data[offset + 8:offset + 16], "big")
header = 16
elif size == 0:
size = len(data) - offset
if size < header or offset + size > len(data):
raise ValueError("invalid top-level MP4 box")
if box_type == b"moof":
first_fragment = offset
break
offset += size
if first_fragment is None:
raise ValueError("ffmpeg output did not contain an fMP4 moof box")
init_bytes = data[:first_fragment]
media_bytes = data[first_fragment:]
if b"ftyp" not in init_bytes or b"moov" not in init_bytes:
raise ValueError("ffmpeg output is missing fMP4 initialization boxes")
return init_bytes, media_bytes
class FMP4BlockWriter:
def __init__(
self,
output_root: str | os.PathLike[str],
*,
fps: float,
ffmpeg_bin: str | None = None,
) -> None:
self.fps = float(fps)
if self.fps <= 0:
raise ValueError("fps must be positive")
resolved_ffmpeg = ffmpeg_bin or shutil.which(os.getenv("WANTRACK_FFMPEG_BIN", "ffmpeg"))
if not resolved_ffmpeg:
raise RuntimeError("ffmpeg is required for WanTrack streaming")
self.ffmpeg_bin = resolved_ffmpeg
root = Path(output_root)
root.mkdir(parents=True, exist_ok=True)
self.session_id = uuid.uuid4().hex
self.session_dir = root / self.session_id
self.session_dir.mkdir()
self._block_paths: list[Path] = []
@property
def block_paths(self) -> tuple[Path, ...]:
return tuple(self._block_paths)
def encode_block(
self,
frames: np.ndarray,
block_index: int,
) -> EncodedMediaSegment:
frames = np.asarray(frames)
if frames.ndim != 4 or frames.shape[-1] != 3:
raise ValueError("frames must have shape [T, H, W, 3]")
if frames.shape[0] <= 0:
raise ValueError("cannot encode an empty frame block")
frames = np.ascontiguousarray(frames, dtype=np.uint8)
height, width = int(frames.shape[1]), int(frames.shape[2])
gop = max(1, int(frames.shape[0]))
command = [
self.ffmpeg_bin,
"-hide_banner",
"-loglevel",
"error",
"-y",
"-f",
"rawvideo",
"-pix_fmt",
"rgb24",
"-s:v",
f"{width}x{height}",
"-r",
f"{self.fps:g}",
"-i",
"pipe:0",
"-an",
"-c:v",
"libx264",
"-preset",
"ultrafast",
"-tune",
"zerolatency",
"-profile:v",
"baseline",
"-pix_fmt",
"yuv420p",
"-g",
str(gop),
"-keyint_min",
str(gop),
"-sc_threshold",
"0",
"-movflags",
"+empty_moov+default_base_moof+frag_keyframe",
"-frag_duration",
str(max(1, round(1_000_000 * frames.shape[0] / self.fps))),
"-f",
"mp4",
"pipe:1",
]
result = subprocess.run(
command,
input=frames.tobytes(),
capture_output=True,
check=False,
)
if result.returncode != 0 or not result.stdout:
stderr = result.stderr.decode("utf-8", errors="replace").strip()
raise RuntimeError(f"ffmpeg failed to encode WanTrack block {block_index}: "
f"{stderr or f'exit {result.returncode}'}")
init_bytes, media_bytes = split_fmp4(result.stdout)
path = self.session_dir / f"block_{int(block_index):06d}.mp4"
path.write_bytes(result.stdout)
self._block_paths.append(path)
return EncodedMediaSegment(
block_index=int(block_index),
init_bytes=init_bytes,
media_bytes=media_bytes,
path=path,
)
def finalize(self) -> Path | None:
if not self._block_paths:
return None
output_path = self.session_dir / "wantrack_control.mp4"
concat_path = self.session_dir / "concat.txt"
concat_path.write_text(
"".join(f"file '{self._concat_escape(path)}'\n" for path in self._block_paths),
encoding="utf-8",
)
command = [
self.ffmpeg_bin,
"-hide_banner",
"-loglevel",
"error",
"-y",
"-f",
"concat",
"-safe",
"0",
"-i",
str(concat_path),
"-an",
"-c",
"copy",
"-movflags",
"+faststart",
str(output_path),
]
result = subprocess.run(
command,
capture_output=True,
check=False,
)
if result.returncode != 0 or not output_path.is_file():
return self._finalize_reencode(output_path)
return output_path if output_path.stat().st_size > 0 else None
@staticmethod
def _concat_escape(path: Path) -> str:
return str(path.resolve()).replace("'", "'\\''")
def _finalize_reencode(self, output_path: Path) -> Path | None:
concat_path = self.session_dir / "concat.txt"
command = [
self.ffmpeg_bin,
"-hide_banner",
"-loglevel",
"error",
"-y",
"-f",
"concat",
"-safe",
"0",
"-i",
str(concat_path),
"-an",
"-c:v",
"libx264",
"-preset",
"fast",
"-pix_fmt",
"yuv420p",
"-movflags",
"+faststart",
str(output_path),
]
result = subprocess.run(
command,
capture_output=True,
check=False,
)
if (result.returncode == 0 and output_path.is_file() and output_path.stat().st_size > 0):
return output_path
return None
def total_size(paths: Iterable[Path]) -> int:
return sum(path.stat().st_size for path in paths if path.is_file())
-422
View File
@@ -1,422 +0,0 @@
"""FastAPI/WebSocket server for causal WanTrack control."""
from __future__ import annotations
import asyncio
import base64
from contextlib import suppress
import io
import os
from pathlib import Path
import tempfile
import time
from typing import Any
from fastapi import FastAPI, HTTPException, WebSocket, WebSocketDisconnect
from fastapi.responses import FileResponse
from fastapi.staticfiles import StaticFiles
from PIL import Image
from apps.wantrack_control.media import FMP4BlockWriter
_STATIC_DIR = Path(__file__).resolve().parent / "static"
def _decode_image(value: str) -> bytes:
if not isinstance(value, str) or not value.strip():
raise ValueError("prepare.image must be a non-empty base64 string")
encoded = value.split(",", 1)[1] if value.startswith("data:") else value
try:
data = base64.b64decode(encoded, validate=True)
except Exception as exc:
raise ValueError("prepare.image is not valid base64") from exc
if not data:
raise ValueError("prepare.image decoded to no bytes")
return data
def _encode_image(image: Image.Image) -> str:
output = io.BytesIO()
image.save(output, format="PNG")
encoded = base64.b64encode(output.getvalue()).decode("ascii")
return f"data:image/png;base64,{encoded}"
def _block_value(block: Any, key: str, default: Any = None) -> Any:
if isinstance(block, dict):
return block.get(key, default)
return getattr(block, key, default)
class _RuntimeProvider:
def __init__(self, runtime: Any | None) -> None:
self._runtime = runtime
self._lock = asyncio.Lock()
async def get(self) -> Any:
if self._runtime is not None:
return self._runtime
async with self._lock:
if self._runtime is None:
model_dir = os.getenv("WANTRACK_MODEL_DIR", "").strip()
yaml_path = os.getenv("WANTRACK_YAML_PATH", "").strip()
taehv_checkpoint = os.getenv("WANTRACK_TAEHV_CHECKPOINT", "").strip()
if not model_dir or not yaml_path or not taehv_checkpoint:
raise RuntimeError("Set WANTRACK_MODEL_DIR, WANTRACK_YAML_PATH, and "
"WANTRACK_TAEHV_CHECKPOINT before preparing a session")
from fastvideo.train.models.wantrack.runtime import (
WanTrackInferenceRuntime, )
self._runtime = await asyncio.to_thread(
WanTrackInferenceRuntime.from_export,
model_dir,
yaml_path,
taehv_checkpoint,
)
return self._runtime
def create_app(
runtime: Any | None = None,
*,
output_dir: str | os.PathLike[str] | None = None,
writer_factory: Any = FMP4BlockWriter,
) -> FastAPI:
app = FastAPI(title="Causal WanTrack Control")
app.state.runtime_provider = _RuntimeProvider(runtime)
app.state.active_generation = asyncio.Lock()
app.state.downloads = {}
if output_dir is None:
resolved_output_dir: str | os.PathLike[str] = os.getenv(
"WANTRACK_OUTPUT_DIR",
str(Path(tempfile.gettempdir()) / "wantrack_control"),
)
else:
resolved_output_dir = output_dir
app.state.output_dir = Path(resolved_output_dir)
app.state.writer_factory = writer_factory
app.mount("/static", StaticFiles(directory=_STATIC_DIR), name="static")
@app.get("/")
async def index() -> FileResponse:
return FileResponse(_STATIC_DIR / "index.html")
@app.get("/healthz")
async def healthz() -> dict[str, Any]:
return {
"status": "ok",
"active": app.state.active_generation.locked(),
}
@app.get("/downloads/{download_id}")
async def download(download_id: str) -> FileResponse:
path = app.state.downloads.get(download_id)
if path is None or not path.is_file():
raise HTTPException(status_code=404, detail="download not found")
return FileResponse(
path,
media_type="video/mp4",
filename="wantrack_control.mp4",
)
@app.websocket("/ws")
async def websocket_endpoint(websocket: WebSocket) -> None:
await websocket.accept()
send_lock = asyncio.Lock()
prepared: Any | None = None
session: Any | None = None
generation_task: asyncio.Task[None] | None = None
stop_event = asyncio.Event()
owns_generation_lock = False
connected = True
async def send_json(payload: dict[str, Any]) -> None:
nonlocal connected
if not connected:
return
async with send_lock:
try:
await websocket.send_json(payload)
except Exception:
connected = False
raise
async def send_bytes(payload: bytes) -> None:
nonlocal connected
if not connected:
return
async with send_lock:
try:
await websocket.send_bytes(payload)
except Exception:
connected = False
raise
async def run_generation(runtime_value: Any) -> None:
nonlocal owns_generation_lock, connected, generation_task
writer = app.state.writer_factory(
app.state.output_dir,
fps=float(getattr(runtime_value, "fps", 16.0)),
)
init_sent = False
last_applied_revision = 0
terminal_error: str | None = None
try:
while not stop_event.is_set():
block_index = int(getattr(session, "block_index", 0))
block_started_at = time.perf_counter()
await send_json({
"type": "block_started",
"block_index": block_index,
"num_inference_steps": len(
getattr(runtime_value, "dmd_denoising_steps", [1000, 750, 500, 250])),
"dmd_denoising_steps": list(
getattr(runtime_value, "dmd_denoising_steps", [1000, 750, 500, 250])),
"cfg_enabled": False,
})
block = await asyncio.to_thread(session.generate_next_block)
generated_at = time.perf_counter()
applied_revision = int(_block_value(block, "applied_revision", 0))
if applied_revision > last_applied_revision:
last_applied_revision = applied_revision
await send_json({
"type": "control_applied",
"revision": applied_revision,
"block_index": int(_block_value(block, "block_index", block_index)),
"radius": float(_block_value(block, "radius", 0.0)),
"active_handle_ids": list(_block_value(block, "active_handle_ids", ())),
})
frames = _block_value(block, "pixel_frames")
await send_json({
"type": "block_encoding",
"block_index": block_index,
})
encoded = await asyncio.to_thread(
writer.encode_block,
frames,
int(_block_value(block, "block_index", block_index)),
)
encoded_at = time.perf_counter()
if not init_sent:
await send_json({
"type": "media_init",
"mime": encoded.mime,
})
await send_bytes(encoded.init_bytes)
init_sent = True
await send_bytes(encoded.media_bytes)
await send_json({
"type": "media_segment_complete",
"block_index": encoded.block_index,
"bytes": len(encoded.media_bytes),
"generation_ms": round((generated_at - block_started_at) * 1000),
"encoding_ms": round((encoded_at - generated_at) * 1000),
})
except Exception as exc:
terminal_error = str(exc) or type(exc).__name__
finally:
if session is not None:
with suppress(Exception):
await asyncio.to_thread(
session.close,
"error" if terminal_error else ("disconnect" if not connected else "stop"),
)
final_path = await asyncio.to_thread(writer.finalize)
download_url = None
if final_path is not None:
app.state.downloads[writer.session_id] = final_path
download_url = f"/downloads/{writer.session_id}"
if owns_generation_lock:
app.state.active_generation.release()
owns_generation_lock = False
if terminal_error:
if connected:
with suppress(Exception):
await send_json({
"type": "error",
"message": terminal_error,
"download_url": download_url,
})
elif connected:
with suppress(Exception):
await send_json({
"type": "stream_complete",
"blocks": len(writer.block_paths),
"download_url": download_url,
})
generation_task = None
async def handle_message(message: dict[str, Any]) -> None:
nonlocal prepared, session, generation_task
nonlocal owns_generation_lock
message_type = str(message.get("type", "")).strip()
if message_type == "prepare":
if generation_task is not None:
raise ValueError("prepare is unavailable during generation")
prepare_started_at = time.perf_counter()
await send_json({
"type": "progress",
"phase": "loading_model",
"message": "Loading Track-v0",
"detail": "The first request loads the SF checkpoint onto the GPU.",
})
runtime_value = await app.state.runtime_provider.get()
await send_json({
"type": "progress",
"phase": "preparing_input",
"message": "Preparing image and prompt",
"detail": "Encoding text, the reference image, and its first-frame latent.",
})
image_bytes = _decode_image(message.get("image", ""))
prompt = str(message.get("prompt", "") or "")
prepared = await asyncio.to_thread(
runtime_value.prepare,
image_bytes,
prompt,
)
processed_image = getattr(prepared, "image", None)
if not isinstance(processed_image, Image.Image):
processed_image = Image.open(io.BytesIO(image_bytes)).convert("RGB")
await send_json({
"type": "prepared",
"image": _encode_image(processed_image),
"width": processed_image.width,
"height": processed_image.height,
"fps": float(getattr(runtime_value, "fps", 16.0)),
"chunk_size": int(getattr(runtime_value, "chunk_size", 3)),
"causal_recipe": getattr(runtime_value, "causal_recipe", {}),
"decoder": str(getattr(runtime_value, "decoder_name", "unknown")),
"prepare_ms": round((time.perf_counter() - prepare_started_at) * 1000),
})
return
if message_type == "start":
if prepared is None:
raise ValueError("prepare must complete before start")
if generation_task is not None:
raise ValueError("session is already generating")
handles = message.get("handles")
if not isinstance(handles, list) or not handles:
raise ValueError("start requires at least one handle")
if app.state.active_generation.locked():
raise RuntimeError("another WanTrack session is already generating")
start_started_at = time.perf_counter()
await send_json({
"type": "progress",
"phase": "starting_session",
"message": "Starting causal session",
"detail": "Initializing controls and the causal KV cache.",
})
await app.state.active_generation.acquire()
owns_generation_lock = True
runtime_value = await app.state.runtime_provider.get()
session = runtime_value.create_session()
dmd_steps = list(
getattr(runtime_value, "dmd_denoising_steps", [1000, 750, 500, 250]))
sampling = {
"seed": int(message.get("seed", 0)),
"num_inference_steps": len(dmd_steps),
"text_guidance_scale": 1.0,
"motion_guidance_scale": 1.0,
"motion_cfg": False,
}
try:
await asyncio.to_thread(
session.start,
prepared,
getattr(prepared, "prompt", str(message.get("prompt", "") or "")),
handles,
sampling,
radius=float(message.get("radius", 0.15)),
)
except Exception:
app.state.active_generation.release()
owns_generation_lock = False
raise
await send_json({
"type": "session_started",
"fps": float(getattr(runtime_value, "fps", 16.0)),
"chunk_size": int(getattr(runtime_value, "chunk_size", 3)),
"num_inference_steps": len(dmd_steps),
"dmd_denoising_steps": dmd_steps,
"cfg_enabled": False,
"causal_recipe": getattr(runtime_value, "causal_recipe", {}),
"decoder": str(getattr(runtime_value, "decoder_name", "unknown")),
"start_ms": round((time.perf_counter() - start_started_at) * 1000),
})
stop_event.clear()
generation_task = asyncio.create_task(run_generation(runtime_value))
return
if message_type == "control_update":
if session is None:
raise ValueError("control_update requires a running session")
revision = int(message.get("revision", 0))
accepted = await asyncio.to_thread(
session.apply_control_revision,
revision,
samples=message.get("samples"),
add=message.get("add"),
remove=message.get("remove"),
handles=message.get("handles"),
radius=message.get("radius"),
)
if not accepted:
await send_json({
"type": "control_applied",
"revision": revision,
"status": "ignored_stale",
})
return
if message_type == "stop":
if generation_task is None:
raise ValueError("stop requires a running session")
stop_event.set()
return
raise ValueError(f"unknown client message type: {message_type!r}")
try:
while True:
receive_task = asyncio.create_task(websocket.receive_json())
waiters: set[asyncio.Task[Any]] = {receive_task}
if generation_task is not None:
waiters.add(generation_task)
done, _ = await asyncio.wait(
waiters,
return_when=asyncio.FIRST_COMPLETED,
)
if generation_task is not None and generation_task in done:
receive_task.cancel()
with suppress(asyncio.CancelledError):
await receive_task
await generation_task
return
message = await receive_task
try:
await handle_message(message)
except Exception as exc:
await send_json({
"type": "error",
"message": str(exc) or type(exc).__name__,
})
if generation_task is not None:
stop_event.set()
except WebSocketDisconnect:
connected = False
except Exception:
connected = False
finally:
stop_event.set()
if generation_task is not None:
with suppress(Exception):
await generation_task
elif owns_generation_lock:
app.state.active_generation.release()
return app
app = create_app()
-308
View File
@@ -1,308 +0,0 @@
(() => {
const $ = (id) => document.getElementById(id);
const imageInput = $("image");
const promptInput = $("prompt");
const prepareButton = $("prepare");
const canvas = $("canvas");
const context = canvas.getContext("2d");
const addButton = $("add");
const removeButton = $("remove");
const gridInput = $("grid");
const radiusInput = $("radius");
const radiusOutput = document.querySelector(".radius output");
const startButton = $("start");
const stopButton = $("stop");
const video = $("video");
const status = $("status");
const statusLabel = $("status-label");
const statusDetail = $("status-detail");
const download = $("download");
let socket;
let preparedImage;
let handles = [];
let selectedId = null;
let addMode = false;
let dragging = false;
let generating = false;
let revision = 0;
let sessionStart = 0;
let mediaSource;
let sourceBuffer;
let mediaQueue = [];
let streamComplete = false;
function setStatus(value, detail = "", busy = false) {
statusLabel.textContent = value;
statusDetail.textContent = detail;
status.dataset.busy = String(busy);
}
function seconds(milliseconds) {
return `${(Number(milliseconds || 0) / 1000).toFixed(1)}s`;
}
function connect() {
const scheme = location.protocol === "https:" ? "wss" : "ws";
socket = new WebSocket(`${scheme}://${location.host}/ws`);
socket.binaryType = "arraybuffer";
socket.onopen = () => setStatus("Ready", "Choose an image to begin.");
socket.onclose = () => setStatus("Disconnected", "Refresh after the server reconnects.");
socket.onerror = () => setStatus("Connection error", "The control server is unreachable.");
socket.onmessage = async (event) => {
if (typeof event.data !== "string") {
mediaQueue.push(event.data);
flushMedia();
return;
}
const message = JSON.parse(event.data);
if (message.type === "progress") {
setStatus(message.message || "Working", message.detail || "", true);
} else if (message.type === "prepared") {
preparedImage = new Image();
preparedImage.onload = () => {
canvas.width = message.width;
canvas.height = message.height;
draw();
};
preparedImage.src = message.image;
handles = [];
selectedId = null;
addButton.disabled = false;
startButton.disabled = true;
prepareButton.disabled = false;
prepareButton.textContent = "Prepare again";
const recipe = message.causal_recipe || {};
setStatus(
`Prepared in ${seconds(message.prepare_ms)}`,
`${message.fps} FPS · ${message.decoder || "unknown decoder"} · ${recipe.rope_cache_policy || "unknown"} RoPE · local ${recipe.local_attn_size ?? "?"} · sink ${recipe.sink_size ?? "?"} · Add a handle, then press Start.`,
);
} else if (message.type === "session_started") {
generating = true;
sessionStart = performance.now();
startButton.disabled = true;
stopButton.disabled = false;
setStatus(
`Session started in ${seconds(message.start_ms)}`,
"4-step SF · CFG off · Building the first causal block.",
true,
);
} else if (message.type === "block_started") {
setStatus(
`Generating block ${message.block_index}`,
"DMD 4-step, single conditional branch. Controls update at the next block.",
true,
);
} else if (message.type === "block_encoding") {
setStatus(
`Encoding block ${message.block_index}`,
"The model is done; packaging frames for immediate playback.",
true,
);
} else if (message.type === "control_applied") {
if (message.status === "ignored_stale") {
setStatus(`Stale revision ${message.revision} ignored`, "Drag again to send a newer control update.");
} else {
setStatus(`Control ${message.revision} applied`, "The updated motion is active in this block.");
}
} else if (message.type === "media_init") {
setupMediaSource(message.mime);
} else if (message.type === "media_segment_complete") {
setStatus(
`Playing through block ${message.block_index}`,
`Generated in ${seconds(message.generation_ms)} · encoded in ${seconds(message.encoding_ms)} · drag handles for the next block.`,
);
} else if (message.type === "stream_complete") {
generating = false;
streamComplete = true;
stopButton.disabled = true;
startButton.disabled = false;
if (message.download_url) {
download.href = message.download_url;
download.hidden = false;
}
flushMedia();
setStatus(`Complete · ${message.blocks} blocks`, "The final MP4 is ready to download.");
} else if (message.type === "error") {
generating = false;
prepareButton.disabled = false;
stopButton.disabled = true;
if (message.download_url) {
download.href = message.download_url;
download.hidden = false;
}
setStatus("Request failed", message.message || "Unknown error");
}
};
}
function setupMediaSource(mime) {
mediaQueue = [];
streamComplete = false;
mediaSource = new MediaSource();
video.src = URL.createObjectURL(mediaSource);
mediaSource.addEventListener("sourceopen", () => {
sourceBuffer = mediaSource.addSourceBuffer(mime);
sourceBuffer.mode = "sequence";
sourceBuffer.addEventListener("updateend", flushMedia);
flushMedia();
}, { once: true });
}
function flushMedia() {
if (!sourceBuffer || sourceBuffer.updating) return;
if (mediaQueue.length) {
sourceBuffer.appendBuffer(mediaQueue.shift());
return;
}
if (streamComplete && mediaSource && mediaSource.readyState === "open") {
mediaSource.endOfStream();
}
video.play().catch(() => {});
}
function draw() {
context.clearRect(0, 0, canvas.width, canvas.height);
if (preparedImage) context.drawImage(preparedImage, 0, 0, canvas.width, canvas.height);
if (gridInput.checked) {
context.strokeStyle = "rgba(255,255,255,.12)";
context.lineWidth = 1;
for (let index = 0; index < 50; index += 1) {
const x = index * canvas.width / 49;
const y = index * canvas.height / 49;
context.beginPath(); context.moveTo(x, 0); context.lineTo(x, canvas.height); context.stroke();
context.beginPath(); context.moveTo(0, y); context.lineTo(canvas.width, y); context.stroke();
}
}
for (const handle of handles) {
const x = handle.x * canvas.width;
const y = handle.y * canvas.height;
context.beginPath();
context.arc(x, y, handle.id === selectedId ? 9 : 7, 0, Math.PI * 2);
context.fillStyle = handle.id === selectedId ? "#ffcf4a" : "#58c7ff";
context.fill();
context.strokeStyle = "#111";
context.lineWidth = 2;
context.stroke();
}
removeButton.disabled = !selectedId;
startButton.disabled = !preparedImage || handles.length === 0 || generating;
}
function canvasPoint(event) {
const rect = canvas.getBoundingClientRect();
return {
x: Math.min(1, Math.max(0, (event.clientX - rect.left) / rect.width)),
y: Math.min(1, Math.max(0, (event.clientY - rect.top) / rect.height)),
};
}
function nearest(point) {
let match = null;
let distance = Infinity;
for (const handle of handles) {
const value = Math.hypot(handle.x - point.x, handle.y - point.y);
if (value < distance && value < 18 / canvas.clientWidth) {
match = handle;
distance = value;
}
}
return match;
}
function sendControl(extra = {}) {
if (!generating) return;
revision += 1;
socket.send(JSON.stringify({ type: "control_update", revision, ...extra }));
}
canvas.addEventListener("pointerdown", (event) => {
if (!preparedImage) return;
const point = canvasPoint(event);
if (addMode) {
const handle = { id: crypto.randomUUID(), ...point };
handles.push(handle);
selectedId = handle.id;
addMode = false;
addButton.textContent = "Add handle";
sendControl({ add: [handle] });
draw();
return;
}
const handle = nearest(point);
selectedId = handle ? handle.id : null;
dragging = Boolean(handle);
if (dragging) canvas.setPointerCapture(event.pointerId);
draw();
});
canvas.addEventListener("pointermove", (event) => {
if (!dragging || !selectedId) return;
const point = canvasPoint(event);
const handle = handles.find((item) => item.id === selectedId);
if (!handle) return;
handle.x = point.x; handle.y = point.y;
sendControl({ samples: [{ id: handle.id, ...point, timestamp_ms: performance.now() - sessionStart }] });
draw();
});
canvas.addEventListener("pointerup", () => { dragging = false; });
addButton.addEventListener("click", () => {
addMode = !addMode;
addButton.textContent = addMode ? "Click canvas" : "Add handle";
});
removeButton.addEventListener("click", () => {
if (!selectedId) return;
const removed = selectedId;
handles = handles.filter((item) => item.id !== removed);
selectedId = null;
sendControl({ remove: [removed] });
draw();
});
gridInput.addEventListener("change", draw);
radiusInput.addEventListener("input", () => {
radiusOutput.textContent = Number(radiusInput.value).toFixed(2);
sendControl({ radius: Number(radiusInput.value) });
});
prepareButton.addEventListener("click", async () => {
const file = imageInput.files[0];
if (!file) {
setStatus("Choose an image", "Prepare needs a reference frame.");
return;
}
prepareButton.disabled = true;
prepareButton.textContent = "Preparing…";
setStatus("Reading image", "Using the browser's native file reader.", true);
try {
const image = await new Promise((resolve, reject) => {
const reader = new FileReader();
reader.onload = () => resolve(reader.result);
reader.onerror = () => reject(reader.error || new Error("Failed to read image"));
reader.readAsDataURL(file);
});
setStatus("Sending image", "The server will encode the prompt and reference frame next.", true);
socket.send(JSON.stringify({
type: "prepare",
image,
prompt: promptInput.value,
}));
} catch (error) {
prepareButton.disabled = false;
prepareButton.textContent = "Prepare";
setStatus("Could not read image", error.message || String(error));
}
});
startButton.addEventListener("click", () => {
revision = 0;
download.hidden = true;
startButton.disabled = true;
setStatus("Sending controls", "Starting the fixed DMD 4-step, CFG-free SF sampler.", true);
socket.send(JSON.stringify({
type: "start",
handles,
radius: Number(radiusInput.value),
seed: Number($("seed").value),
}));
});
stopButton.addEventListener("click", () => {
socket.send(JSON.stringify({ type: "stop" }));
stopButton.disabled = true;
setStatus("Finishing current block");
});
connect();
})();
-56
View File
@@ -1,56 +0,0 @@
<!doctype html>
<html lang="en">
<head>
<meta charset="utf-8">
<meta name="viewport" content="width=device-width,initial-scale=1">
<title>WanTrack Control</title>
<link rel="stylesheet" href="/static/styles.css">
</head>
<body>
<main>
<header>
<div>
<h1>WanTrack Control</h1>
<p>Drag points to steer the next causal block.</p>
</div>
<div id="status" role="status" aria-live="polite">
<span id="status-label">Disconnected</span>
<small id="status-detail">Waiting for the server connection.</small>
</div>
</header>
<section class="controls">
<label class="file">Image <input id="image" type="file" accept="image/*"></label>
<label class="prompt">Prompt <input id="prompt" type="text" placeholder="Optional scene description"></label>
<button id="prepare">Prepare</button>
</section>
<section class="workspace">
<div class="canvas-panel">
<div class="canvas-toolbar">
<button id="add" disabled>Add handle</button>
<button id="remove" disabled>Delete selected</button>
<label><input id="grid" type="checkbox"> Grid</label>
</div>
<canvas id="canvas" width="832" height="480"></canvas>
<label class="radius">Radius <input id="radius" type="range" min="0.03" max="0.45" step="0.01" value="0.15"><output>0.15</output></label>
</div>
<div class="video-panel">
<video id="video" muted autoplay playsinline controls></video>
<a id="download" hidden>Download completed MP4</a>
</div>
</section>
<section class="sampling">
<label>Seed <input id="seed" type="number" value="0" step="1"></label>
<div class="recipe">
<strong>SF checkpoint</strong>
<span>DMD [1000,750,500,250] · CFG off · TAEHV · relativistic RoPE · local 6 · sink 1</span>
</div>
<button id="start" disabled>Start</button>
<button id="stop" disabled>Stop</button>
</section>
</main>
<script src="/static/app.js"></script>
</body>
</html>
-57
View File
@@ -1,57 +0,0 @@
:root {
color-scheme: dark;
font-family: Inter, ui-sans-serif, system-ui, sans-serif;
background: #111315;
color: #edf0f2;
}
* { box-sizing: border-box; }
body { margin: 0; }
main { width: min(1180px, calc(100% - 32px)); margin: 24px auto; }
header { display: flex; align-items: start; justify-content: space-between; gap: 16px; }
h1 { margin: 0; font-size: 22px; }
p { margin: 6px 0 20px; color: #929ba3; }
#status {
display: grid; grid-template-columns: auto 1fr; gap: 2px 9px; min-width: 300px;
padding: 9px 12px; border-radius: 10px; background: #23282d; color: #dce2e7;
}
#status::before {
content: ""; grid-row: 1 / 3; align-self: center; width: 8px; height: 8px;
border-radius: 50%; background: #62d58b;
}
#status[data-busy="true"]::before {
width: 12px; height: 12px; border: 2px solid #65717a; border-top-color: #edf0f2;
background: transparent; animation: spin .8s linear infinite;
}
#status-label { font-size: 12px; font-weight: 700; }
#status-detail { color: #9ea8b0; font-size: 11px; }
@keyframes spin { to { transform: rotate(360deg); } }
.controls, .sampling, .canvas-toolbar { display: flex; align-items: end; gap: 10px; flex-wrap: wrap; }
label { color: #aeb6bd; font-size: 12px; }
input[type="text"], input[type="number"], input[type="file"] {
display: block; margin-top: 5px; min-height: 36px; border: 1px solid #353b41;
border-radius: 7px; background: #1a1e22; color: inherit; padding: 7px 9px;
}
.prompt { flex: 1; }
.prompt input { width: 100%; }
button, #download {
border: 0; border-radius: 7px; background: #e7ebee; color: #111315;
min-height: 36px; padding: 8px 13px; font-weight: 650; cursor: pointer;
}
button:disabled { opacity: .38; cursor: default; }
.workspace { display: grid; grid-template-columns: 1fr 1fr; gap: 16px; margin: 16px 0; }
.canvas-panel, .video-panel { min-width: 0; border: 1px solid #2e3439; background: #181c20; border-radius: 10px; padding: 10px; }
.canvas-toolbar { margin-bottom: 8px; }
canvas, video { display: block; width: 100%; aspect-ratio: 832 / 480; object-fit: contain; background: #090a0b; border-radius: 6px; }
canvas { touch-action: none; cursor: crosshair; }
.radius { display: flex; align-items: center; gap: 9px; margin-top: 10px; }
.radius input { flex: 1; }
#download { display: inline-block; margin-top: 10px; text-decoration: none; }
#download[hidden] { display: none; }
.sampling input { width: 92px; }
.recipe {
display: grid; gap: 2px; min-height: 36px; padding: 6px 10px;
border: 1px solid #353b41; border-radius: 7px; background: #1a1e22;
}
.recipe strong { color: #edf0f2; font-size: 12px; }
.recipe span { color: #8f99a1; font-size: 11px; }
@media (max-width: 800px) { .workspace { grid-template-columns: 1fr; } }
@@ -1,267 +0,0 @@
from __future__ import annotations
from dataclasses import dataclass
from pathlib import Path
import threading
import numpy as np
from PIL import Image
from fastapi.testclient import TestClient
from apps.wantrack_control.media import EncodedMediaSegment
from apps.wantrack_control.server import create_app
@dataclass
class _Prepared:
image: Image.Image
prompt: str
@dataclass
class _Block:
block_index: int
pixel_frames: np.ndarray
applied_revision: int
radius: float = 0.15
active_handle_ids: tuple[str, ...] = ("h", )
class _FakeSession:
def __init__(
self,
gates: list[threading.Event],
*,
fail_at: int | None = None,
) -> None:
self.gates = gates
self.fail_at = fail_at
self.block_index = 0
self.pending_revision = 0
self.closed_reason = None
def start(self, image, prompt, handles, sampling, *, radius):
assert image.prompt == prompt
assert handles
assert sampling == {
"seed": 0,
"num_inference_steps": 4,
"text_guidance_scale": 1.0,
"motion_guidance_scale": 1.0,
"motion_cfg": False,
}
assert radius > 0
def apply_control_revision(self, revision, **kwargs):
del kwargs
if revision <= self.pending_revision:
return False
self.pending_revision = revision
return True
def generate_next_block(self):
index = self.block_index
applied = self.pending_revision
if index < len(self.gates):
assert self.gates[index].wait(timeout=5)
if self.fail_at == index:
raise RuntimeError("fake generation failure")
self.block_index += 1
return _Block(
block_index=index,
pixel_frames=np.zeros((2, 8, 8, 3), dtype=np.uint8),
applied_revision=applied,
)
def close(self, reason):
self.closed_reason = reason
class _FakeRuntime:
fps = 16.0
chunk_size = 3
decoder_name = "TAEHV (taew2_1)"
causal_recipe = {
"local_attn_size": 6,
"sink_size": 1,
"rope_cache_policy": "relativistic",
}
def __init__(self, session):
self.session = session
def prepare(self, image_bytes, prompt):
assert image_bytes
return _Prepared(Image.new("RGB", (16, 12), "blue"), prompt)
def create_session(self):
return self.session
class _FakeWriter:
def __init__(self, output_root, *, fps):
assert fps == 16.0
self.session_id = "fake-download"
self.session_dir = Path(output_root) / self.session_id
self.session_dir.mkdir(parents=True, exist_ok=True)
self.block_paths = []
def encode_block(self, frames, block_index):
assert frames.shape[-1] == 3
path = self.session_dir / f"block-{block_index}.mp4"
path.write_bytes(b"complete-block")
self.block_paths.append(path)
return EncodedMediaSegment(
block_index=block_index,
init_bytes=b"init",
media_bytes=f"media-{block_index}".encode(),
path=path,
)
def finalize(self):
if not self.block_paths:
return None
path = self.session_dir / "wantrack_control.mp4"
path.write_bytes(b"".join(item.read_bytes()
for item in self.block_paths))
return path
def _json(websocket):
message = websocket.receive()
assert message["type"] == "websocket.send"
assert "text" in message
import json
return json.loads(message["text"])
def _bytes(websocket):
message = websocket.receive()
assert message["type"] == "websocket.send"
return message["bytes"]
def _prepare_and_start(websocket):
import base64
import io
image = io.BytesIO()
Image.new("RGB", (8, 8)).save(image, format="PNG")
websocket.send_json({
"type": "prepare",
"image": base64.b64encode(image.getvalue()).decode(),
"prompt": "",
})
assert _json(websocket)["phase"] == "loading_model"
assert _json(websocket)["phase"] == "preparing_input"
prepared = _json(websocket)
assert prepared["type"] == "prepared"
assert prepared["prepare_ms"] >= 0
assert prepared["causal_recipe"]["rope_cache_policy"] == "relativistic"
assert prepared["decoder"] == "TAEHV (taew2_1)"
websocket.send_json({
"type": "start",
"handles": [{
"id": "h",
"x": 0.5,
"y": 0.5,
}],
"radius": 0.15,
# Client overrides are ignored for the fixed SF recipe.
"steps": 99,
"text_guidance": 9.0,
"motion_guidance": 9.0,
})
assert _json(websocket)["phase"] == "starting_session"
started = _json(websocket)
assert started["type"] == "session_started"
assert started["num_inference_steps"] == 4
assert started["cfg_enabled"] is False
assert started["causal_recipe"]["local_attn_size"] == 6
assert started["decoder"] == "TAEHV (taew2_1)"
def test_two_block_binary_order_future_update_and_stop(tmp_path):
gates = [threading.Event(), threading.Event()]
session = _FakeSession(gates)
app = create_app(
_FakeRuntime(session),
output_dir=tmp_path,
writer_factory=_FakeWriter,
)
with TestClient(app) as client:
with client.websocket_connect("/ws") as websocket:
_prepare_and_start(websocket)
started = _json(websocket)
assert started["type"] == "block_started"
assert started["block_index"] == 0
assert started["num_inference_steps"] == 4
assert started["cfg_enabled"] is False
websocket.send_json({
"type": "control_update",
"revision": 1,
"samples": [{
"id": "h",
"x": 0.7,
"y": 0.5,
"timestamp_ms": 10,
}],
})
websocket.send_json({
"type": "control_update",
"revision": 1,
"samples": [],
})
stale = _json(websocket)
assert stale["status"] == "ignored_stale"
gates[0].set()
assert _json(websocket)["type"] == "block_encoding"
assert _json(websocket)["type"] == "media_init"
assert _bytes(websocket) == b"init"
assert _bytes(websocket) == b"media-0"
assert _json(websocket)["type"] == "media_segment_complete"
assert _json(websocket)["type"] == "block_started"
websocket.send_json({"type": "stop"})
gates[1].set()
applied = _json(websocket)
assert applied["type"] == "control_applied"
assert applied["revision"] == 1
assert _json(websocket)["type"] == "block_encoding"
assert _bytes(websocket) == b"media-1"
assert _json(websocket)["type"] == "media_segment_complete"
complete = _json(websocket)
assert complete["type"] == "stream_complete"
assert complete["blocks"] == 2
assert complete["download_url"]
response = client.get(complete["download_url"])
assert response.status_code == 200
assert response.content
assert client.get("/healthz").json()["active"] is False
def test_error_releases_lock_and_preserves_completed_prefix(tmp_path):
gates = [threading.Event(), threading.Event()]
gates[0].set()
gates[1].set()
session = _FakeSession(gates, fail_at=1)
app = create_app(
_FakeRuntime(session),
output_dir=tmp_path,
writer_factory=_FakeWriter,
)
with TestClient(app) as client:
with client.websocket_connect("/ws") as websocket:
_prepare_and_start(websocket)
assert _json(websocket)["type"] == "block_started"
assert _json(websocket)["type"] == "block_encoding"
assert _json(websocket)["type"] == "media_init"
assert _bytes(websocket) == b"init"
assert _bytes(websocket) == b"media-0"
assert _json(websocket)["type"] == "media_segment_complete"
assert _json(websocket)["type"] == "block_started"
error = _json(websocket)
assert error["type"] == "error"
assert error["download_url"]
assert client.get(error["download_url"]).content
assert client.get("/healthz").json()["active"] is False
+1 -3
View File
@@ -8,15 +8,13 @@ FastVideo provides highly optimized custom attention kernels to accelerate video
* **[Sliding Tile Attention (STA)](sta/index.md)**: STA kernel support is kept in
`fastvideo-kernel`; full FastVideo STA pipeline workflow is archived in
`sta_do_not_delete`.
* **[Attn-QAT Training](../training/attn_qat.md)**: Runtime-JIT Triton forward
and backward kernels for role-local quantization-aware training.
* **Backend development guide**: See the developer guide at
[Attention Backend Development](../contributing/attention_backend.md).
## General Build Instructions
These instructions apply to building the `fastvideo-kernel` package from
source, which includes STA, VSA, and Attn-QAT kernels.
source, which includes both STA and VSA kernels.
### Prerequisites
+6 -3
View File
@@ -85,6 +85,7 @@ status.
| Transformer Tests | `transformer` | `fastvideo/models/dits/**`, `fastvideo/models/loader/**`, `fastvideo/tests/transformers/**`, `fastvideo/layers/**`, `fastvideo/attention/**`, `pyproject.toml`, `docker/Dockerfile` |
| Kernel Tests | `kernel_tests` | `fastvideo-kernel/**`, `pyproject.toml`, `docker/Dockerfile` |
| Unit Tests | `unit_test` | `fastvideo/**`, `.buildkite/**`, `.github/**`, `pyproject.toml`, `docker/Dockerfile` |
| Preprocessing Tests | `preprocessing` | Preprocessing pipelines/workflows, Parquet schema/writer, the integration test, and its CI entrypoints |
| DreamVerse App Tests | `dreamverse_app` | `apps/dreamverse/**`, `pyproject.toml` |
### Tier 3: Full Suite
@@ -147,6 +148,7 @@ Valid direct test names:
| `/test transformer` | `transformer` |
| `/test kernel` | `kernel_tests` |
| `/test unit` | `unit_test` |
| `/test preprocessing` | `preprocessing` |
| `/test dreamverse` | `dreamverse_app` |
| `/test ssim` | `ssim` |
| `/test training` | `training` |
@@ -255,9 +257,10 @@ If you add a new CI test category:
### Documentation
`.github/workflows/infra-docs.yml` builds documentation for PRs that touch
`docs/**`, `mkdocs.yml`, `requirements-mkdocs.txt`, or the workflow itself. On
pushes to `main`, it also deploys the built site to GitHub Pages.
`.github/workflows/infra-docs.yml` builds documentation for same-repository PRs
that touch `docs/**`, `mkdocs.yml`, `requirements-mkdocs.txt`, or the workflow
itself. Fork PRs skip this executable build instead of waiting for maintainer
approval. On pushes to `main`, it also deploys the built site to GitHub Pages.
The docs job:
@@ -333,25 +333,15 @@ surfaces:
use_distill:
sources: [fastvideo.configs.pipelines.longcat.LongCatT2V480PConfig, fastvideo.configs.pipelines.longcat.LongCatT2V704PConfig]
scheduler_arch:
sources:
- fastvideo.configs.pipelines.sd35.SD35Config
- fastvideo.configs.pipelines.zimage.ZImagePipelineConfig
sources: [fastvideo.configs.pipelines.sd35.SD35Config]
text_encoder_archs:
sources:
- fastvideo.configs.pipelines.sd35.SD35Config
- fastvideo.configs.pipelines.zimage.ZImagePipelineConfig
sources: [fastvideo.configs.pipelines.sd35.SD35Config]
tokenizer_archs:
sources:
- fastvideo.configs.pipelines.sd35.SD35Config
- fastvideo.configs.pipelines.zimage.ZImagePipelineConfig
sources: [fastvideo.configs.pipelines.sd35.SD35Config]
transformer_arch:
sources:
- fastvideo.configs.pipelines.sd35.SD35Config
- fastvideo.configs.pipelines.zimage.ZImagePipelineConfig
sources: [fastvideo.configs.pipelines.sd35.SD35Config]
vae_arch:
sources:
- fastvideo.configs.pipelines.sd35.SD35Config
- fastvideo.configs.pipelines.zimage.ZImagePipelineConfig
sources: [fastvideo.configs.pipelines.sd35.SD35Config]
expand_timesteps:
sources:
- fastvideo.configs.pipelines.wan.FastWan2_2_TI2V_5B_Config
@@ -431,8 +421,6 @@ surfaces:
frame_receptive_field: "MagiHuman internal data-proxy receptive-field setting."
image_conditioning: "MagiHuman preset variant marker for reference-image conditioning."
ref_audio_offset: "MagiHuman internal data-proxy audio alignment offset."
scheduler_sigma_min: "Z-Image scheduler parity invariant; not part of the public typed inference API."
scheduler_use_reference_discrete_timesteps: "Z-Image scheduler parity invariant; not part of the public typed inference API."
sr_local_attn_layers: "MagiHuman SR internal sparse-attention layer selection."
text_offset: "MagiHuman internal data-proxy text alignment offset."
vae_stride: "MagiHuman internal VAE/data-proxy stride setting."
@@ -450,7 +438,6 @@ surfaces:
grid_sizes: request.inputs.grid_sizes
pose: request.inputs.pose
c2ws_plucker_emb: request.inputs.c2ws_plucker_emb
action_path: request.inputs.action_path
refine_from: request.inputs.refine_from
stage1_video: request.inputs.stage1_video
prompt: request.prompt
@@ -460,7 +447,6 @@ surfaces:
output_video_name: request.output.output_video_name
num_videos_per_prompt: request.sampling.num_videos_per_prompt
seed: request.sampling.seed
max_sequence_length: request.sampling.max_sequence_length
num_frames: request.sampling.num_frames
height: request.sampling.height
width: request.sampling.width
@@ -470,10 +456,7 @@ surfaces:
num_inference_steps: request.sampling.num_inference_steps
num_inference_steps_sr: request.sampling.num_inference_steps_sr
guidance_scale: request.sampling.guidance_scale
batch_cfg: request.sampling.batch_cfg
guidance_scale_2: request.sampling.guidance_scale_2
cfg_normalization: request.sampling.cfg_normalization
cfg_truncation: request.sampling.cfg_truncation
guidance_rescale: request.sampling.guidance_rescale
use_embedded_guidance: request.sampling.use_embedded_guidance
true_cfg_scale: request.sampling.true_cfg_scale
@@ -525,6 +508,7 @@ surfaces:
internal_only:
data_type: "Derived from the request shape and not a public input."
latents: "Pre-generated diffusion latents supplied by parity/debug harnesses; not a public input."
max_sequence_length: "Model-specific text-encoder sequence cap; not part of the public typed inference API."
sampling_param_extensions: {}
-6
View File
@@ -2,12 +2,6 @@
We introduce a new finetuning strategy - **Sparse-distill**, which jointly integrates **[DMD](https://arxiv.org/abs/2405.14867)** and **[VSA](https://arxiv.org/abs/2505.13389)** in a single training process. This approach combines the benefits of both distillation to shorten diffusion steps and sparse attention to reduce attention computation, enabling much faster video generation.
!!! tip "Attn-QAT DMD2 workflow"
The modular trainer also provides a Wan2.1 MixKit recipe that first
fine-tunes with fake-quantized attention, then distills the student to
timesteps `[1000, 757, 522]` while teacher and critic remain on Flash
Attention. See [Attn-QAT Training](../training/attn_qat.md).
## 📊 Model Overview
We provide two distilled models:
-147
View File
@@ -1,147 +0,0 @@
# Attn-QAT Training
Attn-QAT simulates low-bit attention during training while keeping the rest of
the training method unchanged. In the modular `fastvideo/train` framework it is
a per-role model option, not a separate training method: supervised fine-tuning
and DMD2 still own their losses and optimizer cadence.
This guide covers the QAD Wan2.1-T2V-1.3B MixKit workflow:
1. run a 4,000-step supervised Attn-QAT fine-tune;
2. export the stage-1 DCP checkpoint to Diffusers format; and
3. distill the student to three denoising steps with DMD2.
The ready-to-run configs and wrappers are in
`examples/train/scenario/qad_wan2_1_mixkit/`.
## Role-local attention backends
A DMD2 run owns three independent model roles. Configure the attention backend
on each role so fake quantization is applied only to the student:
```yaml
models:
student:
attention_backend: ATTN_QAT_TRAIN
teacher:
attention_backend: FLASH_ATTN
critic:
attention_backend: FLASH_ATTN
```
The override is active only while that role's transformer is constructed, then
the previous process-wide backend is restored. This lets student, teacher, and
critic use different implementations in one process. Invalid role-level names
fail during configuration instead of silently selecting another backend.
See [Training Infrastructure](train_infra.md) for the complete model-role
configuration reference.
## Prerequisites
- Install FastVideo and make the `fastvideo-kernel` Python package importable.
`ATTN_QAT_TRAIN` intentionally fails instead of falling back to dense
attention when its kernel cannot be loaded.
- Prepare the precomputed MixKit VAE latents and text embeddings.
- Run the commands below from the repository root. The supplied recipe expects
four GPUs by default; set `NUM_GPUS` to override it.
Download the published preprocessed dataset:
```bash
bash examples/training/finetune/wan_t2v_1.3B/mixkit/download_mixkit_data.sh
```
## Stage 1: supervised Attn-QAT fine-tuning
The stage-1 config uses `ATTN_QAT_TRAIN` on the student, sequence parallelism
across four GPUs, FP32 master weights, and 4,000 optimizer steps:
```bash
NUM_GPUS=4 \
bash examples/train/scenario/qad_wan2_1_mixkit/run_stage1.sh
```
Pass a dataset directory as the first positional argument when it differs from
the default:
```bash
NUM_GPUS=4 \
bash examples/train/scenario/qad_wan2_1_mixkit/run_stage1.sh \
/path/to/combined_parquet_dataset
```
The wrapper calls `examples/train/run.sh`; the YAML file remains the source of
truth for optimizer, validation, checkpointing, and distributed settings.
## Export the stage-1 checkpoint
Modular training checkpoints use Distributed Checkpoint (DCP) format. Export
the student before using it to initialize stage 2:
```bash
bash examples/train/scenario/qad_wan2_1_mixkit/export_stage1.sh \
checkpoints/wan_t2v_qat_finetune/checkpoint-4000 \
checkpoints/wan_t2v_qat_finetune/diffusers
```
Both arguments are optional; the command above shows their defaults.
## Stage 2: three-step DMD2 distillation
Stage 2 loads the exported student weights, keeps Attn-QAT on the student, and
uses Flash Attention for the teacher and critic:
```bash
NUM_GPUS=4 \
bash examples/train/scenario/qad_wan2_1_mixkit/run_stage2.sh \
data/HD-Mixkit-Finetune-Wan/combined_parquet_dataset \
checkpoints/wan_t2v_qat_finetune/diffusers/transformer/model.safetensors
```
The migrated recipe preserves these behaviors:
| Behavior | Modular configuration |
|---|---|
| Student fake-quantized attention | `models.student.attention_backend: ATTN_QAT_TRAIN` |
| Teacher and critic full-precision attention | Role-local `FLASH_ATTN` |
| Generator update every five critic steps | `method.generator_update_interval: 5` |
| Three-step rollout | `method.dmd_denoising_steps: [1000, 757, 522]` |
| Score timestep range | `method.min_timestep_ratio: 0.02`, `max_timestep_ratio: 0.98` |
| Legacy guidance `cond + 2(cond - uncond)` | Standard CFG scale `3.0` |
| Stage handoff | DCP checkpoint to Diffusers export to student override weights |
The timestep ratios apply to randomly sampled teacher and critic score
timesteps; `dmd_denoising_steps` separately controls the student rollout. See
[DMD Distillation](../distillation/dmd.md) for general DMD concepts.
## Architecture-specific Triton routing
The training kernel is runtime-JIT-compiled Triton code and selects its route on
every call. It supports different query and key/value sequence lengths for
cross-attention; key and value must have the same sequence length.
| Hardware/configuration | Route |
|---|---|
| SM100, validated non-causal BF16 QAT configuration with head dimension 128 | Large-tile forward and split 64x64 backward; optimized backward requires a 16-aligned KV length |
| SM120, including RTX 5090 | Previous forward tiling with joined quantized/STE P@V operations and a shallower backward pipeline for long sequences |
| Unsupported configurations | Previous Triton implementation |
Warp specialization is disabled automatically on SM100 and SM120 because the
Triton 3.7 NVWS compiler pass aborts for this kernel on Blackwell. No user
setting is required.
The available tuning and comparison controls are:
| Environment variable | Default | Effect |
|---|---|---|
| `FASTVIDEO_ATTN_QAT_FWD_MODE` | `fast` | Selects `fast`, `balanced`, or `reference` forward tiling on the SM100 optimized route |
| `FASTVIDEO_ATTN_QAT_FWD_EXACT_M` | `0` | Set to `1` to recompute reference-order softmax statistics and keep `dV` bitwise-compatible on the SM100 optimized route |
| `FASTVIDEO_ATTN_QAT_SM100_OPTIMIZED` | `1` | Set to `0` to force the previous SM100 forward and backward for comparison |
| `FASTVIDEO_ATTN_QAT_SM120_JOIN_QAT_PV` | `1` | Set to `0` to compare SM120 against the split P@V path |
The first invocation JIT-compiles the selected configuration; later calls reuse
the Triton cache. To measure the production shape, run
`python benchmarks/benchmark_attn_qat_train.py` from `fastvideo-kernel/`.
For import and backend-selection failures, see [Debugging](../utilities/debugging.md).
-13
View File
@@ -62,18 +62,6 @@ bash examples/training/finetune/wan_t2v_1.3B/crush_smol/finetune_t2v.sh
- Gradient checkpointing: `--enable_gradient_checkpointing_type "full"`
- Memory scaling: Increase `--sp_size` or reduce `--num_latent_t` to fit in memory
## Attention Quantization-Aware Training
Attn-QAT fine-tunes a model while simulating low-bit attention in the forward
and backward passes. The modular trainer can select the backend per model role,
so a later DMD2 stage can keep fake quantization on the student while the
teacher and critic use Flash Attention.
The ready-to-run Wan2.1 MixKit workflow includes supervised fine-tuning,
checkpoint export, and three-step DMD2 distillation:
**→ [Follow the Attn-QAT training guide](attn_qat.md)**
## LoRA Finetuning
LoRA (Low-Rank Adaptation) trains lightweight adapters while keeping the base model frozen. This significantly reduces memory usage and training time.
@@ -178,7 +166,6 @@ Ready-to-run training scripts are available for multiple models:
| Wan2.1 I2V 14B | I2V | `examples/training/finetune/wan_i2v_14B_480p/crush_smol/` |
| Wan2.1-Fun 1.3B InP | I2V | `examples/training/finetune/Wan2.1-Fun-1.3B-InP/crush_smol/` |
| Wan2.1 VSA | T2V/I2V | `examples/training/finetune/Wan2.1-VSA/Wan-Syn-Data/` |
| Wan2.1 T2V 1.3B Attn-QAT | QAT SFT + DMD2 | `examples/train/scenario/qad_wan2_1_mixkit/` |
Each example includes:
+1 -6
View File
@@ -43,9 +43,6 @@ Ready-to-run examples with preprocessing scripts, training launchers, and valida
**→ [Browse all training examples](examples/examples_training_index.md)**
For the complete two-stage Wan2.1 MixKit quantization-aware workflow, see
**[Attn-QAT Training](attn_qat.md)**.
Each example includes:
- `download_dataset.sh` — download sample data
@@ -62,11 +59,9 @@ FastVideo supports several training approaches:
| **Full finetune** | Adapt entire model to a new domain or style |
| **LoRA finetune** | Lightweight adaptation with frozen base weights |
| **VSA finetune** | Finetune with Variable Sparse Attention for efficiency |
| **Attn-QAT** | Train with fake-quantized attention, optionally followed by DMD2 distillation |
## Next Steps
1. **Get started**: Pick an example from the [training examples index](examples/examples_training_index.md)
2. **Prepare data**: Follow [data preprocessing](data_preprocess.md) for your own dataset
3. **Train with quantized attention**: Follow the [Attn-QAT two-stage recipe](attn_qat.md)
4. **Run inference**: After training, see [inference examples](../inference/examples/examples_inference_index.md)
3. **Run inference**: After training, see [inference examples](../inference/examples/examples_inference_index.md)
-3
View File
@@ -81,7 +81,6 @@ Common model parameters:
| `disable_custom_init_weights` | `false` | Skip custom weight initialization (use for teacher/critic) |
| `flow_shift` | `3.0` | Timestep shifting factor |
| `enable_gradient_checkpointing_type` | `null` | Gradient checkpointing (`"full"` or `null`) |
| `attention_backend` | `null` | Optional role-local backend for Wan models (for example `ATTN_QAT_TRAIN`); overrides the process default only while this role's transformer is built |
Which roles are needed depends on the training method:
@@ -299,8 +298,6 @@ method:
| `dmd_denoising_steps` | *(required)* | Timestep schedule for student rollout |
| `generator_update_interval` | `1` | Update student every N critic steps |
| `real_score_guidance_scale` | `1.0` | CFG scale for teacher predictions |
| `min_timestep_ratio` | `0.0` | Lower bound for randomly sampled teacher/critic score timesteps |
| `max_timestep_ratio` | `1.0` | Upper bound for randomly sampled teacher/critic score timesteps |
| `fake_score_learning_rate` | *(required)* | Critic optimizer learning rate |
| `fake_score_betas` | *(required)* | Critic optimizer Adam betas |
| `fake_score_lr_scheduler` | *(required)* | Critic LR scheduler type |
-188
View File
@@ -1,188 +0,0 @@
# WanTrack training
WanTrack extends Wan I2V with sparse point tracks. The same model wrapper and
conditioning path are used for preprocessing, training, validation, and
standalone inference. Both bidirectional and causal checkpoints are supported.
## Build an initialization checkpoint
WanTrack uses the pretrained control slot from Wan2.1-Fun Control as its track
slot. Convert a VideoX-Fun control checkpoint together with a diffusers-format
Wan2.1-Fun InP base:
```bash
python scripts/checkpoint_conversion/wan_fun_control_to_trackwan.py \
--inp-base models/Wan2.1-Fun-1.3B-InP-Diffusers \
--control-ckpt models/Wan2.1-Fun-1.3B-Control/diffusion_pytorch_model.safetensors \
--out models/wantrack-control-init
```
For the causal transformer, add `--causal` and use a separate output directory:
```bash
python scripts/checkpoint_conversion/wan_fun_control_to_trackwan.py \
--inp-base models/Wan2.1-Fun-1.3B-InP-Diffusers \
--control-ckpt models/Wan2.1-Fun-1.3B-Control/diffusion_pytorch_model.safetensors \
--out models/wantrack-control-causal-init \
--causal
```
`--causal` selects `CausalTrackWanTransformer3DModel` in the generated
transformer config. Without it, the converter selects
`TrackWanTransformer3DModel`.
The converted patch input has 52 channels:
| Channels | Meaning |
|----------|---------|
| `0:16` | Noisy video latent |
| `16:20` | I2V mask |
| `20:36` | First-frame latent |
| `36:52` | Track map |
The first three entries form the usual 36-channel Wan I2V input. A
`TrackEncoder` rasterizes sparse points on the VAE grid and appends the
16-channel track map.
## Prepare point-track data
Each video needs an `.npz` sidecar referenced by `points_path` in its metadata.
The archive must contain:
- `tracks`: floating-point source-pixel coordinates with shape `[T, N, 2]`.
The final dimension is `(x, y)`.
- `visibility`: a boolean or numeric visibility mask with shape `[T, N]`.
`T` is the source-video timeline and `N` is the number of point tracks. The
preprocessor applies the same temporal sample and center-crop/resize transform
to the sidecar as it applies to the video. Coordinates should therefore be in
the original video's pixel space, not normalized or pre-cropped. Points outside
the retained crop are marked invisible.
For example, a merged-dataset annotation can contain:
```json
[
{
"path": "videos/clip_0001.mp4",
"points_path": "tracks/clip_0001.npz",
"cap": ["A cyclist follows a winding road."],
"resolution": {"width": 1920, "height": 1080},
"fps": 24.0,
"duration": 5.0,
"num_frames": 120
}
]
```
The paths are relative to the dataset root named in the merge file:
```text
data/wantrack/raw,data/wantrack/metadata.json
```
All examples combined into one training batch must have a stackable point
dimension. Keeping `N` fixed across the dataset is the simplest option.
Run the I2V-track preprocessor with matching pixel and latent lengths. Wan's
temporal compression is four, so 81 pixel frames produce 21 latent frames:
```bash
torchrun --nproc_per_node=1 fastvideo/pipelines/preprocess/v1_preprocess.py \
--model_path models/wantrack-control-init \
--data_merge_path data/wantrack/data_merge.txt \
--output_dir data/wantrack/preprocessed \
--preprocess_task i2v_track \
--num_frames 81 \
--num_latent_t 21 \
--max_height 480 \
--max_width 832 \
--train_fps 16 \
--preprocess_video_batch_size 1
```
The trainer reads the resulting
`data/wantrack/preprocessed/combined_parquet_dataset` directory with
`preprocessed_data_type: i2v_track`.
## Train the bidirectional model
The bidirectional training wrapper is
`fastvideo.train.models.wantrack.WanTrackModel`; it loads
`TrackWanTransformer3DModel`.
```bash
torchrun --nproc_per_node=8 -m fastvideo.train.entrypoint.train \
--config examples/train/configs/fine_tuning/wantrack/bidirectional_i2v.yaml
```
The example optionally subsamples points and applies temporal track masking.
Track IDs are sampled once for a training sample and then reused by every
denoising call and its conditional/unconditional branches. Re-sampling IDs per
call would change the point embeddings even when the coordinates are
unchanged.
## Train the causal model
The causal wrapper is
`fastvideo.train.models.wantrack.WanTrackCausalModel`; it loads
`CausalTrackWanTransformer3DModel`. The example uses
`TeacherForcingSFTMethod`, so clean history and the noisy current chunk receive
the same I2V and track conditioning:
```bash
torchrun --nproc_per_node=8 -m fastvideo.train.entrypoint.train \
--config examples/train/configs/fine_tuning/wantrack/causal_i2v.yaml
```
Track encoding is causal at the VAE temporal boundary. The implementation
left-pads the complete source track sequence at its global beginning, encodes
it, and then slices the latent track map at the current latent `start_frame`.
It must not independently pad and encode each later chunk: doing so would
reset temporal context and produce a different feature for the same point
history.
The same stable `track_ids` are reused across all chunks, denoising steps, and
CFG branches. Causal sampling follows the RobotWM streaming contract:
`predict_noise_streaming()` owns the KV caches, each denoised block is committed
once as clean context, and caches are cleared at the sample boundary.
## Validation and inference
Both example configs use
`fastvideo.train.callbacks.track_validation.TrackValidationCallback`. It loads
fixed samples from the preprocessed WanTrack parquet, builds conditions through
the student's normal `prepare_batch()` path, generates videos, overlays the
active tracks, and logs them to the configured tracker. Set
`callbacks.track_validation.val_data_path` to use a dedicated validation
parquet; otherwise it samples from the training parquet.
The callback and standalone callers share two helpers:
```python
from fastvideo.train.models.wantrack.inference import (
prepare_wantrack_batch,
sample_wantrack,
)
batch = prepare_wantrack_batch(
model,
raw_batch,
seed=1000,
latents_source="zeros",
)
latents = sample_wantrack(
model,
batch,
num_inference_steps=30,
seed=1000,
text_guidance_scale=3.0,
motion_guidance_scale=1.5,
)
video = model.decode_latents(latents)
```
`sample_wantrack()` denoises the complete clip for `WanTrackModel`. For
`WanTrackCausalModel`, it uses the existing `CausalModelBase` streaming API and
the transformer's configured block size; it does not modify or fork the common
Wan causal denoising stage.
+1 -4
View File
@@ -78,10 +78,7 @@ If forcing a backend fails, verify optional dependencies are installed:
- `SAGE_ATTN_THREE`: upstream `sageattn3` package
- `ATTN_QAT_INFER`: `fastvideo-kernel` checkout/source install that exposes
`attn_qat_infer`
- `ATTN_QAT_TRAIN`: `fastvideo-kernel`; its runtime-JIT Triton implementation
selects an optimized route on SM100, joins the quantized and STE P@V paths on
SM120, and retains the previous route for unsupported configurations. See
[Attn-QAT Training](../training/attn_qat.md) for architecture controls.
- `ATTN_QAT_TRAIN`: `fastvideo-kernel` install exposing `fastvideo_kernel`
As a fallback, use:
-17
View File
@@ -1,17 +0,0 @@
# LingBot World 2 Example Dataset
These files were copied unchanged from the LingBot World 2 repository for the
FastVideo causal-fast inference example.
- Repository: `https://github.com/Robbyant/lingbot-world-v2.git`
- Source commit: `94f43115de8d4a4f9f282126528c300a0b232c5f`
- Source directory: `examples/03`
## Files
- `image.jpg`: source image for image-to-video generation. SHA-256:
`6ee3dacfef32cfef504dd698adb8a660cf15f686535c52fed4903fef27c0edd0`
- `poses.npy`: camera-to-world trajectory matrices. SHA-256:
`bd0a23a696e184b0b43e7767eb432bfe644690560fe327fa96961affc941c404`
- `intrinsics.npy`: camera intrinsic parameters. SHA-256:
`821fca6cf957ae8fbb1181307f02479efb1705e04c9e05734cd02fb43462e082`
Binary file not shown.

Before

Width:  |  Height:  |  Size: 1.2 MiB

Binary file not shown.
Binary file not shown.
@@ -1,52 +0,0 @@
"""Generate a five-second Dense LingBot-Video clip with the official defaults."""
import argparse
from pathlib import Path
from fastvideo import VideoGenerator
def parse_args() -> argparse.Namespace:
"""Parse the converted checkpoint and output paths for the sample."""
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument(
"--model-path",
type=Path,
required=True,
help="Path to a converted Dense LingBot-Video checkpoint.",
)
parser.add_argument(
"--output-path",
type=Path,
default=Path("outputs/lingbot-video/dense-t2v"),
help="Directory for the generated video.",
)
return parser.parse_args()
def main() -> None:
"""Load the converted Dense checkpoint and generate the default T2V sample."""
args = parse_args()
generator = VideoGenerator.from_pretrained(
str(args.model_path),
num_gpus=1,
use_fsdp_inference=False,
text_encoder_cpu_offload=True,
vae_cpu_offload=False,
pin_cpu_memory=True,
)
try:
generator.generate({
"prompt": "A red fox runs through fresh snow at sunrise.",
"output": {
"output_path": str(args.output_path),
"save_video": True,
"return_frames": False,
},
})
finally:
generator.shutdown()
if __name__ == "__main__":
main()
@@ -1,52 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Run LingBot World 2 14B causal-fast I2V generation with FastVideo."""
import os
from pathlib import Path
from fastvideo import VideoGenerator
REPO_ROOT = Path(__file__).resolve().parents[3]
DATASET_DIR = REPO_ROOT / "examples" / "dataset" / "lingbotworld2"
OUTPUT_PATH = REPO_ROOT / "outputs" / "lingbotworld2_causal_fast.mp4"
def main() -> None:
"""Load the native FastVideo LingBot World 2 causal-fast pipeline and generate one video."""
generator = VideoGenerator.from_pretrained(
os.environ["LINGBOTWORLD2_MODEL_PATH"],
num_gpus=8,
sp_size=8,
hsdp_shard_dim=8,
use_fsdp_inference=True,
dit_layerwise_offload=False,
dit_cpu_offload=False,
vae_cpu_offload=False,
text_encoder_cpu_offload=False,
pin_cpu_memory=True,
override_pipeline_cls_name="LingBotWorld2CausalFastPipeline",
)
try:
generator.generate_video(
"A serene lakeside scene with a lone tree standing in calm water, surrounded by distant snow-capped mountains under a bright blue sky with drifting white clouds; gentle ripples reflect the tree and sky, creating a tranquil, meditative atmosphere.",
image_path=str(DATASET_DIR / "image.jpg"),
action_path=str(DATASET_DIR),
output_path=str(OUTPUT_PATH),
save_video=True,
height=480,
width=832,
num_frames=65,
num_inference_steps=4,
guidance_scale=1.0,
negative_prompt="",
fps=16,
seed=42,
)
finally:
generator.shutdown()
if __name__ == "__main__":
main()
@@ -1,156 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Run causal WanTrack Self-Forcing I2V through FastVideo.
"""
import argparse
from pathlib import Path
from fastvideo import VideoGenerator
from fastvideo.api import (
EngineConfig,
GenerationRequest,
GeneratorConfig,
InputConfig,
OffloadConfig,
OutputConfig,
ParallelismConfig,
PipelineSelection,
SamplingConfig,
)
from fastvideo.models.dits.trackwan.utils import load_tracks
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="Run causal WanTrack Self-Forcing I2V.")
parser.add_argument(
"--model-path",
default="path/to/ckpt",
help="Local diffusers-format Track-v0 weights directory (or HF id).",
)
parser.add_argument(
"--image",
default="https://huggingface.co/datasets/YiYiXu/testing-images/resolve/main/wan_i2v_input.JPG",
help="Condition image path or URL.",
)
parser.add_argument(
"--tracks",
default=None,
help=("Track package: .pt/.npz dict with track_points+track_visibility "
"(optional track_ids), or a directory containing those files."),
)
parser.add_argument(
"--track-points",
default=None,
help="Optional standalone track_points tensor file (.pt/.npy/.npz).",
)
parser.add_argument(
"--track-visibility",
default=None,
help="Optional standalone track_visibility tensor file (.pt/.npy/.npz).",
)
parser.add_argument(
"--track-ids",
default=None,
help="Optional standalone track_ids tensor file (.pt/.npy/.npz).",
)
parser.add_argument(
"--output",
default="video_samples_wantrack_causal_sf/output.mp4",
help="Output mp4 path.",
)
parser.add_argument(
"--prompt",
default=("Summer beach vacation style, a white cat wearing sunglasses "
"sits on a surfboard. The fluffy-furred feline gazes directly "
"at the camera with a relaxed expression."),
help="Text prompt.",
)
parser.add_argument("--height", type=int, default=480)
parser.add_argument("--width", type=int, default=832)
parser.add_argument("--num-frames", type=int, default=121)
parser.add_argument("--fps", type=int, default=16)
parser.add_argument("--steps", type=int, default=4)
parser.add_argument("--guidance-scale", type=float, default=1.0)
parser.add_argument("--seed", type=int, default=1000)
parser.add_argument(
"--num-tracks",
type=int,
default=8,
help="Demo track count when no track files are provided.",
)
parser.add_argument("--num-gpus", type=int, default=1)
parser.add_argument("--tp-size", type=int, default=None)
parser.add_argument("--sp-size", type=int, default=None)
return parser.parse_args()
def main() -> None:
args = parse_args()
output = Path(args.output)
output.parent.mkdir(parents=True, exist_ok=True)
tp_size = args.tp_size if args.tp_size is not None else (args.num_gpus if args.num_gpus > 1 else 1)
sp_size = args.sp_size if args.sp_size is not None else (1 if args.num_gpus > 1 else args.num_gpus)
tracks = load_tracks(
tracks_path=args.tracks,
track_points_path=args.track_points,
track_visibility_path=args.track_visibility,
track_ids_path=args.track_ids,
num_frames=args.num_frames,
num_tracks=args.num_tracks,
seed=args.seed,
)
generator_config = GeneratorConfig(
model_path=args.model_path,
engine=EngineConfig(
num_gpus=args.num_gpus,
use_fsdp_inference=False,
parallelism=ParallelismConfig(tp_size=tp_size, sp_size=sp_size),
offload=OffloadConfig(
dit=False,
vae=False,
text_encoder=True,
image_encoder=True,
pin_cpu_memory=True,
),
),
pipeline=PipelineSelection(workload_type="i2v"),
)
generator = VideoGenerator.from_config(generator_config)
try:
request = GenerationRequest(
prompt=args.prompt,
inputs=InputConfig(
image_path=args.image,
track_points=tracks["track_points"],
track_visibility=tracks["track_visibility"],
track_ids=tracks["track_ids"],
),
sampling=SamplingConfig(
height=args.height,
width=args.width,
num_frames=args.num_frames,
fps=args.fps,
num_inference_steps=args.steps,
guidance_scale=args.guidance_scale,
seed=args.seed,
),
output=OutputConfig(
output_path=str(output.parent),
output_video_name=output.stem,
save_video=True,
return_frames=False,
),
)
result = generator.generate(request)
if isinstance(result, list):
result = result[0]
print(f"Saved video to {result.video_path}")
finally:
generator.shutdown()
if __name__ == "__main__":
main()
-96
View File
@@ -1,96 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Run Z-Image-Turbo text-to-image generation through FastVideo.
User story:
"I want the official Z-Image-Turbo defaults and a deterministic PNG from
a local or Hugging Face checkpoint."
"""
import argparse
from pathlib import Path
from fastvideo import VideoGenerator
from fastvideo.api import (
EngineConfig,
GenerationRequest,
GeneratorConfig,
OutputConfig,
ParallelismConfig,
PipelineSelection,
SamplingConfig,
)
DEFAULT_PROMPT = (
"Young Chinese woman in red Hanfu, intricate embroidery. Impeccable makeup, red floral forehead pattern. "
"Elaborate high bun, golden phoenix headdress, red flowers, beads. Holds round folding fan with lady, trees, bird. "
"Neon lightning-bolt lamp (⚡️), bright yellow glow, above extended left palm. Soft-lit outdoor night background, "
"silhouetted tiered pagoda (西安大雁塔), blurred colorful distant lights."
)
DEFAULT_REVISION = "f332072aa78be7aecdf3ee76d5c247082da564a6"
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="Run Z-Image-Turbo text-to-image generation.")
parser.add_argument("--model-path", default="Tongyi-MAI/Z-Image-Turbo")
parser.add_argument("--revision", default=DEFAULT_REVISION)
parser.add_argument("--output", default="outputs/zimage/zimage_turbo.png")
parser.add_argument("--prompt", default=DEFAULT_PROMPT)
parser.add_argument("--negative-prompt", default="")
parser.add_argument("--height", type=int, default=1024)
parser.add_argument("--width", type=int, default=1024)
parser.add_argument("--steps", type=int, default=8)
parser.add_argument("--guidance-scale", type=float, default=0.0)
parser.add_argument("--max-sequence-length", type=int, default=512)
parser.add_argument("--cfg-normalization", action=argparse.BooleanOptionalAction, default=False)
parser.add_argument("--cfg-truncation", type=float, default=1.0)
parser.add_argument("--seed", type=int, default=42)
return parser.parse_args()
def main() -> None:
args = parse_args()
output = Path(args.output)
output.parent.mkdir(parents=True, exist_ok=True)
generator = VideoGenerator.from_config(
GeneratorConfig(
model_path=args.model_path,
revision=args.revision,
engine=EngineConfig(
num_gpus=1,
parallelism=ParallelismConfig(tp_size=1, sp_size=1),
use_fsdp_inference=False,
),
# The model registry selects the native zimage_turbo preset.
pipeline=PipelineSelection(workload_type="t2i"),
))
try:
generator.generate(
GenerationRequest(
prompt=args.prompt,
negative_prompt=args.negative_prompt,
sampling=SamplingConfig(
height=args.height,
width=args.width,
num_frames=1,
fps=1,
num_inference_steps=args.steps,
guidance_scale=args.guidance_scale,
max_sequence_length=args.max_sequence_length,
cfg_normalization=args.cfg_normalization,
cfg_truncation=args.cfg_truncation,
seed=args.seed,
),
output=OutputConfig(
output_path=str(output),
save_video=True,
return_frames=False,
),
))
finally:
generator.shutdown()
if __name__ == "__main__":
main()
-4
View File
@@ -91,7 +91,3 @@ examples/train/
```
See `configs/README.md` and `scenario/README.md` for details.
The featured QAD Wan2.1 MixKit scenario runs Attn-QAT supervised fine-tuning,
checkpoint export, and three-step DMD2. See the
[Attn-QAT training guide](https://haoailab.com/FastVideo/training/attn_qat/).
@@ -1,86 +0,0 @@
models:
student:
_target_: fastvideo.train.models.wan.WanCausalModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
transformer_override_safetensor: /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15/A12/export/tf/transformer/model.safetensors
teacher:
_target_: fastvideo.train.models.wan.WanCausalModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: false
transformer_override_safetensor: /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15/A12/export/tf/transformer/model.safetensors
ema:
_target_: fastvideo.train.models.wan.WanCausalModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: false
transformer_override_safetensor: /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15/A12/export/tf/transformer/model.safetensors
method:
_target_: fastvideo.train.methods.consistency_model.causal_cd.CausalConsistencyDistillationMethod
discrete_cd_N: 48
guidance_scale: 3.0
ema_decay: 0.99
ema_start_step: 200
training:
distributed:
num_gpus: 4
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 4
data:
data_path: /mnt/lustre/vlm-s4duan/openvid_1m/combined_parquet_dataset
dataloader_type: streaming
streaming_manifest_path: /mnt/lustre/vlm-k1kong/dataset-index/openvid/streaming-t2v-v2.json
streaming_read_batch_size: 2
streaming_shuffle_row_groups: true
dataloader_num_workers: 0
train_batch_size: 2
training_cfg_rate: 0.0
seed: 1000
num_latent_t: 21
num_height: 480
num_width: 832
num_frames: 81
optimizer:
learning_rate: 2.0e-06
betas:
- 0.0
- 0.999
weight_decay: 0.01
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 2000
gradient_accumulation_steps: 8
checkpoint:
output_dir: /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15/A12/cd/checkpoints
training_state_checkpointing_steps: 1000
checkpoints_total_limit: 2
tracker:
trackers:
- wandb
project_name: causal_forcing_openvid_a12_a15
run_name: A12_sink1_local6_relativistic_chunk3_cd2k_from_tf3k_openvid_81f21l_gbs64
model:
enable_gradient_checkpointing_type: full
callbacks:
grad_clip:
max_grad_norm: 1.0
validation:
pipeline_target: fastvideo.pipelines.basic.wan.wan_causal_pipeline.WanCausalPipeline
dataset_file: /mnt/nfs/vlm-k1kong/FastVideo-openvid-a12-a15-final-20260717/examples/training/finetune/Wan2.1-VSA/Wan-Syn-Data/validation_4.json
every_steps: 200
sampling_steps:
- 4
guidance_scale: 3.0
num_frames: 81
output_dir: /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15/A12/cd/validation
offload_training_state: true
unload_pipeline_after_validation: true
pipeline:
flow_shift: 5
dit_config:
local_attn_size: 6
sink_size: 1
rope_cache_policy: relativistic
causal_train_attention: triton
@@ -1,110 +0,0 @@
models:
student:
_target_: fastvideo.train.models.wan.WanCausalModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
transformer_override_safetensor: /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15/A12/export/cd/transformer/model.safetensors
teacher:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-14B-Diffusers
trainable: false
disable_custom_init_weights: true
critic:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
disable_custom_init_weights: true
method:
_target_: fastvideo.train.methods.distribution_matching.self_forcing.SelfForcingMethod
rollout_mode: simulate
generator_update_interval: 5
real_score_guidance_scale: 4.0
dmd_denoising_steps:
- 1000
- 750
- 500
- 250
warp_denoising_step: true
chunk_size: 3
student_sample_type: sde
context_noise: 0.0
same_step_across_blocks: true
enable_gradient_in_rollout: true
start_gradient_frame: 0
fake_score_learning_rate: 4.0e-07
fake_score_betas:
- 0.0
- 0.999
fake_score_lr_scheduler: constant
training:
distributed:
num_gpus: 4
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 4
data:
data_path: /mnt/lustre/vlm-s4duan/openvid_1m/combined_parquet_dataset
dataloader_type: streaming
streaming_manifest_path: /mnt/lustre/vlm-k1kong/dataset-index/openvid/streaming-t2v-v2.json
streaming_read_batch_size: 2
streaming_shuffle_row_groups: true
dataloader_num_workers: 0
train_batch_size: 2
training_cfg_rate: 0.0
seed: 1000
num_latent_t: 21
num_height: 480
num_width: 832
num_frames: 81
optimizer:
learning_rate: 2.0e-06
betas:
- 0.0
- 0.999
weight_decay: 0.01
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 1000
gradient_accumulation_steps: 8
checkpoint:
output_dir: /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15/A12/sf/checkpoints
training_state_checkpointing_steps: 1000
checkpoints_total_limit: 1
tracker:
trackers:
- wandb
project_name: causal_forcing_openvid_a12_a15
run_name: A12_sink1_local6_relativistic_chunk3_sf1k_from_cd2k_openvid_81f21l_gbs64
model:
enable_gradient_checkpointing_type: full
callbacks:
ema:
decay: 0.99
start_iter: 200
grad_clip:
max_grad_norm: 1.0
validation:
pipeline_target: fastvideo.pipelines.basic.wan.wan_causal_dmd_pipeline.WanCausalDMDPipeline
dataset_file: /mnt/nfs/vlm-k1kong/FastVideo-openvid-a12-a15-final-20260717/examples/training/finetune/Wan2.1-VSA/Wan-Syn-Data/validation_4.json
every_steps: 200
sampling_steps:
- 4
guidance_scale: 3.0
num_frames: 81
output_dir: /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15/A12/sf/validation
offload_training_state: true
unload_pipeline_after_validation: true
sampling_timesteps:
- 1000
- 750
- 500
- 250
pipeline:
flow_shift: 5
dit_config:
local_attn_size: 6
sink_size: 1
rope_cache_policy: relativistic
causal_train_attention: triton
@@ -1,72 +0,0 @@
models:
student:
_target_: fastvideo.train.models.wan.WanCausalModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
method:
_target_: fastvideo.train.methods.fine_tuning.tfsft.TeacherForcingSFTMethod
chunk_size: 3
training:
distributed:
num_gpus: 4
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 4
data:
data_path: /mnt/lustre/vlm-s4duan/openvid_1m/combined_parquet_dataset
dataloader_type: streaming
streaming_manifest_path: /mnt/lustre/vlm-k1kong/dataset-index/openvid/streaming-t2v-v2.json
streaming_read_batch_size: 2
streaming_shuffle_row_groups: true
dataloader_num_workers: 0
train_batch_size: 2
training_cfg_rate: 0.0
seed: 1000
num_latent_t: 21
num_height: 480
num_width: 832
num_frames: 81
optimizer:
learning_rate: 2.0e-06
betas:
- 0.0
- 0.999
weight_decay: 0.01
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 3000
gradient_accumulation_steps: 8
checkpoint:
output_dir: /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15/A12/tf/checkpoints
training_state_checkpointing_steps: 1000
checkpoints_total_limit: 3
tracker:
trackers:
- wandb
project_name: causal_forcing_openvid_a12_a15
run_name: A12_sink1_local6_relativistic_chunk3_tf3k_openvid_81f21l_gbs64
model:
enable_gradient_checkpointing_type: full
callbacks:
grad_clip:
max_grad_norm: 1.0
validation:
pipeline_target: fastvideo.pipelines.basic.wan.wan_causal_pipeline.WanCausalPipeline
dataset_file: /mnt/nfs/vlm-k1kong/FastVideo-openvid-a12-a15-final-20260717/examples/training/finetune/Wan2.1-VSA/Wan-Syn-Data/validation_4.json
every_steps: 200
sampling_steps:
- 40
guidance_scale: 6.0
num_frames: 81
output_dir: /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15/A12/tf/validation
offload_training_state: true
unload_pipeline_after_validation: true
pipeline:
flow_shift: 5
dit_config:
local_attn_size: 6
sink_size: 1
rope_cache_policy: relativistic
causal_train_attention: triton
@@ -1,86 +0,0 @@
models:
student:
_target_: fastvideo.train.models.wan.WanCausalModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
transformer_override_safetensor: /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15/A13/export/tf/transformer/model.safetensors
teacher:
_target_: fastvideo.train.models.wan.WanCausalModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: false
transformer_override_safetensor: /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15/A13/export/tf/transformer/model.safetensors
ema:
_target_: fastvideo.train.models.wan.WanCausalModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: false
transformer_override_safetensor: /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15/A13/export/tf/transformer/model.safetensors
method:
_target_: fastvideo.train.methods.consistency_model.causal_cd.CausalConsistencyDistillationMethod
discrete_cd_N: 48
guidance_scale: 3.0
ema_decay: 0.99
ema_start_step: 200
training:
distributed:
num_gpus: 4
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 4
data:
data_path: /mnt/lustre/vlm-s4duan/openvid_1m/combined_parquet_dataset
dataloader_type: streaming
streaming_manifest_path: /mnt/lustre/vlm-k1kong/dataset-index/openvid/streaming-t2v-v2.json
streaming_read_batch_size: 2
streaming_shuffle_row_groups: true
dataloader_num_workers: 0
train_batch_size: 2
training_cfg_rate: 0.0
seed: 1000
num_latent_t: 21
num_height: 480
num_width: 832
num_frames: 81
optimizer:
learning_rate: 2.0e-06
betas:
- 0.0
- 0.999
weight_decay: 0.01
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 2000
gradient_accumulation_steps: 8
checkpoint:
output_dir: /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15/A13/cd/checkpoints
training_state_checkpointing_steps: 1000
checkpoints_total_limit: 2
tracker:
trackers:
- wandb
project_name: causal_forcing_openvid_a12_a15
run_name: A13_sink0_local6_relativistic_chunk3_cd2k_from_tf3k_openvid_81f21l_gbs64
model:
enable_gradient_checkpointing_type: full
callbacks:
grad_clip:
max_grad_norm: 1.0
validation:
pipeline_target: fastvideo.pipelines.basic.wan.wan_causal_pipeline.WanCausalPipeline
dataset_file: /mnt/nfs/vlm-k1kong/FastVideo-openvid-a12-a15-final-20260717/examples/training/finetune/Wan2.1-VSA/Wan-Syn-Data/validation_4.json
every_steps: 200
sampling_steps:
- 4
guidance_scale: 3.0
num_frames: 81
output_dir: /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15/A13/cd/validation
offload_training_state: true
unload_pipeline_after_validation: true
pipeline:
flow_shift: 5
dit_config:
local_attn_size: 6
sink_size: 0
rope_cache_policy: relativistic
causal_train_attention: triton
@@ -1,110 +0,0 @@
models:
student:
_target_: fastvideo.train.models.wan.WanCausalModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
transformer_override_safetensor: /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15/A13/export/cd/transformer/model.safetensors
teacher:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-14B-Diffusers
trainable: false
disable_custom_init_weights: true
critic:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
disable_custom_init_weights: true
method:
_target_: fastvideo.train.methods.distribution_matching.self_forcing.SelfForcingMethod
rollout_mode: simulate
generator_update_interval: 5
real_score_guidance_scale: 4.0
dmd_denoising_steps:
- 1000
- 750
- 500
- 250
warp_denoising_step: true
chunk_size: 3
student_sample_type: sde
context_noise: 0.0
same_step_across_blocks: true
enable_gradient_in_rollout: true
start_gradient_frame: 0
fake_score_learning_rate: 4.0e-07
fake_score_betas:
- 0.0
- 0.999
fake_score_lr_scheduler: constant
training:
distributed:
num_gpus: 4
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 4
data:
data_path: /mnt/lustre/vlm-s4duan/openvid_1m/combined_parquet_dataset
dataloader_type: streaming
streaming_manifest_path: /mnt/lustre/vlm-k1kong/dataset-index/openvid/streaming-t2v-v2.json
streaming_read_batch_size: 2
streaming_shuffle_row_groups: true
dataloader_num_workers: 0
train_batch_size: 2
training_cfg_rate: 0.0
seed: 1000
num_latent_t: 21
num_height: 480
num_width: 832
num_frames: 81
optimizer:
learning_rate: 2.0e-06
betas:
- 0.0
- 0.999
weight_decay: 0.01
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 1000
gradient_accumulation_steps: 8
checkpoint:
output_dir: /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15/A13/sf/checkpoints
training_state_checkpointing_steps: 1000
checkpoints_total_limit: 1
tracker:
trackers:
- wandb
project_name: causal_forcing_openvid_a12_a15
run_name: A13_sink0_local6_relativistic_chunk3_sf1k_from_cd2k_openvid_81f21l_gbs64
model:
enable_gradient_checkpointing_type: full
callbacks:
ema:
decay: 0.99
start_iter: 200
grad_clip:
max_grad_norm: 1.0
validation:
pipeline_target: fastvideo.pipelines.basic.wan.wan_causal_dmd_pipeline.WanCausalDMDPipeline
dataset_file: /mnt/nfs/vlm-k1kong/FastVideo-openvid-a12-a15-final-20260717/examples/training/finetune/Wan2.1-VSA/Wan-Syn-Data/validation_4.json
every_steps: 200
sampling_steps:
- 4
guidance_scale: 3.0
num_frames: 81
output_dir: /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15/A13/sf/validation
offload_training_state: true
unload_pipeline_after_validation: true
sampling_timesteps:
- 1000
- 750
- 500
- 250
pipeline:
flow_shift: 5
dit_config:
local_attn_size: 6
sink_size: 0
rope_cache_policy: relativistic
causal_train_attention: triton
@@ -1,72 +0,0 @@
models:
student:
_target_: fastvideo.train.models.wan.WanCausalModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
method:
_target_: fastvideo.train.methods.fine_tuning.tfsft.TeacherForcingSFTMethod
chunk_size: 3
training:
distributed:
num_gpus: 4
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 4
data:
data_path: /mnt/lustre/vlm-s4duan/openvid_1m/combined_parquet_dataset
dataloader_type: streaming
streaming_manifest_path: /mnt/lustre/vlm-k1kong/dataset-index/openvid/streaming-t2v-v2.json
streaming_read_batch_size: 2
streaming_shuffle_row_groups: true
dataloader_num_workers: 0
train_batch_size: 2
training_cfg_rate: 0.0
seed: 1000
num_latent_t: 21
num_height: 480
num_width: 832
num_frames: 81
optimizer:
learning_rate: 2.0e-06
betas:
- 0.0
- 0.999
weight_decay: 0.01
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 3000
gradient_accumulation_steps: 8
checkpoint:
output_dir: /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15/A13/tf/checkpoints
training_state_checkpointing_steps: 1000
checkpoints_total_limit: 3
tracker:
trackers:
- wandb
project_name: causal_forcing_openvid_a12_a15
run_name: A13_sink0_local6_relativistic_chunk3_tf3k_openvid_81f21l_gbs64
model:
enable_gradient_checkpointing_type: full
callbacks:
grad_clip:
max_grad_norm: 1.0
validation:
pipeline_target: fastvideo.pipelines.basic.wan.wan_causal_pipeline.WanCausalPipeline
dataset_file: /mnt/nfs/vlm-k1kong/FastVideo-openvid-a12-a15-final-20260717/examples/training/finetune/Wan2.1-VSA/Wan-Syn-Data/validation_4.json
every_steps: 200
sampling_steps:
- 40
guidance_scale: 6.0
num_frames: 81
output_dir: /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15/A13/tf/validation
offload_training_state: true
unload_pipeline_after_validation: true
pipeline:
flow_shift: 5
dit_config:
local_attn_size: 6
sink_size: 0
rope_cache_policy: relativistic
causal_train_attention: triton
@@ -1,86 +0,0 @@
models:
student:
_target_: fastvideo.train.models.wan.WanCausalModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
transformer_override_safetensor: /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15/A14/export/tf/transformer/model.safetensors
teacher:
_target_: fastvideo.train.models.wan.WanCausalModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: false
transformer_override_safetensor: /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15/A14/export/tf/transformer/model.safetensors
ema:
_target_: fastvideo.train.models.wan.WanCausalModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: false
transformer_override_safetensor: /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15/A14/export/tf/transformer/model.safetensors
method:
_target_: fastvideo.train.methods.consistency_model.causal_cd.CausalConsistencyDistillationMethod
discrete_cd_N: 48
guidance_scale: 3.0
ema_decay: 0.99
ema_start_step: 200
training:
distributed:
num_gpus: 4
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 4
data:
data_path: /mnt/lustre/vlm-s4duan/openvid_1m/combined_parquet_dataset
dataloader_type: streaming
streaming_manifest_path: /mnt/lustre/vlm-k1kong/dataset-index/openvid/streaming-t2v-v2.json
streaming_read_batch_size: 2
streaming_shuffle_row_groups: true
dataloader_num_workers: 0
train_batch_size: 2
training_cfg_rate: 0.0
seed: 1000
num_latent_t: 21
num_height: 480
num_width: 832
num_frames: 81
optimizer:
learning_rate: 2.0e-06
betas:
- 0.0
- 0.999
weight_decay: 0.01
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 2000
gradient_accumulation_steps: 8
checkpoint:
output_dir: /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15/A14/cd/checkpoints
training_state_checkpointing_steps: 1000
checkpoints_total_limit: 2
tracker:
trackers:
- wandb
project_name: causal_forcing_openvid_a12_a15
run_name: A14_sink1_local6_absolute_chunk3_cd2k_from_tf3k_openvid_81f21l_gbs64
model:
enable_gradient_checkpointing_type: full
callbacks:
grad_clip:
max_grad_norm: 1.0
validation:
pipeline_target: fastvideo.pipelines.basic.wan.wan_causal_pipeline.WanCausalPipeline
dataset_file: /mnt/nfs/vlm-k1kong/FastVideo-openvid-a12-a15-final-20260717/examples/training/finetune/Wan2.1-VSA/Wan-Syn-Data/validation_4.json
every_steps: 200
sampling_steps:
- 4
guidance_scale: 3.0
num_frames: 81
output_dir: /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15/A14/cd/validation
offload_training_state: true
unload_pipeline_after_validation: true
pipeline:
flow_shift: 5
dit_config:
local_attn_size: 6
sink_size: 1
rope_cache_policy: absolute
causal_train_attention: triton
@@ -1,110 +0,0 @@
models:
student:
_target_: fastvideo.train.models.wan.WanCausalModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
transformer_override_safetensor: /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15/A14/export/cd/transformer/model.safetensors
teacher:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-14B-Diffusers
trainable: false
disable_custom_init_weights: true
critic:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
disable_custom_init_weights: true
method:
_target_: fastvideo.train.methods.distribution_matching.self_forcing.SelfForcingMethod
rollout_mode: simulate
generator_update_interval: 5
real_score_guidance_scale: 4.0
dmd_denoising_steps:
- 1000
- 750
- 500
- 250
warp_denoising_step: true
chunk_size: 3
student_sample_type: sde
context_noise: 0.0
same_step_across_blocks: true
enable_gradient_in_rollout: true
start_gradient_frame: 0
fake_score_learning_rate: 4.0e-07
fake_score_betas:
- 0.0
- 0.999
fake_score_lr_scheduler: constant
training:
distributed:
num_gpus: 4
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 4
data:
data_path: /mnt/lustre/vlm-s4duan/openvid_1m/combined_parquet_dataset
dataloader_type: streaming
streaming_manifest_path: /mnt/lustre/vlm-k1kong/dataset-index/openvid/streaming-t2v-v2.json
streaming_read_batch_size: 2
streaming_shuffle_row_groups: true
dataloader_num_workers: 0
train_batch_size: 2
training_cfg_rate: 0.0
seed: 1000
num_latent_t: 21
num_height: 480
num_width: 832
num_frames: 81
optimizer:
learning_rate: 2.0e-06
betas:
- 0.0
- 0.999
weight_decay: 0.01
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 1000
gradient_accumulation_steps: 8
checkpoint:
output_dir: /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15/A14/sf/checkpoints
training_state_checkpointing_steps: 1000
checkpoints_total_limit: 1
tracker:
trackers:
- wandb
project_name: causal_forcing_openvid_a12_a15
run_name: A14_sink1_local6_absolute_chunk3_sf1k_from_cd2k_openvid_81f21l_gbs64
model:
enable_gradient_checkpointing_type: full
callbacks:
ema:
decay: 0.99
start_iter: 200
grad_clip:
max_grad_norm: 1.0
validation:
pipeline_target: fastvideo.pipelines.basic.wan.wan_causal_dmd_pipeline.WanCausalDMDPipeline
dataset_file: /mnt/nfs/vlm-k1kong/FastVideo-openvid-a12-a15-final-20260717/examples/training/finetune/Wan2.1-VSA/Wan-Syn-Data/validation_4.json
every_steps: 200
sampling_steps:
- 4
guidance_scale: 3.0
num_frames: 81
output_dir: /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15/A14/sf/validation
offload_training_state: true
unload_pipeline_after_validation: true
sampling_timesteps:
- 1000
- 750
- 500
- 250
pipeline:
flow_shift: 5
dit_config:
local_attn_size: 6
sink_size: 1
rope_cache_policy: absolute
causal_train_attention: triton
@@ -1,72 +0,0 @@
models:
student:
_target_: fastvideo.train.models.wan.WanCausalModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
method:
_target_: fastvideo.train.methods.fine_tuning.tfsft.TeacherForcingSFTMethod
chunk_size: 3
training:
distributed:
num_gpus: 4
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 4
data:
data_path: /mnt/lustre/vlm-s4duan/openvid_1m/combined_parquet_dataset
dataloader_type: streaming
streaming_manifest_path: /mnt/lustre/vlm-k1kong/dataset-index/openvid/streaming-t2v-v2.json
streaming_read_batch_size: 2
streaming_shuffle_row_groups: true
dataloader_num_workers: 0
train_batch_size: 2
training_cfg_rate: 0.0
seed: 1000
num_latent_t: 21
num_height: 480
num_width: 832
num_frames: 81
optimizer:
learning_rate: 2.0e-06
betas:
- 0.0
- 0.999
weight_decay: 0.01
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 3000
gradient_accumulation_steps: 8
checkpoint:
output_dir: /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15/A14/tf/checkpoints
training_state_checkpointing_steps: 1000
checkpoints_total_limit: 3
tracker:
trackers:
- wandb
project_name: causal_forcing_openvid_a12_a15
run_name: A14_sink1_local6_absolute_chunk3_tf3k_openvid_81f21l_gbs64
model:
enable_gradient_checkpointing_type: full
callbacks:
grad_clip:
max_grad_norm: 1.0
validation:
pipeline_target: fastvideo.pipelines.basic.wan.wan_causal_pipeline.WanCausalPipeline
dataset_file: /mnt/nfs/vlm-k1kong/FastVideo-openvid-a12-a15-final-20260717/examples/training/finetune/Wan2.1-VSA/Wan-Syn-Data/validation_4.json
every_steps: 200
sampling_steps:
- 40
guidance_scale: 6.0
num_frames: 81
output_dir: /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15/A14/tf/validation
offload_training_state: true
unload_pipeline_after_validation: true
pipeline:
flow_shift: 5
dit_config:
local_attn_size: 6
sink_size: 1
rope_cache_policy: absolute
causal_train_attention: triton
@@ -1,89 +0,0 @@
models:
student:
_target_: fastvideo.train.models.wan.WanCausalModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
transformer_override_safetensor: /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15/A15/export/tf/transformer/model.safetensors
num_frames_per_block: 1
teacher:
_target_: fastvideo.train.models.wan.WanCausalModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: false
transformer_override_safetensor: /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15/A15/export/tf/transformer/model.safetensors
num_frames_per_block: 1
ema:
_target_: fastvideo.train.models.wan.WanCausalModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: false
transformer_override_safetensor: /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15/A15/export/tf/transformer/model.safetensors
num_frames_per_block: 1
method:
_target_: fastvideo.train.methods.consistency_model.causal_cd.CausalConsistencyDistillationMethod
discrete_cd_N: 48
guidance_scale: 3.0
ema_decay: 0.99
ema_start_step: 200
training:
distributed:
num_gpus: 4
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 4
data:
data_path: /mnt/lustre/vlm-s4duan/openvid_1m/combined_parquet_dataset
dataloader_type: streaming
streaming_manifest_path: /mnt/lustre/vlm-k1kong/dataset-index/openvid/streaming-t2v-v2.json
streaming_read_batch_size: 2
streaming_shuffle_row_groups: true
dataloader_num_workers: 0
train_batch_size: 2
training_cfg_rate: 0.0
seed: 1000
num_latent_t: 21
num_height: 480
num_width: 832
num_frames: 81
optimizer:
learning_rate: 2.0e-06
betas:
- 0.0
- 0.999
weight_decay: 0.01
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 2000
gradient_accumulation_steps: 8
checkpoint:
output_dir: /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15/A15/cd/checkpoints
training_state_checkpointing_steps: 1000
checkpoints_total_limit: 2
tracker:
trackers:
- wandb
project_name: causal_forcing_openvid_a12_a15
run_name: A15_sink1_local6_relativistic_framewise_cd2k_from_tf3k_openvid_81f21l_gbs64
model:
enable_gradient_checkpointing_type: full
callbacks:
grad_clip:
max_grad_norm: 1.0
validation:
pipeline_target: fastvideo.pipelines.basic.wan.wan_causal_pipeline.WanCausalPipeline
dataset_file: /mnt/nfs/vlm-k1kong/FastVideo-openvid-a12-a15-final-20260717/examples/training/finetune/Wan2.1-VSA/Wan-Syn-Data/validation_4.json
every_steps: 200
sampling_steps:
- 4
guidance_scale: 3.0
num_frames: 81
output_dir: /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15/A15/cd/validation
offload_training_state: true
unload_pipeline_after_validation: true
pipeline:
flow_shift: 5
dit_config:
local_attn_size: 6
sink_size: 1
rope_cache_policy: relativistic
causal_train_attention: triton
@@ -1,111 +0,0 @@
models:
student:
_target_: fastvideo.train.models.wan.WanCausalModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
transformer_override_safetensor: /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15/A15/export/cd/transformer/model.safetensors
num_frames_per_block: 1
teacher:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-14B-Diffusers
trainable: false
disable_custom_init_weights: true
critic:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
disable_custom_init_weights: true
method:
_target_: fastvideo.train.methods.distribution_matching.self_forcing.SelfForcingMethod
rollout_mode: simulate
generator_update_interval: 5
real_score_guidance_scale: 4.0
dmd_denoising_steps:
- 1000
- 750
- 500
- 250
warp_denoising_step: true
chunk_size: 1
student_sample_type: sde
context_noise: 0.0
same_step_across_blocks: true
enable_gradient_in_rollout: true
start_gradient_frame: 0
fake_score_learning_rate: 4.0e-07
fake_score_betas:
- 0.0
- 0.999
fake_score_lr_scheduler: constant
training:
distributed:
num_gpus: 4
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 4
data:
data_path: /mnt/lustre/vlm-s4duan/openvid_1m/combined_parquet_dataset
dataloader_type: streaming
streaming_manifest_path: /mnt/lustre/vlm-k1kong/dataset-index/openvid/streaming-t2v-v2.json
streaming_read_batch_size: 2
streaming_shuffle_row_groups: true
dataloader_num_workers: 0
train_batch_size: 2
training_cfg_rate: 0.0
seed: 1000
num_latent_t: 21
num_height: 480
num_width: 832
num_frames: 81
optimizer:
learning_rate: 2.0e-06
betas:
- 0.0
- 0.999
weight_decay: 0.01
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 1000
gradient_accumulation_steps: 8
checkpoint:
output_dir: /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15/A15/sf/checkpoints
training_state_checkpointing_steps: 1000
checkpoints_total_limit: 1
tracker:
trackers:
- wandb
project_name: causal_forcing_openvid_a12_a15
run_name: A15_sink1_local6_relativistic_framewise_sf1k_from_cd2k_openvid_81f21l_gbs64
model:
enable_gradient_checkpointing_type: full
callbacks:
ema:
decay: 0.99
start_iter: 200
grad_clip:
max_grad_norm: 1.0
validation:
pipeline_target: fastvideo.pipelines.basic.wan.wan_causal_dmd_pipeline.WanCausalDMDPipeline
dataset_file: /mnt/nfs/vlm-k1kong/FastVideo-openvid-a12-a15-final-20260717/examples/training/finetune/Wan2.1-VSA/Wan-Syn-Data/validation_4.json
every_steps: 200
sampling_steps:
- 4
guidance_scale: 3.0
num_frames: 81
output_dir: /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15/A15/sf/validation
offload_training_state: true
unload_pipeline_after_validation: true
sampling_timesteps:
- 1000
- 750
- 500
- 250
pipeline:
flow_shift: 5
dit_config:
local_attn_size: 6
sink_size: 1
rope_cache_policy: relativistic
causal_train_attention: triton
@@ -1,73 +0,0 @@
models:
student:
_target_: fastvideo.train.models.wan.WanCausalModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
num_frames_per_block: 1
method:
_target_: fastvideo.train.methods.fine_tuning.tfsft.TeacherForcingSFTMethod
chunk_size: 1
training:
distributed:
num_gpus: 4
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 4
data:
data_path: /mnt/lustre/vlm-s4duan/openvid_1m/combined_parquet_dataset
dataloader_type: streaming
streaming_manifest_path: /mnt/lustre/vlm-k1kong/dataset-index/openvid/streaming-t2v-v2.json
streaming_read_batch_size: 2
streaming_shuffle_row_groups: true
dataloader_num_workers: 0
train_batch_size: 2
training_cfg_rate: 0.0
seed: 1000
num_latent_t: 21
num_height: 480
num_width: 832
num_frames: 81
optimizer:
learning_rate: 2.0e-06
betas:
- 0.0
- 0.999
weight_decay: 0.01
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 3000
gradient_accumulation_steps: 8
checkpoint:
output_dir: /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15/A15/tf/checkpoints
training_state_checkpointing_steps: 1000
checkpoints_total_limit: 3
tracker:
trackers:
- wandb
project_name: causal_forcing_openvid_a12_a15
run_name: A15_sink1_local6_relativistic_framewise_tf3k_openvid_81f21l_gbs64
model:
enable_gradient_checkpointing_type: full
callbacks:
grad_clip:
max_grad_norm: 1.0
validation:
pipeline_target: fastvideo.pipelines.basic.wan.wan_causal_pipeline.WanCausalPipeline
dataset_file: /mnt/nfs/vlm-k1kong/FastVideo-openvid-a12-a15-final-20260717/examples/training/finetune/Wan2.1-VSA/Wan-Syn-Data/validation_4.json
every_steps: 200
sampling_steps:
- 40
guidance_scale: 6.0
num_frames: 81
output_dir: /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15/A15/tf/validation
offload_training_state: true
unload_pipeline_after_validation: true
pipeline:
flow_shift: 5
dit_config:
local_attn_size: 6
sink_size: 1
rope_cache_policy: relativistic
causal_train_attention: triton
@@ -1,52 +0,0 @@
# OpenVid causal A12-A15 plan (NOT STARTED)
Prepared from commit `30ada30e4c6b05aa68cd1eb8940a34d149457147`. This directory contains configuration only;
no training command was launched.
Data source: `/mnt/lustre/vlm-s4duan/openvid_1m/combined_parquet_dataset` (4494 parquet files, about 5.5 TiB). It is owned by
`vlm-s4duan`; filesystem read permission exists, but obtain the owner's consent
before launching and coordinate I/O. Do not write into that directory.
The opt-in streaming loader projects only the 15 T2V columns, reads each
assigned row group sequentially, and stores its JSON manifest at
`/mnt/lustre/vlm-k1kong/dataset-index/openvid/streaming-t2v-v2.json` in user-owned Lustre. It uses zero DataLoader workers and never
writes a cache or index into the shared source tree.
All stages use 4 GPUs, microbatch 2/rank, gradient accumulation 8, hence global
batch = 2 * 4 * 8 = 64. `dataloader_num_workers=0` limits shared-memory and
Lustre prefetch pressure.
All four conditions use exactly 21 latent frames / 81 raw frames in TF, CD,
SF, and validation. Chunk-3 conditions therefore use seven identical
three-latent blocks. A15 is length-matched and uses framewise blocks.
A15 "framewise" means `num_frames_per_block=1` on every causal role, plus
`method.chunk_size=1` in TF and SF. Causal CD has no independent-frame timestep
option: it samples one t/t_next pair and broadcasts it over T, so A15 CD is
framewise causal attention but not framewise diffusion-time sampling.
The requested LR/betas apply to each stage's main optimizer: 2e-6 and
(0.0, 0.999). SF critic keeps the proven DMD value 4e-7 with (0.0, 0.999).
To launch one condition later:
export WANDB_API_KEY=...
bash /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15/scripts/run_condition.sh A12 /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15/A12 29820
To run all four sequentially on one 4-GPU node:
export WANDB_API_KEY=...
bash /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15/scripts/run_all_sequential.sh /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15
Online W&B is the default. For an intentionally offline launch, omit the key
and set `WANDB_MODE=offline`. Training checkpoints still resume from `latest`,
but W&B does not merge separate offline process restarts into one run; sync the
resulting offline runs individually later.
To validate all three configs for one condition without starting training,
creating W&B state, or requiring a key:
PREFLIGHT_ONLY=1 bash /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15/scripts/run_condition.sh A12 /mnt/lustre/vlm-k1kong/experiments/openvid_causal_a12_a15/20260717_014659_openvid_a12_a15/A12 29820
`PREFLIGHT_ONLY=1` is also supported by `run_all_sequential.sh`; it validates
all twelve configs and exits before queue state or checkpoint checks.
@@ -1,40 +0,0 @@
# MixKit-21 causal-kernel teacher-forcing ablation
All runs use the same training and validation settings. Only the three
attention axes listed in `experiment_matrix.tsv` vary.
| ID | sink | local | RoPE | lane | primary contrast |
|---|---:|---:|---|---|---|
| A01 | 0 | 21 | absolute | node0 | rope policy control |
| A02 | 0 | 21 | relativistic | node1 | rope policy control |
| A03 | 1 | 21 | relativistic | node0 | sink at local 21 |
| A04 | 0 | 6 | relativistic | node0 | local window at sink 0 |
| A05 | 1 | 6 | relativistic | node0 | sink at local 6 |
| A06 | 0 | 12 | relativistic | node1 | local window at sink 0 |
| A07 | 1 | 12 | relativistic | node1 | sink at local 12 |
| A08 | 3 | 12 | relativistic | node1 | sink size at local 12 |
Fixed invariants:
- MixKit precomputed data at 480x832 with 21 stored latent frames.
- Teacher forcing for 2,000 steps, batch size 1, chunk size 3.
- Fused training attention: `causal_train_attention: triton`.
- Validation at 249 pixel frames, equivalent to 63 Wan latent frames.
- Validation reuses the training transformer and pipeline config, so sink,
local window, and RoPE policy stay identical between training and inference.
- Four GPUs, full gradient checkpointing, validation/checkpoint every 200/1000
steps, and W&B online tracking in a dedicated project.
Validate the matrix and template:
```bash
python scripts/train/manage_mixkit21_tf_ablation.py validate
```
Run one lane after supplying `WANDB_API_KEY` at runtime:
```bash
SEQUENCE_ID=<timestamp> LANE=node0 \
RUN_CONDITIONS="A01 A03 A04 A05" \
bash scripts/train/run_mixkit21_tf_ablation_lane.sh
```
@@ -1,9 +0,0 @@
id condition sink_size local_attn_size rope_cache_policy lane primary_contrast
A01 sink0_local21_absolute 0 21 absolute node0 A01-vs-A02: rope policy control
A02 sink0_local21_relative 0 21 relativistic node1 A01-vs-A02: rope policy control
A03 sink1_local21_relative 1 21 relativistic node0 A02-vs-A03: sink at local21
A04 sink0_local6_relative 0 6 relativistic node0 A02-vs-A04: local window at sink0
A05 sink1_local6_relative 1 6 relativistic node0 A04-vs-A05: sink at local6
A06 sink0_local12_relative 0 12 relativistic node1 A02-vs-A06: local window at sink0
A07 sink1_local12_relative 1 12 relativistic node1 A06-vs-A07: sink at local12
A08 sink3_local12_relative 3 12 relativistic node1 A06-vs-A07-vs-A08: sink size at local12
1 id condition sink_size local_attn_size rope_cache_policy lane primary_contrast
2 A01 sink0_local21_absolute 0 21 absolute node0 A01-vs-A02: rope policy control
3 A02 sink0_local21_relative 0 21 relativistic node1 A01-vs-A02: rope policy control
4 A03 sink1_local21_relative 1 21 relativistic node0 A02-vs-A03: sink at local21
5 A04 sink0_local6_relative 0 6 relativistic node0 A02-vs-A04: local window at sink0
6 A05 sink1_local6_relative 1 6 relativistic node0 A04-vs-A05: sink at local6
7 A06 sink0_local12_relative 0 12 relativistic node1 A02-vs-A06: local window at sink0
8 A07 sink1_local12_relative 1 12 relativistic node1 A06-vs-A07: sink at local12
9 A08 sink3_local12_relative 3 12 relativistic node1 A06-vs-A07-vs-A08: sink size at local12
@@ -1,74 +0,0 @@
models:
student:
_target_: fastvideo.train.models.wan.WanCausalModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
method:
_target_: fastvideo.train.methods.fine_tuning.tfsft.TeacherForcingSFTMethod
chunk_size: 3
training:
distributed:
num_gpus: 4
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 4
data:
data_path: /mnt/lustre/vlm-k1kong/datasets/HD-Mixkit-Finetune-Wan/combined_parquet_dataset
dataloader_num_workers: 0
train_batch_size: 1
training_cfg_rate: 0.0
seed: 1000
num_latent_t: 21
num_height: 480
num_width: 832
num_frames: 81
optimizer:
learning_rate: 2.0e-6
betas: [0.0, 0.999]
weight_decay: 0.01
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: __MAX_TRAIN_STEPS__
gradient_accumulation_steps: 1
checkpoint:
output_dir: __CHECKPOINT_DIR__
training_state_checkpointing_steps: __CHECKPOINT_STEPS__
checkpoints_total_limit: 2
tracker:
trackers: [wandb]
project_name: __PROJECT_NAME__
run_name: __WANDB_RUN_NAME__
model:
enable_gradient_checkpointing_type: full
callbacks:
grad_clip:
max_grad_norm: 1.0
validation:
pipeline_target: fastvideo.pipelines.basic.wan.wan_causal_pipeline.WanCausalPipeline
dataset_file: __VALIDATION_DATASET_FILE__
every_steps: __VALIDATION_EVERY_STEPS__
sampling_steps: [__VALIDATION_SAMPLING_STEPS__]
guidance_scale: 6.0
num_frames: 249
output_dir: __VALIDATION_DIR__
offload_training_state: true
unload_pipeline_after_validation: true
pipeline:
flow_shift: 5
dit_config:
local_attn_size: __LOCAL_ATTN_SIZE__
sink_size: __SINK_SIZE__
rope_cache_policy: __ROPE_CACHE_POLICY__
causal_train_attention: triton
@@ -1,7 +0,0 @@
{
"data": [
{
"caption": "A lone rider guides a horse across an open field at sunset, with steady motion and a slowly changing background."
}
]
}
@@ -1,43 +0,0 @@
# WanTrack causal synth stage-2 run
This directory archives the exact configuration and launch artifacts used by
the completed WanTrack causal run on 2026-07-23/24.
- Run root:
`/mnt/lustre/vlm-k1kong/experiments/wantrack_causal/20260723_070728/wantrack-ckpt600-block3-local6-sink1-relative-gbs32`
- Initial checkpoint:
`/mnt/lustre/vlm-k1kong/models/wantrack-synth-stage2-ckpt600-bias`
- Dataset:
`/mnt/lustre/vlm-s4duan/data/combined_synth_parquets`
- Validation dataset:
`/mnt/lustre/vlm-s4duan/val_examples_mixed/combined_parquet_dataset`
- Topology: 4 rack-3 nodes, 4 GPUs per node, micro-batch 2, global batch 32
- Attention: 3 latent frames per block, sink 1, local window 6,
relativistic RoPE
- Schedule: TF 3000, CD 2000, SF 1000 optimizer steps
- Validation: every 250 steps; TF uses 30 denoising steps while CD and SF use
4 denoising steps
- Checkpoints: every 500 optimizer steps
The TF job resumed from checkpoint 2000 after its validation configuration was
corrected to multi-step sampling. The YAML files in this directory are the
final files used by the run, including their absolute cluster paths.
`cluster/run_pipeline_node.sh` launches the resumable TF -> CD -> SF pipeline
and exports TF `student`, CD `ema`, and SF `student_ema`. It requires
`WANDB_API_KEY` to be injected through the process environment. The gallery
upload script similarly requires an in-memory `HF_TOKEN`; no credentials are
stored here.
The original Kubernetes manifest retains the historical workload label
`wantrack-causal-framewise-gbs32`. That label is stale metadata: the training
authority is the three YAML files, all of which use
`num_frames_per_block: 3`.
The SF EMA export intentionally contains only the trainable checkpoint role.
For standalone inference, four frozen track-encoder parameters were restored
from the initialization checkpoint. See `receipts/sf-full-export.json` for the
full bundle provenance and hashes.
The `gallery/` scripts produced 64 generated examples plus their 64 matching
ground-truth videos from the full SF EMA bundle.
@@ -1,95 +0,0 @@
models:
student:
_target_: fastvideo.train.models.wantrack.WanTrackCausalModel
init_from: /mnt/lustre/vlm-k1kong/models/wantrack-synth-stage2-ckpt600-bias
transformer_override_safetensor: /mnt/lustre/vlm-k1kong/experiments/wantrack_causal/20260723_070728/wantrack-ckpt600-block3-local6-sink1-relative-gbs32/export/tf/transformer/model.safetensors
trainable: true
num_frames_per_block: 3
freeze_track_encoder: true
track_augmentation:
enabled: true
sparse_object_sampling: true
extra_points: 20
extra_point_sampling: random
track_dropout_probability: 0.5
temporal_mask_probability: 0.2
temporal_mask_chunk_size: 8
motion_dropout_probability: 0.3
text_dropout_probability: 0.0
teacher:
_target_: fastvideo.train.models.wantrack.WanTrackCausalModel
init_from: /mnt/lustre/vlm-k1kong/models/wantrack-synth-stage2-ckpt600-bias
transformer_override_safetensor: /mnt/lustre/vlm-k1kong/experiments/wantrack_causal/20260723_070728/wantrack-ckpt600-block3-local6-sink1-relative-gbs32/export/tf/transformer/model.safetensors
trainable: false
num_frames_per_block: 3
freeze_track_encoder: true
ema:
_target_: fastvideo.train.models.wantrack.WanTrackCausalModel
init_from: /mnt/lustre/vlm-k1kong/models/wantrack-synth-stage2-ckpt600-bias
transformer_override_safetensor: /mnt/lustre/vlm-k1kong/experiments/wantrack_causal/20260723_070728/wantrack-ckpt600-block3-local6-sink1-relative-gbs32/export/tf/transformer/model.safetensors
trainable: false
num_frames_per_block: 3
freeze_track_encoder: true
method:
_target_: fastvideo.train.methods.consistency_model.causal_cd.CausalConsistencyDistillationMethod
discrete_cd_N: 48
guidance_scale: 3.0
ema_decay: 0.99
ema_start_step: 200
training:
distributed:
num_gpus: 16
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 4
hsdp_shard_dim: 4
data:
data_path: /mnt/lustre/vlm-s4duan/data/combined_synth_parquets
preprocessed_data_type: i2v_track
dataloader_num_workers: 4
train_batch_size: 2
training_cfg_rate: 0.0
seed: 1000
num_latent_t: 31
num_height: 480
num_width: 832
num_frames: 121
optimizer:
learning_rate: 2.0e-6
betas: [0.0, 0.999]
weight_decay: 0.01
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 2000
gradient_accumulation_steps: 1
checkpoint:
output_dir: /mnt/lustre/vlm-k1kong/experiments/wantrack_causal/20260723_070728/wantrack-ckpt600-block3-local6-sink1-relative-gbs32/cd/checkpoints
training_state_checkpointing_steps: 500
checkpoints_total_limit: 4
tracker:
trackers: [wandb]
entity: kaiqin_kong_ucsd
project_name: wantrack_causal_synth_stage2
run_name: wantrack_ckpt600_block3_sink1_local6_relative_cd2k_gbs32
model:
enable_gradient_checkpointing_type: full
callbacks:
grad_clip:
max_grad_norm: 1.0
wantrack_validation:
_target_: fastvideo.train.callbacks.track_validation.TrackValidationCallback
every_steps: 250
val_data_path: /mnt/lustre/vlm-s4duan/val_examples_mixed/combined_parquet_dataset
num_val_samples: 4
num_inference_steps: 4
output_dir: /mnt/lustre/vlm-k1kong/experiments/wantrack_causal/20260723_070728/wantrack-ckpt600-block3-local6-sink1-relative-gbs32/cd/validation
seed: 1000
validate_at_start: false
pipeline:
flow_shift: 6
dit_config:
local_attn_size: 6
sink_size: 1
rope_cache_policy: relativistic
causal_train_attention: triton
@@ -1,106 +0,0 @@
models:
student:
_target_: fastvideo.train.models.wantrack.WanTrackCausalModel
init_from: /mnt/lustre/vlm-k1kong/models/wantrack-synth-stage2-ckpt600-bias
transformer_override_safetensor: /mnt/lustre/vlm-k1kong/experiments/wantrack_causal/20260723_070728/wantrack-ckpt600-block3-local6-sink1-relative-gbs32/export/cd/transformer/model.safetensors
trainable: true
num_frames_per_block: 3
freeze_track_encoder: true
track_augmentation:
enabled: true
sparse_object_sampling: true
extra_points: 20
extra_point_sampling: random
track_dropout_probability: 0.5
temporal_mask_probability: 0.2
temporal_mask_chunk_size: 8
motion_dropout_probability: 0.3
text_dropout_probability: 0.0
teacher:
_target_: fastvideo.train.models.wantrack.WanTrackModel
init_from: /mnt/lustre/vlm-k1kong/models/wantrack-synth-stage2-ckpt600-bias
trainable: false
freeze_track_encoder: true
disable_custom_init_weights: true
critic:
_target_: fastvideo.train.models.wantrack.WanTrackModel
init_from: /mnt/lustre/vlm-k1kong/models/wantrack-synth-stage2-ckpt600-bias
trainable: true
freeze_track_encoder: true
disable_custom_init_weights: true
method:
_target_: fastvideo.train.methods.distribution_matching.self_forcing.SelfForcingMethod
rollout_mode: simulate
generator_update_interval: 5
real_score_guidance_scale: 4.0
dmd_denoising_steps: [1000, 750, 500, 250]
warp_denoising_step: true
chunk_size: 3
student_sample_type: sde
context_noise: 0.0
same_step_across_blocks: true
enable_gradient_in_rollout: true
start_gradient_frame: 0
fake_score_learning_rate: 4.0e-7
fake_score_betas: [0.0, 0.999]
fake_score_lr_scheduler: constant
training:
distributed:
num_gpus: 16
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 4
hsdp_shard_dim: 4
data:
data_path: /mnt/lustre/vlm-s4duan/data/combined_synth_parquets
preprocessed_data_type: i2v_track
dataloader_num_workers: 4
train_batch_size: 2
training_cfg_rate: 0.0
seed: 1000
num_latent_t: 31
num_height: 480
num_width: 832
num_frames: 121
optimizer:
learning_rate: 2.0e-6
betas: [0.0, 0.999]
weight_decay: 0.01
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 1000
gradient_accumulation_steps: 1
checkpoint:
output_dir: /mnt/lustre/vlm-k1kong/experiments/wantrack_causal/20260723_070728/wantrack-ckpt600-block3-local6-sink1-relative-gbs32/sf/checkpoints
training_state_checkpointing_steps: 500
checkpoints_total_limit: 2
tracker:
trackers: [wandb]
entity: kaiqin_kong_ucsd
project_name: wantrack_causal_synth_stage2
run_name: wantrack_ckpt600_block3_sink1_local6_relative_sf1k_gbs32
model:
enable_gradient_checkpointing_type: full
callbacks:
ema:
decay: 0.99
start_iter: 200
grad_clip:
max_grad_norm: 1.0
wantrack_validation:
_target_: fastvideo.train.callbacks.track_validation.TrackValidationCallback
every_steps: 250
val_data_path: /mnt/lustre/vlm-s4duan/val_examples_mixed/combined_parquet_dataset
num_val_samples: 4
num_inference_steps: 4
output_dir: /mnt/lustre/vlm-k1kong/experiments/wantrack_causal/20260723_070728/wantrack-ckpt600-block3-local6-sink1-relative-gbs32/sf/validation
seed: 1000
validate_at_start: false
pipeline:
flow_shift: 6
dit_config:
local_attn_size: 6
sink_size: 1
rope_cache_policy: relativistic
causal_train_attention: triton
@@ -1,79 +0,0 @@
models:
student:
_target_: fastvideo.train.models.wantrack.WanTrackCausalModel
init_from: /mnt/lustre/vlm-k1kong/models/wantrack-synth-stage2-ckpt600-bias
trainable: true
num_frames_per_block: 3
freeze_track_encoder: true
track_augmentation:
enabled: true
sparse_object_sampling: true
extra_points: 20
extra_point_sampling: random
track_dropout_probability: 0.5
temporal_mask_probability: 0.2
temporal_mask_chunk_size: 8
motion_dropout_probability: 0.3
text_dropout_probability: 0.0
method:
_target_: fastvideo.train.methods.fine_tuning.tfsft.TeacherForcingSFTMethod
chunk_size: 3
training:
distributed:
num_gpus: 16
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 4
hsdp_shard_dim: 4
data:
data_path: /mnt/lustre/vlm-s4duan/data/combined_synth_parquets
preprocessed_data_type: i2v_track
dataloader_num_workers: 4
train_batch_size: 2
training_cfg_rate: 0.0
seed: 1000
num_latent_t: 31
num_height: 480
num_width: 832
num_frames: 121
optimizer:
learning_rate: 2.0e-6
betas: [0.0, 0.999]
weight_decay: 0.01
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 3000
gradient_accumulation_steps: 1
checkpoint:
output_dir: /mnt/lustre/vlm-k1kong/experiments/wantrack_causal/20260723_070728/wantrack-ckpt600-block3-local6-sink1-relative-gbs32/tf/checkpoints
training_state_checkpointing_steps: 500
checkpoints_total_limit: 6
tracker:
trackers: [wandb]
entity: kaiqin_kong_ucsd
project_name: wantrack_causal_synth_stage2
run_name: wantrack_ckpt600_block3_sink1_local6_relative_tf3k_gbs32_resume2000_multistep
model:
enable_gradient_checkpointing_type: full
callbacks:
grad_clip:
max_grad_norm: 1.0
wantrack_validation:
_target_: fastvideo.train.callbacks.track_validation.TrackValidationCallback
every_steps: 250
val_data_path: /mnt/lustre/vlm-s4duan/val_examples_mixed/combined_parquet_dataset
num_val_samples: 4
num_inference_steps: 30
guidance_scale: 3.0
motion_guidance_scale: 1.5
output_dir: /mnt/lustre/vlm-k1kong/experiments/wantrack_causal/20260723_070728/wantrack-ckpt600-block3-local6-sink1-relative-gbs32/tf/validation
seed: 1000
validate_at_start: false
pipeline:
flow_shift: 6
dit_config:
local_attn_size: 6
sink_size: 1
rope_cache_policy: relativistic
causal_train_attention: triton
-3
View File
@@ -25,7 +25,6 @@ models:
disable_custom_init_weights: false # default: false
flow_shift: 3.0 # default: 3.0
enable_gradient_checkpointing_type: null # default: null (falls back to training.model)
attention_backend: null # default: null (global/default); role-local when set
teacher:
_target_: fastvideo.train.models.wan.WanModel
@@ -53,8 +52,6 @@ method:
rollout_mode: simulate # required: "simulate" or "data_latent"
generator_update_interval: 5 # default: 1
dmd_denoising_steps: [1000, 750, 500, 250] # SDE timestep schedule
min_timestep_ratio: 0.0 # score-model timestep lower bound
max_timestep_ratio: 1.0 # score-model timestep upper bound
# Critic optimizer (all required — no fallback)
fake_score_learning_rate: 8.0e-6
@@ -1,82 +0,0 @@
# WanTrack bidirectional I2V fine-tuning.
#
# Build models/wantrack-control-init with
# scripts/checkpoint_conversion/wan_fun_control_to_trackwan.py, then preprocess
# the dataset with --preprocess_task i2v_track.
models:
student:
_target_: fastvideo.train.models.wantrack.WanTrackModel
init_from: models/wantrack-control-init
trainable: true
track_augmentation:
enabled: true
min_points: 1000
max_points: 2500
temporal_mask_probability: 0.2
temporal_mask_chunk_size: 8
motion_dropout_probability: 0.0
text_dropout_probability: 0.0
method:
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
training:
distributed:
num_gpus: 8
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 8
data:
data_path: data/wantrack/preprocessed/combined_parquet_dataset
preprocessed_data_type: i2v_track
dataloader_num_workers: 4
train_batch_size: 1
training_cfg_rate: 0.0
seed: 1000
num_latent_t: 21
num_height: 480
num_width: 832
num_frames: 81
optimizer:
learning_rate: 2.0e-6
betas: [0.9, 0.999]
weight_decay: 0.01
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 4000
gradient_accumulation_steps: 1
checkpoint:
output_dir: outputs/wantrack_bidirectional_i2v
training_state_checkpointing_steps: 1000
checkpoints_total_limit: 3
resume_from_checkpoint: null
tracker:
project_name: fastvideo
run_name: wantrack_bidirectional_i2v
model:
enable_gradient_checkpointing_type: full
callbacks:
grad_clip:
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
max_grad_norm: 1.0
track_validation:
_target_: fastvideo.train.callbacks.track_validation.TrackValidationCallback
every_steps: 250
num_val_samples: 2
num_inference_steps: 30
guidance_scale: 3.0
motion_guidance_scale: 1.5
validate_at_start: false
pipeline:
flow_shift: 5
@@ -1,88 +0,0 @@
# WanTrack causal I2V teacher-forcing fine-tuning.
#
# Build models/wantrack-control-causal-init with the checkpoint converter's
# --causal flag. Track IDs and the full encoded track map are shared by every
# latent chunk.
models:
student:
_target_: fastvideo.train.models.wantrack.WanTrackCausalModel
init_from: models/wantrack-control-causal-init
trainable: true
track_augmentation:
enabled: true
min_points: 1000
max_points: 2500
temporal_mask_probability: 0.2
temporal_mask_chunk_size: 8
motion_dropout_probability: 0.0
text_dropout_probability: 0.0
method:
_target_: fastvideo.train.methods.fine_tuning.tfsft.TeacherForcingSFTMethod
chunk_size: 3
training:
distributed:
num_gpus: 8
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 8
data:
data_path: data/wantrack/preprocessed/combined_parquet_dataset
preprocessed_data_type: i2v_track
dataloader_num_workers: 4
train_batch_size: 1
training_cfg_rate: 0.0
seed: 1000
num_latent_t: 21
num_height: 480
num_width: 832
num_frames: 81
optimizer:
learning_rate: 2.0e-6
betas: [0.9, 0.999]
weight_decay: 0.01
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 4000
gradient_accumulation_steps: 1
checkpoint:
output_dir: outputs/wantrack_causal_i2v
training_state_checkpointing_steps: 1000
checkpoints_total_limit: 3
resume_from_checkpoint: null
tracker:
project_name: fastvideo
run_name: wantrack_causal_i2v
model:
enable_gradient_checkpointing_type: full
callbacks:
grad_clip:
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
max_grad_norm: 1.0
track_validation:
_target_: fastvideo.train.callbacks.track_validation.TrackValidationCallback
every_steps: 250
num_val_samples: 2
num_inference_steps: 30
guidance_scale: 3.0
motion_guidance_scale: 1.5
validate_at_start: false
pipeline:
flow_shift: 5
dit_config:
local_attn_size: 6
sink_size: 1
rope_cache_policy: relativistic
causal_train_attention: triton
@@ -1,95 +0,0 @@
# LTX-2.3 T2V overfitting test config.
#
# Overfits on a single short video (480x832, 81 frames @ 24fps) to
# verify the LTX-2 training plugin works end-to-end on an LTX-2.3
# checkpoint (gated attention, cross-attention AdaLN, 4096-d
# post-connector text embeddings, no in-DiT caption projection).
# Uses the distilled checkpoint (validation is 8 sampling steps).
#
# Preprocess data first (writes data/ltx2_3_overfit_preprocessed):
# CUDA_VISIBLE_DEVICES=0 \
# LTX2_OVERFIT_MODEL=FastVideo/LTX-2.3-Distilled-Diffusers \
# LTX2_OVERFIT_OUTPUT_DIR=data/ltx2_3_overfit_preprocessed \
# python fastvideo/pipelines/preprocess/preprocess_ltx2_overfit.py
#
# Run:
# NUM_GPUS=4 bash examples/train/run.sh examples/train/configs/overfit_ltx2_3_t2v.yaml
models:
student:
_target_: fastvideo.train.models.ltx2.LTX2Model
init_from: FastVideo/LTX-2.3-Distilled-Diffusers
trainable: true
enable_gradient_checkpointing_type: full
method:
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
training:
distributed:
num_gpus: 4
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 4
data:
data_path: data/ltx2_3_overfit_preprocessed
dataloader_num_workers: 0
train_batch_size: 1
# LTX2Model requires 0.0: CFG dropout would zero post-connector
# embeddings, which is not the model's unconditional input.
training_cfg_rate: 0.0
seed: 42
num_latent_t: 11 # (81 - 1) / 8 + 1
num_height: 480
num_width: 832
num_frames: 81
optimizer:
learning_rate: 5.0e-5
betas: [0.9, 0.999]
weight_decay: 0.0
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 300
gradient_accumulation_steps: 1
checkpoint:
output_dir: outputs/ltx2_3_overfit
# A full training-state checkpoint is ~150GB for the 13B trainable
# video branch; disable saves for the overfit smoke run.
training_state_checkpointing_steps: 0
checkpoints_total_limit: 1
tracker:
trackers: [wandb]
project_name: fastvideo_ltx2
run_name: ltx2_3_overfit
model:
# LTX2Model.predict_noise converts the DiT's x0 output to velocity,
# so the default noise-minus-clean target reproduces the official
# unweighted masked-MSE (mask is all-ones for plain T2V).
precondition_outputs: false
enable_gradient_checkpointing_type: full
callbacks:
grad_clip:
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
max_grad_norm: 1.0
validation:
_target_: fastvideo.train.callbacks.validation.ValidationCallback
pipeline_target: fastvideo.pipelines.basic.ltx2.ltx2_pipeline.LTX2Pipeline
dataset_file: data/ltx2_3_overfit_preprocessed/validation_prompts.json
every_steps: 50
sampling_steps: [8]
guidance_scale: 1.0
num_frames: 81
# Required so the LTX2T2VConfig pipeline config is resolved from
# init_from (without a `pipeline:` key the loader falls back to a
# generic PipelineConfig and the LTX-2 DiT cannot be constructed).
pipeline: {}
@@ -1,90 +0,0 @@
# LTX-2 T2V overfitting test config.
#
# Overfits on a single short video (480x832, 81 frames @ 24fps) to
# verify the LTX-2 training plugin works end-to-end. Uses the
# distilled checkpoint (validation is 8 sampling steps, single pass).
#
# Preprocess data first (writes data/ltx2_overfit_preprocessed):
# CUDA_VISIBLE_DEVICES=0 python fastvideo/pipelines/preprocess/preprocess_ltx2_overfit.py
#
# Run:
# NUM_GPUS=4 bash examples/train/run.sh examples/train/configs/overfit_ltx2_t2v.yaml
models:
student:
_target_: fastvideo.train.models.ltx2.LTX2Model
init_from: FastVideo/LTX2-Distilled-Diffusers
trainable: true
enable_gradient_checkpointing_type: full
method:
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
training:
distributed:
num_gpus: 4
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 4
data:
data_path: data/ltx2_overfit_preprocessed
dataloader_num_workers: 0
train_batch_size: 1
# LTX2Model requires 0.0: CFG dropout would zero post-connector
# embeddings, which is not the model's unconditional input.
training_cfg_rate: 0.0
seed: 42
num_latent_t: 11 # (81 - 1) / 8 + 1
num_height: 480
num_width: 832
num_frames: 81
optimizer:
learning_rate: 5.0e-5
betas: [0.9, 0.999]
weight_decay: 0.0
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 300
gradient_accumulation_steps: 1
checkpoint:
output_dir: outputs/ltx2_overfit
# A full training-state checkpoint is ~150GB for the 13B trainable
# video branch; disable saves for the overfit smoke run.
training_state_checkpointing_steps: 0
checkpoints_total_limit: 1
tracker:
trackers: [wandb]
project_name: fastvideo_ltx2
run_name: ltx2_overfit
model:
# LTX2Model.predict_noise converts the DiT's x0 output to velocity,
# so the default noise-minus-clean target reproduces the official
# unweighted masked-MSE (mask is all-ones for plain T2V).
precondition_outputs: false
enable_gradient_checkpointing_type: full
callbacks:
grad_clip:
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
max_grad_norm: 1.0
validation:
_target_: fastvideo.train.callbacks.validation.ValidationCallback
pipeline_target: fastvideo.pipelines.basic.ltx2.ltx2_pipeline.LTX2Pipeline
dataset_file: data/ltx2_overfit_preprocessed/validation_prompts.json
every_steps: 50
sampling_steps: [8]
guidance_scale: 1.0
num_frames: 81
# Required so the LTX2T2VConfig pipeline config is resolved from
# init_from (without a `pipeline:` key the loader falls back to a
# generic PipelineConfig and the LTX-2 DiT cannot be constructed).
pipeline: {}
+2 -5
View File
@@ -5,12 +5,9 @@ all configs, scripts, and data needed to run a complete workflow.
```
scenario/
├── ode_init_self_forcing_wan_causal/ # KD → export → Self-Forcing
└── qad_wan2_1_mixkit/ # Attn-QAT SFT → export → DMD2
└── ode_init_self_forcing_wan_causal/ # KD → export → Self-Forcing
```
See the `usage.md` inside each scenario for step-by-step instructions. The QAD
workflow is also documented in the website-visible
[Attn-QAT training guide](https://haoailab.com/FastVideo/training/attn_qat/).
See the `usage.md` inside each scenario for step-by-step instructions.
For single-step configs, see `examples/train/configs/`.
@@ -1,15 +0,0 @@
#!/usr/bin/env bash
# Export a modular-trainer DCP checkpoint for stage-2 initialization.
set -euo pipefail
SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
REPO_ROOT="$(cd "${SCRIPT_DIR}/../../../.." && pwd)"
cd "${REPO_ROOT}"
CHECKPOINT_DIR=${1:-checkpoints/wan_t2v_qat_finetune/checkpoint-4000}
OUTPUT_DIR=${2:-checkpoints/wan_t2v_qat_finetune/diffusers}
python -m fastvideo.train.entrypoint.dcp_to_diffusers \
--role student \
--checkpoint "${CHECKPOINT_DIR}" \
--output-dir "${OUTPUT_DIR}"
@@ -1,20 +0,0 @@
#!/usr/bin/env bash
# Run the modular Attn-QAT finetune recipe.
set -euo pipefail
SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
REPO_ROOT="$(cd "${SCRIPT_DIR}/../../../.." && pwd)"
cd "${REPO_ROOT}"
DATA_DIR=${1:-data/HD-Mixkit-Finetune-Wan/combined_parquet_dataset}
NUM_GPUS=${NUM_GPUS:-4}
export NUM_GPUS
export FASTVIDEO_ATTN_QAT_FWD_EXACT_M=${FASTVIDEO_ATTN_QAT_FWD_EXACT_M:-0}
bash "${REPO_ROOT}/examples/train/run.sh" \
"${SCRIPT_DIR}/stage1_attn_qat_finetune.yaml" \
--training.data.data_path "${DATA_DIR}" \
--training.distributed.num_gpus "${NUM_GPUS}" \
--training.distributed.sp_size "${NUM_GPUS}" \
--training.distributed.hsdp_replicate_dim 1 \
--training.distributed.hsdp_shard_dim "${NUM_GPUS}"
@@ -1,27 +0,0 @@
#!/usr/bin/env bash
# Run modular DMD2 with Attn-QAT on the student only.
set -euo pipefail
SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
REPO_ROOT="$(cd "${SCRIPT_DIR}/../../../.." && pwd)"
cd "${REPO_ROOT}"
DATA_DIR=${1:-data/HD-Mixkit-Finetune-Wan/combined_parquet_dataset}
INIT_WEIGHTS=${2:-checkpoints/wan_t2v_qat_finetune/diffusers/transformer/model.safetensors}
NUM_GPUS=${NUM_GPUS:-4}
export NUM_GPUS
if [[ ! -f "${INIT_WEIGHTS}" ]]; then
echo "Missing exported stage-1 weights: ${INIT_WEIGHTS}" >&2
echo "Run export_stage1.sh before stage 2." >&2
exit 1
fi
bash "${REPO_ROOT}/examples/train/run.sh" \
"${SCRIPT_DIR}/stage2_attn_qat_dmd.yaml" \
--models.student.transformer_override_safetensor "${INIT_WEIGHTS}" \
--training.data.data_path "${DATA_DIR}" \
--training.distributed.num_gpus "${NUM_GPUS}" \
--training.distributed.sp_size 1 \
--training.distributed.hsdp_replicate_dim "${NUM_GPUS}" \
--training.distributed.hsdp_shard_dim 1
@@ -1,72 +0,0 @@
# QAD stage 1: Attn-QAT finetune of Wan2.1-T2V-1.3B on MixKit.
models:
student:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
# Role-local: does not change the backend used by other models loaded in
# this process.
attention_backend: ATTN_QAT_TRAIN
method:
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
training:
distributed:
num_gpus: 4
sp_size: 4
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 4
data:
data_path: data/HD-Mixkit-Finetune-Wan/combined_parquet_dataset
dataloader_num_workers: 1
train_batch_size: 1
training_cfg_rate: 0.1
seed: 1000
num_latent_t: 20
num_height: 480
num_width: 832
num_frames: 77
optimizer:
learning_rate: 1.0e-6
betas: [0.9, 0.999]
weight_decay: 0.01
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 4000
gradient_accumulation_steps: 1
checkpoint:
output_dir: checkpoints/wan_t2v_qat_finetune
training_state_checkpointing_steps: 500
checkpoints_total_limit: 0
tracker:
project_name: wan_t2v_qat_finetune
run_name: wan_t2v_qat_finetune
model:
enable_gradient_checkpointing_type: full
dit_precision: fp32
callbacks:
grad_clip:
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
max_grad_norm: 1.0
validation:
_target_: fastvideo.train.callbacks.validation.ValidationCallback
pipeline_target: fastvideo.pipelines.basic.wan.wan_pipeline.WanPipeline
dataset_file: examples/training/finetune/wan_t2v_1.3B/crush_smol/validation.json
every_steps: 50
sampling_steps: [50]
guidance_scale: 5.0
pipeline:
flow_shift: 1
@@ -1,101 +0,0 @@
# QAD stage 2: distill the Attn-QAT student to three sampling steps.
#
# Export the stage-1 DCP checkpoint first, then point
# models.student.transformer_override_safetensor at the exported weight file.
models:
student:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
transformer_override_safetensor: checkpoints/wan_t2v_qat_finetune/diffusers/transformer/model.safetensors
trainable: true
attention_backend: ATTN_QAT_TRAIN
teacher:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: false
disable_custom_init_weights: true
attention_backend: FLASH_ATTN
critic:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
disable_custom_init_weights: true
attention_backend: FLASH_ATTN
method:
_target_: fastvideo.train.methods.distribution_matching.dmd2.DMD2Method
rollout_mode: data_latent
generator_update_interval: 5
dmd_denoising_steps: [1000, 757, 522]
min_timestep_ratio: 0.02
max_timestep_ratio: 0.98
# The modular DMD method uses standard CFG:
# uncond + scale * (cond - uncond). This is equivalent to the legacy
# recipe's cond + 2.0 * (cond - uncond).
real_score_guidance_scale: 3.0
# The legacy recipe inherited these values from its global optimizer.
fake_score_learning_rate: 2.0e-6
fake_score_betas: [0.9, 0.999]
fake_score_lr_scheduler: constant
training:
distributed:
num_gpus: 4
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 4
hsdp_shard_dim: 1
data:
data_path: data/HD-Mixkit-Finetune-Wan/combined_parquet_dataset
dataloader_num_workers: 4
train_batch_size: 1
training_cfg_rate: 0.0
seed: 1000
num_latent_t: 20
num_height: 480
num_width: 832
num_frames: 77
optimizer:
learning_rate: 2.0e-6
betas: [0.9, 0.999]
weight_decay: 0.01
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 2000
gradient_accumulation_steps: 1
checkpoint:
output_dir: checkpoints/wan_t2v_distill_dmd_qat
training_state_checkpointing_steps: 500
checkpoints_total_limit: 3
tracker:
project_name: wan_t2v_distill_dmd_qat
run_name: wan_t2v_distill_dmd_qat
model:
enable_gradient_checkpointing_type: full
dit_precision: fp32
callbacks:
grad_clip:
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
max_grad_norm: 1.0
validation:
_target_: fastvideo.train.callbacks.validation.ValidationCallback
pipeline_target: fastvideo.pipelines.basic.wan.wan_dmd_pipeline.WanDMDPipeline
dataset_file: examples/training/finetune/wan_t2v_1.3B/crush_smol/validation.json
every_steps: 200
sampling_steps: [3]
sampling_timesteps: [1000, 757, 522]
guidance_scale: 6.0
pipeline:
flow_shift: 3
@@ -1,26 +0,0 @@
# QAD Wan2.1 MixKit Attn-QAT
This scenario runs entirely on the modular `fastvideo/train` stack. The
student's attention backend is configured per role, so DMD2 can keep the
teacher and critic on Flash Attention while the student uses the fake-quantized
Attn-QAT kernel.
From the repository root:
```bash
# 1. Download the preprocessed MixKit data.
bash examples/training/finetune/wan_t2v_1.3B/mixkit/download_mixkit_data.sh
# 2. Stage 1: Attn-QAT supervised finetune.
NUM_GPUS=4 bash examples/train/scenario/qad_wan2_1_mixkit/run_stage1.sh
# 3. Export the stage-1 DCP checkpoint to a Diffusers weight file.
bash examples/train/scenario/qad_wan2_1_mixkit/export_stage1.sh
# 4. Stage 2: three-step Attn-QAT DMD2 distillation.
NUM_GPUS=4 bash examples/train/scenario/qad_wan2_1_mixkit/run_stage2.sh
```
The two YAML configs are also directly runnable through `examples/train/run.sh`.
The wrapper scripts only provide dataset/checkpoint paths and derive distributed
dimensions from `NUM_GPUS`.
@@ -52,7 +52,7 @@ for the full parameter reference.
## Train (QAT finetune)
With the data in place, run the quantization-aware finetune. The 4-bit attention
path is **config-driven** and selected by an environment variable:
path is **config-driven** — selected purely by an env var, no monkey-patching:
```bash
bash examples/training/finetune/wan_t2v_1.3B/mixkit/finetune_qat.sh
@@ -60,19 +60,10 @@ bash examples/training/finetune/wan_t2v_1.3B/mixkit/finetune_qat.sh
NUM_GPUS=4 bash .../mixkit/finetune_qat.sh data/HD-Mixkit-Finetune-Wan/combined_parquet_dataset/
```
`FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_TRAIN` keeps the fake-quantized Triton
forward and backward (straight-through estimator). Both kernels ship in
`fastvideo-kernel`: SM100 automatically uses the optimized Triton path for the
production non-causal, head-dimension-128 configuration, while SM120 GPUs such
as RTX 5090 join the quantized and STE P@V operations and use a shallower
backward pipeline for long sequences. Set
`FASTVIDEO_ATTN_QAT_SM120_JOIN_QAT_PV=0` to compare against the split P@V path.
The script defaults
`FASTVIDEO_ATTN_QAT_FWD_EXACT_M=0` for throughput; set it to `1` to reproduce
the previous forward softmax statistic and bitwise-compatible `dV`.
For the website-visible modular SFT-to-DMD2 workflow, see the
[Attn-QAT training guide](https://haoailab.com/FastVideo/training/attn_qat/).
`FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_TRAIN` routes attention through the
fake-quantized Triton kernel (straight-through estimator), so the DiT learns to
absorb FP4 attention error. This kernel is Triton, so it runs on both `sm_100`
(B200/GB200) and `sm_120` (RTX 5090).
## Train stage 2 (QAT DMD distillation to 3 steps)
@@ -80,8 +71,8 @@ Distill the QAT-finetuned generator down to **3 sampling steps**. Only the
generator is quantized (Attn-QAT); the teacher (`real_score`) and critic
(`fake_score`) stay full precision. This is enforced in the loader
(`component_loader.py`, via the `_loading_teacher_critic_model` flag), so the
same global `ATTN_QAT_TRAIN` env reaches **only** the generator, with no
per-model flags.
same global `ATTN_QAT_TRAIN` env reaches **only** the generator — no per-model
flags or monkey-patching.
```bash
# generator init = the stage-1 finetune checkpoint
@@ -2,23 +2,20 @@
# QAD recipe — quantization-aware finetune of Wan2.1-T2V-1.3B with fake-quant
# (Attn-QAT) attention.
#
# The 4-bit attention path is selected by env var:
# FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_TRAIN keeps the fake-quantized Triton
# forward and backward, so the DiT learns to absorb FP4 attention error instead
# of fighting it. SM100 selects the optimized kernels; RTX 5090 keeps the
# previous Triton implementation.
# The 4-bit attention path is selected purely by env var (config-driven, no
# monkey-patching): FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_TRAIN routes attention
# through the fake-quantized Triton kernel (straight-through estimator), so the
# DiT learns to absorb FP4 attention error instead of fighting it.
#
# Data: run download_mixkit_data.sh first (preprocessed Parquet).
#
# Verified end-to-end on Blackwell (GB200/sm_100): the ATTN_QAT_TRAIN backend is
# selected (not a fallback), forward+backward run, loss/grad are healthy, and
# validation generates videos.
# validation generates videos. The kernel is Triton so it runs on sm_100 and
# sm_120 alike (the FP4 inference kernel, by contrast, is sm_120-only).
set -euo pipefail
REPO_ROOT="$(cd "$(dirname "${BASH_SOURCE[0]}")/../../../../../.." && pwd)"
export PYTHONPATH="${REPO_ROOT}/fastvideo-kernel/python${PYTHONPATH:+:${PYTHONPATH}}"
export FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_TRAIN # <-- enables Attn-QAT training
export FASTVIDEO_ATTN_QAT_FWD_EXACT_M=${FASTVIDEO_ATTN_QAT_FWD_EXACT_M:-0}
export WANDB_MODE=${WANDB_MODE:-online}
export TOKENIZERS_PARALLELISM=false
@@ -26,8 +23,6 @@ MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_DIR=${1:-"data/HD-Mixkit-Finetune-Wan/combined_parquet_dataset/"}
VALIDATION_FILE="$(dirname "$0")/../crush_smol/validation.json"
NUM_GPUS=${NUM_GPUS:-4}
MAX_TRAIN_STEPS=${MAX_TRAIN_STEPS:-4000}
VALIDATION_SAMPLING_STEPS=${VALIDATION_SAMPLING_STEPS:-50}
torchrun --nnodes 1 --nproc_per_node "${NUM_GPUS}" \
fastvideo/training/wan_training_pipeline.py \
@@ -35,15 +30,15 @@ torchrun --nnodes 1 --nproc_per_node "${NUM_GPUS}" \
--hsdp_replicate_dim 1 --hsdp_shard_dim "${NUM_GPUS}" \
--model_path "${MODEL_PATH}" --pretrained_model_name_or_path "${MODEL_PATH}" \
--data_path "${DATA_DIR}" --dataloader_num_workers 1 \
--max_train_steps "${MAX_TRAIN_STEPS}" --train_batch_size 1 --train_sp_batch_size 1 \
--max_train_steps 2000 --train_batch_size 1 --train_sp_batch_size 1 \
--gradient_accumulation_steps 1 \
--num_latent_t 20 --num_height 480 --num_width 832 --num_frames 77 \
--enable_gradient_checkpointing_type full \
--log_validation --validation_dataset_file "${VALIDATION_FILE}" \
--validation_steps 50 --validation_sampling_steps "${VALIDATION_SAMPLING_STEPS}" --validation_guidance_scale 5.0 \
--learning_rate 1e-6 --mixed_precision bf16 --weight_decay 0.01 --max_grad_norm 1.0 \
--validation_steps 200 --validation_sampling_steps 50 --validation_guidance_scale 3.0 \
--learning_rate 5e-5 --mixed_precision bf16 --weight_decay 1e-4 --max_grad_norm 1.0 \
--weight_only_checkpointing_steps 500 --training_state_checkpointing_steps 500 \
--tracker_project_name wan_t2v_qat_finetune --output_dir checkpoints/wan_t2v_qat_finetune \
--inference_mode False --training_cfg_rate 0.1 --not_apply_cfg_solver \
--dit_precision fp32 --num_euler_timesteps 50 --ema_start_step 0 --flow_shift 1 \
--dit_precision fp32 --num_euler_timesteps 50 --ema_start_step 0 \
--multi_phased_distill_schedule "4000-1"
-25
View File
@@ -308,21 +308,6 @@ if(BUILD_CXX_KERNELS)
csrc/turbodiffusion/quant/quant.cu
)
# Blackwell block-causal + sink + sliding-window attention (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 / .cta_group instruction. The explicit gencode
# spelling below is required, and matches the 10.0a entry in TORCH_CUDA_ARCH_LIST.
set(ENABLE_BCS_SM100A OFF)
if(TORCH_CUDA_ARCH_LIST MATCHES "(^|[; ,])(10\\.0a|100a|sm_100a)([; ,]|$)")
set(ENABLE_BCS_SM100A ON)
endif()
if(ENABLE_BCS_SM100A)
message(STATUS "fastvideo-kernel: building block_causal_sink_sm100a (Blackwell)")
list(APPEND EXTENSION_SOURCES csrc/attention/block_causal_sink_sm100a.cu)
set_source_files_properties(csrc/attention/block_causal_sink_sm100a.cu PROPERTIES
COMPILE_OPTIONS "-gencode;arch=compute_100a,code=sm_100a")
endif()
# Conditionally add TK kernels
if(ENABLE_TK_KERNELS)
list(APPEND EXTENSION_SOURCES
@@ -351,9 +336,6 @@ if(BUILD_CXX_KERNELS)
if(ENABLE_TK_KERNELS)
list(APPEND COMPILE_DEFS TK_COMPILE_ST_ATTN TK_COMPILE_BLOCK_SPARSE)
endif()
if(ENABLE_BCS_SM100A)
list(APPEND COMPILE_DEFS TK_COMPILE_BLOCK_CAUSAL_SINK_SM100A)
endif()
target_compile_definitions(fastvideo_kernel_ops PRIVATE ${COMPILE_DEFS})
@@ -365,13 +347,6 @@ if(BUILD_CXX_KERNELS)
# (e.g., torch::autograd vtables) when loading the extension module.
target_link_libraries(fastvideo_kernel_ops PRIVATE ${TORCH_LIBRARIES})
# The Blackwell kernel builds its TMA descriptors with cuTensorMapEncodeTiled, a CUDA
# DRIVER API entry point -- it is not in libcudart, so the module fails to import with
# "undefined symbol: cuTensorMapEncodeTiled" unless libcuda is linked explicitly.
if(ENABLE_BCS_SM100A)
target_link_libraries(fastvideo_kernel_ops PRIVATE cuda)
endif()
# Also link against libtorch_python to satisfy Python-binding symbols
# (e.g., torch::PyWarningHandler) required by torch/extension.h.
execute_process(
-27
View File
@@ -108,33 +108,6 @@ out = moba_attn_varlen(q, k, v, cu_seqlens_q, cu_seqlens_k, ...)
## Benchmark
### Attn-QAT training
The default shape matches one sequence-parallel rank of the 4-GPU
Wan2.1-T2V-1.3B MixKit recipe (`B=1, H=3, L=31200, D=128`):
```bash
cd fastvideo-kernel
python benchmarks/benchmark_attn_qat_train.py
```
The benchmark reports both conventional attention FLOPs and the extra matrix
multiplications executed by the QAT straight-through path. Override
`--peak-tflops` when running on a GPU other than RTX 5090.
The QAT kernel is entirely Triton and routes by architecture at runtime. SM100
uses a large-tile forward and split 64x64 backward for the production
non-causal, head-dimension-128 configuration with a 16-aligned KV length. SM120
(including RTX 5090) keeps the previous tiling but joins the quantized and STE
P@V operations and uses a shallower backward pipeline for long sequences. Set
`FASTVIDEO_ATTN_QAT_SM120_JOIN_QAT_PV=0` to compare against the split P@V path.
Unsupported configurations retain the previous implementation. Set
`FASTVIDEO_ATTN_QAT_SM100_OPTIMIZED=0` to benchmark that previous path on SM100. Forward tuning is available through
`FASTVIDEO_ATTN_QAT_FWD_MODE=fast|balanced|reference`; exact reference-order
softmax statistics are controlled by `FASTVIDEO_ATTN_QAT_FWD_EXACT_M` and are
disabled by default for maximum throughput. Set it to `1` for reference-order
statistics and bitwise-compatible `dV`.
### VSA (block-sparse) TFLOPs
After building/installing `fastvideo-kernel`, run:
@@ -1,127 +0,0 @@
#!/usr/bin/env python3
"""Benchmark the Attn-QAT training kernel on a single GPU.
Defaults model one rank of the 4-GPU Wan2.1-T2V-1.3B MixKit recipe:
``B=1, H=12/4, L=20*30*52, D=128``.
"""
from __future__ import annotations
import argparse
import statistics
import time
from collections.abc import Callable
import torch
from fastvideo_kernel.triton_kernels.attn_qat_train import attention
RTX_5090_DENSE_BF16_TFLOPS = 209.5
def _qat_attention(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> torch.Tensor:
consumer_blackwell = torch.cuda.get_device_capability()[0] == 12
return attention(
q,
k,
v,
False,
q.shape[-1]**-0.5,
True, # use_qat_qkv_backward
False, # smooth_k
not consumer_blackwell, # warp_specialize
True, # IS_QAT
False, # two_level_quant_P
True, # fake_quant_P
True, # use_high_prec_o
False, # smooth_q
False, # use_global_sf_P
False, # use_global_sf_QKV
)
def _measure_ms(fn: Callable[[], object], warmup: int, repeat: int) -> tuple[float, float, float]:
for _ in range(warmup):
fn()
torch.cuda.synchronize()
samples = []
for _ in range(repeat):
start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
start.record()
fn()
end.record()
end.synchronize()
samples.append(start.elapsed_time(end))
return statistics.median(samples), min(samples), max(samples)
def _format_result(
label: str,
timing_ms: tuple[float, float, float],
algorithmic_flops: int,
executed_matmul_flops: int,
peak_tflops: float,
) -> str:
median_ms, min_ms, max_ms = timing_ms
algorithmic_tflops = algorithmic_flops / (median_ms * 1e9)
executed_tflops = executed_matmul_flops / (median_ms * 1e9)
return (
f"{label}: {median_ms:.3f} ms (min={min_ms:.3f}, max={max_ms:.3f}), "
f"algorithmic={algorithmic_tflops:.2f} TFLOPS/{100 * algorithmic_tflops / peak_tflops:.2f}% MFU, "
f"executed_matmul={executed_tflops:.2f} TFLOPS/{100 * executed_tflops / peak_tflops:.2f}% MFU"
)
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--batch-size", type=int, default=1)
parser.add_argument("--heads", type=int, default=3, help="Heads per SP rank; Wan 1.3B has 12 total.")
parser.add_argument("--query-length", type=int, default=31_200)
parser.add_argument("--kv-length", type=int, help="Defaults to --query-length.")
parser.add_argument("--head-dim", type=int, default=128)
parser.add_argument("--warmup", type=int, default=3)
parser.add_argument("--repeat", type=int, default=10)
parser.add_argument(
"--peak-tflops",
type=float,
default=RTX_5090_DENSE_BF16_TFLOPS,
help="Dense BF16 Tensor TFLOPS with FP32 accumulation; default is RTX 5090 boost-clock peak.",
)
args = parser.parse_args()
kv_length = args.kv_length or args.query_length
torch.manual_seed(0)
q_shape = (args.batch_size, args.heads, args.query_length, args.head_dim)
kv_shape = (args.batch_size, args.heads, kv_length, args.head_dim)
q = torch.randn(q_shape, device="cuda", dtype=torch.bfloat16, requires_grad=True)
k = torch.randn(kv_shape, device="cuda", dtype=torch.bfloat16, requires_grad=True)
v = torch.randn(kv_shape, device="cuda", dtype=torch.bfloat16, requires_grad=True)
grad_out = torch.randn_like(q)
compile_start = time.perf_counter()
output = _qat_attention(q, k, v)
torch.cuda.synchronize()
compile_seconds = time.perf_counter() - compile_start
forward_ms = _measure_ms(lambda: _qat_attention(q, k, v), args.warmup, args.repeat)
backward_ms = _measure_ms(
lambda: torch.autograd.grad(output, (q, k, v), grad_out, retain_graph=True),
args.warmup,
args.repeat,
)
base_flops = args.batch_size * args.heads * args.query_length * kv_length * args.head_dim
# Conventional attention FLOPs are 4*base forward and 10*base backward.
# QAT additionally computes the STE high-precision P@V path in forward and
# the quantized-P dV path in backward, for 6*base and 14*base matmul FLOPs.
print(f"device: {torch.cuda.get_device_name()}")
print(f"q: {q_shape}; k/v: {kv_shape}; compile+first-forward: {compile_seconds:.3f} s")
print(_format_result("forward", forward_ms, 4 * base_flops, 6 * base_flops, args.peak_tflops))
print(_format_result("backward", backward_ms, 10 * base_flops, 14 * base_flops, args.peak_tflops))
if __name__ == "__main__":
main()
File diff suppressed because it is too large Load Diff
@@ -1,235 +0,0 @@
// block_causal_sink_launch_sm100a.cuh -- callable entry point for the block-causal + sink +
// sliding-window FMHA kernel (sm_100a, forward only). No env vars, no allocation, no host-side
// reordering: everything comes from the caller, so a torch extension and our benchmark harness
// share it.
//
// LAYOUT. Q/K/V/O are [B, H, L, D] with head_dim CONTIGUOUS. Only head_dim contiguity is
// required -- the outer strides are read from the caller, so a torch tensor that has merely
// been permuted (no .contiguous()) is consumed as-is with no copy. V is read MN-major, so it
// needs no transpose either.
//
// SCALE. sm_scale is the caller's, matching FastVideo's qk_scale = sm_scale * LOG2E. Getting
// it wrong corrupts both O and lse and does so plausibly, so it is never derived here.
#pragma once
#include "block_causal_sink_kernel_sm100a.cuh"
struct BlockCausalSinkArgs {
const __nv_bfloat16* q = nullptr; // [B, H_q, L, D]
const __nv_bfloat16* k = nullptr; // [B, H_kv, L, D]
const __nv_bfloat16* v = nullptr; // [B, H_kv, L, D] (MN-major, NOT transposed)
const __nv_bfloat16* q_sink = nullptr; // [B, H_q, L, D], required iff has_delta
__nv_bfloat16* o = nullptr; // [B, H_q, L, D]
float* lse = nullptr; // [B*H_q, L] fp32; nullptr skips the store
int batch = 0, seqlen = 0, num_q_heads = 0, num_kv_heads = 0, head_dim = 128;
int tokens_per_block = 0; // num_frame_per_block * frame_seqlen
int sink_tokens = 0; // sink_size * frame_seqlen
int rolling_window_tokens = 0; // local_attn_size * frame_seqlen
float sm_scale = 0.f; // 0 -> 1/sqrt(head_dim)
bool has_delta = false; // relativistic sink RoPE correction
};
// Returns cudaErrorInvalidValue for an unsupported configuration rather than computing
// silently-wrong results.
inline cudaError_t block_causal_sink_supported(const BlockCausalSinkArgs& a) {
if (a.head_dim != HEAD_DIM) return cudaErrorInvalidValue;
if (a.num_q_heads % a.num_kv_heads) return cudaErrorInvalidValue;
if (a.tokens_per_block <= 0) return cudaErrorInvalidValue;
if (a.seqlen % a.tokens_per_block) return cudaErrorInvalidValue; // no partial last block
if (a.sink_tokens - a.tokens_per_block > K_TILE)
return cudaErrorInvalidValue; // large-sink regime
if (a.has_delta && a.q_sink == nullptr) return cudaErrorInvalidValue;
return cudaSuccess;
}
template <bool MHA, bool LPT, bool HAS_SINK_ROPE_DELTA>
static cudaError_t launch_block_causal_sink_impl(const BlockCausalSinkArgs& a,
cudaStream_t stream) {
const int gqa_group_size = a.num_q_heads / a.num_kv_heads; // q-heads per kv-head
const int q_tokens_per_mtile = M_TILE / gqa_group_size; // q-tokens per M-tile
const int q_tokens_per_cta = 2 * q_tokens_per_mtile; // q-tokens per CTA (2 M-tiles)
CUtensorMap tmap_q, tmap_k, tmap_v_t, tmap_o, tmap_q_sink;
{
uint64_t global_dims[4] = {(uint64_t)a.head_dim, (uint64_t)a.num_q_heads, (uint64_t)a.seqlen,
(uint64_t)a.batch};
// FV [B,H,L,D]: head strides by a whole sequence, token by one head_dim row. Same dims, same
// box, same coordinates as the old [B,L,H,D] map -- only these two strides swap roles.
uint64_t global_strides[3] = {(uint64_t)a.seqlen * a.head_dim * 2u, (uint64_t)a.head_dim * 2u,
(uint64_t)a.num_q_heads * a.seqlen * a.head_dim * 2u};
uint32_t box_dims[4] = {(uint32_t)SUB_COLS_BF16, (uint32_t)gqa_group_size,
(uint32_t)q_tokens_per_mtile, 1u};
uint32_t elem_strides[4] = {1u, 1u, 1u, 1u};
CUresult r = cuTensorMapEncodeTiled(
&tmap_q, CU_TENSOR_MAP_DATA_TYPE_BFLOAT16, 4, const_cast<__nv_bfloat16*>(a.q), global_dims,
global_strides, box_dims, elem_strides, 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_CHECK(r == CUDA_SUCCESS ? cudaSuccess : cudaErrorInvalidValue);
r = cuTensorMapEncodeTiled(&tmap_o, CU_TENSOR_MAP_DATA_TYPE_BFLOAT16, 4,
const_cast<__nv_bfloat16*>(a.o), global_dims, global_strides,
box_dims, elem_strides, 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_CHECK(r == CUDA_SUCCESS ? cudaSuccess : cudaErrorInvalidValue);
// q_sink map: identical 4D layout, base = a.q_sink (or a.q when !has_delta -- unused
// placeholder).
r = cuTensorMapEncodeTiled(
&tmap_q_sink, CU_TENSOR_MAP_DATA_TYPE_BFLOAT16, 4,
const_cast<__nv_bfloat16*>(a.has_delta ? a.q_sink : a.q), global_dims, global_strides,
box_dims, elem_strides, 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_CHECK(r == CUDA_SUCCESS ? cudaSuccess : cudaErrorInvalidValue);
}
// K: ONE 3D TMA copy folds the 2 head-dim swizzle atoms (HEAD_DIM = 2 x SUB_COLS_BF16) into the
// box (vs looping 2 x 2D copies). dims [atom-col SUB_COLS_BF16, token ((long)a.batch * a.seqlen),
// atom (num_kv_heads*head_dim)/SUB_COLS_BF16]; box [SUB_COLS_BF16, K_TILE, K_SUBTILES]; strides
// token=(num_kv_heads*head_dim)*2B, atom=SUB_COLS_BF16*2B. The box dim order (atom outermost)
// reproduces the atom-outer smem layout the MMA reads (atom0 then atom1).
{
// FV [B,H,L,D]: head and sample are adjacent with stride L*D (sample stride = HK * L*D), so the
// two fold into ONE dim indexed sample*HK + h_kv. Atom stays the outermost box dim to keep the
// atom-outer smem order the MMA reads.
uint64_t global_dims[4] = {(uint64_t)SUB_COLS_BF16, (uint64_t)a.seqlen,
(uint64_t)(a.head_dim / SUB_COLS_BF16),
(uint64_t)((long)a.batch * a.num_kv_heads)};
uint64_t global_strides[3] = {(uint64_t)a.head_dim * 2u, (uint64_t)SUB_COLS_BF16 * 2u,
(uint64_t)a.seqlen * a.head_dim * 2u};
uint32_t box_dims[4] = {(uint32_t)SUB_COLS_BF16, (uint32_t)K_TILE, (uint32_t)K_SUBTILES, 1u};
uint32_t elem_strides[4] = {1u, 1u, 1u, 1u};
CUresult r = cuTensorMapEncodeTiled(
&tmap_k, CU_TENSOR_MAP_DATA_TYPE_BFLOAT16, 4, const_cast<__nv_bfloat16*>(a.k), global_dims,
global_strides, box_dims, elem_strides, 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_CHECK(r == CUDA_SUCCESS ? cudaSuccess : cudaErrorInvalidValue);
}
{ // FV SPIKE: V map is now byte-for-byte the K map, just over dV.
uint64_t global_dims[4] = {(uint64_t)SUB_COLS_BF16, (uint64_t)a.seqlen,
(uint64_t)(a.head_dim / SUB_COLS_BF16),
(uint64_t)((long)a.batch * a.num_kv_heads)};
uint64_t global_strides[3] = {(uint64_t)a.head_dim * 2u, (uint64_t)SUB_COLS_BF16 * 2u,
(uint64_t)a.seqlen * a.head_dim * 2u};
uint32_t box_dims[4] = {(uint32_t)SUB_COLS_BF16, (uint32_t)K_TILE, (uint32_t)K_SUBTILES, 1u};
uint32_t elem_strides[4] = {1u, 1u, 1u, 1u};
CUresult r = cuTensorMapEncodeTiled(
&tmap_v_t, CU_TENSOR_MAP_DATA_TYPE_BFLOAT16, 4, const_cast<__nv_bfloat16*>(a.v),
global_dims, global_strides, box_dims, elem_strides, 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_CHECK(r == CUDA_SUCCESS ? cudaSuccess : cudaErrorInvalidValue);
}
// ---- shared memory budget ----
const int packed_mtiles_per_seq =
(a.seqlen + q_tokens_per_cta - 1) / q_tokens_per_cta; // packed-M tiles per (sample, kv-head)
// FastDivmod magics for decode_workitem's divides:
// magic0 = mtiles_per_sample (workitem_id -> sample), magic1 = mtiles_per_seq (rr -> kv_head),
// magic2 = num_kv_heads (swizzle path + non-Q_RASTER rr -> tile_index).
const unsigned long long magic0 = make_magic((unsigned)(packed_mtiles_per_seq * a.num_kv_heads));
const unsigned long long magic1 = make_magic((unsigned)packed_mtiles_per_seq);
const unsigned long long magic2 = make_magic((unsigned)a.num_kv_heads);
int lpt_swz_log2 = 0, lpt_hb_quot = 0, lpt_hb_rem = 1;
unsigned long long lpt_major_magic = 1, lpt_rem_magic = 1;
{
const long kv_head_bytes = (long)a.seqlen * (a.head_dim + a.head_dim) * 2; // K + V per kv-head
const long size_l2 = 100L << 20; // GB200 L2 ~126MB; leave headroom
int swz = 1;
while (((long)swz << 1) * kv_head_bytes <= size_l2) swz <<= 1;
const int hb_total = a.batch * a.num_kv_heads;
while (swz > hb_total && swz > 1) swz >>= 1; // clamp to problem
lpt_swz_log2 = 0;
while ((1 << (lpt_swz_log2 + 1)) <= swz) ++lpt_swz_log2;
lpt_hb_quot = hb_total >> lpt_swz_log2;
lpt_hb_rem = hb_total - (lpt_hb_quot << lpt_swz_log2);
if (lpt_hb_rem == 0) lpt_hb_rem = 1;
lpt_major_magic = make_magic((unsigned)(packed_mtiles_per_seq << lpt_swz_log2));
lpt_rem_magic = make_magic((unsigned)lpt_hb_rem);
}
// block-causal-sink runtime bounds (0 tokens_per_block => plain full/causal path).
const int tokens_per_block_arg = (a.tokens_per_block > 0) ? a.tokens_per_block : 0;
const int sink_tokens_arg = (a.tokens_per_block > 0) ? a.sink_tokens : 0;
const int rolling_window_tokens_arg = (a.tokens_per_block > 0) ? a.rolling_window_tokens : 0;
const size_t smem =
(size_t)2 * Q_TILE_BYTES + NUM_KV_STAGES * K_TILE_BYTES // Q (x2) + shared K/V ring
+ (size_t)2 * M_TILE * HEAD_DIM * sizeof(__nv_bfloat16) // 2 sO bufs for TMA-O
+ (2 * NUM_KV_STAGES + 22) * 8 // mbarriers (incl full/empty_bar_o_epi)
+ (size_t)CLC_STAGES * (2 * 8 + 16) + 16 // CLC: clc_full+clc_empty + response (16B aligned)
+ 8 // tmem_slot
+ (size_t)2 * M_TILE * sizeof(float) // alpha_and_l_smem [2][M_TILE]
+ 512; // slack / alignment + isolated wait_scale bar granule
constexpr bool FULL_NAMED_BAR = true, EX2_EMU = true, SPLIT_P = true, SOFTMAX_THROTTLE = true,
Q_RASTER = true;
constexpr bool USE_CLC = false; // PROBE: static + swizzle (FA4's causal config)
auto kernel_fn =
&block_causal_sink_sm100a_kernel<32, FULL_NAMED_BAR, EX2_EMU, SPLIT_P, SOFTMAX_THROTTLE,
USE_CLC, Q_RASTER, MHA, LPT, 8, HAS_SINK_ROPE_DELTA>;
CUDA_CHECK(
cudaFuncSetAttribute(kernel_fn, cudaFuncAttributeMaxDynamicSharedMemorySize, (int)smem));
// ---- launch geometry: CLC persistent. Launch the FULL problem grid (one CTA per work tile);
// clusterlaunchcontrol keeps only ~#SMs CTAs resident and hands the rest of the CTA-ids out via
// try_cancel (HW work-stealing scheduler), so the grid-size is the tile count, not #SMs. ----
// exp2-domain scale: qk * (sm_scale * log2e), matching FastVideo's qk_scale = sm_scale * LOG2E.
// Wrong here corrupts BOTH O and lse (lse = m*scale_log2 + log2 l), and does so plausibly.
const float sm_scale = (a.sm_scale > 0.f) ? a.sm_scale : (1.0f / sqrtf((float)a.head_dim));
const float scale_log2 = sm_scale * (float)M_LOG2E;
int numSM = 0;
CUDA_CHECK(cudaDeviceGetAttribute(&numSM, cudaDevAttrMultiProcessorCount, 0));
const int total_workitems_host = a.batch * packed_mtiles_per_seq *
a.num_kv_heads; // one CTA per (sample, packed-M tile, kv-head)
const int nblk = USE_CLC ? total_workitems_host : std::min(total_workitems_host, numSM);
(void)numSM;
dim3 grid(nblk, 1, 1), block(N_WARPS * 32, 1, 1);
// CLC must be launched via cudaLaunchKernelEx with a cluster-dimension attribute -- a plain
// <<<grid,block>>> launch does NOT enable clusterlaunchcontrol (try_cancel silently misbehaves
// and tiles get skipped). cluster {1,1,1} (matches __cluster_dims__(1,1,1)); no PSS attribute
// (we don't drive griddepcontrol, so leave the dependent-launch serialization off).
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;
auto launch = [&]() {
if (USE_CLC)
return cudaLaunchKernelEx(
&cfg, kernel_fn, tmap_q, tmap_k, tmap_v_t, tmap_o, tmap_q_sink, a.lse, a.seqlen,
a.num_q_heads, a.num_kv_heads, scale_log2, packed_mtiles_per_seq, a.batch, magic0, magic1,
magic2, lpt_swz_log2, lpt_hb_quot, lpt_hb_rem, lpt_major_magic, lpt_rem_magic,
tokens_per_block_arg, sink_tokens_arg, rolling_window_tokens_arg);
kernel_fn<<<grid, block, smem, stream>>>(
tmap_q, tmap_k, tmap_v_t, tmap_o, tmap_q_sink, a.lse, a.seqlen, a.num_q_heads,
a.num_kv_heads, scale_log2, packed_mtiles_per_seq, a.batch, magic0, magic1, magic2,
lpt_swz_log2, lpt_hb_quot, lpt_hb_rem, lpt_major_magic, lpt_rem_magic, tokens_per_block_arg,
sink_tokens_arg, rolling_window_tokens_arg);
return cudaGetLastError();
};
return launch();
}
// Runtime -> template dispatch. MHA (H_q == H_kv) and the relativistic sink correction are
// compile-time in the kernel; the caller only knows them at runtime, so pick the instantiation
// here. LPT (heaviest-first causal balance) is a fixed tuning choice.
inline cudaError_t launch_block_causal_sink_sm100a(const BlockCausalSinkArgs& a,
cudaStream_t stream) {
const cudaError_t bad = block_causal_sink_supported(a);
if (bad != cudaSuccess) return bad;
constexpr bool LPT = true;
const bool mha = (a.num_q_heads == a.num_kv_heads);
if (mha) {
return a.has_delta ? launch_block_causal_sink_impl<true, LPT, true>(a, stream)
: launch_block_causal_sink_impl<true, LPT, false>(a, stream);
}
return a.has_delta ? launch_block_causal_sink_impl<false, LPT, true>(a, stream)
: launch_block_causal_sink_impl<false, LPT, false>(a, stream);
}
@@ -1,88 +0,0 @@
// block_causal_sink_sm100a.cu -- torch binding for the sm_100a block-causal + sink +
// sliding-window FMHA forward.
//
// Forward only: returns (out, lse) so the existing Triton backward keeps working
// unchanged -- lse is exactly the tensor _fwd_kernel writes today.
//
// Nothing is copied or reordered here. Q/K/V/O only need head_dim contiguous; the outer
// strides are read off the tensors, so a permuted (non-contiguous) view costs nothing.
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAGuard.h>
#include "block_causal_sink_launch_sm100a.cuh"
namespace {
void check_qkv(const torch::Tensor& t, const char* name, int64_t B, int64_t H, int64_t L,
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 [B, H, L, D], got ", t.dim(), " dims");
TORCH_CHECK(t.size(0) == B && t.size(1) == H && t.size(2) == L && t.size(3) == D, name,
" has shape ", t.sizes(), ", expected [", B, ",", H, ",", L, ",", D, "]");
// TODO: the TMA descriptors are built for a contiguous [B, H, L, D] tensor. Plumbing the
// caller's strides through BlockCausalSinkArgs would make any permutation of B/H/L free,
// as long as head_dim stays innermost.
TORCH_CHECK(t.is_contiguous(), name, " must be contiguous [B, H, L, D]");
}
} // namespace
// Returns {out, lse}. lse is [B*H_q, L] float32 -- FastVideo's backward input.
std::vector<torch::Tensor> block_causal_sink_sm100a_fwd(
torch::Tensor q, torch::Tensor k, torch::Tensor v, c10::optional<torch::Tensor> q_sink,
int64_t tokens_per_block, int64_t sink_tokens, int64_t rolling_window_tokens, double sm_scale,
bool need_lse) {
const at::cuda::OptionalCUDAGuard guard(device_of(q));
const int64_t B = q.size(0), Hq = q.size(1), L = q.size(2), D = q.size(3);
const int64_t Hkv = k.size(1);
check_qkv(q, "q", B, Hq, L, D);
check_qkv(k, "k", B, Hkv, L, D);
check_qkv(v, "v", B, Hkv, L, D);
TORCH_CHECK(Hq % Hkv == 0, "num_q_heads (", Hq, ") must be divisible by num_kv_heads (", Hkv,
")");
const bool has_delta = q_sink.has_value();
if (has_delta) check_qkv(*q_sink, "q_sink", B, Hq, L, D);
auto out = torch::empty_like(q);
torch::Tensor lse;
if (need_lse) lse = torch::empty({B * Hq, L}, q.options().dtype(torch::kFloat32));
BlockCausalSinkArgs 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.q_sink = has_delta ? reinterpret_cast<const __nv_bfloat16*>(q_sink->data_ptr()) : nullptr;
a.o = reinterpret_cast<__nv_bfloat16*>(out.data_ptr());
a.lse = need_lse ? lse.data_ptr<float>() : nullptr;
a.batch = (int)B;
a.seqlen = (int)L;
a.num_q_heads = (int)Hq;
a.num_kv_heads = (int)Hkv;
a.head_dim = (int)D;
a.tokens_per_block = (int)tokens_per_block;
a.sink_tokens = (int)sink_tokens;
a.rolling_window_tokens = (int)rolling_window_tokens;
a.sm_scale = (float)sm_scale;
a.has_delta = has_delta;
// Report an unsupported regime loudly. Outside it the decode/masking are out of spec and the
// kernel would return plausible-looking but wrong values.
TORCH_CHECK(
block_causal_sink_supported(a) == cudaSuccess,
"block_causal_sink_sm100a: unsupported configuration -- requires head_dim==", HEAD_DIM,
", seqlen divisible by tokens_per_block (no partial last block), and a sink reaching "
"at most one K_TILE past a block end. Got head_dim=",
D, " seqlen=", L, " tokens_per_block=", tokens_per_block, " sink_tokens=", sink_tokens);
const cudaError_t err = launch_block_causal_sink_sm100a(a, at::cuda::getCurrentCUDAStream());
TORCH_CHECK(err == cudaSuccess,
"block_causal_sink_sm100a launch failed: ", cudaGetErrorString(err));
if (need_lse) return {out, lse};
return {out};
}
@@ -1,738 +0,0 @@
// primitives.cuh -- device primitives for the sm_100a block-causal + sink +
// sliding-window 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 <cassert>
#include <cstring>
#include <vector_types.h>
#include <cmath>
#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__ 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");
}
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));
}
}
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));
}
}
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 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));
}
__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__ 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");
}
__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");
}
__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");
}
__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 fence_mbarrier_init_release_cluster() {
asm volatile("fence.mbarrier_init.release.cluster;\n" ::: "memory");
}
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;
}
__device__ __forceinline__ void bar_sync_dyn(uint32_t barrier_id, uint32_t thread_count) {
asm volatile("bar.sync %0, %1;\n" ::"r"(barrier_id), "r"(thread_count) : "memory");
}
__device__ __forceinline__ void bar_arrive_dyn(uint32_t barrier_id, uint32_t thread_count) {
asm volatile("bar.arrive %0, %1;\n" ::"r"(barrier_id), "r"(thread_count) : "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;
}
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;
}
}
__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));
}
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);
}
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");
}
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;
}
} // namespace
__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__ 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);
}
__device__ __forceinline__ void full_bar_arrive(int stage, int band) {
bar_arrive_dyn((uint32_t)((1 + stage * 4) + band), 64);
}
__device__ __forceinline__ void full_bar_wait(int stage, int band) {
bar_sync_dyn((uint32_t)((1 + stage * 4) + band), 64);
}
@@ -22,14 +22,6 @@ extern std::vector<torch::Tensor> block_sparse_attention_backward(
);
#endif
#ifdef TK_COMPILE_BLOCK_CAUSAL_SINK_SM100A
extern std::vector<torch::Tensor> block_causal_sink_sm100a_fwd(
torch::Tensor q, torch::Tensor k, torch::Tensor v,
c10::optional<torch::Tensor> q_sink,
int64_t tokens_per_block, int64_t sink_tokens, int64_t rolling_window_tokens,
double sm_scale, bool need_lse);
#endif
// TurboDiffusion kernels
void register_quant(pybind11::module_ &);
void register_rms_norm(pybind11::module_ &);
@@ -48,12 +40,6 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("block_sparse_bwd", torch::wrap_pybind_function(block_sparse_attention_backward), "block sparse attention backward (Hopper)");
#endif
#ifdef TK_COMPILE_BLOCK_CAUSAL_SINK_SM100A
m.def("block_causal_sink_sm100a_fwd",
torch::wrap_pybind_function(block_causal_sink_sm100a_fwd),
"block-causal + sink + sliding-window attention forward (Blackwell sm100a)");
#endif
// TurboDiffusion
register_quant(m);
register_rms_norm(m);
@@ -20,68 +20,14 @@ def supports_host_descriptor():
return is_cuda() and torch.cuda.get_device_capability()[0] >= 9
def is_sm100(device=None):
return is_cuda() and torch.cuda.get_device_capability(device) == (10, 0)
def is_blackwell():
return is_cuda() and torch.cuda.get_device_capability()[0] == 10
def is_consumer_blackwell():
return is_cuda() and torch.cuda.get_device_capability()[0] == 12
def is_hopper():
return is_cuda() and torch.cuda.get_device_capability()[0] == 9
def _sm100_optimization_enabled():
return os.environ.get("FASTVIDEO_ATTN_QAT_SM100_OPTIMIZED", "1") != "0"
def _sm100_exact_m_enabled():
return os.environ.get("FASTVIDEO_ATTN_QAT_FWD_EXACT_M", "0") != "0"
def _consumer_blackwell_join_qat_pv_enabled():
"""Return whether SM120 uses the joined quantized/STE P@V path."""
return os.environ.get("FASTVIDEO_ATTN_QAT_SM120_JOIN_QAT_PV", "1") != "0"
def _use_sm100_optimized_qat(
device,
head_dim: int,
causal: bool,
is_qat: bool,
fake_quant_p: bool,
two_level_quant_p: bool,
use_global_sf_p: bool,
) -> bool:
"""Return whether this call matches the validated SM100 fast path."""
return (
_sm100_optimization_enabled()
and is_sm100(device)
and head_dim == 128
and not causal
and is_qat
and fake_quant_p
and not two_level_quant_p
and not use_global_sf_p
)
def _select_sm100_forward_config(n_ctx_q: int, n_ctx_kv: int, mode: str):
n_ctx = max(n_ctx_q, n_ctx_kv)
if mode == "reference":
return 32, 32, 4, 4 if n_ctx >= 16_384 else 5
if n_ctx <= 2_048:
return 32, 32, 4, 5
if mode == "balanced":
return 64, 32, 4, 4
return 128, 128, 8, 3
@triton.jit
def _mul_alpha(acc, alpha, BM: tl.constexpr, BN: tl.constexpr):
acc0, acc1 = acc.reshape([BM, 2, BN // 2]).permute(0, 2, 1).split()
@@ -101,8 +47,7 @@ def _attn_fwd_inner(acc, high_prec_acc, l_i, m_i, q, q_valid,
IS_QAT: tl.constexpr,
fake_quant_P: tl.constexpr = True,
two_level_quant_P: tl.constexpr = False,
use_global_sf_P: tl.constexpr = True,
JOIN_QAT_PV: tl.constexpr = False):
use_global_sf_P: tl.constexpr = True):
# range of values handled by this stage (kv blocks)
if STAGE == 1:
lo, hi = 0, start_m * BLOCK_M
@@ -168,19 +113,10 @@ def _attn_fwd_inner(acc, high_prec_acc, l_i, m_i, q, q_valid,
v = desc_v.load([offsetv_y, 0])
v = tl.where(kv_valid[:, None], v, 0.0)
p = p.to(dtype)
# Keep the quantized and STE paths in one tensor-core operation. They
# share V, so joining along M avoids issuing two small, independent
# dot operations for every KV tile.
if IS_QAT and JOIN_QAT_PV:
joined_p = tl.join(p, high_prec_p).permute(2, 0, 1).reshape([2 * BLOCK_M, BLOCK_N])
joined_acc = tl.join(acc, high_prec_acc).permute(2, 0, 1).reshape([2 * BLOCK_M, HEAD_DIM])
joined_acc = tl.dot(joined_p, v.to(dtype), joined_acc)
acc, high_prec_acc = joined_acc.reshape([2, BLOCK_M, HEAD_DIM]).permute(1, 2, 0).split()
else:
# note that this non transposed v for FP8 is only supported on Blackwell
acc = tl.dot(p, v.to(dtype), acc)
if IS_QAT:
high_prec_acc = tl.dot(high_prec_p, v, high_prec_acc)
# note that this non transposed v for FP8 is only supported on Blackwell
acc = tl.dot(p, v.to(dtype), acc)
if IS_QAT:
high_prec_acc = tl.dot(high_prec_p, v, high_prec_acc)
# update m_i and l_i
# place this at the end of the loop to reduce register pressure
l_i = l_i * alpha + l_ij
@@ -268,7 +204,6 @@ def _attn_fwd(sm_scale, M,
fake_quant_P: tl.constexpr = True,
two_level_quant_P: tl.constexpr = False,
use_global_sf_P: tl.constexpr = True,
JOIN_QAT_PV: tl.constexpr = False,
):
dtype = tl.float8e5 if FP8_OUTPUT else tl.bfloat16
tl.static_assert(BLOCK_N <= HEAD_DIM)
@@ -333,7 +268,7 @@ def _attn_fwd(sm_scale, M,
offset_y_kv, dtype, start_m, qk_scale,
BLOCK_M, HEAD_DIM, BLOCK_N,
4 - STAGE, offs_m, offs_n, N_CTX_KV,
warp_specialize, IS_HOPPER, IS_QAT, fake_quant_P, two_level_quant_P, use_global_sf_P, JOIN_QAT_PV
warp_specialize, IS_HOPPER, IS_QAT, fake_quant_P, two_level_quant_P, use_global_sf_P
)
# stage 2: on-band
if STAGE & 2:
@@ -343,7 +278,7 @@ def _attn_fwd(sm_scale, M,
offset_y_kv, dtype, start_m, qk_scale,
BLOCK_M, HEAD_DIM, BLOCK_N,
2, offs_m, offs_n, N_CTX_KV,
warp_specialize, IS_HOPPER, IS_QAT, fake_quant_P, two_level_quant_P, use_global_sf_P, JOIN_QAT_PV
warp_specialize, IS_HOPPER, IS_QAT, fake_quant_P, two_level_quant_P, use_global_sf_P
)
# epilogue
m_i += tl.math.log2(l_i)
@@ -357,52 +292,6 @@ def _attn_fwd(sm_scale, M,
desc_high_prec_o.store([off_hz, start_m * BLOCK_M, 0], high_prec_acc[None, :, :])
@triton.jit
def _attn_fwd_exact_m(
desc_q,
desc_k,
M,
sm_scale,
N_CTX_Q,
N_CTX_KV,
HEAD_DIM: tl.constexpr,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
):
"""Reproduce the legacy 32x32 forward softmax statistic exactly."""
start_m = tl.program_id(0) * BLOCK_M
off_hz = tl.program_id(1)
offs_m = start_m + tl.arange(0, BLOCK_M)
offs_n = tl.arange(0, BLOCK_N)
q_valid = offs_m < N_CTX_Q
q_base = off_hz * N_CTX_Q
kv_base = off_hz * N_CTX_KV
q = desc_q.load([q_base + start_m, 0])
q = tl.where(q_valid[:, None], q, 0.0)
m_i = tl.full([BLOCK_M], -float("inf"), tl.float32)
l_i = tl.full([BLOCK_M], 1.0, tl.float32)
qk_scale = sm_scale * 1.44269504
for start_n in tl.range(0, N_CTX_KV, BLOCK_N):
start_n = tl.multiple_of(start_n, BLOCK_N)
kv_valid = start_n + offs_n < N_CTX_KV
k = desc_k.load([kv_base + start_n, 0])
k = tl.where(kv_valid[:, None], k, 0.0)
qk = tl.dot(q, tl.trans(k))
qk = tl.where(kv_valid[None, :], qk, -1.0e6)
m_ij = tl.maximum(m_i, tl.max(qk, axis=1) * qk_scale)
p = tl.math.exp2(qk * qk_scale - m_ij[:, None])
l_ij = tl.sum(p.to(tl.bfloat16), axis=1)
alpha = tl.math.exp2(m_i - m_ij)
l_i = l_i * alpha + l_ij
m_i = m_ij
m_i += tl.math.log2(l_i)
tl.store(M + off_hz * N_CTX_Q + offs_m, m_i, mask=q_valid)
@triton.jit
def _attn_bwd_preprocess(O, DO,
Delta,
@@ -946,33 +835,10 @@ class _attention(torch.autograd.Function):
assert HEAD_DIM_Q == HEAD_DIM_K and HEAD_DIM_K == HEAD_DIM_V
assert HEAD_DIM_K in {16, 32, 64, 128, 256}
# Triton 3.7's NVWS pass aborts for this kernel on Blackwell. Keep the
# architecture guard next to the kernel so direct callers and the
# FastVideo backend follow the same supported path on sm_100/sm_120.
consumer_blackwell = is_consumer_blackwell()
blackwell = is_blackwell() or consumer_blackwell
warp_specialize = warp_specialize and not blackwell
# Support different sequence lengths for q and k/v (needed for cross attention)
N_CTX_Q = q.shape[2] # Query sequence length
N_CTX_KV = k.shape[2] # Key/Value sequence length (may differ from query)
assert k.shape[2] == v.shape[2], "k and v must have the same sequence length"
sm100_optimized = (
q.dtype == torch.bfloat16
and k.dtype == q.dtype
and v.dtype == q.dtype
and k.device == q.device
and v.device == q.device
and _use_sm100_optimized_qat(
q.device,
HEAD_DIM_K,
causal,
IS_QAT,
fake_quant_P,
two_level_quant_P,
use_global_sf_P,
)
)
# smoothing k from SageAttn
ctx.k_mean = None
@@ -1058,20 +924,7 @@ class _attention(torch.autograd.Function):
else:
extra_kern_args["maxnreg"] = 80
qkv_block_m, qkv_block_n = 32, 32
fwd_block_m, fwd_block_n = 32, 32
fwd_num_warps, fwd_num_stages = 4, 2
fwd_mode = "legacy"
if sm100_optimized:
fwd_mode = os.environ.get("FASTVIDEO_ATTN_QAT_FWD_MODE", "fast").lower()
if fwd_mode not in {"fast", "balanced", "reference"}:
raise ValueError(
f"FASTVIDEO_ATTN_QAT_FWD_MODE={fwd_mode!r} "
"(want fast|balanced|reference)"
)
fwd_block_m, fwd_block_n, fwd_num_warps, fwd_num_stages = _select_sm100_forward_config(
N_CTX_Q, N_CTX_KV, fwd_mode
)
BLOCK_M, BLOCK_N = 32, 32
if IS_QAT:
fake_q = torch.empty_like(q)
fake_k = torch.empty_like(k)
@@ -1089,8 +942,8 @@ class _attention(torch.autograd.Function):
desc_v = fake_v
H = q.shape[1]
grid_1 = (triton.cdiv(q.shape[2], qkv_block_m), q.shape[0] * q.shape[1], 1)
grid_2 = (triton.cdiv(k.shape[2], qkv_block_n), q.shape[0] * q.shape[1], 1)
grid_1 = (triton.cdiv(q.shape[2], BLOCK_M), q.shape[0] * q.shape[1], 1)
grid_2 = (triton.cdiv(k.shape[2], BLOCK_N), q.shape[0] * q.shape[1], 1)
fake_quantize_q[grid_1](
q, fake_q,
@@ -1099,7 +952,7 @@ class _attention(torch.autograd.Function):
fake_q.stride(0), fake_q.stride(1),
fake_q.stride(2), fake_q.stride(3),
H, N_CTX_Q,
BLOCK_M=qkv_block_m, HEAD_DIM=HEAD_DIM_K,
BLOCK_M=BLOCK_M, HEAD_DIM=HEAD_DIM_K,
use_global_sf=use_global_sf_QKV,
)
fake_quantize_kv[grid_2](
@@ -1109,14 +962,14 @@ class _attention(torch.autograd.Function):
fake_k.stride(0), fake_k.stride(1),
fake_k.stride(2), fake_k.stride(3),
H, N_CTX_KV,
BLOCK_N=qkv_block_n, HEAD_DIM=HEAD_DIM_K,
BLOCK_N=BLOCK_N, HEAD_DIM=HEAD_DIM_K,
use_global_sf=use_global_sf_QKV,
)
# Apply pre-hook to set block shapes on tensor descriptors
_host_descriptor_pre_hook({
"BLOCK_M": fwd_block_m,
"BLOCK_N": fwd_block_n,
"BLOCK_M": BLOCK_M,
"BLOCK_N": BLOCK_N,
"HEAD_DIM": HEAD_DIM_K,
"desc_q": desc_q,
"desc_k": desc_k,
@@ -1133,7 +986,7 @@ class _attention(torch.autograd.Function):
N_CTX_Q=N_CTX_Q,
N_CTX_KV=N_CTX_KV,
HEAD_DIM=HEAD_DIM_K,
BLOCK_M=fwd_block_m, BLOCK_N=fwd_block_n,
BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N,
FP8_OUTPUT=q.dtype == torch.float8_e5m2,
STAGE=stage,
warp_specialize=warp_specialize,
@@ -1142,40 +995,10 @@ class _attention(torch.autograd.Function):
fake_quant_P=fake_quant_P,
two_level_quant_P=two_level_quant_P,
use_global_sf_P=use_global_sf_P,
JOIN_QAT_PV=(consumer_blackwell and _consumer_blackwell_join_qat_pv_enabled()),
num_warps=fwd_num_warps,
num_stages=fwd_num_stages,
num_warps=4,
num_stages=2,
**extra_kern_args
)
exact_m = _sm100_exact_m_enabled()
if (
sm100_optimized
and fwd_mode != "reference"
and exact_m
and (fwd_block_m != 32 or fwd_block_n != 32)
):
# The large forward tile changes a legal reduction order. Restore
# the legacy statistic so dV remains bitwise-compatible while the
# two output paths retain the faster large-tile PV computation.
assert isinstance(desc_q, TensorDescriptor)
assert isinstance(desc_k, TensorDescriptor)
desc_q.block_shape = [32, HEAD_DIM_K]
desc_k.block_shape = [32, HEAD_DIM_K]
stats_grid = (triton.cdiv(N_CTX_Q, 32), q.shape[0] * q.shape[1])
_attn_fwd_exact_m[stats_grid](
desc_q,
desc_k,
M,
sm_scale,
N_CTX_Q,
N_CTX_KV,
HEAD_DIM=HEAD_DIM_K,
BLOCK_M=32,
BLOCK_N=32,
num_warps=8,
num_stages=4,
)
o_for_bwd = high_prec_o if IS_QAT and use_high_prec_o else o
if IS_QAT:
@@ -1195,7 +1018,6 @@ class _attention(torch.autograd.Function):
ctx.smooth_q = smooth_q
ctx.use_global_sf_P = use_global_sf_P
ctx.warp_specialize = warp_specialize
ctx.sm100_optimized = sm100_optimized
return o
@staticmethod
@@ -1210,10 +1032,7 @@ class _attention(torch.autograd.Function):
N_CTX_KV = k.shape[2]
assert k.shape[2] == v.shape[2], "k and v must have the same sequence length"
PRE_BLOCK = 128
# Long video sequences are occupancy-bound on consumer Blackwell: a
# third software-pipeline stage consumes shared memory without hiding
# additional latency. Shorter sequences retain the deeper pipeline.
NUM_STAGES = 2 if is_consumer_blackwell() and max(N_CTX_Q, N_CTX_KV) >= 8192 else 3
NUM_STAGES = 3
NUM_WARPS = 4
BLOCK_M1, BLOCK_N1, BLOCK_M2, BLOCK_N2 = 32, 32, 32, 32
if not ctx.use_qat_qkv_backward:
@@ -1238,85 +1057,7 @@ class _attention(torch.autograd.Function):
# _, q_m = triton_group_mean(q)
q_m = q_m.repeat_interleave(q.shape[2] // q_m.shape[2], dim=2) # B,H,L,D
sm100_optimized_backward = (
getattr(ctx, "sm100_optimized", False)
and ctx.use_qat_qkv_backward
and not ctx.smooth_k
and not ctx.smooth_q
and N_CTX_KV % 16 == 0
)
if sm100_optimized_backward:
# Keeping dQ and dK/dV in separate programs allows 64x64 tiles
# without carrying all three fp32 accumulators at once. On SM100
# this is substantially faster than the legacy 32x32 combined
# self-attention program with the same math and BF16 parity bounds.
block_m, block_n = 64, 64
grid_dq = ((N_CTX_Q + block_m - 1) // block_m, 1, BATCH * N_HEAD)
_attn_bwd_dq_cross[grid_dq](
q,
arg_k,
v,
ctx.sm_scale,
do,
dq,
M,
delta,
q.stride(0),
k.stride(0),
q.stride(1),
k.stride(1),
q.stride(2),
k.stride(2),
q.stride(3),
k.stride(3),
N_HEAD,
N_CTX_Q,
N_CTX_KV,
ctx.k_mean,
BLOCK_M2=block_m,
BLOCK_N2=block_n,
HEAD_DIM=ctx.HEAD_DIM,
SMOOTH_K=False,
warp_specialize=False,
num_warps=8,
num_stages=2,
)
grid_dkdv = ((N_CTX_KV + block_n - 1) // block_n, 1, BATCH * N_HEAD)
_attn_bwd_dkdv_cross[grid_dkdv](
q,
arg_k,
v,
ctx.sm_scale,
do,
dk,
dv,
M,
delta,
q_m,
q.stride(0),
k.stride(0),
q.stride(1),
k.stride(1),
q.stride(2),
k.stride(2),
q.stride(3),
k.stride(3),
N_HEAD,
N_CTX_Q,
N_CTX_KV,
BLOCK_M1=block_m,
BLOCK_N1=block_n,
HEAD_DIM=ctx.HEAD_DIM,
IS_QAT=True,
two_level_quant_P=False,
fake_quant_P=True,
SMOOTH_Q=False,
use_global_sf_P=False,
warp_specialize=False,
num_warps=8,
num_stages=3,
)
elif N_CTX_Q == N_CTX_KV:
if N_CTX_Q == N_CTX_KV:
# Use existing kernel for self-attention (same sequence lengths)
grid = ((N_CTX_KV + BLOCK_N1 - 1) // BLOCK_N1, 1, BATCH * N_HEAD)
_attn_bwd[grid](
@@ -1333,10 +1074,10 @@ class _attention(torch.autograd.Function):
IS_QAT=ctx.IS_QAT,
SMOOTH_K=ctx.smooth_k,
two_level_quant_P=ctx.two_level_quant_P,
fake_quant_P=ctx.fake_quant_P,
SMOOTH_Q=ctx.smooth_q,
use_global_sf_P=ctx.use_global_sf_P,
warp_specialize=ctx.warp_specialize,
fake_quant_P=ctx.fake_quant_P,
SMOOTH_Q=ctx.smooth_q,
use_global_sf_P=ctx.use_global_sf_P,
warp_specialize=ctx.warp_specialize,
num_warps=NUM_WARPS,
num_stages=NUM_STAGES
)
@@ -1,190 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
import math
import pytest
import torch
from fastvideo_kernel.triton_kernels import attn_qat_train as kernel
def _production_route_kwargs():
return {
"device": torch.device("cuda"),
"head_dim": 128,
"causal": False,
"is_qat": True,
"fake_quant_p": True,
"two_level_quant_p": False,
"use_global_sf_p": False,
}
def test_sm100_production_configuration_uses_optimized_route(monkeypatch):
monkeypatch.setattr(kernel, "is_sm100", lambda device=None: True)
monkeypatch.delenv("FASTVIDEO_ATTN_QAT_SM100_OPTIMIZED", raising=False)
assert kernel._use_sm100_optimized_qat(**_production_route_kwargs())
@pytest.mark.parametrize(
("override", "value"),
[
("head_dim", 64),
("causal", True),
("is_qat", False),
("fake_quant_p", False),
("two_level_quant_p", True),
("use_global_sf_p", True),
],
)
def test_unsupported_configuration_keeps_legacy_route(monkeypatch, override, value):
monkeypatch.setattr(kernel, "is_sm100", lambda device=None: True)
kwargs = _production_route_kwargs()
kwargs[override] = value
assert not kernel._use_sm100_optimized_qat(**kwargs)
def test_non_sm100_and_debug_switch_keep_legacy_route(monkeypatch):
kwargs = _production_route_kwargs()
monkeypatch.setattr(kernel, "is_sm100", lambda device=None: False)
assert not kernel._use_sm100_optimized_qat(**kwargs)
monkeypatch.setattr(kernel, "is_sm100", lambda device=None: True)
monkeypatch.setenv("FASTVIDEO_ATTN_QAT_SM100_OPTIMIZED", "0")
assert not kernel._use_sm100_optimized_qat(**kwargs)
def test_exact_m_is_opt_in(monkeypatch):
monkeypatch.delenv("FASTVIDEO_ATTN_QAT_FWD_EXACT_M", raising=False)
assert not kernel._sm100_exact_m_enabled()
monkeypatch.setenv("FASTVIDEO_ATTN_QAT_FWD_EXACT_M", "1")
assert kernel._sm100_exact_m_enabled()
def test_sm120_joined_pv_is_enabled_by_default_and_can_be_disabled(monkeypatch):
monkeypatch.delenv("FASTVIDEO_ATTN_QAT_SM120_JOIN_QAT_PV", raising=False)
assert kernel._consumer_blackwell_join_qat_pv_enabled()
monkeypatch.setenv("FASTVIDEO_ATTN_QAT_SM120_JOIN_QAT_PV", "0")
assert not kernel._consumer_blackwell_join_qat_pv_enabled()
@pytest.mark.parametrize(
("n_ctx", "mode", "expected"),
[
(2_048, "fast", (32, 32, 4, 5)),
(4_096, "fast", (128, 128, 8, 3)),
(4_096, "balanced", (64, 32, 4, 4)),
(4_096, "reference", (32, 32, 4, 5)),
(31_200, "reference", (32, 32, 4, 4)),
],
)
def test_sm100_forward_config_selection(n_ctx, mode, expected):
assert kernel._select_sm100_forward_config(n_ctx, n_ctx, mode) == expected
@pytest.mark.skipif(
not torch.cuda.is_available() or torch.cuda.get_device_capability() != (10, 0),
reason="SM100 parity test",
)
@pytest.mark.parametrize(("q_length", "kv_length"), [(2_112, 2_112), (2_112, 2_080)])
def test_sm100_optimized_forward_backward_matches_legacy(monkeypatch, q_length, kv_length):
torch.manual_seed(7)
q_shape = (1, 1, q_length, 128)
kv_shape = (1, 1, kv_length, 128)
inputs = [
torch.randn(q_shape, device="cuda", dtype=torch.bfloat16),
torch.randn(kv_shape, device="cuda", dtype=torch.bfloat16),
torch.randn(kv_shape, device="cuda", dtype=torch.bfloat16),
]
grad_out = torch.randn(q_shape, device="cuda", dtype=torch.bfloat16)
flags = (
True, # use_qat_qkv_backward
False, # smooth_k
True, # warp_specialize (disabled internally on Blackwell)
True, # IS_QAT
False, # two_level_quant_P
True, # fake_quant_P
True, # use_high_prec_o
False, # smooth_q
False, # use_global_sf_P
False, # use_global_sf_QKV
)
def run(optimized: bool):
monkeypatch.setenv("FASTVIDEO_ATTN_QAT_SM100_OPTIMIZED", "1" if optimized else "0")
monkeypatch.setenv("FASTVIDEO_ATTN_QAT_FWD_MODE", "fast")
monkeypatch.setenv("FASTVIDEO_ATTN_QAT_FWD_EXACT_M", "1")
q, k, v = [tensor.clone().requires_grad_(True) for tensor in inputs]
output = kernel.attention(
q,
k,
v,
False,
1.0 / math.sqrt(q_shape[-1]),
*flags,
)
output.backward(grad_out)
return output.detach(), q.grad, k.grad, v.grad
legacy = run(False)
optimized = run(True)
assert (optimized[0].float() - legacy[0].float()).abs().max().item() <= 1e-2
assert (optimized[1].float() - legacy[1].float()).abs().max().item() <= 4e-3
assert (optimized[2].float() - legacy[2].float()).abs().max().item() <= 4e-3
assert torch.equal(optimized[3], legacy[3])
@pytest.mark.skipif(
not torch.cuda.is_available() or torch.cuda.get_device_capability()[0] != 12,
reason="SM120 parity test",
)
@pytest.mark.parametrize(("q_length", "kv_length"), [(2_112, 2_112), (2_112, 2_080)])
def test_sm120_joined_pv_forward_backward_matches_split_path(monkeypatch, q_length, kv_length):
torch.manual_seed(11)
q_shape = (1, 1, q_length, 128)
kv_shape = (1, 1, kv_length, 128)
inputs = [
torch.randn(q_shape, device="cuda", dtype=torch.bfloat16),
torch.randn(kv_shape, device="cuda", dtype=torch.bfloat16),
torch.randn(kv_shape, device="cuda", dtype=torch.bfloat16),
]
grad_out = torch.randn(q_shape, device="cuda", dtype=torch.bfloat16)
flags = (
True, # use_qat_qkv_backward
False, # smooth_k
True, # warp_specialize (disabled internally on Blackwell)
True, # IS_QAT
False, # two_level_quant_P
True, # fake_quant_P
True, # use_high_prec_o
False, # smooth_q
False, # use_global_sf_P
False, # use_global_sf_QKV
)
def run(joined_pv: bool):
monkeypatch.setenv("FASTVIDEO_ATTN_QAT_SM120_JOIN_QAT_PV", "1" if joined_pv else "0")
q, k, v = [tensor.clone().requires_grad_(True) for tensor in inputs]
output = kernel.attention(
q,
k,
v,
False,
1.0 / math.sqrt(q_shape[-1]),
*flags,
)
output.backward(grad_out)
return output.detach(), q.grad, k.grad, v.grad
split = run(False)
joined = run(True)
assert torch.equal(joined[0], split[0])
assert torch.equal(joined[1], split[1])
assert torch.equal(joined[2], split[2])
assert torch.equal(joined[3], split[3])
-11
View File
@@ -304,8 +304,6 @@ def generator_config_to_fastvideo_args(config: GeneratorConfig | Mapping[str, An
for key in _LTX2_REFINE_FLAT_KEYS:
if key in refine:
kwargs[f"ltx2_refine_{key}"] = refine[key]
if "enabled" in refine:
kwargs["refine_enabled"] = refine["enabled"]
kwargs.update(preset_overrides)
kwargs.update(deepcopy(normalized.pipeline.experimental))
return FastVideoArgs.from_kwargs(**kwargs)
@@ -328,9 +326,6 @@ def legacy_generate_call_to_request(
mouse_cond: Any | None = None,
keyboard_cond: Any | None = None,
grid_sizes: Any | None = None,
track_points: Any | None = None,
track_visibility: Any | None = None,
track_ids: Any | None = None,
legacy_kwargs: Mapping[str, Any] | None = None,
) -> GenerationRequest:
raw = _sampling_param_to_request_raw(sampling_param)
@@ -346,12 +341,6 @@ def legacy_generate_call_to_request(
raw.setdefault("inputs", {})["keyboard_cond"] = keyboard_cond
if grid_sizes is not None:
raw.setdefault("inputs", {})["grid_sizes"] = grid_sizes
if track_points is not None:
raw.setdefault("inputs", {})["track_points"] = track_points
if track_visibility is not None:
raw.setdefault("inputs", {})["track_visibility"] = track_visibility
if track_ids is not None:
raw.setdefault("inputs", {})["track_ids"] = track_ids
normalized = parse_config(GenerationRequest, raw)
bind_generation_request_raw(normalized, raw)
+1 -30
View File
@@ -38,10 +38,6 @@ class SamplingParam:
keyboard_cond: Any | None = None # Shape: (B, T, K)
grid_sizes: Any | None = None # Shape: (3,) [F,H,W]
track_points: Any | None = None # Shape: (B, T, N, 2)
track_visibility: Any | None = None # Shape: (B, T, N)
track_ids: Any | None = None # Shape: (B, N)
# Camera control inputs (HYWorld)
pose: str | None = None # Camera trajectory: pose string (e.g., 'w-31') or JSON file path
prompt_attention_mask: list = field(default_factory=list)
@@ -55,9 +51,8 @@ class SamplingParam:
gt_latents: Any | None = None # Ground truth latents [B, 16, T, H, W]
conditioning_mask: Any | None = None # Mask [B, 1, T, H, W]
# Camera control inputs (LingBotWorld and LingBotWorld2)
# Camera control inputs (LingBotWorld)
c2ws_plucker_emb: Any | None = None # Plucker embedding: [B, C, F_lat, H_lat, W_lat]
action_path: str | None = None # Directory containing poses.npy and intrinsics.npy
# Refine inputs (LongCat 480p->720p upscaling)
# Path-based refine (load stage1 video from disk, e.g. MP4)
@@ -94,13 +89,7 @@ class SamplingParam:
num_inference_steps: int = 50
num_inference_steps_sr: int = 50
guidance_scale: float = 1.0
batch_cfg: bool = False
guidance_scale_2: float | None = None
# Z-Image CFG controls. ``cfg_normalization=True`` caps the guided
# prediction norm at the positive-prediction norm; ``cfg_truncation``
# disables CFG above the normalized-noise threshold.
cfg_normalization: bool = False
cfg_truncation: float | None = 1.0
# Embedded guidance (FLUX): do not treat ``guidance_scale > 1`` as classic CFG.
use_embedded_guidance: bool = False
# Diffusers-style true CFG for FLUX when > 1 (requires negative prompt encoding).
@@ -334,24 +323,6 @@ class SamplingParam:
default=SamplingParam.guidance_scale,
help="Classifier-free guidance scale",
)
parser.add_argument(
"--cfg-normalization",
action=StoreBoolean,
default=SamplingParam.cfg_normalization,
help="Cap Z-Image CFG prediction norm to the positive-prediction norm",
)
parser.add_argument(
"--cfg-truncation",
type=float,
default=SamplingParam.cfg_truncation,
help="Disable Z-Image CFG above this normalized-noise threshold",
)
parser.add_argument(
"--batch-cfg",
action=StoreBoolean,
default=SamplingParam.batch_cfg,
help="Evaluate conditional and unconditional CFG branches in one batch",
)
parser.add_argument(
"--guidance-rescale",
type=float,
-8
View File
@@ -129,11 +129,7 @@ class InputConfig:
mouse_cond: Any | None = None
keyboard_cond: Any | None = None
grid_sizes: Any | None = None
track_points: Any | None = None
track_visibility: Any | None = None
track_ids: Any | None = None
c2ws_plucker_emb: Any | None = None
action_path: str | None = None
refine_from: str | None = None
stage1_video: Any | None = None
@@ -142,7 +138,6 @@ class InputConfig:
class SamplingConfig:
num_videos_per_prompt: int = 1
seed: int = 1024
max_sequence_length: int | None = None
num_frames: int = 125
height: int = 720
width: int = 1280
@@ -152,10 +147,7 @@ class SamplingConfig:
num_inference_steps: int = 50
num_inference_steps_sr: int = 50
guidance_scale: float = 1.0
batch_cfg: bool = False
guidance_scale_2: float | None = None
cfg_normalization: bool = False
cfg_truncation: float | None = 1.0
guidance_rescale: float = 0.0
true_cfg_scale: float | None = None
use_embedded_guidance: bool | None = None
-17
View File
@@ -1,17 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass
from fastvideo.api.sampling_param import SamplingParam
@dataclass
class WanTrackSamplingParam(SamplingParam):
"""Sampling defaults for causal WanTrack Self-Forcing I2V."""
height: int = 480
width: int = 832
num_frames: int = 121
fps: int = 16
guidance_scale: float = 1.0
num_inference_steps: int = 4
negative_prompt: str | None = None
+9 -18
View File
@@ -18,14 +18,14 @@ from fastvideo.logger import init_logger
logger = init_logger(__name__)
_project_root = Path(__file__).resolve().parent.parent.parent.parent
_kernel_python_root = _project_root / "fastvideo-kernel" / "python"
_kernel_root = _project_root / "fastvideo-kernel"
_kernel_python_root = _kernel_root / "python"
_attn_qat_train_attention: Callable[..., torch.Tensor] | None = None
_attn_qat_train_import_attempted = False
_attn_qat_train_import_error: ImportError | None = None
def _ensure_kernel_paths() -> None:
for path in (_project_root, _kernel_python_root):
for path in (_project_root, _kernel_root, _kernel_python_root):
path_str = str(path)
if path_str not in sys.path:
sys.path.insert(0, path_str)
@@ -34,7 +34,6 @@ def _ensure_kernel_paths() -> None:
def _get_attn_qat_train_attention() -> Callable[..., torch.Tensor] | None:
global _attn_qat_train_attention
global _attn_qat_train_import_attempted
global _attn_qat_train_import_error
if _attn_qat_train_import_attempted:
return _attn_qat_train_attention
@@ -43,11 +42,8 @@ def _get_attn_qat_train_attention() -> Callable[..., torch.Tensor] | None:
_ensure_kernel_paths()
try:
triton_qat = importlib.import_module("fastvideo_kernel.triton_kernels.attn_qat_train")
_attn_qat_train_attention = triton_qat.attention
logger.info("ATTN_QAT_TRAIN loaded FastVideo's architecture-optimized Triton kernel")
except ImportError as exc:
_attn_qat_train_import_error = exc
_attn_qat_train_attention = importlib.import_module("fastvideo_kernel.triton_kernels.attn_qat_train").attention
except ImportError:
_attn_qat_train_attention = None
return _attn_qat_train_attention
@@ -64,9 +60,8 @@ def attn_qat_train(q_BLHD: torch.Tensor,
sm_scale: float | None = None) -> torch.Tensor:
attention = _get_attn_qat_train_attention()
if attention is None:
detail = f" Original import error: {_attn_qat_train_import_error}" if _attn_qat_train_import_error else ""
raise ImportError("ATTN_QAT_TRAIN requires FastVideo's fastvideo-kernel package. Install it or make "
f"fastvideo-kernel/python importable.{detail}")
raise ImportError("fastvideo_kernel.triton_kernels.attn_qat_train is not available. "
"Please ensure the FastVideo kernel package is installed.")
q_BHLD = q_BLHD.permute(0, 2, 1, 3).contiguous()
k_BHLD = k_BLHD.permute(0, 2, 1, 3).contiguous()
@@ -74,11 +69,7 @@ def attn_qat_train(q_BLHD: torch.Tensor,
use_qat_qkv_backward = True
smooth_k = False
# Triton 3.7's NVWS pass aborts while compiling this kernel on Blackwell.
# The kernel has a supported non-warp-specialized path, so use it on both
# datacenter (sm_100) and consumer (sm_120) Blackwell GPUs.
capability_major = torch.cuda.get_device_capability()[0]
warp_specialize = capability_major not in (10, 12)
warp_specialize = True
is_qat = True
two_level_quant_p_sage3 = False
fake_quant_p_bwd = True
@@ -115,7 +106,7 @@ class AttnQatTrainBackend(AttentionBackend):
@staticmethod
def get_supported_head_sizes() -> list[int]:
return [128]
return [64, 96, 128, 160, 192, 224, 256]
@staticmethod
def get_name() -> str:
+2 -5
View File
@@ -180,11 +180,8 @@ class FlashAttentionImpl(AttentionImpl):
# SP). Cast through bf16 and restore, matching TORCH_SDPA's tolerance.
orig_dtype = query.dtype
if orig_dtype not in (torch.float16, torch.bfloat16):
# Keep logging (and its Python lru_cache) out of compiled traces
# so that the AOTAutograd cache can serialize.
if not torch.compiler.is_compiling():
logger.warning_once(f"FLASH_ATTN received {orig_dtype} inputs; casting to "
f"bfloat16 for the kernel and restoring on output.")
logger.warning_once(f"FLASH_ATTN received {orig_dtype} inputs; casting to "
f"bfloat16 for the kernel and restoring on output.")
query = query.to(torch.bfloat16)
key = key.to(torch.bfloat16)
value = value.to(torch.bfloat16)
-6
View File
@@ -1,6 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Standalone Triton kernels used by specific model code paths.
Unlike ``fastvideo/attention/backends``, these are not registered with the
attention-backend selector; models import them directly.
"""
File diff suppressed because it is too large Load Diff
@@ -1,74 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""CUDA (sm_100a) forward for the block-causal + sink + sliding-window training attention.
Goes in FastVideo as ``fastvideo/attention/kernels/block_causal_sink_cuda.py``.
This replaces ONLY the forward. It returns ``out`` and ``lse`` in exactly the form
``_fwd_kernel`` produces them, so ``_BlockCausalSinkAttention.backward`` and its Triton
kernels are reused untouched -- see INTEGRATION.md for the ~6-line patch to
``block_causal_sink.py``.
Scope: ``kind="blockwise"``, Blackwell (sm_100a), bf16, ``head_dim == 128``, uniform
sequence length, ``num_frames % num_frame_per_block == 0``, and a sink reaching at most one
128-token tile past a block end. Everything else must fall back to Triton -- outside that
regime the reference is self-inconsistent, so returning an answer at all would be wrong.
``is_supported()`` is the predicate; callers should consult it rather than assume.
"""
from __future__ import annotations
import torch
try:
from fastvideo_kernel._C import fastvideo_kernel_ops as _C
_HAS_CUDA_BCS = hasattr(_C, "block_causal_sink_sm100a_fwd")
except ImportError: # pragma: no cover - extension not built
_C = None
_HAS_CUDA_BCS = False
_SM100 = (10, 0) # compiled for Blackwell only
HEAD_DIM = 128
K_TILE = 128
def is_supported(plan, q: torch.Tensor) -> bool:
"""True iff this backend can run `plan` on `q`; otherwise the caller uses Triton."""
if not _HAS_CUDA_BCS or not q.is_cuda:
return False
if torch.cuda.get_device_capability(q.device) != _SM100:
return False
if plan.kind != "blockwise":
return False # teacher_forcing is a separate kernel
if q.dtype != torch.bfloat16 or q.shape[-1] != HEAD_DIM:
return False
# TODO: the TMA descriptors are built for a contiguous [B, H, L, D] tensor, so only that
# layout is accepted today. Plumbing the caller's strides through BlockCausalSinkArgs would
# make any permutation of B/H/L free (head_dim must stay innermost).
if not (q.is_contiguous() and q.dim() == 4):
return False
if plan.local_attn_size is None or plan.local_attn_size < 0:
return False
if plan.num_frames % plan.num_frame_per_block != 0:
return False # partial last block is out of spec
if (plan.sink_size - plan.num_frame_per_block) * plan.frame_seqlen > K_TILE:
return False # large-sink regime is out of spec
return plan.num_frame_per_block * plan.frame_seqlen > 0
def block_causal_sink_forward_cuda(q, k, v, q_sink, plan):
"""Forward pass. q/k/v/q_sink: ``[B, H, L, D]`` bf16. Returns ``(out, lse)``.
Tensors are consumed with whatever strides they arrive with -- a permuted view is fine
and is NOT copied; only ``head_dim`` has to be the contiguous axis. ``lse`` is
``[B*H, L]`` float32, identical in meaning to the Triton forward's.
"""
out, lse = _C.block_causal_sink_sm100a_fwd(
q, k, v,
q_sink if q_sink is not None and q_sink is not q else None,
plan.num_frame_per_block * plan.frame_seqlen, # tokens_per_block
plan.sink_size * plan.frame_seqlen, # sink_tokens
max(plan.local_attn_size - plan.sink_size, 0) * plan.frame_seqlen,
# local_attn_size is the TOTAL budget: sink frames plus the trailing window.
float(plan.sm_scale),
True, # need_lse (backward consumes it)
)
return out, lse
-26
View File
@@ -32,27 +32,6 @@ def backend_name_to_enum(backend_name: str) -> AttentionBackendEnum | None:
None
def coerce_attn_backend(attn_backend: AttentionBackendEnum | str | None, ) -> AttentionBackendEnum | None:
"""Normalize an explicit backend selection.
Environment-variable parsing remains permissive via
:func:`backend_name_to_enum`, but typed/config-driven call sites should
fail fast on typos instead of silently falling back to another backend.
"""
if attn_backend is None or isinstance(attn_backend, AttentionBackendEnum):
return attn_backend
if not isinstance(attn_backend, str) or not attn_backend.strip():
raise ValueError("attention backend must be a non-empty string, "
f"an AttentionBackendEnum, or None; got {attn_backend!r}")
backend_name = attn_backend.strip().upper()
backend = backend_name_to_enum(backend_name)
if backend is None:
raise ValueError(f"Unknown attention backend {attn_backend!r}. "
f"Expected one of {sorted(AttentionBackendEnum.__members__)}")
return backend
def get_env_variable_attn_backend() -> AttentionBackendEnum | None:
'''
Get the backend override specified by the FastVideo attention
@@ -90,11 +69,6 @@ def global_force_attn_backend(attn_backend: AttentionBackendEnum | None) -> None
'''
global forced_attn_backend
forced_attn_backend = attn_backend
# Backend selection is cached by tensor shape/dtype, while the global
# override is intentionally not part of that cache key. Invalidate cached
# resolutions whenever the override changes so independently constructed
# role models can bind different attention implementations.
_cached_get_attn_backend.cache_clear()
def get_global_forced_attn_backend() -> AttentionBackendEnum | None:
+5 -11
View File
@@ -1,5 +1,6 @@
from __future__ import annotations
import functools
from collections.abc import Callable
import torch
@@ -42,26 +43,19 @@ except ImportError:
_flash_attn_2_func = None
_flash_attn_2_varlen_func = None
# Dynamo unwraps functools caches, so keep this cache inside the opaque helper.
_SM90_OR_NEWER_BY_DEVICE: dict[int, bool] = {}
def _check_dropout(dropout_p: float) -> None:
if dropout_p != 0.0:
raise NotImplementedError(f"flash_attn.cute does not support dropout (got dropout_p={dropout_p})")
@torch.compiler.assume_constant_result
def _sm90_or_newer(device_id: int) -> bool:
if device_id not in _SM90_OR_NEWER_BY_DEVICE:
_SM90_OR_NEWER_BY_DEVICE[device_id] = (current_platform.has_device_capability(90, device_id))
return _SM90_OR_NEWER_BY_DEVICE[device_id]
@functools.cache
def _sm90_or_newer() -> bool:
return current_platform.has_device_capability(90)
def _use_fa2(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> bool:
device_id = q.device.index
assert device_id is not None
if _sm90_or_newer(device_id):
if _sm90_or_newer():
return False
# Pre-sm90 FA4 cute limitations, both served by FA2 (deterministic
# capability gate, not a runtime fallback):
+1 -6
View File
@@ -11,18 +11,13 @@ from fastvideo.configs.models.dits.longcat import LongCatVideoConfig
from fastvideo.configs.models.dits.ltx2 import LTX2VideoConfig
from fastvideo.configs.models.dits.magi_human import MagiHumanVideoConfig
from fastvideo.configs.models.dits.stable_audio import StableAudioConfig
from fastvideo.configs.models.dits.trackwan import CausalTrackWanVideoConfig, TrackWanVideoConfig
from fastvideo.configs.models.dits.wanvideo import WanVideoConfig
from fastvideo.configs.models.dits.zimage import ZImageDiTConfig
from fastvideo.configs.models.dits.hyworld import HYWorldConfig
from fastvideo.configs.models.dits.kandinsky5 import Kandinsky5VideoConfig
from fastvideo.configs.models.dits.lingbotworld2 import LingBotWorld2CausalFastVideoConfig
from fastvideo.configs.models.dits.lingbot_video import LingBotVideoConfig
__all__ = [
"HunyuanVideoConfig", "HunyuanVideo15Config", "HunyuanGameCraftConfig", "WanVideoConfig", "DreamXWorldConfig",
"DreamXWorldARConfig", "CosmosVideoConfig", "Cosmos25VideoConfig", "FluxDiTConfig", "Flux2Config",
"LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig", "Kandinsky5VideoConfig", "MagiHumanVideoConfig",
"StableAudioConfig", "GlmImageDiTConfig", "LingBotWorld2CausalFastVideoConfig", "LingBotVideoConfig",
"ZImageDiTConfig", "TrackWanVideoConfig", "CausalTrackWanVideoConfig"
"StableAudioConfig", "GlmImageDiTConfig"
]
@@ -1,70 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Architecture configuration for LingBot-Video Dense and MoE DiTs."""
from __future__ import annotations
from dataclasses import dataclass, field
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
from fastvideo.platforms import AttentionBackendEnum
def _is_lingbot_video_block(name: str, module: object) -> bool:
"""Select top-level transformer blocks for FSDP and compilation."""
del module
parts = name.split(".")
return len(parts) == 2 and parts[0] == "blocks" and parts[1].isdigit()
@dataclass
class LingBotVideoArchConfig(DiTArchConfig):
"""One-to-one representation of the released transformer config JSON."""
_fsdp_shard_conditions: list = field(default_factory=lambda: [_is_lingbot_video_block])
_supported_attention_backends: tuple[AttentionBackendEnum, ...] = (
AttentionBackendEnum.TORCH_SDPA,
AttentionBackendEnum.FLASH_ATTN,
)
param_names_mapping: dict = field(default_factory=lambda: {r"^(.*)$": r"\1"})
patch_size: tuple[int, int, int] = (1, 2, 2)
in_channels: int = 16
out_channels: int = 16
hidden_size: int = 2048
num_attention_heads: int = 16
depth: int = 24
intermediate_size: int = 6144
text_dim: int = 2560
freq_dim: int = 256
norm_eps: float = 1e-6
rope_theta: float = 256.0
axes_dims: tuple[int, int, int] = (32, 48, 48)
axes_lens: tuple[int, int, int] = (8192, 1024, 1024)
qkv_bias: bool = False
out_bias: bool = True
patch_embed_bias: bool = True
timestep_mlp_bias: bool = True
num_experts: int = 0
num_experts_per_tok: int = 8
moe_intermediate_size: int = 512
decoder_sparse_step: int = 1
mlp_only_layers: tuple[int, ...] = ()
n_shared_experts: int | None = None
score_func: str = "sigmoid"
norm_topk_prob: bool = True
n_group: int | None = None
topk_group: int | None = None
routed_scaling_factor: float = 1.0
def __post_init__(self) -> None:
"""Populate FastVideo loader fields from the released architecture."""
super().__post_init__()
self.num_channels_latents = self.in_channels
self.attention_head_dim = self.hidden_size // self.num_attention_heads
@dataclass
class LingBotVideoConfig(DiTConfig):
"""FastVideo component configuration for LingBot-Video transformers."""
arch_config: DiTArchConfig = field(default_factory=LingBotVideoArchConfig)
@@ -1,55 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
def is_blocks(n: str, m) -> bool:
return "blocks" in n and str.isdigit(n.split(".")[-1])
@dataclass
class LingBotWorld2CausalFastArchConfig(DiTArchConfig):
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_blocks])
param_names_mapping: dict = field(default_factory=dict)
reverse_param_names_mapping: dict = field(default_factory=dict)
lora_param_names_mapping: dict = field(default_factory=dict)
model_type: str = "i2v"
patch_size: tuple[int, int, int] = (1, 2, 2)
text_len: int = 512
in_dim: int = 36
dim: int = 5120
ffn_dim: int = 13824
freq_dim: int = 256
text_dim: int = 4096
out_dim: int = 16
num_heads: int = 40
num_layers: int = 40
qk_norm: bool = True
cross_attn_norm: bool = True
eps: float = 1e-6
local_attn_size: int = 18
sink_size: int = 6
chunk_size: int = 4
sample_shift: float = 10.0
num_train_timesteps: int = 1000
timesteps_index: tuple[int, int, int, int] = (0, 250, 500, 750)
max_area: int = 480 * 832
def __post_init__(self):
super().__post_init__()
self.hidden_size = self.dim
self.num_attention_heads = self.num_heads
self.attention_head_dim = self.dim // self.num_heads
self.in_channels = self.in_dim
self.out_channels = self.out_dim
self.num_channels_latents = self.out_dim
@dataclass
class LingBotWorld2CausalFastVideoConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=LingBotWorld2CausalFastArchConfig)
prefix: str = "Wan"
-61
View File
@@ -1,61 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Configuration for track-conditioned Wan transformers."""
from dataclasses import dataclass, field
from fastvideo.configs.models.dits.base import DiTArchConfig
from fastvideo.configs.models.dits.wanvideo import WanVideoArchConfig, WanVideoConfig
def _default_track_config() -> dict[str, int | bool]:
return {
"id_dim": 64,
"track_channels": 16,
"vae_spatial_compression": 8,
"vae_temporal_compression": 4,
"max_track_id": 100_000,
"zero_init_head": False,
"use_bias": False,
}
@dataclass
class TrackWanVideoArchConfig(WanVideoArchConfig):
"""Wan I2V input widened with a latent-aligned point-track map.
Channel order is fixed and checkpoint-visible:
noisy latent (16), I2V mask (4), first-frame latent (16), track map (16).
"""
in_channels: int = 52
out_channels: int = 16
image_dim: int = 1280
added_kv_proj_dim: int | None = 5120
track_config: dict[str, int | bool] = field(default_factory=_default_track_config)
def __post_init__(self) -> None:
super().__post_init__()
track_channels = int(self.track_config.get("track_channels", 0))
expected_in_channels = self.num_channels_latents + 20 + track_channels
if track_channels <= 0:
raise ValueError("track_config.track_channels must be positive")
if self.in_channels != expected_in_channels:
raise ValueError("TrackWan in_channels must equal latent channels + 20 I2V "
f"channels + track channels; got {self.in_channels}, expected "
f"{expected_in_channels}")
@dataclass
class TrackWanVideoConfig(WanVideoConfig):
arch_config: DiTArchConfig = field(default_factory=TrackWanVideoArchConfig)
prefix: str = "Wan"
@dataclass
class CausalTrackWanVideoArchConfig(TrackWanVideoArchConfig):
"""Explicit config type for the causal TrackWan architecture."""
@dataclass
class CausalTrackWanVideoConfig(TrackWanVideoConfig):
arch_config: DiTArchConfig = field(default_factory=CausalTrackWanVideoArchConfig)
@@ -91,11 +91,6 @@ class WanVideoArchConfig(DiTArchConfig):
# RoPE policy for the causal-rollout paths (causal Wan / MatrixGame2).
# "relativistic" keeps long rollouts in-distribution; a no-op unless sink_size > 0 and local_attn_size > 0.
rope_cache_policy: str = "absolute"
# Full-sequence training attention implementation for causal models:
# "flex" (FlexAttention BlockMask, default), "triton" (fused sink +
# rolling-window kernel with exact relativistic sink RoPE correction), or
# "reference" (slow pure-PyTorch, for tests).
causal_train_attention: str = "flex"
# AnyFlow dual-timestep conditioning. Defaults preserve bit-identity with
# the legacy single-timestep forward (no delta_embedder allocated, no
-60
View File
@@ -1,60 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
from fastvideo.platforms import AttentionBackendEnum
def is_zimage_block(name: str, module) -> bool:
parts = name.split(".")
return len(parts) >= 2 and parts[-2] in {"noise_refiner", "context_refiner", "layers"} and parts[-1].isdigit()
@dataclass
class ZImageDiTArchConfig(DiTArchConfig):
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_zimage_block])
_supported_attention_backends: tuple[AttentionBackendEnum, ...] = (AttentionBackendEnum.TORCH_SDPA, )
all_patch_size: tuple[int, ...] = (2, )
all_f_patch_size: tuple[int, ...] = (1, )
in_channels: int = 16
dim: int = 3840
n_layers: int = 30
n_refiner_layers: int = 2
n_heads: int = 30
n_kv_heads: int = 30
norm_eps: float = 1e-5
qk_norm: bool = True
cap_feat_dim: int = 2560
rope_theta: float = 256.0
t_scale: float = 1000.0
axes_dims: tuple[int, ...] = (32, 48, 48)
axes_lens: tuple[int, ...] = (1536, 512, 512)
adaln_embed_dim: int = 256
frequency_embedding_size: int = 256
timestep_mid_size: int = 1024
max_period: int = 10000
seq_multi_of: int = 32
def __post_init__(self) -> None:
super().__post_init__()
if len(self.all_patch_size) != len(self.all_f_patch_size):
raise ValueError("all_patch_size and all_f_patch_size must have equal length")
if self.dim % self.n_heads:
raise ValueError("dim must be divisible by n_heads")
if self.dim // self.n_heads != sum(self.axes_dims):
raise ValueError("attention head dimension must equal sum(axes_dims)")
if len(self.axes_dims) != len(self.axes_lens) or any(dim % 2 for dim in self.axes_dims):
raise ValueError("RoPE axes require matching lengths and even dimensions")
self.hidden_size = self.dim
self.num_attention_heads = self.n_heads
self.num_channels_latents = self.in_channels
self.out_channels = self.in_channels
@dataclass
class ZImageDiTConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=ZImageDiTArchConfig)
prefix: str = "ZImage"
@@ -2,7 +2,6 @@ from fastvideo.configs.models.encoders.base import (BaseEncoderOutput, EncoderCo
TextEncoderConfig)
from fastvideo.configs.models.encoders.clip import (CLIPTextConfig, CLIPVisionConfig, WAN2_1ControlCLIPVisionConfig)
from fastvideo.configs.models.encoders.llama import LlamaConfig
from fastvideo.configs.models.encoders.lingbotworld2_t5 import LingBotWorld2UMT5ArchConfig, LingBotWorld2UMT5Config
from fastvideo.configs.models.encoders.t5 import T5Config, T5LargeConfig
from fastvideo.configs.models.encoders.qwen2_5 import Qwen2_5_VLConfig
from fastvideo.configs.models.encoders.siglip import SiglipVisionConfig
@@ -10,7 +9,6 @@ from fastvideo.configs.models.encoders.reason1 import Reason1ArchConfig, Reason1
from fastvideo.configs.models.encoders.gemma import LTX2GemmaConfig
from fastvideo.configs.models.encoders.mistral3 import Mistral3TextConfig
from fastvideo.configs.models.encoders.qwen3 import Qwen3TextConfig
from fastvideo.configs.models.encoders.lingbot_video import LingBotVideoQwen3VLTextConfig
from fastvideo.configs.models.encoders.stable_audio_conditioner import (StableAudioConditionerArchConfig,
StableAudioConditionerConfig)
from fastvideo.configs.models.encoders.t5gemma import T5GemmaEncoderConfig
@@ -19,6 +17,5 @@ __all__ = [
"EncoderConfig", "TextEncoderConfig", "ImageEncoderConfig", "BaseEncoderOutput", "CLIPTextConfig",
"CLIPVisionConfig", "WAN2_1ControlCLIPVisionConfig", "LlamaConfig", "T5Config", "T5LargeConfig", "Qwen2_5_VLConfig",
"Reason1ArchConfig", "Reason1Config", "LTX2GemmaConfig", "SiglipVisionConfig", "StableAudioConditionerArchConfig",
"StableAudioConditionerConfig", "T5GemmaEncoderConfig", "Qwen3TextConfig", "Mistral3TextConfig",
"LingBotWorld2UMT5ArchConfig", "LingBotWorld2UMT5Config", "LingBotVideoQwen3VLTextConfig"
"StableAudioConditionerConfig", "T5GemmaEncoderConfig", "Qwen3TextConfig", "Mistral3TextConfig"
]
@@ -83,7 +83,6 @@ class TextEncoderConfig(EncoderConfig):
arch_config: ArchConfig = field(default_factory=TextEncoderArchConfig)
is_chat_model: bool = False
treat_empty_as_dot: bool = False
chat_template_enable_thinking: bool = field(default=False, kw_only=True)
@dataclass

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