revert and fix again...

I'll clean this later
This commit is contained in:
kijai
2025-03-03 01:38:51 +02:00
parent 24fb1b42fe
commit 0183a03bf9
2 changed files with 9 additions and 9 deletions
+4 -5
View File
@@ -1266,7 +1266,6 @@ 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, :, :]
@@ -1276,7 +1275,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,
@@ -1286,7 +1285,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,
@@ -1323,14 +1322,14 @@ class WanVideoSampler:
t=timestep,
current_step=i,
is_uncond=False,
y=image_embeds.get("image_embeds", None),
y=[image_embeds["image_embeds"]] if "image_embeds" in image_embeds else [],
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["image_embeds"]] if "image_embeds" in image_embeds else [],
t=timestep,
current_step=i,
is_uncond=True,
+5 -4
View File
@@ -609,12 +609,13 @@ 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([x, y], dim=0)
if y:
x = [torch.cat([u, v], dim=0) for u, v in zip(x, y)]
# embeddings
x = [self.patch_embedding(u.unsqueeze(0)) for u in x]