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.": ""}) sd = comfy.utils.state_dict_prefix_replace(sd, {"wav2vec2.": ""})
embed_dim = sd["encoder.layer_norm.bias"].shape[0] embed_dim = sd["encoder.layer_norm.bias"].shape[0]
if embed_dim == 1024:# large if embed_dim == 1024:# large
config = { config = {
"embed_dim": 1024, "embed_dim": 1024,
"num_heads": 16, "num_heads": 16,
"num_layers": 24, "num_layers": 24,

View File

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