trim trailing whitespace

This commit is contained in:
kijai 2025-08-30 21:06:54 +03:00
parent b64004fc5e
commit 66e10910dc
2 changed files with 5 additions and 5 deletions

View File

@ -43,7 +43,7 @@ def load_audio_encoder_from_sd(sd, prefix=""):
sd = comfy.utils.state_dict_prefix_replace(sd, {"wav2vec2.": ""})
embed_dim = sd["encoder.layer_norm.bias"].shape[0]
if embed_dim == 1024:# large
config = {
config = {
"embed_dim": 1024,
"num_heads": 16,
"num_layers": 24,

View File

@ -12,7 +12,7 @@ class LayerNormConv(nn.Module):
def forward(self, x):
x = self.conv(x)
return torch.nn.functional.gelu(self.layer_norm(x.transpose(-2, -1)).transpose(-2, -1))
class LayerGroupNormConv(nn.Module):
def __init__(self, in_channels, out_channels, kernel_size, stride, bias=False, dtype=None, device=None, operations=None):
super().__init__()
@ -22,7 +22,7 @@ class LayerGroupNormConv(nn.Module):
def forward(self, x):
x = self.conv(x)
return torch.nn.functional.gelu(self.layer_norm(x))
class ConvNoNorm(nn.Module):
def __init__(self, in_channels, out_channels, kernel_size, stride, bias=False, dtype=None, device=None, operations=None):
super().__init__()
@ -250,13 +250,13 @@ class Wav2Vec2Model(nn.Module):
audio_duration = x.shape[1] / sr
target_seq_len = int(audio_duration * fps)
features = self.linear_interpolation(features, target_seq_len)
features = self.feature_projection(features)
batch_size, seq_len, _ = features.shape
x, all_x = self.encoder(features)
return x, all_x
def linear_interpolation(self, features, target_seq_len):
features = features.transpose(1, 2)
output_features = torch.nn.functional.interpolate(features, size=target_seq_len, align_corners=True, mode='linear')