Compare commits
2
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
5671ebc718 | ||
|
|
57a2fa9ec1 |
@@ -0,0 +1,278 @@
|
||||
# LTX-2.3 distilled — NVFP4 / FA4 / CUDA-graphs benchmark & reproduction
|
||||
|
||||
End-to-end instructions to reproduce the LTX-2.3 distilled inference speed
|
||||
sweep on a **Blackwell GB300 (sm_103)**, including the FP4 attention (FA4-FP4)
|
||||
path and the CUDA-graphs optimization that makes NVFP4 the fastest config.
|
||||
|
||||
## TL;DR result (t2v, 832×1280, 121 frames, 8 denoise + 3 refine, GB300)
|
||||
|
||||
| config | denoise | refine | **e2e** |
|
||||
|--------------------------------|---------|--------|---------|
|
||||
| bf16 + FA2 (baseline) compile | 1.476 | 2.759 | 5.15 s |
|
||||
| bf16 + FA4 compile | 1.423 | 1.754 | 4.06 s |
|
||||
| nvfp4 + FA4-FP4 compile | 3.777 | 1.655 | 6.35 s |
|
||||
| bf16 + FA4 + **cudagraphs** | 1.105 | 1.728 | 3.81 s |
|
||||
| **nvfp4 + FA4-FP4 + cudagraphs** | **0.859** | **1.193** | **3.00 s** ⭐ |
|
||||
|
||||
Key findings:
|
||||
- **FA4-bf16** (CuTeDSL attention kernel) is a big win over FlashAttention-2.
|
||||
- The NVFP4 denoise penalty is **per-step launch-bound** (per-layer FP4 quant +
|
||||
mm_fp4 = hundreds of tiny kernel launches/step; ~flat across resolution).
|
||||
- **CUDA graphs** removes that launch overhead → NVFP4 denoise 3.78→0.86 s, and
|
||||
nvfp4+FA4+cudagraphs becomes the fastest end-to-end config.
|
||||
|
||||
---
|
||||
|
||||
## 1. Hardware
|
||||
|
||||
- NVIDIA **GB300** (Grace-Blackwell, compute capability **sm_103**, 256 GB).
|
||||
- On a mixed box, pin it: this repo's box has GB300 as **GPU 1**
|
||||
(`CUDA_VISIBLE_DEVICES=1`), GPU 0 is an RTX PRO 6000 workstation card.
|
||||
- ⚠️ The installed torch wheel lists archs `sm_80/90/100/120` (no `sm_103`).
|
||||
Compute works on GB300, but importing torch/fastvideo with **both** GPUs
|
||||
visible crashes in the capability check — always set `CUDA_VISIBLE_DEVICES=1`.
|
||||
|
||||
## 2. Conda env + PyTorch
|
||||
|
||||
Python 3.12, CUDA 12.8 torch build:
|
||||
|
||||
```bash
|
||||
# torch 2.11.0+cu128 (aarch64/sbsa wheel for Grace) + torchvision
|
||||
pip install --index-url https://download.pytorch.org/whl/cu128 \
|
||||
'torch==2.11.0+cu128' 'torchvision==0.26.0+cu128'
|
||||
```
|
||||
|
||||
Verified versions on the reference machine:
|
||||
|
||||
```
|
||||
torch==2.11.0+cu128 torchvision==0.26.0+cu128
|
||||
flashinfer-python==0.6.12 nvidia-cutlass-dsl==4.5.2
|
||||
flash-attn-4==0.0.1.dev1330+gd268d2b86 quack-kernels==0.5.0
|
||||
ninja==1.13.0
|
||||
```
|
||||
|
||||
## 3. System CUDA toolkit (for flashinfer / FA4 JIT)
|
||||
|
||||
flashinfer and the FA4-FP4 CuTeDSL kernels JIT-compile on first use and need a
|
||||
full CUDA toolkit with `nvcc` that supports `compute_103`. Use **CUDA 13.2**:
|
||||
|
||||
```bash
|
||||
export CUDA_HOME=/usr/local/cuda-13.2 # must contain bin/nvcc + include/
|
||||
# ensure the conda env bin (with `ninja`) AND cuda bin are on PATH:
|
||||
export PATH="$CONDA_PREFIX/bin:/usr/local/cuda-13.2/bin:$PATH"
|
||||
nvcc --version # -> release 13.2 ; supports compute_103
|
||||
```
|
||||
|
||||
## 4. NVFP4 (DiT-linear FP4) — flashinfer
|
||||
|
||||
```bash
|
||||
pip install flashinfer-python
|
||||
```
|
||||
|
||||
> ⚠️ `flashinfer-python` may try to **downgrade torch** to a CPU build. If it
|
||||
> does, reinstall torch from step 2 afterwards (use `--no-deps` on flashinfer or
|
||||
> reinstall torch). NVFP4 linears JIT-build an sm_103 module on first use; this
|
||||
> needs `nvcc` (CUDA_HOME) and `ninja` (conda env bin) on PATH.
|
||||
|
||||
## 5. FA4-FP4 attention (quantized Q/K) — extra deps + a one-line shim
|
||||
|
||||
FP4 attention (`nvfp4_fa4=True`) needs the hao-ai-lab flash-attention-fp4 fork
|
||||
plus QuACK, and a tiny cutlass-dsl compatibility shim.
|
||||
|
||||
```bash
|
||||
# the FP4 CuTeDSL kernels (installs pkg `flash-attn-4`, providing flash_attn.cute)
|
||||
pip install --no-deps \
|
||||
"git+https://github.com/hao-ai-lab/flash-attention-fp4.git@fp4#subdirectory=flash_attn/cute"
|
||||
# QuACK kernels (imported as `quack`)
|
||||
pip install --no-deps quack-kernels
|
||||
```
|
||||
|
||||
`flash_attn.cute` imports `cutlass.utils.ampere_helpers`, which
|
||||
`nvidia-cutlass-dsl >= 4.5` removed (and we can't downgrade — flashinfer needs
|
||||
>=4.5). Restore just the one symbol it uses (`SMEM_CAPACITY`):
|
||||
|
||||
```bash
|
||||
python - <<'PY'
|
||||
import os, cutlass
|
||||
d = os.path.join(os.path.dirname(cutlass.__file__), "utils")
|
||||
open(os.path.join(d, "ampere_helpers.py"), "w").write(
|
||||
"SMEM_CAPACITY = {"
|
||||
"'sm80': (164-1)*1024, 'sm86': (100-1)*1024, "
|
||||
"'sm87': (164-1)*1024, 'sm89': (100-1)*1024}\n")
|
||||
print("wrote", os.path.join(d, "ampere_helpers.py"))
|
||||
PY
|
||||
```
|
||||
|
||||
Sanity check the whole chain:
|
||||
|
||||
```bash
|
||||
CUDA_VISIBLE_DEVICES=1 python -c "
|
||||
import flash_attn; from flash_attn import flash_attn_func # FA2 core
|
||||
import flash_attn.cute.interface # FA4 cute
|
||||
from fastvideo.attention.utils.flash_attn_cute import flash_attn_fp4_func
|
||||
import fastvideo.attention.backends.flash_attn as fa
|
||||
print('FA4-FP4 available:', fa._FA4_FP4_AVAILABLE)" # -> True
|
||||
```
|
||||
|
||||
## 6. fastvideo code change (already in this branch)
|
||||
|
||||
`fastvideo/attention/backends/flash_attn.py`:
|
||||
- `_nvfp4_quantize_for_fa4` is wrapped as a `torch.library.custom_op` +
|
||||
`register_fake` so torch.compile treats the flashinfer FP4 quant as an opaque
|
||||
leaf (dynamo can't trace its `@functools.cache` lock / JIT `subprocess`).
|
||||
- `FlashAttentionImpl.__init__` pre-builds the FP4 quant module (warms the cache
|
||||
before compile).
|
||||
|
||||
These are required for **compile + FA4-FP4** to work at all. Nothing to do if
|
||||
you're on this branch.
|
||||
|
||||
## 7. Model
|
||||
|
||||
Auto-downloaded on first run (~67 GB; skips the redundant `download/` raw
|
||||
checkpoints). To pre-stage and avoid re-fetching:
|
||||
|
||||
```bash
|
||||
python -c "from huggingface_hub import snapshot_download; \
|
||||
print(snapshot_download('FastVideo/LTX-2.3-Distilled-Diffusers', \
|
||||
ignore_patterns=['download/*','*.onnx','*.msgpack'], max_workers=8))"
|
||||
# point runs at the printed snapshot dir via LTX23_MODEL_PATH (optional)
|
||||
```
|
||||
|
||||
## 8. Running the benchmark
|
||||
|
||||
The parametrized harness is `examples/inference/basic/bench_ltx2_3_distilled_t2v.py`.
|
||||
All knobs are env vars:
|
||||
|
||||
| env var | default | meaning |
|
||||
|-----------------------|---------|---------|
|
||||
| `LTX23_COMPILE` | 0 | torch.compile DiT+TE+VAE |
|
||||
| `LTX23_NVFP4` | 0 | NVFP4 quant the DiT linears |
|
||||
| `LTX23_FA4` | 0 | FA4-FP4 attention (quant Q/K, sets `nvfp4_fa4`) |
|
||||
| `LTX23_COMPILE_MODE` | default | DiT compile mode (`default` / `reduce-overhead`) |
|
||||
| `LTX23_NO_CGTREES` | 0 | set 1 with reduce-overhead (see caveats) |
|
||||
| `LTX23_COMPILE_TE` | 1 | compile text encoder (set 0 for cudagraphs) |
|
||||
| `LTX23_COMPILE_VAE` | 1 | compile VAE (kept on `default` mode) |
|
||||
| `LTX23_HEIGHT`/`_WIDTH` | 1280/832 | resolution (H×W) |
|
||||
| `LTX23_NUM_FRAMES` | 121 | 121≈5 s, 481≈20 s @24fps |
|
||||
| `LTX23_STEPS` | 8 | denoise steps (refine fixed at 3) |
|
||||
| `LTX23_VAE_TILING` | 0 | tile the VAE decode — set 1 for long/high-res (avoids the 32-bit decode overflow; see §11) |
|
||||
| `LTX23_DECODE` | 1 | set 0 → `output_type="latent"`: skip VAE decode + video save to measure DiT (denoise/refine) timing cleanly on long runs |
|
||||
| `LTX23_WARMUP` | 2 / 1 | warmup runs (default 2 if compile else 1) |
|
||||
| `LTX23_MEASURED` | 3 | measured runs to average |
|
||||
| `LTX23_MODEL_PATH` | hub id | local snapshot dir (optional) |
|
||||
|
||||
Common env block for every run:
|
||||
|
||||
```bash
|
||||
export CUDA_VISIBLE_DEVICES=1 CUDA_DEVICE_ORDER=PCI_BUS_ID
|
||||
export CUDA_HOME=/usr/local/cuda-13.2
|
||||
export PATH="$CONDA_PREFIX/bin:/usr/local/cuda-13.2/bin:$PATH"
|
||||
export TORCHINDUCTOR_CACHE_DIR=$HOME/.cache/torchinductor_ltx23 # caches cold compiles
|
||||
cd <repo root>
|
||||
SCRIPT=examples/inference/basic/bench_ltx2_3_distilled_t2v.py
|
||||
```
|
||||
|
||||
### The 5 headline configs (832×1280, 8+3)
|
||||
|
||||
```bash
|
||||
# bf16 + FA2 baseline (force FA2 by hiding flash_attn.cute is not needed; FA4 is
|
||||
# auto-selected once cute is installed — to get the FA2 number, benchmark before
|
||||
# installing the FA4 deps, or compare against the table above)
|
||||
|
||||
# bf16 + FA4, compile
|
||||
LTX23_COMPILE=1 LTX23_NVFP4=0 LTX23_FA4=0 python $SCRIPT # ~4.06 s
|
||||
|
||||
# nvfp4 + FA4-FP4, compile
|
||||
LTX23_COMPILE=1 LTX23_NVFP4=1 LTX23_FA4=1 python $SCRIPT # ~6.35 s
|
||||
|
||||
# bf16 + FA4 + CUDA graphs
|
||||
LTX23_COMPILE=1 LTX23_NVFP4=0 LTX23_FA4=0 \
|
||||
LTX23_COMPILE_MODE=reduce-overhead LTX23_NO_CGTREES=1 LTX23_COMPILE_TE=0 \
|
||||
python $SCRIPT # ~3.81 s
|
||||
|
||||
# nvfp4 + FA4-FP4 + CUDA graphs ⭐ fastest
|
||||
LTX23_COMPILE=1 LTX23_NVFP4=1 LTX23_FA4=1 \
|
||||
LTX23_COMPILE_MODE=reduce-overhead LTX23_NO_CGTREES=1 LTX23_COMPILE_TE=0 \
|
||||
python $SCRIPT # ~3.00 s
|
||||
```
|
||||
|
||||
Each run prints a per-stage breakdown and an averaged e2e over 3 measured runs
|
||||
(after warmups). First run pays a one-time cold compile (~15 min) cached in
|
||||
`$TORCHINDUCTOR_CACHE_DIR`.
|
||||
|
||||
i2v variants (anchor an image at frame 0) live in the same folder:
|
||||
`basic_ltx2_3_distilled_i2v_{uncompiled,compiled}{,_nvfp4}.py` — set
|
||||
`LTX23_I2V_IMAGE=/path/to.jpg`.
|
||||
|
||||
## 9. CUDA graphs caveats (why the extra flags)
|
||||
|
||||
Plain `mode="reduce-overhead"` fails with the FP4 pipeline; two workarounds are
|
||||
needed and are wired to the env flags above:
|
||||
|
||||
1. **`LTX23_NO_CGTREES=1`** → sets `torch._inductor.config.triton.cudagraph_trees
|
||||
= False`. cudagraph_**trees** rejects the flashinfer FP4 custom ops because
|
||||
they allocate tensors inside the captured region that inductor doesn't track
|
||||
(`Detected N tensor(s) in the cudagraph pool not tracked as outputs`).
|
||||
2. **`LTX23_COMPILE_TE=0`** → don't cudagraph the text encoder. cudagraphs on
|
||||
`gemma.py` triggers cross-module static-buffer aliasing
|
||||
(`accessing tensor output of CUDAGraphs that has been overwritten`). VAE stays
|
||||
compiled on `mode="default"`. So cudagraphs is applied to the **DiT only**.
|
||||
|
||||
Correctness: with `cudagraph_trees=False` the trees safety check is off, so we
|
||||
validated the output is a real, coherent video (121 frames, healthy stats) and
|
||||
that the cudagraph result is the *closest* match to its own non-cudagraph output
|
||||
(PSNR 17.95 dB — higher than any cross-config pair). The per-frame difference is
|
||||
ordinary diffusion sensitivity to kernel/numeric changes, not corruption.
|
||||
|
||||
## 10. Long-token / longer-video scaling (241 / 481 frames)
|
||||
|
||||
DiT token length is `N = T_lat × H_lat × W_lat` with VAE compression 32×
|
||||
spatial / 8× temporal and DiT patch 1: `T_lat = (frames−1)/8 + 1`,
|
||||
`H_lat = H/32`, `W_lat = W/32`. At 832×1280 each extra latent frame adds
|
||||
40×26 = 1040 tokens, so **frame count is the clean linear axis** for long-context
|
||||
tests — use `frames = 8k+1` (121 / 241 / 481), else `(frames−1)//8` truncates.
|
||||
|
||||
Measured on GB300 (t2v, 832×1280, 8+3, nvfp4+FA4-FP4+CUDA graphs vs bf16+FA4+CG),
|
||||
seconds:
|
||||
|
||||
| frames (~dur) | N (refine) | nvfp4 denoise / refine / **e2e** | bf16 denoise / refine / e2e |
|
||||
|---------------|-----------:|-----------------------------------|------------------------------|
|
||||
| 121 (~5 s) | 16,640 | 0.859 / 1.193 / **3.00** | 1.082 / 1.725 / 3.77 |
|
||||
| 241 (~10 s) | 32,240 | 1.549 / 2.956 / **6.15** | 2.051 / 4.073 / 7.85 |
|
||||
| 481 (~20 s) | 63,440 | 3.261 / 8.259 / **11.9*** | — |
|
||||
|
||||
*481-frame e2e is **DiT-only** (`LTX23_DECODE=0`): the full-res tiled VAE decode
|
||||
is CPU-bound and dominates wall-time without saying anything about DiT scaling,
|
||||
so it is excluded. The 121/241 e2e include decode + video save.
|
||||
|
||||
Finding: the low-res **denoise** stays ~linear in frames (launch-bound), but the
|
||||
full-res **refine** goes super-linear — attention (N²) is ~29 % of refine at
|
||||
121 f, ~44 % at 241 f, ~61 % at 481 f. nvfp4's ~20 % e2e win over bf16 holds
|
||||
across all lengths.
|
||||
|
||||
```bash
|
||||
# 241-frame full e2e (decode still fits 32-bit untiled at 241 f)
|
||||
LTX23_COMPILE=1 LTX23_NVFP4=1 LTX23_FA4=1 \
|
||||
LTX23_COMPILE_MODE=reduce-overhead LTX23_NO_CGTREES=1 LTX23_COMPILE_TE=0 \
|
||||
LTX23_NUM_FRAMES=241 python $SCRIPT # nvfp4 ~6.15 s
|
||||
|
||||
# 481-frame DiT-only (skip the slow tiled decode, measure denoise/refine)
|
||||
LTX23_COMPILE=1 LTX23_NVFP4=1 LTX23_FA4=1 \
|
||||
LTX23_COMPILE_MODE=reduce-overhead LTX23_NO_CGTREES=1 LTX23_COMPILE_TE=0 \
|
||||
LTX23_NUM_FRAMES=481 LTX23_DECODE=0 python $SCRIPT # DiT ~11.9 s
|
||||
|
||||
# 481-frame WITH decode (needs tiling; decode is slow/CPU-bound)
|
||||
LTX23_COMPILE=1 LTX23_NVFP4=1 LTX23_FA4=1 \
|
||||
LTX23_COMPILE_MODE=reduce-overhead LTX23_NO_CGTREES=1 LTX23_COMPILE_TE=0 \
|
||||
LTX23_NUM_FRAMES=481 LTX23_VAE_TILING=1 python $SCRIPT
|
||||
```
|
||||
|
||||
## 11. Known limitations / TODO
|
||||
|
||||
- **Long / high-res VAE decode**: past ~241 frames at 832×1280 (and ≥1920×1088),
|
||||
untiled decode overflows `input tensor must fit into 32-bit index math`. It runs
|
||||
with `LTX23_VAE_TILING=1`, but the tiled decode is **CPU-bound and very slow**
|
||||
(the per-tile trapezoidal-blend stitch dominates) — not yet optimized. Use
|
||||
`LTX23_DECODE=0` to benchmark DiT timing without it.
|
||||
- cudagraphs is DiT-only here; text-encoder/VAE graph capture needs
|
||||
`cudagraph_mark_step_begin()` plumbing to be safe.
|
||||
@@ -0,0 +1,241 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""LTX-2.3 distilled image-to-video — bf16, torch.compile ON + timing.
|
||||
|
||||
Compiled sibling of ``basic_ltx2_3_distilled_i2v_uncompiled.py``. Same
|
||||
generation recipe (8 denoise + 3 refine steps, CFG=1, no refine LoRA — the
|
||||
distilled production recipe), same input/measurement methodology, but with
|
||||
torch.compile fully ENABLED (DiT + text encoder + VAE). Use this to fill the
|
||||
bf16/compile cell of the 2x2 (bf16 vs NVFP4) x (eager vs compile) table.
|
||||
|
||||
NOTE: first warmup pays a long cold-compile (tens of minutes on Blackwell);
|
||||
results are cached in ``$TORCHINDUCTOR_CACHE_DIR`` for later runs.
|
||||
|
||||
CUDA_VISIBLE_DEVICES=1 python \
|
||||
examples/inference/basic/basic_ltx2_3_distilled_i2v_compiled.py
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import time
|
||||
from collections import OrderedDict
|
||||
from pathlib import Path
|
||||
|
||||
import torch._inductor.config as _inductor
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
from fastvideo.utils import maybe_download_model
|
||||
|
||||
os.environ.setdefault("FASTVIDEO_ATTENTION_BACKEND", "FLASH_ATTN")
|
||||
os.environ.setdefault("FASTVIDEO_STAGE_LOGGING", "1")
|
||||
|
||||
# Inductor knobs. ``shape_padding=False`` is mandatory on Blackwell to avoid a
|
||||
# cuBLAS INVALID_VALUE crash inside pad_mm during the refine path. The rest are
|
||||
# autotune-friendliness flags (match the canonical compiled example).
|
||||
_inductor.shape_padding = False
|
||||
_inductor.conv_1x1_as_mm = True
|
||||
_inductor.coordinate_descent_tuning = True
|
||||
_inductor.coordinate_descent_check_all_directions = True
|
||||
_inductor.epilogue_fusion = False
|
||||
|
||||
MODEL_ID = os.path.expandvars(
|
||||
os.path.expanduser(
|
||||
os.getenv("LTX23_MODEL_PATH", "FastVideo/LTX-2.3-Distilled-Diffusers")
|
||||
)
|
||||
)
|
||||
OUTPUT_DIR = Path(
|
||||
os.getenv(
|
||||
"LTX23_OUTPUT_DIR", "outputs_video/ltx2_3_distilled_i2v_compiled"
|
||||
)
|
||||
)
|
||||
I2V_IMAGE = os.getenv("LTX23_I2V_IMAGE", "")
|
||||
DEFAULT_PROMPT = (
|
||||
"A fashion model takes a slow step forward and shifts her weight, "
|
||||
"the soft fabric of her clothing swaying and rippling with the "
|
||||
"motion, her hair shifting gently, soft even studio lighting on a "
|
||||
"clean light background, elegant slow-motion runway feel."
|
||||
)
|
||||
PROMPT = os.getenv("LTX23_I2V_PROMPT", DEFAULT_PROMPT)
|
||||
|
||||
|
||||
def _print_stage_breakdown(result: dict, label: str) -> float | None:
|
||||
logging_info = result.get("logging_info")
|
||||
stages = getattr(logging_info, "stages", None) if logging_info else None
|
||||
if not stages:
|
||||
print(f" [{label}] stage breakdown unavailable")
|
||||
return None
|
||||
print(f" [{label}] stage breakdown:")
|
||||
total = 0.0
|
||||
for name, metrics in stages.items():
|
||||
exec_s = float(metrics.get("execution_time", 0.0))
|
||||
total += exec_s
|
||||
print(f" - {name}: {exec_s:.3f}s")
|
||||
print(f" - stage_sum: {total:.3f}s")
|
||||
return total
|
||||
|
||||
|
||||
def _collect_stage_times(
|
||||
result: dict,
|
||||
stage_times: dict[str, list[float]],
|
||||
stage_order: OrderedDict[str, None],
|
||||
) -> None:
|
||||
logging_info = result.get("logging_info")
|
||||
stages = getattr(logging_info, "stages", None) if logging_info else None
|
||||
if not stages:
|
||||
return
|
||||
for name, metrics in stages.items():
|
||||
stage_order.setdefault(name, None)
|
||||
stage_times.setdefault(name, []).append(
|
||||
float(metrics.get("execution_time", 0.0))
|
||||
)
|
||||
|
||||
|
||||
def _resolve_refine_upsampler(model_root: str) -> Path:
|
||||
"""LTX-2.3 distilled snapshots ship a `spatial_upscaler/` subdir."""
|
||||
for name in ("spatial_upscaler", "spatial_upsampler"):
|
||||
cand = Path(model_root) / name
|
||||
if (cand / "config.json").is_file():
|
||||
return cand
|
||||
raise FileNotFoundError(
|
||||
f"No refine upsampler directory under {model_root}. "
|
||||
f"Expected `{model_root}/spatial_upscaler/config.json`."
|
||||
)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
if not I2V_IMAGE:
|
||||
raise SystemExit(
|
||||
"LTX23_I2V_IMAGE is required for i2v. Example:\n"
|
||||
" export LTX23_I2V_IMAGE=/path/to/portrait_or_product.jpg\n"
|
||||
" python examples/inference/basic/"
|
||||
"basic_ltx2_3_distilled_i2v_compiled.py"
|
||||
)
|
||||
if not Path(I2V_IMAGE).is_file():
|
||||
raise SystemExit(f"LTX23_I2V_IMAGE not found: {I2V_IMAGE}")
|
||||
|
||||
OUTPUT_DIR.mkdir(parents=True, exist_ok=True)
|
||||
model_root = maybe_download_model(MODEL_ID)
|
||||
refine_upsampler_path = _resolve_refine_upsampler(model_root)
|
||||
print(f"Model: {model_root}")
|
||||
print(f"Refine upsampler: {refine_upsampler_path}")
|
||||
print(f"i2v image: {I2V_IMAGE}")
|
||||
print(f"Output dir: {OUTPUT_DIR.resolve()}")
|
||||
print("torch.compile: ENABLED (DiT + text encoder + VAE)")
|
||||
print("DiT quant: bf16")
|
||||
|
||||
pipeline_config = PipelineConfig.from_pretrained(model_root)
|
||||
pipeline_config.dit_config.quant_config = None
|
||||
|
||||
torch_compile_kwargs = {
|
||||
"backend": "inductor",
|
||||
"fullgraph": True,
|
||||
"mode": "default",
|
||||
"dynamic": False,
|
||||
}
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_root,
|
||||
num_gpus=1,
|
||||
ltx2_refine_enabled=True,
|
||||
ltx2_refine_upsampler_path=str(refine_upsampler_path),
|
||||
ltx2_refine_lora_path="",
|
||||
ltx2_refine_num_inference_steps=3,
|
||||
ltx2_refine_guidance_scale=1.0,
|
||||
ltx2_refine_add_noise=True,
|
||||
pipeline_config=pipeline_config,
|
||||
enable_torch_compile=True,
|
||||
enable_torch_compile_text_encoder=True,
|
||||
enable_torch_compile_vae=True,
|
||||
torch_compile_kwargs=torch_compile_kwargs,
|
||||
torch_compile_kwargs_vae=torch_compile_kwargs,
|
||||
dit_cpu_offload=False,
|
||||
text_encoder_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
ltx2_vae_tiling=False,
|
||||
)
|
||||
|
||||
common_kwargs = dict(
|
||||
prompt=PROMPT,
|
||||
negative_prompt="",
|
||||
guidance_scale=1.0,
|
||||
height=1280, width=832,
|
||||
num_frames=121, fps=24,
|
||||
num_inference_steps=8,
|
||||
ltx2_images=[(I2V_IMAGE, 0, 1.0)],
|
||||
ltx2_image_crf=0.0,
|
||||
save_video=True,
|
||||
)
|
||||
|
||||
# Compile → 2 warmups: first pays cold compile + first-shape guards, the
|
||||
# second lets any residual recompiles settle before we measure.
|
||||
warmup_runs = 2
|
||||
measured_runs = 3
|
||||
warmup_secs: list[float] = []
|
||||
measured_secs: list[float] = []
|
||||
stage_times: dict[str, list[float]] = {}
|
||||
stage_order: OrderedDict[str, None] = OrderedDict()
|
||||
|
||||
try:
|
||||
for w in range(warmup_runs):
|
||||
t0 = time.perf_counter()
|
||||
print(f"\n[warmup {w + 1}/{warmup_runs}] compiling + generating…")
|
||||
generator.generate_video(
|
||||
output_path=str(OUTPUT_DIR / f"_warmup_{w + 1}.mp4"),
|
||||
seed=7,
|
||||
**common_kwargs,
|
||||
)
|
||||
dt = time.perf_counter() - t0
|
||||
warmup_secs.append(dt)
|
||||
print(f"[warmup {w + 1}/{warmup_runs}] wall={dt:.1f}s")
|
||||
|
||||
for w in range(warmup_runs):
|
||||
(OUTPUT_DIR / f"_warmup_{w + 1}.mp4").unlink(missing_ok=True)
|
||||
|
||||
for m in range(measured_runs):
|
||||
out_path = (
|
||||
OUTPUT_DIR
|
||||
/ f"output_ltx2_3_distilled_i2v_compiled_run_{m + 1}.mp4"
|
||||
)
|
||||
print(f"\n[measured {m + 1}/{measured_runs}] generating: {out_path}")
|
||||
t0 = time.perf_counter()
|
||||
result = generator.generate_video(
|
||||
output_path=str(out_path),
|
||||
seed=2002 + m,
|
||||
**common_kwargs,
|
||||
)
|
||||
wall = time.perf_counter() - t0
|
||||
e2e = (
|
||||
result.get("e2e_latency")
|
||||
if isinstance(result, dict) else None
|
||||
) or wall
|
||||
measured_secs.append(e2e)
|
||||
print(f"[measured {m + 1}/{measured_runs}] e2e={e2e:.2f}s wall={wall:.2f}s")
|
||||
if isinstance(result, dict):
|
||||
_print_stage_breakdown(result, f"measured {m + 1}")
|
||||
_collect_stage_times(result, stage_times, stage_order)
|
||||
|
||||
print("\n=== summary (bf16 / COMPILE) ===")
|
||||
print(f"warmup wall-times: {[round(x, 1) for x in warmup_secs]}")
|
||||
if measured_secs:
|
||||
avg = sum(measured_secs) / len(measured_secs)
|
||||
print(
|
||||
f"measured e2e (n={len(measured_secs)}): "
|
||||
f"{[round(x, 2) for x in measured_secs]} -> avg {avg:.2f}s"
|
||||
)
|
||||
if stage_times:
|
||||
print(f"average stage times over {measured_runs} measured runs:")
|
||||
avg_total = 0.0
|
||||
for name in stage_order:
|
||||
vals = stage_times.get(name) or []
|
||||
if not vals:
|
||||
continue
|
||||
avg_v = sum(vals) / len(vals)
|
||||
avg_total += avg_v
|
||||
print(f" - {name}: {avg_v:.3f}s")
|
||||
print(f" - stage_sum_avg: {avg_total:.3f}s")
|
||||
finally:
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,244 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""LTX-2.3 distilled image-to-video — NVFP4 DiT, torch.compile ON + timing.
|
||||
|
||||
Compiled + NVFP4 sibling of ``basic_ltx2_3_distilled_i2v_uncompiled.py``. Same
|
||||
generation recipe (8 denoise + 3 refine steps, CFG=1, no refine LoRA — the
|
||||
distilled production recipe), same input/measurement methodology, but with
|
||||
torch.compile fully ENABLED (DiT + text encoder + VAE) AND the DiT quantized
|
||||
to NVFP4. Fills the NVFP4/compile cell of the 2x2 (bf16 vs NVFP4) x
|
||||
(eager vs compile) table. Requires flashinfer + CUDA_HOME=/usr/local/cuda-13.2.
|
||||
|
||||
NOTE: first warmup pays a long cold-compile (tens of minutes on Blackwell);
|
||||
results are cached in ``$TORCHINDUCTOR_CACHE_DIR`` for later runs.
|
||||
|
||||
CUDA_VISIBLE_DEVICES=1 python \
|
||||
examples/inference/basic/basic_ltx2_3_distilled_i2v_compiled.py
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import time
|
||||
from collections import OrderedDict
|
||||
from pathlib import Path
|
||||
|
||||
import torch._inductor.config as _inductor
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
from fastvideo.layers.quantization.nvfp4_config import NVFP4Config
|
||||
from fastvideo.utils import maybe_download_model
|
||||
|
||||
os.environ.setdefault("FASTVIDEO_ATTENTION_BACKEND", "FLASH_ATTN")
|
||||
os.environ.setdefault("FASTVIDEO_STAGE_LOGGING", "1")
|
||||
|
||||
# Inductor knobs. ``shape_padding=False`` is mandatory on Blackwell to avoid a
|
||||
# cuBLAS INVALID_VALUE crash inside pad_mm during the refine path. The rest are
|
||||
# autotune-friendliness flags (match the canonical compiled example).
|
||||
_inductor.shape_padding = False
|
||||
_inductor.conv_1x1_as_mm = True
|
||||
_inductor.coordinate_descent_tuning = True
|
||||
_inductor.coordinate_descent_check_all_directions = True
|
||||
_inductor.epilogue_fusion = False
|
||||
|
||||
MODEL_ID = os.path.expandvars(
|
||||
os.path.expanduser(
|
||||
os.getenv("LTX23_MODEL_PATH", "FastVideo/LTX-2.3-Distilled-Diffusers")
|
||||
)
|
||||
)
|
||||
OUTPUT_DIR = Path(
|
||||
os.getenv(
|
||||
"LTX23_OUTPUT_DIR",
|
||||
"outputs_video/ltx2_3_distilled_i2v_compiled_nvfp4",
|
||||
)
|
||||
)
|
||||
I2V_IMAGE = os.getenv("LTX23_I2V_IMAGE", "")
|
||||
DEFAULT_PROMPT = (
|
||||
"A fashion model takes a slow step forward and shifts her weight, "
|
||||
"the soft fabric of her clothing swaying and rippling with the "
|
||||
"motion, her hair shifting gently, soft even studio lighting on a "
|
||||
"clean light background, elegant slow-motion runway feel."
|
||||
)
|
||||
PROMPT = os.getenv("LTX23_I2V_PROMPT", DEFAULT_PROMPT)
|
||||
|
||||
|
||||
def _print_stage_breakdown(result: dict, label: str) -> float | None:
|
||||
logging_info = result.get("logging_info")
|
||||
stages = getattr(logging_info, "stages", None) if logging_info else None
|
||||
if not stages:
|
||||
print(f" [{label}] stage breakdown unavailable")
|
||||
return None
|
||||
print(f" [{label}] stage breakdown:")
|
||||
total = 0.0
|
||||
for name, metrics in stages.items():
|
||||
exec_s = float(metrics.get("execution_time", 0.0))
|
||||
total += exec_s
|
||||
print(f" - {name}: {exec_s:.3f}s")
|
||||
print(f" - stage_sum: {total:.3f}s")
|
||||
return total
|
||||
|
||||
|
||||
def _collect_stage_times(
|
||||
result: dict,
|
||||
stage_times: dict[str, list[float]],
|
||||
stage_order: OrderedDict[str, None],
|
||||
) -> None:
|
||||
logging_info = result.get("logging_info")
|
||||
stages = getattr(logging_info, "stages", None) if logging_info else None
|
||||
if not stages:
|
||||
return
|
||||
for name, metrics in stages.items():
|
||||
stage_order.setdefault(name, None)
|
||||
stage_times.setdefault(name, []).append(
|
||||
float(metrics.get("execution_time", 0.0))
|
||||
)
|
||||
|
||||
|
||||
def _resolve_refine_upsampler(model_root: str) -> Path:
|
||||
"""LTX-2.3 distilled snapshots ship a `spatial_upscaler/` subdir."""
|
||||
for name in ("spatial_upscaler", "spatial_upsampler"):
|
||||
cand = Path(model_root) / name
|
||||
if (cand / "config.json").is_file():
|
||||
return cand
|
||||
raise FileNotFoundError(
|
||||
f"No refine upsampler directory under {model_root}. "
|
||||
f"Expected `{model_root}/spatial_upscaler/config.json`."
|
||||
)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
if not I2V_IMAGE:
|
||||
raise SystemExit(
|
||||
"LTX23_I2V_IMAGE is required for i2v. Example:\n"
|
||||
" export LTX23_I2V_IMAGE=/path/to/portrait_or_product.jpg\n"
|
||||
" python examples/inference/basic/"
|
||||
"basic_ltx2_3_distilled_i2v_compiled.py"
|
||||
)
|
||||
if not Path(I2V_IMAGE).is_file():
|
||||
raise SystemExit(f"LTX23_I2V_IMAGE not found: {I2V_IMAGE}")
|
||||
|
||||
OUTPUT_DIR.mkdir(parents=True, exist_ok=True)
|
||||
model_root = maybe_download_model(MODEL_ID)
|
||||
refine_upsampler_path = _resolve_refine_upsampler(model_root)
|
||||
print(f"Model: {model_root}")
|
||||
print(f"Refine upsampler: {refine_upsampler_path}")
|
||||
print(f"i2v image: {I2V_IMAGE}")
|
||||
print(f"Output dir: {OUTPUT_DIR.resolve()}")
|
||||
print("torch.compile: ENABLED (DiT + text encoder + VAE)")
|
||||
print("DiT quant: NVFP4")
|
||||
|
||||
pipeline_config = PipelineConfig.from_pretrained(model_root)
|
||||
pipeline_config.dit_config.quant_config = NVFP4Config()
|
||||
|
||||
torch_compile_kwargs = {
|
||||
"backend": "inductor",
|
||||
"fullgraph": True,
|
||||
"mode": "default",
|
||||
"dynamic": False,
|
||||
}
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_root,
|
||||
num_gpus=1,
|
||||
ltx2_refine_enabled=True,
|
||||
ltx2_refine_upsampler_path=str(refine_upsampler_path),
|
||||
ltx2_refine_lora_path="",
|
||||
ltx2_refine_num_inference_steps=3,
|
||||
ltx2_refine_guidance_scale=1.0,
|
||||
ltx2_refine_add_noise=True,
|
||||
pipeline_config=pipeline_config,
|
||||
enable_torch_compile=True,
|
||||
enable_torch_compile_text_encoder=True,
|
||||
enable_torch_compile_vae=True,
|
||||
torch_compile_kwargs=torch_compile_kwargs,
|
||||
torch_compile_kwargs_vae=torch_compile_kwargs,
|
||||
dit_cpu_offload=False,
|
||||
text_encoder_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
ltx2_vae_tiling=False,
|
||||
)
|
||||
|
||||
common_kwargs = dict(
|
||||
prompt=PROMPT,
|
||||
negative_prompt="",
|
||||
guidance_scale=1.0,
|
||||
height=1280, width=832,
|
||||
num_frames=121, fps=24,
|
||||
num_inference_steps=8,
|
||||
ltx2_images=[(I2V_IMAGE, 0, 1.0)],
|
||||
ltx2_image_crf=0.0,
|
||||
save_video=True,
|
||||
)
|
||||
|
||||
# Compile → 2 warmups: first pays cold compile + first-shape guards, the
|
||||
# second lets any residual recompiles settle before we measure.
|
||||
warmup_runs = 2
|
||||
measured_runs = 3
|
||||
warmup_secs: list[float] = []
|
||||
measured_secs: list[float] = []
|
||||
stage_times: dict[str, list[float]] = {}
|
||||
stage_order: OrderedDict[str, None] = OrderedDict()
|
||||
|
||||
try:
|
||||
for w in range(warmup_runs):
|
||||
t0 = time.perf_counter()
|
||||
print(f"\n[warmup {w + 1}/{warmup_runs}] compiling + generating…")
|
||||
generator.generate_video(
|
||||
output_path=str(OUTPUT_DIR / f"_warmup_{w + 1}.mp4"),
|
||||
seed=7,
|
||||
**common_kwargs,
|
||||
)
|
||||
dt = time.perf_counter() - t0
|
||||
warmup_secs.append(dt)
|
||||
print(f"[warmup {w + 1}/{warmup_runs}] wall={dt:.1f}s")
|
||||
|
||||
for w in range(warmup_runs):
|
||||
(OUTPUT_DIR / f"_warmup_{w + 1}.mp4").unlink(missing_ok=True)
|
||||
|
||||
for m in range(measured_runs):
|
||||
out_path = (
|
||||
OUTPUT_DIR
|
||||
/ f"output_ltx2_3_distilled_i2v_compiled_nvfp4_run_{m + 1}.mp4"
|
||||
)
|
||||
print(f"\n[measured {m + 1}/{measured_runs}] generating: {out_path}")
|
||||
t0 = time.perf_counter()
|
||||
result = generator.generate_video(
|
||||
output_path=str(out_path),
|
||||
seed=2002 + m,
|
||||
**common_kwargs,
|
||||
)
|
||||
wall = time.perf_counter() - t0
|
||||
e2e = (
|
||||
result.get("e2e_latency")
|
||||
if isinstance(result, dict) else None
|
||||
) or wall
|
||||
measured_secs.append(e2e)
|
||||
print(f"[measured {m + 1}/{measured_runs}] e2e={e2e:.2f}s wall={wall:.2f}s")
|
||||
if isinstance(result, dict):
|
||||
_print_stage_breakdown(result, f"measured {m + 1}")
|
||||
_collect_stage_times(result, stage_times, stage_order)
|
||||
|
||||
print("\n=== summary (NVFP4 / COMPILE) ===")
|
||||
print(f"warmup wall-times: {[round(x, 1) for x in warmup_secs]}")
|
||||
if measured_secs:
|
||||
avg = sum(measured_secs) / len(measured_secs)
|
||||
print(
|
||||
f"measured e2e (n={len(measured_secs)}): "
|
||||
f"{[round(x, 2) for x in measured_secs]} -> avg {avg:.2f}s"
|
||||
)
|
||||
if stage_times:
|
||||
print(f"average stage times over {measured_runs} measured runs:")
|
||||
avg_total = 0.0
|
||||
for name in stage_order:
|
||||
vals = stage_times.get(name) or []
|
||||
if not vals:
|
||||
continue
|
||||
avg_v = sum(vals) / len(vals)
|
||||
avg_total += avg_v
|
||||
print(f" - {name}: {avg_v:.3f}s")
|
||||
print(f" - stage_sum_avg: {avg_total:.3f}s")
|
||||
finally:
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,237 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""LTX-2.3 distilled image-to-video — EAGER (no torch.compile) + timing.
|
||||
|
||||
Uncompiled sibling of ``basic_ltx2_3_distilled_i2v.py``. Identical generation
|
||||
recipe (8 denoise + 3 refine steps, CFG=1, no refine LoRA — the distilled
|
||||
production recipe) but with torch.compile fully DISABLED, so there is no
|
||||
cold-compile cost and the per-stage timings reflect plain eager execution.
|
||||
Use this to get a quick speed baseline before paying the compile tax.
|
||||
|
||||
Quick start
|
||||
-----------
|
||||
export LTX23_I2V_IMAGE=/path/to/your/portrait_or_product.jpg
|
||||
# optional overrides:
|
||||
# export LTX23_I2V_PROMPT="a fashion model walks toward camera..."
|
||||
# export LTX23_OUTPUT_DIR=outputs_video/ltx2_3_distilled_i2v_uncompiled
|
||||
# export LTX23_MODEL_PATH=/local/path/to/LTX-2.3-Distilled-Diffusers
|
||||
CUDA_VISIBLE_DEVICES=1 python \
|
||||
examples/inference/basic/basic_ltx2_3_distilled_i2v_uncompiled.py
|
||||
|
||||
Hardware notes
|
||||
--------------
|
||||
- Single-GPU example. On this box pin the GB300 with
|
||||
``CUDA_VISIBLE_DEVICES=1`` (GPU 0 is the RTX PRO 6000 workstation card).
|
||||
- No compile → no Inductor cold-start, no ``shape_padding`` landmine, so the
|
||||
Blackwell-specific Inductor knobs from the compiled example are omitted.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import time
|
||||
from collections import OrderedDict
|
||||
from pathlib import Path
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
from fastvideo.utils import maybe_download_model
|
||||
|
||||
os.environ.setdefault("FASTVIDEO_ATTENTION_BACKEND", "FLASH_ATTN")
|
||||
os.environ.setdefault("FASTVIDEO_STAGE_LOGGING", "1")
|
||||
|
||||
MODEL_ID = os.path.expandvars(
|
||||
os.path.expanduser(
|
||||
os.getenv("LTX23_MODEL_PATH", "FastVideo/LTX-2.3-Distilled-Diffusers")
|
||||
)
|
||||
)
|
||||
OUTPUT_DIR = Path(
|
||||
os.getenv(
|
||||
"LTX23_OUTPUT_DIR", "outputs_video/ltx2_3_distilled_i2v_uncompiled"
|
||||
)
|
||||
)
|
||||
I2V_IMAGE = os.getenv("LTX23_I2V_IMAGE", "")
|
||||
DEFAULT_PROMPT = (
|
||||
"A fashion model takes a slow step forward and shifts her weight, "
|
||||
"the soft fabric of her clothing swaying and rippling with the "
|
||||
"motion, her hair shifting gently, soft even studio lighting on a "
|
||||
"clean light background, elegant slow-motion runway feel."
|
||||
)
|
||||
PROMPT = os.getenv("LTX23_I2V_PROMPT", DEFAULT_PROMPT)
|
||||
|
||||
|
||||
def _print_stage_breakdown(result: dict, label: str) -> float | None:
|
||||
logging_info = result.get("logging_info")
|
||||
stages = getattr(logging_info, "stages", None) if logging_info else None
|
||||
if not stages:
|
||||
print(f" [{label}] stage breakdown unavailable")
|
||||
return None
|
||||
print(f" [{label}] stage breakdown:")
|
||||
total = 0.0
|
||||
for name, metrics in stages.items():
|
||||
exec_s = float(metrics.get("execution_time", 0.0))
|
||||
total += exec_s
|
||||
print(f" - {name}: {exec_s:.3f}s")
|
||||
print(f" - stage_sum: {total:.3f}s")
|
||||
return total
|
||||
|
||||
|
||||
def _collect_stage_times(
|
||||
result: dict,
|
||||
stage_times: dict[str, list[float]],
|
||||
stage_order: OrderedDict[str, None],
|
||||
) -> None:
|
||||
logging_info = result.get("logging_info")
|
||||
stages = getattr(logging_info, "stages", None) if logging_info else None
|
||||
if not stages:
|
||||
return
|
||||
for name, metrics in stages.items():
|
||||
stage_order.setdefault(name, None)
|
||||
stage_times.setdefault(name, []).append(
|
||||
float(metrics.get("execution_time", 0.0))
|
||||
)
|
||||
|
||||
|
||||
def _resolve_refine_upsampler(model_root: str) -> Path:
|
||||
"""LTX-2.3 distilled snapshots ship a `spatial_upscaler/` subdir."""
|
||||
for name in ("spatial_upscaler", "spatial_upsampler"):
|
||||
cand = Path(model_root) / name
|
||||
if (cand / "config.json").is_file():
|
||||
return cand
|
||||
raise FileNotFoundError(
|
||||
f"No refine upsampler directory under {model_root}. "
|
||||
f"Expected `{model_root}/spatial_upscaler/config.json`."
|
||||
)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
if not I2V_IMAGE:
|
||||
raise SystemExit(
|
||||
"LTX23_I2V_IMAGE is required for i2v. Example:\n"
|
||||
" export LTX23_I2V_IMAGE=/path/to/portrait_or_product.jpg\n"
|
||||
" python examples/inference/basic/"
|
||||
"basic_ltx2_3_distilled_i2v_uncompiled.py"
|
||||
)
|
||||
if not Path(I2V_IMAGE).is_file():
|
||||
raise SystemExit(f"LTX23_I2V_IMAGE not found: {I2V_IMAGE}")
|
||||
|
||||
OUTPUT_DIR.mkdir(parents=True, exist_ok=True)
|
||||
model_root = maybe_download_model(MODEL_ID)
|
||||
refine_upsampler_path = _resolve_refine_upsampler(model_root)
|
||||
print(f"Model: {model_root}")
|
||||
print(f"Refine upsampler: {refine_upsampler_path}")
|
||||
print(f"i2v image: {I2V_IMAGE}")
|
||||
print(f"Output dir: {OUTPUT_DIR.resolve()}")
|
||||
print("torch.compile: DISABLED (eager)")
|
||||
|
||||
# Loading the pipeline config *with model_path* binds model-specific
|
||||
# tuning (notably VAE precision/decoder defaults) into the config.
|
||||
pipeline_config = PipelineConfig.from_pretrained(model_root)
|
||||
pipeline_config.dit_config.quant_config = None
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_root,
|
||||
num_gpus=1,
|
||||
# LTX-2.3 distilled uses the two-stage refine pipeline; the refine
|
||||
# LoRA is intentionally empty for the distilled student.
|
||||
ltx2_refine_enabled=True,
|
||||
ltx2_refine_upsampler_path=str(refine_upsampler_path),
|
||||
ltx2_refine_lora_path="",
|
||||
ltx2_refine_num_inference_steps=3,
|
||||
ltx2_refine_guidance_scale=1.0,
|
||||
ltx2_refine_add_noise=True,
|
||||
pipeline_config=pipeline_config,
|
||||
# --- eager: every compile switch off ---
|
||||
enable_torch_compile=False,
|
||||
enable_torch_compile_text_encoder=False,
|
||||
enable_torch_compile_vae=False,
|
||||
# Keep everything resident — no CPU offload for serving-style runs.
|
||||
dit_cpu_offload=False,
|
||||
text_encoder_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
ltx2_vae_tiling=False,
|
||||
)
|
||||
|
||||
common_kwargs = dict(
|
||||
prompt=PROMPT,
|
||||
negative_prompt="", # distilled is CFG-free; no negative needed
|
||||
guidance_scale=1.0, # CFG=1 for distilled
|
||||
height=1280, width=832, # portrait runway aspect
|
||||
num_frames=121, fps=24, # ~5s clip
|
||||
num_inference_steps=8, # distilled denoise steps
|
||||
ltx2_images=[(I2V_IMAGE, 0, 1.0)],
|
||||
ltx2_image_crf=0.0,
|
||||
save_video=True,
|
||||
)
|
||||
|
||||
# No compile → a single warmup is enough to settle allocator / cuDNN /
|
||||
# autotune before we measure.
|
||||
warmup_runs = 1
|
||||
measured_runs = 3
|
||||
warmup_secs: list[float] = []
|
||||
measured_secs: list[float] = []
|
||||
stage_times: dict[str, list[float]] = {}
|
||||
stage_order: OrderedDict[str, None] = OrderedDict()
|
||||
|
||||
try:
|
||||
for w in range(warmup_runs):
|
||||
t0 = time.perf_counter()
|
||||
print(f"\n[warmup {w + 1}/{warmup_runs}] generating…")
|
||||
generator.generate_video(
|
||||
output_path=str(OUTPUT_DIR / f"_warmup_{w + 1}.mp4"),
|
||||
seed=7,
|
||||
**common_kwargs,
|
||||
)
|
||||
dt = time.perf_counter() - t0
|
||||
warmup_secs.append(dt)
|
||||
print(f"[warmup {w + 1}/{warmup_runs}] wall={dt:.1f}s")
|
||||
|
||||
for w in range(warmup_runs):
|
||||
(OUTPUT_DIR / f"_warmup_{w + 1}.mp4").unlink(missing_ok=True)
|
||||
|
||||
for m in range(measured_runs):
|
||||
out_path = (
|
||||
OUTPUT_DIR
|
||||
/ f"output_ltx2_3_distilled_i2v_uncompiled_run_{m + 1}.mp4"
|
||||
)
|
||||
print(f"\n[measured {m + 1}/{measured_runs}] generating: {out_path}")
|
||||
t0 = time.perf_counter()
|
||||
result = generator.generate_video(
|
||||
output_path=str(out_path),
|
||||
seed=2002 + m,
|
||||
**common_kwargs,
|
||||
)
|
||||
wall = time.perf_counter() - t0
|
||||
e2e = (
|
||||
result.get("e2e_latency")
|
||||
if isinstance(result, dict) else None
|
||||
) or wall
|
||||
measured_secs.append(e2e)
|
||||
print(f"[measured {m + 1}/{measured_runs}] e2e={e2e:.2f}s wall={wall:.2f}s")
|
||||
if isinstance(result, dict):
|
||||
_print_stage_breakdown(result, f"measured {m + 1}")
|
||||
_collect_stage_times(result, stage_times, stage_order)
|
||||
|
||||
print("\n=== summary (EAGER / no compile) ===")
|
||||
print(f"warmup wall-times: {[round(x, 1) for x in warmup_secs]}")
|
||||
if measured_secs:
|
||||
avg = sum(measured_secs) / len(measured_secs)
|
||||
print(
|
||||
f"measured e2e (n={len(measured_secs)}): "
|
||||
f"{[round(x, 2) for x in measured_secs]} -> avg {avg:.2f}s"
|
||||
)
|
||||
if stage_times:
|
||||
print(f"average stage times over {measured_runs} measured runs:")
|
||||
avg_total = 0.0
|
||||
for name in stage_order:
|
||||
vals = stage_times.get(name) or []
|
||||
if not vals:
|
||||
continue
|
||||
avg_v = sum(vals) / len(vals)
|
||||
avg_total += avg_v
|
||||
print(f" - {name}: {avg_v:.3f}s")
|
||||
print(f" - stage_sum_avg: {avg_total:.3f}s")
|
||||
finally:
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,241 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""LTX-2.3 distilled image-to-video — NVFP4 DiT, EAGER (no compile) + timing.
|
||||
|
||||
NVFP4 sibling of ``basic_ltx2_3_distilled_i2v_uncompiled.py``. Identical
|
||||
generation recipe (8 denoise + 3 refine steps, CFG=1, no refine LoRA — the
|
||||
distilled production recipe) and torch.compile still fully DISABLED, but the
|
||||
DiT runs with an ``NVFP4Config`` quant config (FP4 linear layers) instead of
|
||||
bf16. Use this to measure NVFP4 eager speed vs. the bf16 eager baseline.
|
||||
|
||||
Quick start
|
||||
-----------
|
||||
export LTX23_I2V_IMAGE=/path/to/your/portrait_or_product.jpg
|
||||
# optional overrides:
|
||||
# export LTX23_I2V_PROMPT="a fashion model walks toward camera..."
|
||||
# export LTX23_OUTPUT_DIR=outputs_video/ltx2_3_distilled_i2v_uncompiled
|
||||
# export LTX23_MODEL_PATH=/local/path/to/LTX-2.3-Distilled-Diffusers
|
||||
CUDA_VISIBLE_DEVICES=1 python \
|
||||
examples/inference/basic/basic_ltx2_3_distilled_i2v_uncompiled.py
|
||||
|
||||
Hardware notes
|
||||
--------------
|
||||
- Single-GPU example. On this box pin the GB300 with
|
||||
``CUDA_VISIBLE_DEVICES=1`` (GPU 0 is the RTX PRO 6000 workstation card).
|
||||
- No compile → no Inductor cold-start, no ``shape_padding`` landmine, so the
|
||||
Blackwell-specific Inductor knobs from the compiled example are omitted.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import time
|
||||
from collections import OrderedDict
|
||||
from pathlib import Path
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
from fastvideo.layers.quantization.nvfp4_config import NVFP4Config
|
||||
from fastvideo.utils import maybe_download_model
|
||||
|
||||
os.environ.setdefault("FASTVIDEO_ATTENTION_BACKEND", "FLASH_ATTN")
|
||||
os.environ.setdefault("FASTVIDEO_STAGE_LOGGING", "1")
|
||||
|
||||
MODEL_ID = os.path.expandvars(
|
||||
os.path.expanduser(
|
||||
os.getenv("LTX23_MODEL_PATH", "FastVideo/LTX-2.3-Distilled-Diffusers")
|
||||
)
|
||||
)
|
||||
OUTPUT_DIR = Path(
|
||||
os.getenv(
|
||||
"LTX23_OUTPUT_DIR",
|
||||
"outputs_video/ltx2_3_distilled_i2v_uncompiled_nvfp4",
|
||||
)
|
||||
)
|
||||
I2V_IMAGE = os.getenv("LTX23_I2V_IMAGE", "")
|
||||
DEFAULT_PROMPT = (
|
||||
"A fashion model takes a slow step forward and shifts her weight, "
|
||||
"the soft fabric of her clothing swaying and rippling with the "
|
||||
"motion, her hair shifting gently, soft even studio lighting on a "
|
||||
"clean light background, elegant slow-motion runway feel."
|
||||
)
|
||||
PROMPT = os.getenv("LTX23_I2V_PROMPT", DEFAULT_PROMPT)
|
||||
|
||||
|
||||
def _print_stage_breakdown(result: dict, label: str) -> float | None:
|
||||
logging_info = result.get("logging_info")
|
||||
stages = getattr(logging_info, "stages", None) if logging_info else None
|
||||
if not stages:
|
||||
print(f" [{label}] stage breakdown unavailable")
|
||||
return None
|
||||
print(f" [{label}] stage breakdown:")
|
||||
total = 0.0
|
||||
for name, metrics in stages.items():
|
||||
exec_s = float(metrics.get("execution_time", 0.0))
|
||||
total += exec_s
|
||||
print(f" - {name}: {exec_s:.3f}s")
|
||||
print(f" - stage_sum: {total:.3f}s")
|
||||
return total
|
||||
|
||||
|
||||
def _collect_stage_times(
|
||||
result: dict,
|
||||
stage_times: dict[str, list[float]],
|
||||
stage_order: OrderedDict[str, None],
|
||||
) -> None:
|
||||
logging_info = result.get("logging_info")
|
||||
stages = getattr(logging_info, "stages", None) if logging_info else None
|
||||
if not stages:
|
||||
return
|
||||
for name, metrics in stages.items():
|
||||
stage_order.setdefault(name, None)
|
||||
stage_times.setdefault(name, []).append(
|
||||
float(metrics.get("execution_time", 0.0))
|
||||
)
|
||||
|
||||
|
||||
def _resolve_refine_upsampler(model_root: str) -> Path:
|
||||
"""LTX-2.3 distilled snapshots ship a `spatial_upscaler/` subdir."""
|
||||
for name in ("spatial_upscaler", "spatial_upsampler"):
|
||||
cand = Path(model_root) / name
|
||||
if (cand / "config.json").is_file():
|
||||
return cand
|
||||
raise FileNotFoundError(
|
||||
f"No refine upsampler directory under {model_root}. "
|
||||
f"Expected `{model_root}/spatial_upscaler/config.json`."
|
||||
)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
if not I2V_IMAGE:
|
||||
raise SystemExit(
|
||||
"LTX23_I2V_IMAGE is required for i2v. Example:\n"
|
||||
" export LTX23_I2V_IMAGE=/path/to/portrait_or_product.jpg\n"
|
||||
" python examples/inference/basic/"
|
||||
"basic_ltx2_3_distilled_i2v_uncompiled.py"
|
||||
)
|
||||
if not Path(I2V_IMAGE).is_file():
|
||||
raise SystemExit(f"LTX23_I2V_IMAGE not found: {I2V_IMAGE}")
|
||||
|
||||
OUTPUT_DIR.mkdir(parents=True, exist_ok=True)
|
||||
model_root = maybe_download_model(MODEL_ID)
|
||||
refine_upsampler_path = _resolve_refine_upsampler(model_root)
|
||||
print(f"Model: {model_root}")
|
||||
print(f"Refine upsampler: {refine_upsampler_path}")
|
||||
print(f"i2v image: {I2V_IMAGE}")
|
||||
print(f"Output dir: {OUTPUT_DIR.resolve()}")
|
||||
print("torch.compile: DISABLED (eager)")
|
||||
print("DiT quant: NVFP4")
|
||||
|
||||
# Loading the pipeline config *with model_path* binds model-specific
|
||||
# tuning (notably VAE precision/decoder defaults) into the config.
|
||||
pipeline_config = PipelineConfig.from_pretrained(model_root)
|
||||
# NVFP4: quantize the DiT linear layers to FP4 (vs. None=bf16 baseline).
|
||||
pipeline_config.dit_config.quant_config = NVFP4Config()
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_root,
|
||||
num_gpus=1,
|
||||
# LTX-2.3 distilled uses the two-stage refine pipeline; the refine
|
||||
# LoRA is intentionally empty for the distilled student.
|
||||
ltx2_refine_enabled=True,
|
||||
ltx2_refine_upsampler_path=str(refine_upsampler_path),
|
||||
ltx2_refine_lora_path="",
|
||||
ltx2_refine_num_inference_steps=3,
|
||||
ltx2_refine_guidance_scale=1.0,
|
||||
ltx2_refine_add_noise=True,
|
||||
pipeline_config=pipeline_config,
|
||||
# --- eager: every compile switch off ---
|
||||
enable_torch_compile=False,
|
||||
enable_torch_compile_text_encoder=False,
|
||||
enable_torch_compile_vae=False,
|
||||
# Keep everything resident — no CPU offload for serving-style runs.
|
||||
dit_cpu_offload=False,
|
||||
text_encoder_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
ltx2_vae_tiling=False,
|
||||
)
|
||||
|
||||
common_kwargs = dict(
|
||||
prompt=PROMPT,
|
||||
negative_prompt="", # distilled is CFG-free; no negative needed
|
||||
guidance_scale=1.0, # CFG=1 for distilled
|
||||
height=1280, width=832, # portrait runway aspect
|
||||
num_frames=121, fps=24, # ~5s clip
|
||||
num_inference_steps=8, # distilled denoise steps
|
||||
ltx2_images=[(I2V_IMAGE, 0, 1.0)],
|
||||
ltx2_image_crf=0.0,
|
||||
save_video=True,
|
||||
)
|
||||
|
||||
# No compile → a single warmup is enough to settle allocator / cuDNN /
|
||||
# autotune before we measure.
|
||||
warmup_runs = 1
|
||||
measured_runs = 3
|
||||
warmup_secs: list[float] = []
|
||||
measured_secs: list[float] = []
|
||||
stage_times: dict[str, list[float]] = {}
|
||||
stage_order: OrderedDict[str, None] = OrderedDict()
|
||||
|
||||
try:
|
||||
for w in range(warmup_runs):
|
||||
t0 = time.perf_counter()
|
||||
print(f"\n[warmup {w + 1}/{warmup_runs}] generating…")
|
||||
generator.generate_video(
|
||||
output_path=str(OUTPUT_DIR / f"_warmup_{w + 1}.mp4"),
|
||||
seed=7,
|
||||
**common_kwargs,
|
||||
)
|
||||
dt = time.perf_counter() - t0
|
||||
warmup_secs.append(dt)
|
||||
print(f"[warmup {w + 1}/{warmup_runs}] wall={dt:.1f}s")
|
||||
|
||||
for w in range(warmup_runs):
|
||||
(OUTPUT_DIR / f"_warmup_{w + 1}.mp4").unlink(missing_ok=True)
|
||||
|
||||
for m in range(measured_runs):
|
||||
out_path = (
|
||||
OUTPUT_DIR
|
||||
/ f"output_ltx2_3_distilled_i2v_uncompiled_run_{m + 1}.mp4"
|
||||
)
|
||||
print(f"\n[measured {m + 1}/{measured_runs}] generating: {out_path}")
|
||||
t0 = time.perf_counter()
|
||||
result = generator.generate_video(
|
||||
output_path=str(out_path),
|
||||
seed=2002 + m,
|
||||
**common_kwargs,
|
||||
)
|
||||
wall = time.perf_counter() - t0
|
||||
e2e = (
|
||||
result.get("e2e_latency")
|
||||
if isinstance(result, dict) else None
|
||||
) or wall
|
||||
measured_secs.append(e2e)
|
||||
print(f"[measured {m + 1}/{measured_runs}] e2e={e2e:.2f}s wall={wall:.2f}s")
|
||||
if isinstance(result, dict):
|
||||
_print_stage_breakdown(result, f"measured {m + 1}")
|
||||
_collect_stage_times(result, stage_times, stage_order)
|
||||
|
||||
print("\n=== summary (NVFP4 / EAGER / no compile) ===")
|
||||
print(f"warmup wall-times: {[round(x, 1) for x in warmup_secs]}")
|
||||
if measured_secs:
|
||||
avg = sum(measured_secs) / len(measured_secs)
|
||||
print(
|
||||
f"measured e2e (n={len(measured_secs)}): "
|
||||
f"{[round(x, 2) for x in measured_secs]} -> avg {avg:.2f}s"
|
||||
)
|
||||
if stage_times:
|
||||
print(f"average stage times over {measured_runs} measured runs:")
|
||||
avg_total = 0.0
|
||||
for name in stage_order:
|
||||
vals = stage_times.get(name) or []
|
||||
if not vals:
|
||||
continue
|
||||
avg_v = sum(vals) / len(vals)
|
||||
avg_total += avg_v
|
||||
print(f" - {name}: {avg_v:.3f}s")
|
||||
print(f" - stage_sum_avg: {avg_total:.3f}s")
|
||||
finally:
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,280 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""LTX-2.3 distilled TEXT-to-video — parametrized 2x2 speed benchmark.
|
||||
|
||||
Pure t2v (no image conditioning). Same recipe/resolution/measurement as the
|
||||
i2v scripts so the two are directly comparable: 8 denoise + 3 refine, CFG=1,
|
||||
no refine LoRA, 832x1280 portrait, 121 frames @ 24fps.
|
||||
|
||||
Toggle the 2x2 cell via env vars:
|
||||
LTX23_COMPILE = 1|0 -> torch.compile on/off (DiT + text encoder + VAE)
|
||||
LTX23_NVFP4 = 1|0 -> DiT NVFP4 quant on/off (needs flashinfer +
|
||||
CUDA_HOME=/usr/local/cuda-13.2)
|
||||
|
||||
CUDA_VISIBLE_DEVICES=1 LTX23_COMPILE=1 LTX23_NVFP4=0 \
|
||||
python examples/inference/basic/bench_ltx2_3_distilled_t2v.py
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import time
|
||||
from collections import OrderedDict
|
||||
from pathlib import Path
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
from fastvideo.utils import maybe_download_model
|
||||
|
||||
os.environ.setdefault("FASTVIDEO_ATTENTION_BACKEND", "FLASH_ATTN")
|
||||
os.environ.setdefault("FASTVIDEO_STAGE_LOGGING", "1")
|
||||
|
||||
|
||||
def _env_flag(name: str, default: str = "0") -> bool:
|
||||
return os.getenv(name, default).strip().lower() in ("1", "true", "yes", "on")
|
||||
|
||||
|
||||
COMPILE = _env_flag("LTX23_COMPILE", "0")
|
||||
NVFP4 = _env_flag("LTX23_NVFP4", "0")
|
||||
# FA4-FP4 attention: quantize Q/K to NVFP4 for Blackwell block-scaled MMA
|
||||
# (needs the cutlass.utils.ampere_helpers shim; auto-selects FlashAttention-4).
|
||||
FA4 = _env_flag("LTX23_FA4", "0")
|
||||
|
||||
if COMPILE:
|
||||
# Inductor knobs — shape_padding=False mandatory on Blackwell (pad_mm
|
||||
# cuBLAS landmine); rest are autotune-friendliness flags.
|
||||
import torch._inductor.config as _inductor
|
||||
_inductor.shape_padding = False
|
||||
_inductor.conv_1x1_as_mm = True
|
||||
_inductor.coordinate_descent_tuning = True
|
||||
_inductor.coordinate_descent_check_all_directions = True
|
||||
_inductor.epilogue_fusion = False
|
||||
if _env_flag("LTX23_NO_CGTREES", "0"):
|
||||
# Fall back to the simpler (non-tree) cudagraph impl: cudagraph_trees
|
||||
# rejects FP4 flashinfer custom ops that allocate untracked tensors
|
||||
# inside the captured region.
|
||||
_inductor.triton.cudagraph_trees = False
|
||||
|
||||
if NVFP4:
|
||||
from fastvideo.layers.quantization.nvfp4_config import NVFP4Config
|
||||
|
||||
HEIGHT = int(os.getenv("LTX23_HEIGHT", "1280"))
|
||||
WIDTH = int(os.getenv("LTX23_WIDTH", "832"))
|
||||
NUM_FRAMES = int(os.getenv("LTX23_NUM_FRAMES", "121"))
|
||||
STEPS = int(os.getenv("LTX23_STEPS", "8"))
|
||||
|
||||
_CELL = f"{'nvfp4' if NVFP4 else 'bf16'}_{'compile' if COMPILE else 'eager'}"
|
||||
if FA4:
|
||||
_CELL += "_fa4"
|
||||
_CELL += f"_{WIDTH}x{HEIGHT}_{NUM_FRAMES}f_{STEPS}st"
|
||||
|
||||
MODEL_ID = os.path.expandvars(
|
||||
os.path.expanduser(
|
||||
os.getenv("LTX23_MODEL_PATH", "FastVideo/LTX-2.3-Distilled-Diffusers")
|
||||
)
|
||||
)
|
||||
OUTPUT_DIR = Path(
|
||||
os.getenv("LTX23_OUTPUT_DIR", f"outputs_video/ltx2_3_distilled_t2v_{_CELL}")
|
||||
)
|
||||
DEFAULT_PROMPT = (
|
||||
"A fashion model takes a slow step forward and shifts her weight, "
|
||||
"the soft fabric of her clothing swaying and rippling with the "
|
||||
"motion, her hair shifting gently, soft even studio lighting on a "
|
||||
"clean light background, elegant slow-motion runway feel."
|
||||
)
|
||||
PROMPT = os.getenv("LTX23_T2V_PROMPT", DEFAULT_PROMPT)
|
||||
|
||||
|
||||
def _print_stage_breakdown(result: dict, label: str) -> float | None:
|
||||
logging_info = result.get("logging_info")
|
||||
stages = getattr(logging_info, "stages", None) if logging_info else None
|
||||
if not stages:
|
||||
print(f" [{label}] stage breakdown unavailable")
|
||||
return None
|
||||
print(f" [{label}] stage breakdown:")
|
||||
total = 0.0
|
||||
for name, metrics in stages.items():
|
||||
exec_s = float(metrics.get("execution_time", 0.0))
|
||||
total += exec_s
|
||||
print(f" - {name}: {exec_s:.3f}s")
|
||||
print(f" - stage_sum: {total:.3f}s")
|
||||
return total
|
||||
|
||||
|
||||
def _collect_stage_times(
|
||||
result: dict,
|
||||
stage_times: dict[str, list[float]],
|
||||
stage_order: OrderedDict[str, None],
|
||||
) -> None:
|
||||
logging_info = result.get("logging_info")
|
||||
stages = getattr(logging_info, "stages", None) if logging_info else None
|
||||
if not stages:
|
||||
return
|
||||
for name, metrics in stages.items():
|
||||
stage_order.setdefault(name, None)
|
||||
stage_times.setdefault(name, []).append(
|
||||
float(metrics.get("execution_time", 0.0))
|
||||
)
|
||||
|
||||
|
||||
def _resolve_refine_upsampler(model_root: str) -> Path:
|
||||
for name in ("spatial_upscaler", "spatial_upsampler"):
|
||||
cand = Path(model_root) / name
|
||||
if (cand / "config.json").is_file():
|
||||
return cand
|
||||
raise FileNotFoundError(
|
||||
f"No refine upsampler directory under {model_root}. "
|
||||
f"Expected `{model_root}/spatial_upscaler/config.json`."
|
||||
)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
OUTPUT_DIR.mkdir(parents=True, exist_ok=True)
|
||||
model_root = maybe_download_model(MODEL_ID)
|
||||
refine_upsampler_path = _resolve_refine_upsampler(model_root)
|
||||
print(f"Cell: {_CELL} (compile={COMPILE}, nvfp4={NVFP4}, fa4={FA4})")
|
||||
print(f"Task: t2v (no image conditioning)")
|
||||
print(f"Model: {model_root}")
|
||||
print(f"Refine upsampler: {refine_upsampler_path}")
|
||||
print(f"Output dir: {OUTPUT_DIR.resolve()}")
|
||||
|
||||
pipeline_config = PipelineConfig.from_pretrained(model_root)
|
||||
pipeline_config.dit_config.quant_config = NVFP4Config() if NVFP4 else None
|
||||
|
||||
gen_kwargs = dict(
|
||||
num_gpus=1,
|
||||
ltx2_refine_enabled=True,
|
||||
ltx2_refine_upsampler_path=str(refine_upsampler_path),
|
||||
ltx2_refine_lora_path="",
|
||||
ltx2_refine_num_inference_steps=3,
|
||||
ltx2_refine_guidance_scale=1.0,
|
||||
ltx2_refine_add_noise=True,
|
||||
pipeline_config=pipeline_config,
|
||||
dit_cpu_offload=False,
|
||||
text_encoder_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
# Long / high-res runs (e.g. 968 frames) overflow VAE decode's 32-bit
|
||||
# index math; enable tiling via LTX23_VAE_TILING=1 to run them.
|
||||
ltx2_vae_tiling=_env_flag("LTX23_VAE_TILING", "0"),
|
||||
)
|
||||
if FA4:
|
||||
# Enables FlashAttention-4 + FP4 Q/K quant (sets FASTVIDEO_NVFP4_FA4=1).
|
||||
gen_kwargs["nvfp4_fa4"] = True
|
||||
# output_type="latent" bypasses the VAE decode in DecodingStage (generator-
|
||||
# level fastvideo arg, not a per-request sampling param).
|
||||
if not _env_flag("LTX23_DECODE", "1"):
|
||||
gen_kwargs["output_type"] = "latent"
|
||||
if COMPILE:
|
||||
dit_mode = os.getenv("LTX23_COMPILE_MODE", "default")
|
||||
torch_compile_kwargs = {
|
||||
"backend": "inductor",
|
||||
"fullgraph": True,
|
||||
"mode": dit_mode,
|
||||
"dynamic": False,
|
||||
}
|
||||
# VAE keeps "default" mode: cudagraphs (reduce-overhead) on the VAE/text
|
||||
# encoder triggers cross-module static-buffer aliasing errors in this
|
||||
# pipeline. Apply cudagraphs to the DiT only.
|
||||
vae_kwargs = {**torch_compile_kwargs, "mode": "default"}
|
||||
compile_te = _env_flag("LTX23_COMPILE_TE", "1")
|
||||
compile_vae = _env_flag("LTX23_COMPILE_VAE", "1")
|
||||
gen_kwargs.update(
|
||||
enable_torch_compile=True,
|
||||
enable_torch_compile_text_encoder=compile_te,
|
||||
enable_torch_compile_vae=compile_vae,
|
||||
torch_compile_kwargs=torch_compile_kwargs,
|
||||
torch_compile_kwargs_vae=vae_kwargs,
|
||||
)
|
||||
else:
|
||||
gen_kwargs.update(
|
||||
enable_torch_compile=False,
|
||||
enable_torch_compile_text_encoder=False,
|
||||
enable_torch_compile_vae=False,
|
||||
)
|
||||
|
||||
generator = VideoGenerator.from_pretrained(model_root, **gen_kwargs)
|
||||
|
||||
# Pure t2v: no ltx2_images / ltx2_image_crf.
|
||||
# LTX23_DECODE=0 -> output_type="latent": bypass the VAE decode (+ video
|
||||
# save) so denoise/refine DiT timing is measured cleanly. Needed for long /
|
||||
# high-res runs where tiled VAE decode dominates wall-time and would
|
||||
# otherwise pay ~1h per run while telling us nothing about DiT scaling.
|
||||
decode = _env_flag("LTX23_DECODE", "1")
|
||||
common_kwargs = dict(
|
||||
prompt=PROMPT,
|
||||
negative_prompt="",
|
||||
guidance_scale=1.0,
|
||||
height=HEIGHT, width=WIDTH,
|
||||
num_frames=NUM_FRAMES, fps=24,
|
||||
num_inference_steps=STEPS,
|
||||
save_video=decode,
|
||||
)
|
||||
|
||||
# Compile needs 2 warmups (cold compile + settle); eager needs 1.
|
||||
# Overridable for long runs where 5 full generations is too costly.
|
||||
warmup_runs = int(os.getenv("LTX23_WARMUP", "2" if COMPILE else "1"))
|
||||
measured_runs = int(os.getenv("LTX23_MEASURED", "3"))
|
||||
warmup_secs: list[float] = []
|
||||
measured_secs: list[float] = []
|
||||
stage_times: dict[str, list[float]] = {}
|
||||
stage_order: OrderedDict[str, None] = OrderedDict()
|
||||
|
||||
try:
|
||||
for w in range(warmup_runs):
|
||||
t0 = time.perf_counter()
|
||||
print(f"\n[warmup {w + 1}/{warmup_runs}] generating…")
|
||||
generator.generate_video(
|
||||
output_path=str(OUTPUT_DIR / f"_warmup_{w + 1}.mp4"),
|
||||
seed=7,
|
||||
**common_kwargs,
|
||||
)
|
||||
dt = time.perf_counter() - t0
|
||||
warmup_secs.append(dt)
|
||||
print(f"[warmup {w + 1}/{warmup_runs}] wall={dt:.1f}s")
|
||||
|
||||
for w in range(warmup_runs):
|
||||
(OUTPUT_DIR / f"_warmup_{w + 1}.mp4").unlink(missing_ok=True)
|
||||
|
||||
for m in range(measured_runs):
|
||||
out_path = OUTPUT_DIR / f"output_t2v_{_CELL}_run_{m + 1}.mp4"
|
||||
print(f"\n[measured {m + 1}/{measured_runs}] generating: {out_path}")
|
||||
t0 = time.perf_counter()
|
||||
result = generator.generate_video(
|
||||
output_path=str(out_path),
|
||||
seed=2002 + m,
|
||||
**common_kwargs,
|
||||
)
|
||||
wall = time.perf_counter() - t0
|
||||
e2e = (
|
||||
result.get("e2e_latency")
|
||||
if isinstance(result, dict) else None
|
||||
) or wall
|
||||
measured_secs.append(e2e)
|
||||
print(f"[measured {m + 1}/{measured_runs}] e2e={e2e:.2f}s wall={wall:.2f}s")
|
||||
if isinstance(result, dict):
|
||||
_print_stage_breakdown(result, f"measured {m + 1}")
|
||||
_collect_stage_times(result, stage_times, stage_order)
|
||||
|
||||
print(f"\n=== summary (t2v / {_CELL}) ===")
|
||||
print(f"warmup wall-times: {[round(x, 1) for x in warmup_secs]}")
|
||||
if measured_secs:
|
||||
avg = sum(measured_secs) / len(measured_secs)
|
||||
print(
|
||||
f"measured e2e (n={len(measured_secs)}): "
|
||||
f"{[round(x, 2) for x in measured_secs]} -> avg {avg:.2f}s"
|
||||
)
|
||||
if stage_times:
|
||||
print(f"average stage times over {measured_runs} measured runs:")
|
||||
avg_total = 0.0
|
||||
for name in stage_order:
|
||||
vals = stage_times.get(name) or []
|
||||
if not vals:
|
||||
continue
|
||||
avg_v = sum(vals) / len(vals)
|
||||
avg_total += avg_v
|
||||
print(f" - {name}: {avg_v:.3f}s")
|
||||
print(f" - stage_sum_avg: {avg_total:.3f}s")
|
||||
finally:
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -122,9 +122,16 @@ except ImportError:
|
||||
_FA4_FP4_AVAILABLE = False
|
||||
|
||||
|
||||
def _nvfp4_quantize_for_fa4(tensor_4d: torch.Tensor, ) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
@torch.library.custom_op("fastvideo::_nvfp4_quantize_for_fa4", mutates_args=())
|
||||
def _nvfp4_quantize_for_fa4(tensor_4d: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Quantize a (batch, seqlen, nheads, headdim) BF16 tensor to FP4.
|
||||
|
||||
Registered as a ``torch.library.custom_op`` so torch.compile/dynamo treats
|
||||
it as an opaque leaf and never traces into flashinfer ``nvfp4_quantize``
|
||||
(whose ``@functools.cache``d module getter takes a ``_thread.allocate_lock``
|
||||
and JIT-builds via ``subprocess`` on first call — neither is dynamo-traceable).
|
||||
The paired ``register_fake`` below supplies output shapes/strides for tracing.
|
||||
|
||||
Returns:
|
||||
fp4_tensor: torch.float4_e2m1fn_x2, shape (batch, seqlen_padded, nheads, headdim//2)
|
||||
where seqlen_padded is seqlen rounded up to multiple of 128.
|
||||
@@ -171,6 +178,27 @@ def _nvfp4_quantize_for_fa4(tensor_4d: torch.Tensor, ) -> tuple[torch.Tensor, to
|
||||
return fp4_tensor, sf_mma
|
||||
|
||||
|
||||
@_nvfp4_quantize_for_fa4.register_fake
|
||||
def _nvfp4_quantize_for_fa4_fake(tensor_4d: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
# Shapes/strides must mirror the real impl above so the compiled graph
|
||||
# threads the right metadata into flash_attn_fp4_func. No flashinfer call —
|
||||
# pure meta-tensor allocation.
|
||||
batch, seqlen, nheads, headdim = tensor_4d.shape
|
||||
tile_m, sf_vec_size = 128, 16
|
||||
seqlen_padded = (seqlen + tile_m - 1) // tile_m * tile_m
|
||||
fp4_tensor = torch.empty((batch, seqlen_padded, nheads, headdim // 2),
|
||||
dtype=torch.float4_e2m1fn_x2, device=tensor_4d.device)
|
||||
atom_m0, atom_m1, atom_k = 32, 4, 4
|
||||
rest_m = seqlen_padded // tile_m
|
||||
rest_k = (headdim // sf_vec_size) // atom_k
|
||||
# Build contiguous then apply the same final permute so sf_mma's strides
|
||||
# (notably stride[3]==1, required by FA4) match the real tensor.
|
||||
sf_canonical = torch.empty((batch, nheads, rest_m, rest_k, atom_m0, atom_m1, atom_k),
|
||||
dtype=torch.uint8, device=tensor_4d.device)
|
||||
sf_mma = sf_canonical.permute(4, 5, 2, 6, 3, 1, 0)
|
||||
return fp4_tensor, sf_mma
|
||||
|
||||
|
||||
class FlashAttentionBackend(AttentionBackend):
|
||||
accept_output_buffer: bool = True
|
||||
|
||||
@@ -255,6 +283,20 @@ class FlashAttentionImpl(AttentionImpl):
|
||||
assert cap in [(10, 0), (10, 3)], (f"NVFP4 FA4 requires Blackwell (sm100a/sm103a), got sm{cap[0]}{cap[1]}")
|
||||
assert _FA4_FP4_AVAILABLE, ("NVFP4 FA4 requires flash-attention-fp4 (flash_attn.cute). "
|
||||
"Install via instructions in docs/inference/optimizations.md")
|
||||
# Pre-build the flashinfer FP4 quant module eagerly (here, in the
|
||||
# worker process, before torch.compile). The Q/K quant in
|
||||
# `_nvfp4_quantize_for_fa4` calls flashinfer `nvfp4_quantize`, whose
|
||||
# first invocation JIT-builds the sm10x module via `subprocess`
|
||||
# (nvcc --version). If that first call lands inside a compiled
|
||||
# region, dynamo tries to trace the subprocess and crashes. The
|
||||
# builder is `@functools.cache`d, so warming it once here makes the
|
||||
# later compiled call a pure cache hit.
|
||||
try:
|
||||
from flashinfer.quantization.fp4_quantization import (
|
||||
get_fp4_quantization_module)
|
||||
get_fp4_quantization_module(f"{cap[0]}{cap[1]}")
|
||||
except Exception as exc: # pragma: no cover - best-effort warmup
|
||||
logger.warning("FP4 quant module pre-build failed: %s", exc)
|
||||
logger.info("NVFP4 FA4 enabled for FlashAttentionImpl (quant_qk only)")
|
||||
|
||||
def forward(
|
||||
|
||||
Reference in New Issue
Block a user