Compare commits
2
Commits
main
...
klin/5090-fixes
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
bfe8c0d484 | ||
|
|
7a4a071986 |
@@ -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
@@ -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]
|
||||
|
||||
Reference in New Issue
Block a user