dtype fixes

This commit is contained in:
kijai
2025-09-30 18:33:05 +03:00
parent 1ba1a1662b
commit ed167a0135
+5 -5
View File
@@ -2251,7 +2251,7 @@ class WanModel(torch.nn.Module):
if ip_image.dim() == 6 and ip_image.shape[3] == 1:
ip_image = ip_image.squeeze(1)
ip_image_patch = self.original_patch_embedding(ip_image.float()).to(x.dtype)
ip_image_patch = self.original_patch_embedding(ip_image.float()).to(self.base_dtype)
f_ip, h_ip, w_ip = ip_image_patch.shape[2:]
x_ip = ip_image_patch.flatten(2).transpose(1, 2) # [B, N, D]
freq_offset = standin_input["freq_offset"]
@@ -2404,7 +2404,7 @@ class WanModel(torch.nn.Module):
# MultiTalk
if multitalk_audio is not None:
self.multitalk_audio_proj.to(self.main_device)
audio_cond = multitalk_audio.to(device=x.device, dtype=x.dtype)
audio_cond = multitalk_audio.to(device=x.device, dtype=self.base_dtype)
first_frame_audio_emb_s = audio_cond[:, :1, ...]
latter_frame_audio_emb = audio_cond[:, 1:, ...]
latter_frame_audio_emb = rearrange(latter_frame_audio_emb, "b (n_t n) w s c -> b n_t n w s c", n=4)
@@ -2418,7 +2418,7 @@ class WanModel(torch.nn.Module):
latter_frame_audio_emb_s = torch.concat([latter_first_frame_audio_emb, latter_middle_frame_audio_emb, latter_last_frame_audio_emb], dim=2)
multitalk_audio_embedding = self.multitalk_audio_proj(first_frame_audio_emb_s, latter_frame_audio_emb_s)
human_num = len(multitalk_audio_embedding)
multitalk_audio_embedding = torch.concat(multitalk_audio_embedding.split(1), dim=2).to(x.dtype)
multitalk_audio_embedding = torch.concat(multitalk_audio_embedding.split(1), dim=2).to(self.base_dtype)
self.multitalk_audio_proj.to(self.offload_device)
# convert ref_target_masks to token_ref_target_masks
@@ -2429,7 +2429,7 @@ class WanModel(torch.nn.Module):
token_ref_target_masks = token_ref_target_masks.squeeze(0)
token_ref_target_masks = (token_ref_target_masks > 0)
token_ref_target_masks = token_ref_target_masks.view(token_ref_target_masks.shape[0], -1)
token_ref_target_masks = token_ref_target_masks.to(x.dtype).to(device)
token_ref_target_masks = token_ref_target_masks.to(device, self.base_dtype)
humo_audio_input = None
if humo_audio is not None:
@@ -2629,7 +2629,7 @@ class WanModel(torch.nn.Module):
if (uni3c_data["start"] <= current_step_percentage <= uni3c_data["end"]) or \
(uni3c_data["end"] > 0 and current_step == 0 and current_step_percentage >= uni3c_data["start"]):
self.controlnet.to(self.main_device)
with torch.autocast(device_type=mm.get_autocast_device(device), dtype=x.dtype, enabled=True):
with torch.autocast(device_type=mm.get_autocast_device(device), dtype=self.base_dtype, enabled=True):
uni3c_controlnet_states = self.controlnet(
render_latent=render_latent.to(self.main_device, self.controlnet.dtype),
render_mask=uni3c_data["render_mask"],