Use fp32 norms for other models too and other fixes
This commit is contained in:
+3
-3
@@ -16,10 +16,10 @@ class WanRMSNorm(nn.Module):
|
|||||||
Args:
|
Args:
|
||||||
x(Tensor): Shape [B, L, C]
|
x(Tensor): Shape [B, L, C]
|
||||||
"""
|
"""
|
||||||
return self._norm(x.float()).type_as(x) * self.weight
|
return self._norm(x.to(self.weight.dtype)) * self.weight
|
||||||
|
|
||||||
def _norm(self, x):
|
def _norm(self, x):
|
||||||
return x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps)
|
return x * (torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps)).to(x.dtype)
|
||||||
|
|
||||||
|
|
||||||
class DummyAdapterLayer(nn.Module):
|
class DummyAdapterLayer(nn.Module):
|
||||||
@@ -81,7 +81,7 @@ class AudioProjModel(nn.Module):
|
|||||||
|
|
||||||
context_tokens = self.audio_proj_glob_3(audio_embeds).reshape(batch_size, self.context_tokens, self.output_dim)
|
context_tokens = self.audio_proj_glob_3(audio_embeds).reshape(batch_size, self.context_tokens, self.output_dim)
|
||||||
|
|
||||||
context_tokens = self.audio_proj_glob_norm(context_tokens)
|
context_tokens = self.audio_proj_glob_norm(context_tokens.to(self.audio_proj_glob_norm.layer.weight.dtype)).to(audio_embeds.dtype)
|
||||||
context_tokens = rearrange(context_tokens, "(bz f) m c -> bz f m c", f=video_length)
|
context_tokens = rearrange(context_tokens, "(bz f) m c -> bz f m c", f=video_length)
|
||||||
|
|
||||||
return context_tokens
|
return context_tokens
|
||||||
+1
-2
@@ -995,12 +995,11 @@ class WanVideoSampler:
|
|||||||
# Experimental args
|
# Experimental args
|
||||||
use_cfg_zero_star = use_tangential = use_fresca = bidirectional_sampling = use_tsr = False
|
use_cfg_zero_star = use_tangential = use_fresca = bidirectional_sampling = use_tsr = False
|
||||||
raag_alpha = 0.0
|
raag_alpha = 0.0
|
||||||
|
transformer.video_attention_split_steps = []
|
||||||
if experimental_args is not None:
|
if experimental_args is not None:
|
||||||
video_attention_split_steps = experimental_args.get("video_attention_split_steps", [])
|
video_attention_split_steps = experimental_args.get("video_attention_split_steps", [])
|
||||||
if video_attention_split_steps:
|
if video_attention_split_steps:
|
||||||
transformer.video_attention_split_steps = [int(x.strip()) for x in video_attention_split_steps.split(",")]
|
transformer.video_attention_split_steps = [int(x.strip()) for x in video_attention_split_steps.split(",")]
|
||||||
else:
|
|
||||||
transformer.video_attention_split_steps = []
|
|
||||||
|
|
||||||
use_zero_init = experimental_args.get("use_zero_init", True)
|
use_zero_init = experimental_args.get("use_zero_init", True)
|
||||||
use_cfg_zero_star = experimental_args.get("cfg_zero_star", False)
|
use_cfg_zero_star = experimental_args.get("cfg_zero_star", False)
|
||||||
|
|||||||
+66
-94
@@ -345,10 +345,10 @@ class WanRMSNorm(nn.Module):
|
|||||||
if use_chunked:
|
if use_chunked:
|
||||||
return self.forward_chunked(x, num_chunks)
|
return self.forward_chunked(x, num_chunks)
|
||||||
else:
|
else:
|
||||||
return self._norm(x.float()).type_as(x) * self.weight
|
return self._norm(x.to(self.weight.dtype)) * self.weight
|
||||||
|
|
||||||
def _norm(self, x):
|
def _norm(self, x):
|
||||||
return x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps)
|
return x * (torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps)).to(x.dtype)
|
||||||
|
|
||||||
def forward_chunked(self, x, num_chunks=4):
|
def forward_chunked(self, x, num_chunks=4):
|
||||||
output = torch.empty_like(x)
|
output = torch.empty_like(x)
|
||||||
@@ -398,23 +398,23 @@ class WanLayerNorm(nn.LayerNorm):
|
|||||||
def __init__(self, dim, eps=1e-6, elementwise_affine=False):
|
def __init__(self, dim, eps=1e-6, elementwise_affine=False):
|
||||||
super().__init__(dim, elementwise_affine=elementwise_affine, eps=eps)
|
super().__init__(dim, elementwise_affine=elementwise_affine, eps=eps)
|
||||||
|
|
||||||
def forward(self, inputs: torch.Tensor) -> torch.Tensor:
|
# def forward(self, inputs: torch.Tensor) -> torch.Tensor:
|
||||||
origin_dtype = inputs.dtype
|
# origin_dtype = inputs.dtype
|
||||||
out = F.layer_norm(
|
# out = F.layer_norm(
|
||||||
inputs.float(),
|
# inputs.float(),
|
||||||
self.normalized_shape,
|
# self.normalized_shape,
|
||||||
None if self.weight is None else self.weight.float(),
|
# None if self.weight is None else self.weight.float(),
|
||||||
None if self.bias is None else self.bias.float() ,
|
# None if self.bias is None else self.bias.float() ,
|
||||||
self.eps
|
# self.eps
|
||||||
).to(origin_dtype)
|
# ).to(origin_dtype)
|
||||||
return out
|
# return out
|
||||||
|
|
||||||
# def forward(self, x):
|
def forward(self, x):
|
||||||
# r"""
|
r"""
|
||||||
# Args:
|
Args:
|
||||||
# x(Tensor): Shape [B, L, C]
|
x(Tensor): Shape [B, L, C]
|
||||||
# """
|
"""
|
||||||
# return super().forward(x)
|
return super().forward(x)
|
||||||
|
|
||||||
|
|
||||||
class WanSelfAttention(nn.Module):
|
class WanSelfAttention(nn.Module):
|
||||||
@@ -465,8 +465,8 @@ class WanSelfAttention(nn.Module):
|
|||||||
|
|
||||||
def qkv_fn(self, x):
|
def qkv_fn(self, x):
|
||||||
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
|
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
|
||||||
q = self.norm_q(self.q(x)).view(b, s, n, d)
|
q = self.norm_q(self.q(x).to(self.norm_q.weight.dtype)).to(x.dtype).view(b, s, n, d)
|
||||||
k = self.norm_k(self.k(x)).view(b, s, n, d)
|
k = self.norm_k(self.k(x).to(self.norm_k.weight.dtype)).to(x.dtype).view(b, s, n, d)
|
||||||
v = self.v(x).view(b, s, n, d)
|
v = self.v(x).view(b, s, n, d)
|
||||||
return q, k, v
|
return q, k, v
|
||||||
|
|
||||||
@@ -481,8 +481,8 @@ class WanSelfAttention(nn.Module):
|
|||||||
|
|
||||||
def qkv_fn_ip(self, x):
|
def qkv_fn_ip(self, x):
|
||||||
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
|
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
|
||||||
q = self.norm_q(self.q(x) + self.q_loras(x)).view(b, s, n, d)
|
q = self.norm_q(self.q(x) + self.q_loras(x).to(self.norm_q.weight.dtype)).to(x.dtype).view(b, s, n, d)
|
||||||
k = self.norm_k(self.k(x) + self.k_loras(x)).view(b, s, n, d)
|
k = self.norm_k(self.k(x) + self.k_loras(x).to(self.norm_k.weight.dtype)).to(x.dtype).view(b, s, n, d)
|
||||||
v = (self.v(x) + self.v_loras(x)).view(b, s, n, d)
|
v = (self.v(x) + self.v_loras(x)).view(b, s, n, d)
|
||||||
return q, k, v
|
return q, k, v
|
||||||
|
|
||||||
@@ -549,7 +549,7 @@ class WanSelfAttention(nn.Module):
|
|||||||
return x, x_ref_attn_map
|
return x, x_ref_attn_map
|
||||||
|
|
||||||
|
|
||||||
def forward_split(self, q, k, v, seq_lens, grid_sizes, freqs, seq_chunks=1,current_step=0, video_attention_split_steps = []):
|
def forward_split(self, q, k, v, seq_lens, grid_sizes, seq_chunks):
|
||||||
r"""
|
r"""
|
||||||
Args:
|
Args:
|
||||||
x(Tensor): Shape [B, L, num_heads, C / num_heads]
|
x(Tensor): Shape [B, L, num_heads, C / num_heads]
|
||||||
@@ -559,62 +559,36 @@ class WanSelfAttention(nn.Module):
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
# Split by frames if multiple prompts are provided
|
# Split by frames if multiple prompts are provided
|
||||||
if seq_chunks > 1 and current_step in video_attention_split_steps:
|
frames, height, width = grid_sizes[0]
|
||||||
outputs = []
|
|
||||||
# Extract frame, height, width from grid_sizes
|
|
||||||
frames = grid_sizes[0][0]
|
|
||||||
height = grid_sizes[0][1]
|
|
||||||
width = grid_sizes[0][2]
|
|
||||||
tokens_per_frame = height * width
|
tokens_per_frame = height * width
|
||||||
|
|
||||||
actual_chunks = torch.min(torch.tensor(seq_chunks, device=frames.device), frames)
|
seq_chunks_tensor = torch.tensor(seq_chunks, device=q.device, dtype=frames.dtype)
|
||||||
|
actual_chunks = torch.minimum(seq_chunks_tensor, frames)
|
||||||
base_frames_per_chunk = frames // actual_chunks
|
base_frames_per_chunk = frames // actual_chunks
|
||||||
extra_frames = frames % actual_chunks
|
extra_frames = frames % actual_chunks
|
||||||
|
|
||||||
# Calculate all chunk boundaries
|
chunk_indices = torch.arange(actual_chunks, device=q.device)
|
||||||
chunk_indices = torch.arange(actual_chunks, device=frames.device)
|
chunk_sizes = base_frames_per_chunk + (chunk_indices < extra_frames)
|
||||||
chunk_sizes = base_frames_per_chunk + (chunk_indices < extra_frames).long()
|
chunk_starts = torch.cumsum(torch.cat([torch.zeros(1, device=q.device, dtype=torch.long), chunk_sizes[:-1]]), dim=0)
|
||||||
chunk_starts = torch.cumsum(torch.cat([torch.zeros(1, device=frames.device), chunk_sizes[:-1]]), dim=0).long()
|
|
||||||
chunk_ends = chunk_starts + chunk_sizes
|
chunk_ends = chunk_starts + chunk_sizes
|
||||||
|
|
||||||
# Process each chunk using tensor indexing
|
outputs = []
|
||||||
for i in range(actual_chunks.item()):
|
for i in chunk_indices:
|
||||||
start_frame = chunk_starts[i]
|
start_idx = chunk_starts[i] * tokens_per_frame
|
||||||
end_frame = chunk_ends[i]
|
end_idx = chunk_ends[i] * tokens_per_frame
|
||||||
|
|
||||||
# Convert to token indices using tensor operations
|
|
||||||
start_idx = start_frame * tokens_per_frame
|
|
||||||
end_idx = end_frame * tokens_per_frame
|
|
||||||
|
|
||||||
chunk_q = q[:, start_idx:end_idx, :, :]
|
|
||||||
chunk_k = k[:, start_idx:end_idx, :, :]
|
|
||||||
chunk_v = v[:, start_idx:end_idx, :, :]
|
|
||||||
|
|
||||||
chunk_out = attention(
|
chunk_out = attention(
|
||||||
q=chunk_q,
|
q[:, start_idx:end_idx, :, :],
|
||||||
k=chunk_k,
|
k[:, start_idx:end_idx, :, :],
|
||||||
v=chunk_v,
|
v[:, start_idx:end_idx, :, :],
|
||||||
k_lens=seq_lens,
|
k_lens=seq_lens,
|
||||||
attention_mode=self.attention_mode)
|
attention_mode=self.attention_mode
|
||||||
|
)
|
||||||
outputs.append(chunk_out)
|
outputs.append(chunk_out)
|
||||||
|
|
||||||
# Concatenate outputs along the sequence dimension
|
|
||||||
x = torch.cat(outputs, dim=1)
|
x = torch.cat(outputs, dim=1)
|
||||||
else:
|
|
||||||
# Original attention computation
|
|
||||||
x = attention(
|
|
||||||
q=q,
|
|
||||||
k=k,
|
|
||||||
v=v,
|
|
||||||
k_lens=seq_lens,
|
|
||||||
attention_mode=self.attention_mode)
|
|
||||||
|
|
||||||
# output
|
# output
|
||||||
x = x.flatten(2)
|
return self.o(x.flatten(2))
|
||||||
x = self.o(x)
|
|
||||||
|
|
||||||
return x
|
|
||||||
|
|
||||||
def normalized_attention_guidance(self, b, n, d, q, context, nag_context=None, nag_params={}):
|
def normalized_attention_guidance(self, b, n, d, q, context, nag_context=None, nag_params={}):
|
||||||
# NAG text attention
|
# NAG text attention
|
||||||
@@ -624,9 +598,9 @@ class WanSelfAttention(nn.Module):
|
|||||||
nag_alpha = nag_params['nag_alpha']
|
nag_alpha = nag_params['nag_alpha']
|
||||||
nag_tau = nag_params['nag_tau']
|
nag_tau = nag_params['nag_tau']
|
||||||
|
|
||||||
k_positive = self.norm_k(self.k(context_positive)).view(b, -1, n, d)
|
k_positive = self.norm_k(self.k(context_positive).to(self.norm_k.weight.dtype)).view(b, -1, n, d).to(q.dtype)
|
||||||
v_positive = self.v(context_positive).view(b, -1, n, d)
|
v_positive = self.v(context_positive).view(b, -1, n, d)
|
||||||
k_negative = self.norm_k(self.k(context_negative)).view(b, -1, n, d)
|
k_negative = self.norm_k(self.k(context_negative).to(self.norm_k.weight.dtype)).view(b, -1, n, d).to(q.dtype)
|
||||||
v_negative = self.v(context_negative).view(b, -1, n, d)
|
v_negative = self.v(context_negative).view(b, -1, n, d)
|
||||||
|
|
||||||
x_positive = attention(q, k_positive, v_positive, attention_mode=self.attention_mode)
|
x_positive = attention(q, k_positive, v_positive, attention_mode=self.attention_mode)
|
||||||
@@ -703,15 +677,15 @@ class WanT2VCrossAttention(WanSelfAttention):
|
|||||||
x = x[:, num_cond_latents_thw:]
|
x = x[:, num_cond_latents_thw:]
|
||||||
q = self.norm_q(self.q(x).view(b, -1, n, d))
|
q = self.norm_q(self.q(x).view(b, -1, n, d))
|
||||||
else:
|
else:
|
||||||
q = self.norm_q(self.q(x),num_chunks=2 if rope_func == "comfy_chunked" else 1).view(b, -1, n, d)
|
q = self.norm_q(self.q(x).to(self.norm_q.weight.dtype),num_chunks=2 if rope_func == "comfy_chunked" else 1).to(x.dtype).view(b, -1, n, d)
|
||||||
|
|
||||||
if nag_context is not None and not is_uncond:
|
if nag_context is not None and not is_uncond:
|
||||||
x = self.normalized_attention_guidance(b, n, d, q, context, nag_context, nag_params)
|
x = self.normalized_attention_guidance(b, n, d, q, context, nag_context, nag_params)
|
||||||
else:
|
else:
|
||||||
if is_longcat:
|
if is_longcat:
|
||||||
k = self.norm_k(self.k(context).view(b, -1, n, d))
|
k = self.norm_k(self.k(context).to(self.norm_k.weight.dtype)).to(x.dtype).view(b, -1, n, d)
|
||||||
else:
|
else:
|
||||||
k = self.norm_k(self.k(context)).view(b, -1, n, d)
|
k = self.norm_k(self.k(context).to(self.norm_k.weight.dtype)).to(x.dtype).view(b, -1, n, d)
|
||||||
|
|
||||||
v = self.v(context).view(b, -1, n, d)
|
v = self.v(context).view(b, -1, n, d)
|
||||||
|
|
||||||
@@ -794,19 +768,19 @@ class WanI2VCrossAttention(WanSelfAttention):
|
|||||||
"""
|
"""
|
||||||
b, n, d = x.size(0), self.num_heads, self.head_dim
|
b, n, d = x.size(0), self.num_heads, self.head_dim
|
||||||
# compute query
|
# compute query
|
||||||
q = self.norm_q(self.q(x),num_chunks=2 if rope_func == "comfy_chunked" else 1).view(b, -1, n, d)
|
q = self.norm_q(self.q(x).to(self.norm_q.weight.dtype),num_chunks=2 if rope_func == "comfy_chunked" else 1).view(b, -1, n, d).to(x.dtype)
|
||||||
|
|
||||||
if nag_context is not None and not is_uncond:
|
if nag_context is not None and not is_uncond:
|
||||||
x_text = self.normalized_attention_guidance(b, n, d, q, context, nag_context, nag_params)
|
x_text = self.normalized_attention_guidance(b, n, d, q, context, nag_context, nag_params)
|
||||||
else:
|
else:
|
||||||
# text attention
|
# text attention
|
||||||
k = self.norm_k(self.k(context)).view(b, -1, n, d)
|
k = self.norm_k(self.k(context).to(self.norm_k.weight.dtype)).view(b, -1, n, d).to(x.dtype)
|
||||||
v = self.v(context).view(b, -1, n, d)
|
v = self.v(context).view(b, -1, n, d)
|
||||||
x_text = attention(q, k, v, attention_mode=self.attention_mode).flatten(2)
|
x_text = attention(q, k, v, attention_mode=self.attention_mode).flatten(2)
|
||||||
|
|
||||||
#img attention
|
#img attention
|
||||||
if clip_embed is not None:
|
if clip_embed is not None:
|
||||||
k_img = self.norm_k_img(self.k_img(clip_embed)).view(b, -1, n, d)
|
k_img = self.norm_k_img(self.k_img(clip_embed).to(self.norm_k_img.weight.dtype)).view(b, -1, n, d).to(x.dtype)
|
||||||
v_img = self.v_img(clip_embed).view(b, -1, n, d)
|
v_img = self.v_img(clip_embed).view(b, -1, n, d)
|
||||||
img_x = attention(q, k_img, v_img, attention_mode=self.attention_mode).flatten(2)
|
img_x = attention(q, k_img, v_img, attention_mode=self.attention_mode).flatten(2)
|
||||||
x = x_text + img_x
|
x = x_text + img_x
|
||||||
@@ -1068,8 +1042,9 @@ class WanAttentionBlock(nn.Module):
|
|||||||
def forward(
|
def forward(
|
||||||
self, x, e, seq_lens, grid_sizes, freqs, context, current_step,
|
self, x, e, seq_lens, grid_sizes, freqs, context, current_step,
|
||||||
last_step=False,
|
last_step=False,
|
||||||
video_attention_split_steps=[],
|
|
||||||
clip_embed=None,
|
clip_embed=None,
|
||||||
|
seq_chunks=0, #comfy chunked cross-attn
|
||||||
|
chunked_self_attention=False,
|
||||||
camera_embed=None, #ReCamMaster
|
camera_embed=None, #ReCamMaster
|
||||||
audio_proj=None, audio_scale=1.0, #fantasytalking
|
audio_proj=None, audio_scale=1.0, #fantasytalking
|
||||||
num_latent_frames=21,
|
num_latent_frames=21,
|
||||||
@@ -1111,9 +1086,9 @@ class WanAttentionBlock(nn.Module):
|
|||||||
T = num_latent_frames
|
T = num_latent_frames
|
||||||
is_longcat = C == 4096
|
is_longcat = C == 4096
|
||||||
if is_longcat:
|
if is_longcat:
|
||||||
input_x = self.modulate(self.norm1(x.view(B, T, -1, C).float()), shift_msa, scale_msa, seg_idx=self.seg_idx).to(input_dtype).view(B, N, C)
|
input_x = self.modulate(self.norm1(x.view(B, T, -1, C).to(shift_msa.dtype)), shift_msa, scale_msa, seg_idx=self.seg_idx).to(input_dtype).view(B, N, C)
|
||||||
else:
|
else:
|
||||||
input_x = self.modulate(self.norm1(x), shift_msa, scale_msa, seg_idx=self.seg_idx)
|
input_x = self.modulate(self.norm1(x.to(shift_msa.dtype)), shift_msa, scale_msa, seg_idx=self.seg_idx).to(input_dtype)
|
||||||
|
|
||||||
del shift_msa, scale_msa
|
del shift_msa, scale_msa
|
||||||
|
|
||||||
@@ -1199,14 +1174,8 @@ class WanAttentionBlock(nn.Module):
|
|||||||
and inner_t is None
|
and inner_t is None
|
||||||
and x_ip is None # Don't split when using IP-Adapter
|
and x_ip is None # Don't split when using IP-Adapter
|
||||||
)
|
)
|
||||||
if split_attn:
|
if split_attn and chunked_self_attention:
|
||||||
y = self.self_attn.forward_split(
|
y = self.self_attn.forward_split(q, k, v, seq_lens, grid_sizes, seq_chunks)
|
||||||
q, k, v,
|
|
||||||
seq_lens, grid_sizes, freqs,
|
|
||||||
seq_chunks=max(context.shape[0], clip_embed.shape[0] if clip_embed is not None else 0),
|
|
||||||
current_step=current_step,
|
|
||||||
video_attention_split_steps=video_attention_split_steps
|
|
||||||
)
|
|
||||||
elif ref_target_masks is not None: #multi/infinite talk
|
elif ref_target_masks is not None: #multi/infinite talk
|
||||||
y, x_ref_attn_map = self.self_attn.forward_multitalk(q, k, v, seq_lens, grid_sizes, ref_target_masks)
|
y, x_ref_attn_map = self.self_attn.forward_multitalk(q, k, v, seq_lens, grid_sizes, ref_target_masks)
|
||||||
elif self.attention_mode == "radial_sage_attention":
|
elif self.attention_mode == "radial_sage_attention":
|
||||||
@@ -1301,14 +1270,13 @@ class WanAttentionBlock(nn.Module):
|
|||||||
target_seq_lens=seq_lens_ovi,
|
target_seq_lens=seq_lens_ovi,
|
||||||
target_grid_sizes=grid_sizes_ovi,
|
target_grid_sizes=grid_sizes_ovi,
|
||||||
target_freqs=freqs_ovi)
|
target_freqs=freqs_ovi)
|
||||||
y = self.ffn(torch.addcmul(shift_mlp, self.norm2(x), 1 + scale_mlp))
|
|
||||||
x = x.addcmul(y, gate_mlp)
|
|
||||||
elif split_attn:
|
elif split_attn:
|
||||||
if nag_context is not None:
|
if nag_context is not None:
|
||||||
raise NotImplementedError("nag_context is not supported in split_cross_attn_ffn")
|
raise NotImplementedError("nag_context is not supported in split_cross_attn_ffn")
|
||||||
x = self.split_cross_attn_ffn(x, context, shift_mlp, scale_mlp, gate_mlp, clip_embed, grid_sizes)
|
x = self.split_cross_attn_ffn(x, context, shift_mlp, scale_mlp, gate_mlp, clip_embed, grid_sizes)
|
||||||
|
return x, x_ip, lynx_ref_feature, x_ovi
|
||||||
else:
|
else:
|
||||||
x = x + self.cross_attn(self.norm3(x), context, grid_sizes, clip_embed=clip_embed, audio_proj=audio_proj, audio_scale=audio_scale,
|
x = x + self.cross_attn(self.norm3(x.to(self.norm3.weight.dtype)).to(input_dtype), context, grid_sizes, clip_embed=clip_embed, audio_proj=audio_proj, audio_scale=audio_scale,
|
||||||
num_latent_frames=num_latent_frames, nag_params=nag_params, nag_context=nag_context, is_uncond=is_uncond,
|
num_latent_frames=num_latent_frames, nag_params=nag_params, nag_context=nag_context, is_uncond=is_uncond,
|
||||||
rope_func=self.rope_func, inner_t=inner_t, inner_c=inner_c, cross_freqs=cross_freqs,
|
rope_func=self.rope_func, inner_t=inner_t, inner_c=inner_c, cross_freqs=cross_freqs,
|
||||||
adapter_proj=adapter_proj, ip_scale=ip_scale, orig_seq_len=original_seq_len, lynx_x_ip=lynx_x_ip, lynx_ip_scale=lynx_ip_scale, num_cond_latents=num_cond_latents)
|
adapter_proj=adapter_proj, ip_scale=ip_scale, orig_seq_len=original_seq_len, lynx_x_ip=lynx_x_ip, lynx_ip_scale=lynx_ip_scale, num_cond_latents=num_cond_latents)
|
||||||
@@ -1343,10 +1311,10 @@ class WanAttentionBlock(nn.Module):
|
|||||||
x_ffn = self.ffn(norm2_x)
|
x_ffn = self.ffn(norm2_x)
|
||||||
else:
|
else:
|
||||||
if not is_longcat:
|
if not is_longcat:
|
||||||
mod_x = torch.addcmul(shift_mlp, self.norm2(x), 1 + scale_mlp)
|
mod_x = torch.addcmul(shift_mlp, self.norm2(x.to(shift_mlp.dtype)), 1 + scale_mlp)
|
||||||
else:
|
else:
|
||||||
mod_x = torch.addcmul(shift_mlp, self.norm2(x.view(B, -1, N//T, C).float()), 1 + scale_mlp).view(B, -1, C).to(input_dtype)
|
mod_x = torch.addcmul(shift_mlp, self.norm2(x.view(B, -1, N//T, C).float()), 1 + scale_mlp).view(B, -1, C)
|
||||||
x_ffn = self.ffn(mod_x)
|
x_ffn = self.ffn(mod_x.to(input_dtype))
|
||||||
del shift_mlp, scale_mlp
|
del shift_mlp, scale_mlp
|
||||||
|
|
||||||
# gate_mlp
|
# gate_mlp
|
||||||
@@ -1358,7 +1326,7 @@ class WanAttentionBlock(nn.Module):
|
|||||||
x = x.add(x_ffn)
|
x = x.add(x_ffn)
|
||||||
else:
|
else:
|
||||||
if not is_longcat:
|
if not is_longcat:
|
||||||
x = x.addcmul(x_ffn, gate_mlp)
|
x = x.addcmul(x_ffn.to(gate_mlp.dtype), gate_mlp).to(input_dtype)
|
||||||
else:
|
else:
|
||||||
x = x + (gate_mlp * x_ffn.view(B, -1, N//T, C).float()).to(input_dtype).view(B, -1, C)
|
x = x + (gate_mlp * x_ffn.view(B, -1, N//T, C).float()).to(input_dtype).view(B, -1, C)
|
||||||
del gate_mlp
|
del gate_mlp
|
||||||
@@ -1407,7 +1375,7 @@ class WanAttentionBlock(nn.Module):
|
|||||||
segment_clip_embed = clip_embed[clip_idx:clip_idx+1]
|
segment_clip_embed = clip_embed[clip_idx:clip_idx+1]
|
||||||
|
|
||||||
# Get tensor segment
|
# Get tensor segment
|
||||||
x_segment = x[:, segment_indices, :]
|
x_segment = x[:, segment_indices, :].to(self.norm3.weight.dtype)
|
||||||
|
|
||||||
# Process segment with its prompt and clip embedding
|
# Process segment with its prompt and clip embedding
|
||||||
processed_segment = self.cross_attn(self.norm3(x_segment), segment_context, clip_embed=segment_clip_embed)
|
processed_segment = self.cross_attn(self.norm3(x_segment), segment_context, clip_embed=segment_clip_embed)
|
||||||
@@ -1520,7 +1488,7 @@ class Head(nn.Module):
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
e = self.get_mod(e.to(x.device))
|
e = self.get_mod(e.to(x.device))
|
||||||
x = self.head(self.norm(x).mul_(1 + e[1]).add_(e[0]))
|
x = self.head(self.norm(x.float()).to(x.dtype).mul_(1 + e[1]).add_(e[0]))
|
||||||
return x
|
return x
|
||||||
|
|
||||||
class Head_adaLN(nn.Module):
|
class Head_adaLN(nn.Module):
|
||||||
@@ -2844,6 +2812,9 @@ class WanModel(torch.nn.Module):
|
|||||||
if unianim_data['start_percent'] <= current_step_percentage <= unianim_data['end_percent']:
|
if unianim_data['start_percent'] <= current_step_percentage <= unianim_data['end_percent']:
|
||||||
dwpose_emb = rearrange(unianim_data['dwpose'], 'b c f h w -> b (f h w) c').contiguous()
|
dwpose_emb = rearrange(unianim_data['dwpose'], 'b c f h w -> b (f h w) c').contiguous()
|
||||||
x.add_(dwpose_emb, alpha=unianim_data['strength'])
|
x.add_(dwpose_emb, alpha=unianim_data['strength'])
|
||||||
|
|
||||||
|
seq_chunks = max(context.shape[0], clip_embed.shape[0] if clip_embed is not None else 0)
|
||||||
|
chunked_self_attention = seq_chunks > 1 and current_step in self.video_attention_split_steps
|
||||||
# arguments
|
# arguments
|
||||||
kwargs = dict(
|
kwargs = dict(
|
||||||
e=e0,
|
e=e0,
|
||||||
@@ -2854,7 +2825,8 @@ class WanModel(torch.nn.Module):
|
|||||||
clip_embed=clip_embed,
|
clip_embed=clip_embed,
|
||||||
current_step=torch.tensor(current_step),
|
current_step=torch.tensor(current_step),
|
||||||
last_step=torch.tensor(last_step, dtype=torch.bool),
|
last_step=torch.tensor(last_step, dtype=torch.bool),
|
||||||
video_attention_split_steps=self.video_attention_split_steps,
|
chunked_self_attention=chunked_self_attention,
|
||||||
|
seq_chunks=seq_chunks,
|
||||||
camera_embed=camera_embed,
|
camera_embed=camera_embed,
|
||||||
audio_proj=audio_proj,
|
audio_proj=audio_proj,
|
||||||
num_latent_frames = F,
|
num_latent_frames = F,
|
||||||
|
|||||||
Reference in New Issue
Block a user