Compare commits
37
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
9f94673c34 | ||
|
|
c150889330 | ||
|
|
7f783e30b7 | ||
|
|
ec9db97431 | ||
|
|
fd236a06a7 | ||
|
|
54fe0be488 | ||
|
|
3a398a2bbd | ||
|
|
4e76897f1f | ||
|
|
95d8e74d65 | ||
|
|
bcd0ad7dee | ||
|
|
9acbf92455 | ||
|
|
ce4fa8ef2d | ||
|
|
41dcd065ae | ||
|
|
c0cb2228ac | ||
|
|
2fa0f70ee8 | ||
|
|
e172215880 | ||
|
|
23e3b325a4 | ||
|
|
44d0990199 | ||
|
|
59b6578906 | ||
|
|
fd8372d172 | ||
|
|
33eb7ea13c | ||
|
|
051d21958d | ||
|
|
bf29e20bf5 | ||
|
|
291fa5d9a6 | ||
|
|
954daf7fb3 | ||
|
|
e05c04a2f6 | ||
|
|
9fb74b9732 | ||
|
|
521dee0e82 | ||
|
|
229419208e | ||
|
|
65f3b946b9 | ||
|
|
191fcbf46c | ||
|
|
755a4e4470 | ||
|
|
9709b7513b | ||
|
|
32cd603515 | ||
|
|
d4bdd3621a | ||
|
|
e2f8322842 | ||
|
|
6966f9e0bc |
@@ -1,6 +1,8 @@
|
||||
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:
|
||||
|
||||
@@ -27,14 +27,25 @@ 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.
|
||||
- name: Save trusted hook config
|
||||
# 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
|
||||
if: github.event_name == 'pull_request_target'
|
||||
run: cp .pre-commit-config.yaml "$RUNNER_TEMP/trusted-pre-commit-config.yaml"
|
||||
- uses: actions/checkout@v4
|
||||
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
|
||||
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
|
||||
@@ -48,5 +59,7 @@ 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 .github/scripts/test_gate_full_suite.sh
|
||||
run: bash "${GATE_SCRIPTS_DIR:-.github/scripts}/test_gate_full_suite.sh"
|
||||
|
||||
@@ -11,9 +11,7 @@ on:
|
||||
- 'requirements-mkdocs.txt'
|
||||
- 'scripts/check_docs_links.py'
|
||||
- '.github/workflows/infra-docs.yml'
|
||||
# Run the trusted base-branch workflow so fork PRs can be skipped without
|
||||
# waiting for maintainer approval.
|
||||
pull_request_target:
|
||||
pull_request:
|
||||
branches: [ main ]
|
||||
paths:
|
||||
- 'docs/**'
|
||||
@@ -26,19 +24,21 @@ 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,7 +52,6 @@ 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
|
||||
@@ -62,22 +61,17 @@ 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.event_name == 'push'
|
||||
if: github.ref == 'refs/heads/main'
|
||||
steps:
|
||||
- name: Deploy to GitHub Pages
|
||||
id: deployment
|
||||
|
||||
@@ -6,6 +6,7 @@ results/
|
||||
wandb/
|
||||
*.ipynb
|
||||
*.jpg
|
||||
!examples/dataset/lingbotworld2/image.jpg
|
||||
*.safetensors
|
||||
*.mp4
|
||||
*.png
|
||||
@@ -34,6 +35,7 @@ env
|
||||
*.log
|
||||
weights/
|
||||
logs/
|
||||
/Z-Image/
|
||||
official_weights/
|
||||
converted_weights/
|
||||
|
||||
|
||||
@@ -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. [FastWan-QAD models](https://huggingface.co/FastVideo/FastWan-QAD-FP8-1.3B), check out the [Blog](https://haoailab.com/blogs/fastwan-qad/).
|
||||
- `2026/06/23`: Release FastWan-QAD: 5s of Video generated in 1.8s E2E. See the [FastWan-QAD models](https://huggingface.co/FastVideo/FastWan-QAD-FP8-1.3B), [Attn-QAT training guide](https://haoailab.com/FastVideo/training/attn_qat/), and [blog](https://haoailab.com/blogs/fastwan-qad/).
|
||||
- `2026/03/17`: Release demo: Into the Dreamverse: Vibe Directing in FastVideo, check out the [Blog](https://haoailab.com/blogs/dreamverse/).
|
||||
- `2026/03/13`: Release demo: Create a 5s 1080p Video in 4.5s with FastVideo on a Single GPU, check out the [Blog](https://haoailab.com/blogs/fastvideo_realtime_1080p/).
|
||||
- `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,6 +39,17 @@ 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
|
||||
|
||||
@@ -0,0 +1,35 @@
|
||||
# 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.
|
||||
@@ -0,0 +1,5 @@
|
||||
"""Standalone causal WanTrack control prototype."""
|
||||
|
||||
from apps.wantrack_control.server import create_app
|
||||
|
||||
__all__ = ["create_app"]
|
||||
@@ -0,0 +1,18 @@
|
||||
"""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()
|
||||
@@ -0,0 +1,236 @@
|
||||
"""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())
|
||||
@@ -0,0 +1,422 @@
|
||||
"""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()
|
||||
@@ -0,0 +1,308 @@
|
||||
(() => {
|
||||
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();
|
||||
})();
|
||||
@@ -0,0 +1,56 @@
|
||||
<!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>
|
||||
@@ -0,0 +1,57 @@
|
||||
: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; } }
|
||||
@@ -0,0 +1,267 @@
|
||||
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
|
||||
@@ -8,13 +8,15 @@ 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 both STA and VSA kernels.
|
||||
source, which includes STA, VSA, and Attn-QAT kernels.
|
||||
|
||||
### Prerequisites
|
||||
|
||||
|
||||
@@ -255,10 +255,9 @@ If you add a new CI test category:
|
||||
|
||||
### Documentation
|
||||
|
||||
`.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.
|
||||
`.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.
|
||||
|
||||
The docs job:
|
||||
|
||||
|
||||
@@ -333,15 +333,25 @@ surfaces:
|
||||
use_distill:
|
||||
sources: [fastvideo.configs.pipelines.longcat.LongCatT2V480PConfig, fastvideo.configs.pipelines.longcat.LongCatT2V704PConfig]
|
||||
scheduler_arch:
|
||||
sources: [fastvideo.configs.pipelines.sd35.SD35Config]
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.sd35.SD35Config
|
||||
- fastvideo.configs.pipelines.zimage.ZImagePipelineConfig
|
||||
text_encoder_archs:
|
||||
sources: [fastvideo.configs.pipelines.sd35.SD35Config]
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.sd35.SD35Config
|
||||
- fastvideo.configs.pipelines.zimage.ZImagePipelineConfig
|
||||
tokenizer_archs:
|
||||
sources: [fastvideo.configs.pipelines.sd35.SD35Config]
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.sd35.SD35Config
|
||||
- fastvideo.configs.pipelines.zimage.ZImagePipelineConfig
|
||||
transformer_arch:
|
||||
sources: [fastvideo.configs.pipelines.sd35.SD35Config]
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.sd35.SD35Config
|
||||
- fastvideo.configs.pipelines.zimage.ZImagePipelineConfig
|
||||
vae_arch:
|
||||
sources: [fastvideo.configs.pipelines.sd35.SD35Config]
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.sd35.SD35Config
|
||||
- fastvideo.configs.pipelines.zimage.ZImagePipelineConfig
|
||||
expand_timesteps:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.wan.FastWan2_2_TI2V_5B_Config
|
||||
@@ -421,6 +431,8 @@ 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."
|
||||
@@ -438,6 +450,7 @@ 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
|
||||
@@ -447,6 +460,7 @@ 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
|
||||
@@ -456,7 +470,10 @@ 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
|
||||
@@ -508,7 +525,6 @@ 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: {}
|
||||
|
||||
|
||||
@@ -2,6 +2,12 @@
|
||||
|
||||
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:
|
||||
|
||||
@@ -0,0 +1,147 @@
|
||||
# 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).
|
||||
@@ -62,6 +62,18 @@ 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.
|
||||
@@ -166,6 +178,7 @@ 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:
|
||||
|
||||
|
||||
@@ -43,6 +43,9 @@ 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
|
||||
@@ -59,9 +62,11 @@ 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. **Run inference**: After training, see [inference examples](../inference/examples/examples_inference_index.md)
|
||||
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)
|
||||
|
||||
@@ -81,6 +81,7 @@ 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:
|
||||
|
||||
@@ -298,6 +299,8 @@ 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 |
|
||||
|
||||
@@ -0,0 +1,188 @@
|
||||
# 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.
|
||||
@@ -78,7 +78,10 @@ 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` install exposing `fastvideo_kernel`
|
||||
- `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.
|
||||
|
||||
As a fallback, use:
|
||||
|
||||
|
||||
@@ -0,0 +1,17 @@
|
||||
# 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.
|
After Width: | Height: | Size: 1.2 MiB |
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,52 @@
|
||||
"""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()
|
||||
@@ -0,0 +1,52 @@
|
||||
# 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()
|
||||
@@ -0,0 +1,156 @@
|
||||
# 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()
|
||||
@@ -0,0 +1,96 @@
|
||||
# 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()
|
||||
@@ -91,3 +91,7 @@ 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/).
|
||||
|
||||
@@ -0,0 +1,86 @@
|
||||
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
|
||||
@@ -0,0 +1,110 @@
|
||||
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
|
||||
@@ -0,0 +1,72 @@
|
||||
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
|
||||
@@ -0,0 +1,86 @@
|
||||
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
|
||||
@@ -0,0 +1,110 @@
|
||||
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
|
||||
@@ -0,0 +1,72 @@
|
||||
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
|
||||
@@ -0,0 +1,86 @@
|
||||
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
|
||||
@@ -0,0 +1,110 @@
|
||||
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
|
||||
@@ -0,0 +1,72 @@
|
||||
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
|
||||
@@ -0,0 +1,89 @@
|
||||
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
|
||||
@@ -0,0 +1,111 @@
|
||||
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
|
||||
@@ -0,0 +1,73 @@
|
||||
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
|
||||
@@ -0,0 +1,52 @@
|
||||
# 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.
|
||||
@@ -0,0 +1,40 @@
|
||||
# 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
|
||||
```
|
||||
@@ -0,0 +1,9 @@
|
||||
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
|
||||
|
@@ -0,0 +1,74 @@
|
||||
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
|
||||
@@ -0,0 +1,7 @@
|
||||
{
|
||||
"data": [
|
||||
{
|
||||
"caption": "A lone rider guides a horse across an open field at sunset, with steady motion and a slowly changing background."
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -0,0 +1,43 @@
|
||||
# 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.
|
||||
@@ -0,0 +1,95 @@
|
||||
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
|
||||
@@ -0,0 +1,106 @@
|
||||
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
|
||||
@@ -0,0 +1,79 @@
|
||||
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
|
||||
@@ -25,6 +25,7 @@ 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
|
||||
@@ -52,6 +53,8 @@ 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
|
||||
|
||||
@@ -0,0 +1,82 @@
|
||||
# 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
|
||||
@@ -0,0 +1,88 @@
|
||||
# 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
|
||||
@@ -0,0 +1,95 @@
|
||||
# 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: {}
|
||||
@@ -0,0 +1,90 @@
|
||||
# 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: {}
|
||||
@@ -5,9 +5,12 @@ all configs, scripts, and data needed to run a complete workflow.
|
||||
|
||||
```
|
||||
scenario/
|
||||
└── ode_init_self_forcing_wan_causal/ # KD → export → Self-Forcing
|
||||
├── ode_init_self_forcing_wan_causal/ # KD → export → Self-Forcing
|
||||
└── qad_wan2_1_mixkit/ # Attn-QAT SFT → export → DMD2
|
||||
```
|
||||
|
||||
See the `usage.md` inside each scenario for step-by-step instructions.
|
||||
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/).
|
||||
|
||||
For single-step configs, see `examples/train/configs/`.
|
||||
|
||||
@@ -0,0 +1,15 @@
|
||||
#!/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}"
|
||||
+20
@@ -0,0 +1,20 @@
|
||||
#!/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}"
|
||||
+27
@@ -0,0 +1,27 @@
|
||||
#!/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
|
||||
@@ -0,0 +1,72 @@
|
||||
# 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
|
||||
@@ -0,0 +1,101 @@
|
||||
# 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
|
||||
@@ -0,0 +1,26 @@
|
||||
# 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** — selected purely by an env var, no monkey-patching:
|
||||
path is **config-driven** and selected by an environment variable:
|
||||
|
||||
```bash
|
||||
bash examples/training/finetune/wan_t2v_1.3B/mixkit/finetune_qat.sh
|
||||
@@ -60,10 +60,19 @@ 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` 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).
|
||||
`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/).
|
||||
|
||||
## Train stage 2 (QAT DMD distillation to 3 steps)
|
||||
|
||||
@@ -71,8 +80,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 — no per-model
|
||||
flags or monkey-patching.
|
||||
same global `ATTN_QAT_TRAIN` env reaches **only** the generator, with no
|
||||
per-model flags.
|
||||
|
||||
```bash
|
||||
# generator init = the stage-1 finetune checkpoint
|
||||
|
||||
@@ -2,20 +2,23 @@
|
||||
# QAD recipe — quantization-aware finetune of Wan2.1-T2V-1.3B with fake-quant
|
||||
# (Attn-QAT) attention.
|
||||
#
|
||||
# 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.
|
||||
# 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.
|
||||
#
|
||||
# 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. The kernel is Triton so it runs on sm_100 and
|
||||
# sm_120 alike (the FP4 inference kernel, by contrast, is sm_120-only).
|
||||
# validation generates videos.
|
||||
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
|
||||
|
||||
@@ -23,6 +26,8 @@ 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 \
|
||||
@@ -30,15 +35,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 2000 --train_batch_size 1 --train_sp_batch_size 1 \
|
||||
--max_train_steps "${MAX_TRAIN_STEPS}" --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 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 \
|
||||
--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 \
|
||||
--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 \
|
||||
--dit_precision fp32 --num_euler_timesteps 50 --ema_start_step 0 --flow_shift 1 \
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
|
||||
@@ -308,6 +308,21 @@ 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
|
||||
@@ -336,6 +351,9 @@ 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})
|
||||
|
||||
@@ -347,6 +365,13 @@ 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(
|
||||
|
||||
@@ -108,6 +108,33 @@ 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:
|
||||
|
||||
@@ -0,0 +1,127 @@
|
||||
#!/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
@@ -0,0 +1,235 @@
|
||||
// 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);
|
||||
}
|
||||
@@ -0,0 +1,88 @@
|
||||
// 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};
|
||||
}
|
||||
@@ -0,0 +1,738 @@
|
||||
// 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,6 +22,14 @@ 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_ &);
|
||||
@@ -40,6 +48,12 @@ 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,14 +20,68 @@ 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()
|
||||
@@ -47,7 +101,8 @@ 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):
|
||||
use_global_sf_P: tl.constexpr = True,
|
||||
JOIN_QAT_PV: tl.constexpr = False):
|
||||
# range of values handled by this stage (kv blocks)
|
||||
if STAGE == 1:
|
||||
lo, hi = 0, start_m * BLOCK_M
|
||||
@@ -113,10 +168,19 @@ 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)
|
||||
# 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)
|
||||
# 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)
|
||||
# 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
|
||||
@@ -204,6 +268,7 @@ 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)
|
||||
@@ -268,7 +333,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
|
||||
warp_specialize, IS_HOPPER, IS_QAT, fake_quant_P, two_level_quant_P, use_global_sf_P, JOIN_QAT_PV
|
||||
)
|
||||
# stage 2: on-band
|
||||
if STAGE & 2:
|
||||
@@ -278,7 +343,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
|
||||
warp_specialize, IS_HOPPER, IS_QAT, fake_quant_P, two_level_quant_P, use_global_sf_P, JOIN_QAT_PV
|
||||
)
|
||||
# epilogue
|
||||
m_i += tl.math.log2(l_i)
|
||||
@@ -292,6 +357,52 @@ 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,
|
||||
@@ -835,10 +946,33 @@ 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
|
||||
@@ -924,7 +1058,20 @@ class _attention(torch.autograd.Function):
|
||||
else:
|
||||
extra_kern_args["maxnreg"] = 80
|
||||
|
||||
BLOCK_M, BLOCK_N = 32, 32
|
||||
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
|
||||
)
|
||||
if IS_QAT:
|
||||
fake_q = torch.empty_like(q)
|
||||
fake_k = torch.empty_like(k)
|
||||
@@ -942,8 +1089,8 @@ class _attention(torch.autograd.Function):
|
||||
desc_v = fake_v
|
||||
|
||||
H = q.shape[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)
|
||||
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)
|
||||
|
||||
fake_quantize_q[grid_1](
|
||||
q, fake_q,
|
||||
@@ -952,7 +1099,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=BLOCK_M, HEAD_DIM=HEAD_DIM_K,
|
||||
BLOCK_M=qkv_block_m, HEAD_DIM=HEAD_DIM_K,
|
||||
use_global_sf=use_global_sf_QKV,
|
||||
)
|
||||
fake_quantize_kv[grid_2](
|
||||
@@ -962,14 +1109,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=BLOCK_N, HEAD_DIM=HEAD_DIM_K,
|
||||
BLOCK_N=qkv_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": BLOCK_M,
|
||||
"BLOCK_N": BLOCK_N,
|
||||
"BLOCK_M": fwd_block_m,
|
||||
"BLOCK_N": fwd_block_n,
|
||||
"HEAD_DIM": HEAD_DIM_K,
|
||||
"desc_q": desc_q,
|
||||
"desc_k": desc_k,
|
||||
@@ -986,7 +1133,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=BLOCK_M, BLOCK_N=BLOCK_N,
|
||||
BLOCK_M=fwd_block_m, BLOCK_N=fwd_block_n,
|
||||
FP8_OUTPUT=q.dtype == torch.float8_e5m2,
|
||||
STAGE=stage,
|
||||
warp_specialize=warp_specialize,
|
||||
@@ -995,10 +1142,40 @@ 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,
|
||||
num_warps=4,
|
||||
num_stages=2,
|
||||
JOIN_QAT_PV=(consumer_blackwell and _consumer_blackwell_join_qat_pv_enabled()),
|
||||
num_warps=fwd_num_warps,
|
||||
num_stages=fwd_num_stages,
|
||||
**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:
|
||||
@@ -1018,6 +1195,7 @@ 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
|
||||
@@ -1032,7 +1210,10 @@ 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
|
||||
NUM_STAGES = 3
|
||||
# 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_WARPS = 4
|
||||
BLOCK_M1, BLOCK_N1, BLOCK_M2, BLOCK_N2 = 32, 32, 32, 32
|
||||
if not ctx.use_qat_qkv_backward:
|
||||
@@ -1057,7 +1238,85 @@ 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
|
||||
|
||||
if N_CTX_Q == N_CTX_KV:
|
||||
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:
|
||||
# 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](
|
||||
@@ -1074,10 +1333,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
|
||||
)
|
||||
|
||||
@@ -0,0 +1,190 @@
|
||||
# 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])
|
||||
@@ -304,6 +304,8 @@ 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)
|
||||
@@ -326,6 +328,9 @@ 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)
|
||||
@@ -341,6 +346,12 @@ 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)
|
||||
|
||||
@@ -38,6 +38,10 @@ 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)
|
||||
@@ -51,8 +55,9 @@ 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)
|
||||
# Camera control inputs (LingBotWorld and LingBotWorld2)
|
||||
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)
|
||||
@@ -89,7 +94,13 @@ 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).
|
||||
@@ -323,6 +334,24 @@ 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,
|
||||
|
||||
@@ -129,7 +129,11 @@ 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
|
||||
|
||||
@@ -138,6 +142,7 @@ 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
|
||||
@@ -147,7 +152,10 @@ 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
|
||||
|
||||
@@ -0,0 +1,17 @@
|
||||
# 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
|
||||
@@ -18,14 +18,14 @@ from fastvideo.logger import init_logger
|
||||
logger = init_logger(__name__)
|
||||
|
||||
_project_root = Path(__file__).resolve().parent.parent.parent.parent
|
||||
_kernel_root = _project_root / "fastvideo-kernel"
|
||||
_kernel_python_root = _kernel_root / "python"
|
||||
_kernel_python_root = _project_root / "fastvideo-kernel" / "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_root, _kernel_python_root):
|
||||
for path in (_project_root, _kernel_python_root):
|
||||
path_str = str(path)
|
||||
if path_str not in sys.path:
|
||||
sys.path.insert(0, path_str)
|
||||
@@ -34,6 +34,7 @@ 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
|
||||
@@ -42,8 +43,11 @@ def _get_attn_qat_train_attention() -> Callable[..., torch.Tensor] | None:
|
||||
_ensure_kernel_paths()
|
||||
|
||||
try:
|
||||
_attn_qat_train_attention = importlib.import_module("fastvideo_kernel.triton_kernels.attn_qat_train").attention
|
||||
except ImportError:
|
||||
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 = None
|
||||
|
||||
return _attn_qat_train_attention
|
||||
@@ -60,8 +64,9 @@ 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:
|
||||
raise ImportError("fastvideo_kernel.triton_kernels.attn_qat_train is not available. "
|
||||
"Please ensure the FastVideo kernel package is installed.")
|
||||
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}")
|
||||
|
||||
q_BHLD = q_BLHD.permute(0, 2, 1, 3).contiguous()
|
||||
k_BHLD = k_BLHD.permute(0, 2, 1, 3).contiguous()
|
||||
@@ -69,7 +74,11 @@ def attn_qat_train(q_BLHD: torch.Tensor,
|
||||
|
||||
use_qat_qkv_backward = True
|
||||
smooth_k = False
|
||||
warp_specialize = True
|
||||
# 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)
|
||||
is_qat = True
|
||||
two_level_quant_p_sage3 = False
|
||||
fake_quant_p_bwd = True
|
||||
@@ -106,7 +115,7 @@ class AttnQatTrainBackend(AttentionBackend):
|
||||
|
||||
@staticmethod
|
||||
def get_supported_head_sizes() -> list[int]:
|
||||
return [64, 96, 128, 160, 192, 224, 256]
|
||||
return [128]
|
||||
|
||||
@staticmethod
|
||||
def get_name() -> str:
|
||||
|
||||
@@ -180,8 +180,11 @@ 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):
|
||||
logger.warning_once(f"FLASH_ATTN received {orig_dtype} inputs; casting to "
|
||||
f"bfloat16 for the kernel and restoring on output.")
|
||||
# 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.")
|
||||
query = query.to(torch.bfloat16)
|
||||
key = key.to(torch.bfloat16)
|
||||
value = value.to(torch.bfloat16)
|
||||
|
||||
@@ -0,0 +1,6 @@
|
||||
# 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
@@ -0,0 +1,74 @@
|
||||
# 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
|
||||
@@ -32,6 +32,27 @@ 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
|
||||
@@ -69,6 +90,11 @@ 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:
|
||||
|
||||
@@ -1,6 +1,5 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import functools
|
||||
from collections.abc import Callable
|
||||
|
||||
import torch
|
||||
@@ -43,19 +42,26 @@ 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})")
|
||||
|
||||
|
||||
@functools.cache
|
||||
def _sm90_or_newer() -> bool:
|
||||
return current_platform.has_device_capability(90)
|
||||
@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]
|
||||
|
||||
|
||||
def _use_fa2(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> bool:
|
||||
if _sm90_or_newer():
|
||||
device_id = q.device.index
|
||||
assert device_id is not None
|
||||
if _sm90_or_newer(device_id):
|
||||
return False
|
||||
# Pre-sm90 FA4 cute limitations, both served by FA2 (deterministic
|
||||
# capability gate, not a runtime fallback):
|
||||
|
||||
@@ -11,13 +11,18 @@ 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"
|
||||
"StableAudioConfig", "GlmImageDiTConfig", "LingBotWorld2CausalFastVideoConfig", "LingBotVideoConfig",
|
||||
"ZImageDiTConfig", "TrackWanVideoConfig", "CausalTrackWanVideoConfig"
|
||||
]
|
||||
|
||||
@@ -0,0 +1,70 @@
|
||||
# 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)
|
||||
@@ -0,0 +1,55 @@
|
||||
# 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"
|
||||
@@ -0,0 +1,61 @@
|
||||
# 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,6 +91,11 @@ 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
|
||||
|
||||
@@ -0,0 +1,60 @@
|
||||
# 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,6 +2,7 @@ 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
|
||||
@@ -9,6 +10,7 @@ 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
|
||||
@@ -17,5 +19,6 @@ __all__ = [
|
||||
"EncoderConfig", "TextEncoderConfig", "ImageEncoderConfig", "BaseEncoderOutput", "CLIPTextConfig",
|
||||
"CLIPVisionConfig", "WAN2_1ControlCLIPVisionConfig", "LlamaConfig", "T5Config", "T5LargeConfig", "Qwen2_5_VLConfig",
|
||||
"Reason1ArchConfig", "Reason1Config", "LTX2GemmaConfig", "SiglipVisionConfig", "StableAudioConditionerArchConfig",
|
||||
"StableAudioConditionerConfig", "T5GemmaEncoderConfig", "Qwen3TextConfig", "Mistral3TextConfig"
|
||||
"StableAudioConditionerConfig", "T5GemmaEncoderConfig", "Qwen3TextConfig", "Mistral3TextConfig",
|
||||
"LingBotWorld2UMT5ArchConfig", "LingBotWorld2UMT5Config", "LingBotVideoQwen3VLTextConfig"
|
||||
]
|
||||
|
||||
@@ -83,6 +83,7 @@ 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
|
||||
|
||||
@@ -0,0 +1,50 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Qwen3-VL text-only encoder configuration used by LingBot-Video."""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.encoders.base import TextEncoderArchConfig
|
||||
from fastvideo.configs.models.encoders.qwen3 import Qwen3TextArchConfig, Qwen3TextConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class LingBotVideoQwen3VLTextArchConfig(Qwen3TextArchConfig):
|
||||
"""Exact Qwen3-VL language-model architecture released with LingBot-Video."""
|
||||
|
||||
architectures: list[str] = field(default_factory=lambda: ["LingBotVideoQwen3VLTextModel"])
|
||||
vocab_size: int = 151936
|
||||
hidden_size: int = 2560
|
||||
intermediate_size: int = 9728
|
||||
num_hidden_layers: int = 36
|
||||
num_attention_heads: int = 32
|
||||
num_key_value_heads: int = 8
|
||||
max_position_embeddings: int = 262144
|
||||
rms_norm_eps: float = 1e-6
|
||||
rope_theta: float = 5000000.0
|
||||
rope_scaling: dict | None = None
|
||||
mrope_interleaved: bool = True
|
||||
mrope_section: tuple[int, int, int] = (24, 20, 20)
|
||||
bos_token_id: int = 151643
|
||||
eos_token_id: int = 151645
|
||||
pad_token_id: int = 151643
|
||||
text_len: int = 37698
|
||||
output_hidden_states: bool = True
|
||||
require_processor: bool = True
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Match the official processor call used by LingBotVideoPipeline."""
|
||||
self.tokenizer_kwargs = {
|
||||
"truncation": True,
|
||||
"max_length": self.text_len,
|
||||
"padding": "longest",
|
||||
"return_tensors": "pt",
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class LingBotVideoQwen3VLTextConfig(Qwen3TextConfig):
|
||||
"""FastVideo loader config for the LingBot-Video text-only Qwen3-VL path."""
|
||||
|
||||
arch_config: TextEncoderArchConfig = field(default_factory=LingBotVideoQwen3VLTextArchConfig)
|
||||
prefix: str = "language_model"
|
||||
is_chat_model: bool = False
|
||||
@@ -0,0 +1,40 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.encoders.base import (
|
||||
TextEncoderArchConfig,
|
||||
TextEncoderConfig,
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class LingBotWorld2UMT5ArchConfig(TextEncoderArchConfig):
|
||||
architectures: list[str] = field(default_factory=lambda: ["LingBotWorld2T5EncoderModel"])
|
||||
vocab_size: int = 256384
|
||||
dim: int = 4096
|
||||
dim_attn: int = 4096
|
||||
dim_ffn: int = 10240
|
||||
num_heads: int = 64
|
||||
num_layers: int = 24
|
||||
num_buckets: int = 32
|
||||
text_len: int = 512
|
||||
hidden_size: int = 4096
|
||||
dropout: float = 0.1
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
super().__post_init__()
|
||||
self.tokenizer_kwargs = {
|
||||
"padding": "max_length",
|
||||
"truncation": True,
|
||||
"max_length": self.text_len,
|
||||
"add_special_tokens": True,
|
||||
"return_attention_mask": True,
|
||||
"return_tensors": "pt",
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class LingBotWorld2UMT5Config(TextEncoderConfig):
|
||||
arch_config: TextEncoderArchConfig = field(default_factory=LingBotWorld2UMT5ArchConfig)
|
||||
|
||||
prefix: str = "text_encoder"
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user