Merge branch 'main' into steadydancer
This commit is contained in:
@@ -96,6 +96,7 @@ class CustomLinear(nn.Linear):
|
||||
|
||||
if not allow_compile:
|
||||
self._get_weight_with_lora = torch.compiler.disable()(self._get_weight_with_lora)
|
||||
self.forward = torch.compiler.disable()(self.forward)
|
||||
|
||||
def set_lora_diffs(self, lora_diffs, device=torch.device("cpu")):
|
||||
self.lora_diffs = []
|
||||
|
||||
+106
-131
@@ -23,7 +23,8 @@ from ...multitalk.multitalk import get_attn_map_with_target
|
||||
from ...echoshot.echoshot import rope_apply_z, rope_apply_c, rope_apply_echoshot
|
||||
|
||||
from ...MTV.mtv import apply_rotary_emb
|
||||
|
||||
from comfy.ldm.flux.math import apply_rope1 as apply_rope_comfy1
|
||||
from comfy.ldm.flux.math import apply_rope as apply_rope_comfy
|
||||
from comfy import model_management as mm
|
||||
|
||||
__all__ = ['WanModel']
|
||||
@@ -120,70 +121,6 @@ def torch_dfs(model: nn.Module, parent_name='root'):
|
||||
modules += child_modules
|
||||
return modules, module_names
|
||||
|
||||
#from comfy.ldm.flux.math import apply_rope as apply_rope_comfy
|
||||
def apply_rope_comfy(xq, xk, freqs_cis):
|
||||
xq_ = xq.to(dtype=freqs_cis.dtype).reshape(*xq.shape[:-1], -1, 1, 2)
|
||||
xk_ = xk.to(dtype=freqs_cis.dtype).reshape(*xk.shape[:-1], -1, 1, 2)
|
||||
xq_out = freqs_cis[..., 0] * xq_[..., 0] + freqs_cis[..., 1] * xq_[..., 1]
|
||||
xk_out = freqs_cis[..., 0] * xk_[..., 0] + freqs_cis[..., 1] * xk_[..., 1]
|
||||
return xq_out.reshape(*xq.shape).type_as(xq), xk_out.reshape(*xk.shape).type_as(xk)
|
||||
|
||||
def apply_rope_comfy_chunked(xq, xk, freqs_cis, num_chunks=4):
|
||||
seq_dim = 1
|
||||
|
||||
# Initialize output tensors
|
||||
xq_out = torch.empty_like(xq)
|
||||
xk_out = torch.empty_like(xk)
|
||||
|
||||
# Calculate chunks
|
||||
seq_len = xq.shape[seq_dim]
|
||||
chunk_sizes = [seq_len // num_chunks + (1 if i < seq_len % num_chunks else 0)
|
||||
for i in range(num_chunks)]
|
||||
|
||||
# First pass: process xq completely
|
||||
start_idx = 0
|
||||
for size in chunk_sizes:
|
||||
end_idx = start_idx + size
|
||||
|
||||
slices = [slice(None)] * len(xq.shape)
|
||||
slices[seq_dim] = slice(start_idx, end_idx)
|
||||
|
||||
freq_slices = [slice(None)] * len(freqs_cis.shape)
|
||||
if seq_dim < len(freqs_cis.shape):
|
||||
freq_slices[seq_dim] = slice(start_idx, end_idx)
|
||||
freqs_chunk = freqs_cis[tuple(freq_slices)]
|
||||
|
||||
xq_chunk = xq[tuple(slices)]
|
||||
xq_chunk_ = xq_chunk.to(dtype=freqs_cis.dtype).reshape(*xq_chunk.shape[:-1], -1, 1, 2)
|
||||
xq_out[tuple(slices)] = (freqs_chunk[..., 0] * xq_chunk_[..., 0] +
|
||||
freqs_chunk[..., 1] * xq_chunk_[..., 1]).reshape(*xq_chunk.shape).type_as(xq)
|
||||
|
||||
del xq_chunk, xq_chunk_, freqs_chunk
|
||||
start_idx = end_idx
|
||||
|
||||
# Second pass: process xk completely
|
||||
start_idx = 0
|
||||
for size in chunk_sizes:
|
||||
end_idx = start_idx + size
|
||||
|
||||
slices = [slice(None)] * len(xk.shape)
|
||||
slices[seq_dim] = slice(start_idx, end_idx)
|
||||
|
||||
freq_slices = [slice(None)] * len(freqs_cis.shape)
|
||||
if seq_dim < len(freqs_cis.shape):
|
||||
freq_slices[seq_dim] = slice(start_idx, end_idx)
|
||||
freqs_chunk = freqs_cis[tuple(freq_slices)]
|
||||
|
||||
xk_chunk = xk[tuple(slices)]
|
||||
xk_chunk_ = xk_chunk.to(dtype=freqs_cis.dtype).reshape(*xk_chunk.shape[:-1], -1, 1, 2)
|
||||
xk_out[tuple(slices)] = (freqs_chunk[..., 0] * xk_chunk_[..., 0] +
|
||||
freqs_chunk[..., 1] * xk_chunk_[..., 1]).reshape(*xk_chunk.shape).type_as(xk)
|
||||
|
||||
del xk_chunk, xk_chunk_, freqs_chunk
|
||||
start_idx = end_idx
|
||||
|
||||
return xq_out, xk_out
|
||||
|
||||
def rope_riflex(pos, dim, i, theta, L_test, k, ntk_factor=1.0):
|
||||
assert dim % 2 == 0
|
||||
if mm.is_device_mps(pos.device) or mm.is_intel_xpu() or mm.is_directml_enabled():
|
||||
@@ -462,21 +399,62 @@ class WanSelfAttention(nn.Module):
|
||||
self.norm_q = WanRMSNorm(norm_dim, eps=eps) if qk_norm else nn.Identity()
|
||||
self.norm_k = WanRMSNorm(norm_dim, eps=eps) if qk_norm else nn.Identity()
|
||||
|
||||
def qkv_fn(self, x):
|
||||
def qkv_fn(self, x, is_longcat=False):
|
||||
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
|
||||
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).to(self.norm_k.weight.dtype)).to(x.dtype).view(b, s, n, d)
|
||||
if is_longcat:
|
||||
q = self.q(x).view(b, s, n, d)
|
||||
q = self.norm_q(q.float()).to(x.dtype)
|
||||
k = self.k(x).view(b, s, n, d)
|
||||
k = self.norm_k(k.float()).to(x.dtype)
|
||||
else:
|
||||
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).to(self.norm_k.weight.dtype)).to(x.dtype).view(b, s, n, d)
|
||||
v = self.v(x).view(b, s, n, d)
|
||||
return q, k, v
|
||||
|
||||
def qkv_fn_longcat(self, x):
|
||||
|
||||
def _qkv_fn_with_rope(self, x, linear_layer, norm_layer, freqs, num_chunks=1, is_longcat=False):
|
||||
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
|
||||
q = self.q(x).view(b, s, n, d)
|
||||
q = self.norm_q(q.float()).to(x.dtype)
|
||||
k = self.k(x).view(b, s, n, d)
|
||||
k = self.norm_k(k.float()).to(x.dtype)
|
||||
v = self.v(x).view(b, s, n, d)
|
||||
return q, k, v
|
||||
|
||||
use_chunked = num_chunks > 1
|
||||
if use_chunked:
|
||||
chunk_sizes = [s // num_chunks + (1 if i < s % num_chunks else 0)
|
||||
for i in range(num_chunks)]
|
||||
|
||||
out = torch.empty(b, s, n, d, dtype=x.dtype, device=x.device)
|
||||
start_idx = 0
|
||||
for size in chunk_sizes:
|
||||
end_idx = start_idx + size
|
||||
|
||||
x_chunk = x[:, start_idx:end_idx]
|
||||
|
||||
if is_longcat:
|
||||
chunk = linear_layer(x_chunk).view(b, size, n, d)
|
||||
chunk = norm_layer(chunk.float()).to(x.dtype)
|
||||
else:
|
||||
chunk = norm_layer(linear_layer(x_chunk).to(norm_layer.weight.dtype)).to(x.dtype).view(b, size, n, d)
|
||||
|
||||
freqs_chunk = freqs[:, start_idx:end_idx] if freqs.shape[1] > 1 else freqs
|
||||
out[:, start_idx:end_idx] = apply_rope_comfy1(chunk, freqs_chunk)
|
||||
|
||||
start_idx = end_idx
|
||||
return out
|
||||
else:
|
||||
if is_longcat:
|
||||
result = linear_layer(x).view(b, s, n, d)
|
||||
result = norm_layer(result.float()).to(x.dtype)
|
||||
else:
|
||||
result = norm_layer(linear_layer(x).to(norm_layer.weight.dtype)).to(x.dtype).view(b, s, n, d)
|
||||
return apply_rope_comfy1(result, freqs)
|
||||
|
||||
def qkv_fn_q_with_rope(self, x, freqs, num_chunks=1, is_longcat=False):
|
||||
return self._qkv_fn_with_rope(x, self.q, self.norm_q, freqs, num_chunks, is_longcat)
|
||||
|
||||
def qkv_fn_k_with_rope(self, x, freqs, num_chunks=1, is_longcat=False):
|
||||
return self._qkv_fn_with_rope(x, self.k, self.norm_k, freqs, num_chunks, is_longcat)
|
||||
|
||||
def qkv_fn_v(self, x):
|
||||
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
|
||||
return self.v(x).view(b, s, n, d)
|
||||
|
||||
def qkv_fn_ip(self, x):
|
||||
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
|
||||
@@ -662,7 +640,7 @@ class WanT2VCrossAttention(WanSelfAttention):
|
||||
self.k_fusion = None
|
||||
|
||||
def forward(self, x, context, grid_sizes=None, clip_embed=None, audio_proj=None, audio_scale=1.0,
|
||||
num_latent_frames=21, nag_params={}, nag_context=None, is_uncond=False, rope_func="comfy",
|
||||
num_latent_frames=21, nag_params={}, nag_context=None, rope_func="comfy",
|
||||
inner_t=None, inner_c=None, cross_freqs=None,
|
||||
adapter_proj=None, adapter_attn_mask=None, ip_scale=1.0, orig_seq_len=None, lynx_x_ip=None, lynx_ip_scale=1.0, num_cond_latents=None, **kwargs):
|
||||
b, n, d = x.size(0), self.num_heads, self.head_dim
|
||||
@@ -678,7 +656,7 @@ class WanT2VCrossAttention(WanSelfAttention):
|
||||
else:
|
||||
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:
|
||||
x = self.normalized_attention_guidance(b, n, d, q, context, nag_context, nag_params)
|
||||
else:
|
||||
if is_longcat:
|
||||
@@ -689,7 +667,7 @@ class WanT2VCrossAttention(WanSelfAttention):
|
||||
v = self.v(context).view(b, -1, n, d)
|
||||
|
||||
#EchoShot rope
|
||||
if inner_t is not None and cross_freqs is not None and not is_uncond:
|
||||
if inner_t is not None and cross_freqs is not None:
|
||||
q = rope_apply_z(q, grid_sizes, cross_freqs, inner_t).to(q)
|
||||
k = rope_apply_c(k, cross_freqs, inner_c).to(q)
|
||||
|
||||
@@ -758,7 +736,7 @@ class WanI2VCrossAttention(WanSelfAttention):
|
||||
self.attention_mode = attention_mode
|
||||
|
||||
def forward(self, x, context, grid_sizes=None, clip_embed=None, audio_proj=None,
|
||||
audio_scale=1.0, num_latent_frames=21, nag_params={}, nag_context=None, is_uncond=False, rope_func="comfy",
|
||||
audio_scale=1.0, num_latent_frames=21, nag_params={}, nag_context=None, rope_func="comfy",
|
||||
adapter_proj=None, adapter_attn_mask=None, ip_scale=1.0, orig_seq_len=None, **kwargs):
|
||||
r"""
|
||||
Args:
|
||||
@@ -769,7 +747,7 @@ class WanI2VCrossAttention(WanSelfAttention):
|
||||
# compute query
|
||||
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:
|
||||
x_text = self.normalized_attention_guidance(b, n, d, q, context, nag_context, nag_params)
|
||||
else:
|
||||
# text attention
|
||||
@@ -935,9 +913,7 @@ class WanAttentionBlock(nn.Module):
|
||||
self.norm2 = WanLayerNorm(self.dim, eps)
|
||||
|
||||
if not is_longcat:
|
||||
self.ffn = nn.Sequential(
|
||||
nn.Linear(in_features, ffn_dim), nn.GELU(approximate='tanh'),
|
||||
nn.Linear(ffn2_dim, out_features))
|
||||
self.ffn = nn.Sequential(nn.Linear(in_features, ffn_dim), nn.GELU(approximate='tanh'), nn.Linear(ffn2_dim, out_features))
|
||||
else:
|
||||
from ...LongCat.layers import FeedForwardSwiGLU
|
||||
mlp_ratio = 4
|
||||
@@ -1002,23 +978,17 @@ class WanAttentionBlock(nn.Module):
|
||||
else:
|
||||
return torch.addcmul(shift_msa, norm_x, 1 + scale_msa)
|
||||
|
||||
def ffn_chunked(self, x, shift_mlp, scale_mlp, num_chunks=4):
|
||||
modulated_input = torch.addcmul(shift_mlp, self.norm2(x.to(shift_mlp.dtype)), 1 + scale_mlp).to(x.dtype)
|
||||
def ffn_chunked(self, mod_x, num_chunks=4):
|
||||
seq_len = mod_x.shape[1]
|
||||
if seq_len <= 8192 or num_chunks <= 1:
|
||||
return self.ffn(mod_x)
|
||||
|
||||
result = torch.empty_like(x)
|
||||
seq_len = modulated_input.shape[1]
|
||||
chunk_size = (seq_len + num_chunks - 1) // num_chunks
|
||||
for i in range(0, seq_len, chunk_size):
|
||||
end_idx = min(i + chunk_size, seq_len)
|
||||
mod_x[:, i:end_idx] = self.ffn(mod_x[:, i:end_idx].contiguous())
|
||||
|
||||
chunk_sizes = [seq_len // num_chunks + (1 if i < seq_len % num_chunks else 0)
|
||||
for i in range(num_chunks)]
|
||||
|
||||
start_idx = 0
|
||||
for size in chunk_sizes:
|
||||
end_idx = start_idx + size
|
||||
chunk = modulated_input[:, start_idx:end_idx, :]
|
||||
result[:, start_idx:end_idx, :] = self.ffn(chunk)
|
||||
start_idx = end_idx
|
||||
|
||||
return result
|
||||
return mod_x
|
||||
|
||||
#region attention forward
|
||||
def forward(
|
||||
@@ -1033,7 +1003,6 @@ class WanAttentionBlock(nn.Module):
|
||||
original_seq_len=None,
|
||||
enhance_enabled=False, #feta
|
||||
nag_params={}, nag_context=None, #normalized attention guidance
|
||||
is_uncond=False,
|
||||
multitalk_audio_embedding=None, ref_target_masks=None, human_num=0, #multitalk
|
||||
inner_t=None, inner_c=None, cross_freqs=None, #echoshot
|
||||
x_ip=None, e_ip=None, freqs_ip=None, ip_scale=1.0, #stand-in
|
||||
@@ -1099,6 +1068,9 @@ class WanAttentionBlock(nn.Module):
|
||||
# self-attention variables
|
||||
q_ip = k_ip = v_ip = None
|
||||
|
||||
if lynx_ref_feature is None and self.self_attn.ref_adapter is not None:
|
||||
lynx_ref_feature = input_x
|
||||
|
||||
#RoPE and QKV computation
|
||||
if inner_t is not None:
|
||||
#query, key, value
|
||||
@@ -1109,33 +1081,35 @@ class WanAttentionBlock(nn.Module):
|
||||
# First pass - separate main and IP components
|
||||
x_main, x_ip_input = input_x[:, : -self.cond_size], input_x[:, -self.cond_size :]
|
||||
# Compute QKV for main content
|
||||
q, k, v = self.self_attn.qkv_fn(x_main)
|
||||
if self.rope_func == "comfy":
|
||||
q, k = apply_rope_comfy(q, k, freqs)
|
||||
q = self.self_attn.qkv_fn_q_with_rope(x_main, freqs)
|
||||
k = self.self_attn.qkv_fn_k_with_rope(x_main, freqs)
|
||||
v = self.self_attn.qkv_fn_v(x_main)
|
||||
elif self.rope_func == "comfy_chunked":
|
||||
q, k = apply_rope_comfy_chunked(q, k, freqs)
|
||||
q = self.self_attn.qkv_fn_q_with_rope(x_main, freqs, num_chunks=2)
|
||||
k = self.self_attn.qkv_fn_k_with_rope(x_main, freqs, num_chunks=2)
|
||||
v = self.self_attn.qkv_fn_v(x_main)
|
||||
# Compute QKV for IP content
|
||||
q_ip, k_ip, v_ip = self.self_attn.qkv_fn_ip(x_ip_input)
|
||||
if self.rope_func == "comfy":
|
||||
if "comfy" in self.rope_func:
|
||||
q_ip, k_ip, v_ip = self.self_attn.qkv_fn_ip(x_ip_input)
|
||||
q_ip, k_ip = apply_rope_comfy(q_ip, k_ip, freqs_ip)
|
||||
elif self.rope_func == "comfy_chunked":
|
||||
q_ip, k_ip = apply_rope_comfy_chunked(q_ip, k_ip, freqs_ip)
|
||||
else:
|
||||
if is_longcat:
|
||||
q, k, v = self.self_attn.qkv_fn_longcat(input_x)
|
||||
if "comfy" in self.rope_func:
|
||||
num_chunks = 2 if self.rope_func == "comfy_chunked" else 1
|
||||
q = self.self_attn.qkv_fn_q_with_rope(input_x, freqs, num_chunks=num_chunks, is_longcat=is_longcat)
|
||||
k = self.self_attn.qkv_fn_k_with_rope(input_x, freqs, num_chunks=num_chunks, is_longcat=is_longcat)
|
||||
v = self.self_attn.qkv_fn_v(input_x)
|
||||
else:
|
||||
q, k, v = self.self_attn.qkv_fn(input_x)
|
||||
if self.rope_func == "comfy":
|
||||
q, k = apply_rope_comfy(q, k, freqs)
|
||||
elif self.rope_func == "comfy_chunked":
|
||||
q, k = apply_rope_comfy_chunked(q, k, freqs)
|
||||
elif self.rope_func == "mocha":
|
||||
from ...mocha.nodes import rope_apply_mocha
|
||||
q=rope_apply_mocha(q, grid_sizes, freqs)
|
||||
k=rope_apply_mocha(k, grid_sizes, freqs)
|
||||
else:
|
||||
q = rope_apply(q, grid_sizes, freqs, reverse_time=reverse_time)
|
||||
k = rope_apply(k, grid_sizes, freqs, reverse_time=reverse_time)
|
||||
if self.rope_func == "mocha":
|
||||
from ...mocha.nodes import rope_apply_mocha
|
||||
q = rope_apply_mocha(q, grid_sizes, freqs)
|
||||
k = rope_apply_mocha(k, grid_sizes, freqs)
|
||||
else:
|
||||
q = rope_apply(q, grid_sizes, freqs, reverse_time=reverse_time)
|
||||
k = rope_apply(k, grid_sizes, freqs, reverse_time=reverse_time)
|
||||
|
||||
del input_x
|
||||
|
||||
if x_ovi is not None:
|
||||
q_ovi, k_ovi, v_ovi = self.audio_block.self_attn.qkv_fn(input_x_ovi)
|
||||
@@ -1143,7 +1117,7 @@ class WanAttentionBlock(nn.Module):
|
||||
k_ovi = rope_apply(k_ovi, grid_sizes_ovi, freqs_ovi)
|
||||
y_ovi = self.audio_block.self_attn.forward(q_ovi, k_ovi, v_ovi, seq_lens_ovi)
|
||||
x_ovi = x_ovi.addcmul(y_ovi, gate_msa_ovi)
|
||||
|
||||
del input_x_ovi, y_ovi, gate_msa_ovi
|
||||
|
||||
# FETA
|
||||
if enhance_enabled:
|
||||
@@ -1198,9 +1172,8 @@ class WanAttentionBlock(nn.Module):
|
||||
y = torch.cat([x_cond, x_noise], dim=1).contiguous()
|
||||
else:
|
||||
y = self.self_attn.forward(q, k, v, seq_lens, lynx_ref_feature=lynx_ref_feature, lynx_ref_scale=lynx_ref_scale)
|
||||
|
||||
if lynx_ref_feature is None and self.self_attn.ref_adapter is not None:
|
||||
lynx_ref_feature = input_x
|
||||
|
||||
del q, k, v
|
||||
|
||||
# FETA
|
||||
if enhance_enabled:
|
||||
@@ -1259,7 +1232,7 @@ class WanAttentionBlock(nn.Module):
|
||||
return x, x_ip, lynx_ref_feature, x_ovi
|
||||
else:
|
||||
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,
|
||||
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)
|
||||
x = x.to(input_dtype)
|
||||
@@ -1281,7 +1254,8 @@ class WanAttentionBlock(nn.Module):
|
||||
|
||||
# ffn
|
||||
if self.rope_func == "comfy_chunked":
|
||||
x_ffn = self.ffn_chunked(x, shift_mlp, scale_mlp)
|
||||
mod_x = torch.addcmul(shift_mlp, self.norm2(x.to(shift_mlp.dtype)), 1 + scale_mlp)
|
||||
x_ffn = self.ffn_chunked(mod_x)
|
||||
else:
|
||||
if zero_timestep:
|
||||
norm2_x = self.norm2(x)
|
||||
@@ -1296,8 +1270,9 @@ class WanAttentionBlock(nn.Module):
|
||||
mod_x = torch.addcmul(shift_mlp, self.norm2(x.to(shift_mlp.dtype)), 1 + scale_mlp)
|
||||
else:
|
||||
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.to(input_dtype))
|
||||
del shift_mlp, scale_mlp
|
||||
del shift_mlp, scale_mlp
|
||||
x_ffn = self.ffn_chunked(mod_x.to(input_dtype), num_chunks=1)
|
||||
del mod_x
|
||||
|
||||
# gate_mlp
|
||||
if zero_timestep:
|
||||
@@ -2828,13 +2803,13 @@ class WanModel(torch.nn.Module):
|
||||
original_seq_len=self.original_seq_len,
|
||||
enhance_enabled=enhance_enabled,
|
||||
audio_scale=audio_scale,
|
||||
nag_params=nag_params, nag_context=nag_context,
|
||||
is_uncond = is_uncond,
|
||||
nag_params=nag_params,
|
||||
nag_context=nag_context if not is_uncond else None,
|
||||
multitalk_audio_embedding=multitalk_audio_embedding if multitalk_audio is not None else None,
|
||||
ref_target_masks=token_ref_target_masks if multitalk_audio is not None else None,
|
||||
human_num=human_num if multitalk_audio is not None else 0,
|
||||
inner_t=inner_t, inner_c=inner_c,
|
||||
cross_freqs=self.cross_freqs if inner_t is not None else None,
|
||||
cross_freqs=self.cross_freqs if inner_t is not None and not is_uncond else None,
|
||||
freqs_ip=freqs_ip if x_ip is not None else None,
|
||||
e_ip=e0_ip if x_ip is not None else None,
|
||||
adapter_proj=adapter_proj,
|
||||
|
||||
Reference in New Issue
Block a user