Compare commits

...
Author SHA1 Message Date
H1yori233 9f94673c34 fix 2026-08-06 21:13:14 -07:00
H1yori233 c150889330 update server 2026-08-06 20:47:35 -07:00
H1yori233 7f783e30b7 add inference script 2026-08-06 18:18:36 -07:00
lpc0220 ec9db97431 [kernel] attention: sm_100a CUDA forward for block-causal-sink (Blackwell) (#1677) 2026-08-06 16:36:47 -07:00
H1yori233 fd236a06a7 add config 2026-07-24 20:29:03 -07:00
H1yori233 54fe0be488 app test 2026-07-24 20:04:31 -07:00
H1yori233 3a398a2bbd [bugfix]: load only exported role from DCP 2026-07-24 08:44:33 -07:00
H1yori233 4e76897f1f [bugfix]: restore lazy optimizer checkpoints 2026-07-23 12:19:14 -07:00
H1yori233 95d8e74d65 [bugfix]: flatten WanTrack validation logging 2026-07-23 11:37:43 -07:00
H1yori233 bcd0ad7dee [bugfix]: preserve WanTrack sampling dtype 2026-07-23 03:05:32 -07:00
H1yori233 9acbf92455 [bugfix]: support leading I2V frame in WanTrack sampling 2026-07-23 02:59:46 -07:00
H1yori233 ce4fa8ef2d [feat]: add bidirectional and causal WanTrack support 2026-07-23 02:23:10 -07:00
H1yori233 41dcd065ae Fix rank-local dataloader checkpoint resume 2026-07-22 21:29:48 -07:00
H1yori233 c0cb2228ac Fix causal EMA update and export semantics 2026-07-22 21:29:48 -07:00
H1yori233 2fa0f70ee8 Use 21 latent frames for OpenVid causal training 2026-07-22 21:29:47 -07:00
H1yori233 e172215880 Harden OpenVid launcher modes 2026-07-22 21:29:46 -07:00
H1yori233 23e3b325a4 harden-openvid-launch-safety 2026-07-22 21:29:46 -07:00
H1yori233 44d0990199 streaming-config-parser 2026-07-22 21:29:45 -07:00
H1yori233 59b6578906 feat: add OpenVid causal A12-A15 runs 2026-07-22 21:29:45 -07:00
H1yori233 fd8372d172 Add OpenVid causal A12-A15 experiment plan 2026-07-22 21:29:44 -07:00
H1yori233 33eb7ea13c feat(dataset): harden parquet streaming loader 2026-07-22 21:29:43 -07:00
H1yori233 051d21958d [feat]: 2026-07-22 21:29:43 -07:00
H1yori233 bf29e20bf5 [bugfix]: fix causal Wan batch and padding dimensions 2026-07-22 21:29:42 -07:00
H1yori233 291fa5d9a6 fix(train): save and export exact EMA checkpoints 2026-07-22 21:29:42 -07:00
H1yori233 954daf7fb3 Add 2026-07-22 21:29:40 -07:00
H1yori233 e05c04a2f6 [feat]: import block causal sink attention kernel 2026-07-22 21:29:05 -07:00
121 changed files with 17111 additions and 247 deletions
+35
View File
@@ -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.
+5
View File
@@ -0,0 +1,5 @@
"""Standalone causal WanTrack control prototype."""
from apps.wantrack_control.server import create_app
__all__ = ["create_app"]
+18
View File
@@ -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()
+236
View File
@@ -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())
+422
View 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()
+308
View File
@@ -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();
})();
+56
View File
@@ -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>
+57
View File
@@ -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
+188
View File
@@ -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.
@@ -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,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
1 id condition sink_size local_attn_size rope_cache_policy lane primary_contrast
2 A01 sink0_local21_absolute 0 21 absolute node0 A01-vs-A02: rope policy control
3 A02 sink0_local21_relative 0 21 relativistic node1 A01-vs-A02: rope policy control
4 A03 sink1_local21_relative 1 21 relativistic node0 A02-vs-A03: sink at local21
5 A04 sink0_local6_relative 0 6 relativistic node0 A02-vs-A04: local window at sink0
6 A05 sink1_local6_relative 1 6 relativistic node0 A04-vs-A05: sink at local6
7 A06 sink0_local12_relative 0 12 relativistic node1 A02-vs-A06: local window at sink0
8 A07 sink1_local12_relative 1 12 relativistic node1 A06-vs-A07: sink at local12
9 A08 sink3_local12_relative 3 12 relativistic node1 A06-vs-A07-vs-A08: sink size at local12
@@ -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
@@ -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
+25
View File
@@ -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(
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);
+9
View File
@@ -328,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)
@@ -343,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)
+4
View File
@@ -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)
+3
View File
@@ -129,6 +129,9 @@ 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
+17
View File
@@ -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
+6
View File
@@ -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
+2 -1
View File
@@ -11,6 +11,7 @@ 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
@@ -23,5 +24,5 @@ __all__ = [
"DreamXWorldARConfig", "CosmosVideoConfig", "Cosmos25VideoConfig", "FluxDiTConfig", "Flux2Config",
"LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig", "Kandinsky5VideoConfig", "MagiHumanVideoConfig",
"StableAudioConfig", "GlmImageDiTConfig", "LingBotWorld2CausalFastVideoConfig", "LingBotVideoConfig",
"ZImageDiTConfig"
"ZImageDiTConfig", "TrackWanVideoConfig", "CausalTrackWanVideoConfig"
]
+61
View File
@@ -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
+29
View File
@@ -0,0 +1,29 @@
# SPDX-License-Identifier: Apache-2.0
"""Pipeline configs for causal WanTrack Self-Forcing inference."""
from dataclasses import dataclass, field
from fastvideo.configs.models import DiTConfig
from fastvideo.configs.models.dits.trackwan import CausalTrackWanVideoConfig
from fastvideo.configs.pipelines.wan import WanI2V480PConfig
@dataclass
class CausalTrackWanSFI2VConfig(WanI2V480PConfig):
"""Causal TrackWan I2V distilled with Self-Forcing (DMD timesteps)."""
dit_config: DiTConfig = field(default_factory=CausalTrackWanVideoConfig)
is_causal: bool = True
flow_shift: float | None = 6.0
dmd_denoising_steps: list[int] | None = field(default_factory=lambda: [1000, 750, 500, 250])
warp_denoising_step: bool = True
context_noise: int = 0
def __post_init__(self) -> None:
super().__post_init__()
# Match the SF recipe used for Track-v0 / wantrack causal synth stage2.
arch = self.dit_config.arch_config
arch.local_attn_size = 6
arch.sink_size = 1
arch.rope_cache_policy = "relativistic"
arch.num_frames_per_block = 3
+39 -25
View File
@@ -2,13 +2,10 @@
from torchvision import transforms
from torchvision.transforms import Lambda
from fastvideo.dataset.parquet_dataset_map_style import (
build_parquet_map_style_dataloader)
from fastvideo.dataset.ltx2_precomputed_dataset import (
build_ltx2_precomputed_dataloader, LTX2PrecomputedDataset)
from fastvideo.dataset.parquet_dataset_map_style import build_parquet_map_style_dataloader
from fastvideo.dataset.ltx2_precomputed_dataset import build_ltx2_precomputed_dataloader, LTX2PrecomputedDataset
from fastvideo.dataset.preprocessing_datasets import VideoCaptionMergedDataset, TextDataset
from fastvideo.dataset.transform import (CenterCropResizeVideo, Normalize255,
TemporalRandomCrop)
from fastvideo.dataset.transform import CenterCropResizeVideo, Normalize255, TemporalRandomCrop
from fastvideo.dataset.validation_dataset import ValidationDataset
@@ -21,34 +18,51 @@ def getdataset(args) -> VideoCaptionMergedDataset:
resize = [
CenterCropResizeVideo((args.max_height, args.max_width)),
]
transform = transforms.Compose([
# Normalize255(),
*resize,
])
transform_topcrop = transforms.Compose([
Normalize255(),
*resize_topcrop,
norm_fun,
])
return VideoCaptionMergedDataset(data_merge_path=args.data_merge_path,
args=args,
transform=transform,
temporal_sample=temporal_sample,
transform_topcrop=transform_topcrop,
seed=args.seed)
transform = transforms.Compose(
[
# Normalize255(),
*resize,
]
)
transform_topcrop = transforms.Compose(
[
Normalize255(),
*resize_topcrop,
norm_fun,
]
)
return VideoCaptionMergedDataset(
data_merge_path=args.data_merge_path,
args=args,
transform=transform,
temporal_sample=temporal_sample,
transform_topcrop=transform_topcrop,
seed=args.seed,
)
def gettextdataset(args) -> TextDataset:
return TextDataset(data_merge_path=args.data_merge_path,
args=args,
seed=args.seed)
return TextDataset(data_merge_path=args.data_merge_path, args=args, seed=args.seed)
__all__ = [
"build_parquet_map_style_dataloader",
"build_parquet_streaming_style_dataloader",
"LatentsParquetStreamingDataset",
"build_ltx2_precomputed_dataloader",
"LTX2PrecomputedDataset",
"ValidationDataset",
"VideoCaptionMergedDataset",
"TextDataset",
]
def __getattr__(name: str):
if name in {
"LatentsParquetStreamingDataset",
"build_parquet_streaming_style_dataloader",
}:
from fastvideo.dataset import parquet_dataset_streaming_style
return getattr(parquet_dataset_streaming_style, name)
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
@@ -0,0 +1,59 @@
# SPDX-License-Identifier: Apache-2.0
"""Small benchmark for projected sequential row-group reads.
Example:
python -m fastvideo.dataset.benchmarks.benchmark_parquet_streaming_style \
--data-path /path/to/parquet --batches 4
"""
from __future__ import annotations
import argparse
import os
import tempfile
import time
from fastvideo.dataset.dataloader.schema import pyarrow_schema_t2v
from fastvideo.dataset.parquet_dataset_streaming_style import (
LatentsParquetStreamingDataset,
)
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--data-path", required=True)
parser.add_argument("--manifest-path", default="")
parser.add_argument("--batch-size", type=int, default=8)
parser.add_argument("--read-batch-size", type=int, default=4)
parser.add_argument("--batches", type=int, default=4)
args = parser.parse_args()
manifest_path = args.manifest_path or os.path.join(tempfile.gettempdir(), "fastvideo-streaming-benchmark.json")
dataset = LatentsParquetStreamingDataset(
args.data_path,
args.batch_size,
pyarrow_schema_t2v,
manifest_path,
num_workers=0,
read_batch_size=args.read_batch_size,
shuffle_row_groups=False,
)
started = time.perf_counter()
rows = 0
for index, batch in enumerate(dataset):
rows += len(batch["info_list"])
if index + 1 >= args.batches:
break
elapsed = time.perf_counter() - started
print(
{
"batches": index + 1,
"rows": rows,
"seconds": elapsed,
"rows_per_second": rows / elapsed,
"projected_columns": len(pyarrow_schema_t2v.names),
}
)
if __name__ == "__main__":
main()
+41 -1
View File
@@ -50,6 +50,46 @@ pyarrow_schema_i2v = pa.schema([
pa.field("fps", pa.float64()),
])
pyarrow_schema_i2v_track = pa.schema([
pa.field("id", pa.string()),
pa.field("vae_latent_bytes", pa.binary()),
pa.field("vae_latent_shape", pa.list_(pa.int64())),
pa.field("vae_latent_dtype", pa.string()),
pa.field("text_embedding_bytes", pa.binary()),
pa.field("text_embedding_shape", pa.list_(pa.int64())),
pa.field("text_embedding_dtype", pa.string()),
pa.field("first_frame_latent_bytes", pa.binary()),
pa.field("first_frame_latent_shape", pa.list_(pa.int64())),
pa.field("first_frame_latent_dtype", pa.string()),
pa.field("clip_feature_bytes", pa.binary()),
pa.field("clip_feature_shape", pa.list_(pa.int64())),
pa.field("clip_feature_dtype", pa.string()),
# MotionStream point tracks, normalized to [0, 1].
pa.field("track_points_bytes", pa.binary()),
pa.field("track_points_shape", pa.list_(pa.int64())), # [T, N, 2]
pa.field("track_points_dtype", pa.string()),
pa.field("track_visibility_bytes", pa.binary()),
pa.field("track_visibility_shape", pa.list_(pa.int64())), # [T, N]
pa.field("track_visibility_dtype", pa.string()),
# Optional preprocessing signals used by object-covering track sampling.
# New preprocessing writes dense fallback values when segmentation or
# informativeness weights are unavailable.
pa.field("object_ids_bytes", pa.binary()),
pa.field("object_ids_shape", pa.list_(pa.int64())), # [N]
pa.field("object_ids_dtype", pa.string()),
pa.field("track_weights_bytes", pa.binary()),
pa.field("track_weights_shape", pa.list_(pa.int64())), # [N]
pa.field("track_weights_dtype", pa.string()),
pa.field("file_name", pa.string()),
pa.field("caption", pa.string()),
pa.field("media_type", pa.string()),
pa.field("width", pa.int64()),
pa.field("height", pa.int64()),
pa.field("num_frames", pa.int64()),
pa.field("duration_sec", pa.float64()),
pa.field("fps", pa.float64()),
])
pyarrow_schema_t2v = pa.schema([
pa.field("id", pa.string()),
@@ -189,4 +229,4 @@ pyarrow_schema_matrixgame2_ode_trajectory = pa.schema([
pa.field("file_name", pa.string()),
pa.field("caption", pa.string()),
pa.field("media_type", pa.string()), # 'image' or 'video'
])
])
@@ -0,0 +1,576 @@
# SPDX-License-Identifier: Apache-2.0
"""Low-amplification Parquet streaming for large, shared T2V datasets.
The map-style loader remains the default. This opt-in loader keeps the source
dataset read-only, stores only JSON metadata in a caller-owned cache, projects
the requested T2V columns, and consumes each assigned row group sequentially.
"""
from __future__ import annotations
import hashlib
import json
import os
import random
import tempfile
from collections.abc import Iterator
from typing import Any
import pyarrow as pa
import pyarrow.parquet as pq
import torch.distributed as dist
from torch.utils.data import IterableDataset, get_worker_info
from torchdata.stateful_dataloader import StatefulDataLoader
from fastvideo.dataset.utils import collate_rows_from_parquet_schema
from fastvideo.distributed import get_sp_world_size, get_world_rank, get_world_size
from fastvideo.logger import init_logger
logger = init_logger(__name__)
_MANIFEST_VERSION = 2
_STATE_VERSION = 2
def _barrier() -> None:
if dist.is_available() and dist.is_initialized():
dist.barrier()
def _real(path: str | os.PathLike[str]) -> str:
return os.path.realpath(os.path.expanduser(os.fspath(path)))
def _assert_user_owned_manifest(dataset_root: str, manifest_path: str) -> None:
manifest_parent = _real(os.path.dirname(manifest_path) or ".")
if os.path.commonpath([dataset_root, manifest_parent]) == dataset_root:
raise ValueError(
f"streaming_manifest_path must be outside the read-only dataset root ({dataset_root}); got {manifest_path}"
)
def _manifest_fingerprint(row_groups: list[dict[str, Any]]) -> str:
payload = json.dumps(row_groups, sort_keys=True, separators=(",", ":")).encode("utf-8")
return hashlib.sha256(payload).hexdigest()
def _manifest_file_path(dataset_root: str, relative_path: str) -> str:
file_path = _real(os.path.join(dataset_root, relative_path))
if os.path.commonpath([dataset_root, file_path]) != dataset_root:
raise ValueError(f"Manifest file escapes dataset root: {relative_path!r}")
return file_path
def _manifest_matches_source(
dataset_root: str,
manifest: dict[str, Any],
) -> bool:
"""Validate cached metadata without reopening every Parquet footer."""
row_groups = manifest.get("row_groups")
if not isinstance(row_groups, list) or not row_groups:
return False
expected_sample_start = 0
stats: dict[str, os.stat_result] = {}
try:
for item in row_groups:
relative_path = str(item["file"])
rows = int(item["rows"])
row_group = int(item["row_group"])
sample_start = int(item["sample_start"])
if rows <= 0 or row_group < 0 or sample_start != expected_sample_start:
return False
expected_sample_start += rows
if relative_path not in stats:
file_path = _manifest_file_path(dataset_root, relative_path)
stats[relative_path] = os.stat(file_path)
stat = stats[relative_path]
if int(item["file_size"]) != stat.st_size or int(item["file_mtime_ns"]) != stat.st_mtime_ns:
return False
except (KeyError, OSError, TypeError, ValueError):
return False
return int(manifest.get("total_rows", -1)) == expected_sample_start and manifest.get(
"fingerprint"
) == _manifest_fingerprint(row_groups)
def _scan_manifest(dataset_root: str, columns: tuple[str, ...]) -> dict[str, Any]:
parquet_files: list[str] = []
for root, _, files in os.walk(dataset_root):
for name in sorted(files):
if name.endswith(".parquet"):
parquet_files.append(_real(os.path.join(root, name)))
parquet_files.sort()
if not parquet_files:
raise FileNotFoundError(f"No parquet files found under dataset path: {dataset_root}")
row_groups: list[dict[str, Any]] = []
sample_start = 0
for file_path in parquet_files:
file_stat = os.stat(file_path)
parquet_file = pq.ParquetFile(file_path)
missing = sorted(set(columns) - set(parquet_file.schema_arrow.names))
if missing:
raise ValueError(f"Parquet file {file_path} is missing projected T2V columns: {missing}")
relative_path = os.path.relpath(file_path, dataset_root)
for row_group_index in range(parquet_file.num_row_groups):
rows = parquet_file.metadata.row_group(row_group_index).num_rows
row_groups.append(
{
"file": relative_path,
"row_group": row_group_index,
"rows": rows,
"sample_start": sample_start,
"file_size": file_stat.st_size,
"file_mtime_ns": file_stat.st_mtime_ns,
}
)
sample_start += rows
return {
"version": _MANIFEST_VERSION,
"dataset_root": dataset_root,
"columns": list(columns),
"total_rows": sample_start,
"row_groups": row_groups,
"fingerprint": _manifest_fingerprint(row_groups),
}
def _write_json_atomic(path: str, payload: dict[str, Any]) -> None:
parent = os.path.dirname(path)
os.makedirs(parent, exist_ok=True)
fd, temporary_path = tempfile.mkstemp(dir=parent, prefix=f".{os.path.basename(path)}.", suffix=".tmp")
try:
with os.fdopen(fd, "w", encoding="utf-8") as handle:
json.dump(payload, handle, sort_keys=True)
handle.flush()
os.fsync(handle.fileno())
os.replace(temporary_path, path)
finally:
if os.path.exists(temporary_path):
os.unlink(temporary_path)
def load_or_create_streaming_manifest(
dataset_root: str, manifest_path: str, columns: tuple[str, ...]
) -> dict[str, Any]:
if os.name != "posix":
raise RuntimeError("The streaming Parquet loader currently requires POSIX file locks")
import fcntl
dataset_root = _real(dataset_root)
manifest_path = _real(manifest_path)
_assert_user_owned_manifest(dataset_root, manifest_path)
os.makedirs(os.path.dirname(manifest_path), exist_ok=True)
lock_path = f"{manifest_path}.lock"
with open(lock_path, "a", encoding="utf-8") as lock_handle:
fcntl.flock(lock_handle.fileno(), fcntl.LOCK_EX)
manifest: dict[str, Any] | None = None
if os.path.exists(manifest_path):
with open(manifest_path, encoding="utf-8") as handle:
candidate = json.load(handle)
if (
candidate.get("version") == _MANIFEST_VERSION
and candidate.get("dataset_root") == dataset_root
and candidate.get("columns") == list(columns)
and _manifest_matches_source(dataset_root, candidate)
):
manifest = candidate
if manifest is None:
logger.info("Building safe JSON Parquet manifest at %s", manifest_path)
manifest = _scan_manifest(dataset_root, columns)
_write_json_atomic(manifest_path, manifest)
fcntl.flock(lock_handle.fileno(), fcntl.LOCK_UN)
return manifest
def _contiguous_row_group_shards(row_groups: list[dict[str, Any]], num_shards: int) -> list[list[dict[str, Any]]]:
"""Split ordered row groups into row-balanced contiguous shards."""
if num_shards < 1:
raise ValueError("num_shards must be positive")
shards: list[list[dict[str, Any]]] = [[] for _ in range(num_shards)]
remaining_rows = sum(int(item["rows"]) for item in row_groups)
cursor = 0
for shard_index in range(num_shards):
shards_left = num_shards - shard_index
target = remaining_rows / shards_left if shards_left else 0
shard_rows = 0
while cursor < len(row_groups):
item = row_groups[cursor]
item_rows = int(item["rows"])
items_left = len(row_groups) - cursor
if shard_rows and shard_rows + item_rows > target and items_left >= shards_left - 1:
break
shards[shard_index].append(item)
shard_rows += item_rows
remaining_rows -= item_rows
cursor += 1
if cursor >= len(row_groups):
break
return shards
def reconstruct_streaming_dataset_state(
manifest: dict[str, Any],
*,
global_rank: int,
world_size: int,
sp_world_size: int,
num_workers: int,
batch_size: int,
read_batch_size: int,
seed: int,
shuffle_row_groups: bool,
yielded_samples: int,
worker_id: int = 0,
) -> dict[str, Any]:
"""Reconstruct an epoch-zero streaming cursor without reading Parquet.
This is intentionally strict and is meant for audited migration of legacy
checkpoints whose rank-local dataloader state was incorrectly stored in a
shared DCP key. It refuses epoch rollover because the exact wrapper state
at that boundary depends on whether the iterator has observed StopIteration.
"""
world_size = int(world_size)
sp_world_size = int(sp_world_size)
global_rank = int(global_rank)
num_workers = max(int(num_workers), 1)
batch_size = int(batch_size)
read_batch_size = int(read_batch_size)
yielded_samples = int(yielded_samples)
worker_id = int(worker_id)
if world_size < 1 or sp_world_size < 1 or world_size % sp_world_size:
raise ValueError("world_size must be positive and divisible by sp_world_size")
if not 0 <= global_rank < world_size:
raise ValueError("global_rank is outside the requested world")
if not 0 <= worker_id < num_workers:
raise ValueError("worker_id is outside the requested worker layout")
if batch_size < 1 or read_batch_size < 1:
raise ValueError("batch_size and read_batch_size must be positive")
if yielded_samples < 0 or yielded_samples % batch_size:
raise ValueError("yielded_samples must be a non-negative batch multiple")
row_groups = manifest.get("row_groups")
if not isinstance(row_groups, list) or not row_groups:
raise ValueError("Streaming manifest has no row groups")
num_sp_groups = world_size // sp_world_size
shards = _contiguous_row_group_shards(
row_groups,
num_sp_groups * num_workers,
)
shard_rows = [sum(int(item["rows"]) for item in shard) for shard in shards]
samples_per_worker = min(shard_rows) // batch_size * batch_size
if samples_per_worker == 0:
raise ValueError("Dataset is too small for the requested topology")
if yielded_samples >= samples_per_worker:
raise ValueError(
"Legacy cursor reconstruction only supports the first epoch and "
"requires yielded_samples < samples_per_worker"
)
sp_group_index = global_rank // sp_world_size
shard_index = sp_group_index * num_workers + worker_id
shard = list(shards[shard_index])
if shuffle_row_groups:
random.Random(int(seed)).shuffle(shard)
row_group_position = 0
row_offset = 0
remaining = yielded_samples
if remaining:
for position, item in enumerate(shard):
rows = int(item["rows"])
if remaining <= rows:
row_group_position = position
row_offset = remaining
break
remaining -= rows
else:
raise ValueError("yielded_samples exceeds the selected streaming shard")
return {
"version": _STATE_VERSION,
"manifest_fingerprint": str(manifest["fingerprint"]),
"topology": {
"global_rank": global_rank,
"world_size": world_size,
"sp_world_size": sp_world_size,
"sp_group_index": sp_group_index,
"num_workers": num_workers,
"batch_size": batch_size,
"read_batch_size": read_batch_size,
"seed": int(seed),
"shuffle_row_groups": bool(shuffle_row_groups),
"samples_per_worker": samples_per_worker,
},
"worker_id": worker_id,
"epoch": 0,
"row_group_position": row_group_position,
"row_offset": row_offset,
"yielded_samples": yielded_samples,
}
class LatentsParquetStreamingDataset(IterableDataset):
"""Sequential row-group loader with deterministic DP/SP sharding.
Ranks in the same sequence-parallel group receive identical samples.
Different data-parallel groups receive disjoint contiguous row-group
shards. Worker zero is a real worker shard, so ``num_workers=0`` is fully
supported and is the recommended shared-filesystem setting.
"""
def __init__(
self,
path: str,
batch_size: int,
parquet_schema: pa.Schema,
manifest_path: str,
cfg_rate: float = 0.0,
num_workers: int = 0,
drop_last: bool = True,
text_padding_length: int = 512,
seed: int = 42,
read_batch_size: int = 8,
shuffle_row_groups: bool = True,
) -> None:
super().__init__()
if not isinstance(path, str):
raise TypeError("The streaming loader currently accepts one dataset root")
if not drop_last:
raise ValueError("drop_last must be True for distributed streaming")
if batch_size < 1 or read_batch_size < 1:
raise ValueError("batch_size and read_batch_size must be positive")
if not manifest_path:
raise ValueError(
"streaming_manifest_path is required; place it in user-owned storage, never inside the shared dataset"
)
self.path = _real(path)
self.batch_size = int(batch_size)
self.parquet_schema = parquet_schema
self.columns = tuple(parquet_schema.names)
self.manifest_path = _real(manifest_path)
self.cfg_rate = float(cfg_rate)
self.num_workers = max(int(num_workers), 1)
self.text_padding_length = int(text_padding_length)
self.seed = int(seed)
self.read_batch_size = int(read_batch_size)
self.shuffle_row_groups = bool(shuffle_row_groups)
if dist.is_available() and dist.is_initialized():
self.global_rank = get_world_rank()
self.world_size = get_world_size()
self.sp_world_size = get_sp_world_size()
else:
# Keep standalone tests and small I/O benchmarks useful without
# requiring the full FastVideo distributed bootstrap.
self.global_rank = 0
self.world_size = 1
self.sp_world_size = 1
if self.world_size % self.sp_world_size:
raise ValueError("world_size must be divisible by sp_world_size")
self.num_sp_groups = self.world_size // self.sp_world_size
self.sp_group_index = self.global_rank // self.sp_world_size
if self.global_rank == 0:
load_or_create_streaming_manifest(self.path, self.manifest_path, self.columns)
_barrier()
with open(self.manifest_path, encoding="utf-8") as handle:
self.manifest = json.load(handle)
if self.manifest.get("dataset_root") != self.path:
raise ValueError("Streaming manifest belongs to a different dataset root")
if self.manifest.get("columns") != list(self.columns):
raise ValueError("Streaming manifest uses a different column projection")
self.manifest_fingerprint = str(self.manifest["fingerprint"])
num_shards = self.num_sp_groups * self.num_workers
self.shards = _contiguous_row_group_shards(self.manifest["row_groups"], num_shards)
shard_rows = [sum(int(item["rows"]) for item in shard) for shard in self.shards]
self.samples_per_worker = min(shard_rows) // self.batch_size * self.batch_size
if self.samples_per_worker == 0:
raise ValueError("Dataset is too small for the requested DP/SP/worker layout")
self._epoch = 0
self._row_group_position = 0
self._row_offset = 0
self._yielded_samples = 0
self._loaded_state: dict[str, Any] | None = None
logger.info(
"Streaming parquet: %d rows, %d row groups, %d samples per worker",
int(self.manifest["total_rows"]),
len(self.manifest["row_groups"]),
self.samples_per_worker,
)
def _worker_id(self) -> int:
info = get_worker_info()
return 0 if info is None else int(info.id)
def _worker_shard(self, worker_id: int) -> list[dict[str, Any]]:
shard_index = self.sp_group_index * self.num_workers + worker_id
shard = list(self.shards[shard_index])
if self.shuffle_row_groups:
random.Random(self.seed + self._epoch).shuffle(shard)
return shard
def __iter__(self) -> Iterator[dict[str, Any]]:
worker_id = self._worker_id()
if worker_id >= self.num_workers:
raise RuntimeError(f"Unexpected DataLoader worker id: {worker_id}")
if self._loaded_state is not None:
saved_worker = int(self._loaded_state.get("worker_id", worker_id))
if saved_worker != worker_id:
raise ValueError(f"Dataloader state belongs to worker {saved_worker}, not {worker_id}")
self._loaded_state = None
shard = self._worker_shard(worker_id)
rows: list[dict[str, Any]] = []
while self._yielded_samples < self.samples_per_worker and self._row_group_position < len(shard):
item = shard[self._row_group_position]
file_path = os.path.join(self.path, item["file"])
parquet_file = pq.ParquetFile(file_path)
consumed_in_group = 0
record_batches = parquet_file.iter_batches(
batch_size=self.read_batch_size,
row_groups=[int(item["row_group"])],
columns=list(self.columns),
use_threads=False,
)
for record_batch in record_batches:
record_count = record_batch.num_rows
record_end = consumed_in_group + record_count
if record_end <= self._row_offset:
consumed_in_group = record_end
continue
start = max(self._row_offset - consumed_in_group, 0)
for local_index, row in enumerate(record_batch.slice(start).to_pylist(), start=start):
global_offset = consumed_in_group + local_index
row["_sample_index"] = int(item["sample_start"]) + global_offset
rows.append(row)
self._row_offset = global_offset + 1
if len(rows) == self.batch_size:
self._yielded_samples += self.batch_size
yield collate_rows_from_parquet_schema(
rows,
self.parquet_schema,
self.text_padding_length,
cfg_rate=self.cfg_rate,
seed=self.seed,
)
rows = []
if self._yielded_samples >= self.samples_per_worker:
break
consumed_in_group = record_end
if self._yielded_samples >= self.samples_per_worker:
break
if self._yielded_samples >= self.samples_per_worker:
break
self._row_group_position += 1
self._row_offset = 0
if self._yielded_samples != self.samples_per_worker:
raise RuntimeError("Streaming shard ended before producing the balanced sample count")
self._epoch += 1
self._row_group_position = 0
self._row_offset = 0
self._yielded_samples = 0
def __len__(self) -> int:
return self.samples_per_worker * self.num_workers // self.batch_size
def state_dict(self) -> dict[str, Any]:
return {
"version": _STATE_VERSION,
"manifest_fingerprint": self.manifest_fingerprint,
"topology": self._resume_topology(),
"worker_id": self._worker_id(),
"epoch": self._epoch,
"row_group_position": self._row_group_position,
"row_offset": self._row_offset,
"yielded_samples": self._yielded_samples,
}
def load_state_dict(self, state_dict: dict[str, Any]) -> None:
if int(state_dict.get("version", -1)) != _STATE_VERSION:
raise ValueError("Cannot resume: unsupported streaming dataset state")
if state_dict.get("manifest_fingerprint") != self.manifest_fingerprint:
raise ValueError("Cannot resume: streaming dataset manifest changed")
if state_dict.get("topology") != self._resume_topology():
raise ValueError("Cannot resume: streaming loader topology or sampling config changed")
self._epoch = int(state_dict["epoch"])
self._row_group_position = int(state_dict["row_group_position"])
self._row_offset = int(state_dict["row_offset"])
self._yielded_samples = int(state_dict["yielded_samples"])
self._loaded_state = dict(state_dict)
def _resume_topology(self) -> dict[str, Any]:
return {
"global_rank": self.global_rank,
"world_size": self.world_size,
"sp_world_size": self.sp_world_size,
"sp_group_index": self.sp_group_index,
"num_workers": self.num_workers,
"batch_size": self.batch_size,
"read_batch_size": self.read_batch_size,
"seed": self.seed,
"shuffle_row_groups": self.shuffle_row_groups,
"samples_per_worker": self.samples_per_worker,
}
def get_validation_negative_prompt(self) -> tuple[Any, Any, str]:
first = self.manifest["row_groups"][0]
parquet_file = pq.ParquetFile(os.path.join(self.path, first["file"]))
first_batch = next(
parquet_file.iter_batches(
batch_size=1,
row_groups=[int(first["row_group"])],
columns=list(self.columns),
use_threads=False,
)
)
row = first_batch.to_pylist()[0]
batch = collate_rows_from_parquet_schema(
[row], self.parquet_schema, self.text_padding_length, cfg_rate=0.0, seed=self.seed
)
embedding = batch["text_embedding"]
mask = batch["text_attention_mask"]
prompt = batch["info_list"][0]["prompt"]
return embedding, mask, prompt
def build_parquet_streaming_style_dataloader(
path: str,
batch_size: int,
num_data_workers: int,
parquet_schema: pa.Schema,
manifest_path: str,
cfg_rate: float = 0.0,
drop_last: bool = True,
text_padding_length: int = 512,
seed: int = 42,
read_batch_size: int = 8,
shuffle_row_groups: bool = True,
) -> tuple[LatentsParquetStreamingDataset, StatefulDataLoader]:
dataset = LatentsParquetStreamingDataset(
path=path,
batch_size=batch_size,
parquet_schema=parquet_schema,
manifest_path=manifest_path,
cfg_rate=cfg_rate,
num_workers=num_data_workers,
drop_last=drop_last,
text_padding_length=text_padding_length,
seed=seed,
read_batch_size=read_batch_size,
shuffle_row_groups=shuffle_row_groups,
)
loader = StatefulDataLoader(
dataset,
batch_size=None,
num_workers=num_data_workers,
pin_memory=True,
persistent_workers=num_data_workers > 0,
)
return dataset, loader
+15 -1
View File
@@ -41,6 +41,7 @@ class PreprocessBatch:
sample_frame_index: list[int] | None = None
sample_num_frames: int | None = None
action_path: str | None = None
points_path: str | None = None
# Processed data
pixel_values: torch.Tensor | None = None
@@ -471,6 +472,8 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
item["path"] = opj(folder, item["path"])
if "action_path" in item and item["action_path"]:
item["action_path"] = opj(folder, item["action_path"])
if "points_path" in item and item["points_path"]:
item["points_path"] = opj(folder, item["points_path"])
return data_items
@@ -492,7 +495,8 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
resolution=item.get("resolution"),
fps=item.get("fps"),
duration=item.get("duration"),
action_path=item.get("action_path"))
action_path=item.get("action_path"),
points_path=item.get("points_path"))
# Apply filtering stages
if not self._apply_filter_stages(batch, filter_counts):
@@ -570,6 +574,16 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
# Add action_path
if batch.action_path:
result["action_path"] = batch.action_path
if batch.points_path:
result["points_path"] = batch.points_path
if batch.sample_frame_index is not None:
result["sample_frame_index"] = torch.tensor(batch.sample_frame_index,
dtype=torch.long)
if batch.resolution is not None:
if "width" in batch.resolution:
result["source_width"] = int(batch.resolution["width"])
if "height" in batch.resolution:
result["source_height"] = int(batch.resolution["height"])
return result
+15
View File
@@ -399,6 +399,9 @@ class VideoGenerator:
keyboard_cond: torch.Tensor | None = None,
grid_sizes: tuple[int, int, int] | list[int] | torch.Tensor
| None = None,
track_points: torch.Tensor | None = None,
track_visibility: torch.Tensor | None = None,
track_ids: torch.Tensor | None = None,
**kwargs,
) -> dict[str, Any] | list[dict[str, Any]]:
"""
@@ -447,6 +450,9 @@ class VideoGenerator:
mouse_cond=mouse_cond,
keyboard_cond=keyboard_cond,
grid_sizes=grid_sizes,
track_points=track_points,
track_visibility=track_visibility,
track_ids=track_ids,
legacy_kwargs=kwargs,
)
@@ -528,6 +534,9 @@ class VideoGenerator:
keyboard_cond: torch.Tensor | None = None,
grid_sizes: tuple[int, int, int] | list[int] | torch.Tensor
| None = None,
track_points: torch.Tensor | None = None,
track_visibility: torch.Tensor | None = None,
track_ids: torch.Tensor | None = None,
fastvideo_args: FastVideoArgs | None = None,
**kwargs,
) -> dict[str, Any] | list[np.ndarray] | list[dict[str, Any]]:
@@ -546,6 +555,12 @@ class VideoGenerator:
kwargs['keyboard_cond'] = keyboard_cond
if grid_sizes is not None:
kwargs['grid_sizes'] = grid_sizes
if track_points is not None:
kwargs['track_points'] = track_points
if track_visibility is not None:
kwargs['track_visibility'] = track_visibility
if track_ids is not None:
kwargs['track_ids'] = track_ids
extra_overrides: dict[str, Any] = {}
for _ek in _BATCH_EXTRA_PASSTHROUGH_KEYS:
@@ -0,0 +1,384 @@
# SPDX-License-Identifier: Apache-2.0
"""Training-time attention spec for causal video DiTs.
The rolling KV cache used at inference (``sink_size`` pinned frames plus a
trailing ``local_attn_size - sink_size`` frame window, optionally with
relativistic RoPE re-indexing) must be mirrored by the full-sequence attention
pattern used in teacher-forcing / causal-distillation training, otherwise the
student trains against a context layout it never sees when streaming.
This module is the single source of truth for that pattern:
- frame-level visibility rules shared by the FlexAttention mask builders, the
Triton kernel and the pure-PyTorch reference implementation;
- ``CausalTrainAttentionPlan``: a lightweight, cacheable description of one
training attention layout (``blockwise`` or ``teacher_forcing``);
- the relativistic sink RoPE correction. Within the rolling window,
relativistic re-indexing shifts every position by the same amount, so
absolute training RoPE already matches it exactly (RoPE phases depend only
on position differences). The only pairs whose phase differs are
query -> sink pairs once the window has scrolled past the sink: streaming
inference re-indexes the sink to a fixed small distance while absolute
training RoPE lets it recede. Matching that in training requires the query
to be rotated back by ``delta = max(0, block_end - local_attn_size)``
frames for sink columns only - a per-query-block key repositioning that
FlexAttention cannot express. The Triton / reference implementations apply
it exactly via ``delta_cos`` / ``delta_sin``.
"""
from __future__ import annotations
from dataclasses import dataclass, replace
import torch
from fastvideo.logger import init_logger
logger = init_logger(__name__)
def blockwise_frame_visible(
q_frame: int,
kv_frame: int,
*,
num_frame_per_block: int,
local_attn_size: int,
sink_size: int,
) -> bool:
"""Frame-level visibility for the blockwise-causal (single-sequence) layout.
``q_frame`` attends everything up to the end of its own block, restricted
(when ``local_attn_size != -1``) to the sink frames plus the trailing
``local_attn_size - sink_size`` frame window - the same set a streaming
step reads from the rolled KV cache.
"""
block_end = (q_frame // num_frame_per_block + 1) * num_frame_per_block
if kv_frame >= block_end:
return False
if local_attn_size == -1:
return True
rolling = max(0, int(local_attn_size) - int(sink_size))
if kv_frame >= block_end - rolling:
return True
return kv_frame < sink_size
def teacher_forcing_frame_visible(
q_half: int,
q_frame: int,
kv_half: int,
kv_frame: int,
*,
num_frame_per_block: int,
local_attn_size: int,
sink_size: int,
) -> bool:
"""Frame-level visibility for the ``[clean | noisy]`` teacher-forcing layout.
``*_half`` is 0 for the clean half and 1 for the noisy half. Clean rows are
blockwise-causal over the clean half (they mirror the context write-back
forward at inference); a noisy row attends its own noisy block plus the
clean context of strictly-previous blocks. With a rolling window the clean
context is restricted to the sink plus the trailing window, whose budget
includes the noisy block itself (it occupies the newest ``num_frame_per_block``
cache slots while being denoised at inference).
"""
block = q_frame // num_frame_per_block
block_end = (block + 1) * num_frame_per_block
if q_half == 0:
if kv_half != 0:
return False
return blockwise_frame_visible(
q_frame,
kv_frame,
num_frame_per_block=num_frame_per_block,
local_attn_size=local_attn_size,
sink_size=sink_size,
)
if kv_half == 1:
return kv_frame // num_frame_per_block == block
context_end = block_end - num_frame_per_block
if kv_frame >= context_end:
return False
if local_attn_size == -1:
return True
rolling = max(0, int(local_attn_size) - int(sink_size))
if kv_frame >= block_end - rolling:
return True
return kv_frame < sink_size
def sink_delta_frames(
num_frames: int,
*,
num_frame_per_block: int,
local_attn_size: int,
) -> list[int]:
"""Relativistic sink phase offset per query frame-block, in frames.
``delta[b] = max(0, block_end - local_attn_size)``: once the rolling
window has scrolled past the sink, streaming inference re-indexes the sink
to sit ``local_attn_size`` frames behind the newest block end instead of
at its absolute distance.
"""
num_blocks = (num_frames + num_frame_per_block - 1) // num_frame_per_block
return [
max(0, (b + 1) * num_frame_per_block - int(local_attn_size))
for b in range(num_blocks)
]
def build_sink_delta_tables(
*,
num_frames: int,
num_frame_per_block: int,
local_attn_size: int,
sink_size: int,
hidden_size: int,
num_attention_heads: int,
rope_theta: float = 10000,
dtype: torch.dtype = torch.float64,
device: torch.device | str = "cpu",
) -> tuple[torch.Tensor, torch.Tensor] | None:
"""Per-query-block rotation tables for the relativistic sink correction.
Returns ``(delta_cos, delta_sin)`` of shape ``[num_blocks, head_dim]`` in
float32, or ``None`` when every delta is zero (window never scrolls past
the sink) or the correction does not apply. ``rope(x, pos + delta)``
equals the rope rotation at temporal position ``delta`` (h = w = 0, so the
spatial dims get an identity rotation) applied to ``rope(x, pos)``;
sampling one rope row per delta reuses the model's exact frequencies.
"""
if local_attn_size == -1 or sink_size <= 0:
return None
deltas = sink_delta_frames(
num_frames,
num_frame_per_block=num_frame_per_block,
local_attn_size=local_attn_size,
)
max_delta = max(deltas)
if max_delta == 0:
return None
from fastvideo.layers.rotary_embedding import get_rotary_pos_embed
d = hidden_size // num_attention_heads
rope_dim_list = [d - 4 * (d // 6), 2 * (d // 6), 2 * (d // 6)]
cos, sin = get_rotary_pos_embed(
(max_delta + 1, 1, 1),
hidden_size,
num_attention_heads,
rope_dim_list,
rope_theta=rope_theta,
dtype=dtype,
)
idx = torch.tensor(deltas, dtype=torch.long)
delta_cos = cos.index_select(0, idx).to(device=device, dtype=torch.float32)
delta_sin = sin.index_select(0, idx).to(device=device, dtype=torch.float32)
return delta_cos, delta_sin
@dataclass
class CausalTrainAttentionPlan:
"""Cacheable description of one full-sequence training attention layout."""
kind: str # "blockwise" | "teacher_forcing"
impl: str # "triton" | "reference"
num_frames: int # latent frames (per half for teacher_forcing)
frame_seqlen: int
num_frame_per_block: int
local_attn_size: int
sink_size: int
sm_scale: float
# Relativistic sink correction ([num_blocks, head_dim] float32 each);
# None means the correction is a no-op and absolute RoPE is exact.
delta_cos: torch.Tensor | None = None
delta_sin: torch.Tensor | None = None
def __post_init__(self) -> None:
if self.kind not in ("blockwise", "teacher_forcing"):
raise ValueError(f"Unknown plan kind: {self.kind!r}")
if self.impl not in ("triton", "reference"):
raise ValueError(f"Unknown plan impl: {self.impl!r}")
if (self.delta_cos is None) != (self.delta_sin is None):
raise ValueError("delta_cos and delta_sin must be set together")
@property
def seq_len(self) -> int:
halves = 2 if self.kind == "teacher_forcing" else 1
return self.num_frames * self.frame_seqlen * halves
def without_delta(self) -> CausalTrainAttentionPlan:
"""Mask-only copy for attention streams that do not carry RoPE (PRoPE)."""
if self.delta_cos is None:
return self
return replace(self, delta_cos=None, delta_sin=None)
def _rotate_half_gptj(x: torch.Tensor) -> torch.Tensor:
x_real, x_imag = x.reshape(*x.shape[:-1], -1, 2).unbind(-1)
return torch.stack([-x_imag, x_real], dim=-1).flatten(-2)
def apply_sink_delta_to_query(
q: torch.Tensor,
delta_cos: torch.Tensor,
delta_sin: torch.Tensor,
) -> torch.Tensor:
"""Rotate roped queries back by their block's sink delta (``R(-delta)``).
``score = <rope(k, f + delta), rope(q, t)> = <rope(k, f), R(-delta) rope(q, t)>``,
so the per-query-block sink repositioning is applied on the query side.
Args:
q: ``[..., L, D]`` roped queries.
delta_cos / delta_sin: ``[L, D]`` rows already gathered per query token.
"""
return (q.float() * delta_cos - _rotate_half_gptj(q.float()) * delta_sin).type_as(q)
def _plan_row_geometry(
plan: CausalTrainAttentionPlan,
device: torch.device,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""Per-token (half, frame, block) indices for a plan's sequence layout."""
tokens_per_half = plan.num_frames * plan.frame_seqlen
idx = torch.arange(plan.seq_len, device=device)
half = (idx // tokens_per_half).clamp(max=1)
frame = (idx % tokens_per_half) // plan.frame_seqlen
block = frame // plan.num_frame_per_block
return half, frame, block
def build_plan_visibility(
plan: CausalTrainAttentionPlan,
device: torch.device | str = "cpu",
) -> torch.Tensor:
"""Dense boolean visibility matrix ``[seq_len, seq_len]`` for a plan."""
device = torch.device(device)
half, frame, block = _plan_row_geometry(plan, device)
q_half, kv_half = half[:, None], half[None, :]
q_frame, kv_frame = frame[:, None], frame[None, :]
q_block, kv_block = block[:, None], block[None, :]
block_end = (q_block + 1) * plan.num_frame_per_block
if plan.local_attn_size == -1:
in_window = torch.ones_like(kv_frame < 0)
in_sink = torch.zeros_like(in_window)
else:
rolling = max(0, plan.local_attn_size - plan.sink_size)
in_window = kv_frame >= (block_end - rolling)
in_sink = kv_frame < plan.sink_size
windowed = in_window | in_sink
if plan.kind == "blockwise":
return (kv_frame < block_end) & windowed
clean_rows = (q_half == 0) & (kv_half == 0) & (kv_frame < block_end) & windowed
noisy_self = (q_half == 1) & (kv_half == 1) & (kv_block == q_block)
context_end = block_end - plan.num_frame_per_block
noisy_context = ((q_half == 1) & (kv_half == 0) &
(kv_frame < context_end) & windowed)
return clean_rows | noisy_self | noisy_context
def reference_causal_train_attention(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
plan: CausalTrainAttentionPlan,
) -> torch.Tensor:
"""O(L^2) float32 reference for the plan semantics.
Args:
q / k / v: ``[B, H, L, D]``.
Returns:
``[B, H, L, D]`` in ``v.dtype``.
"""
seq_len = q.shape[-2]
if seq_len != plan.seq_len:
raise ValueError(f"Plan expects seq_len={plan.seq_len}, got {seq_len}")
device = q.device
visible = build_plan_visibility(plan, device=device)
scores = torch.einsum("bhld,bhmd->bhlm", q.float(), k.float()) * plan.sm_scale
if plan.delta_cos is not None:
_, _, block = _plan_row_geometry(plan, device)
delta_cos = plan.delta_cos.to(device=device)[block]
delta_sin = plan.delta_sin.to(device=device)[block]
q_delta = apply_sink_delta_to_query(q.float(), delta_cos, delta_sin)
sink_scores = torch.einsum("bhld,bhmd->bhlm", q_delta, k.float()) * plan.sm_scale
sink_tokens = plan.sink_size * plan.frame_seqlen
sink_col = torch.arange(seq_len, device=device)[None, None, None, :] < sink_tokens
scores = torch.where(sink_col, sink_scores, scores)
scores = scores.masked_fill(~visible, float("-inf"))
probs = torch.softmax(scores, dim=-1)
out = torch.einsum("bhlm,bhmd->bhld", probs, v.float())
return out.type_as(v)
def run_causal_train_attention(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
plan: CausalTrainAttentionPlan,
) -> torch.Tensor:
"""Dispatch a plan to its implementation. q/k/v: ``[B, H, L, D]``."""
if plan.impl == "triton":
from fastvideo.attention.kernels.block_causal_sink import (
block_causal_sink_attention, )
return block_causal_sink_attention(q, k, v, plan)
return reference_causal_train_attention(q, k, v, plan)
def validate_causal_attention_geometry(
*,
local_attn_size: int,
sink_size: int,
num_frame_per_block: int,
where: str,
) -> None:
"""Shared validation for the sink + rolling-window geometry."""
if sink_size < 0:
raise ValueError(f"{where}: sink_size must be non-negative, got {sink_size}")
if local_attn_size == -1 or sink_size == 0:
return
if sink_size + num_frame_per_block > local_attn_size:
raise ValueError(
f"{where}: local_attn_size ({local_attn_size}) must cover the sink "
f"plus at least one frame block (sink_size={sink_size} + "
f"num_frame_per_block={num_frame_per_block}); the rolling KV cache "
"otherwise cannot hold a new block after pinning the sink")
def approx_relativistic_delta_max(
*,
num_frames: int,
num_frame_per_block: int,
local_attn_size: int,
) -> int:
"""Largest sink phase offset the training sequence would need."""
if local_attn_size == -1:
return 0
deltas = sink_delta_frames(
num_frames,
num_frame_per_block=num_frame_per_block,
local_attn_size=local_attn_size,
)
return max(deltas) if deltas else 0
__all__ = [
"CausalTrainAttentionPlan",
"apply_sink_delta_to_query",
"approx_relativistic_delta_max",
"blockwise_frame_visible",
"build_plan_visibility",
"build_sink_delta_tables",
"reference_causal_train_attention",
"run_causal_train_attention",
"sink_delta_frames",
"teacher_forcing_frame_visible",
"validate_causal_attention_geometry",
]
+242 -66
View File
@@ -31,9 +31,16 @@ from fastvideo.layers.rotary_embedding import (_apply_rotary_emb,
get_rotary_pos_embed)
from fastvideo.layers.visual_embedding import (PatchEmbed)
from fastvideo.logger import init_logger
from fastvideo.models.dits._causal_train_attention import (
CausalTrainAttentionPlan,
approx_relativistic_delta_max,
build_sink_delta_tables,
run_causal_train_attention,
validate_causal_attention_geometry,
)
from fastvideo.models.dits._relative_rope import relativistic_window_offsets
from fastvideo.models.dits.base import BaseDiT
from fastvideo.models.dits.wanvideo import WanT2VCrossAttention, WanTimeTextImageEmbedding
from fastvideo.models.dits.wanvideo import WanI2VCrossAttention, WanT2VCrossAttention, WanTimeTextImageEmbedding
from fastvideo.platforms import AttentionBackendEnum, current_platform
logger = init_logger(__name__)
@@ -41,6 +48,25 @@ logger = init_logger(__name__)
GLOBAL_ATTN_COMPAT_MAX_LATENT_FRAMES = 21
def _blockwise_causal_attention_visible(q_idx,
kv_idx,
block_end,
frame_seqlen: int,
local_attn_size: int,
sink_size: int = 0):
"""Token-level visibility shared by blockwise-causal training attention."""
visible_before_block_end = kv_idx < block_end
if local_attn_size == -1:
return visible_before_block_end | (q_idx == kv_idx)
rolling_size = max(0, int(local_attn_size) - int(sink_size))
visible_in_window = kv_idx >= (block_end - rolling_size * frame_seqlen)
visible_in_sink = kv_idx < 0
if int(sink_size) > 0:
visible_in_sink = kv_idx < int(sink_size) * frame_seqlen
return (visible_before_block_end & (visible_in_sink | visible_in_window)) | (q_idx == kv_idx)
class CausalWanSelfAttention(nn.Module):
def __init__(self,
@@ -79,7 +105,7 @@ class CausalWanSelfAttention(nn.Module):
k: torch.Tensor,
v: torch.Tensor,
freqs_cis: tuple[torch.Tensor, torch.Tensor],
block_mask: BlockMask,
block_mask: BlockMask | CausalTrainAttentionPlan,
kv_cache: dict | None = None,
current_start: int = 0,
cache_start: int | None = None,
@@ -104,6 +130,18 @@ class CausalWanSelfAttention(nn.Module):
roped_key = _apply_rotary_emb(k, cos, sin, is_neox_style=False).type_as(v)
if kv_cache is None:
if isinstance(block_mask, CausalTrainAttentionPlan):
# Fused sink + rolling-window blockwise attention (Triton or
# reference); no 128-padding needed and the relativistic sink
# RoPE correction is applied exactly.
x = run_causal_train_attention(
roped_query.transpose(2, 1),
roped_key.transpose(2, 1),
v.transpose(2, 1),
block_mask,
).transpose(2, 1)
return x
# Padding for flex attention
padded_length = math.ceil(q.shape[1] / 128) * 128 - q.shape[1]
padded_roped_query = torch.cat(
@@ -130,7 +168,7 @@ class CausalWanSelfAttention(nn.Module):
key=padded_roped_key.transpose(2, 1),
value=padded_v.transpose(2, 1),
block_mask=block_mask
)[:, :, :-padded_length].transpose(2, 1)
)[:, :, :q.shape[1]].transpose(2, 1)
else:
current_end = current_start + q.shape[1]
sink_tokens = self.sink_size * frame_seqlen
@@ -185,8 +223,19 @@ class CausalWanSelfAttention(nn.Module):
# logger.info("kv_cache['k'] is in comp graph: %s", kv_cache["k"].requires_grad or kv_cache["k"].grad_fn is not None)
kv_cache["k"][:, local_start_index:local_end_index] = stored_key
kv_cache["v"][:, local_start_index:local_end_index] = v
key_window = kv_cache["k"][:, max(0, local_end_index - max_attention_size):local_end_index]
value_window = kv_cache["v"][:, max(0, local_end_index - max_attention_size):local_end_index]
window_start = max(0, local_end_index - max_attention_size)
if sink_tokens > 0 and window_start > 0:
if sink_tokens >= max_attention_size:
raise ValueError(f"sink_size tokens ({sink_tokens}) must be smaller than "
f"the attention budget ({max_attention_size})")
local_start = local_end_index - (max_attention_size - sink_tokens)
key_window = torch.cat(
[kv_cache["k"][:, :sink_tokens], kv_cache["k"][:, local_start:local_end_index]], dim=1)
value_window = torch.cat(
[kv_cache["v"][:, :sink_tokens], kv_cache["v"][:, local_start:local_end_index]], dim=1)
else:
key_window = kv_cache["k"][:, window_start:local_end_index]
value_window = kv_cache["v"][:, window_start:local_end_index]
if relativistic:
window_len, query_lo, query_hi = relativistic_window_offsets(
local_end_index, num_new_tokens, max_attention_size)
@@ -262,12 +311,25 @@ class CausalWanTransformerBlock(nn.Module):
elementwise_affine=True,
dtype=torch.float32)
# 2. Cross-attention
# Only T2V for now
self.attn2 = WanT2VCrossAttention(dim,
num_heads,
qk_norm=qk_norm,
eps=eps)
# 2. Cross-attention. I2V checkpoints prepend CLIP image tokens and
# carry a second pair of projections, which must be preserved during
# causal rollouts.
if added_kv_proj_dim is not None:
self.attn2 = WanI2VCrossAttention(
dim,
num_heads,
qk_norm=qk_norm,
eps=eps,
prefix=f"{prefix}.attn2",
)
else:
self.attn2 = WanT2VCrossAttention(
dim,
num_heads,
qk_norm=qk_norm,
eps=eps,
prefix=f"{prefix}.attn2",
)
self.cross_attn_residual_norm = ScaleResidualLayerNormScaleShift(
dim,
norm_type="layer",
@@ -287,7 +349,7 @@ class CausalWanTransformerBlock(nn.Module):
encoder_hidden_states: torch.Tensor,
temb: torch.Tensor,
freqs_cis: tuple[torch.Tensor, torch.Tensor],
block_mask: BlockMask,
block_mask: BlockMask | CausalTrainAttentionPlan,
kv_cache: dict | None = None,
crossattn_cache: dict | None = None,
current_start: int = 0,
@@ -389,7 +451,14 @@ class CausalWanTransformer3DModel(BaseDiT):
self.patch_size = config.patch_size
self.text_len = config.text_len
self.local_attn_size = config.local_attn_size
self.sink_size = config.sink_size
self.rope_cache_policy = config.arch_config.rope_cache_policy
self.causal_train_attention = getattr(config.arch_config,
"causal_train_attention", "flex")
if self.causal_train_attention not in ("flex", "triton", "reference"):
raise ValueError(
"causal_train_attention must be one of 'flex', 'triton', "
f"'reference'; got {self.causal_train_attention!r}")
# 1. Patch & position embedding
self.patch_embedding = PatchEmbed(in_chans=config.in_channels,
@@ -441,13 +510,21 @@ class CausalWanTransformer3DModel(BaseDiT):
self.num_frame_per_block = config.arch_config.num_frames_per_block
assert self.num_frame_per_block <= 3
self.independent_first_frame = False
validate_causal_attention_geometry(
local_attn_size=self.local_attn_size,
sink_size=self.sink_size,
num_frame_per_block=self.num_frame_per_block,
where=type(self).__name__,
)
self._relativistic_train_rope_warned = False
self.__post_init__()
@staticmethod
def _prepare_blockwise_causal_attn_mask(
device: torch.device | str, num_frames: int = 21,
frame_seqlen: int = 1560, num_frame_per_block=1, local_attn_size=-1
frame_seqlen: int = 1560, num_frame_per_block=1, local_attn_size=-1,
sink_size: int = 0
) -> BlockMask:
"""
we will divide the token sequence into the following format
@@ -475,10 +552,14 @@ class CausalWanTransformer3DModel(BaseDiT):
frame_seqlen * num_frame_per_block
def attention_mask(b, h, q_idx, kv_idx):
if local_attn_size == -1:
return (kv_idx < ends[q_idx]) | (q_idx == kv_idx)
else:
return ((kv_idx < ends[q_idx]) & (kv_idx >= (ends[q_idx] - local_attn_size * frame_seqlen))) | (q_idx == kv_idx)
return _blockwise_causal_attention_visible(
q_idx,
kv_idx,
ends[q_idx],
frame_seqlen,
local_attn_size,
sink_size,
)
# return ((kv_idx < total_length) & (q_idx < total_length)) | (q_idx == kv_idx) # bidirectional mask
block_mask = create_block_mask(attention_mask, B=None, H=None, Q_LEN=total_length + padded_length,
@@ -504,24 +585,23 @@ class CausalWanTransformer3DModel(BaseDiT):
@staticmethod
def _prepare_teacher_forcing_mask(
device: torch.device | str, num_frames: int = 21,
frame_seqlen: int = 1560, num_frame_per_block=1, local_attn_size=-1
frame_seqlen: int = 1560, num_frame_per_block=1,
local_attn_size: int = -1, sink_size: int = 0
) -> BlockMask:
"""Attention mask for the teacher-forcing ``[clean | noisy]`` sequence.
A noisy token attends to its own block plus the clean context of all
strictly previous blocks; clean tokens are block-wise causal.
A noisy token attends to its own block plus the clean context of
strictly previous blocks; clean tokens are block-wise causal. With
``local_attn_size != -1`` the visible context mirrors the rolling KV
cache: the ``sink_size`` leading frames plus the trailing
``local_attn_size - sink_size`` frame window, whose budget includes
the noisy block itself.
"""
if local_attn_size != -1:
raise NotImplementedError(
f"Teacher forcing ignores local_attn_size={local_attn_size}: "
"unlike the block-wise causal mask, this mask always attends "
"to the full clean context. Windowed teacher forcing is not "
"implemented; use local_attn_size=-1 for teacher-forcing "
"training.")
total_length = num_frames * frame_seqlen * 2
padded_length = math.ceil(total_length / 128) * 128 - total_length
clean_ends = num_frames * frame_seqlen
context_starts = torch.zeros(total_length + padded_length, device=device, dtype=torch.long)
context_ends = torch.zeros(total_length + padded_length, device=device, dtype=torch.long)
noise_context_starts = torch.zeros(total_length + padded_length, device=device, dtype=torch.long)
noise_context_ends = torch.zeros(total_length + padded_length, device=device, dtype=torch.long)
@@ -529,12 +609,18 @@ class CausalWanTransformer3DModel(BaseDiT):
noise_noise_ends = torch.zeros(total_length + padded_length, device=device, dtype=torch.long)
attention_block_size = frame_seqlen * num_frame_per_block
rolling_tokens = (max(0, int(local_attn_size) - int(sink_size)) *
frame_seqlen if local_attn_size != -1 else 0)
sink_tokens = int(sink_size) * frame_seqlen if local_attn_size != -1 else 0
frame_indices = torch.arange(
start=0, end=num_frames * frame_seqlen,
step=attention_block_size, device=device, dtype=torch.long
)
for start in frame_indices:
context_ends[start:start + attention_block_size] = start + attention_block_size
end = start + attention_block_size
context_ends[start:end] = end
if local_attn_size != -1:
context_starts[start:end] = max(0, int(end) - rolling_tokens)
noisy_image_start_list = torch.arange(
num_frames * frame_seqlen, total_length,
@@ -545,11 +631,17 @@ class CausalWanTransformer3DModel(BaseDiT):
noise_noise_starts[start:end] = start
noise_noise_ends[start:end] = end
noise_context_ends[start:end] = block_index * attention_block_size
if local_attn_size != -1:
block_end_tokens = (block_index + 1) * attention_block_size
noise_context_starts[start:end] = max(0, block_end_tokens - rolling_tokens)
def attention_mask(b, h, q_idx, kv_idx):
clean_mask = (q_idx < clean_ends) & (kv_idx < context_ends[q_idx])
in_sink = kv_idx < sink_tokens
clean_mask = (q_idx < clean_ends) & (kv_idx < context_ends[q_idx]) & (
(kv_idx >= context_starts[q_idx]) | in_sink)
c1 = (kv_idx < noise_noise_ends[q_idx]) & (kv_idx >= noise_noise_starts[q_idx])
c2 = (kv_idx < noise_context_ends[q_idx]) & (kv_idx >= noise_context_starts[q_idx])
c2 = (kv_idx < noise_context_ends[q_idx]) & (
(kv_idx >= noise_context_starts[q_idx]) | in_sink)
noise_mask = (q_idx >= clean_ends) & (c1 | c2)
eye_mask = q_idx == kv_idx
return eye_mask | clean_mask | noise_mask
@@ -565,6 +657,96 @@ class CausalWanTransformer3DModel(BaseDiT):
return block_mask
def _get_train_attention_spec(
self,
*,
device: torch.device | str,
num_frames: int,
frame_seqlen: int,
teacher_forcing: bool,
) -> BlockMask | CausalTrainAttentionPlan:
"""Build and cache the full-sequence causal training attention spec."""
attr = "teacher_forcing_block_mask" if teacher_forcing else "block_mask"
spec = getattr(self, attr)
if spec is not None:
return spec
relativistic_sinks = (
self.rope_cache_policy == "relativistic" and self.sink_size > 0
and approx_relativistic_delta_max(
num_frames=num_frames,
num_frame_per_block=self.num_frame_per_block,
local_attn_size=self.local_attn_size) > 0)
if self.causal_train_attention == "flex":
if relativistic_sinks and not self._relativistic_train_rope_warned:
logger.warning(
"rope_cache_policy='relativistic' with sink_size=%d: the "
"FlexAttention training path keeps absolute sink RoPE, so "
"query->sink phases differ from the re-indexed streaming "
"cache once the rolling window scrolls past the sink. Set "
"pipeline.dit_config.causal_train_attention: triton for "
"the exact training-time correction.", self.sink_size)
self._relativistic_train_rope_warned = True
if teacher_forcing:
spec = self._prepare_teacher_forcing_mask(
device=device,
num_frames=num_frames,
frame_seqlen=frame_seqlen,
num_frame_per_block=self.num_frame_per_block,
local_attn_size=self.local_attn_size,
sink_size=self.sink_size,
)
else:
spec = self._prepare_blockwise_causal_attn_mask(
device=device,
num_frames=num_frames,
frame_seqlen=frame_seqlen,
num_frame_per_block=self.num_frame_per_block,
local_attn_size=self.local_attn_size,
sink_size=self.sink_size,
)
else:
delta_cos = delta_sin = None
if relativistic_sinks:
tables = build_sink_delta_tables(
num_frames=num_frames,
num_frame_per_block=self.num_frame_per_block,
local_attn_size=self.local_attn_size,
sink_size=self.sink_size,
hidden_size=self.hidden_size,
num_attention_heads=self.num_attention_heads,
dtype=torch.float32 if current_platform.is_mps() else torch.float64,
device=device,
)
if tables is not None:
delta_cos, delta_sin = tables
spec = CausalTrainAttentionPlan(
kind="teacher_forcing" if teacher_forcing else "blockwise",
impl=self.causal_train_attention,
num_frames=num_frames,
frame_seqlen=frame_seqlen,
num_frame_per_block=self.num_frame_per_block,
local_attn_size=self.local_attn_size,
sink_size=self.sink_size,
sm_scale=1.0 / math.sqrt(self.hidden_size // self.num_attention_heads),
delta_cos=delta_cos,
delta_sin=delta_sin,
)
if not dist.is_initialized() or dist.get_rank() == 0:
logger.info(
"cache a %s causal train-attention plan (impl=%s, "
"local_attn_size=%d, sink_size=%d, relativistic_sinks=%s)",
spec.kind,
spec.impl,
self.local_attn_size,
self.sink_size,
relativistic_sinks,
)
setattr(self, attr, spec)
return spec
def _forward_inference(
self,
hidden_states: torch.Tensor,
@@ -588,11 +770,10 @@ class CausalWanTransformer3DModel(BaseDiT):
orig_dtype = hidden_states.dtype
if not isinstance(encoder_hidden_states, torch.Tensor):
encoder_hidden_states = encoder_hidden_states[0]
if isinstance(encoder_hidden_states_image,
list) and len(encoder_hidden_states_image) > 0:
encoder_hidden_states_image = encoder_hidden_states_image[0]
else:
encoder_hidden_states_image = None
if isinstance(encoder_hidden_states_image, list):
encoder_hidden_states_image = (
encoder_hidden_states_image[0]
if encoder_hidden_states_image else None)
batch_size, num_channels, num_frames, height, width = hidden_states.shape
p_t, p_h, p_w = self.patch_size
@@ -628,11 +809,16 @@ class CausalWanTransformer3DModel(BaseDiT):
freqs_sin) if freqs_cos is not None else None
hidden_states = self.patch_embedding(hidden_states)
grid_sizes = torch.stack(
[torch.tensor(hidden_states[0].shape[1:], dtype=torch.long)])
grid_sizes = torch.tensor(
hidden_states.shape[2:], dtype=torch.long).unsqueeze(0).repeat(
batch_size, 1)
hidden_states = hidden_states.flatten(2).transpose(1, 2)
encoder_hidden_states = torch.cat([encoder_hidden_states, encoder_hidden_states.new_zeros(1, self.text_len - encoder_hidden_states.size(1), encoder_hidden_states.size(2))], dim=1)
encoder_hidden_states_padding = encoder_hidden_states.new_zeros(
batch_size, self.text_len - encoder_hidden_states.size(1),
encoder_hidden_states.size(2))
encoder_hidden_states = torch.cat(
[encoder_hidden_states, encoder_hidden_states_padding], dim=1)
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
timestep.flatten(), encoder_hidden_states, encoder_hidden_states_image)
@@ -702,11 +888,10 @@ class CausalWanTransformer3DModel(BaseDiT):
teacher_forcing = clean_x is not None
if not isinstance(encoder_hidden_states, torch.Tensor):
encoder_hidden_states = encoder_hidden_states[0]
if isinstance(encoder_hidden_states_image,
list) and len(encoder_hidden_states_image) > 0:
encoder_hidden_states_image = encoder_hidden_states_image[0]
else:
encoder_hidden_states_image = None
if isinstance(encoder_hidden_states_image, list):
encoder_hidden_states_image = (
encoder_hidden_states_image[0]
if encoder_hidden_states_image else None)
batch_size, num_channels, num_frames, height, width = hidden_states.shape
p_t, p_h, p_w = self.patch_size
@@ -732,33 +917,24 @@ class CausalWanTransformer3DModel(BaseDiT):
freqs_cis = (freqs_cos,
freqs_sin) if freqs_cos is not None else None
if teacher_forcing:
if self.teacher_forcing_block_mask is None:
self.teacher_forcing_block_mask = self._prepare_teacher_forcing_mask(
device=hidden_states.device,
num_frames=num_frames,
frame_seqlen=post_patch_height * post_patch_width,
num_frame_per_block=self.num_frame_per_block,
local_attn_size=self.local_attn_size,
)
block_mask = self.teacher_forcing_block_mask
else:
if self.block_mask is None:
self.block_mask = self._prepare_blockwise_causal_attn_mask(
device=hidden_states.device,
num_frames=num_frames,
frame_seqlen=post_patch_height * post_patch_width,
num_frame_per_block=self.num_frame_per_block,
local_attn_size=self.local_attn_size
)
block_mask = self.block_mask
block_mask = self._get_train_attention_spec(
device=hidden_states.device,
num_frames=num_frames,
frame_seqlen=post_patch_height * post_patch_width,
teacher_forcing=teacher_forcing,
)
hidden_states = self.patch_embedding(hidden_states)
grid_sizes = torch.stack(
[torch.tensor(hidden_states[0].shape[1:], dtype=torch.long)])
grid_sizes = torch.tensor(
hidden_states.shape[2:], dtype=torch.long).unsqueeze(0).repeat(
batch_size, 1)
hidden_states = hidden_states.flatten(2).transpose(1, 2)
encoder_hidden_states = torch.cat([encoder_hidden_states, encoder_hidden_states.new_zeros(1, self.text_len - encoder_hidden_states.size(1), encoder_hidden_states.size(2))], dim=1)
encoder_hidden_states_padding = encoder_hidden_states.new_zeros(
batch_size, self.text_len - encoder_hidden_states.size(1),
encoder_hidden_states.size(2))
encoder_hidden_states = torch.cat(
[encoder_hidden_states, encoder_hidden_states_padding], dim=1)
encoder_hidden_states_text = encoder_hidden_states
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
@@ -0,0 +1,17 @@
from fastvideo.models.dits.trackwan.model import (
CausalTrackWanTransformer3DModel,
TrackWanTransformer3DModel,
)
from fastvideo.models.dits.trackwan.track_encoder import TrackEncoder
__all__ = [
"TrackEncoder",
"TrackWanTransformer3DModel",
"CausalTrackWanTransformer3DModel",
]
# Entry points for model registry discovery.
EntryClass = [
TrackWanTransformer3DModel,
CausalTrackWanTransformer3DModel,
]
+238
View File
@@ -0,0 +1,238 @@
# SPDX-License-Identifier: Apache-2.0
"""Bidirectional and causal track-conditioned Wan transformers."""
from __future__ import annotations
from typing import Any
import torch
from fastvideo.configs.models.dits.trackwan import TrackWanVideoConfig
from fastvideo.models.dits.causal_wanvideo import CausalWanTransformer3DModel
from fastvideo.models.dits.trackwan.track_encoder import TrackEncoder
from fastvideo.models.dits.wanvideo import WanTransformer3DModel
class _TrackConditioningMixin:
supports_track_input = True
track_channels: int
track_encoder: TrackEncoder
def _init_track_conditioning(
self,
config: TrackWanVideoConfig,
hf_config: dict[str, Any],
) -> None:
track_config = dict(getattr(config, "track_config", None) or {})
# Checkpoint config is authoritative for checkpoint-visible TrackEncoder
# shapes and legacy bias tensors. Defaults only fill missing fields.
if isinstance(hf_config, dict):
track_config.update(dict(hf_config.get("track_config", {}) or {}))
self.track_channels = int(
track_config.get("track_channels", self.in_channels - 36))
base_channels = self.in_channels - self.track_channels
expected_base_channels = self.num_channels_latents + 20
if self.track_channels <= 0:
raise ValueError("TrackWan requires positive track_channels")
if base_channels != expected_base_channels:
raise ValueError(
"TrackWan expects noisy latent + 20 I2V channels before "
f"track conditioning; got {base_channels}, expected "
f"{expected_base_channels}")
self.track_encoder = TrackEncoder(
id_dim=int(track_config.get("id_dim", 64)),
track_channels=self.track_channels,
vae_spatial_compression=int(
track_config.get("vae_spatial_compression", 8)),
vae_temporal_compression=int(
track_config.get("vae_temporal_compression", 4)),
max_track_id=int(track_config.get("max_track_id", 100_000)),
zero_init=bool(track_config.get("zero_init_head", False)),
use_bias=bool(track_config.get("use_bias", False)),
)
def _append_track_conditioning(
self,
hidden_states: torch.Tensor,
*,
track_points: torch.Tensor | None,
track_visibility: torch.Tensor | None,
track_ids: torch.Tensor | None,
track_map: torch.Tensor | None,
start_frame: int,
) -> torch.Tensor:
expected_channels = self.in_channels - self.track_channels
if hidden_states.ndim != 5:
raise ValueError(
"TrackWan hidden_states must be [B, C, T, H, W], got "
f"{tuple(hidden_states.shape)}")
if hidden_states.shape[1] != expected_channels:
raise ValueError(
"TrackWan received an unexpected pre-track channel count: "
f"{hidden_states.shape[1]} (expected {expected_channels})")
if (track_points is None) != (track_visibility is None):
raise ValueError(
"track_points and track_visibility must be provided together")
if track_map is not None and track_points is not None:
raise ValueError(
"track_map is mutually exclusive with raw track_points/"
"track_visibility")
batch_size, _, latent_t, latent_h, latent_w = hidden_states.shape
if track_map is not None:
expected_shape = (
batch_size,
self.track_channels,
latent_t,
latent_h,
latent_w,
)
if tuple(track_map.shape) != expected_shape:
raise ValueError(
"track_map must be latent-aligned with shape "
f"{expected_shape}, got {tuple(track_map.shape)}")
track_map = track_map.to(
device=hidden_states.device,
dtype=hidden_states.dtype,
)
elif track_points is None:
track_map = hidden_states.new_zeros(
batch_size,
self.track_channels,
latent_t,
latent_h,
latent_w,
)
else:
if track_points.shape[0] != batch_size:
raise ValueError(
"track batch size must match hidden_states batch size")
temporal_ratio = self.track_encoder.vae_temporal_compression
full_latent_t = ((track_points.shape[1] - 1) //
temporal_ratio + 1)
end_frame = int(start_frame) + latent_t
if start_frame < 0 or end_frame > full_latent_t:
raise ValueError(
"track sequence does not cover the requested latent "
f"window [{start_frame}, {end_frame}); available "
f"latent frames: {full_latent_t}")
# Encode from the global beginning so the causal temporal conv is
# left-padded exactly once, then take the requested latent chunk.
full_track_map = self.track_encoder(
track_points,
track_visibility,
full_latent_t,
latent_h,
latent_w,
track_ids=track_ids,
)
track_map = full_track_map[:, :, start_frame:end_frame]
track_map = track_map.to(dtype=hidden_states.dtype)
return torch.cat([hidden_states, track_map], dim=1)
class TrackWanTransformer3DModel(
_TrackConditioningMixin, WanTransformer3DModel):
"""Bidirectional Wan I2V transformer conditioned on sparse point tracks."""
def __init__(
self,
config: TrackWanVideoConfig,
hf_config: dict[str, Any],
) -> None:
super().__init__(config=config, hf_config=hf_config)
self._init_track_conditioning(config, hf_config)
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor | list[torch.Tensor],
timestep: torch.LongTensor,
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor]
| None = None,
guidance: torch.Tensor | None = None,
r_timestep: torch.Tensor | None = None,
track_points: torch.Tensor | None = None,
track_visibility: torch.Tensor | None = None,
track_ids: torch.Tensor | None = None,
track_map: torch.Tensor | None = None,
**kwargs: Any,
) -> torch.Tensor:
hidden_states = self._append_track_conditioning(
hidden_states,
track_points=track_points,
track_visibility=track_visibility,
track_ids=track_ids,
track_map=track_map,
start_frame=0,
)
return super().forward(
hidden_states,
encoder_hidden_states,
timestep,
encoder_hidden_states_image=encoder_hidden_states_image,
guidance=guidance,
r_timestep=r_timestep,
**kwargs,
)
class CausalTrackWanTransformer3DModel(
_TrackConditioningMixin, CausalWanTransformer3DModel):
"""Causal Wan I2V transformer with globally aligned track chunks."""
def __init__(
self,
config: TrackWanVideoConfig,
hf_config: dict[str, Any],
) -> None:
super().__init__(config=config, hf_config=hf_config)
self._init_track_conditioning(config, hf_config)
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor | list[torch.Tensor],
timestep: torch.LongTensor,
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor]
| None = None,
start_frame: int = 0,
clean_x: torch.Tensor | None = None,
aug_t: torch.Tensor | None = None,
track_points: torch.Tensor | None = None,
track_visibility: torch.Tensor | None = None,
track_ids: torch.Tensor | None = None,
track_map: torch.Tensor | None = None,
**kwargs: Any,
) -> torch.Tensor:
hidden_states = self._append_track_conditioning(
hidden_states,
track_points=track_points,
track_visibility=track_visibility,
track_ids=track_ids,
track_map=track_map,
start_frame=start_frame,
)
if clean_x is not None:
clean_x = self._append_track_conditioning(
clean_x,
track_points=track_points,
track_visibility=track_visibility,
track_ids=track_ids,
track_map=track_map,
start_frame=start_frame,
)
return super().forward(
hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
timestep=timestep,
encoder_hidden_states_image=encoder_hidden_states_image,
start_frame=start_frame,
clean_x=clean_x,
aug_t=aug_t,
**kwargs,
)
@@ -0,0 +1,259 @@
# SPDX-License-Identifier: Apache-2.0
"""Sparse point tracks to a dense, latent-aligned conditioning map."""
from __future__ import annotations
import math
import torch
import torch.nn.functional as F
from torch import nn
def sinusoidal_embedding(
ids: torch.Tensor,
dim: int,
max_period: int = 10_000,
) -> torch.Tensor:
"""Map integer identity labels to sinusoidal embeddings."""
half = dim // 2
frequencies = torch.exp(
-math.log(max_period) *
torch.arange(half, device=ids.device, dtype=torch.float32) /
max(half, 1))
angles = ids.float().unsqueeze(-1) * frequencies
embedding = torch.cat([torch.cos(angles), torch.sin(angles)], dim=-1)
if dim % 2:
embedding = F.pad(embedding, (0, 1))
return embedding
class TrackEncoder(nn.Module):
"""Encode ``[B, T, N, 2]`` normalized tracks on the Wan latent grid."""
def __init__(
self,
id_dim: int,
track_channels: int,
vae_spatial_compression: int = 8,
vae_temporal_compression: int = 4,
max_track_id: int = 100_000,
zero_init: bool = False,
use_bias: bool = False,
) -> None:
super().__init__()
if id_dim <= 0 or track_channels <= 0:
raise ValueError("id_dim and track_channels must be positive")
if vae_spatial_compression <= 0 or vae_temporal_compression <= 0:
raise ValueError("VAE compression ratios must be positive")
if max_track_id <= 0:
raise ValueError("max_track_id must be positive")
self.id_dim = int(id_dim)
self.track_channels = int(track_channels)
self.vae_spatial_compression = int(vae_spatial_compression)
self.vae_temporal_compression = int(vae_temporal_compression)
self.max_track_id = int(max_track_id)
self.use_bias = bool(use_bias)
# Bias-free convolutions keep empty cells exactly zero. Left-padding
# below gives the same temporal mapping as the causal Wan VAE:
# T_latent = (T_pixel - 1) // ratio + 1.
self.temporal_conv = nn.Conv3d(
self.id_dim,
self.track_channels,
kernel_size=(self.vae_temporal_compression, 1, 1),
stride=(self.vae_temporal_compression, 1, 1),
bias=self.use_bias,
)
self.proj = nn.Conv3d(
self.track_channels,
self.track_channels,
kernel_size=1,
bias=self.use_bias,
)
if self.temporal_conv.bias is not None:
nn.init.zeros_(self.temporal_conv.bias)
if self.proj.bias is not None:
nn.init.zeros_(self.proj.bias)
if zero_init:
nn.init.zeros_(self.proj.weight)
def sample_ids(
self,
batch_size: int,
num_tracks: int,
device: torch.device,
generator: torch.Generator | None = None,
) -> torch.Tensor:
"""Sample IDs once at a batch/session boundary."""
if num_tracks > self.max_track_id:
raise ValueError(
f"num_tracks ({num_tracks}) exceeds max_track_id "
f"({self.max_track_id}); unique IDs are unavailable")
return torch.stack([
torch.randperm(
self.max_track_id,
device=device,
generator=generator,
)[:num_tracks] for _ in range(batch_size)
])
def forward(
self,
coords: torch.Tensor,
visibility: torch.Tensor,
latent_t: int,
latent_h: int,
latent_w: int,
track_ids: torch.Tensor | None = None,
) -> torch.Tensor:
dense = self._rasterize(
coords,
visibility,
latent_h=latent_h,
latent_w=latent_w,
track_ids=track_ids,
)
dense = F.pad(
dense,
(0, 0, 0, 0, self.vae_temporal_compression - 1, 0),
)
dense = self.temporal_conv(dense)
if dense.shape[2:] != (latent_t, latent_h, latent_w):
dense = F.interpolate(
dense,
size=(latent_t, latent_h, latent_w),
mode="nearest",
)
return self.proj(dense)
def forward_window(
self,
coords: torch.Tensor,
visibility: torch.Tensor,
*,
latent_start: int,
latent_t: int,
latent_h: int,
latent_w: int,
pixel_start: int = 0,
track_ids: torch.Tensor | None = None,
) -> torch.Tensor:
"""Encode only the pixel window needed for a causal latent block.
``coords`` may be a complete sequence or a cropped sequence whose first
frame has global pixel index ``pixel_start``. The result is exactly the
same as slicing ``forward(full)[..., latent_start:latent_start+latent_t]``.
"""
if latent_start < 0:
raise ValueError("latent_start must be non-negative")
if pixel_start < 0:
raise ValueError("pixel_start must be non-negative")
if latent_t <= 0:
raise ValueError("latent_t must be positive")
ratio = self.vae_temporal_compression
required_start = max(0, latent_start * ratio - (ratio - 1))
required_end = (latent_start + latent_t - 1) * ratio + 1
provided_end = pixel_start + coords.shape[1]
if pixel_start > required_start or provided_end < required_end:
raise ValueError(
"track window does not cover required global pixel frames "
f"[{required_start}, {required_end}); provided "
f"[{pixel_start}, {provided_end})")
local_start = required_start - pixel_start
local_end = required_end - pixel_start
dense = self._rasterize(
coords[:, local_start:local_end],
visibility[:, local_start:local_end],
latent_h=latent_h,
latent_w=latent_w,
track_ids=track_ids,
)
if required_start == 0:
dense = F.pad(
dense,
(0, 0, 0, 0, ratio - 1, 0),
)
dense = self.temporal_conv(dense)
if dense.shape[2] != latent_t:
raise RuntimeError(
"TrackEncoder window produced an unexpected latent length: "
f"{dense.shape[2]} (expected {latent_t})")
return self.proj(dense)
def _rasterize(
self,
coords: torch.Tensor,
visibility: torch.Tensor,
*,
latent_h: int,
latent_w: int,
track_ids: torch.Tensor | None,
) -> torch.Tensor:
if coords.ndim != 4 or coords.shape[-1] != 2:
raise ValueError(
"track coords must have shape [B, T, N, 2], got "
f"{tuple(coords.shape)}")
if visibility.shape != coords.shape[:-1]:
raise ValueError(
"track visibility must have shape [B, T, N] matching "
f"coords, got {tuple(visibility.shape)}")
if latent_h <= 0 or latent_w <= 0:
raise ValueError("latent spatial dimensions must be positive")
batch_size, _, num_tracks, _ = coords.shape
device = coords.device
if num_tracks > self.max_track_id:
raise ValueError(
f"num_tracks ({num_tracks}) exceeds max_track_id "
f"({self.max_track_id})")
if track_ids is None:
# A deterministic fallback is required for repeated denoising and
# causal chunks. Training samples random IDs once in its wrapper.
track_ids = torch.arange(
num_tracks,
device=device,
).unsqueeze(0).expand(batch_size, num_tracks)
elif track_ids.shape != (batch_size, num_tracks):
raise ValueError(
"track_ids must have shape [B, N], got "
f"{tuple(track_ids.shape)}")
else:
track_ids = track_ids.to(device=device)
identities = sinusoidal_embedding(track_ids, self.id_dim)
coords_float = coords.float()
x = (coords_float[..., 0].clamp(0.0, 1.0) *
(latent_w - 1)).round().long()
y = (coords_float[..., 1].clamp(0.0, 1.0) *
(latent_h - 1)).round().long()
visible = visibility.to(device=device, dtype=torch.float32)
contribution = visible.unsqueeze(-1) * identities.unsqueeze(1)
cell_index = y * latent_w + x
dense = torch.zeros(
(
batch_size,
coords.shape[1],
latent_h * latent_w,
self.id_dim,
),
device=device,
dtype=torch.float32,
)
dense.scatter_add_(
2,
cell_index.unsqueeze(-1).expand(-1, -1, -1, self.id_dim),
contribution,
)
dense = dense.view(
batch_size,
coords.shape[1],
latent_h,
latent_w,
self.id_dim,
).permute(0, 4, 1, 2, 3).contiguous()
return dense.to(dtype=self.temporal_conv.weight.dtype)
+173
View File
@@ -0,0 +1,173 @@
# SPDX-License-Identifier: Apache-2.0
"""Helpers for WanTrack inference demos."""
from __future__ import annotations
from pathlib import Path
from typing import Any
import numpy as np
import torch
def create_track_presets(
num_frames: int,
num_tracks: int = 8,
seed: int = 0,
) -> dict[str, torch.Tensor]:
"""Build simple linear tracks in normalized ``[0, 1]`` coordinates.
Returns unbatched tensors:
``track_points`` ``[T, N, 2]``, ``track_visibility`` ``[T, N]``,
``track_ids`` ``[N]``.
"""
if num_frames < 1 or num_tracks < 1:
raise ValueError("num_frames and num_tracks must be positive")
generator = torch.Generator(device="cpu").manual_seed(int(seed))
starts = torch.rand(num_tracks, 2, generator=generator)
ends = torch.rand(num_tracks, 2, generator=generator)
alphas = torch.linspace(0.0, 1.0, num_frames).view(num_frames, 1, 1)
points = starts.unsqueeze(0) * (1.0 - alphas) + ends.unsqueeze(0) * alphas
visibility = torch.ones(num_frames, num_tracks, dtype=torch.float32)
track_ids = torch.arange(num_tracks, dtype=torch.long)
return {
"track_points": points.to(torch.float32),
"track_visibility": visibility,
"track_ids": track_ids,
}
def _as_tensor(value: Any, *, dtype: torch.dtype | None = None) -> torch.Tensor:
if isinstance(value, torch.Tensor):
tensor = value.detach().cpu()
else:
tensor = torch.as_tensor(np.asarray(value))
if dtype is not None:
tensor = tensor.to(dtype=dtype)
return tensor
def _ensure_batch(track_points: torch.Tensor, track_visibility: torch.Tensor,
track_ids: torch.Tensor | None) -> dict[str, torch.Tensor]:
"""Normalize shapes to ``[B, T, N, 2]`` / ``[B, T, N]`` / ``[B, N]``."""
points = _as_tensor(track_points, dtype=torch.float32)
visibility = _as_tensor(track_visibility, dtype=torch.float32)
if points.ndim == 3:
points = points.unsqueeze(0)
if visibility.ndim == 2:
visibility = visibility.unsqueeze(0)
if points.ndim != 4 or points.shape[-1] != 2:
raise ValueError(f"track_points must be [B,T,N,2] or [T,N,2], got {tuple(points.shape)}")
if visibility.shape != points.shape[:-1]:
raise ValueError("track_visibility shape must match track_points[:-1], got "
f"{tuple(visibility.shape)} vs {tuple(points.shape)}")
batch_size, _t, num_tracks = points.shape[:3]
if track_ids is None:
ids = torch.arange(num_tracks, dtype=torch.long).unsqueeze(0).expand(batch_size, -1).contiguous()
else:
ids = _as_tensor(track_ids, dtype=torch.long)
if ids.ndim == 1:
ids = ids.unsqueeze(0)
if ids.shape != (batch_size, num_tracks):
raise ValueError(f"track_ids must be [B,N] or [N], got {tuple(ids.shape)}")
return {
"track_points": points,
"track_visibility": visibility,
"track_ids": ids,
}
def _load_array_file(path: Path) -> Any:
suffix = path.suffix.lower()
if suffix in {".pt", ".pth"}:
return torch.load(path, map_location="cpu", weights_only=False)
if suffix == ".npz":
return np.load(path, allow_pickle=True)
if suffix == ".npy":
return np.load(path)
raise ValueError(f"Unsupported track file type: {path} (use .pt/.pth/.npz/.npy)")
def load_tracks(
*,
tracks_path: str | Path | None = None,
track_points_path: str | Path | None = None,
track_visibility_path: str | Path | None = None,
track_ids_path: str | Path | None = None,
num_frames: int | None = None,
num_tracks: int = 8,
seed: int = 0,
) -> dict[str, torch.Tensor]:
"""Load track tensors from disk, or synthesize a demo preset.
Accepted inputs:
- ``tracks_path``: ``.pt``/``.npz`` dict with ``track_points`` +
``track_visibility`` (optional ``track_ids``), or a directory containing
``track_points.*`` / ``track_visibility.*`` / ``track_ids.*``
- or the three individual file paths
- if nothing is provided, ``create_track_presets`` is used
"""
if tracks_path is not None:
path = Path(tracks_path)
if path.is_dir():
def _find(stem: str) -> Path | None:
for suffix in (".pt", ".pth", ".npy", ".npz"):
candidate = path / f"{stem}{suffix}"
if candidate.exists():
return candidate
return None
points_file = _find("track_points")
vis_file = _find("track_visibility")
ids_file = _find("track_ids")
if points_file is None or vis_file is None:
raise FileNotFoundError(
f"{path} must contain track_points.(pt|npy|npz) and "
"track_visibility.(pt|npy|npz)")
return load_tracks(
track_points_path=points_file,
track_visibility_path=vis_file,
track_ids_path=ids_file,
num_frames=num_frames,
)
payload = _load_array_file(path)
if isinstance(payload, np.lib.npyio.NpzFile):
payload = {key: payload[key] for key in payload.files}
if not isinstance(payload, dict):
raise ValueError(f"{path} must contain a dict with track_points/track_visibility")
points = payload.get("track_points")
visibility = payload.get("track_visibility")
if points is None or visibility is None:
raise ValueError(f"{path} missing required keys track_points/track_visibility")
tracks = _ensure_batch(points, visibility, payload.get("track_ids"))
elif track_points_path is not None and track_visibility_path is not None:
points = _load_array_file(Path(track_points_path))
visibility = _load_array_file(Path(track_visibility_path))
if isinstance(points, dict):
points = points["track_points"] if "track_points" in points else next(iter(points.values()))
if isinstance(visibility, dict):
visibility = (visibility["track_visibility"]
if "track_visibility" in visibility else next(iter(visibility.values())))
ids = None
if track_ids_path is not None:
ids = _load_array_file(Path(track_ids_path))
if isinstance(ids, dict):
ids = ids["track_ids"] if "track_ids" in ids else next(iter(ids.values()))
tracks = _ensure_batch(points, visibility, ids)
elif track_points_path is None and track_visibility_path is None and track_ids_path is None:
if num_frames is None:
raise ValueError("num_frames is required when synthesizing demo tracks")
raw = create_track_presets(num_frames, num_tracks=num_tracks, seed=seed)
tracks = _ensure_batch(raw["track_points"], raw["track_visibility"], raw["track_ids"])
else:
raise ValueError("Provide --tracks, or both --track-points and --track-visibility")
if num_frames is not None:
tracks["track_points"] = tracks["track_points"][:, :num_frames]
tracks["track_visibility"] = tracks["track_visibility"][:, :num_frames]
return tracks
+27 -12
View File
@@ -258,7 +258,7 @@ class WanI2VCrossAttention(WanSelfAttention):
self.norm_added_k = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
self.norm_added_q = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
def forward(self, x, context, context_lens):
def forward(self, x, context, context_lens, crossattn_cache=None):
r"""
Args:
x(Tensor): Shape [B, L1, C]
@@ -269,13 +269,29 @@ class WanI2VCrossAttention(WanSelfAttention):
context = context[:, 257:]
b, n, d = x.size(0), self.num_heads, self.head_dim
# compute query, key, value
# Query depends on the current latent tokens. Text/image keys and
# values only depend on conditioning and can be reused by causal
# streaming rollouts.
q = self.norm_q(self.to_q(x)[0]).view(b, -1, n, d)
k = self.norm_k(self.to_k(context)[0]).view(b, -1, n, d)
v = self.to_v(context)[0].view(b, -1, n, d)
k_img = self.norm_added_k(self.add_k_proj(context_img)[0]).view(
b, -1, n, d)
v_img = self.add_v_proj(context_img)[0].view(b, -1, n, d)
if crossattn_cache is not None and crossattn_cache.get("is_init", False):
k = crossattn_cache["k"]
v = crossattn_cache["v"]
k_img = crossattn_cache["k_img"]
v_img = crossattn_cache["v_img"]
else:
k = self.norm_k(self.to_k(context)[0]).view(b, -1, n, d)
v = self.to_v(context)[0].view(b, -1, n, d)
k_img = self.norm_added_k(self.add_k_proj(context_img)[0]).view(
b, -1, n, d)
v_img = self.add_v_proj(context_img)[0].view(b, -1, n, d)
if crossattn_cache is not None:
crossattn_cache.update({
"is_init": True,
"k": k,
"v": v,
"k_img": k_img,
"v_img": v_img,
})
img_x = self.attn(q, k_img, v_img)
# compute attention
x = self.attn(q, k, v) if k.size(1) > 0 else torch.zeros_like(q)
@@ -694,11 +710,10 @@ class WanTransformer3DModel(BaseDiT):
orig_dtype = hidden_states.dtype
if encoder_hidden_states is not None and not isinstance(encoder_hidden_states, torch.Tensor):
encoder_hidden_states = encoder_hidden_states[0]
if isinstance(encoder_hidden_states_image,
list) and len(encoder_hidden_states_image) > 0:
encoder_hidden_states_image = encoder_hidden_states_image[0]
else:
encoder_hidden_states_image = None
if isinstance(encoder_hidden_states_image, list):
encoder_hidden_states_image = (
encoder_hidden_states_image[0]
if encoder_hidden_states_image else None)
batch_size, num_channels, num_frames, height, width = hidden_states.shape
p_t, p_h, p_w = self.patch_size
+8
View File
@@ -32,9 +32,13 @@ _TEXT_TO_VIDEO_DIT_MODELS = {
"HYWorldTransformer3DModel":
("dits", "hyworld", "HYWorldTransformer3DModel"),
"WanTransformer3DModel": ("dits", "wanvideo", "WanTransformer3DModel"),
"TrackWanTransformer3DModel":
("dits", "trackwan", "TrackWanTransformer3DModel"),
"DreamXWorldTransformer3DModel": ("dits", "dreamx_world", "DreamXWorldTransformer3DModel"),
"DreamXWorldARTransformer3DModel": ("dits", "dreamx_world_ar", "DreamXWorldARTransformer3DModel"),
"CausalWanTransformer3DModel": ("dits", "causal_wanvideo", "CausalWanTransformer3DModel"),
"CausalTrackWanTransformer3DModel":
("dits", "trackwan", "CausalTrackWanTransformer3DModel"),
"CosmosTransformer3DModel": ("dits", "cosmos", "CosmosTransformer3DModel"),
"Cosmos25Transformer3DModel": ("dits", "cosmos2_5", "Cosmos25Transformer3DModel"),
"LongCatVideoTransformer3DModel": ("dits", "longcat_video_dit", "LongCatVideoTransformer3DModel"), # Wrapper (Phase 1)
@@ -55,9 +59,13 @@ _TEXT_TO_VIDEO_DIT_MODELS = {
_IMAGE_TO_VIDEO_DIT_MODELS = {
# "HunyuanVideoTransformer3DModel": ("dits", "hunyuanvideo", "HunyuanVideoDiT"),
"WanTransformer3DModel": ("dits", "wanvideo", "WanTransformer3DModel"),
"TrackWanTransformer3DModel":
("dits", "trackwan", "TrackWanTransformer3DModel"),
"DreamXWorldTransformer3DModel": ("dits", "dreamx_world", "DreamXWorldTransformer3DModel"),
"DreamXWorldARTransformer3DModel": ("dits", "dreamx_world_ar", "DreamXWorldARTransformer3DModel"),
"CausalWanTransformer3DModel": ("dits", "causal_wanvideo", "CausalWanTransformer3DModel"),
"CausalTrackWanTransformer3DModel":
("dits", "trackwan", "CausalTrackWanTransformer3DModel"),
"LingBotWorld2CausalFastTransformer3DModel": (
"dits",
"lingbotworld2",
@@ -0,0 +1 @@
# SPDX-License-Identifier: Apache-2.0
@@ -0,0 +1,34 @@
# SPDX-License-Identifier: Apache-2.0
"""WanTrack pipeline presets."""
from fastvideo.api.presets import InferencePreset, PresetStageSpec
_DENOISE_STAGE = PresetStageSpec(
name="denoise",
kind="denoising",
description="Causal TrackWan Self-Forcing denoising",
allowed_overrides=frozenset({
"num_inference_steps",
"guidance_scale",
}),
)
SF_WANTRACK_CAUSAL_I2V = InferencePreset(
name="sf_wantrack_causal_i2v",
version=1,
model_family="wantrack",
description="Causal WanTrack Self-Forcing I2V (Track-v0)",
workload_type="i2v",
stage_schemas=(_DENOISE_STAGE, ),
defaults={
"height": 480,
"width": 832,
"num_frames": 121,
"fps": 16,
"guidance_scale": 1.0,
"num_inference_steps": 4,
"negative_prompt": "",
},
)
ALL_PRESETS = (SF_WANTRACK_CAUSAL_I2V, )
@@ -0,0 +1,77 @@
# SPDX-License-Identifier: Apache-2.0
"""Causal Self-Forcing WanTrack I2V pipeline."""
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler, )
from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
from fastvideo.pipelines.stages import (
ConditioningStage,
DecodingStage,
ImageEncodingStage,
ImageVAEEncodingStage,
InputValidationStage,
LatentPreparationStage,
TextEncodingStage,
)
from fastvideo.pipelines.stages.wantrack_causal_denoising import (
WanTrackCausalDenoisingStage, )
class WanTrackCausalDMDPipeline(LoRAPipeline, ComposedPipelineBase):
"""Wan I2V encode stages + causal TrackWan DMD denoising."""
_required_config_modules = [
"text_encoder",
"tokenizer",
"vae",
"transformer",
"scheduler",
"image_encoder",
"image_processor",
]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs) -> None:
self.modules["scheduler"] = FlowMatchEulerDiscreteScheduler(
shift=fastvideo_args.pipeline_config.flow_shift)
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
self.add_stage(stage_name="input_validation_stage", stage=InputValidationStage())
self.add_stage(
stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
),
)
self.add_stage(
stage_name="image_encoding_stage",
stage=ImageEncodingStage(
image_encoder=self.get_module("image_encoder"),
image_processor=self.get_module("image_processor"),
),
)
self.add_stage(stage_name="conditioning_stage", stage=ConditioningStage())
self.add_stage(
stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer", None),
),
)
self.add_stage(
stage_name="image_latent_preparation_stage",
stage=ImageVAEEncodingStage(vae=self.get_module("vae")),
)
self.add_stage(
stage_name="denoising_stage",
stage=WanTrackCausalDenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler"),
vae=self.get_module("vae"),
),
)
self.add_stage(stage_name="decoding_stage", stage=DecodingStage(vae=self.get_module("vae")))
EntryClass = WanTrackCausalDMDPipeline
@@ -134,6 +134,11 @@ class ForwardBatch:
mouse_cond: torch.Tensor | None = None # Shape: (B, T, 2)
keyboard_cond: torch.Tensor | None = None # Shape: (B, T, K)
grid_sizes: torch.Tensor | None = None # Shape: (3,) [F,H,W]
track_points: torch.Tensor | None = None # Shape: (B, T, N, 2)
track_visibility: torch.Tensor | None = None # Shape: (B, T, N)
track_ids: torch.Tensor | None = None # Shape: (B, N)
track_map: torch.Tensor | None = None
num_iterations: int | None = None
use_base_model: bool = False
@@ -301,6 +306,14 @@ class TrainingBatch:
image_latents: torch.Tensor | None = None
infos: list[dict[str, Any]] | None = None
mask_lat_size: torch.Tensor | None = None
# WanTrack point trajectories. IDs are sampled once per batch/session
# and reused by all denoising steps and causal chunks.
track_points: torch.Tensor | None = None
track_visibility: torch.Tensor | None = None
track_ids: torch.Tensor | None = None
# Optional latent-aligned result of TrackEncoder. Runtime callers use this
# to encode only the causal pixel window required by the current block.
track_map: torch.Tensor | None = None
# ODE trajectory supervision
trajectory_latents: torch.Tensor | None = None
@@ -286,6 +286,14 @@ class BasePreprocessPipeline(ComposedPipelineBase):
if "action_path" in data:
valid_data["action_path"] = [data["action_path"][i] for i in valid_indices]
if "points_path" in data:
valid_data["points_path"] = [data["points_path"][i] for i in valid_indices]
if "sample_frame_index" in data:
valid_data["sample_frame_index"] = [data["sample_frame_index"][i] for i in valid_indices]
if "source_width" in data:
valid_data["source_width"] = [data["source_width"][i] for i in valid_indices]
if "source_height" in data:
valid_data["source_height"] = [data["source_height"][i] for i in valid_indices]
# VAE
with torch.autocast("cuda", dtype=torch.float32):
@@ -0,0 +1,214 @@
# SPDX-License-Identifier: Apache-2.0
"""I2V preprocessing with point-track conditioning."""
from typing import Any
import numpy as np
from fastvideo.dataset.dataloader.schema import pyarrow_schema_i2v_track
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.pipelines.preprocess.preprocess_pipeline_i2v import (
PreprocessPipeline_I2V, )
class PreprocessPipeline_I2V_Track(PreprocessPipeline_I2V):
"""Store I2V features and spatially aligned MotionStream point tracks."""
def get_pyarrow_schema(self):
return pyarrow_schema_i2v_track
@staticmethod
def _get_source_size(
sidecar: Any,
valid_data: dict[str, Any],
idx: int,
points_path: str,
) -> tuple[float, float]:
if "width" in sidecar and "height" in sidecar:
width = float(np.asarray(sidecar["width"]).item())
height = float(np.asarray(sidecar["height"]).item())
elif "source_width" in valid_data and "source_height" in valid_data:
width = float(valid_data["source_width"][idx])
height = float(valid_data["source_height"][idx])
else:
raise ValueError(f"{points_path}: source width and height are required to align "
"pixel-space tracks. Store them in the sidecar or manifest resolution.")
if width <= 0 or height <= 0:
raise ValueError(f"{points_path}: invalid source size ({width}, {height})")
return width, height
@staticmethod
def _normalize_after_center_crop(
tracks: np.ndarray,
visibility: np.ndarray,
*,
source_width: float,
source_height: float,
target_width: int,
target_height: int,
) -> tuple[np.ndarray, np.ndarray]:
"""Apply ``CenterCropResizeVideo`` geometry and normalize coordinates."""
if target_width <= 0 or target_height <= 0:
raise ValueError(f"Invalid target size ({target_width}, {target_height})")
target_ratio = target_height / target_width
if source_height / source_width > target_ratio:
crop_height = int(source_width * target_ratio)
crop_width = int(source_width)
else:
crop_height = int(source_height)
crop_width = int(source_height / target_ratio)
if crop_width <= 0 or crop_height <= 0:
raise ValueError(f"Center crop is empty for source size "
f"({source_width}, {source_height}) and target size "
f"({target_width}, {target_height})")
crop_top = int(round((source_height - crop_height) / 2.0))
crop_left = int(round((source_width - crop_width) / 2.0))
x = tracks[..., 0]
y = tracks[..., 1]
finite = np.isfinite(x) & np.isfinite(y)
in_crop = (finite
& (x >= crop_left)
& (x < crop_left + crop_width)
& (y >= crop_top)
& (y < crop_top + crop_height))
normalized = np.empty_like(tracks, dtype=np.float32)
normalized[..., 0] = (x - crop_left) / crop_width
normalized[..., 1] = (y - crop_top) / crop_height
normalized = np.nan_to_num(normalized, nan=0.0, posinf=1.0, neginf=0.0)
np.clip(normalized, 0.0, 1.0, out=normalized)
aligned_visibility = np.asarray(visibility, dtype=np.float32) * in_crop
aligned_visibility = np.nan_to_num(aligned_visibility, nan=0.0, posinf=1.0, neginf=0.0)
np.clip(aligned_visibility, 0.0, 1.0, out=aligned_visibility)
return (
np.ascontiguousarray(normalized, dtype=np.float32),
np.ascontiguousarray(aligned_visibility, dtype=np.float32),
)
def get_extra_features(
self,
valid_data: dict[str, Any],
fastvideo_args: FastVideoArgs,
) -> dict[str, Any]:
features = super().get_extra_features(valid_data, fastvideo_args)
points_paths = valid_data.get("points_path")
frame_indices_batch = valid_data.get("sample_frame_index")
if not points_paths or frame_indices_batch is None:
raise ValueError("i2v_track preprocessing requires points_path and sampled frame "
"indices for every video.")
non_video_paths = [path for path in valid_data["path"] if not str(path).lower().endswith(".mp4")]
if non_video_paths:
raise ValueError("i2v_track preprocessing currently supports videos only; "
f"received {non_video_paths[0]!r}")
target_height = int(valid_data["pixel_values"].shape[-2])
target_width = int(valid_data["pixel_values"].shape[-1])
expected_frames = int(valid_data["pixel_values"].shape[2])
track_points: list[np.ndarray] = []
track_visibility: list[np.ndarray] = []
object_ids: list[np.ndarray] = []
track_weights: list[np.ndarray] = []
samples = zip(points_paths, frame_indices_batch, strict=True)
for idx, (points_path, frame_indices) in enumerate(samples):
indices = np.asarray(frame_indices, dtype=np.int64)
if indices.ndim != 1 or indices.size != expected_frames:
raise ValueError(f"{points_path}: expected {expected_frames} sampled frame "
f"indices, got shape {indices.shape}")
if indices.size == 0 or indices.min() < 0:
raise ValueError(f"{points_path}: sampled frame indices must be non-negative")
with np.load(points_path) as sidecar:
if "tracks" not in sidecar or "visibility" not in sidecar:
raise ValueError(f"{points_path}: sidecar must contain 'tracks' and 'visibility'")
tracks = np.asarray(sidecar["tracks"], dtype=np.float32)
visibility = np.asarray(sidecar["visibility"], dtype=np.float32)
if tracks.ndim != 3 or tracks.shape[-1] != 2:
raise ValueError(f"{points_path}: tracks must have shape [T, N, 2], got "
f"{tracks.shape}")
if visibility.shape != tracks.shape[:2]:
raise ValueError(f"{points_path}: visibility shape {visibility.shape} does "
f"not match tracks {tracks.shape[:2]}")
if indices.max() >= tracks.shape[0]:
raise ValueError(f"{points_path}: sampled frame {indices.max()} exceeds "
f"{tracks.shape[0]} track frames")
source_width, source_height = self._get_source_size(sidecar, valid_data, idx, points_path)
num_tracks = tracks.shape[1]
sample_object_ids = np.asarray(
sidecar["object_ids"] if "object_ids" in sidecar else np.full(num_tracks, -1, dtype=np.float32),
dtype=np.float32,
)
sample_track_weights = np.asarray(
sidecar["track_weights"] if "track_weights" in sidecar else np.zeros(num_tracks, dtype=np.float32),
dtype=np.float32,
)
if sample_object_ids.shape != (num_tracks, ):
raise ValueError(f"{points_path}: object_ids must have shape "
f"[{num_tracks}], got {sample_object_ids.shape}")
if sample_track_weights.shape != (num_tracks, ):
raise ValueError(f"{points_path}: track_weights must have shape "
f"[{num_tracks}], got {sample_track_weights.shape}")
tracks = tracks[indices]
visibility = visibility[indices]
tracks, visibility = self._normalize_after_center_crop(
tracks,
visibility,
source_width=source_width,
source_height=source_height,
target_width=target_width,
target_height=target_height,
)
track_points.append(tracks)
track_visibility.append(visibility)
object_ids.append(np.ascontiguousarray(sample_object_ids, dtype=np.float32))
track_weights.append(np.ascontiguousarray(sample_track_weights, dtype=np.float32))
features["track_points"] = track_points
features["track_visibility"] = track_visibility
features["object_ids"] = object_ids
features["track_weights"] = track_weights
return features
def create_record(
self,
video_name: str,
vae_latent: np.ndarray,
text_embedding: np.ndarray,
valid_data: dict[str, Any],
idx: int,
extra_features: dict[str, Any] | None = None,
) -> dict[str, Any]:
record = super().create_record(
video_name=video_name,
vae_latent=vae_latent,
text_embedding=text_embedding,
valid_data=valid_data,
idx=idx,
extra_features=extra_features,
)
for name in (
"track_points",
"track_visibility",
"object_ids",
"track_weights",
):
if extra_features is None or name not in extra_features:
raise ValueError(f"Missing required WanTrack feature {name!r} for {video_name}")
array = np.ascontiguousarray(extra_features[name], dtype=np.float32)
record[f"{name}_bytes"] = array.tobytes()
record[f"{name}_shape"] = list(array.shape)
record[f"{name}_dtype"] = str(array.dtype)
record.pop(name, None)
return record
EntryClass = PreprocessPipeline_I2V_Track
@@ -8,6 +8,7 @@ from fastvideo.distributed import (get_world_size, maybe_init_distributed_enviro
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines.preprocess.preprocess_pipeline_i2v import (PreprocessPipeline_I2V)
from fastvideo.pipelines.preprocess.preprocess_pipeline_i2v_track import (PreprocessPipeline_I2V_Track)
from fastvideo.pipelines.preprocess.preprocess_pipeline_ode_trajectory import (PreprocessPipeline_ODE_Trajectory)
from fastvideo.pipelines.preprocess.preprocess_pipeline_t2v import (PreprocessPipeline_T2V)
from fastvideo.pipelines.preprocess.preprocess_pipeline_text import (PreprocessPipeline_Text)
@@ -52,6 +53,8 @@ def main(args) -> None:
PreprocessPipeline = PreprocessPipeline_T2V
elif args.preprocess_task == "i2v":
PreprocessPipeline = PreprocessPipeline_I2V
elif args.preprocess_task == "i2v_track":
PreprocessPipeline = PreprocessPipeline_I2V_Track
elif args.preprocess_task == "text_only":
PreprocessPipeline = PreprocessPipeline_Text
elif args.preprocess_task == "ode_trajectory":
@@ -72,7 +75,8 @@ def main(args) -> None:
PreprocessPipeline = PreprocessPipeline_MatrixGame2_ODE_Trajectory
else:
raise ValueError(f"Invalid preprocess task: {args.preprocess_task}. "
f"Valid options: t2v, i2v, ode_trajectory, text_only, matrixgame2, matrixgame2_ode_trajectory")
f"Valid options: t2v, i2v, i2v_track, ode_trajectory, text_only, "
f"matrixgame2, matrixgame2_ode_trajectory")
logger.info("Preprocess task: %s using %s", args.preprocess_task, PreprocessPipeline.__name__)
@@ -115,6 +119,7 @@ if __name__ == "__main__":
choices=[
"t2v",
"i2v",
"i2v_track",
"text_only",
"ode_trajectory",
"matrixgame2",
+2
View File
@@ -32,6 +32,7 @@ from fastvideo.pipelines.basic.ltx2.stages import ( # noqa: F401
)
from fastvideo.pipelines.stages.matrixgame2_denoising import MatrixGame2CausalDenoisingStage
from fastvideo.pipelines.stages.matrixgame3_denoising import MatrixGame3DenoisingStage
from fastvideo.pipelines.stages.wantrack_causal_denoising import WanTrackCausalDenoisingStage
from fastvideo.pipelines.stages.hyworld_denoising import HYWorldDenoisingStage
from fastvideo.pipelines.stages.kandinsky5 import (Kandinsky5DecodingStage, Kandinsky5DenoisingStage,
Kandinsky5LatentPreparationStage)
@@ -66,6 +67,7 @@ __all__ = [
"CausalDenoisingStage",
"MatrixGame2CausalDenoisingStage",
"MatrixGame3DenoisingStage",
"WanTrackCausalDenoisingStage",
"HYWorldDenoisingStage",
"Kandinsky5DecodingStage",
"Kandinsky5DenoisingStage",
+2 -1
View File
@@ -460,7 +460,8 @@ class ImageVAEEncodingStage(PipelineStage):
generator = batch.generator
if generator is None:
raise ValueError("Generator must be provided")
latent_condition = self.retrieve_latents(encoder_output, generator)
sample_mode = "argmax" if fastvideo_args.pipeline_config.is_causal else "sample"
latent_condition = self.retrieve_latents(encoder_output, generator, sample_mode=sample_mode)
# Apply shifting if needed
if (hasattr(self.vae, "shift_factor") and self.vae.shift_factor is not None):
@@ -99,11 +99,11 @@ class InputValidationStage(PipelineStage):
ih, iw = img.height, img.width
pipeline_class_name = type(fastvideo_args.pipeline_config).__name__
if 'MatrixGame' in pipeline_class_name or 'MatrixCausal' in pipeline_class_name:
is_matrix_game = ('MatrixGame' in pipeline_class_name or 'MatrixCausal' in pipeline_class_name)
if is_matrix_game or fastvideo_args.pipeline_config.is_causal:
oh, ow = batch.height, batch.width
img = img.resize((ow, oh), Image.LANCZOS)
else:
# Standard Wan logic
patch_size = fastvideo_args.pipeline_config.dit_config.arch_config.patch_size
vae_stride = fastvideo_args.pipeline_config.vae_config.arch_config.scale_factor_spatial
dh, dw = patch_size[1] * vae_stride, patch_size[2] * vae_stride
@@ -121,12 +121,13 @@ class InputValidationStage(PipelineStage):
assert img.width == ow and img.height == oh
logger.info("final processed img height: %s, img width: %s", img.height, img.width)
# to tensor
img = TF.to_tensor(img).sub_(0.5).div_(0.5).to(self.device).unsqueeze(1)
img = img.unsqueeze(0)
batch.height = oh
batch.width = ow
batch.pil_image = img
if is_matrix_game or fastvideo_args.pipeline_config.ti2v_task:
img = TF.to_tensor(img).sub_(0.5).div_(0.5).to(self.device).unsqueeze(1)
batch.pil_image = img.unsqueeze(0)
else:
batch.pil_image = img
# for v2v, get control video from video path
if batch.video_path is not None:
@@ -0,0 +1,236 @@
# SPDX-License-Identifier: Apache-2.0
"""Causal Self-Forcing denoising for TrackWan I2V."""
from __future__ import annotations
import torch
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.models.utils import pred_noise_to_pred_video
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.causal_denoising import CausalDMDDenosingStage
from fastvideo.pipelines.stages.validators import StageValidators as V
from fastvideo.pipelines.stages.validators import VerificationResult
logger = init_logger(__name__)
class WanTrackCausalDenoisingStage(CausalDMDDenosingStage):
"""Causal DMD loop with Wan I2V 20ch concat + sparse track conditioning.
Differs from :class:`CausalDMDDenosingStage` in three ways:
1. Latent clip geometry is ``1 + N * num_frames_per_block`` (e.g. 31 = 1+10*3).
2. Each block cats ``[noise_16, image_latent_20]`` before the DiT forward.
3. Passes ``track_points`` / ``track_visibility`` / ``track_ids`` and CLIP image embeds.
"""
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
target_dtype = torch.bfloat16
autocast_enabled = (target_dtype != torch.float32) and not fastvideo_args.disable_autocast
assert batch.latents is not None, "latents must be provided"
assert batch.image_latent is not None, "WanTrack requires image_latent (20ch I2V)"
assert batch.track_points is not None and batch.track_visibility is not None, (
"WanTrack requires track_points and track_visibility")
latents = batch.latents # [B, C, T, H, W]
b, _c, t, h, w = latents.shape
image_latent = batch.image_latent.to(device=latents.device, dtype=target_dtype)
if image_latent.shape[1] != 20:
raise ValueError(f"WanTrack image_latent must have 20 channels, got {image_latent.shape[1]}")
if image_latent.shape[2] < t:
raise ValueError(f"image_latent temporal length {image_latent.shape[2]} < latents {t}")
latent_seq_length = h * w
patch_ratio = (self.transformer.config.arch_config.patch_size[-1] *
self.transformer.config.arch_config.patch_size[-2])
self.frame_seq_length = latent_seq_length // patch_ratio
timesteps = torch.tensor(fastvideo_args.pipeline_config.dmd_denoising_steps, dtype=torch.long).cpu()
if fastvideo_args.pipeline_config.warp_denoising_step:
# Ensure the scheduler has a 1000-step grid for warping.
self.scheduler.set_timesteps(1000, device="cpu")
scheduler_timesteps = torch.cat(
(self.scheduler.timesteps.cpu(), torch.tensor([0], dtype=torch.float32)))
timesteps = scheduler_timesteps[1000 - timesteps]
timesteps = timesteps.to(get_local_torch_device())
prompt_embeds = batch.prompt_embeds
assert torch.isnan(prompt_embeds[0]).sum() == 0
image_embeds = batch.image_embeds
if not image_embeds:
raise ValueError("WanTrack requires CLIP image embeds")
image_kwargs = {
"encoder_hidden_states_image": [embed.to(target_dtype) for embed in image_embeds],
}
track_points = batch.track_points.to(device=latents.device, dtype=torch.float32)
track_visibility = batch.track_visibility.to(device=latents.device, dtype=torch.float32)
track_ids = batch.track_ids
if track_ids is None:
track_ids = self._sample_track_ids(track_points, batch)
else:
track_ids = track_ids.to(device=latents.device, dtype=torch.long)
track_kwargs = {
"track_points": track_points,
"track_visibility": track_visibility,
"track_ids": track_ids,
}
pos_cond_kwargs = self.prepare_extra_func_kwargs(
self.transformer.forward,
{"encoder_attention_mask": batch.prompt_attention_mask},
)
kv_cache1 = self._initialize_kv_cache(batch_size=b, dtype=target_dtype, device=latents.device)
crossattn_cache = self._initialize_crossattn_cache(
batch_size=b,
max_text_len=fastvideo_args.pipeline_config.text_encoder_configs[0].arch_config.text_len,
dtype=target_dtype,
device=latents.device,
)
block_sizes = self._block_sizes(t)
start_index = 0
pos_start_base = 0
context_noise = int(getattr(fastvideo_args.pipeline_config, "context_noise", 0))
with self.progress_bar(total=len(block_sizes) * len(timesteps)) as progress_bar:
for current_num_frames in block_sizes:
current_latents = latents[:, :, start_index:start_index + current_num_frames, :, :]
noise_latents_btchw = current_latents.permute(0, 2, 1, 3, 4)
video_raw_latent_shape = noise_latents_btchw.shape
image_block = image_latent[:, :, start_index:start_index + current_num_frames, :, :]
for i, t_cur in enumerate(timesteps):
noise_latents = noise_latents_btchw.clone()
latent_model_input = torch.cat(
[current_latents.to(target_dtype), image_block],
dim=1,
)
t_expand = t_cur.repeat(b)
with torch.autocast(device_type="cuda", dtype=target_dtype,
enabled=autocast_enabled), set_forward_context(
current_timestep=i,
attn_metadata=None,
forward_batch=batch,
):
t_expanded_noise = t_cur * torch.ones(
(b, 1), device=latent_model_input.device, dtype=torch.long)
pred_noise_btchw = self.transformer(
latent_model_input,
prompt_embeds,
t_expanded_noise,
kv_cache=kv_cache1,
crossattn_cache=crossattn_cache,
current_start=(pos_start_base + start_index) * self.frame_seq_length,
start_frame=start_index,
**image_kwargs,
**track_kwargs,
**pos_cond_kwargs,
).permute(0, 2, 1, 3, 4)
pred_video_btchw = pred_noise_to_pred_video(
pred_noise=pred_noise_btchw.flatten(0, 1),
noise_input_latent=noise_latents.flatten(0, 1),
timestep=t_expand,
scheduler=self.scheduler,
).unflatten(0, pred_noise_btchw.shape[:2])
if i < len(timesteps) - 1:
next_timestep = timesteps[i + 1] * torch.ones(
[1], dtype=torch.long, device=pred_video_btchw.device)
noise = torch.randn(
video_raw_latent_shape,
dtype=pred_video_btchw.dtype,
generator=(batch.generator[0]
if isinstance(batch.generator, list) else batch.generator),
).to(self.device)
noise_latents_btchw = self.scheduler.add_noise(
pred_video_btchw.flatten(0, 1),
noise.flatten(0, 1),
next_timestep,
).unflatten(0, pred_video_btchw.shape[:2])
current_latents = noise_latents_btchw.permute(0, 2, 1, 3, 4)
else:
current_latents = pred_video_btchw.permute(0, 2, 1, 3, 4)
if progress_bar is not None:
progress_bar.update()
latents[:, :, start_index:start_index + current_num_frames, :, :] = current_latents
# Commit clean context into the KV cache for later blocks.
t_context = torch.ones([b], device=latents.device, dtype=torch.long) * context_noise
context_input = torch.cat(
[current_latents.to(target_dtype), image_block],
dim=1,
)
with torch.autocast(device_type="cuda", dtype=target_dtype,
enabled=autocast_enabled), set_forward_context(
current_timestep=0,
attn_metadata=None,
forward_batch=batch,
):
self.transformer(
context_input,
prompt_embeds,
t_context.unsqueeze(1),
kv_cache=kv_cache1,
crossattn_cache=crossattn_cache,
current_start=(pos_start_base + start_index) * self.frame_seq_length,
start_frame=start_index,
**image_kwargs,
**track_kwargs,
**pos_cond_kwargs,
)
start_index += current_num_frames
batch.latents = latents
return batch
def _block_sizes(self, latent_t: int) -> list[int]:
chunk = int(self.num_frames_per_block)
if chunk <= 0:
raise ValueError("num_frames_per_block must be positive")
if latent_t % chunk == 0:
return [chunk] * (latent_t // chunk)
if (latent_t - 1) % chunk == 0:
return [1] + [chunk] * ((latent_t - 1) // chunk)
raise ValueError("Causal WanTrack requires latent frames that form complete "
f"blocks (optionally after one leading I2V frame); got "
f"latent_t={latent_t}, num_frames_per_block={chunk}")
def _sample_track_ids(
self,
track_points: torch.Tensor,
batch: ForwardBatch,
) -> torch.Tensor:
max_track_id = int(getattr(getattr(self.transformer, "track_encoder", None), "max_track_id", 100_000))
batch_size, _t, num_tracks = track_points.shape[:3]
if num_tracks > max_track_id:
raise ValueError(f"num_tracks ({num_tracks}) exceeds max_track_id ({max_track_id})")
generator = batch.generator[0] if isinstance(batch.generator, list) else batch.generator
# randperm on CUDA does not accept a CPU generator; sample on CPU then move.
ids = []
for _ in range(batch_size):
perm = torch.randperm(max_track_id, generator=generator)[:num_tracks]
ids.append(perm)
return torch.stack(ids, dim=0).to(device=track_points.device, dtype=torch.long)
def verify_input(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> VerificationResult:
result = super().verify_input(batch, fastvideo_args)
result.add_check("track_points", batch.track_points, [V.is_tensor, V.with_dims(4)])
result.add_check("track_visibility", batch.track_visibility, [V.is_tensor, V.with_dims(3)])
result.add_check("image_latent", batch.image_latent, [V.is_tensor, V.with_dims(5)])
return result
+23
View File
@@ -40,6 +40,7 @@ from fastvideo.configs.pipelines.flux_2 import (
)
from fastvideo.configs.pipelines.matrixgame2 import MatrixGame2I2V480PConfig
from fastvideo.configs.pipelines.matrixgame3 import MatrixGame3I2V720PConfig
from fastvideo.configs.pipelines.wantrack import CausalTrackWanSFI2VConfig
from fastvideo.configs.pipelines.turbodiffusion import (
TurboDiffusionI2V_A14B_Config,
TurboDiffusionT2V_14B_Config,
@@ -68,6 +69,7 @@ from fastvideo.configs.pipelines.zimage import ZImagePipelineConfig
from fastvideo.api.sampling_param import SamplingParam
from fastvideo.api.matrixgame2 import MatrixGame2SamplingParam
from fastvideo.api.matrixgame3 import MatrixGame3SamplingParam
from fastvideo.api.wantrack import WanTrackSamplingParam
from fastvideo.api.flux import FluxSamplingParam
from fastvideo.fastvideo_args import WorkloadType
@@ -782,6 +784,24 @@ def _register_configs() -> None:
model_family="matrixgame",
default_preset="matrixgame2_i2v",
)
# Causal WanTrack Self-Forcing I2V (Track-v0)
register_configs(
sampling_param_cls=WanTrackSamplingParam,
pipeline_config_cls=CausalTrackWanSFI2VConfig,
workload_types=(WorkloadType.I2V, ),
hf_model_paths=[],
model_detectors=[
lambda path: any(token in path.lower() for token in (
"track-v0",
"wantrack",
"causaltrack",
)),
],
model_family="wantrack",
default_preset="sf_wantrack_causal_i2v",
pipeline_cls_name="WanTrackCausalDMDPipeline",
)
# MatrixGame 3.0 (I2V)
register_configs(
sampling_param_cls=MatrixGame3SamplingParam,
@@ -1279,6 +1299,8 @@ def _register_presets() -> None:
ALL_PRESETS as MATRIXGAME2_PRESETS, )
from fastvideo.pipelines.basic.matrixgame3.presets import (
ALL_PRESETS as MATRIXGAME3_PRESETS, )
from fastvideo.pipelines.basic.wantrack.presets import (
ALL_PRESETS as WANTRACK_PRESETS, )
from fastvideo.pipelines.basic.sd35.presets import (
ALL_PRESETS as SD35_PRESETS, )
from fastvideo.pipelines.basic.stable_audio.presets import (
@@ -1309,6 +1331,7 @@ def _register_presets() -> None:
LTX2_PRESETS,
MATRIXGAME2_PRESETS,
MATRIXGAME3_PRESETS,
WANTRACK_PRESETS,
SD35_PRESETS,
STABLE_AUDIO_PRESETS,
TURBODIFFUSION_PRESETS,
@@ -0,0 +1,264 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import json
from pathlib import Path
import numpy as np
import pyarrow as pa
import pyarrow.parquet as pq
import pytest
import fastvideo.dataset.parquet_dataset_streaming_style as streaming
from fastvideo.dataset.dataloader.schema import pyarrow_schema_t2v
def _row(index: int) -> dict:
latent = np.full((1, 1, 1, 1), index, dtype=np.float32)
embedding = np.full((2, 3), index, dtype=np.float32)
return {
"id": f"row-{index}",
"vae_latent_bytes": latent.tobytes(),
"vae_latent_shape": list(latent.shape),
"vae_latent_dtype": "float32",
"text_embedding_bytes": embedding.tobytes(),
"text_embedding_shape": list(embedding.shape),
"text_embedding_dtype": "float32",
"file_name": f"{index}.mp4",
"caption": f"caption {index}",
"media_type": "video",
"width": 832,
"height": 480,
"num_frames": 121,
"duration_sec": 5.0,
"fps": 24.0,
"unused_large_column": b"x" * 1024,
}
@pytest.fixture()
def parquet_root(tmp_path: Path) -> Path:
root = tmp_path / "shared-read-only-dataset"
root.mkdir()
schema = pyarrow_schema_t2v.append(pa.field("unused_large_column", pa.binary()))
pq.write_table(
pa.Table.from_pylist([_row(i) for i in range(12)], schema=schema), root / "part-0.parquet", row_group_size=3
)
pq.write_table(
pa.Table.from_pylist([_row(i) for i in range(12, 24)], schema=schema), root / "part-1.parquet", row_group_size=3
)
return root
def _patch_dist(monkeypatch: pytest.MonkeyPatch, rank: int = 0, world_size: int = 1, sp_size: int = 1) -> None:
monkeypatch.setattr(streaming.dist, "is_initialized", lambda: True)
monkeypatch.setattr(streaming, "get_world_rank", lambda: rank)
monkeypatch.setattr(streaming, "get_world_size", lambda: world_size)
monkeypatch.setattr(streaming, "get_sp_world_size", lambda: sp_size)
monkeypatch.setattr(streaming, "_barrier", lambda: None)
def _dataset(parquet_root: Path, cache_root: Path, **kwargs):
return streaming.LatentsParquetStreamingDataset(
path=str(parquet_root),
batch_size=2,
parquet_schema=pyarrow_schema_t2v,
manifest_path=str(cache_root / "openvid-manifest.json"),
num_workers=0,
text_padding_length=4,
read_batch_size=2,
shuffle_row_groups=False,
**kwargs,
)
def test_workers_zero_projects_columns_and_writes_json_manifest(
parquet_root: Path, tmp_path: Path, monkeypatch: pytest.MonkeyPatch
) -> None:
_patch_dist(monkeypatch)
recorded_columns = []
original = pq.ParquetFile.iter_batches
def recording_iter_batches(self, *args, **kwargs):
recorded_columns.append(tuple(kwargs["columns"]))
return original(self, *args, **kwargs)
monkeypatch.setattr(pq.ParquetFile, "iter_batches", recording_iter_batches)
dataset = _dataset(parquet_root, tmp_path / "owned-cache")
batch = next(iter(dataset))
assert batch["vae_latent"].shape[0] == 2
assert batch["info_list"][0]["id"] == "row-0"
assert recorded_columns == [tuple(pyarrow_schema_t2v.names)]
manifest = json.loads((tmp_path / "owned-cache" / "openvid-manifest.json").read_text())
assert manifest["total_rows"] == 24
assert manifest["columns"] == pyarrow_schema_t2v.names
assert not list((tmp_path / "owned-cache").glob("*.pkl"))
def test_state_dict_resumes_at_exact_next_batch(
parquet_root: Path, tmp_path: Path, monkeypatch: pytest.MonkeyPatch
) -> None:
_patch_dist(monkeypatch)
first_dataset = _dataset(parquet_root, tmp_path / "owned-cache")
iterator = iter(first_dataset)
assert next(iterator)["info_list"][0]["id"] == "row-0"
state = first_dataset.state_dict()
expected = next(iterator)["info_list"]
resumed_dataset = _dataset(parquet_root, tmp_path / "owned-cache")
resumed_dataset.load_state_dict(state)
actual = next(iter(resumed_dataset))["info_list"]
assert [item["id"] for item in actual] == [item["id"] for item in expected]
def test_reconstructed_state_matches_continuous_cursor(
parquet_root: Path, tmp_path: Path, monkeypatch: pytest.MonkeyPatch
) -> None:
_patch_dist(monkeypatch)
continuous = _dataset(parquet_root, tmp_path / "owned-cache")
iterator = iter(continuous)
for _ in range(3):
next(iterator)
expected_state = continuous.state_dict()
expected_next = next(iterator)["info_list"]
reconstructed_state = streaming.reconstruct_streaming_dataset_state(
continuous.manifest,
global_rank=0,
world_size=1,
sp_world_size=1,
num_workers=0,
batch_size=2,
read_batch_size=2,
seed=42,
shuffle_row_groups=False,
yielded_samples=6,
)
assert reconstructed_state == expected_state
resumed = _dataset(parquet_root, tmp_path / "owned-cache")
resumed.load_state_dict(reconstructed_state)
actual_next = next(iter(resumed))["info_list"]
assert [item["id"] for item in actual_next] == [item["id"] for item in expected_next]
def test_reconstructed_state_refuses_epoch_boundary(
parquet_root: Path, tmp_path: Path, monkeypatch: pytest.MonkeyPatch
) -> None:
_patch_dist(monkeypatch)
dataset = _dataset(parquet_root, tmp_path / "owned-cache")
with pytest.raises(ValueError, match="first epoch"):
streaming.reconstruct_streaming_dataset_state(
dataset.manifest,
global_rank=0,
world_size=1,
sp_world_size=1,
num_workers=0,
batch_size=2,
read_batch_size=2,
seed=42,
shuffle_row_groups=False,
yielded_samples=dataset.samples_per_worker,
)
def test_odd_row_group_tails_are_carried_into_next_group(
parquet_root: Path, tmp_path: Path, monkeypatch: pytest.MonkeyPatch
) -> None:
_patch_dist(monkeypatch)
dataset = _dataset(parquet_root, tmp_path / "owned-cache")
ids = [item["id"] for batch in dataset for item in batch["info_list"]]
assert ids == [f"row-{index}" for index in range(24)]
def test_dp_sp_shards_are_identical_within_sp_and_disjoint_across_dp(
parquet_root: Path, tmp_path: Path, monkeypatch: pytest.MonkeyPatch
) -> None:
ids_by_rank = {}
for rank in range(4):
_patch_dist(monkeypatch, rank=rank, world_size=4, sp_size=2)
dataset = _dataset(parquet_root, tmp_path / "owned-cache")
ids_by_rank[rank] = {item["id"] for batch in dataset for item in batch["info_list"]}
assert ids_by_rank[0] == ids_by_rank[1]
assert ids_by_rank[2] == ids_by_rank[3]
assert ids_by_rank[0].isdisjoint(ids_by_rank[2])
assert len(ids_by_rank[0] | ids_by_rank[2]) == 24
def test_equal_number_of_row_groups_and_dp_shards_stays_nonempty(
parquet_root: Path, tmp_path: Path, monkeypatch: pytest.MonkeyPatch
) -> None:
ids_by_rank = []
for rank in range(4):
_patch_dist(monkeypatch, rank=rank, world_size=4, sp_size=1)
dataset = _dataset(parquet_root, tmp_path / "owned-cache")
ids_by_rank.append({item["id"] for batch in dataset for item in batch["info_list"]})
assert all(ids for ids in ids_by_rank)
assert len(set().union(*ids_by_rank)) == 24
def test_manifest_must_not_be_written_inside_dataset(parquet_root: Path, monkeypatch: pytest.MonkeyPatch) -> None:
_patch_dist(monkeypatch)
with pytest.raises(ValueError, match="outside the read-only dataset"):
_dataset(parquet_root, parquet_root / "cache")
def test_stateful_dataloader_restores_dataset_cursor(
parquet_root: Path, tmp_path: Path, monkeypatch: pytest.MonkeyPatch
) -> None:
_, loader = streaming.build_parquet_streaming_style_dataloader(
path=str(parquet_root),
batch_size=2,
num_data_workers=0,
parquet_schema=pyarrow_schema_t2v,
manifest_path=str(tmp_path / "owned-cache" / "manifest.json"),
text_padding_length=4,
read_batch_size=2,
shuffle_row_groups=False,
)
iterator = iter(loader)
assert next(iterator)["info_list"][0]["id"] == "row-0"
state = loader.state_dict()
expected = next(iterator)["info_list"]
_, resumed = streaming.build_parquet_streaming_style_dataloader(
path=str(parquet_root),
batch_size=2,
num_data_workers=0,
parquet_schema=pyarrow_schema_t2v,
manifest_path=str(tmp_path / "owned-cache" / "manifest.json"),
text_padding_length=4,
read_batch_size=2,
shuffle_row_groups=False,
)
resumed.load_state_dict(state)
actual = next(iter(resumed))["info_list"]
assert [item["id"] for item in actual] == [item["id"] for item in expected]
def test_resume_rejects_changed_topology(parquet_root: Path, tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None:
_patch_dist(monkeypatch)
dataset = _dataset(parquet_root, tmp_path / "owned-cache")
next(iter(dataset))
state = dataset.state_dict()
state["topology"]["batch_size"] = 99
resumed = _dataset(parquet_root, tmp_path / "owned-cache")
with pytest.raises(ValueError, match="topology or sampling config changed"):
resumed.load_state_dict(state)
def test_manifest_is_rebuilt_when_source_file_changes(
parquet_root: Path, tmp_path: Path, monkeypatch: pytest.MonkeyPatch
) -> None:
_patch_dist(monkeypatch)
dataset = _dataset(parquet_root, tmp_path / "owned-cache")
old_fingerprint = dataset.manifest_fingerprint
source_file = parquet_root / "part-1.parquet"
table = pq.read_table(source_file)
pq.write_table(table, source_file, row_group_size=2)
refreshed = _dataset(parquet_root, tmp_path / "owned-cache")
assert refreshed.manifest_fingerprint != old_fingerprint
assert refreshed.manifest["total_rows"] == 24
+89 -41
View File
@@ -9,6 +9,7 @@ falls through to raw tensors for non-DTensor inputs.
"""
from __future__ import annotations
from pathlib import Path
from typing import Any
import pytest
@@ -16,7 +17,6 @@ import torch
from fastvideo.train.callbacks.ema import EMACallback
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
@@ -43,9 +43,16 @@ class _Method:
self,
transformer: torch.nn.Module | None,
tracker: Any | None = None,
ema_update_iterations: set[int] | None = None,
) -> None:
self.student = _Student(transformer)
self.tracker = tracker
self.ema_update_iterations = ema_update_iterations
def should_update_ema(self, iteration: int) -> bool:
if self.ema_update_iterations is None:
return True
return iteration in self.ema_update_iterations
def _tiny_transformer(*, fill: float = 0.0) -> torch.nn.Module:
@@ -89,9 +96,7 @@ class TestOnTrainingStepEnd:
def test_no_op_before_train_start(self) -> None:
cb = EMACallback()
# student_ema is None until on_train_start.
cb.on_training_step_end(
_Method(transformer=None), loss_dict={}, iteration=0
)
cb.on_training_step_end(_Method(transformer=None), loss_dict={}, iteration=0)
assert not cb._ema_started
def test_skipped_until_start_iter(self) -> None:
@@ -103,9 +108,7 @@ class TestOnTrainingStepEnd:
with torch.no_grad():
transformer.weight.fill_(7.0)
cb.on_training_step_end(
_Method(transformer), loss_dict={}, iteration=5
)
cb.on_training_step_end(_Method(transformer), loss_dict={}, iteration=5)
# Below start_iter: shadow is untouched, _ema_started False.
assert not cb._ema_started
assert torch.allclose(
@@ -122,9 +125,7 @@ class TestOnTrainingStepEnd:
with torch.no_grad():
transformer.weight.fill_(5.0)
cb.on_training_step_end(
_Method(transformer), loss_dict={}, iteration=10
)
cb.on_training_step_end(_Method(transformer), loss_dict={}, iteration=10)
# First active step: shadow is re-initialized from the
# current transformer (5.0) and *then* update() applies decay
# against the same value, so shadow stays at 5.0.
@@ -140,16 +141,12 @@ class TestOnTrainingStepEnd:
cb.on_train_start(_Method(transformer), iteration=0)
# Step 0: re-init at 2.0, then update against 2.0 → still 2.0.
cb.on_training_step_end(
_Method(transformer), loss_dict={}, iteration=0
)
cb.on_training_step_end(_Method(transformer), loss_dict={}, iteration=0)
# Step 1: drift transformer to 12.0, expect
# shadow = 0.9 * 2.0 + 0.1 * 12.0 = 3.0.
with torch.no_grad():
transformer.weight.fill_(12.0)
cb.on_training_step_end(
_Method(transformer), loss_dict={}, iteration=1
)
cb.on_training_step_end(_Method(transformer), loss_dict={}, iteration=1)
assert torch.allclose(
cb.student_ema.shadow["weight"],
torch.full((2, 4), 3.0),
@@ -164,9 +161,35 @@ class TestOnTrainingStepEnd:
cb.on_train_start(method, iteration=0)
cb.on_training_step_end(method, loss_dict={}, iteration=0)
assert any(
payload.get("ema/decay") == 0.99 and step == 0
for payload, step in tracker.entries
assert any(payload.get("ema/decay") == 0.99 and step == 0 for payload, step in tracker.entries)
def test_skips_steps_without_student_optimizer_update(self) -> None:
transformer = _tiny_transformer(fill=1.0)
method = _Method(
transformer,
ema_update_iterations={1, 5},
)
cb = EMACallback(decay=0.5, start_iter=0)
cb.on_train_start(method, iteration=0)
# First generator update initializes the active EMA at 1.0.
cb.on_training_step_end(method, loss_dict={}, iteration=1)
with torch.no_grad():
transformer.weight.fill_(9.0)
# Critic-only outer steps must not repeatedly decay EMA toward the
# unchanged generator weights.
for iteration in (2, 3, 4):
cb.on_training_step_end(method, loss_dict={}, iteration=iteration)
assert torch.allclose(
cb.student_ema.shadow["weight"],
torch.full((2, 4), 1.0),
)
cb.on_training_step_end(method, loss_dict={}, iteration=5)
assert torch.allclose(
cb.student_ema.shadow["weight"],
torch.full((2, 4), 5.0),
)
@@ -183,9 +206,7 @@ class TestEmaContext:
# No on_train_start → student_ema is None.
with cb.ema_context(transformer) as t:
assert t is transformer
assert torch.allclose(
t.weight, torch.full((2, 4), 3.0)
)
assert torch.allclose(t.weight, torch.full((2, 4), 3.0))
def test_swaps_weights_then_restores(self) -> None:
transformer = _tiny_transformer(fill=1.0)
@@ -203,9 +224,7 @@ class TestEmaContext:
with cb.ema_context(transformer) as t:
assert torch.allclose(t.weight, torch.full((2, 4), 1.0))
assert torch.allclose(
transformer.weight, torch.full((2, 4), 9.0)
)
assert torch.allclose(transformer.weight, torch.full((2, 4), 9.0))
# ---------------------------------------------------------------------------
@@ -219,7 +238,7 @@ class TestStateDict:
cb = EMACallback()
assert cb.state_dict() == {}
def test_round_trip_preserves_shadow_and_started_flag(self) -> None:
def test_state_dict_only_contains_dcp_safe_metadata(self) -> None:
transformer = _tiny_transformer(fill=4.0)
cb = EMACallback(decay=0.5, start_iter=0)
method = _Method(transformer)
@@ -227,24 +246,53 @@ class TestStateDict:
cb.on_training_step_end(method, loss_dict={}, iteration=0)
state = cb.state_dict()
assert "student_ema" in state
assert state["ema_started"] is True
assert state == {"ema_started": True}
def test_checkpoint_hooks_round_trip_rank_local_shadow(
self,
tmp_path: Path,
monkeypatch: pytest.MonkeyPatch,
) -> None:
transformer = _tiny_transformer(fill=4.0)
cb = EMACallback(decay=0.5, start_iter=0)
method = _Method(transformer)
cb.on_train_start(method, iteration=0)
cb.on_training_step_end(method, loss_dict={}, iteration=0)
expected = cb.student_ema.shadow["weight"].clone()
consolidated: list[tuple[int, str]] = []
def fake_consolidate(ema, module, rank, save_dir, base_name):
consolidated.append((rank, base_name))
monkeypatch.setattr(
"fastvideo.train.callbacks.ema."
"save_consolidated_local_shard_ema_safetensors",
fake_consolidate,
)
cb.on_checkpoint_save(method, tmp_path, iteration=10)
assert (tmp_path / "ema" / "local_shards" / "rank-0.pt").is_file()
assert consolidated == [(0, "student")]
# Build a fresh callback and load.
fresh = EMACallback(decay=0.5, start_iter=0)
fresh.on_train_start(_Method(_tiny_transformer(fill=0.0)),
iteration=0)
# Sanity: fresh shadow != saved shadow before load.
assert not torch.allclose(
fresh.student_ema.shadow["weight"],
cb.student_ema.shadow["weight"],
)
fresh.load_state_dict(state)
fresh_method = _Method(_tiny_transformer(fill=0.0))
fresh.on_train_start(fresh_method, iteration=0)
fresh.load_state_dict(cb.state_dict())
fresh.on_checkpoint_load(fresh_method, tmp_path, iteration=10)
assert fresh._ema_started is True
assert torch.allclose(
fresh.student_ema.shadow["weight"],
cb.student_ema.shadow["weight"],
)
assert torch.allclose(fresh.student_ema.shadow["weight"], expected)
def test_legacy_started_checkpoint_has_clear_error(
self,
tmp_path: Path,
) -> None:
method = _Method(_tiny_transformer())
cb = EMACallback()
cb.on_train_start(method, iteration=0)
cb.load_state_dict({"ema_started": True})
with pytest.raises(RuntimeError, match="predates the consolidated EMA"):
cb.on_checkpoint_load(method, tmp_path, iteration=10)
def test_load_without_student_ema_only_sets_flag(self) -> None:
cb = EMACallback()
@@ -165,6 +165,51 @@ class TestConstructor:
assert cb.metrics_config.log_prefix == "custom/validation"
class TestCausalTransformerConfig:
@staticmethod
def _configured_callback() -> ValidationCallback:
cb = _make_callback()
cb.training_config = SimpleNamespace(
pipeline_config=SimpleNamespace(
dit_config=SimpleNamespace(
local_attn_size=6,
sink_size=1,
rope_cache_policy="relativistic",
causal_train_attention="triton",
)
)
)
return cb
def test_matching_runtime_config_is_accepted(self) -> None:
cb = self._configured_callback()
transformer = SimpleNamespace(
local_attn_size=6,
sink_size=1,
rope_cache_policy="relativistic",
causal_train_attention="triton",
)
cb._assert_validation_transformer_config(transformer) # type: ignore[arg-type]
def test_mismatched_runtime_config_is_rejected(self) -> None:
cb = self._configured_callback()
transformer = SimpleNamespace(
local_attn_size=6,
sink_size=0,
rope_cache_policy="absolute",
causal_train_attention="flex",
)
with pytest.raises(ValueError, match="sink_size.*rope_cache_policy.*causal_train_attention"):
cb._assert_validation_transformer_config(transformer) # type: ignore[arg-type]
def test_non_causal_transformer_is_ignored(self) -> None:
cb = self._configured_callback()
cb._assert_validation_transformer_config(torch.nn.Linear(2, 2))
# ---------------------------------------------------------------------------
# B. on_validation_begin gating
# ---------------------------------------------------------------------------
@@ -0,0 +1,56 @@
# SPDX-License-Identifier: Apache-2.0
from pathlib import Path
import pytest
from fastvideo.train.entrypoint.dcp_to_diffusers import (
_export_consolidated_ema, )
def _make_base_model(path: Path) -> None:
transformer = path / "transformer"
transformer.mkdir(parents=True)
(path / "model_index.json").write_text("{}", encoding="utf-8")
(transformer / "config.json").write_text("{}", encoding="utf-8")
(transformer / "old.safetensors").write_bytes(b"old")
(transformer / "model.safetensors.index.json").write_text(
"{}",
encoding="utf-8",
)
def test_export_consolidated_ema_replaces_transformer_weights(
tmp_path: Path,
monkeypatch: pytest.MonkeyPatch,
) -> None:
base = tmp_path / "base"
_make_base_model(base)
ema_path = tmp_path / "student.safetensors"
ema_path.write_bytes(b"exact-ema")
monkeypatch.setattr(
"fastvideo.utils.maybe_download_model",
lambda _: str(base),
)
output = tmp_path / "output"
result = _export_consolidated_ema(
ema_path=str(ema_path),
base_model_path="unused",
output_dir=str(output),
)
assert result == str(output.resolve())
assert (output / "model_index.json").is_file()
assert (output / "transformer" / "config.json").is_file()
assert (output / "transformer" / "model.safetensors").read_bytes() == b"exact-ema"
assert not (output / "transformer" / "old.safetensors").exists()
assert not (output / "transformer" / "model.safetensors.index.json").exists()
def test_export_consolidated_ema_rejects_legacy_checkpoint(tmp_path: Path, ) -> None:
with pytest.raises(FileNotFoundError, match="Legacy checkpoints"):
_export_consolidated_ema(
ema_path=str(tmp_path / "missing.safetensors"),
base_model_path="unused",
output_dir=str(tmp_path / "output"),
)
@@ -0,0 +1,31 @@
# SPDX-License-Identifier: Apache-2.0
"""EMA cadence regression tests for alternating DMD2 optimization."""
from fastvideo.train.methods.distribution_matching.dmd2 import DMD2Method
def test_dmd2_ema_follows_generator_update_interval() -> None:
method = object.__new__(DMD2Method)
object.__setattr__(
method,
"method_config",
{"generator_update_interval": 5},
)
observed = [
DMD2Method.should_update_ema(method, iteration)
for iteration in range(1, 11)
]
assert observed == [
False,
False,
False,
False,
True,
False,
False,
False,
False,
True,
]
@@ -21,6 +21,7 @@ from pathlib import Path
import pytest
import torch
from fastvideo.train.methods.base import TrainingMethod
from fastvideo.train.methods.consistency_model.causal_cd import (
CausalConsistencyDistillationMethod, )
from fastvideo.train.models.wan import WanCausalModel
@@ -32,6 +33,34 @@ _FIXTURE = str(
/ "wan_causal_t2v_causal_cd_min.yaml")
def test_causal_cd_updates_ema_from_first_optimizer_step(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""``ema_start_step`` must not freeze the online CD target."""
method = object.__new__(CausalConsistencyDistillationMethod)
updates: list[int] = []
current_iteration = -1
def fake_super_step(_self: object, iteration: int) -> None:
nonlocal current_iteration
current_iteration = iteration
def fake_update_ema() -> None:
updates.append(current_iteration)
monkeypatch.setattr(
TrainingMethod,
"optimizers_schedulers_step",
fake_super_step,
)
object.__setattr__(method, "_update_ema", fake_update_ema)
for iteration in (1, 199, 200):
method.optimizers_schedulers_step(iteration)
assert updates == [1, 199, 200]
def _build_synthetic_batch(
device: torch.device,
dtype: torch.dtype,
@@ -20,6 +20,8 @@ from pathlib import Path
import pytest
import torch
from fastvideo.models.dits.causal_wanvideo import (
CausalWanSelfAttention, CausalWanTransformer3DModel)
from fastvideo.train.methods.fine_tuning.tfsft import (
TeacherForcingSFTMethod, )
from fastvideo.train.models.wan import WanCausalModel
@@ -31,6 +33,116 @@ _FIXTURE = str(
/ "wan_causal_t2v_tfsft_min.yaml")
class _BatchShapeConditionEmbedder(torch.nn.Module):
def __init__(self, hidden_size: int) -> None:
super().__init__()
self.hidden_size = hidden_size
self.encoder_batch_sizes: list[int] = []
def forward(self, timestep, encoder_hidden_states,
encoder_hidden_states_image):
del encoder_hidden_states_image
self.encoder_batch_sizes.append(int(encoder_hidden_states.shape[0]))
count = int(timestep.numel())
temb = timestep.new_zeros((count, self.hidden_size),
dtype=torch.float32)
timestep_proj = timestep.new_zeros(
(count, 6 * self.hidden_size), dtype=torch.float32)
return temb, timestep_proj, encoder_hidden_states, None
class _IdentityNormOut(torch.nn.Module):
def forward(self, hidden_states, shift, scale):
del shift, scale
return hidden_states
@pytest.mark.parametrize("forward_name", ["_forward_train", "_forward_inference"])
def test_causal_wan_forward_preserves_batch_dimension(
monkeypatch: pytest.MonkeyPatch, forward_name: str) -> None:
"""Both causal forward paths must pad and unpatchify every sample."""
batch_size = 2
hidden_size = 6
model = CausalWanTransformer3DModel.__new__(
CausalWanTransformer3DModel)
torch.nn.Module.__init__(model)
model.patch_size = (1, 1, 1)
model.hidden_size = hidden_size
model.num_attention_heads = 1
model.rope_cache_policy = "absolute"
model.text_len = 5
model.patch_embedding = torch.nn.Identity()
model.condition_embedder = _BatchShapeConditionEmbedder(hidden_size)
model.blocks = torch.nn.ModuleList()
model.gradient_checkpointing = False
model.scale_shift_table = torch.nn.Parameter(
torch.zeros(1, 2, hidden_size), requires_grad=False)
model.norm_out = _IdentityNormOut()
model.proj_out = torch.nn.Identity()
monkeypatch.setattr(model, "_get_train_attention_spec",
lambda **_kwargs: None)
monkeypatch.setattr(
"fastvideo.models.dits.causal_wanvideo.get_sp_world_size",
lambda: 1,
)
seen_grid_sizes: list[torch.Tensor] = []
def _unpatchify(tokens: torch.Tensor,
grid_sizes: torch.Tensor) -> list[torch.Tensor]:
seen_grid_sizes.append(grid_sizes.detach().clone())
return [tokens[i] for i in range(grid_sizes.shape[0])]
monkeypatch.setattr(model, "unpatchify", _unpatchify)
hidden_states = torch.randn(batch_size, hidden_size, 2, 2, 2)
encoder_hidden_states = torch.randn(batch_size, 3, 4)
timestep = torch.ones(batch_size, 2)
output = getattr(model, forward_name)(
hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
timestep=timestep,
)
assert model.condition_embedder.encoder_batch_sizes == [batch_size]
assert len(seen_grid_sizes) == 1
assert tuple(seen_grid_sizes[0].shape) == (batch_size, 3)
assert output.shape[0] == batch_size
@pytest.mark.parametrize("sequence_length", [127, 128])
def test_causal_wan_flex_attention_preserves_unpadded_length(
monkeypatch: pytest.MonkeyPatch, sequence_length: int) -> None:
"""FlexAttention must preserve lengths with zero or nonzero padding."""
attention = CausalWanSelfAttention.__new__(CausalWanSelfAttention)
torch.nn.Module.__init__(attention)
attention.rope_cache_policy = "absolute"
monkeypatch.setattr(
"fastvideo.models.dits.causal_wanvideo._apply_rotary_emb",
lambda tensor, *_args, **_kwargs: tensor,
)
monkeypatch.setattr(
"fastvideo.models.dits.causal_wanvideo.flex_attention",
lambda query, key, value, block_mask: query,
)
query = torch.randn(2, sequence_length, 1, 4)
output = attention(
q=query,
k=query,
v=query,
freqs_cis=(torch.empty(0), torch.empty(0)),
block_mask=object(),
)
assert output.shape == query.shape
torch.testing.assert_close(output, query)
def _build_synthetic_batch(
device: torch.device,
dtype: torch.dtype,
@@ -0,0 +1,323 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
from dataclasses import dataclass
from types import SimpleNamespace
import numpy as np
from PIL import Image
import pytest
import torch
from fastvideo.models.dits.trackwan.track_encoder import TrackEncoder
from fastvideo.models.dits.trackwan.model import _TrackConditioningMixin
from fastvideo.pipelines import TrainingBatch
from fastvideo.train.models.wantrack.control import (
HandleState,
StableGridController,
cosine_radius_weights,
deform_grid,
)
from fastvideo.train.models.wantrack.runtime import (
CausalWanTrackSession,
PreparedWanTrackInput,
WanTrackInferenceRuntime,
)
def test_track_encoder_window_matches_full_sequence() -> None:
torch.manual_seed(7)
encoder = TrackEncoder(
id_dim=8,
track_channels=4,
vae_temporal_compression=4,
)
coords = torch.rand(1, 17, 6, 2)
visibility = (torch.rand(1, 17, 6) > 0.2).float()
track_ids = torch.arange(6).unsqueeze(0)
full = encoder(
coords,
visibility,
latent_t=5,
latent_h=8,
latent_w=10,
track_ids=track_ids,
)
leading = encoder.forward_window(
coords[:, :1],
visibility[:, :1],
latent_start=0,
latent_t=1,
latent_h=8,
latent_w=10,
pixel_start=0,
track_ids=track_ids,
)
later = encoder.forward_window(
coords[:, 1:17],
visibility[:, 1:17],
latent_start=1,
latent_t=4,
latent_h=8,
latent_w=10,
pixel_start=1,
track_ids=track_ids,
)
torch.testing.assert_close(leading, full[:, :, :1])
torch.testing.assert_close(later, full[:, :, 1:5])
def test_preencoded_track_map_is_latent_aligned_and_exclusive() -> None:
condition = _TrackConditioningMixin()
condition.in_channels = 5
condition.track_channels = 1
hidden = torch.zeros(1, 4, 2, 3, 3)
track_map = torch.ones(1, 1, 2, 3, 3)
combined = condition._append_track_conditioning(
hidden,
track_points=None,
track_visibility=None,
track_ids=None,
track_map=track_map,
start_frame=0,
)
assert combined.shape == (1, 5, 2, 3, 3)
torch.testing.assert_close(combined[:, -1:], track_map)
with pytest.raises(ValueError, match="mutually exclusive"):
condition._append_track_conditioning(
hidden,
track_points=torch.zeros(1, 5, 1, 2),
track_visibility=torch.ones(1, 5, 1),
track_ids=torch.zeros(1, 1, dtype=torch.long),
track_map=track_map,
start_frame=0,
)
def test_radius_falloff_and_overlap_blending() -> None:
points = np.array([[0.5, 0.5], [0.6, 0.5], [0.9, 0.5]],
dtype=np.float32)
weights = cosine_radius_weights(points, (0.5, 0.5), 0.2)
assert weights[0] == pytest.approx(1.0)
assert 0.0 < weights[1] < 1.0
assert weights[2] == 0.0
handles = (
HandleState("left", (0.5, 0.5), (0.7, 0.5)),
HandleState("right", (0.5, 0.5), (0.3, 0.5)),
)
deformed = deform_grid(points[:1], handles, 0.2)
np.testing.assert_allclose(deformed, points[:1], atol=1e-6)
def test_control_revisions_are_monotonic_resampled_and_boundary_only() -> None:
controller = StableGridController(
[{"id": "main", "x": 0.5, "y": 0.5}],
radius=0.2,
)
initial = controller.render_constant(3)
assert controller.queue_revision(
2,
samples=[{
"id": "main",
"x": 0.7,
"y": 0.5,
"timestamp_ms": 50,
}],
add=[{
"id": "new",
"x": 0.2,
"y": 0.2,
}],
)
assert not controller.queue_revision(1, samples=[])
# Queuing does not mutate the already-rendered/committed prefix.
np.testing.assert_array_equal(initial.tracks,
controller.render_constant(3).tracks)
applied = controller.apply_pending(
3,
interval_start_ms=0,
interval_end_ms=100,
)
assert applied.revision == 2
assert applied.active_handle_ids == ("main", "new")
center_index = int(np.argmin(
np.linalg.norm(controller.grid - np.array([0.5, 0.5]), axis=1)))
assert applied.tracks[0, center_index, 0] == pytest.approx(
controller.grid[center_index, 0], abs=2e-2)
assert applied.tracks[-1, center_index, 0] > applied.tracks[
0, center_index, 0]
assert controller.queue_revision(3, remove=["main"], radius=0.1)
removed = controller.apply_pending(
2,
interval_start_ms=100,
interval_end_ms=200,
)
assert removed.active_handle_ids == ("new", )
assert removed.radius == pytest.approx(0.1)
@dataclass
class _FakeModel:
device: torch.device = torch.device("cpu")
class _FakeRuntime:
fps = 16.0
chunk_size = 3
temporal_compression = 4
dmd_denoising_steps = [1000, 750, 500, 250]
warp_denoising_step = True
def __init__(self) -> None:
self.model = _FakeModel()
self.clear_count = 0
self.encoded_windows: list[np.ndarray] = []
def prepare(self, image, prompt):
del image
batch = TrainingBatch(
latents=torch.zeros(1, 1, 1, 2, 2),
conditional_dict={
"track_ids": None,
"track_map": None,
"track_points": None,
"track_visibility": None,
},
)
return PreparedWanTrackInput(
image=Image.new("RGB", (16, 16)),
prompt=prompt,
batch=batch,
latent_channels=1,
latent_height=2,
latent_width=2,
)
def clear_state(self):
self.clear_count += 1
def new_vae_cache(self):
return []
def encode_track_window(self, **kwargs):
self.encoded_windows.append(kwargs["points"].copy())
return torch.zeros(1, 1, kwargs["latent_t"], 2, 2)
def decode_block(self, latents, *, cache, first):
del first
frames = np.repeat(
latents[:, :, :1, :1, :1].float().cpu().numpy().reshape(
-1, 1, 1, 1),
3,
axis=-1,
)
return frames.astype(np.float32), cache
def test_session_noise_is_deterministic_and_edits_are_future_only(
monkeypatch: pytest.MonkeyPatch,
) -> None:
import fastvideo.train.models.wantrack.runtime as runtime_module
monkeypatch.setattr(
runtime_module,
"sample_wantrack_block",
lambda model, batch, latents, **kwargs: latents,
)
def run_two_blocks(runtime: _FakeRuntime):
session = CausalWanTrackSession(runtime)
session.start(
runtime.prepare(b"", ""),
"",
[{"id": "h", "x": 0.5, "y": 0.5}],
{"seed": 9, "steps": 1},
radius=0.2,
)
first = session.generate_next_block()
first_history = session.committed_control_history
assert session.apply_control_revision(
1,
samples=[{
"id": "h",
"x": 0.7,
"y": 0.5,
"timestamp_ms": 1,
}],
received_at_ms=1,
)
second = session.generate_next_block()
np.testing.assert_array_equal(
first_history[0],
session.committed_control_history[0],
)
session.close()
return first.pixel_frames, second.pixel_frames
first_runtime = _FakeRuntime()
second_runtime = _FakeRuntime()
first_outputs = run_two_blocks(first_runtime)
second_outputs = run_two_blocks(second_runtime)
np.testing.assert_array_equal(first_outputs[0], second_outputs[0])
np.testing.assert_array_equal(first_outputs[1], second_outputs[1])
assert first_runtime.encoded_windows[1].shape[0] == 12
assert first_runtime.clear_count == 2
def test_session_sampling_error_clears_model_and_vae_state(
monkeypatch: pytest.MonkeyPatch,
) -> None:
import fastvideo.train.models.wantrack.runtime as runtime_module
def fail_sampling(*args, **kwargs):
raise RuntimeError("sampling failed")
monkeypatch.setattr(runtime_module, "sample_wantrack_block",
fail_sampling)
runtime = _FakeRuntime()
session = CausalWanTrackSession(runtime)
session.start(
runtime.prepare(b"", ""),
"",
[{"id": "h", "x": 0.5, "y": 0.5}],
{"seed": 9, "steps": 1},
)
with pytest.raises(RuntimeError, match="sampling failed"):
session.generate_next_block()
assert session.state == "failed"
assert runtime.clear_count == 2
def test_image_preprocessing_empty_prompt_fps_and_invalid_export(
tmp_path,
) -> None:
image = Image.new("RGB", (100, 50), "red")
processed = WanTrackInferenceRuntime.preprocess_image(
image,
width=32,
height=32,
)
assert processed.size == (32, 32)
assert WanTrackInferenceRuntime._fps_from_yaml({}) == 16.0
assert WanTrackInferenceRuntime._fps_from_yaml({
"callbacks": {
"track_validation": {
"fps": 24
}
}
}) == 24.0
model_dir = tmp_path / "model"
model_dir.mkdir()
(model_dir / "model_index.json").write_text(
'{"transformer": ["diffusers", "TrackWanTransformer3DModel"]}',
encoding="utf-8",
)
yaml_path = tmp_path / "config.yaml"
yaml_path.write_text("models: {}\n", encoding="utf-8")
with pytest.raises(ValueError, match="causal"):
WanTrackInferenceRuntime.from_export(model_dir, yaml_path)
+198 -9
View File
@@ -14,6 +14,10 @@ from pathlib import Path
from typing import Any
import pytest
import torch
import torch.distributed.checkpoint as dcp
import fastvideo.train.utils.checkpoint as checkpoint_module
from fastvideo.train.utils.checkpoint import (
CheckpointConfig,
@@ -24,7 +28,6 @@ from fastvideo.train.utils.checkpoint import (
_resolve_resume_checkpoint,
)
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
@@ -87,6 +90,44 @@ class _MissingLoad:
return {}
class _Method:
def checkpoint_state(self) -> dict[str, Any]:
return {"method": _Full()}
class _RecordingStateful:
def __init__(self, state: dict[str, Any] | None = None) -> None:
self.state = dict(state or {})
self.loaded: dict[str, Any] | None = None
def state_dict(self) -> dict[str, Any]:
return dict(self.state)
def load_state_dict(self, state: dict[str, Any]) -> None:
self.loaded = dict(state)
class _LoaderWithDataset(_RecordingStateful):
def __init__(self) -> None:
super().__init__()
self.dataset = _RecordingStateful()
class _TensorStateful:
def __init__(self, value: int) -> None:
self.value = torch.tensor([value], dtype=torch.int64)
def state_dict(self) -> dict[str, Any]:
return {"value": self.value}
def load_state_dict(self, state: dict[str, Any]) -> None:
self.value = state["value"]
def test_is_stateful_true_for_full_object() -> None:
assert _is_stateful(_Full()) is True
@@ -99,6 +140,150 @@ def test_is_stateful_false_when_missing_load_state_dict() -> None:
assert _is_stateful(_MissingLoad()) is False
def test_dataloader_is_not_part_of_shared_dcp_state(tmp_path: Path) -> None:
manager = CheckpointManager(
method=_Method(),
dataloader=_RecordingStateful({"cursor": 7}),
output_dir=str(tmp_path),
config=CheckpointConfig(save_steps=1, keep_last=0),
)
assert set(manager._build_states()) == {"method"}
def test_dcp_partial_load_ignores_legacy_dataloader_key(tmp_path: Path) -> None:
checkpoint_dir = tmp_path / "dcp"
dcp.save(
{
"method": _TensorStateful(7),
"dataloader": _TensorStateful(99),
},
checkpoint_id=str(checkpoint_dir),
)
restored = _TensorStateful(-1)
dcp.load({"method": restored}, checkpoint_id=str(checkpoint_dir))
assert restored.value.item() == 7
def test_legacy_single_rank_dataloader_dcp_fallback(tmp_path: Path) -> None:
checkpoint_dir = _make_checkpoint_dir(tmp_path, 10)
dcp.save(
{"dataloader": _TensorStateful(42)},
checkpoint_id=str(checkpoint_dir / "dcp"),
)
restored = _TensorStateful(-1)
manager = CheckpointManager(
method=_Method(),
dataloader=restored,
output_dir=str(tmp_path),
config=CheckpointConfig(save_steps=1, keep_last=0),
)
manager._load_dataloader_snapshot(checkpoint_dir, 10)
assert restored.value.item() == 42
def test_rank_local_dataloader_sidecar_roundtrip(
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.setattr(checkpoint_module, "_rank", lambda: 3)
monkeypatch.setattr(checkpoint_module, "_world_size", lambda: 8)
checkpoint_dir = _make_checkpoint_dir(tmp_path, 20)
original = _RecordingStateful({"cursor": 11})
manager = CheckpointManager(
method=_Method(),
dataloader=original,
output_dir=str(tmp_path),
config=CheckpointConfig(save_steps=1, keep_last=0),
)
manager._save_dataloader_snapshot(checkpoint_dir, 20)
path = checkpoint_dir / "dataloader_state_rank3.pt"
assert path.is_file()
assert not list(checkpoint_dir.glob(".*.tmp"))
restored = _RecordingStateful()
restored_manager = CheckpointManager(
method=_Method(),
dataloader=restored,
output_dir=str(tmp_path),
config=CheckpointConfig(save_steps=1, keep_last=0),
)
restored_manager._load_dataloader_snapshot(checkpoint_dir, 20)
assert restored.loaded == {"cursor": 11}
def test_dataset_only_migration_sidecar_uses_loader_dataset(
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.setattr(checkpoint_module, "_rank", lambda: 1)
monkeypatch.setattr(checkpoint_module, "_world_size", lambda: 4)
checkpoint_dir = _make_checkpoint_dir(tmp_path, 10)
migration_dir = tmp_path / "migration"
migration_dir.mkdir()
torch.save(
{
"version": 1,
"rank": 1,
"world_size": 4,
"step": 10,
"state_kind": "dataset",
"state": {"cursor": 99},
},
migration_dir / "dataloader_state_rank1.pt",
)
monkeypatch.setenv("FASTVIDEO_DATALOADER_STATE_DIR", str(migration_dir))
loader = _LoaderWithDataset()
manager = CheckpointManager(
method=_Method(),
dataloader=loader,
output_dir=str(tmp_path),
config=CheckpointConfig(save_steps=1, keep_last=0),
)
manager._load_dataloader_snapshot(checkpoint_dir, 10)
assert loader.loaded is None
assert loader.dataset.loaded == {"cursor": 99}
def test_dataloader_sidecar_metadata_mismatch_is_rejected(
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.setattr(checkpoint_module, "_rank", lambda: 0)
monkeypatch.setattr(checkpoint_module, "_world_size", lambda: 2)
checkpoint_dir = _make_checkpoint_dir(tmp_path, 10)
torch.save(
{
"version": 1,
"rank": 1,
"world_size": 2,
"step": 10,
"state_kind": "dataloader",
"state": {},
},
checkpoint_dir / "dataloader_state_rank0.pt",
)
manager = CheckpointManager(
method=_Method(),
dataloader=_RecordingStateful(),
output_dir=str(tmp_path),
config=CheckpointConfig(save_steps=1, keep_last=0),
)
with pytest.raises(ValueError, match="rank mismatch"):
manager._load_dataloader_snapshot(checkpoint_dir, 10)
def test_missing_rank_local_dataloader_sidecar_fails_closed(
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.setattr(checkpoint_module, "_world_size", lambda: 2)
checkpoint_dir = _make_checkpoint_dir(tmp_path, 10)
manager = CheckpointManager(
method=_Method(),
dataloader=_RecordingStateful(),
output_dir=str(tmp_path),
config=CheckpointConfig(save_steps=1, keep_last=0),
)
with pytest.raises(RuntimeError, match="Legacy multi-rank DCP dataloader state is unsafe"):
manager._load_dataloader_snapshot(checkpoint_dir, 10)
# ---------------------------------------------------------------------------
# B. _parse_step_from_dir
# ---------------------------------------------------------------------------
@@ -163,8 +348,7 @@ def test_find_latest_skips_non_checkpoint_dirs(tmp_path: Path) -> None:
# ---------------------------------------------------------------------------
def test_resolve_latest_with_no_checkpoints_returns_none(
tmp_path: Path) -> None:
def test_resolve_latest_with_no_checkpoints_returns_none(tmp_path: Path) -> None:
out = tmp_path / "outputs"
out.mkdir()
assert _resolve_resume_checkpoint("latest", output_dir=str(out)) is None
@@ -180,8 +364,7 @@ def test_resolve_latest_returns_latest_checkpoint(tmp_path: Path) -> None:
def test_resolve_explicit_checkpoint_dir(tmp_path: Path) -> None:
ckpt = _make_checkpoint_dir(tmp_path, 42)
resolved = _resolve_resume_checkpoint(str(ckpt),
output_dir=str(tmp_path))
resolved = _resolve_resume_checkpoint(str(ckpt), output_dir=str(tmp_path))
assert resolved is not None
assert resolved.name == "checkpoint-42"
@@ -189,8 +372,7 @@ def test_resolve_explicit_checkpoint_dir(tmp_path: Path) -> None:
def test_resolve_dcp_subdir_returns_parent_checkpoint(tmp_path: Path) -> None:
ckpt = _make_checkpoint_dir(tmp_path, 42)
dcp_path = ckpt / "dcp"
resolved = _resolve_resume_checkpoint(str(dcp_path),
output_dir=str(tmp_path))
resolved = _resolve_resume_checkpoint(str(dcp_path), output_dir=str(tmp_path))
assert resolved is not None
assert resolved.name == "checkpoint-42"
@@ -207,8 +389,7 @@ def test_resolve_output_dir_returns_latest(tmp_path: Path) -> None:
def test_resolve_nonexistent_path_raises(tmp_path: Path) -> None:
with pytest.raises(FileNotFoundError):
_resolve_resume_checkpoint(str(tmp_path / "missing"),
output_dir=str(tmp_path))
_resolve_resume_checkpoint(str(tmp_path / "missing"), output_dir=str(tmp_path))
def test_resolve_checkpoint_without_dcp_raises(tmp_path: Path) -> None:
@@ -358,3 +539,11 @@ def test_maybe_save_triggers_on_each_interval(tmp_path: Path) -> None:
for step in range(1, 41):
mgr.maybe_save(step=step)
assert calls == [10, 20, 30, 40]
def test_save_final_skips_step_already_saved(tmp_path: Path) -> None:
mgr = _make_manager(tmp_path, save_steps=10)
calls = _record_save_calls(mgr)
mgr.maybe_save(step=20)
mgr.save_final(step=20)
assert calls == [20]
+24 -13
View File
@@ -1,5 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
"""CPU-only unit tests for :func:`load_run_config`."""
from __future__ import annotations
from pathlib import Path
@@ -27,8 +28,7 @@ def _minimal_yaml() -> dict[str, Any]:
},
},
"method": {
"_target_":
"fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod",
"_target_": "fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod",
},
"training": {},
}
@@ -44,8 +44,7 @@ def test_minimal_yaml_loads_happy_path(tmp_path: Path) -> None:
assert isinstance(cfg, RunConfig)
assert isinstance(cfg.training, TrainingConfig)
assert cfg.method["_target_"] == (
"fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod")
assert cfg.method["_target_"] == ("fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod")
assert "student" in cfg.models
assert cfg.callbacks == {}
# raw retains the original YAML dict for downstream logging.
@@ -64,6 +63,10 @@ def test_minimal_yaml_applies_all_defaults(tmp_path: Path) -> None:
assert t.data.train_batch_size == 1
assert t.data.dataloader_num_workers == 0
assert t.data.dataloader_type == "map"
assert t.data.streaming_manifest_path == ""
assert t.data.streaming_read_batch_size == 8
assert t.data.streaming_shuffle_row_groups is True
assert t.data.training_cfg_rate == 0.0
assert t.data.seed == 0
@@ -107,6 +110,10 @@ def test_full_yaml_populates_all_training_fields(tmp_path: Path) -> None:
"data_path": "/some/path",
"train_batch_size": 2,
"dataloader_num_workers": 4,
"dataloader_type": "streaming",
"streaming_manifest_path": "/owned/index/manifest.json",
"streaming_read_batch_size": 2,
"streaming_shuffle_row_groups": False,
"training_cfg_rate": 0.1,
"seed": 42,
"num_height": 256,
@@ -136,9 +143,7 @@ def test_full_yaml_populates_all_training_fields(tmp_path: Path) -> None:
"project_name": "myproj",
"run_name": "myrun",
},
"vsa": {
"sparsity": 0.5
},
"vsa": {"sparsity": 0.5},
"model": {
"weighting_scheme": "logit_normal",
"logit_mean": 0.5,
@@ -156,6 +161,10 @@ def test_full_yaml_populates_all_training_fields(tmp_path: Path) -> None:
assert t.data.train_batch_size == 2
assert t.data.data_path == "/some/path"
assert t.data.dataloader_type == "streaming"
assert t.data.streaming_manifest_path == "/owned/index/manifest.json"
assert t.data.streaming_read_batch_size == 2
assert t.data.streaming_shuffle_row_groups is False
assert t.data.num_frames == 33
assert t.data.seed == 42
@@ -230,10 +239,13 @@ def test_missing_config_file_raises(tmp_path: Path) -> None:
# ---------------------------------------------------------------------------
@pytest.mark.parametrize("betas_value, expected", [
([0.8, 0.9], (0.8, 0.9)),
("0.9,0.999", (0.9, 0.999)),
])
@pytest.mark.parametrize(
"betas_value, expected",
[
([0.8, 0.9], (0.8, 0.9)),
("0.9,0.999", (0.9, 0.999)),
],
)
def test_betas_parses_list_and_string_forms(
tmp_path: Path,
betas_value: Any,
@@ -347,8 +359,7 @@ def test_callbacks_passed_through_when_present(tmp_path: Path) -> None:
data = _minimal_yaml()
data["callbacks"] = {
"grad_clip": {
"_target_":
"fastvideo.train.callbacks.grad_clip.GradNormClipCallback",
"_target_": "fastvideo.train.callbacks.grad_clip.GradNormClipCallback",
"max_grad_norm": 1.0,
},
}
+17
View File
@@ -7,6 +7,7 @@ Adapted from FastGen's callback pattern to FastVideo's types.
from __future__ import annotations
from collections.abc import Callable
from pathlib import Path
from typing import Any, TYPE_CHECKING
from fastvideo.logger import init_logger
@@ -87,6 +88,22 @@ class Callback:
) -> None:
pass
def on_checkpoint_save(
self,
method: TrainingMethod,
checkpoint_dir: Path,
iteration: int = 0,
) -> None:
pass
def on_checkpoint_load(
self,
method: TrainingMethod,
checkpoint_dir: Path,
iteration: int = 0,
) -> None:
pass
def state_dict(self) -> dict[str, Any]:
return {}
+79 -5
View File
@@ -9,14 +9,20 @@ lives under ``callbacks.ema`` in the YAML file.
from __future__ import annotations
import contextlib
import os
from collections.abc import Generator
from pathlib import Path
from typing import Any, TYPE_CHECKING
import torch
import torch.distributed as dist
from fastvideo.logger import init_logger
from fastvideo.train.callbacks.callback import Callback
from fastvideo.training.training_utils import EMA_FSDP
from fastvideo.training.training_utils import (
EMA_FSDP,
save_consolidated_local_shard_ema_safetensors,
)
if TYPE_CHECKING:
from fastvideo.train.methods.base import TrainingMethod
@@ -94,6 +100,8 @@ class EMACallback(Callback):
if iteration < self._start_iter:
return
if not method.should_update_ema(iteration):
return
if not self._ema_started:
logger.info(
"Starting EMA updates at iteration %d "
@@ -140,7 +148,6 @@ class EMACallback(Callback):
if self.student_ema is None:
return {}
return {
"student_ema": self.student_ema.state_dict(),
"ema_started": self._ema_started,
}
@@ -148,7 +155,74 @@ class EMACallback(Callback):
self,
state_dict: dict[str, Any],
) -> None:
ema_state = state_dict.get("student_ema")
if (ema_state is not None and self.student_ema is not None):
self.student_ema.load_state_dict(ema_state)
self._ema_started = bool(state_dict.get("ema_started", False), )
def on_checkpoint_save(
self,
method: TrainingMethod,
checkpoint_dir: Path,
iteration: int = 0,
) -> None:
if self.student_ema is None or not self._ema_started:
return
rank = int(dist.get_rank()) if dist.is_initialized() else 0
world_size = (int(dist.get_world_size()) if dist.is_initialized() else 1)
shard_dir = checkpoint_dir / "ema" / "local_shards"
shard_dir.mkdir(parents=True, exist_ok=True)
shard_path = shard_dir / f"rank-{rank}.pt"
temp_path = shard_path.with_suffix(".pt.tmp")
torch.save(
{
"version": 1,
"rank": rank,
"world_size": world_size,
"student_ema": self.student_ema.state_dict(),
},
temp_path,
)
os.replace(temp_path, shard_path)
save_consolidated_local_shard_ema_safetensors(
self.student_ema,
method.student.transformer,
rank,
str(checkpoint_dir),
"student",
)
def on_checkpoint_load(
self,
method: TrainingMethod,
checkpoint_dir: Path,
iteration: int = 0,
) -> None:
if not self._ema_started:
return
if self.student_ema is None:
raise RuntimeError("EMA is active in the checkpoint but was not initialized")
rank = int(dist.get_rank()) if dist.is_initialized() else 0
world_size = (int(dist.get_world_size()) if dist.is_initialized() else 1)
shard_path = checkpoint_dir / "ema" / "local_shards" / f"rank-{rank}.pt"
if not shard_path.is_file():
raise RuntimeError("Checkpoint EMA cannot be resumed: missing rank-local EMA "
f"state at {shard_path}. This checkpoint predates the "
"consolidated EMA format; its DCP callback tensors contain "
"incomplete rank-local shards and cannot be reconstructed "
"losslessly.")
payload = torch.load(
shard_path,
map_location="cpu",
weights_only=True,
)
saved_world_size = int(payload.get("world_size", -1))
if saved_world_size != world_size:
raise RuntimeError("EMA checkpoint topology mismatch: "
f"saved world_size={saved_world_size}, current world_size={world_size}.")
ema_state = payload.get("student_ema")
if not isinstance(ema_state, dict):
raise RuntimeError(f"Invalid EMA shard payload: {shard_path}")
self.student_ema.load_state_dict(ema_state)
logger.info("Restored rank-local EMA state from %s", shard_path)
@@ -0,0 +1,305 @@
# SPDX-License-Identifier: Apache-2.0
"""Track-conditioned validation shared by bidirectional and causal WanTrack."""
from __future__ import annotations
import colorsys
import glob
import os
from typing import Any
import imageio.v2 as imageio
import numpy as np
import pyarrow.parquet as pq
import torch
from PIL import Image, ImageDraw
from fastvideo.distributed import get_world_group
from fastvideo.logger import init_logger
from fastvideo.train.callbacks.callback import Callback
from fastvideo.train.models.wantrack.inference import (
prepare_wantrack_batch,
sample_wantrack,
)
from fastvideo.training.trackers import DummyTracker
logger = init_logger(__name__)
def _track_colors(count: int) -> np.ndarray:
colors = np.empty((count, 3), dtype=np.uint8)
for index in range(count):
hue = index / max(count, 1)
red, green, blue = colorsys.hsv_to_rgb(hue, 0.9, 1.0)
colors[index] = (
int(red * 255),
int(green * 255),
int(blue * 255),
)
return colors
def _overlay_tracks(
frames: np.ndarray,
track_points: torch.Tensor,
track_visibility: torch.Tensor,
*,
stride: int,
tail: int,
radius: int,
visibility_threshold: float,
) -> list[np.ndarray]:
points = track_points[0].float().cpu().numpy()
visibility = track_visibility[0].float().cpu().numpy()
frame_count = min(
int(frames.shape[0]),
int(points.shape[0]),
int(visibility.shape[0]),
)
frames = frames[:frame_count]
points = points[:frame_count, ::stride]
visibility = visibility[:frame_count, ::stride]
colors = _track_colors(int(points.shape[1]))
height, width = int(frames.shape[1]), int(frames.shape[2])
points = points.copy()
points[..., 0] *= width
points[..., 1] *= height
output: list[np.ndarray] = []
for frame_index, frame in enumerate(frames):
image = Image.fromarray(frame).convert("RGB")
draw = ImageDraw.Draw(image)
start = max(0, frame_index - tail)
for track_index, color_value in enumerate(colors):
color = tuple(int(value) for value in color_value)
trail: list[tuple[float, float]] = []
for history_index in range(start, frame_index + 1):
if visibility[history_index, track_index] < visibility_threshold:
trail = []
continue
x, y = points[history_index, track_index]
if 0 <= x < width and 0 <= y < height:
trail.append((float(x), float(y)))
if len(trail) >= 2:
draw.line(trail, fill=color, width=1)
if visibility[frame_index, track_index] >= visibility_threshold:
x, y = points[frame_index, track_index]
if 0 <= x < width and 0 <= y < height:
draw.ellipse(
(
float(x - radius),
float(y - radius),
float(x + radius),
float(y + radius),
),
fill=color,
)
output.append(np.asarray(image))
return output
class TrackValidationCallback(Callback):
"""Generate fixed track-conditioned samples through the live student."""
def __init__(
self,
*,
every_steps: int = 250,
val_data_path: str | None = None,
num_val_samples: int = 1,
num_inference_steps: int = 30,
guidance_scale: float = 1.0,
motion_guidance_scale: float = 1.0,
output_dir: str | None = None,
fps: int = 24,
grid_stride: int = 3,
tail: int = 12,
radius: int = 2,
visibility_threshold: float = 0.5,
seed: int = 1000,
validate_at_start: bool = False,
) -> None:
self.every_steps = int(every_steps)
self.val_data_path = (str(val_data_path) if val_data_path is not None else None)
self.num_val_samples = int(num_val_samples)
self.num_inference_steps = int(num_inference_steps)
self.guidance_scale = float(guidance_scale)
self.motion_guidance_scale = float(motion_guidance_scale)
self.output_dir = (str(output_dir) if output_dir is not None else None)
self.fps = int(fps)
self.grid_stride = int(grid_stride)
self.tail = int(tail)
self.radius = int(radius)
self.visibility_threshold = float(visibility_threshold)
self.seed = int(seed)
self.validate_at_start = bool(validate_at_start)
if self.every_steps <= 0:
raise ValueError("every_steps must be positive")
if self.num_val_samples <= 0:
raise ValueError("num_val_samples must be positive")
if self.num_inference_steps <= 0:
raise ValueError("num_inference_steps must be positive")
if self.grid_stride <= 0:
raise ValueError("grid_stride must be positive")
self.tracker: Any = DummyTracker()
self._samples: list[dict[str, Any]] = []
self._is_main = False
self._did_start_validation = False
def on_train_start(self, method: Any, iteration: int = 0) -> None:
del iteration
self.tracker = getattr(method, "tracker", None) or DummyTracker()
self._is_main = int(get_world_group().rank) == 0
try:
self._samples = self._load_samples()
logger.info(
"WanTrack validation loaded %d fixed sample(s)",
len(self._samples),
)
except Exception as error: # noqa: BLE001
logger.warning(
"WanTrack validation setup failed; disabling callback: %s",
error,
)
self._samples = []
def on_validation_begin(self, method: Any, iteration: int = 0) -> None:
if not self._samples:
return
run_at_start = (self.validate_at_start and not self._did_start_validation and iteration == 0)
run_periodic = iteration > 0 and iteration % self.every_steps == 0
if not run_at_start and not run_periodic:
return
self._did_start_validation = True
try:
self._run(method, iteration)
except Exception as error: # noqa: BLE001
logger.warning(
"WanTrack validation failed at step %d: %s",
iteration,
error,
)
def _load_samples(self) -> list[dict[str, Any]]:
from fastvideo.dataset.dataloader.schema import (
pyarrow_schema_i2v_track, )
from fastvideo.dataset.utils import (
collate_rows_from_parquet_schema, )
data_path = self.val_data_path or str(self.training_config.data.data_path)
files = sorted(glob.glob(
os.path.join(data_path, "**", "*.parquet"),
recursive=True,
))
if not files:
raise FileNotFoundError(f"No WanTrack validation parquet under {data_path}")
rows: list[dict[str, Any]] = []
for path in files:
table = pq.read_table(
path,
columns=pyarrow_schema_i2v_track.names,
)
rows.extend(table.to_pylist())
if len(rows) >= self.num_val_samples:
break
text_len = int(self.training_config.pipeline_config.text_encoder_configs[0].arch_config.text_len)
collated = collate_rows_from_parquet_schema(
rows[:self.num_val_samples],
pyarrow_schema_i2v_track,
text_padding_length=text_len,
cfg_rate=0.0,
seed=self.seed,
)
infos = collated.get("info_list") or [{} for _ in rows]
samples: list[dict[str, Any]] = []
batch_size = int(collated["text_embedding"].shape[0])
for index in range(batch_size):
sample = {
name: value[index:index + 1].clone()
for name, value in collated.items() if torch.is_tensor(value)
}
sample["info_list"] = [infos[index]]
samples.append(sample)
return samples
@torch.no_grad()
def _run(self, method: Any, iteration: int) -> None:
student = method.student
transformer = student.transformer
was_training = bool(transformer.training)
transformer.eval()
try:
artifacts: list[Any] = []
for sample_index, raw_sample in enumerate(self._samples):
seed = self.seed + sample_index
batch = prepare_wantrack_batch(
student,
raw_sample,
seed=seed,
latents_source="zeros",
)
sampled = sample_wantrack(
student,
batch,
num_inference_steps=self.num_inference_steps,
seed=seed,
text_guidance_scale=self.guidance_scale,
motion_guidance_scale=self.motion_guidance_scale,
)
decoded = student.decode_latents(sampled)
if not self._is_main:
continue
output_dir = self.output_dir or os.path.join(
self.training_config.checkpoint.output_dir,
"track_validation",
)
os.makedirs(output_dir, exist_ok=True)
output_path = os.path.join(
output_dir,
f"step{iteration:06d}_sample{sample_index}.mp4",
)
video = (decoded[0].clamp(0, 1).float().cpu().numpy() * 255.0).astype(np.uint8)
frames = video.transpose(1, 2, 3, 0)
conditional = batch.conditional_dict
if conditional is None:
raise RuntimeError("WanTrack validation batch has no conditioning")
frames_with_tracks = _overlay_tracks(
frames,
conditional["track_points"],
conditional["track_visibility"],
stride=self.grid_stride,
tail=self.tail,
radius=self.radius,
visibility_threshold=self.visibility_threshold,
)
imageio.mimsave(
output_path,
frames_with_tracks,
fps=self.fps,
macro_block_size=1,
)
info = raw_sample.get("info_list") or [{}]
caption = str(info[0].get("caption", ""))
artifact = self.tracker.video(
output_path,
caption=(f"WanTrack step {iteration}: "
f"{caption[:120]}"),
)
if artifact is not None:
artifacts.append(artifact)
if self._is_main and artifacts:
self.tracker.log_artifacts(
{"track_validation": artifacts},
iteration,
)
finally:
if was_training:
transformer.train()
+44
View File
@@ -1071,6 +1071,7 @@ class ValidationCallback(Callback):
return self._pipeline
tc = self.training_config
self._assert_validation_transformer_config(transformer)
PipelineCls = resolve_target(self.pipeline_target)
flow_shift = getattr(
tc.pipeline_config,
@@ -1116,6 +1117,49 @@ class ValidationCallback(Callback):
self._pipeline_key = key
return self._pipeline
def _assert_validation_transformer_config(
self,
transformer: torch.nn.Module,
) -> None:
"""Ensure causal validation reuses the exact training attention setup."""
dit_config = self.training_config.pipeline_config.dit_config
fields = (
"local_attn_size",
"sink_size",
"rope_cache_policy",
"causal_train_attention",
)
runtime_values = {
field: getattr(transformer, field)
for field in fields
if hasattr(transformer, field)
}
if not runtime_values:
return
expected_values = {field: getattr(dit_config, field) for field in fields}
mismatches = {
field: (expected_values[field], runtime_values.get(field, "<missing>"))
for field in fields
if runtime_values.get(field, "<missing>") != expected_values[field]
}
if mismatches:
details = ", ".join(
f"{field}: config={expected!r}, transformer={runtime!r}"
for field, (expected, runtime) in mismatches.items()
)
raise ValueError(
"Validation transformer does not match the training causal "
f"attention config ({details})."
)
logger.info(
"Validation reuses training transformer config: "
"local_attn_size=%s, sink_size=%s, rope_cache_policy=%s, "
"causal_train_attention=%s",
*(runtime_values[field] for field in fields),
)
# ----------------------------------------------------------
# Batch preparation
# ----------------------------------------------------------
+85 -11
View File
@@ -51,6 +51,57 @@ def _ensure_distributed() -> None:
os.environ.setdefault(key, default)
def _export_consolidated_ema(
*,
ema_path: str,
base_model_path: str,
output_dir: str,
overwrite: bool = False,
) -> str:
"""Build a diffusers directory using a checkpoint's full EMA weights."""
import shutil
from pathlib import Path
from fastvideo.utils import maybe_download_model
src_ema = Path(ema_path).resolve()
if not src_ema.is_file():
raise FileNotFoundError("Consolidated student EMA is missing from this checkpoint: "
f"{src_ema}. Legacy checkpoints only contain incomplete "
"rank-local EMA callback tensors and cannot export exact EMA weights.")
local_base = Path(maybe_download_model(str(base_model_path))).resolve()
dst = Path(os.path.expanduser(str(output_dir))).resolve()
if dst.exists():
if overwrite:
shutil.rmtree(dst, ignore_errors=True)
else:
raise FileExistsError(f"Refusing to overwrite existing directory: {dst}. "
"Pass --overwrite to replace it.")
def _copy_or_link(src: str, dest: str) -> None:
try:
os.link(src, dest)
except OSError:
shutil.copy2(src, dest)
shutil.copytree(
local_base,
dst,
symlinks=True,
copy_function=_copy_or_link,
)
transformer_dir = dst / "transformer"
if not transformer_dir.is_dir():
raise FileNotFoundError(f"Base model is missing transformer component: {transformer_dir}")
for pattern in ("*.safetensors", "*.bin", "*.index.json"):
for path in transformer_dir.glob(pattern):
path.unlink(missing_ok=True)
shutil.copy2(src_ema, transformer_dir / "model.safetensors")
logger.info("Exported consolidated student EMA to %s", dst)
return str(dst)
def _save_role_pretrained(
*,
role: str,
@@ -214,6 +265,7 @@ def convert(
from fastvideo.train.utils.builder import build_from_config
from fastvideo.train.utils.checkpoint import (
CheckpointManager,
_FullModelState,
_resolve_resume_checkpoint,
)
from fastvideo.train.utils.config import (
@@ -247,6 +299,20 @@ def convert(
tc = cfg.training
base_model_path = str(tc.model_path)
if not base_model_path:
raise ValueError("Cannot determine base_model_path from "
"config. Ensure models.student.init_from "
"is set.")
if role == "student_ema":
return _export_consolidated_ema(
ema_path=str(resolved / "ema" / "student.safetensors"),
base_model_path=base_model_path,
output_dir=output_dir,
overwrite=overwrite,
)
# -- Init distributed (1 GPU is enough; DCP reshards) --
maybe_init_distributed_environment_and_model_parallel(
tp_size=1,
@@ -263,22 +329,29 @@ def convert(
# -- Build model (loads pretrained weights + FSDP) --
_, method, _, _ = build_from_config(cfg)
# -- Load DCP weights into the model --
states = method.checkpoint_state()
# -- Load only the role being exported --
if role not in method._role_models:
raise KeyError(f"Role {role!r} is not configured. "
f"Available: {sorted(method._role_models)}")
model = method._role_models[role]
if model.transformer is None:
raise ValueError(f"Role {role!r} does not have a transformer to export")
# Export is a model-only operation. Building the complete training
# checkpoint state also includes optimizers and schedulers, whose lazy
# state may not exist until the first optimizer step. Loading just the
# requested role avoids coupling stage handoff to unrelated training
# runtime state.
role_state_key = f"roles.{role}.transformer"
states = {role_state_key: _FullModelState(model.transformer)}
logger.info(
"Loading DCP checkpoint from %s",
"Loading %s from DCP checkpoint %s",
role_state_key,
resolved,
)
dcp.load(states, checkpoint_id=str(dcp_dir))
# -- Export to diffusers format --
model = method._role_models[role]
base_model_path = str(tc.model_path)
if not base_model_path:
raise ValueError("Cannot determine base_model_path from "
"config. Ensure models.student.init_from "
"is set.")
logger.info(
"Exporting role=%s to %s (base=%s)",
role,
@@ -394,7 +467,8 @@ def main() -> None:
"--role",
type=str,
default="student",
help="Role to export (default: student).",
help=("Role to export (default: student). Use student_ema to export "
"the exact EMA weights used by validation."),
)
parser.add_argument(
"--overwrite",
+10
View File
@@ -249,6 +249,16 @@ class TrainingMethod(torch.nn.Module, ABC):
"""
return {"student": self.student.transformer}
def should_update_ema(self, iteration: int) -> bool:
"""Whether the student optimizer updated on this outer step.
EMA callbacks use this hook to follow the optimizer cadence instead
of the trainer's outer-loop cadence. Methods that update the student
conditionally (for example DMD2) should override it.
"""
del iteration
return True
def on_train_start(self) -> None:
from fastvideo.distributed import (
get_world_group, )
@@ -42,6 +42,9 @@ class CausalConsistencyDistillationMethod(TrainingMethod):
if self._discrete_cd_n < 2:
raise ValueError("method.discrete_cd_N must be >= 2")
self._ema_decay = float(self.method_config.get("ema_decay", 0.99))
# Kept in the parsed method state for reference-config provenance.
# It is an EMA checkpoint-selection threshold in the official code,
# not an online-target update gate.
self._ema_start_step = int(self.method_config.get("ema_start_step", 200))
shift = getattr(self.training_config.pipeline_config, "flow_shift", None)
self._flow_shift = float(shift) if shift else 5.0
@@ -186,8 +189,11 @@ class CausalConsistencyDistillationMethod(TrainingMethod):
def optimizers_schedulers_step(self, iteration: int) -> None:
super().optimizers_schedulers_step(iteration)
if iteration >= self._ema_start_step:
self._update_ema()
# The reference Causal-CD trainer updates its target EMA after every
# student optimizer step. ``ema_start_step`` controls which weights
# are selected for checkpoint/export there; it does not freeze the
# online consistency target during the first iterations.
self._update_ema()
# ------------------------------------------------------------------
# Internal helpers
@@ -385,6 +385,10 @@ class DMD2Method(TrainingMethod):
return True
return iteration % interval == 0
def should_update_ema(self, iteration: int) -> bool:
"""Keep EMA cadence aligned with generator optimizer updates."""
return self._should_update_student(iteration)
def _get_denoising_step_list(
self,
device: torch.device,
+2
View File
@@ -592,7 +592,9 @@ class WanModel(ModelBase):
text_dict: dict[str, torch.Tensor] | None,
clean_x: torch.Tensor | None = None,
aug_t: torch.Tensor | None = None,
start_frame: int = 0,
) -> dict[str, Any]:
del start_frame
if text_dict is None:
raise ValueError("text_dict cannot be None for "
"Wan distillation")
+1
View File
@@ -172,6 +172,7 @@ class WanCausalModel(WanModel, CausalModelBase):
noisy_latents,
timestep_full,
text_dict,
start_frame=cur_start_frame,
))
input_kwargs["timestep"] = (timestep_full.to(
device=self.device,
@@ -0,0 +1,19 @@
from fastvideo.train.models.wantrack.wantrack import (
TrackAugmentationConfig,
WanTrackModel,
)
from fastvideo.train.models.wantrack.wantrack_causal import WanTrackCausalModel
from fastvideo.train.models.wantrack.runtime import (
CausalWanTrackSession,
WanTrackInferenceRuntime,
WanTrackSamplingSettings,
)
__all__ = [
"TrackAugmentationConfig",
"WanTrackModel",
"WanTrackCausalModel",
"WanTrackInferenceRuntime",
"WanTrackSamplingSettings",
"CausalWanTrackSession",
]
+312
View File
@@ -0,0 +1,312 @@
# SPDX-License-Identifier: Apache-2.0
"""Thread-safe, future-only handle control for causal WanTrack sessions."""
from __future__ import annotations
from copy import deepcopy
from dataclasses import dataclass
import math
import threading
from typing import Any
import numpy as np
@dataclass(frozen=True, slots=True)
class HandleState:
handle_id: str
anchor: tuple[float, float]
position: tuple[float, float]
@dataclass(frozen=True, slots=True)
class AppliedControl:
revision: int
radius: float
active_handle_ids: tuple[str, ...]
tracks: np.ndarray
visibility: np.ndarray
def make_normalized_grid(size: int = 50) -> np.ndarray:
if size <= 1:
raise ValueError("grid size must be greater than one")
axis = np.linspace(0.0, 1.0, size, dtype=np.float32)
x, y = np.meshgrid(axis, axis, indexing="xy")
return np.stack([x, y], axis=-1).reshape(-1, 2)
def cosine_radius_weights(
points: np.ndarray,
center: tuple[float, float],
radius: float,
) -> np.ndarray:
if not 0.0 < radius <= math.sqrt(2.0):
raise ValueError("radius must be in (0, sqrt(2)]")
distance = np.linalg.norm(
np.asarray(points, dtype=np.float32) - np.asarray(center, dtype=np.float32),
axis=-1,
)
normalized = np.clip(distance / float(radius), 0.0, 1.0)
weights = 0.5 * (1.0 + np.cos(np.pi * normalized))
weights[distance >= radius] = 0.0
return weights.astype(np.float32, copy=False)
def deform_grid(
grid: np.ndarray,
handles: tuple[HandleState, ...],
radius: float,
) -> np.ndarray:
"""Blend overlapping handle displacements without amplifying overlap."""
grid = np.asarray(grid, dtype=np.float32)
if not handles:
return grid.copy()
weighted_displacement = np.zeros_like(grid)
total_weight = np.zeros((grid.shape[0], 1), dtype=np.float32)
for handle in handles:
weights = cosine_radius_weights(grid, handle.anchor, radius)[:, None]
displacement = (np.asarray(handle.position, dtype=np.float32) - np.asarray(handle.anchor, dtype=np.float32))
weighted_displacement += weights * displacement
total_weight += weights
# A lone handle retains its cosine falloff. Overlap is normalized once the
# combined influence exceeds one, avoiding motion amplification.
blended = weighted_displacement / np.maximum(1.0, total_weight)
return np.clip(grid + blended, 0.0, 1.0).astype(np.float32, copy=False)
class StableGridController:
"""Own stable grid identities and boundary-applied control revisions."""
def __init__(
self,
handles: list[dict[str, Any]],
*,
radius: float,
grid_size: int = 50,
) -> None:
if not handles:
raise ValueError("at least one handle is required")
self.grid = make_normalized_grid(grid_size)
self.radius = self._validate_radius(radius)
self._handles: dict[str, HandleState] = {}
for raw in handles:
state = self._coerce_handle(raw)
if state.handle_id in self._handles:
raise ValueError(f"duplicate handle id: {state.handle_id!r}")
self._handles[state.handle_id] = state
self._accepted_revision = 0
self._applied_revision = 0
self._pending: list[dict[str, Any]] = []
self._lock = threading.RLock()
@staticmethod
def _validate_radius(radius: float) -> float:
value = float(radius)
if not 0.0 < value <= math.sqrt(2.0):
raise ValueError("radius must be in (0, sqrt(2)]")
return value
@staticmethod
def _coerce_point(raw: dict[str, Any]) -> tuple[float, float]:
point = (float(raw["x"]), float(raw["y"]))
if not all(math.isfinite(v) and 0.0 <= v <= 1.0 for v in point):
raise ValueError("handle coordinates must be finite and normalized")
return point
@classmethod
def _coerce_handle(cls, raw: dict[str, Any]) -> HandleState:
handle_id = str(raw.get("id", "")).strip()
if not handle_id:
raise ValueError("handle id must be non-empty")
point = cls._coerce_point(raw)
return HandleState(handle_id, point, point)
@property
def accepted_revision(self) -> int:
with self._lock:
return self._accepted_revision
@property
def applied_revision(self) -> int:
with self._lock:
return self._applied_revision
@property
def active_handle_ids(self) -> tuple[str, ...]:
with self._lock:
return tuple(sorted(self._handles))
def queue_revision(
self,
revision: int,
*,
samples: list[dict[str, Any]] | None = None,
add: list[dict[str, Any]] | None = None,
remove: list[str] | None = None,
handles: list[dict[str, Any]] | None = None,
radius: float | None = None,
received_at_ms: float | None = None,
) -> bool:
"""Queue a monotonic revision; stale revisions are ignored."""
revision = int(revision)
with self._lock:
if revision <= self._accepted_revision:
return False
payload = {
"revision": revision,
"samples": deepcopy(samples or []),
"add": deepcopy(add or []),
"remove": [str(item) for item in (remove or [])],
"handles": deepcopy(handles),
"radius": (None if radius is None else self._validate_radius(radius)),
"received_at_ms": (None if received_at_ms is None else float(received_at_ms)),
}
self._pending.append(payload)
self._accepted_revision = revision
return True
def render_constant(self, frame_count: int) -> AppliedControl:
if frame_count <= 0:
raise ValueError("frame_count must be positive")
with self._lock:
frame = deform_grid(
self.grid,
tuple(self._handles.values()),
self.radius,
)
tracks = np.repeat(frame[None], frame_count, axis=0)
visibility = np.ones(tracks.shape[:2], dtype=np.float32)
return AppliedControl(
revision=self._applied_revision,
radius=self.radius,
active_handle_ids=tuple(sorted(self._handles)),
tracks=tracks,
visibility=visibility,
)
def apply_pending(
self,
frame_count: int,
*,
interval_start_ms: float,
interval_end_ms: float,
) -> AppliedControl:
"""Atomically apply queued lifecycle state and resample the interval."""
if frame_count <= 0:
raise ValueError("frame_count must be positive")
if interval_end_ms < interval_start_ms:
raise ValueError("control interval end precedes its start")
with self._lock:
if not self._pending:
return self.render_constant(frame_count)
pending = self._pending
self._pending = []
old_handles = dict(self._handles)
new_handles = dict(self._handles)
sample_map: dict[str, list[tuple[float, tuple[float, float]]]] = {}
latest_radius = self.radius
for payload in pending:
if payload["radius"] is not None:
latest_radius = payload["radius"]
explicit_handles = payload["handles"]
if explicit_handles is not None:
desired: dict[str, HandleState] = {}
for raw in explicit_handles:
candidate = self._coerce_handle(raw)
previous = new_handles.get(candidate.handle_id)
if previous is not None:
candidate = HandleState(
candidate.handle_id,
previous.anchor,
previous.position,
)
desired[candidate.handle_id] = candidate
new_handles = desired
for handle_id in payload["remove"]:
new_handles.pop(handle_id, None)
for raw in payload["add"]:
candidate = self._coerce_handle(raw)
new_handles[candidate.handle_id] = candidate
for raw in payload["samples"]:
handle_id = str(raw.get("id", "")).strip()
if not handle_id:
raise ValueError("control sample id must be non-empty")
point = self._coerce_point(raw)
timestamp = raw.get("timestamp_ms", raw.get("time_ms"))
if timestamp is None:
timestamp = payload["received_at_ms"]
if timestamp is None:
timestamp = interval_end_ms
sample_map.setdefault(handle_id, []).append((float(timestamp), point))
frame_times = np.linspace(
interval_start_ms,
interval_end_ms,
frame_count,
dtype=np.float64,
)
paths: dict[str, np.ndarray] = {}
for handle_id, state in new_handles.items():
prior = old_handles.get(handle_id, state)
samples = sorted(sample_map.get(handle_id, []), key=lambda item: item[0])
times = [float(interval_start_ms)]
points = [prior.position]
for timestamp, point in samples:
timestamp = min(
max(timestamp, interval_start_ms),
interval_end_ms,
)
if timestamp == times[-1]:
points[-1] = point
else:
times.append(timestamp)
points.append(point)
if times[-1] < interval_end_ms:
times.append(float(interval_end_ms))
points.append(points[-1])
x = np.interp(frame_times, times, [point[0] for point in points])
y = np.interp(frame_times, times, [point[1] for point in points])
paths[handle_id] = np.stack([x, y], axis=-1).astype(np.float32)
final_position = (
float(paths[handle_id][-1, 0]),
float(paths[handle_id][-1, 1]),
)
new_handles[handle_id] = HandleState(
handle_id,
state.anchor,
final_position,
)
tracks = np.empty(
(frame_count, self.grid.shape[0], 2),
dtype=np.float32,
)
for frame_index in range(frame_count):
frame_handles = tuple(
HandleState(
handle_id,
state.anchor,
(
float(paths[handle_id][frame_index, 0]),
float(paths[handle_id][frame_index, 1]),
),
) for handle_id, state in sorted(new_handles.items()))
tracks[frame_index] = deform_grid(
self.grid,
frame_handles,
latest_radius,
)
self._handles = new_handles
self.radius = latest_radius
self._applied_revision = pending[-1]["revision"]
visibility = np.ones(tracks.shape[:2], dtype=np.float32)
return AppliedControl(
revision=self._applied_revision,
radius=self.radius,
active_handle_ids=tuple(sorted(self._handles)),
tracks=tracks,
visibility=visibility,
)
@@ -0,0 +1,537 @@
# SPDX-License-Identifier: Apache-2.0
"""Shared bidirectional and causal WanTrack sampling helpers."""
from __future__ import annotations
from dataclasses import replace
from typing import Any, Literal
import torch
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler, )
from fastvideo.models.utils import pred_noise_to_pred_video
from fastvideo.pipelines import TrainingBatch
from fastvideo.train.models.base import CausalModelBase, ModelBase
_Branch = Literal["full", "no_text", "no_motion"]
_CACHE_BRANCHES: tuple[_Branch, ...] = (
"full",
"no_text",
"no_motion",
)
def prepare_wantrack_batch(
model: ModelBase,
raw_batch: dict[str, Any],
*,
seed: int,
latents_source: Literal["data", "zeros"] = "zeros",
) -> TrainingBatch:
"""Build inference conditions through the same path used by training."""
augmentation = getattr(model, "track_augmentation", None)
if augmentation is None:
raise TypeError("WanTrack inference requires a WanTrack training model")
generator = torch.Generator(device=model.device).manual_seed(int(seed))
model.track_augmentation = replace(
augmentation,
track_dropout_probability=0.0,
temporal_mask_probability=0.0,
motion_dropout_probability=0.0,
text_dropout_probability=0.0,
)
try:
batch = model.prepare_batch(
raw_batch,
generator=generator,
latents_source=latents_source,
)
finally:
model.track_augmentation = augmentation
# Streaming inference owns its cache/mask geometry. Training attention
# metadata must not leak into validation or standalone sampling.
batch.attn_metadata = None
batch.attn_metadata_vsa = None
return batch
def _branch_args(branch: _Branch) -> tuple[bool, dict[str, Any] | None]:
if branch == "full":
return True, None
if branch == "no_text":
return False, {
"text": "zero",
"track": "keep",
"on_missing": "ignore",
}
if branch == "no_motion":
return False, {
"text": "keep",
"track": "drop",
"on_missing": "ignore",
}
raise ValueError(f"Unknown WanTrack CFG branch: {branch!r}")
def _predict(
model: ModelBase,
latents: torch.Tensor,
timestep: torch.Tensor,
batch: TrainingBatch,
*,
branch: _Branch,
cache_tag: str,
start_frame: int,
store_kv: bool,
) -> torch.Tensor | None:
conditional, cfg_uncond = _branch_args(branch)
if isinstance(model, CausalModelBase):
return model.predict_noise_streaming(
latents,
timestep,
batch,
conditional=conditional,
cache_tag=cache_tag,
store_kv=store_kv,
cur_start_frame=start_frame,
cfg_uncond=cfg_uncond,
attn_kind="dense",
)
if store_kv:
return None
return model.predict_noise(
latents,
timestep,
batch,
conditional=conditional,
cfg_uncond=cfg_uncond,
attn_kind="dense",
)
def _guided_prediction(
model: ModelBase,
latents: torch.Tensor,
timestep: torch.Tensor,
batch: TrainingBatch,
*,
start_frame: int,
text_guidance_scale: float,
motion_guidance_scale: float,
motion_cfg: bool,
) -> tuple[torch.Tensor, tuple[_Branch, ...]]:
full = _predict(
model,
latents,
timestep,
batch,
branch="full",
cache_tag="wantrack_full",
start_frame=start_frame,
store_kv=False,
)
if full is None:
raise RuntimeError("WanTrack prediction unexpectedly returned None")
text_scale = float(text_guidance_scale)
motion_scale = float(motion_guidance_scale)
if text_scale == 1.0 and motion_scale == 1.0:
return full, ("full", )
no_text = _predict(
model,
latents,
timestep,
batch,
branch="no_text",
cache_tag="wantrack_no_text",
start_frame=start_frame,
store_kv=False,
)
if no_text is None:
raise RuntimeError("WanTrack no-text prediction returned None")
if not motion_cfg:
return no_text + text_scale * (full - no_text), (
"full",
"no_text",
)
no_motion = _predict(
model,
latents,
timestep,
batch,
branch="no_motion",
cache_tag="wantrack_no_motion",
start_frame=start_frame,
store_kv=False,
)
if no_motion is None:
raise RuntimeError("WanTrack no-motion prediction returned None")
denominator = text_scale + motion_scale
alpha = text_scale / denominator if denominator > 0 else 0.5
base = alpha * no_text + (1.0 - alpha) * no_motion
guided = (base + text_scale * (full - no_text) + motion_scale * (full - no_motion))
return guided, ("full", "no_text", "no_motion")
def _store_causal_context(
model: CausalModelBase,
latents: torch.Tensor,
batch: TrainingBatch,
*,
start_frame: int,
branches: tuple[_Branch, ...],
) -> None:
timestep = torch.zeros(
latents.shape[:2],
device=latents.device,
dtype=torch.float32,
)
batch.timesteps = timestep
for branch in branches:
_predict(
model,
latents,
timestep,
batch,
branch=branch,
cache_tag=f"wantrack_{branch}",
start_frame=start_frame,
store_kv=True,
)
def resolve_dmd_timesteps(
scheduler: FlowMatchEulerDiscreteScheduler,
dmd_denoising_steps: list[int] | tuple[int, ...],
*,
warp_denoising_step: bool = True,
device: torch.device | str | None = None,
) -> torch.Tensor:
"""Map Self-Forcing DMD indices onto the flow-matching schedule."""
if not dmd_denoising_steps:
raise ValueError("dmd_denoising_steps must be a non-empty sequence")
steps = torch.tensor([int(s) for s in dmd_denoising_steps], dtype=torch.long)
if bool(warp_denoising_step):
# Match WanTrackCausalDenoisingStage / Self-Forcing training warp.
scheduler.set_timesteps(1000, device="cpu")
schedule = torch.cat(
(scheduler.timesteps.cpu(), torch.tensor([0], dtype=torch.float32)))
num_train = int(getattr(scheduler.config, "num_train_timesteps", 1000))
steps = schedule[num_train - steps]
if device is None:
return steps
return steps.to(device=device)
def _sample_block_euler(
model: ModelBase,
latents: torch.Tensor,
batch: TrainingBatch,
*,
scheduler: FlowMatchEulerDiscreteScheduler,
num_inference_steps: int,
start_frame: int,
text_guidance_scale: float,
motion_guidance_scale: float,
motion_cfg: bool,
) -> tuple[torch.Tensor, tuple[_Branch, ...]]:
latent_dtype = latents.dtype
scheduler.set_timesteps(int(num_inference_steps), device=latents.device)
branches: tuple[_Branch, ...] = ("full", )
for current_timestep in scheduler.timesteps:
timestep = torch.full(
latents.shape[:2],
float(current_timestep.item()),
device=latents.device,
dtype=torch.float32,
)
batch.timesteps = timestep
prediction, branches = _guided_prediction(
model,
latents,
timestep,
batch,
start_frame=start_frame,
text_guidance_scale=text_guidance_scale,
motion_guidance_scale=motion_guidance_scale,
motion_cfg=motion_cfg,
)
latents = scheduler.step(
prediction.float(),
current_timestep,
latents.float(),
return_dict=False,
)[0].to(dtype=latent_dtype)
return latents, branches
def _sample_block_dmd(
model: ModelBase,
latents: torch.Tensor,
batch: TrainingBatch,
*,
scheduler: FlowMatchEulerDiscreteScheduler,
dmd_timesteps: torch.Tensor,
start_frame: int,
text_guidance_scale: float,
motion_guidance_scale: float,
motion_cfg: bool,
) -> tuple[torch.Tensor, tuple[_Branch, ...]]:
"""Self-Forcing / DMD multistep: predict x0, then re-noise to the next t."""
latent_dtype = latents.dtype
branches: tuple[_Branch, ...] = ("full", )
for step_idx, current_timestep in enumerate(dmd_timesteps):
timestep = torch.full(
latents.shape[:2],
float(current_timestep.item()),
device=latents.device,
dtype=torch.float32,
)
batch.timesteps = timestep
prediction, branches = _guided_prediction(
model,
latents,
timestep,
batch,
start_frame=start_frame,
text_guidance_scale=text_guidance_scale,
motion_guidance_scale=motion_guidance_scale,
motion_cfg=motion_cfg,
)
pred_x0 = pred_noise_to_pred_video(
pred_noise=prediction.flatten(0, 1).float(),
noise_input_latent=latents.flatten(0, 1).float(),
timestep=timestep,
scheduler=scheduler,
).unflatten(0, prediction.shape[:2]).to(dtype=latent_dtype)
if step_idx + 1 >= len(dmd_timesteps):
latents = pred_x0
break
next_timestep = dmd_timesteps[step_idx + 1]
next_t = torch.full(
latents.shape[:2],
float(next_timestep.item()),
device=latents.device,
dtype=torch.float32,
)
noise = torch.randn_like(pred_x0, dtype=torch.float32)
latents = scheduler.add_noise(
pred_x0.flatten(0, 1).float(),
noise.flatten(0, 1),
next_t,
).unflatten(0, pred_x0.shape[:2]).to(dtype=latent_dtype)
return latents, branches
def _sample_block(
model: ModelBase,
latents: torch.Tensor,
batch: TrainingBatch,
*,
scheduler: FlowMatchEulerDiscreteScheduler,
num_inference_steps: int,
start_frame: int,
text_guidance_scale: float,
motion_guidance_scale: float,
motion_cfg: bool,
dmd_denoising_steps: list[int] | tuple[int, ...] | None = None,
warp_denoising_step: bool = True,
) -> tuple[torch.Tensor, tuple[_Branch, ...]]:
if dmd_denoising_steps is not None:
# Ensure the scheduler owns a 1000-step grid before x0 / add_noise.
scheduler.set_timesteps(1000, device=latents.device)
dmd_timesteps = resolve_dmd_timesteps(
scheduler,
dmd_denoising_steps,
warp_denoising_step=warp_denoising_step,
device=latents.device,
)
return _sample_block_dmd(
model,
latents,
batch,
scheduler=scheduler,
dmd_timesteps=dmd_timesteps,
start_frame=start_frame,
text_guidance_scale=text_guidance_scale,
motion_guidance_scale=motion_guidance_scale,
motion_cfg=motion_cfg,
)
return _sample_block_euler(
model,
latents,
batch,
scheduler=scheduler,
num_inference_steps=num_inference_steps,
start_frame=start_frame,
text_guidance_scale=text_guidance_scale,
motion_guidance_scale=motion_guidance_scale,
motion_cfg=motion_cfg,
)
def clear_wantrack_caches(model: ModelBase) -> None:
"""Clear every CFG-tagged causal cache owned by WanTrack sampling."""
if not isinstance(model, CausalModelBase):
return
for branch in _CACHE_BRANCHES:
model.clear_caches(cache_tag=f"wantrack_{branch}")
@torch.no_grad()
def sample_wantrack_block(
model: CausalModelBase,
batch: TrainingBatch,
latents: torch.Tensor,
*,
start_frame: int,
num_inference_steps: int = 30,
text_guidance_scale: float = 1.0,
motion_guidance_scale: float = 1.0,
motion_cfg: bool = True,
scheduler: FlowMatchEulerDiscreteScheduler | None = None,
commit: bool = True,
dmd_denoising_steps: list[int] | tuple[int, ...] | None = None,
warp_denoising_step: bool = True,
) -> torch.Tensor:
"""Denoise one causal block and optionally commit its context once."""
if latents.ndim != 5:
raise ValueError("WanTrack block latents must be [B, T, C, H, W]")
if start_frame < 0:
raise ValueError("start_frame must be non-negative")
if num_inference_steps <= 0:
raise ValueError("num_inference_steps must be positive")
if scheduler is None:
scheduler = FlowMatchEulerDiscreteScheduler(shift=float(getattr(model, "timestep_shift", 5.0)), )
sampled, branches = _sample_block(
model,
latents,
batch,
scheduler=scheduler,
num_inference_steps=num_inference_steps,
start_frame=start_frame,
text_guidance_scale=text_guidance_scale,
motion_guidance_scale=motion_guidance_scale,
motion_cfg=motion_cfg,
dmd_denoising_steps=dmd_denoising_steps,
warp_denoising_step=warp_denoising_step,
)
if commit:
_store_causal_context(
model,
sampled,
batch,
start_frame=start_frame,
branches=branches,
)
return sampled
@torch.no_grad()
def sample_wantrack(
model: ModelBase,
batch: TrainingBatch,
*,
num_inference_steps: int = 30,
seed: int = 0,
text_guidance_scale: float = 1.0,
motion_guidance_scale: float = 1.0,
motion_cfg: bool = True,
chunk_size: int | None = None,
dmd_denoising_steps: list[int] | tuple[int, ...] | None = None,
warp_denoising_step: bool = True,
) -> torch.Tensor:
"""Generate normalized latents in ``[B, T, C, H, W]`` layout.
Bidirectional models denoise the complete clip. Causal models reuse the
same streaming prediction and cache API as RobotWM.
"""
if batch.latents is None or batch.latents.ndim != 5:
raise ValueError("WanTrack inference requires [B, T, C, H, W] latents")
if num_inference_steps <= 0:
raise ValueError("num_inference_steps must be positive")
generator = torch.Generator(device="cpu").manual_seed(int(seed))
latents = torch.randn(
tuple(batch.latents.shape),
generator=generator,
dtype=torch.float32,
).to(device=batch.latents.device, dtype=batch.latents.dtype)
scheduler = FlowMatchEulerDiscreteScheduler(shift=float(getattr(model, "timestep_shift", 5.0)), )
if not isinstance(model, CausalModelBase):
sampled, _ = _sample_block(
model,
latents,
batch,
scheduler=scheduler,
num_inference_steps=num_inference_steps,
start_frame=0,
text_guidance_scale=text_guidance_scale,
motion_guidance_scale=motion_guidance_scale,
motion_cfg=motion_cfg,
dmd_denoising_steps=dmd_denoising_steps,
warp_denoising_step=warp_denoising_step,
)
return sampled
if chunk_size is None:
transformer = getattr(model, "transformer", None)
chunk_size = int(
getattr(
transformer,
"num_frame_per_block",
getattr(transformer.config.arch_config, "num_frames_per_block", 3),
))
chunk_size = int(chunk_size)
if chunk_size <= 0:
raise ValueError("chunk_size must be positive")
num_frames = int(latents.shape[1])
if num_frames % chunk_size == 0:
block_sizes = [chunk_size] * (num_frames // chunk_size)
elif (num_frames - 1) % chunk_size == 0:
# Wan I2V clips have one leading latent frame followed by regular
# causal blocks (for example, 31 = 1 + 10 * 3).
block_sizes = [1] + [chunk_size] * ((num_frames - 1) // chunk_size)
else:
raise ValueError("Causal WanTrack inference requires latent frames "
"to form complete blocks, optionally after one "
f"leading I2V frame; got {num_frames} and "
f"{chunk_size}")
clear_wantrack_caches(model)
sampled_blocks: list[torch.Tensor] = []
try:
start_frame = 0
for block_size in block_sizes:
block = latents[:, start_frame:start_frame + block_size]
block = sample_wantrack_block(
model,
batch,
block,
start_frame=start_frame,
num_inference_steps=num_inference_steps,
text_guidance_scale=text_guidance_scale,
motion_guidance_scale=motion_guidance_scale,
motion_cfg=motion_cfg,
scheduler=scheduler,
commit=True,
dmd_denoising_steps=dmd_denoising_steps,
warp_denoising_step=warp_denoising_step,
)
sampled_blocks.append(block)
start_frame += block_size
finally:
clear_wantrack_caches(model)
return torch.cat(sampled_blocks, dim=1)

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