Compare commits
26
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
9f94673c34 | ||
|
|
c150889330 | ||
|
|
7f783e30b7 | ||
|
|
ec9db97431 | ||
|
|
fd236a06a7 | ||
|
|
54fe0be488 | ||
|
|
3a398a2bbd | ||
|
|
4e76897f1f | ||
|
|
95d8e74d65 | ||
|
|
bcd0ad7dee | ||
|
|
9acbf92455 | ||
|
|
ce4fa8ef2d | ||
|
|
41dcd065ae | ||
|
|
c0cb2228ac | ||
|
|
2fa0f70ee8 | ||
|
|
e172215880 | ||
|
|
23e3b325a4 | ||
|
|
44d0990199 | ||
|
|
59b6578906 | ||
|
|
fd8372d172 | ||
|
|
33eb7ea13c | ||
|
|
051d21958d | ||
|
|
bf29e20bf5 | ||
|
|
291fa5d9a6 | ||
|
|
954daf7fb3 | ||
|
|
e05c04a2f6 |
@@ -0,0 +1,35 @@
|
||||
# Causal WanTrack Control
|
||||
|
||||
This standalone prototype prepares one image and continuously generates causal
|
||||
WanTrack blocks. Handle updates received during block N are committed only at
|
||||
the next block boundary, so generated frames and control history are immutable.
|
||||
One GPU session may generate at a time. The SF checkpoint uses its fixed
|
||||
DMD four-step schedule (`method.dmd_denoising_steps`, default
|
||||
`[1000, 750, 500, 250]` with warp) without classifier-free guidance; the
|
||||
server does not accept client overrides for steps or guidance.
|
||||
|
||||
Set a Diffusers-format causal WanTrack export and its Self-Forcing training
|
||||
YAML, then launch:
|
||||
|
||||
```bash
|
||||
export WANTRACK_MODEL_DIR=/path/to/wantrack-causal-export
|
||||
export WANTRACK_YAML_PATH=/path/to/sf/config/run.yaml
|
||||
export WANTRACK_TAEHV_CHECKPOINT=/path/to/taew2_1.pth
|
||||
python -m apps.wantrack_control
|
||||
```
|
||||
|
||||
Open `http://127.0.0.1:8010`. FFmpeg with `libx264` is required. Completed
|
||||
blocks are preserved under `WANTRACK_OUTPUT_DIR` (or the system temporary
|
||||
directory) and concatenated into a downloadable MP4 on Stop, disconnect, or a
|
||||
recoverable failure.
|
||||
|
||||
The interactive preview uses the official TAEHV `StreamingTAEHV` decoder with
|
||||
the Wan 2.1 `taew2_1.pth` weights. The full Wan VAE remains loaded only because
|
||||
input preparation still needs its encoder.
|
||||
|
||||
The `/ws` endpoint accepts `prepare`, `start`, `control_update`, and `stop`
|
||||
JSON messages. It emits phase-specific `progress`, `prepared`,
|
||||
`session_started`, `block_started`, `block_encoding`, `control_applied`,
|
||||
`media_init`, `media_segment_complete`, `stream_complete`, and terminal
|
||||
`error` events. Binary frames carry one fMP4 initialization section followed
|
||||
by ordered block fragments.
|
||||
@@ -0,0 +1,5 @@
|
||||
"""Standalone causal WanTrack control prototype."""
|
||||
|
||||
from apps.wantrack_control.server import create_app
|
||||
|
||||
__all__ = ["create_app"]
|
||||
@@ -0,0 +1,18 @@
|
||||
"""Launch the standalone WanTrack control server."""
|
||||
|
||||
import os
|
||||
|
||||
import uvicorn
|
||||
|
||||
|
||||
def main() -> None:
|
||||
uvicorn.run(
|
||||
"apps.wantrack_control.server:app",
|
||||
host=os.getenv("WANTRACK_HOST", "127.0.0.1"),
|
||||
port=int(os.getenv("WANTRACK_PORT", "8010")),
|
||||
reload=False,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,236 @@
|
||||
"""Video-only fragmented-MP4 encoding and completed-prefix finalization."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
import os
|
||||
from pathlib import Path
|
||||
import shutil
|
||||
import subprocess
|
||||
from collections.abc import Iterable
|
||||
import uuid
|
||||
|
||||
import numpy as np
|
||||
|
||||
MEDIA_MIME = 'video/mp4; codecs="avc1.42E01E"'
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class EncodedMediaSegment:
|
||||
block_index: int
|
||||
init_bytes: bytes
|
||||
media_bytes: bytes
|
||||
path: Path
|
||||
mime: str = MEDIA_MIME
|
||||
|
||||
|
||||
def split_fmp4(data: bytes) -> tuple[bytes, bytes]:
|
||||
"""Split top-level fMP4 initialization boxes from moof/mdat media."""
|
||||
offset = 0
|
||||
first_fragment: int | None = None
|
||||
while offset + 8 <= len(data):
|
||||
size = int.from_bytes(data[offset:offset + 4], "big")
|
||||
box_type = data[offset + 4:offset + 8]
|
||||
header = 8
|
||||
if size == 1:
|
||||
if offset + 16 > len(data):
|
||||
break
|
||||
size = int.from_bytes(data[offset + 8:offset + 16], "big")
|
||||
header = 16
|
||||
elif size == 0:
|
||||
size = len(data) - offset
|
||||
if size < header or offset + size > len(data):
|
||||
raise ValueError("invalid top-level MP4 box")
|
||||
if box_type == b"moof":
|
||||
first_fragment = offset
|
||||
break
|
||||
offset += size
|
||||
if first_fragment is None:
|
||||
raise ValueError("ffmpeg output did not contain an fMP4 moof box")
|
||||
init_bytes = data[:first_fragment]
|
||||
media_bytes = data[first_fragment:]
|
||||
if b"ftyp" not in init_bytes or b"moov" not in init_bytes:
|
||||
raise ValueError("ffmpeg output is missing fMP4 initialization boxes")
|
||||
return init_bytes, media_bytes
|
||||
|
||||
|
||||
class FMP4BlockWriter:
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
output_root: str | os.PathLike[str],
|
||||
*,
|
||||
fps: float,
|
||||
ffmpeg_bin: str | None = None,
|
||||
) -> None:
|
||||
self.fps = float(fps)
|
||||
if self.fps <= 0:
|
||||
raise ValueError("fps must be positive")
|
||||
resolved_ffmpeg = ffmpeg_bin or shutil.which(os.getenv("WANTRACK_FFMPEG_BIN", "ffmpeg"))
|
||||
if not resolved_ffmpeg:
|
||||
raise RuntimeError("ffmpeg is required for WanTrack streaming")
|
||||
self.ffmpeg_bin = resolved_ffmpeg
|
||||
root = Path(output_root)
|
||||
root.mkdir(parents=True, exist_ok=True)
|
||||
self.session_id = uuid.uuid4().hex
|
||||
self.session_dir = root / self.session_id
|
||||
self.session_dir.mkdir()
|
||||
self._block_paths: list[Path] = []
|
||||
|
||||
@property
|
||||
def block_paths(self) -> tuple[Path, ...]:
|
||||
return tuple(self._block_paths)
|
||||
|
||||
def encode_block(
|
||||
self,
|
||||
frames: np.ndarray,
|
||||
block_index: int,
|
||||
) -> EncodedMediaSegment:
|
||||
frames = np.asarray(frames)
|
||||
if frames.ndim != 4 or frames.shape[-1] != 3:
|
||||
raise ValueError("frames must have shape [T, H, W, 3]")
|
||||
if frames.shape[0] <= 0:
|
||||
raise ValueError("cannot encode an empty frame block")
|
||||
frames = np.ascontiguousarray(frames, dtype=np.uint8)
|
||||
height, width = int(frames.shape[1]), int(frames.shape[2])
|
||||
gop = max(1, int(frames.shape[0]))
|
||||
command = [
|
||||
self.ffmpeg_bin,
|
||||
"-hide_banner",
|
||||
"-loglevel",
|
||||
"error",
|
||||
"-y",
|
||||
"-f",
|
||||
"rawvideo",
|
||||
"-pix_fmt",
|
||||
"rgb24",
|
||||
"-s:v",
|
||||
f"{width}x{height}",
|
||||
"-r",
|
||||
f"{self.fps:g}",
|
||||
"-i",
|
||||
"pipe:0",
|
||||
"-an",
|
||||
"-c:v",
|
||||
"libx264",
|
||||
"-preset",
|
||||
"ultrafast",
|
||||
"-tune",
|
||||
"zerolatency",
|
||||
"-profile:v",
|
||||
"baseline",
|
||||
"-pix_fmt",
|
||||
"yuv420p",
|
||||
"-g",
|
||||
str(gop),
|
||||
"-keyint_min",
|
||||
str(gop),
|
||||
"-sc_threshold",
|
||||
"0",
|
||||
"-movflags",
|
||||
"+empty_moov+default_base_moof+frag_keyframe",
|
||||
"-frag_duration",
|
||||
str(max(1, round(1_000_000 * frames.shape[0] / self.fps))),
|
||||
"-f",
|
||||
"mp4",
|
||||
"pipe:1",
|
||||
]
|
||||
result = subprocess.run(
|
||||
command,
|
||||
input=frames.tobytes(),
|
||||
capture_output=True,
|
||||
check=False,
|
||||
)
|
||||
if result.returncode != 0 or not result.stdout:
|
||||
stderr = result.stderr.decode("utf-8", errors="replace").strip()
|
||||
raise RuntimeError(f"ffmpeg failed to encode WanTrack block {block_index}: "
|
||||
f"{stderr or f'exit {result.returncode}'}")
|
||||
init_bytes, media_bytes = split_fmp4(result.stdout)
|
||||
path = self.session_dir / f"block_{int(block_index):06d}.mp4"
|
||||
path.write_bytes(result.stdout)
|
||||
self._block_paths.append(path)
|
||||
return EncodedMediaSegment(
|
||||
block_index=int(block_index),
|
||||
init_bytes=init_bytes,
|
||||
media_bytes=media_bytes,
|
||||
path=path,
|
||||
)
|
||||
|
||||
def finalize(self) -> Path | None:
|
||||
if not self._block_paths:
|
||||
return None
|
||||
output_path = self.session_dir / "wantrack_control.mp4"
|
||||
concat_path = self.session_dir / "concat.txt"
|
||||
concat_path.write_text(
|
||||
"".join(f"file '{self._concat_escape(path)}'\n" for path in self._block_paths),
|
||||
encoding="utf-8",
|
||||
)
|
||||
command = [
|
||||
self.ffmpeg_bin,
|
||||
"-hide_banner",
|
||||
"-loglevel",
|
||||
"error",
|
||||
"-y",
|
||||
"-f",
|
||||
"concat",
|
||||
"-safe",
|
||||
"0",
|
||||
"-i",
|
||||
str(concat_path),
|
||||
"-an",
|
||||
"-c",
|
||||
"copy",
|
||||
"-movflags",
|
||||
"+faststart",
|
||||
str(output_path),
|
||||
]
|
||||
result = subprocess.run(
|
||||
command,
|
||||
capture_output=True,
|
||||
check=False,
|
||||
)
|
||||
if result.returncode != 0 or not output_path.is_file():
|
||||
return self._finalize_reencode(output_path)
|
||||
return output_path if output_path.stat().st_size > 0 else None
|
||||
|
||||
@staticmethod
|
||||
def _concat_escape(path: Path) -> str:
|
||||
return str(path.resolve()).replace("'", "'\\''")
|
||||
|
||||
def _finalize_reencode(self, output_path: Path) -> Path | None:
|
||||
concat_path = self.session_dir / "concat.txt"
|
||||
command = [
|
||||
self.ffmpeg_bin,
|
||||
"-hide_banner",
|
||||
"-loglevel",
|
||||
"error",
|
||||
"-y",
|
||||
"-f",
|
||||
"concat",
|
||||
"-safe",
|
||||
"0",
|
||||
"-i",
|
||||
str(concat_path),
|
||||
"-an",
|
||||
"-c:v",
|
||||
"libx264",
|
||||
"-preset",
|
||||
"fast",
|
||||
"-pix_fmt",
|
||||
"yuv420p",
|
||||
"-movflags",
|
||||
"+faststart",
|
||||
str(output_path),
|
||||
]
|
||||
result = subprocess.run(
|
||||
command,
|
||||
capture_output=True,
|
||||
check=False,
|
||||
)
|
||||
if (result.returncode == 0 and output_path.is_file() and output_path.stat().st_size > 0):
|
||||
return output_path
|
||||
return None
|
||||
|
||||
|
||||
def total_size(paths: Iterable[Path]) -> int:
|
||||
return sum(path.stat().st_size for path in paths if path.is_file())
|
||||
@@ -0,0 +1,422 @@
|
||||
"""FastAPI/WebSocket server for causal WanTrack control."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import base64
|
||||
from contextlib import suppress
|
||||
import io
|
||||
import os
|
||||
from pathlib import Path
|
||||
import tempfile
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
from fastapi import FastAPI, HTTPException, WebSocket, WebSocketDisconnect
|
||||
from fastapi.responses import FileResponse
|
||||
from fastapi.staticfiles import StaticFiles
|
||||
from PIL import Image
|
||||
|
||||
from apps.wantrack_control.media import FMP4BlockWriter
|
||||
|
||||
_STATIC_DIR = Path(__file__).resolve().parent / "static"
|
||||
|
||||
|
||||
def _decode_image(value: str) -> bytes:
|
||||
if not isinstance(value, str) or not value.strip():
|
||||
raise ValueError("prepare.image must be a non-empty base64 string")
|
||||
encoded = value.split(",", 1)[1] if value.startswith("data:") else value
|
||||
try:
|
||||
data = base64.b64decode(encoded, validate=True)
|
||||
except Exception as exc:
|
||||
raise ValueError("prepare.image is not valid base64") from exc
|
||||
if not data:
|
||||
raise ValueError("prepare.image decoded to no bytes")
|
||||
return data
|
||||
|
||||
|
||||
def _encode_image(image: Image.Image) -> str:
|
||||
output = io.BytesIO()
|
||||
image.save(output, format="PNG")
|
||||
encoded = base64.b64encode(output.getvalue()).decode("ascii")
|
||||
return f"data:image/png;base64,{encoded}"
|
||||
|
||||
|
||||
def _block_value(block: Any, key: str, default: Any = None) -> Any:
|
||||
if isinstance(block, dict):
|
||||
return block.get(key, default)
|
||||
return getattr(block, key, default)
|
||||
|
||||
|
||||
class _RuntimeProvider:
|
||||
|
||||
def __init__(self, runtime: Any | None) -> None:
|
||||
self._runtime = runtime
|
||||
self._lock = asyncio.Lock()
|
||||
|
||||
async def get(self) -> Any:
|
||||
if self._runtime is not None:
|
||||
return self._runtime
|
||||
async with self._lock:
|
||||
if self._runtime is None:
|
||||
model_dir = os.getenv("WANTRACK_MODEL_DIR", "").strip()
|
||||
yaml_path = os.getenv("WANTRACK_YAML_PATH", "").strip()
|
||||
taehv_checkpoint = os.getenv("WANTRACK_TAEHV_CHECKPOINT", "").strip()
|
||||
if not model_dir or not yaml_path or not taehv_checkpoint:
|
||||
raise RuntimeError("Set WANTRACK_MODEL_DIR, WANTRACK_YAML_PATH, and "
|
||||
"WANTRACK_TAEHV_CHECKPOINT before preparing a session")
|
||||
from fastvideo.train.models.wantrack.runtime import (
|
||||
WanTrackInferenceRuntime, )
|
||||
|
||||
self._runtime = await asyncio.to_thread(
|
||||
WanTrackInferenceRuntime.from_export,
|
||||
model_dir,
|
||||
yaml_path,
|
||||
taehv_checkpoint,
|
||||
)
|
||||
return self._runtime
|
||||
|
||||
|
||||
def create_app(
|
||||
runtime: Any | None = None,
|
||||
*,
|
||||
output_dir: str | os.PathLike[str] | None = None,
|
||||
writer_factory: Any = FMP4BlockWriter,
|
||||
) -> FastAPI:
|
||||
app = FastAPI(title="Causal WanTrack Control")
|
||||
app.state.runtime_provider = _RuntimeProvider(runtime)
|
||||
app.state.active_generation = asyncio.Lock()
|
||||
app.state.downloads = {}
|
||||
if output_dir is None:
|
||||
resolved_output_dir: str | os.PathLike[str] = os.getenv(
|
||||
"WANTRACK_OUTPUT_DIR",
|
||||
str(Path(tempfile.gettempdir()) / "wantrack_control"),
|
||||
)
|
||||
else:
|
||||
resolved_output_dir = output_dir
|
||||
app.state.output_dir = Path(resolved_output_dir)
|
||||
app.state.writer_factory = writer_factory
|
||||
|
||||
app.mount("/static", StaticFiles(directory=_STATIC_DIR), name="static")
|
||||
|
||||
@app.get("/")
|
||||
async def index() -> FileResponse:
|
||||
return FileResponse(_STATIC_DIR / "index.html")
|
||||
|
||||
@app.get("/healthz")
|
||||
async def healthz() -> dict[str, Any]:
|
||||
return {
|
||||
"status": "ok",
|
||||
"active": app.state.active_generation.locked(),
|
||||
}
|
||||
|
||||
@app.get("/downloads/{download_id}")
|
||||
async def download(download_id: str) -> FileResponse:
|
||||
path = app.state.downloads.get(download_id)
|
||||
if path is None or not path.is_file():
|
||||
raise HTTPException(status_code=404, detail="download not found")
|
||||
return FileResponse(
|
||||
path,
|
||||
media_type="video/mp4",
|
||||
filename="wantrack_control.mp4",
|
||||
)
|
||||
|
||||
@app.websocket("/ws")
|
||||
async def websocket_endpoint(websocket: WebSocket) -> None:
|
||||
await websocket.accept()
|
||||
send_lock = asyncio.Lock()
|
||||
prepared: Any | None = None
|
||||
session: Any | None = None
|
||||
generation_task: asyncio.Task[None] | None = None
|
||||
stop_event = asyncio.Event()
|
||||
owns_generation_lock = False
|
||||
connected = True
|
||||
|
||||
async def send_json(payload: dict[str, Any]) -> None:
|
||||
nonlocal connected
|
||||
if not connected:
|
||||
return
|
||||
async with send_lock:
|
||||
try:
|
||||
await websocket.send_json(payload)
|
||||
except Exception:
|
||||
connected = False
|
||||
raise
|
||||
|
||||
async def send_bytes(payload: bytes) -> None:
|
||||
nonlocal connected
|
||||
if not connected:
|
||||
return
|
||||
async with send_lock:
|
||||
try:
|
||||
await websocket.send_bytes(payload)
|
||||
except Exception:
|
||||
connected = False
|
||||
raise
|
||||
|
||||
async def run_generation(runtime_value: Any) -> None:
|
||||
nonlocal owns_generation_lock, connected, generation_task
|
||||
writer = app.state.writer_factory(
|
||||
app.state.output_dir,
|
||||
fps=float(getattr(runtime_value, "fps", 16.0)),
|
||||
)
|
||||
init_sent = False
|
||||
last_applied_revision = 0
|
||||
terminal_error: str | None = None
|
||||
try:
|
||||
while not stop_event.is_set():
|
||||
block_index = int(getattr(session, "block_index", 0))
|
||||
block_started_at = time.perf_counter()
|
||||
await send_json({
|
||||
"type": "block_started",
|
||||
"block_index": block_index,
|
||||
"num_inference_steps": len(
|
||||
getattr(runtime_value, "dmd_denoising_steps", [1000, 750, 500, 250])),
|
||||
"dmd_denoising_steps": list(
|
||||
getattr(runtime_value, "dmd_denoising_steps", [1000, 750, 500, 250])),
|
||||
"cfg_enabled": False,
|
||||
})
|
||||
block = await asyncio.to_thread(session.generate_next_block)
|
||||
generated_at = time.perf_counter()
|
||||
applied_revision = int(_block_value(block, "applied_revision", 0))
|
||||
if applied_revision > last_applied_revision:
|
||||
last_applied_revision = applied_revision
|
||||
await send_json({
|
||||
"type": "control_applied",
|
||||
"revision": applied_revision,
|
||||
"block_index": int(_block_value(block, "block_index", block_index)),
|
||||
"radius": float(_block_value(block, "radius", 0.0)),
|
||||
"active_handle_ids": list(_block_value(block, "active_handle_ids", ())),
|
||||
})
|
||||
frames = _block_value(block, "pixel_frames")
|
||||
await send_json({
|
||||
"type": "block_encoding",
|
||||
"block_index": block_index,
|
||||
})
|
||||
encoded = await asyncio.to_thread(
|
||||
writer.encode_block,
|
||||
frames,
|
||||
int(_block_value(block, "block_index", block_index)),
|
||||
)
|
||||
encoded_at = time.perf_counter()
|
||||
if not init_sent:
|
||||
await send_json({
|
||||
"type": "media_init",
|
||||
"mime": encoded.mime,
|
||||
})
|
||||
await send_bytes(encoded.init_bytes)
|
||||
init_sent = True
|
||||
await send_bytes(encoded.media_bytes)
|
||||
await send_json({
|
||||
"type": "media_segment_complete",
|
||||
"block_index": encoded.block_index,
|
||||
"bytes": len(encoded.media_bytes),
|
||||
"generation_ms": round((generated_at - block_started_at) * 1000),
|
||||
"encoding_ms": round((encoded_at - generated_at) * 1000),
|
||||
})
|
||||
except Exception as exc:
|
||||
terminal_error = str(exc) or type(exc).__name__
|
||||
finally:
|
||||
if session is not None:
|
||||
with suppress(Exception):
|
||||
await asyncio.to_thread(
|
||||
session.close,
|
||||
"error" if terminal_error else ("disconnect" if not connected else "stop"),
|
||||
)
|
||||
final_path = await asyncio.to_thread(writer.finalize)
|
||||
download_url = None
|
||||
if final_path is not None:
|
||||
app.state.downloads[writer.session_id] = final_path
|
||||
download_url = f"/downloads/{writer.session_id}"
|
||||
if owns_generation_lock:
|
||||
app.state.active_generation.release()
|
||||
owns_generation_lock = False
|
||||
if terminal_error:
|
||||
if connected:
|
||||
with suppress(Exception):
|
||||
await send_json({
|
||||
"type": "error",
|
||||
"message": terminal_error,
|
||||
"download_url": download_url,
|
||||
})
|
||||
elif connected:
|
||||
with suppress(Exception):
|
||||
await send_json({
|
||||
"type": "stream_complete",
|
||||
"blocks": len(writer.block_paths),
|
||||
"download_url": download_url,
|
||||
})
|
||||
generation_task = None
|
||||
|
||||
async def handle_message(message: dict[str, Any]) -> None:
|
||||
nonlocal prepared, session, generation_task
|
||||
nonlocal owns_generation_lock
|
||||
message_type = str(message.get("type", "")).strip()
|
||||
if message_type == "prepare":
|
||||
if generation_task is not None:
|
||||
raise ValueError("prepare is unavailable during generation")
|
||||
prepare_started_at = time.perf_counter()
|
||||
await send_json({
|
||||
"type": "progress",
|
||||
"phase": "loading_model",
|
||||
"message": "Loading Track-v0",
|
||||
"detail": "The first request loads the SF checkpoint onto the GPU.",
|
||||
})
|
||||
runtime_value = await app.state.runtime_provider.get()
|
||||
await send_json({
|
||||
"type": "progress",
|
||||
"phase": "preparing_input",
|
||||
"message": "Preparing image and prompt",
|
||||
"detail": "Encoding text, the reference image, and its first-frame latent.",
|
||||
})
|
||||
image_bytes = _decode_image(message.get("image", ""))
|
||||
prompt = str(message.get("prompt", "") or "")
|
||||
prepared = await asyncio.to_thread(
|
||||
runtime_value.prepare,
|
||||
image_bytes,
|
||||
prompt,
|
||||
)
|
||||
processed_image = getattr(prepared, "image", None)
|
||||
if not isinstance(processed_image, Image.Image):
|
||||
processed_image = Image.open(io.BytesIO(image_bytes)).convert("RGB")
|
||||
await send_json({
|
||||
"type": "prepared",
|
||||
"image": _encode_image(processed_image),
|
||||
"width": processed_image.width,
|
||||
"height": processed_image.height,
|
||||
"fps": float(getattr(runtime_value, "fps", 16.0)),
|
||||
"chunk_size": int(getattr(runtime_value, "chunk_size", 3)),
|
||||
"causal_recipe": getattr(runtime_value, "causal_recipe", {}),
|
||||
"decoder": str(getattr(runtime_value, "decoder_name", "unknown")),
|
||||
"prepare_ms": round((time.perf_counter() - prepare_started_at) * 1000),
|
||||
})
|
||||
return
|
||||
|
||||
if message_type == "start":
|
||||
if prepared is None:
|
||||
raise ValueError("prepare must complete before start")
|
||||
if generation_task is not None:
|
||||
raise ValueError("session is already generating")
|
||||
handles = message.get("handles")
|
||||
if not isinstance(handles, list) or not handles:
|
||||
raise ValueError("start requires at least one handle")
|
||||
if app.state.active_generation.locked():
|
||||
raise RuntimeError("another WanTrack session is already generating")
|
||||
start_started_at = time.perf_counter()
|
||||
await send_json({
|
||||
"type": "progress",
|
||||
"phase": "starting_session",
|
||||
"message": "Starting causal session",
|
||||
"detail": "Initializing controls and the causal KV cache.",
|
||||
})
|
||||
await app.state.active_generation.acquire()
|
||||
owns_generation_lock = True
|
||||
runtime_value = await app.state.runtime_provider.get()
|
||||
session = runtime_value.create_session()
|
||||
dmd_steps = list(
|
||||
getattr(runtime_value, "dmd_denoising_steps", [1000, 750, 500, 250]))
|
||||
sampling = {
|
||||
"seed": int(message.get("seed", 0)),
|
||||
"num_inference_steps": len(dmd_steps),
|
||||
"text_guidance_scale": 1.0,
|
||||
"motion_guidance_scale": 1.0,
|
||||
"motion_cfg": False,
|
||||
}
|
||||
try:
|
||||
await asyncio.to_thread(
|
||||
session.start,
|
||||
prepared,
|
||||
getattr(prepared, "prompt", str(message.get("prompt", "") or "")),
|
||||
handles,
|
||||
sampling,
|
||||
radius=float(message.get("radius", 0.15)),
|
||||
)
|
||||
except Exception:
|
||||
app.state.active_generation.release()
|
||||
owns_generation_lock = False
|
||||
raise
|
||||
await send_json({
|
||||
"type": "session_started",
|
||||
"fps": float(getattr(runtime_value, "fps", 16.0)),
|
||||
"chunk_size": int(getattr(runtime_value, "chunk_size", 3)),
|
||||
"num_inference_steps": len(dmd_steps),
|
||||
"dmd_denoising_steps": dmd_steps,
|
||||
"cfg_enabled": False,
|
||||
"causal_recipe": getattr(runtime_value, "causal_recipe", {}),
|
||||
"decoder": str(getattr(runtime_value, "decoder_name", "unknown")),
|
||||
"start_ms": round((time.perf_counter() - start_started_at) * 1000),
|
||||
})
|
||||
stop_event.clear()
|
||||
generation_task = asyncio.create_task(run_generation(runtime_value))
|
||||
return
|
||||
|
||||
if message_type == "control_update":
|
||||
if session is None:
|
||||
raise ValueError("control_update requires a running session")
|
||||
revision = int(message.get("revision", 0))
|
||||
accepted = await asyncio.to_thread(
|
||||
session.apply_control_revision,
|
||||
revision,
|
||||
samples=message.get("samples"),
|
||||
add=message.get("add"),
|
||||
remove=message.get("remove"),
|
||||
handles=message.get("handles"),
|
||||
radius=message.get("radius"),
|
||||
)
|
||||
if not accepted:
|
||||
await send_json({
|
||||
"type": "control_applied",
|
||||
"revision": revision,
|
||||
"status": "ignored_stale",
|
||||
})
|
||||
return
|
||||
|
||||
if message_type == "stop":
|
||||
if generation_task is None:
|
||||
raise ValueError("stop requires a running session")
|
||||
stop_event.set()
|
||||
return
|
||||
raise ValueError(f"unknown client message type: {message_type!r}")
|
||||
|
||||
try:
|
||||
while True:
|
||||
receive_task = asyncio.create_task(websocket.receive_json())
|
||||
waiters: set[asyncio.Task[Any]] = {receive_task}
|
||||
if generation_task is not None:
|
||||
waiters.add(generation_task)
|
||||
done, _ = await asyncio.wait(
|
||||
waiters,
|
||||
return_when=asyncio.FIRST_COMPLETED,
|
||||
)
|
||||
if generation_task is not None and generation_task in done:
|
||||
receive_task.cancel()
|
||||
with suppress(asyncio.CancelledError):
|
||||
await receive_task
|
||||
await generation_task
|
||||
return
|
||||
message = await receive_task
|
||||
try:
|
||||
await handle_message(message)
|
||||
except Exception as exc:
|
||||
await send_json({
|
||||
"type": "error",
|
||||
"message": str(exc) or type(exc).__name__,
|
||||
})
|
||||
if generation_task is not None:
|
||||
stop_event.set()
|
||||
except WebSocketDisconnect:
|
||||
connected = False
|
||||
except Exception:
|
||||
connected = False
|
||||
finally:
|
||||
stop_event.set()
|
||||
if generation_task is not None:
|
||||
with suppress(Exception):
|
||||
await generation_task
|
||||
elif owns_generation_lock:
|
||||
app.state.active_generation.release()
|
||||
|
||||
return app
|
||||
|
||||
|
||||
app = create_app()
|
||||
@@ -0,0 +1,308 @@
|
||||
(() => {
|
||||
const $ = (id) => document.getElementById(id);
|
||||
const imageInput = $("image");
|
||||
const promptInput = $("prompt");
|
||||
const prepareButton = $("prepare");
|
||||
const canvas = $("canvas");
|
||||
const context = canvas.getContext("2d");
|
||||
const addButton = $("add");
|
||||
const removeButton = $("remove");
|
||||
const gridInput = $("grid");
|
||||
const radiusInput = $("radius");
|
||||
const radiusOutput = document.querySelector(".radius output");
|
||||
const startButton = $("start");
|
||||
const stopButton = $("stop");
|
||||
const video = $("video");
|
||||
const status = $("status");
|
||||
const statusLabel = $("status-label");
|
||||
const statusDetail = $("status-detail");
|
||||
const download = $("download");
|
||||
|
||||
let socket;
|
||||
let preparedImage;
|
||||
let handles = [];
|
||||
let selectedId = null;
|
||||
let addMode = false;
|
||||
let dragging = false;
|
||||
let generating = false;
|
||||
let revision = 0;
|
||||
let sessionStart = 0;
|
||||
let mediaSource;
|
||||
let sourceBuffer;
|
||||
let mediaQueue = [];
|
||||
let streamComplete = false;
|
||||
|
||||
function setStatus(value, detail = "", busy = false) {
|
||||
statusLabel.textContent = value;
|
||||
statusDetail.textContent = detail;
|
||||
status.dataset.busy = String(busy);
|
||||
}
|
||||
function seconds(milliseconds) {
|
||||
return `${(Number(milliseconds || 0) / 1000).toFixed(1)}s`;
|
||||
}
|
||||
function connect() {
|
||||
const scheme = location.protocol === "https:" ? "wss" : "ws";
|
||||
socket = new WebSocket(`${scheme}://${location.host}/ws`);
|
||||
socket.binaryType = "arraybuffer";
|
||||
socket.onopen = () => setStatus("Ready", "Choose an image to begin.");
|
||||
socket.onclose = () => setStatus("Disconnected", "Refresh after the server reconnects.");
|
||||
socket.onerror = () => setStatus("Connection error", "The control server is unreachable.");
|
||||
socket.onmessage = async (event) => {
|
||||
if (typeof event.data !== "string") {
|
||||
mediaQueue.push(event.data);
|
||||
flushMedia();
|
||||
return;
|
||||
}
|
||||
const message = JSON.parse(event.data);
|
||||
if (message.type === "progress") {
|
||||
setStatus(message.message || "Working", message.detail || "", true);
|
||||
} else if (message.type === "prepared") {
|
||||
preparedImage = new Image();
|
||||
preparedImage.onload = () => {
|
||||
canvas.width = message.width;
|
||||
canvas.height = message.height;
|
||||
draw();
|
||||
};
|
||||
preparedImage.src = message.image;
|
||||
handles = [];
|
||||
selectedId = null;
|
||||
addButton.disabled = false;
|
||||
startButton.disabled = true;
|
||||
prepareButton.disabled = false;
|
||||
prepareButton.textContent = "Prepare again";
|
||||
const recipe = message.causal_recipe || {};
|
||||
setStatus(
|
||||
`Prepared in ${seconds(message.prepare_ms)}`,
|
||||
`${message.fps} FPS · ${message.decoder || "unknown decoder"} · ${recipe.rope_cache_policy || "unknown"} RoPE · local ${recipe.local_attn_size ?? "?"} · sink ${recipe.sink_size ?? "?"} · Add a handle, then press Start.`,
|
||||
);
|
||||
} else if (message.type === "session_started") {
|
||||
generating = true;
|
||||
sessionStart = performance.now();
|
||||
startButton.disabled = true;
|
||||
stopButton.disabled = false;
|
||||
setStatus(
|
||||
`Session started in ${seconds(message.start_ms)}`,
|
||||
"4-step SF · CFG off · Building the first causal block.",
|
||||
true,
|
||||
);
|
||||
} else if (message.type === "block_started") {
|
||||
setStatus(
|
||||
`Generating block ${message.block_index}`,
|
||||
"DMD 4-step, single conditional branch. Controls update at the next block.",
|
||||
true,
|
||||
);
|
||||
} else if (message.type === "block_encoding") {
|
||||
setStatus(
|
||||
`Encoding block ${message.block_index}`,
|
||||
"The model is done; packaging frames for immediate playback.",
|
||||
true,
|
||||
);
|
||||
} else if (message.type === "control_applied") {
|
||||
if (message.status === "ignored_stale") {
|
||||
setStatus(`Stale revision ${message.revision} ignored`, "Drag again to send a newer control update.");
|
||||
} else {
|
||||
setStatus(`Control ${message.revision} applied`, "The updated motion is active in this block.");
|
||||
}
|
||||
} else if (message.type === "media_init") {
|
||||
setupMediaSource(message.mime);
|
||||
} else if (message.type === "media_segment_complete") {
|
||||
setStatus(
|
||||
`Playing through block ${message.block_index}`,
|
||||
`Generated in ${seconds(message.generation_ms)} · encoded in ${seconds(message.encoding_ms)} · drag handles for the next block.`,
|
||||
);
|
||||
} else if (message.type === "stream_complete") {
|
||||
generating = false;
|
||||
streamComplete = true;
|
||||
stopButton.disabled = true;
|
||||
startButton.disabled = false;
|
||||
if (message.download_url) {
|
||||
download.href = message.download_url;
|
||||
download.hidden = false;
|
||||
}
|
||||
flushMedia();
|
||||
setStatus(`Complete · ${message.blocks} blocks`, "The final MP4 is ready to download.");
|
||||
} else if (message.type === "error") {
|
||||
generating = false;
|
||||
prepareButton.disabled = false;
|
||||
stopButton.disabled = true;
|
||||
if (message.download_url) {
|
||||
download.href = message.download_url;
|
||||
download.hidden = false;
|
||||
}
|
||||
setStatus("Request failed", message.message || "Unknown error");
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
function setupMediaSource(mime) {
|
||||
mediaQueue = [];
|
||||
streamComplete = false;
|
||||
mediaSource = new MediaSource();
|
||||
video.src = URL.createObjectURL(mediaSource);
|
||||
mediaSource.addEventListener("sourceopen", () => {
|
||||
sourceBuffer = mediaSource.addSourceBuffer(mime);
|
||||
sourceBuffer.mode = "sequence";
|
||||
sourceBuffer.addEventListener("updateend", flushMedia);
|
||||
flushMedia();
|
||||
}, { once: true });
|
||||
}
|
||||
|
||||
function flushMedia() {
|
||||
if (!sourceBuffer || sourceBuffer.updating) return;
|
||||
if (mediaQueue.length) {
|
||||
sourceBuffer.appendBuffer(mediaQueue.shift());
|
||||
return;
|
||||
}
|
||||
if (streamComplete && mediaSource && mediaSource.readyState === "open") {
|
||||
mediaSource.endOfStream();
|
||||
}
|
||||
video.play().catch(() => {});
|
||||
}
|
||||
|
||||
function draw() {
|
||||
context.clearRect(0, 0, canvas.width, canvas.height);
|
||||
if (preparedImage) context.drawImage(preparedImage, 0, 0, canvas.width, canvas.height);
|
||||
if (gridInput.checked) {
|
||||
context.strokeStyle = "rgba(255,255,255,.12)";
|
||||
context.lineWidth = 1;
|
||||
for (let index = 0; index < 50; index += 1) {
|
||||
const x = index * canvas.width / 49;
|
||||
const y = index * canvas.height / 49;
|
||||
context.beginPath(); context.moveTo(x, 0); context.lineTo(x, canvas.height); context.stroke();
|
||||
context.beginPath(); context.moveTo(0, y); context.lineTo(canvas.width, y); context.stroke();
|
||||
}
|
||||
}
|
||||
for (const handle of handles) {
|
||||
const x = handle.x * canvas.width;
|
||||
const y = handle.y * canvas.height;
|
||||
context.beginPath();
|
||||
context.arc(x, y, handle.id === selectedId ? 9 : 7, 0, Math.PI * 2);
|
||||
context.fillStyle = handle.id === selectedId ? "#ffcf4a" : "#58c7ff";
|
||||
context.fill();
|
||||
context.strokeStyle = "#111";
|
||||
context.lineWidth = 2;
|
||||
context.stroke();
|
||||
}
|
||||
removeButton.disabled = !selectedId;
|
||||
startButton.disabled = !preparedImage || handles.length === 0 || generating;
|
||||
}
|
||||
|
||||
function canvasPoint(event) {
|
||||
const rect = canvas.getBoundingClientRect();
|
||||
return {
|
||||
x: Math.min(1, Math.max(0, (event.clientX - rect.left) / rect.width)),
|
||||
y: Math.min(1, Math.max(0, (event.clientY - rect.top) / rect.height)),
|
||||
};
|
||||
}
|
||||
function nearest(point) {
|
||||
let match = null;
|
||||
let distance = Infinity;
|
||||
for (const handle of handles) {
|
||||
const value = Math.hypot(handle.x - point.x, handle.y - point.y);
|
||||
if (value < distance && value < 18 / canvas.clientWidth) {
|
||||
match = handle;
|
||||
distance = value;
|
||||
}
|
||||
}
|
||||
return match;
|
||||
}
|
||||
function sendControl(extra = {}) {
|
||||
if (!generating) return;
|
||||
revision += 1;
|
||||
socket.send(JSON.stringify({ type: "control_update", revision, ...extra }));
|
||||
}
|
||||
|
||||
canvas.addEventListener("pointerdown", (event) => {
|
||||
if (!preparedImage) return;
|
||||
const point = canvasPoint(event);
|
||||
if (addMode) {
|
||||
const handle = { id: crypto.randomUUID(), ...point };
|
||||
handles.push(handle);
|
||||
selectedId = handle.id;
|
||||
addMode = false;
|
||||
addButton.textContent = "Add handle";
|
||||
sendControl({ add: [handle] });
|
||||
draw();
|
||||
return;
|
||||
}
|
||||
const handle = nearest(point);
|
||||
selectedId = handle ? handle.id : null;
|
||||
dragging = Boolean(handle);
|
||||
if (dragging) canvas.setPointerCapture(event.pointerId);
|
||||
draw();
|
||||
});
|
||||
canvas.addEventListener("pointermove", (event) => {
|
||||
if (!dragging || !selectedId) return;
|
||||
const point = canvasPoint(event);
|
||||
const handle = handles.find((item) => item.id === selectedId);
|
||||
if (!handle) return;
|
||||
handle.x = point.x; handle.y = point.y;
|
||||
sendControl({ samples: [{ id: handle.id, ...point, timestamp_ms: performance.now() - sessionStart }] });
|
||||
draw();
|
||||
});
|
||||
canvas.addEventListener("pointerup", () => { dragging = false; });
|
||||
addButton.addEventListener("click", () => {
|
||||
addMode = !addMode;
|
||||
addButton.textContent = addMode ? "Click canvas" : "Add handle";
|
||||
});
|
||||
removeButton.addEventListener("click", () => {
|
||||
if (!selectedId) return;
|
||||
const removed = selectedId;
|
||||
handles = handles.filter((item) => item.id !== removed);
|
||||
selectedId = null;
|
||||
sendControl({ remove: [removed] });
|
||||
draw();
|
||||
});
|
||||
gridInput.addEventListener("change", draw);
|
||||
radiusInput.addEventListener("input", () => {
|
||||
radiusOutput.textContent = Number(radiusInput.value).toFixed(2);
|
||||
sendControl({ radius: Number(radiusInput.value) });
|
||||
});
|
||||
|
||||
prepareButton.addEventListener("click", async () => {
|
||||
const file = imageInput.files[0];
|
||||
if (!file) {
|
||||
setStatus("Choose an image", "Prepare needs a reference frame.");
|
||||
return;
|
||||
}
|
||||
prepareButton.disabled = true;
|
||||
prepareButton.textContent = "Preparing…";
|
||||
setStatus("Reading image", "Using the browser's native file reader.", true);
|
||||
try {
|
||||
const image = await new Promise((resolve, reject) => {
|
||||
const reader = new FileReader();
|
||||
reader.onload = () => resolve(reader.result);
|
||||
reader.onerror = () => reject(reader.error || new Error("Failed to read image"));
|
||||
reader.readAsDataURL(file);
|
||||
});
|
||||
setStatus("Sending image", "The server will encode the prompt and reference frame next.", true);
|
||||
socket.send(JSON.stringify({
|
||||
type: "prepare",
|
||||
image,
|
||||
prompt: promptInput.value,
|
||||
}));
|
||||
} catch (error) {
|
||||
prepareButton.disabled = false;
|
||||
prepareButton.textContent = "Prepare";
|
||||
setStatus("Could not read image", error.message || String(error));
|
||||
}
|
||||
});
|
||||
startButton.addEventListener("click", () => {
|
||||
revision = 0;
|
||||
download.hidden = true;
|
||||
startButton.disabled = true;
|
||||
setStatus("Sending controls", "Starting the fixed DMD 4-step, CFG-free SF sampler.", true);
|
||||
socket.send(JSON.stringify({
|
||||
type: "start",
|
||||
handles,
|
||||
radius: Number(radiusInput.value),
|
||||
seed: Number($("seed").value),
|
||||
}));
|
||||
});
|
||||
stopButton.addEventListener("click", () => {
|
||||
socket.send(JSON.stringify({ type: "stop" }));
|
||||
stopButton.disabled = true;
|
||||
setStatus("Finishing current block");
|
||||
});
|
||||
connect();
|
||||
})();
|
||||
@@ -0,0 +1,56 @@
|
||||
<!doctype html>
|
||||
<html lang="en">
|
||||
<head>
|
||||
<meta charset="utf-8">
|
||||
<meta name="viewport" content="width=device-width,initial-scale=1">
|
||||
<title>WanTrack Control</title>
|
||||
<link rel="stylesheet" href="/static/styles.css">
|
||||
</head>
|
||||
<body>
|
||||
<main>
|
||||
<header>
|
||||
<div>
|
||||
<h1>WanTrack Control</h1>
|
||||
<p>Drag points to steer the next causal block.</p>
|
||||
</div>
|
||||
<div id="status" role="status" aria-live="polite">
|
||||
<span id="status-label">Disconnected</span>
|
||||
<small id="status-detail">Waiting for the server connection.</small>
|
||||
</div>
|
||||
</header>
|
||||
|
||||
<section class="controls">
|
||||
<label class="file">Image <input id="image" type="file" accept="image/*"></label>
|
||||
<label class="prompt">Prompt <input id="prompt" type="text" placeholder="Optional scene description"></label>
|
||||
<button id="prepare">Prepare</button>
|
||||
</section>
|
||||
|
||||
<section class="workspace">
|
||||
<div class="canvas-panel">
|
||||
<div class="canvas-toolbar">
|
||||
<button id="add" disabled>Add handle</button>
|
||||
<button id="remove" disabled>Delete selected</button>
|
||||
<label><input id="grid" type="checkbox"> Grid</label>
|
||||
</div>
|
||||
<canvas id="canvas" width="832" height="480"></canvas>
|
||||
<label class="radius">Radius <input id="radius" type="range" min="0.03" max="0.45" step="0.01" value="0.15"><output>0.15</output></label>
|
||||
</div>
|
||||
<div class="video-panel">
|
||||
<video id="video" muted autoplay playsinline controls></video>
|
||||
<a id="download" hidden>Download completed MP4</a>
|
||||
</div>
|
||||
</section>
|
||||
|
||||
<section class="sampling">
|
||||
<label>Seed <input id="seed" type="number" value="0" step="1"></label>
|
||||
<div class="recipe">
|
||||
<strong>SF checkpoint</strong>
|
||||
<span>DMD [1000,750,500,250] · CFG off · TAEHV · relativistic RoPE · local 6 · sink 1</span>
|
||||
</div>
|
||||
<button id="start" disabled>Start</button>
|
||||
<button id="stop" disabled>Stop</button>
|
||||
</section>
|
||||
</main>
|
||||
<script src="/static/app.js"></script>
|
||||
</body>
|
||||
</html>
|
||||
@@ -0,0 +1,57 @@
|
||||
:root {
|
||||
color-scheme: dark;
|
||||
font-family: Inter, ui-sans-serif, system-ui, sans-serif;
|
||||
background: #111315;
|
||||
color: #edf0f2;
|
||||
}
|
||||
* { box-sizing: border-box; }
|
||||
body { margin: 0; }
|
||||
main { width: min(1180px, calc(100% - 32px)); margin: 24px auto; }
|
||||
header { display: flex; align-items: start; justify-content: space-between; gap: 16px; }
|
||||
h1 { margin: 0; font-size: 22px; }
|
||||
p { margin: 6px 0 20px; color: #929ba3; }
|
||||
#status {
|
||||
display: grid; grid-template-columns: auto 1fr; gap: 2px 9px; min-width: 300px;
|
||||
padding: 9px 12px; border-radius: 10px; background: #23282d; color: #dce2e7;
|
||||
}
|
||||
#status::before {
|
||||
content: ""; grid-row: 1 / 3; align-self: center; width: 8px; height: 8px;
|
||||
border-radius: 50%; background: #62d58b;
|
||||
}
|
||||
#status[data-busy="true"]::before {
|
||||
width: 12px; height: 12px; border: 2px solid #65717a; border-top-color: #edf0f2;
|
||||
background: transparent; animation: spin .8s linear infinite;
|
||||
}
|
||||
#status-label { font-size: 12px; font-weight: 700; }
|
||||
#status-detail { color: #9ea8b0; font-size: 11px; }
|
||||
@keyframes spin { to { transform: rotate(360deg); } }
|
||||
.controls, .sampling, .canvas-toolbar { display: flex; align-items: end; gap: 10px; flex-wrap: wrap; }
|
||||
label { color: #aeb6bd; font-size: 12px; }
|
||||
input[type="text"], input[type="number"], input[type="file"] {
|
||||
display: block; margin-top: 5px; min-height: 36px; border: 1px solid #353b41;
|
||||
border-radius: 7px; background: #1a1e22; color: inherit; padding: 7px 9px;
|
||||
}
|
||||
.prompt { flex: 1; }
|
||||
.prompt input { width: 100%; }
|
||||
button, #download {
|
||||
border: 0; border-radius: 7px; background: #e7ebee; color: #111315;
|
||||
min-height: 36px; padding: 8px 13px; font-weight: 650; cursor: pointer;
|
||||
}
|
||||
button:disabled { opacity: .38; cursor: default; }
|
||||
.workspace { display: grid; grid-template-columns: 1fr 1fr; gap: 16px; margin: 16px 0; }
|
||||
.canvas-panel, .video-panel { min-width: 0; border: 1px solid #2e3439; background: #181c20; border-radius: 10px; padding: 10px; }
|
||||
.canvas-toolbar { margin-bottom: 8px; }
|
||||
canvas, video { display: block; width: 100%; aspect-ratio: 832 / 480; object-fit: contain; background: #090a0b; border-radius: 6px; }
|
||||
canvas { touch-action: none; cursor: crosshair; }
|
||||
.radius { display: flex; align-items: center; gap: 9px; margin-top: 10px; }
|
||||
.radius input { flex: 1; }
|
||||
#download { display: inline-block; margin-top: 10px; text-decoration: none; }
|
||||
#download[hidden] { display: none; }
|
||||
.sampling input { width: 92px; }
|
||||
.recipe {
|
||||
display: grid; gap: 2px; min-height: 36px; padding: 6px 10px;
|
||||
border: 1px solid #353b41; border-radius: 7px; background: #1a1e22;
|
||||
}
|
||||
.recipe strong { color: #edf0f2; font-size: 12px; }
|
||||
.recipe span { color: #8f99a1; font-size: 11px; }
|
||||
@media (max-width: 800px) { .workspace { grid-template-columns: 1fr; } }
|
||||
@@ -0,0 +1,267 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
import threading
|
||||
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from apps.wantrack_control.media import EncodedMediaSegment
|
||||
from apps.wantrack_control.server import create_app
|
||||
|
||||
|
||||
@dataclass
|
||||
class _Prepared:
|
||||
image: Image.Image
|
||||
prompt: str
|
||||
|
||||
|
||||
@dataclass
|
||||
class _Block:
|
||||
block_index: int
|
||||
pixel_frames: np.ndarray
|
||||
applied_revision: int
|
||||
radius: float = 0.15
|
||||
active_handle_ids: tuple[str, ...] = ("h", )
|
||||
|
||||
|
||||
class _FakeSession:
|
||||
def __init__(
|
||||
self,
|
||||
gates: list[threading.Event],
|
||||
*,
|
||||
fail_at: int | None = None,
|
||||
) -> None:
|
||||
self.gates = gates
|
||||
self.fail_at = fail_at
|
||||
self.block_index = 0
|
||||
self.pending_revision = 0
|
||||
self.closed_reason = None
|
||||
|
||||
def start(self, image, prompt, handles, sampling, *, radius):
|
||||
assert image.prompt == prompt
|
||||
assert handles
|
||||
assert sampling == {
|
||||
"seed": 0,
|
||||
"num_inference_steps": 4,
|
||||
"text_guidance_scale": 1.0,
|
||||
"motion_guidance_scale": 1.0,
|
||||
"motion_cfg": False,
|
||||
}
|
||||
assert radius > 0
|
||||
|
||||
def apply_control_revision(self, revision, **kwargs):
|
||||
del kwargs
|
||||
if revision <= self.pending_revision:
|
||||
return False
|
||||
self.pending_revision = revision
|
||||
return True
|
||||
|
||||
def generate_next_block(self):
|
||||
index = self.block_index
|
||||
applied = self.pending_revision
|
||||
if index < len(self.gates):
|
||||
assert self.gates[index].wait(timeout=5)
|
||||
if self.fail_at == index:
|
||||
raise RuntimeError("fake generation failure")
|
||||
self.block_index += 1
|
||||
return _Block(
|
||||
block_index=index,
|
||||
pixel_frames=np.zeros((2, 8, 8, 3), dtype=np.uint8),
|
||||
applied_revision=applied,
|
||||
)
|
||||
|
||||
def close(self, reason):
|
||||
self.closed_reason = reason
|
||||
|
||||
|
||||
class _FakeRuntime:
|
||||
fps = 16.0
|
||||
chunk_size = 3
|
||||
decoder_name = "TAEHV (taew2_1)"
|
||||
causal_recipe = {
|
||||
"local_attn_size": 6,
|
||||
"sink_size": 1,
|
||||
"rope_cache_policy": "relativistic",
|
||||
}
|
||||
|
||||
def __init__(self, session):
|
||||
self.session = session
|
||||
|
||||
def prepare(self, image_bytes, prompt):
|
||||
assert image_bytes
|
||||
return _Prepared(Image.new("RGB", (16, 12), "blue"), prompt)
|
||||
|
||||
def create_session(self):
|
||||
return self.session
|
||||
|
||||
|
||||
class _FakeWriter:
|
||||
def __init__(self, output_root, *, fps):
|
||||
assert fps == 16.0
|
||||
self.session_id = "fake-download"
|
||||
self.session_dir = Path(output_root) / self.session_id
|
||||
self.session_dir.mkdir(parents=True, exist_ok=True)
|
||||
self.block_paths = []
|
||||
|
||||
def encode_block(self, frames, block_index):
|
||||
assert frames.shape[-1] == 3
|
||||
path = self.session_dir / f"block-{block_index}.mp4"
|
||||
path.write_bytes(b"complete-block")
|
||||
self.block_paths.append(path)
|
||||
return EncodedMediaSegment(
|
||||
block_index=block_index,
|
||||
init_bytes=b"init",
|
||||
media_bytes=f"media-{block_index}".encode(),
|
||||
path=path,
|
||||
)
|
||||
|
||||
def finalize(self):
|
||||
if not self.block_paths:
|
||||
return None
|
||||
path = self.session_dir / "wantrack_control.mp4"
|
||||
path.write_bytes(b"".join(item.read_bytes()
|
||||
for item in self.block_paths))
|
||||
return path
|
||||
|
||||
|
||||
def _json(websocket):
|
||||
message = websocket.receive()
|
||||
assert message["type"] == "websocket.send"
|
||||
assert "text" in message
|
||||
import json
|
||||
return json.loads(message["text"])
|
||||
|
||||
|
||||
def _bytes(websocket):
|
||||
message = websocket.receive()
|
||||
assert message["type"] == "websocket.send"
|
||||
return message["bytes"]
|
||||
|
||||
|
||||
def _prepare_and_start(websocket):
|
||||
import base64
|
||||
import io
|
||||
|
||||
image = io.BytesIO()
|
||||
Image.new("RGB", (8, 8)).save(image, format="PNG")
|
||||
websocket.send_json({
|
||||
"type": "prepare",
|
||||
"image": base64.b64encode(image.getvalue()).decode(),
|
||||
"prompt": "",
|
||||
})
|
||||
assert _json(websocket)["phase"] == "loading_model"
|
||||
assert _json(websocket)["phase"] == "preparing_input"
|
||||
prepared = _json(websocket)
|
||||
assert prepared["type"] == "prepared"
|
||||
assert prepared["prepare_ms"] >= 0
|
||||
assert prepared["causal_recipe"]["rope_cache_policy"] == "relativistic"
|
||||
assert prepared["decoder"] == "TAEHV (taew2_1)"
|
||||
websocket.send_json({
|
||||
"type": "start",
|
||||
"handles": [{
|
||||
"id": "h",
|
||||
"x": 0.5,
|
||||
"y": 0.5,
|
||||
}],
|
||||
"radius": 0.15,
|
||||
# Client overrides are ignored for the fixed SF recipe.
|
||||
"steps": 99,
|
||||
"text_guidance": 9.0,
|
||||
"motion_guidance": 9.0,
|
||||
})
|
||||
assert _json(websocket)["phase"] == "starting_session"
|
||||
started = _json(websocket)
|
||||
assert started["type"] == "session_started"
|
||||
assert started["num_inference_steps"] == 4
|
||||
assert started["cfg_enabled"] is False
|
||||
assert started["causal_recipe"]["local_attn_size"] == 6
|
||||
assert started["decoder"] == "TAEHV (taew2_1)"
|
||||
|
||||
|
||||
def test_two_block_binary_order_future_update_and_stop(tmp_path):
|
||||
gates = [threading.Event(), threading.Event()]
|
||||
session = _FakeSession(gates)
|
||||
app = create_app(
|
||||
_FakeRuntime(session),
|
||||
output_dir=tmp_path,
|
||||
writer_factory=_FakeWriter,
|
||||
)
|
||||
with TestClient(app) as client:
|
||||
with client.websocket_connect("/ws") as websocket:
|
||||
_prepare_and_start(websocket)
|
||||
started = _json(websocket)
|
||||
assert started["type"] == "block_started"
|
||||
assert started["block_index"] == 0
|
||||
assert started["num_inference_steps"] == 4
|
||||
assert started["cfg_enabled"] is False
|
||||
websocket.send_json({
|
||||
"type": "control_update",
|
||||
"revision": 1,
|
||||
"samples": [{
|
||||
"id": "h",
|
||||
"x": 0.7,
|
||||
"y": 0.5,
|
||||
"timestamp_ms": 10,
|
||||
}],
|
||||
})
|
||||
websocket.send_json({
|
||||
"type": "control_update",
|
||||
"revision": 1,
|
||||
"samples": [],
|
||||
})
|
||||
stale = _json(websocket)
|
||||
assert stale["status"] == "ignored_stale"
|
||||
gates[0].set()
|
||||
assert _json(websocket)["type"] == "block_encoding"
|
||||
assert _json(websocket)["type"] == "media_init"
|
||||
assert _bytes(websocket) == b"init"
|
||||
assert _bytes(websocket) == b"media-0"
|
||||
assert _json(websocket)["type"] == "media_segment_complete"
|
||||
|
||||
assert _json(websocket)["type"] == "block_started"
|
||||
websocket.send_json({"type": "stop"})
|
||||
gates[1].set()
|
||||
applied = _json(websocket)
|
||||
assert applied["type"] == "control_applied"
|
||||
assert applied["revision"] == 1
|
||||
assert _json(websocket)["type"] == "block_encoding"
|
||||
assert _bytes(websocket) == b"media-1"
|
||||
assert _json(websocket)["type"] == "media_segment_complete"
|
||||
complete = _json(websocket)
|
||||
assert complete["type"] == "stream_complete"
|
||||
assert complete["blocks"] == 2
|
||||
assert complete["download_url"]
|
||||
response = client.get(complete["download_url"])
|
||||
assert response.status_code == 200
|
||||
assert response.content
|
||||
assert client.get("/healthz").json()["active"] is False
|
||||
|
||||
|
||||
def test_error_releases_lock_and_preserves_completed_prefix(tmp_path):
|
||||
gates = [threading.Event(), threading.Event()]
|
||||
gates[0].set()
|
||||
gates[1].set()
|
||||
session = _FakeSession(gates, fail_at=1)
|
||||
app = create_app(
|
||||
_FakeRuntime(session),
|
||||
output_dir=tmp_path,
|
||||
writer_factory=_FakeWriter,
|
||||
)
|
||||
with TestClient(app) as client:
|
||||
with client.websocket_connect("/ws") as websocket:
|
||||
_prepare_and_start(websocket)
|
||||
assert _json(websocket)["type"] == "block_started"
|
||||
assert _json(websocket)["type"] == "block_encoding"
|
||||
assert _json(websocket)["type"] == "media_init"
|
||||
assert _bytes(websocket) == b"init"
|
||||
assert _bytes(websocket) == b"media-0"
|
||||
assert _json(websocket)["type"] == "media_segment_complete"
|
||||
assert _json(websocket)["type"] == "block_started"
|
||||
error = _json(websocket)
|
||||
assert error["type"] == "error"
|
||||
assert error["download_url"]
|
||||
assert client.get(error["download_url"]).content
|
||||
assert client.get("/healthz").json()["active"] is False
|
||||
@@ -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
|
||||
|
@@ -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
|
||||
@@ -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);
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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"
|
||||
]
|
||||
|
||||
@@ -0,0 +1,61 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Configuration for track-conditioned Wan transformers."""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig
|
||||
from fastvideo.configs.models.dits.wanvideo import WanVideoArchConfig, WanVideoConfig
|
||||
|
||||
|
||||
def _default_track_config() -> dict[str, int | bool]:
|
||||
return {
|
||||
"id_dim": 64,
|
||||
"track_channels": 16,
|
||||
"vae_spatial_compression": 8,
|
||||
"vae_temporal_compression": 4,
|
||||
"max_track_id": 100_000,
|
||||
"zero_init_head": False,
|
||||
"use_bias": False,
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class TrackWanVideoArchConfig(WanVideoArchConfig):
|
||||
"""Wan I2V input widened with a latent-aligned point-track map.
|
||||
|
||||
Channel order is fixed and checkpoint-visible:
|
||||
noisy latent (16), I2V mask (4), first-frame latent (16), track map (16).
|
||||
"""
|
||||
|
||||
in_channels: int = 52
|
||||
out_channels: int = 16
|
||||
image_dim: int = 1280
|
||||
added_kv_proj_dim: int | None = 5120
|
||||
track_config: dict[str, int | bool] = field(default_factory=_default_track_config)
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
super().__post_init__()
|
||||
track_channels = int(self.track_config.get("track_channels", 0))
|
||||
expected_in_channels = self.num_channels_latents + 20 + track_channels
|
||||
if track_channels <= 0:
|
||||
raise ValueError("track_config.track_channels must be positive")
|
||||
if self.in_channels != expected_in_channels:
|
||||
raise ValueError("TrackWan in_channels must equal latent channels + 20 I2V "
|
||||
f"channels + track channels; got {self.in_channels}, expected "
|
||||
f"{expected_in_channels}")
|
||||
|
||||
|
||||
@dataclass
|
||||
class TrackWanVideoConfig(WanVideoConfig):
|
||||
arch_config: DiTArchConfig = field(default_factory=TrackWanVideoArchConfig)
|
||||
prefix: str = "Wan"
|
||||
|
||||
|
||||
@dataclass
|
||||
class CausalTrackWanVideoArchConfig(TrackWanVideoArchConfig):
|
||||
"""Explicit config type for the causal TrackWan architecture."""
|
||||
|
||||
|
||||
@dataclass
|
||||
class CausalTrackWanVideoConfig(TrackWanVideoConfig):
|
||||
arch_config: DiTArchConfig = field(default_factory=CausalTrackWanVideoArchConfig)
|
||||
@@ -91,6 +91,11 @@ class WanVideoArchConfig(DiTArchConfig):
|
||||
# RoPE policy for the causal-rollout paths (causal Wan / MatrixGame2).
|
||||
# "relativistic" keeps long rollouts in-distribution; a no-op unless sink_size > 0 and local_attn_size > 0.
|
||||
rope_cache_policy: str = "absolute"
|
||||
# Full-sequence training attention implementation for causal models:
|
||||
# "flex" (FlexAttention BlockMask, default), "triton" (fused sink +
|
||||
# rolling-window kernel with exact relativistic sink RoPE correction), or
|
||||
# "reference" (slow pure-PyTorch, for tests).
|
||||
causal_train_attention: str = "flex"
|
||||
|
||||
# AnyFlow dual-timestep conditioning. Defaults preserve bit-identity with
|
||||
# the legacy single-timestep forward (no delta_embedder allocated, no
|
||||
|
||||
@@ -0,0 +1,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
|
||||
@@ -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()
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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,
|
||||
]
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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]
|
||||
|
||||
@@ -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,
|
||||
},
|
||||
}
|
||||
|
||||
@@ -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 {}
|
||||
|
||||
|
||||
@@ -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()
|
||||
@@ -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
|
||||
# ----------------------------------------------------------
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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
Reference in New Issue
Block a user