Fix batched_cfg and some possible dtype tweaks to avoid recompiles
This commit is contained in:
@@ -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")
|
||||
|
||||
|
||||
@@ -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
@@ -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]
|
||||
|
||||
Reference in New Issue
Block a user