Merge branch 'main' into longcat_avatar
This commit is contained in:
@@ -1513,6 +1513,8 @@ class WanVideoModelLoader:
|
||||
)
|
||||
transformer.multitalk_audio_proj = multitalk_proj_model
|
||||
|
||||
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...")
|
||||
|
||||
+1
-8
@@ -2301,14 +2301,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
|
||||
|
||||
|
||||
@@ -2213,10 +2213,10 @@ class WanModel(torch.nn.Module):
|
||||
img_ids[:, :, :, 0] = img_ids[:, :, :, 0] + grid_t.reshape(-1, 1, 1)
|
||||
else:
|
||||
# Standard temporal encoding
|
||||
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[:, :, :, 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, 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[:, :, :, 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
|
||||
@@ -2540,6 +2540,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]
|
||||
|
||||
@@ -2667,12 +2668,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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user