From 24fb1b42fe2150c1e21f430c10f10cd7b94f81a3 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Mon, 3 Mar 2025 01:35:36 +0200 Subject: [PATCH] bug fixes --- nodes.py | 9 +++++---- wanvideo/modules/model.py | 7 +++---- 2 files changed, 8 insertions(+), 8 deletions(-) diff --git a/nodes.py b/nodes.py index edaf8d9..6b721ee 100644 --- a/nodes.py +++ b/nodes.py @@ -1266,6 +1266,7 @@ class WanVideoSampler: positive_prompt = text_embeds["prompt_embeds"][prompt_index] img_emb = image_embeds.get("image_embeds", None) + partial_img_emb = None if img_emb is not None: print("img_emb shape", img_emb.shape) partial_img_emb = img_emb[:, c, :, :] @@ -1275,7 +1276,7 @@ class WanVideoSampler: # Model inference - returns [frames, channels, height, width] noise_pred_cond = transformer( partial_latent_model_input, - y=[partial_img_emb], + y=partial_img_emb, t=timestep, current_step=i, is_uncond=False, @@ -1285,7 +1286,7 @@ class WanVideoSampler: if cfg[i] != 1.0: noise_pred_uncond = transformer( partial_latent_model_input, - y=[partial_img_emb], + y=partial_img_emb, t=timestep, current_step=i, is_uncond=True, @@ -1322,14 +1323,14 @@ class WanVideoSampler: t=timestep, current_step=i, is_uncond=False, - y=[image_embeds.get("image_embeds", None)], + y=image_embeds.get("image_embeds", None), context = [text_embeds["prompt_embeds"][0]], **args )[0].to(intermediate_device) if cfg[i] != 1.0: noise_pred_uncond = transformer( latent_model_input, - y=[image_embeds.get("image_embeds", None)], + y=image_embeds.get("image_embeds", None), t=timestep, current_step=i, is_uncond=True, diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index dce2a0e..a5e7bf5 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -609,13 +609,12 @@ class WanModel(ModelMixin, ConfigMixin): """ if self.model_type == 'i2v': assert clip_fea is not None and y is not None - # params - #device = self.patch_embedding.weight.device + if freqs.device != device: freqs = freqs.to(device) - + if y is not None: - x = [torch.cat([u, v], dim=0) for u, v in zip(x, y)] + x = torch.cat([x, y], dim=0) # embeddings x = [self.patch_embedding(u.unsqueeze(0)) for u in x]