wan

def prepare_mask_latents(
self, mask, masked_image, batch_size, height, width, dtype, device, generator, do_classifier_free_guidance, noise_aug_strength
):
# resize the mask to latents shape as we concatenate the mask to the latents
# we do that before converting to dtype to avoid breaking in case we're using cpu_offload
# and half precision

if mask is not None:
mask = mask.to(device=device, dtype=self.vae.dtype)
bs = 1
new_mask = []
for i in range(0, mask.shape[0], bs):
mask_bs = mask[i : i + bs]
mask_bs = self.vae.encode(mask_bs)[0]
mask_bs = mask_bs.mode()
new_mask.append(mask_bs)
mask = torch.cat(new_mask, dim = 0)
# mask = mask * self.vae.config.scaling_factor

if masked_image is not None:
masked_image = masked_image.to(device=device, dtype=self.vae.dtype)
bs = 1
new_mask_pixel_values = []
for i in range(0, masked_image.shape[0], bs):
mask_pixel_values_bs = masked_image[i : i + bs]
mask_pixel_values_bs = self.vae.encode(mask_pixel_values_bs)[0]
mask_pixel_values_bs = mask_pixel_values_bs.mode()
new_mask_pixel_values.append(mask_pixel_values_bs)
masked_image_latents = torch.cat(new_mask_pixel_values, dim = 0)
# masked_image_latents = masked_image_latents * self.vae.config.scaling_factor
else:
masked_image_latents = None

return mask, masked_image_latents

if clip_image is not None:
clip_image = TF.to_tensor(clip_image).sub_(0.5).div_(0.5).to(device, weight_dtype)
clip_context = self.clip_image_encoder([clip_image[:, None, :, :]])
else:
clip_image = Image.new("RGB", (512, 512), color=(0, 0, 0))
clip_image = TF.to_tensor(clip_image).sub_(0.5).div_(0.5).to(device, weight_dtype)
clip_context = self.clip_image_encoder([clip_image[:, None, :, :]])
clip_context = torch.zeros_like(clip_context)

with torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=device):
noise_pred = self.transformer(
x=latent_model_input,
context=in_prompt_embeds,
t=timestep,
seq_len=seq_len,
y=y,
clip_fea=clip_context_input,
)

posted @ 2026-06-22 16:55  ruiw123  阅读(6)  评论(0)    收藏  举报