File size: 12,132 Bytes
bdce33b | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 | diff --git a/generate_infinitetalk.py b/generate_infinitetalk.py
index e83daa8..77426ce 100644
--- a/generate_infinitetalk.py
+++ b/generate_infinitetalk.py
@@ -20,7 +20,10 @@ import wan
from wan.configs import SIZE_CONFIGS, SUPPORTED_SIZES, WAN_CONFIGS
from wan.utils.utils import str2bool, is_video, split_wav_librosa
from wan.utils.multitalk_utils import save_video_ffmpeg
-from kokoro import KPipeline
+try:
+ from kokoro import KPipeline
+except ImportError: # TTS stack unused when driving from a local audio file
+ KPipeline = None
from transformers import Wav2Vec2FeatureExtractor
from src.audio_analysis.wav2vec2 import Wav2Vec2Model
from wan.utils.segvideo import shot_detect
@@ -544,6 +547,89 @@ def generate(args):
num_persistent_param_in_dit=args.num_persistent_param_in_dit
)
+ if os.environ.get("WAN_TRT") == "1":
+ # The engines own the transformer stack now: keep the PyTorch blocks off
+ # the GPU (~27GB) so the 18.6GB of engines fit alongside the rest.
+ import torch as _t
+ _n = len(wan_i2v.model.blocks)
+ wan_i2v.model.blocks = _t.nn.ModuleList([]) # the engines own these now
+ _t.cuda.empty_cache()
+ logging.info(f"WAN_TRT: dropped {_n} PyTorch blocks (~27GB); engines own the stack")
+
+ if os.environ.get("CAPTURE") == "1":
+ import sys as _sys
+ _sys.path.insert(0, "/workspace/trt")
+ from capture_calib import install_capture
+ install_capture(wan_i2v.model)
+
+ # --- FP8 (H100 native) + torch.compile toggles, applied after LoRA merge ---
+ if os.environ.get("WAN_FP8") == "1":
+ import torch as _t
+ from torchao.quantization import quantize_, Float8DynamicActivationFloat8WeightConfig
+ _dit = wan_i2v.model
+ # Only the transformer blocks: the tiny embedders/head stay bf16.
+ _n = 0
+ for _blk in _dit.blocks:
+ quantize_(_blk, Float8DynamicActivationFloat8WeightConfig())
+ _n += 1
+ logging.info(f"FP8: quantized {_n} transformer blocks (fp8 dynamic act + fp8 weight)")
+
+ if os.environ.get("WAN_COMPILE") == "1":
+ import torch as _t
+ _dit = wan_i2v.model
+ for _i, _blk in enumerate(_dit.blocks):
+ _dit.blocks[_i] = _t.compile(_blk, dynamic=False)
+ logging.info(f"compiled {len(_dit.blocks)} transformer blocks")
+
+ # --- profiling hook: set PROFILE=1 to time every DiT forward ---
+ if os.environ.get("PROFILE") == "1":
+ import time as _time, atexit as _atexit
+ import torch as _torch
+ _stats = {"n": 0, "t": 0.0}
+ _wall0 = _time.perf_counter()
+ _orig_fwd = wan_i2v.model.forward
+
+ _deep = os.environ.get("PROFILE_DEEP") == "1"
+
+ def _timed_fwd(*a, **kw):
+ _torch.cuda.synchronize()
+ _t0 = _time.perf_counter()
+ # kernel-level trace of a single steady-state forward (the 2nd)
+ if _deep and _stats["n"] == 1:
+ from torch.profiler import profile as _tp, ProfilerActivity as _PA
+ with _tp(activities=[_PA.CPU, _PA.CUDA], record_shapes=False) as _prof:
+ out = _orig_fwd(*a, **kw)
+ _torch.cuda.synchronize()
+ print("\n======= TOP CUDA KERNELS (one forward) =======")
+ print(_prof.key_averages().table(
+ sort_by="self_cuda_time_total", row_limit=28,
+ max_name_column_width=55))
+ else:
+ out = _orig_fwd(*a, **kw)
+ _torch.cuda.synchronize()
+ _stats["t"] += _time.perf_counter() - _t0
+ _stats["n"] += 1
+ return out
+
+ wan_i2v.model.forward = _timed_fwd
+
+ def _report():
+ n, t = _stats["n"], _stats["t"]
+ wall = _time.perf_counter() - _wall0
+ peak = _torch.cuda.max_memory_allocated() / 1e9
+ print("\n================ PROFILE ================")
+ print(f"DiT forwards : {n}")
+ print(f"DiT total time : {t:.1f} s")
+ if n:
+ print(f"DiT per forward : {t / n * 1000:.0f} ms")
+ print(f"wall (post-load) : {wall:.1f} s")
+ if wall > 0:
+ print(f"DiT share of wall : {t / wall * 100:.0f} %")
+ print(f"peak VRAM allocated : {peak:.1f} GB")
+ print("=========================================")
+
+ _atexit.register(_report)
+
generated_list = []
with open(args.input_json, 'r', encoding='utf-8') as f:
input_data = json.load(f)
diff --git a/wan/modules/multitalk_model.py b/wan/modules/multitalk_model.py
index 958e930..6eef6f3 100644
--- a/wan/modules/multitalk_model.py
+++ b/wan/modules/multitalk_model.py
@@ -21,6 +21,29 @@ try:
except:
USE_SAGEATTN = False
+from torch.nn.attention import sdpa_kernel, SDPBackend
+
+# cuDNN's Hopper attention beat sageattn 44.6ms vs 58.5ms at our shape.
+USE_CUDNN_SDPA = os.environ.get("WAN_CUDNN_ATTN", "1") == "1"
+
+# --- FP8 TensorRT block stack (WAN_TRT=1) -------------------------------------
+_TRT_STACK = None
+_TRT_SEQ_LEN = int(os.environ.get("WAN_TRT_SEQ_LEN", "30576")) # 81 frames @ 448x832
+
+
+def _trt_stack():
+ """Lazily load the engines; returns None when TRT is off."""
+ global _TRT_STACK
+ if os.environ.get("WAN_TRT") != "1":
+ return None
+ if _TRT_STACK is None:
+ import sys
+ sys.path.insert(0, "/workspace/trt")
+ from trt_runner import TRTStack
+ _TRT_STACK = TRTStack()
+ return _TRT_STACK
+
+
__all__ = ['WanModel']
@@ -50,30 +73,52 @@ def rope_params(max_seq_len, dim, theta=10000):
return freqs
+_ROPE_CACHE = {}
+
+
+def _rope_cos_sin(grid_sizes, freqs, device, dtype):
+ """cos/sin RoPE table for this latent grid, built once and kept on-device."""
+ f, h, w = (int(v) for v in grid_sizes[0].tolist())
+ key = (f, h, w, str(device), dtype)
+ hit = _ROPE_CACHE.get(key)
+ if hit is not None:
+ return hit
+
+ c = freqs.size(1)
+ fr = freqs.split([c - 2 * (c // 3), c // 3, c // 3], dim=1)
+ freqs_i = torch.cat([
+ fr[0][:f].view(f, 1, 1, -1).expand(f, h, w, -1),
+ fr[1][:h].view(1, h, 1, -1).expand(f, h, w, -1),
+ fr[2][:w].view(1, 1, w, -1).expand(f, h, w, -1),
+ ], dim=-1).reshape(f * h * w, 1, -1) # [L, 1, C/2] complex
+
+ cos = freqs_i.real.to(device=device, dtype=dtype).contiguous()
+ sin = freqs_i.imag.to(device=device, dtype=dtype).contiguous()
+ _ROPE_CACHE[key] = (cos, sin)
+ return cos, sin
+
+
@amp.autocast(enabled=False)
def rope_apply(x, grid_sizes, freqs):
- s, n, c = x.size(1), x.size(2), x.size(3) // 2
+ """Real-valued RoPE. Mathematically identical to the complex fp64 version,
+ minus the per-layer host->device copy and the float64 traffic."""
+ b, s, n, d = x.shape
+ cos, sin = _rope_cos_sin(grid_sizes, freqs, x.device, torch.float32)
+ L = cos.size(0)
- freqs = freqs.split([c - 2 * (c // 3), c // 3, c // 3], dim=1)
+ xf = x[:, :L].float().reshape(b, L, n, d // 2, 2)
+ x_r, x_i = xf[..., 0], xf[..., 1]
- output = []
- for i, (f, h, w) in enumerate(grid_sizes.tolist()):
- seq_len = f * h * w
+ cos_ = cos.unsqueeze(0) # [1, L, 1, C/2]
+ sin_ = sin.unsqueeze(0)
- x_i = torch.view_as_complex(x[i, :s].to(torch.float64).reshape(
- s, n, -1, 2))
- freqs_i = torch.cat([
- freqs[0][:f].view(f, 1, 1, -1).expand(f, h, w, -1),
- freqs[1][:h].view(1, h, 1, -1).expand(f, h, w, -1),
- freqs[2][:w].view(1, 1, w, -1).expand(f, h, w, -1)
- ],
- dim=-1).reshape(seq_len, 1, -1)
- freqs_i = freqs_i.to(device=x_i.device)
- x_i = torch.view_as_real(x_i * freqs_i).flatten(2)
- x_i = torch.cat([x_i, x[i, seq_len:]])
+ o_r = x_r * cos_ - x_i * sin_
+ o_i = x_r * sin_ + x_i * cos_
+ out = torch.stack([o_r, o_i], dim=-1).flatten(3)
- output.append(x_i)
- return torch.stack(output).float()
+ if L < s: # keep any padding untouched
+ out = torch.cat([out, x[:, L:].float()], dim=1)
+ return out.float()
class WanRMSNorm(nn.Module):
@@ -137,7 +182,7 @@ class WanSelfAttention(nn.Module):
self.norm_q = WanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
self.norm_k = WanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
- def forward(self, x, seq_lens, grid_sizes, freqs, ref_target_masks=None):
+ def forward(self, x, seq_lens, grid_sizes, freqs, ref_target_masks=None, human_num=None):
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
# query, key, value function
@@ -151,7 +196,15 @@ class WanSelfAttention(nn.Module):
q = rope_apply(q, grid_sizes, freqs)
k = rope_apply(k, grid_sizes, freqs)
- if USE_SAGEATTN:
+ if USE_CUDNN_SDPA:
+ # [B, L, H, D] -> [B, H, L, D] for SDPA, and back
+ qb = q.transpose(1, 2).to(torch.bfloat16)
+ kb = k.transpose(1, 2).to(torch.bfloat16)
+ vb = v.transpose(1, 2).to(torch.bfloat16)
+ with sdpa_kernel(SDPBackend.CUDNN_ATTENTION):
+ x = F.scaled_dot_product_attention(qb, kb, vb)
+ x = x.transpose(1, 2).type_as(v)
+ elif USE_SAGEATTN:
x = sageattn(q.to(torch.bfloat16), k.to(torch.bfloat16), v, tensor_layout='NHD')
else:
x = flash_attention(
@@ -165,9 +218,15 @@ class WanSelfAttention(nn.Module):
# output
x = x.flatten(2)
x = self.o(x)
- with torch.no_grad():
- x_ref_attn_map = get_attn_map_with_target(q.type_as(x), k.type_as(x), grid_sizes[0],
- ref_target_masks=ref_target_masks)
+ # The ref-attn map only feeds SingleStreamMutiAttention's multi-speaker routing;
+ # with one speaker that branch short-circuits and never reads it, so skip building
+ # a [heads, seq, ref_seq] map (GBs) in every layer.
+ if human_num == 1:
+ x_ref_attn_map = None
+ else:
+ with torch.no_grad():
+ x_ref_attn_map = get_attn_map_with_target(q.type_as(x), k.type_as(x), grid_sizes[0],
+ ref_target_masks=ref_target_masks)
return x, x_ref_attn_map
@@ -294,7 +353,7 @@ class WanAttentionBlock(nn.Module):
# self-attention
y, x_ref_attn_map = self.self_attn(
(self.norm1(x).float() * (1 + e[1]) + e[0]).type_as(x), seq_lens, grid_sizes,
- freqs, ref_target_masks=ref_target_masks)
+ freqs, ref_target_masks=ref_target_masks, human_num=human_num)
with amp.autocast(dtype=torch.float32):
x = x + y * e[2]
@@ -757,7 +816,14 @@ class WanModel(ModelMixin, ConfigMixin):
for block in self.blocks:
x = block(x, **kwargs)
self.previous_residual_uncond = x - ori_x
+ elif _trt_stack() is not None and x.shape[1] == _TRT_SEQ_LEN:
+ cos, sin = _rope_cos_sin(grid_sizes, self.freqs, x.device, torch.float32)
+ x = _trt_stack()(x, e0, context, audio_embedding.squeeze(0), cos, sin)
else:
+ if os.environ.get("WAN_TRT") == "1":
+ raise RuntimeError(
+ f"WAN_TRT=1 but seq_len {x.shape[1]} != engine seq_len {_TRT_SEQ_LEN}. "
+ f"Engines are static; rebuild them for this frame count/resolution.")
for block in self.blocks:
x = block(x, **kwargs)
diff --git a/wan/multitalk.py b/wan/multitalk.py
index be7819e..1e9513b 100644
--- a/wan/multitalk.py
+++ b/wan/multitalk.py
@@ -1,6 +1,5 @@
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
import gc
-from inspect import ArgSpec
import logging
import json
import math
|