Compare commits

...
2 Commits
Author SHA1 Message Date
kevin314 bfe8c0d484 Revert redundant fixes 2026-06-24 23:25:42 +00:00
kevin314 7a4a071986 Init 2026-06-24 23:11:47 +00:00
4 changed files with 203 additions and 88 deletions
@@ -23,14 +23,12 @@ Usage:
python fp4_linear_taehv_wan2_1_1_3b.py --no-compile # eager
python fp4_linear_taehv_wan2_1_1_3b.py --baseline # dense bf16 reference
python fp4_linear_taehv_wan2_1_1_3b.py --distilled_model '' # base Wan2.1 weights
python fp4_linear_taehv_wan2_1_1_3b.py --warmups 5 --benchmark-runs 20 # timing stats (default)
"""
import argparse
import contextlib
import logging
import os
import statistics
import time
import imageio
@@ -202,10 +200,6 @@ def main() -> None:
parser.add_argument("--num_gpus", type=int, default=1)
parser.add_argument("--infer_steps", type=int, default=3)
parser.add_argument("--guidance_scale", type=float, default=1.0)
parser.add_argument("--warmups", type=int, default=5,
help="Warmup runs before timing (default: 5).")
parser.add_argument("--benchmark-runs", type=int, default=20,
help="Timed runs to collect min/max/mean/std over (default: 20).")
args = parser.parse_args()
if not torch.cuda.is_available():
@@ -232,11 +226,14 @@ def main() -> None:
os.makedirs(OUTPUT_PATH, exist_ok=True)
# Warmup: pay the DiT torch.compile cost and warm TAEHV's cuDNN algo
# selection + allocator growth (~0.2s on the first decode) so the timed
# runs below measure steady-state latency only.
# Warmup: with compile enabled the first call(s) pay the DiT compilation
# cost. When using TAEHV we also decode the warmup latents so the timed
# decode below is warm -- TAEHV's decoder is all conv/upsample, so the
# first call otherwise pays cuDNN algo selection + allocator growth
# (~0.2s), which is exactly the cold-start overhead we want to exclude.
n_warmup = 2 if not args.no_compile else 1
with silence_request_log():
for _ in range(args.warmups):
for _ in range(n_warmup):
warm = generator.generate(request={
"prompt": PROMPT,
"sampling": {"num_inference_steps": 2, "guidance_scale": args.guidance_scale},
@@ -246,71 +243,44 @@ def main() -> None:
taehv.decode(warm.samples)
output_path = os.path.join(OUTPUT_PATH, f"raccoon_{mode}.mp4")
torch.cuda.synchronize()
start = time.perf_counter()
result = generator.generate(request={
"prompt": PROMPT,
"sampling": {
"num_inference_steps": args.infer_steps,
"guidance_scale": args.guidance_scale,
},
# When using TAEHV we need the latents back and save manually; the Wan
# VAE path lets the pipeline decode and save the mp4 itself.
"output": {
"save_video": not args.taehv,
"return_frames": args.taehv,
"output_path": output_path,
},
})
torch.cuda.synchronize()
denoise_elapsed = result.generation_time
if args.taehv:
torch.cuda.synchronize()
decode_start = time.perf_counter()
frames = taehv.decode(result.samples)
torch.cuda.synchronize()
decode_elapsed = time.perf_counter() - decode_start
# Benchmark: time each run end-to-end. ``denoise`` is the generator's own
# generation_time; for TAEHV we add the in-script decode. The mp4 is not
# written inside the loop (only the final run's frames are saved below) so
# disk I/O never pollutes the timings.
denoise_times: list[float] = []
decode_times: list[float] = []
totals: list[float] = []
frames = None
with silence_request_log():
for i in range(args.benchmark_runs):
result = generator.generate(request={
"prompt": PROMPT,
"sampling": {
"num_inference_steps": args.infer_steps,
"guidance_scale": args.guidance_scale,
},
"output": {
"save_video": False,
"return_frames": args.taehv,
"output_path": output_path,
},
})
denoise_elapsed = result.generation_time
denoise_times.append(denoise_elapsed)
if args.taehv:
torch.cuda.synchronize()
decode_start = time.perf_counter()
frames = taehv.decode(result.samples)
torch.cuda.synchronize()
decode_elapsed = time.perf_counter() - decode_start
decode_times.append(decode_elapsed)
total = denoise_elapsed + decode_elapsed
else:
total = denoise_elapsed
totals.append(total)
line = f" run {i + 1:02d}/{args.benchmark_runs}: {total:.3f}s"
if args.taehv:
line += (f" (denoise {denoise_elapsed:.3f}s + "
f"decode {decode_elapsed:.3f}s)")
print(line)
if args.taehv and frames is not None:
imageio.mimsave(output_path, frames, fps=16, format="mp4")
total = denoise_elapsed + decode_elapsed
print(f"[{mode.upper()}] denoise {denoise_elapsed:.2f}s + TAEHV decode "
f"{decode_elapsed:.2f}s = {total:.2f}s "
f"({frames.shape[0]} frames @ {tuple(frames.shape[1:3])})")
print(f"Saved video to {output_path}")
# Report min / max / mean / std over the timed runs.
def _stat_row(name: str, xs: list[float]) -> str:
std = statistics.stdev(xs) if len(xs) > 1 else 0.0
return (f" {name:<11}{min(xs):>8.3f}{max(xs):>9.3f}"
f"{statistics.mean(xs):>9.3f}{std:>9.3f}")
if totals:
print(f"\n[{mode.upper()}] {args.benchmark_runs} runs, {args.warmups} warmup, "
f"{args.infer_steps} steps:")
print(f" {'metric':<11}{'min':>8}{'max':>9}{'mean':>9}{'std':>9} (s)")
print(_stat_row("total", totals))
if args.taehv:
print(_stat_row("denoise", denoise_times))
print(_stat_row("decode", decode_times))
else:
print(f"[{mode.upper()}] {args.infer_steps} steps in {denoise_elapsed:.2f}s "
f"({args.infer_steps / denoise_elapsed:.2f} it/s)")
generator.shutdown()
if __name__ == "__main__":
main()
main()
@@ -0,0 +1,150 @@
"""NVFP4 QAD inference example.
Runs Wan2.1-T2V-1.3B with the FastWan-QAD-1.3B distilled checkpoint and
NVFP4QATConfig quantization. Uses ATTN_QAT_INFER attention backend.
Requirements:
- GPU: Blackwell (B200/B300, sm100a+) for the FP4 linear path
- TAEHV (optional): Follow install instructions at https://github.com/madebyollin/taehv
Usage:
python fp4_qad_wan2_1_1_3b.py # NVFP4 QAD (default)
python fp4_qad_wan2_1_1_3b.py --bf16 # BF16 baseline
python fp4_qad_wan2_1_1_3b.py --taehv-checkpoint /path/to/taew2_1.pth
"""
import argparse
import os
import sys
import time
import torch
OUTPUT_PATH = "video_samples"
def load_taehv(checkpoint_path, device="cuda", dtype=torch.float16):
repo_dir = os.path.dirname(checkpoint_path)
if repo_dir not in sys.path:
sys.path.insert(0, repo_dir)
from taehv import TAEHV
print(f"Loading TAEHV from {checkpoint_path}...")
model = TAEHV(checkpoint_path=checkpoint_path).to(device, dtype)
print("TAEHV loaded.")
return model
@torch.no_grad() # type: ignore[misc]
def decode_with_taehv(taehv_model, latents):
latents = latents.permute(0, 2, 1, 3, 4)
latents = latents.to(device=next(taehv_model.parameters()).device,
dtype=next(taehv_model.parameters()).dtype)
decoded = taehv_model.decode_video(latents, parallel=False, show_progress_bar=False)
frames = []
for frame in decoded[0]:
frame_np = (frame.clamp(0, 1) * 255).byte().cpu().permute(1, 2, 0).numpy()
frames.append(frame_np)
return frames
def main():
parser = argparse.ArgumentParser(description="NVFP4 QAD video generation benchmark")
parser.add_argument("--bf16", action="store_true",
help="BF16 baseline (no NVFP4 quantization)")
parser.add_argument("--taehv-checkpoint", default=None, metavar="PATH",
help="Path to taew2_1.pth; enables TAEHV tiny autoencoder decoding")
parser.add_argument("--model", default="FastVideo/FastWan-QAD-1.3B",
help="Model path or HuggingFace ID")
parser.add_argument("--no-compile", action="store_true", help="Disable torch.compile for the DiT")
parser.add_argument("--num_gpus", type=int, default=1)
parser.add_argument("--infer_steps", type=int, default=3)
args = parser.parse_args()
os.environ.setdefault("FASTVIDEO_ATTENTION_BACKEND", "ATTN_QAT_INFER")
os.environ["FASTVIDEO_DISABLE_ATTENTION_COMPILE"] = "0"
os.environ["FLASHINFER_CUDA_ARCH_LIST"] = "12.0a"
# os.environ["FLASHINFER_EXTRA_CFLAGS"] = "-DCCCL_DISABLE_CTK_COMPATIBILITY_CHECK"
# os.environ["FLASHINFER_EXTRA_CUDAFLAGS"] = "-DCCCL_DISABLE_CTK_COMPATIBILITY_CHECK"
# os.environ["CUDA_HOME"] = "/root/miniconda3/envs/fastvideo/lib/python3.12/site-packages/nvidia/cu13"
from fastvideo import VideoGenerator
from fastvideo.configs.pipelines.base import PipelineConfig
mode = "bf16" if args.bf16 else "nvfp4_qad"
if not args.no_compile:
mode += "_compile"
use_taehv = args.taehv_checkpoint is not None
print(f"Mode: {mode.upper()}" + (" decoder=TAEHV" if use_taehv else " decoder=VAE"))
taehv_model = load_taehv(args.taehv_checkpoint) if use_taehv else None
pipeline_config = PipelineConfig.from_pretrained(args.model)
pipeline_config.text_encoder_precisions = ("bf16",)
if not args.bf16:
from fastvideo.layers.quantization.nvfp4_qat_config import NVFP4QATConfig
pipeline_config.dit_config.quant_config = NVFP4QATConfig()
generator = VideoGenerator.from_pretrained(
args.model,
pipeline_config=pipeline_config,
num_gpus=args.num_gpus,
use_fsdp_inference=False,
dit_cpu_offload=False,
dit_layerwise_offload=False,
vae_cpu_offload=use_taehv,
text_encoder_cpu_offload=False,
pin_cpu_memory=False,
enable_torch_compile=not args.no_compile,
enable_torch_compile_text_encoder=not args.no_compile,
enable_torch_compile_vae=not args.no_compile and not use_taehv,
output_type="latent" if use_taehv else "pil",
)
prompt = (
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
)
n_warmup = 2 if not args.no_compile else 0
for _ in range(n_warmup):
warmup_result = generator.generate(request={"prompt": prompt, "sampling": {"num_inference_steps": 3, "guidance_scale": 1.0},
"output": {"save_video": False}})
if use_taehv:
decode_with_taehv(taehv_model, warmup_result.samples)
os.makedirs(OUTPUT_PATH, exist_ok=True)
video_path = os.path.join(OUTPUT_PATH, f"raccoon_{mode}.mp4")
if use_taehv:
import imageio
result = generator.generate(request={
"prompt": prompt,
"sampling": {"num_inference_steps": args.infer_steps, "guidance_scale": 1.0},
"output": {"save_video": False},
})
denoise_elapsed = result.generation_time
torch.cuda.synchronize()
t_decode = time.perf_counter()
frames = decode_with_taehv(taehv_model, result.samples)
torch.cuda.synchronize()
decode_elapsed = time.perf_counter() - t_decode
total = denoise_elapsed + decode_elapsed
imageio.mimsave(video_path, frames, fps=16, format="mp4")
print(f"Saved TAEHV-decoded video to: {video_path}")
print(f"[{mode.upper()}] {args.infer_steps} steps in {total:.3f}s "
f"(denoise {denoise_elapsed:.3f}s + decode {decode_elapsed:.3f}s)")
else:
result = generator.generate(request={
"prompt": prompt,
"sampling": {"num_inference_steps": args.infer_steps, "guidance_scale": 1.0},
"output": {"save_video": True, "output_path": video_path},
})
elapsed = result.generation_time
print(f"[{mode.upper()}] {args.infer_steps} steps in {elapsed:.3f}s "
f"({args.infer_steps / elapsed:.2f} it/s)")
generator.shutdown()
if __name__ == "__main__":
main()
@@ -286,24 +286,19 @@
// load input
const int token_id = token_block_id * BLOCK_SIZE + threadIdx.x / NUM_THREADS_PER_TOKEN;
// Permute V rows within each 32-element block so the PV MMA K-indexed
// access reads the correct CLayout N-indexed values (Edenzzzz causal fix).
const int k_intra = token_id & 31;
const int load_token_id = (token_id & ~31)
| ((k_intra & 6) << 2) | ((k_intra & 24) >> 2) | (k_intra & 1);
PackedVec in_vec;
#pragma unroll
for (int i = 0; i < CVT_FP4_ELTS_PER_THREAD / 2; i++) {
reinterpret_cast<uint32_t&>(in_vec.elts[i]) = 0;
}
if (load_token_id < num_tokens) {
in_vec = reinterpret_cast<PackedVec const*>(input +
if (token_id < num_tokens) {
in_vec = reinterpret_cast<PackedVec const*>(input +
batch_id * stride_bz_input + // batch dim
head_id * stride_h_input + // head dim
load_token_id * stride_seq_input + // seq dim (permuted)
token_id * stride_seq_input + // seq dim
(threadIdx.x % NUM_THREADS_PER_TOKEN) * CVT_FP4_ELTS_PER_THREAD)[0]; // feature dim
}
+7 -7
View File
@@ -27,9 +27,10 @@ dependencies = [
"timm>=1.0.11",
"peft>=0.15.0",
"diffusers>=0.38.0",
"torch==2.11.0",
"torch==2.12.0",
"torchvision",
"torchaudio",
"flashinfer-python",
# Acceleration & Optimization
"accelerate==1.0.1",
@@ -67,7 +68,6 @@ dependencies = [
"remote-pdb",
# Kernel & Packaging
"fastvideo-kernel==0.3.0",
"wheel",
# Training Dependencies
@@ -95,15 +95,15 @@ members = ["apps/dreamverse"]
[tool.uv.sources]
torch = [
{ index = "pytorch-cpu", marker = "sys_platform != 'linux'" },
{ index = "pytorch-cu128", marker = "sys_platform == 'linux'" },
{ index = "pytorch-cu130", marker = "sys_platform == 'linux'" },
]
torchvision = [
{ index = "pytorch-cpu", marker = "sys_platform != 'linux'" },
{ index = "pytorch-cu128", marker = "sys_platform == 'linux'" },
{ index = "pytorch-cu130", marker = "sys_platform == 'linux'" },
]
torchaudio = [
{ index = "pytorch-cpu", marker = "sys_platform != 'linux'" },
{ index = "pytorch-cu128", marker = "sys_platform == 'linux'" },
{ index = "pytorch-cu130", marker = "sys_platform == 'linux'" },
]
imagebind = { git = "https://github.com/facebookresearch/ImageBind.git", rev = "53680b02d7e37b19b124fa37bae4b6c98c38f5be" }
flash-attn-cute = { git = "https://github.com/XOR-op/flash-attention.git", branch = "fa4-compile", subdirectory = "flash_attn/cute" }
@@ -114,8 +114,8 @@ url = "https://download.pytorch.org/whl/cpu"
explicit = true
[[tool.uv.index]]
name = "pytorch-cu128"
url = "https://download.pytorch.org/whl/cu128"
name = "pytorch-cu130"
url = "https://download.pytorch.org/whl/cu130"
explicit = true
[project.optional-dependencies]