From bfa45fb0e674537f9d5833f9490c4a94c9c39c07 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Thu, 4 Sep 2025 20:08:26 +0300 Subject: [PATCH] Do interpolation after --- comfy/audio_encoders/audio_encoders.py | 9 +++------ comfy/audio_encoders/wav2vec2.py | 13 +------------ 2 files changed, 4 insertions(+), 18 deletions(-) diff --git a/comfy/audio_encoders/audio_encoders.py b/comfy/audio_encoders/audio_encoders.py index b452ff447..d1ec78f69 100644 --- a/comfy/audio_encoders/audio_encoders.py +++ b/comfy/audio_encoders/audio_encoders.py @@ -7,7 +7,7 @@ import torchaudio class AudioEncoderModel(): - def __init__(self, config, fps=50): + def __init__(self, config): self.load_device = comfy.model_management.text_encoder_device() offload_device = comfy.model_management.text_encoder_offload_device() self.dtype = comfy.model_management.text_encoder_dtype(self.load_device) @@ -21,7 +21,6 @@ class AudioEncoderModel(): self.model.eval() self.patcher = comfy.model_patcher.ModelPatcher(self.model, load_device=self.load_device, offload_device=offload_device) self.model_sample_rate = 16000 - self.fps = fps def load_sd(self, sd): return self.model.load_state_dict(sd, strict=False) @@ -32,7 +31,7 @@ class AudioEncoderModel(): def encode_audio(self, audio, sample_rate): comfy.model_management.load_model_gpu(self.patcher) audio = torchaudio.functional.resample(audio, sample_rate, self.model_sample_rate) - out, all_layers = self.model(audio.to(self.load_device), fps=self.fps, sr=self.model_sample_rate) + out, all_layers = self.model(audio.to(self.load_device), sr=self.model_sample_rate) outputs = {} outputs["encoded_audio"] = out outputs["encoded_audio_all_layers"] = all_layers @@ -52,7 +51,6 @@ def load_audio_encoder_from_sd(sd, prefix=""): "do_normalize": True, "do_stable_layer_norm": True } - fps = 50 elif embed_dim == 768: # base config = { "embed_dim": 768, @@ -63,11 +61,10 @@ def load_audio_encoder_from_sd(sd, prefix=""): "do_normalize": False, # chinese-wav2vec2-base has this False "do_stable_layer_norm": False } - fps = 25 else: raise RuntimeError("ERROR: audio encoder file is invalid or unsupported embed_dim: {}".format(embed_dim)) - audio_encoder = AudioEncoderModel(config, fps=fps) + audio_encoder = AudioEncoderModel(config) m, u = audio_encoder.load_sd(sd) if len(m) > 0: logging.warning("missing audio encoder: {}".format(m)) diff --git a/comfy/audio_encoders/wav2vec2.py b/comfy/audio_encoders/wav2vec2.py index 4933b7501..ef10dcd2a 100644 --- a/comfy/audio_encoders/wav2vec2.py +++ b/comfy/audio_encoders/wav2vec2.py @@ -238,26 +238,15 @@ class Wav2Vec2Model(nn.Module): device=device, dtype=dtype, operations=operations ) - def forward(self, x, fps=50, sr=16000, mask_time_indices=None, return_dict=False): + def forward(self, x, sr=16000, mask_time_indices=None, return_dict=False): x = torch.mean(x, dim=1) if self.do_normalize: x = (x - x.mean()) / torch.sqrt(x.var() + 1e-7) features = self.feature_extractor(x) - - if fps != 50: - 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') - return output_features.transpose(1, 2)