From 0ac7401271d77345c8adafa83dcc0a798e02172f Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Mon, 22 Dec 2025 18:56:04 +0200 Subject: [PATCH 1/3] Stand-in RoPE adjustments The offset was a bit wrong --- wanvideo/modules/model.py | 15 +++++++-------- 1 file changed, 7 insertions(+), 8 deletions(-) diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index b2d19f8..c920440 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -2144,9 +2144,9 @@ class WanModel(torch.nn.Module): # Main frames position IDs img_ids = torch.zeros((steps_t, steps_h, steps_w, 3), device=device, dtype=dtype) - img_ids[:, :, :, 0] = img_ids[:, :, :, 0] + torch.linspace(t_start+freq_offset, t_start + (t_len - 1), steps=steps_t, device=device, dtype=dtype).reshape(-1, 1, 1) - img_ids[:, :, :, 1] = img_ids[:, :, :, 1] + torch.linspace(freq_offset, h_len - 1, steps=steps_h, device=device, dtype=dtype).reshape(1, -1, 1) - img_ids[:, :, :, 2] = img_ids[:, :, :, 2] + torch.linspace(freq_offset, w_len - 1, steps=steps_w, device=device, dtype=dtype).reshape(1, 1, -1) + img_ids[:, :, :, 0] = img_ids[:, :, :, 0] + torch.linspace(t_start+freq_offset, t_start+freq_offset + (t_len - 1), steps=steps_t, device=device, dtype=dtype).reshape(-1, 1, 1) + img_ids[:, :, :, 1] = img_ids[:, :, :, 1] + torch.linspace(freq_offset, freq_offset + (h_len - 1), steps=steps_h, device=device, dtype=dtype).reshape(1, -1, 1) + img_ids[:, :, :, 2] = img_ids[:, :, :, 2] + torch.linspace(freq_offset, freq_offset + (w_len - 1), steps=steps_w, device=device, dtype=dtype).reshape(1, 1, -1) img_ids = img_ids.reshape(1, -1, img_ids.shape[-1]) segments = [img_ids] # Start with main frames @@ -2470,6 +2470,7 @@ class WanModel(torch.nn.Module): # grid sizes and seq len grid_sizes = torch.stack([torch.tensor(u.shape[2:], device=device, dtype=torch.long) for u in x]) original_grid_sizes = grid_sizes.clone() + f, h, w = x[0].shape[2:] x = [u.flatten(2).transpose(1, 2) for u in x] self.original_seq_len = x[0].shape[1] @@ -2595,12 +2596,10 @@ class WanModel(torch.nn.Module): # Stand-In RoPE frequencies if x_ip is not None: # Generate RoPE frequencies for x_ip - h_len = (H + 1) // 2 - w_len = (W + 1) // 2 ip_img_ids = torch.zeros((f_ip, h_ip, w_ip, 3), device=x.device, dtype=x.dtype) - ip_img_ids[:, :, :, 0] = ip_img_ids[:, :, :, 0] + torch.linspace(0, f_ip - 1, steps=f_ip, device=x.device, dtype=x.dtype).reshape(-1, 1, 1) - ip_img_ids[:, :, :, 1] = ip_img_ids[:, :, :, 1] + torch.linspace(h_len + freq_offset, h_len + freq_offset + h_ip - 1, steps=h_ip, device=x.device, dtype=x.dtype).reshape(1, -1, 1) - ip_img_ids[:, :, :, 2] = ip_img_ids[:, :, :, 2] + torch.linspace(w_len + freq_offset, w_len + freq_offset + w_ip - 1, steps=w_ip, device=x.device, dtype=x.dtype).reshape(1, 1, -1) + ip_img_ids[:, :, :, 0] = -1 + ip_img_ids[:, :, :, 1] = ip_img_ids[:, :, :, 1] + torch.linspace(h + freq_offset, h + freq_offset + (h_ip - 1), steps=h_ip, device=x.device, dtype=x.dtype).reshape(1, -1, 1) + ip_img_ids[:, :, :, 2] = ip_img_ids[:, :, :, 2] + torch.linspace(w + freq_offset, w + freq_offset + (w_ip - 1), steps=w_ip, device=x.device, dtype=x.dtype).reshape(1, 1, -1) ip_img_ids = repeat(ip_img_ids, "t h w c -> b (t h w) c", b=1) freqs_ip = self.rope_embedder(ip_img_ids).movedim(1, 2) From 41683a042311d9939ce031fe0d50d0add71f9983 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Mon, 22 Dec 2025 19:44:53 +0200 Subject: [PATCH 2/3] Possibly fix uni3c + multitalk --- nodes_sampler.py | 9 +-------- 1 file changed, 1 insertion(+), 8 deletions(-) diff --git a/nodes_sampler.py b/nodes_sampler.py index d9943f0..6078a5e 100644 --- a/nodes_sampler.py +++ b/nodes_sampler.py @@ -2278,14 +2278,7 @@ class WanVideoSampler: uni3c_data = uni3c_data_input = None if uni3c_embeds is not None: transformer.controlnet = uni3c_embeds["controlnet"] - uni3c_data = { - "render_latent": uni3c_embeds["render_latent"], - "render_mask": uni3c_embeds["render_mask"], - "camera_embedding": uni3c_embeds["camera_embedding"], - "controlnet_weight": uni3c_embeds["controlnet_weight"], - "start": uni3c_embeds["start"], - "end": uni3c_embeds["end"], - } + uni3c_data = uni3c_embeds.copy() encoded_silence = None From 3e450214225c2934d269f2398378f0af477fd239 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Mon, 22 Dec 2025 21:40:14 +0200 Subject: [PATCH 3/3] Support mixed precision fp8 scaled models --- nodes_model_loading.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/nodes_model_loading.py b/nodes_model_loading.py index d07a71a..e2c24d7 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -1480,6 +1480,8 @@ class WanVideoModelLoader: sd.update(extra_sd) del extra_sd + sd = {k.replace(".weight_scale", ".scale_weight"): v for k, v in sd.items()} + # FlashVSR if "LQ_proj_in.norm1.gamma" in sd: log.info("FlashVSR model detected, patching model...")