Fix batched_cfg and some possible dtype tweaks to avoid recompiles

This commit is contained in:
kijai
2025-04-08 20:08:33 +03:00
parent 7de8ac25a6
commit f3f12aa6cf
3 changed files with 34 additions and 37 deletions
+17 -16
View File
@@ -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")
+6 -6
View File
@@ -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)
+11 -15
View File
@@ -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]