diff --git a/nodes.py b/nodes.py index a6e04de..e8ac3d2 100644 --- a/nodes.py +++ b/nodes.py @@ -642,6 +642,8 @@ class WanVideoModelLoader: total=param_count, leave=True): dtype_to_use = base_dtype if any(keyword in name for keyword in params_to_keep) else dtype + if "modulation" in name: + dtype_to_use = torch.float32 set_module_tensor_to_device(transformer, name, device=transformer_load_device, dtype=dtype_to_use, value=sd[name]) comfy_model.diffusion_model = transformer @@ -2388,7 +2390,6 @@ class WanVideoSampler: input_samples = samples["samples"].squeeze(0).to(noise) if input_samples.shape[1] != noise.shape[1]: input_samples = torch.cat([input_samples[:, :1].repeat(1, noise.shape[1] - input_samples.shape[1], 1, 1), input_samples], dim=1) - print("input_samples shape:", input_samples.shape) noise = noise * latent_timestep / 1000 + (1 - latent_timestep / 1000) * input_samples if samples is not None: @@ -2608,28 +2609,28 @@ class WanVideoSampler: pred_id=teacache_state[1] if teacache_state else None, **base_params ) - noise_pred_uncond=noise_pred_uncond[0].to(intermediate_device) - - #https://github.com/WeichenFan/CFG-Zero-star/ - if use_cfg_zero_star: - alpha = optimized_scale( - noise_pred_cond.view(batch_size, -1), - noise_pred_uncond.view(batch_size, -1) - ).view(batch_size, 1, 1, 1) - noise_pred = noise_pred_uncond * alpha + cfg_scale * (noise_pred_cond - noise_pred_uncond * alpha) - else: - noise_pred = noise_pred_uncond + cfg_scale * (noise_pred_cond - noise_pred_uncond) - return noise_pred, [teacache_state_cond, teacache_state_uncond] + noise_pred_uncond = noise_pred_uncond[0].to(intermediate_device) #batched else: + teacache_state_uncond = None [noise_pred_cond, noise_pred_uncond], teacache_state_cond = transformer( - [z] + [z], context= positive_embeds + negative_embeds, clip_fea=clip_fea, is_uncond=False, current_step_percentage=current_step_percentage, + [z] + [z], context=positive_embeds + negative_embeds, clip_fea=clip_fea, is_uncond=False, current_step_percentage=current_step_percentage, pred_id=teacache_state[0] if teacache_state else None, **base_params ) - noise_pred_uncond=noise_pred_uncond.to(intermediate_device) + #cfg - return noise_pred_uncond + cfg_scale * (noise_pred_cond - noise_pred_uncond), [teacache_state_cond] + #https://github.com/WeichenFan/CFG-Zero-star/ + if use_cfg_zero_star: + alpha = optimized_scale( + noise_pred_cond.view(batch_size, -1), + noise_pred_uncond.view(batch_size, -1) + ).view(batch_size, 1, 1, 1) + noise_pred = noise_pred_uncond * alpha + cfg_scale * (noise_pred_cond - noise_pred_uncond * alpha) + else: + noise_pred = noise_pred_uncond + cfg_scale * (noise_pred_cond - noise_pred_uncond) + + return noise_pred, [teacache_state_cond, teacache_state_uncond] log.info(f"Sampling {(latent_video_length-1) * 4 + 1} frames at {latent.shape[3]*8}x{latent.shape[2]*8} with {steps} steps") diff --git a/wanvideo/modules/attention.py b/wanvideo/modules/attention.py index ca966c8..12a6cd4 100644 --- a/wanvideo/modules/attention.py +++ b/wanvideo/modules/attention.py @@ -184,9 +184,9 @@ def attention( # ) attn_mask = None - q = q.transpose(1, 2).to(dtype) - k = k.transpose(1, 2).to(dtype) - v = v.transpose(1, 2).to(dtype) + q = q.transpose(1, 2)#.to(dtype) + k = k.transpose(1, 2)#.to(dtype) + v = v.transpose(1, 2)#.to(dtype) out = torch.nn.functional.scaled_dot_product_attention( q, k, v, attn_mask=attn_mask, is_causal=causal, dropout_p=dropout_p) @@ -196,9 +196,9 @@ def attention( elif attention_mode == 'sageattn': attn_mask = None - q = q.transpose(1, 2).to(dtype) - k = k.transpose(1, 2).to(dtype) - v = v.transpose(1, 2).to(dtype) + q = q.transpose(1, 2)#.to(dtype) + k = k.transpose(1, 2)#.to(dtype) + v = v.transpose(1, 2)#.to(dtype) out = sageattn_func( q, k, v, attn_mask=attn_mask, is_causal=causal, dropout_p=dropout_p) diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 7b6f27b..a6edda4 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -420,13 +420,11 @@ class WanAttentionBlock(nn.Module): grid_sizes(Tensor): Shape [B, 3], the second dimension contains (F, H, W) freqs(Tensor): Rope freqs, shape [1024, C / num_heads / 2] """ - assert e.dtype == torch.float32 - e = (self.modulation.to(torch.float32).to(e.device) + e.to(torch.float32)).chunk(6, dim=1) - assert e[0].dtype == torch.float32 + e = (self.modulation + e).chunk(6, dim=1) # self-attention y = self.self_attn( - self.norm1(x).float() * (1 + e[1]) + e[0], + self.norm1(x) * (1 + e[1]) + e[0], seq_lens, grid_sizes, freqs, rope_func=rope_func, seq_chunks=max(context.shape[0], clip_embed.shape[0] if clip_embed is not None else 0), @@ -438,7 +436,7 @@ class WanAttentionBlock(nn.Module): # cross-attention & ffn function def cross_attn_ffn(x, context, context_lens, e, clip_embed=None, grid_sizes=None): - if context.shape[0] > 1 or (clip_embed is not None and clip_embed.shape[0] > 1): + if (context.shape[0] > 1 or (clip_embed is not None and clip_embed.shape[0] > 1)) and x.shape[0] == 1: # Get number of prompts num_prompts = context.shape[0] num_clip_embeds = 0 if clip_embed is None else clip_embed.shape[0] @@ -496,7 +494,7 @@ class WanAttentionBlock(nn.Module): else: x = x + self.cross_attn(self.norm3(x), context, context_lens, clip_embed=clip_embed) y = self.ffn(self.norm2(x).float() * (1 + e[4]) + e[3]) - x = x.to(torch.float32) + (y.to(torch.float32) * e[5].to(torch.float32)) + x = x.to(torch.float32) + (y.to(torch.float32) * e[5]) return x x = cross_attn_ffn(x, context, context_lens, e, clip_embed=clip_embed, grid_sizes=grid_sizes) @@ -591,10 +589,9 @@ class Head(nn.Module): e(Tensor): Shape [B, C] """ assert e.dtype == torch.float32 - e_unsqueezed = e.unsqueeze(1).to(torch.float32) - e = (self.modulation.to(torch.float32).to(e.device) + e_unsqueezed).chunk(2, dim=1) - normed = self.norm(x).to(torch.float32) - x = self.head(normed * (1 + e[1].to(torch.float32)) + e[0].to(torch.float32)) + e = (self.modulation + e.unsqueeze(1)).chunk(2, dim=1) + normed = self.norm(x) + x = self.head(normed * (1 + e[1]) + e[0]) return x @@ -1029,7 +1026,7 @@ class WanModel(ModelMixin, ConfigMixin): previous_modulated_input = e.clone() if (self.teacache_use_coefficients and self.teacache_mode == 'e') else e0.clone() if not should_calc: - x += previous_residual.to(x.device) + x = x.to(previous_residual.dtype) + previous_residual.to(x.device) #log.info(f"TeaCache: Skipping uncond step {current_step+1}") self.teacache_state.update( pred_id, @@ -1063,11 +1060,11 @@ class WanModel(ModelMixin, ConfigMixin): if (data["start"] <= current_step_percentage <= data["end"]) or \ (data["end"] > 0 and current_step == 0 and current_step_percentage >= data["start"]): - vace_hints = self.forward_vace(x, data["context"], data["seq_len"], kwargs) + vace_hints = self.forward_vace(x.to(torch.float32), data["context"], data["seq_len"], kwargs) vace_hint_list.append(vace_hints) vace_scale_list.append(data["scale"]) else: - vace_hints = self.forward_vace(x, vace_data, seq_len, kwargs) + vace_hints = self.forward_vace(x.to(torch.float32), vace_data, seq_len, kwargs) vace_hint_list.append(vace_hints) vace_scale_list.append(1.0) @@ -1081,7 +1078,7 @@ class WanModel(ModelMixin, ConfigMixin): continue if b <= self.blocks_to_swap and self.blocks_to_swap >= 0: block.to(self.main_device) - x = block(x, **kwargs) + x = block(x.to(torch.float32), **kwargs) if b <= self.blocks_to_swap and self.blocks_to_swap >= 0: block.to(self.offload_device, non_blocking=self.use_non_blocking) @@ -1092,7 +1089,6 @@ class WanModel(ModelMixin, ConfigMixin): accumulated_rel_l1_distance=accumulated_rel_l1_distance.to(self.teacache_cache_device, non_blocking=self.use_non_blocking), previous_modulated_input=previous_modulated_input.to(self.teacache_cache_device, non_blocking=self.use_non_blocking) ) - x = self.head(x, e) x = self.unpatchify(x, grid_sizes) # type: ignore[arg-type] x = [u.float() for u in x]