Fix device mismatch error with I2V + Lynx embeds

Fixed RuntimeError when using I2V embeds with Lynx embeds where tensors
were on different devices (cuda:0 and cpu).

The issue occurred because lynx_ref_text_embed["prompt_embeds"] were not
explicitly moved to the GPU device before being passed to the transformer
during Lynx reference buffer extraction.

Changes:
- Move lynx text embeddings to device in both conditional and unconditional
  buffer extraction calls (lines 1114 and 1128)
- Ensures all tensors are on the same device during cross-attention operations
This commit is contained in:
mossmatrix
2025-12-04 00:34:54 -05:00
parent b06c7d2d6d
commit 4e2bf0f1fa
+2 -2
View File
@@ -1111,7 +1111,7 @@ class WanVideoSampler:
lynx_ref_buffer = transformer(
[lynx_ref_input.to(device, dtype)],
torch.tensor([0], device=device),
lynx_ref_text_embed["prompt_embeds"],
[emb.to(device) for emb in lynx_ref_text_embed["prompt_embeds"]],
seq_len=math.ceil((lynx_ref_latent.shape[2] * lynx_ref_latent.shape[3]) / 4 * lynx_ref_latent.shape[1]),
lynx_embeds=lynx_embeds
)
@@ -1125,7 +1125,7 @@ class WanVideoSampler:
lynx_ref_buffer_uncond = transformer(
[lynx_ref_input_uncond.to(device, dtype)],
torch.tensor([0], device=device),
lynx_ref_text_embed["prompt_embeds"],
[emb.to(device) for emb in lynx_ref_text_embed["prompt_embeds"]],
seq_len=math.ceil((lynx_ref_latent.shape[2] * lynx_ref_latent.shape[3]) / 4 * lynx_ref_latent.shape[1]),
lynx_embeds=lynx_embeds,
is_uncond=True