diff --git a/nodes.py b/nodes.py index b66ce63..7c355ae 100644 --- a/nodes.py +++ b/nodes.py @@ -1927,8 +1927,8 @@ class WanVideoSampler: # UniAnimate if unianimate_poses is not None: - transformer.dwpose_embedding.to(device, model["dtype"]) - dwpose_data = unianimate_poses["pose"].to(device, model["dtype"]) + transformer.dwpose_embedding.to(device, dtype) + dwpose_data = unianimate_poses["pose"].to(device, dtype) dwpose_data = torch.cat([dwpose_data[:,:,:1].repeat(1,1,3,1,1), dwpose_data], dim=2) dwpose_data = transformer.dwpose_embedding(dwpose_data) log.info(f"UniAnimate pose embed shape: {dwpose_data.shape}") @@ -1941,19 +1941,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"], @@ -3001,9 +3001,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"], @@ -3301,9 +3300,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"], diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index ad2f642..8e4a003 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -1603,7 +1603,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 @@ -1978,8 +1978,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, @@ -2077,7 +2077,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