Merge branch 'main' into dev

This commit is contained in:
kijai
2025-08-25 19:35:13 +03:00
2 changed files with 10 additions and 12 deletions
+6 -8
View File
@@ -2070,19 +2070,19 @@ class WanVideoSampler:
pad_len = latent_video_length - dwpose_data.shape[2]
pad = dwpose_data[:,:,:1].repeat(1,1,pad_len,1,1)
dwpose_data = torch.cat([dwpose_data, pad], dim=2)
dwpose_data_flat = rearrange(dwpose_data, 'b c f h w -> b (f h w) c').contiguous()
random_ref_dwpose_data = None
if image_cond is not None:
transformer.randomref_embedding_pose.to(device)
transformer.randomref_embedding_pose.to(device, dtype)
random_ref_dwpose = unianimate_poses.get("ref", None)
if random_ref_dwpose is not None:
random_ref_dwpose_data = transformer.randomref_embedding_pose(
random_ref_dwpose.to(device)
random_ref_dwpose.to(device, dtype)
).unsqueeze(2).to(model["dtype"]) # [1, 20, 104, 60]
del random_ref_dwpose
unianim_data = {
"dwpose": dwpose_data_flat,
"dwpose": dwpose_data,
"random_ref": random_ref_dwpose_data.squeeze(0) if random_ref_dwpose_data is not None else None,
"strength": unianimate_poses["strength"],
"start_percent": unianimate_poses["start_percent"],
@@ -3162,9 +3162,8 @@ class WanVideoSampler:
partial_unianim_data = None
if unianim_data is not None:
partial_dwpose = dwpose_data[:, :, c]
partial_dwpose_flat=rearrange(partial_dwpose, 'b c f h w -> b (f h w) c')
partial_unianim_data = {
"dwpose": partial_dwpose_flat,
"dwpose": partial_dwpose,
"random_ref": unianim_data["random_ref"],
"strength": unianimate_poses["strength"],
"start_percent": unianimate_poses["start_percent"],
@@ -3472,9 +3471,8 @@ class WanVideoSampler:
partial_unianim_data = None
if unianim_data is not None:
partial_dwpose = dwpose_data[:, :, latent_start_idx:latent_end_idx]
partial_dwpose_flat=rearrange(partial_dwpose, 'b c f h w -> b (f h w) c')
partial_unianim_data = {
"dwpose": partial_dwpose_flat,
"dwpose": partial_dwpose,
"random_ref": unianim_data["random_ref"],
"strength": unianimate_poses["strength"],
"start_percent": unianimate_poses["start_percent"],
+4 -4
View File
@@ -1674,7 +1674,7 @@ class WanModel(torch.nn.Module):
if unianim_data['start_percent'] <= current_step_percentage <= unianim_data['end_percent']:
random_ref_emb = unianim_data["random_ref"]
if random_ref_emb is not None:
y[0] = y[0] + random_ref_emb * unianim_data["strength"]
y[0].add_(random_ref_emb, alpha=unianim_data["strength"])
x = [torch.cat([u, v], dim=0) for u, v in zip(x, y)]
#uni3c controlnet
@@ -2058,8 +2058,8 @@ class WanModel(torch.nn.Module):
if hasattr(self, "dwpose_embedding") and unianim_data is not None:
if unianim_data['start_percent'] <= current_step_percentage <= unianim_data['end_percent']:
dwpose_emb = unianim_data['dwpose']
x += dwpose_emb * unianim_data['strength']
dwpose_emb = rearrange(unianim_data['dwpose'], 'b c f h w -> b (f h w) c').contiguous()
x.add_(dwpose_emb, alpha=unianim_data['strength'])
# arguments
kwargs = dict(
e=e0,
@@ -2162,7 +2162,7 @@ class WanModel(torch.nn.Module):
if b in self.slg_blocks and is_uncond:
if self.slg_start_percent <= current_step_percentage <= self.slg_end_percent:
continue
x, x_ip = block(x, x_ip=x_ip, **kwargs)
x, x_ip = block(x, x_ip=x_ip, **kwargs) #run block
if self.block_swap_debug:
compute_end = time.perf_counter()
compute_time = compute_end - compute_start