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:
+2
-2
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user