fix: write and read audio with soundfile instead of torchaudio

torchaudio 2.11 routes torchaudio.save()/load() through TorchCodec and
ignores the `backend` argument, so every generation died at the save
step with "ImportError: TorchCodec is required for save_with_torchcodec"
after the diffusion had already finished. Reference-audio loading
(audio2audio, repaint, extend) and the training dataset loader hit the
same wall.

soundfile is already a required dependency and covers all four output
formats the UI offers, so use it directly rather than pulling in
TorchCodec and its native FFmpeg stack:

- pipeline_ace_step.save_wav_file(): sf.write(), transposing
  (channels, samples) -> (samples, channels); drops the now-unused
  torchaudio import
- MusicDCAE.load_audio() and text2music_dataset: sf.read(dtype=float32,
  always_2d=True), transposed back to (channels, samples)

torchaudio is still used for Resample/MelScale transforms, which are
unaffected.

Verified end to end: 10s generation on an RTX 3060 in 9.7s, output is
valid non-silent 48kHz stereo; load_audio round-trips it; wav/mp3/ogg/
flac all write and read back.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
Leonid Pershin
2026-09-08 19:03:29 +03:00
co-authored by Claude Opus 5
parent 6e3273d049
commit 0584397884
3 changed files with 26 additions and 12 deletions
+9 -3
View File
@@ -10,6 +10,7 @@ import os
import torch
from diffusers import AutoencoderDC
import torchaudio
import soundfile as sf
import torchvision.transforms as transforms
from diffusers.models.modeling_utils import ModelMixin
from diffusers.loaders import FromOriginalModelMixin
@@ -60,7 +61,11 @@ class MusicDCAE(ModelMixin, ConfigMixin, FromOriginalModelMixin):
self.shift_factor = -1.9091
def load_audio(self, audio_path):
audio, sr = torchaudio.load(audio_path)
# Read with soundfile rather than torchaudio.load(): since torchaudio
# 2.11 the latter routes I/O through TorchCodec, an extra native
# dependency we do not require.
data, sr = sf.read(audio_path, dtype="float32", always_2d=True)
audio = torch.from_numpy(data.T)
if audio.shape[0] == 1:
audio = audio.repeat(2, 1)
return audio, sr
@@ -362,7 +367,8 @@ class MusicDCAE(ModelMixin, ConfigMixin, FromOriginalModelMixin):
if __name__ == "__main__":
audio, sr = torchaudio.load("test.wav")
_data, sr = sf.read("test.wav", dtype="float32", always_2d=True)
audio = torch.from_numpy(_data.T)
audio_lengths = torch.tensor([audio.shape[1]])
audios = audio.unsqueeze(0)
@@ -378,5 +384,5 @@ if __name__ == "__main__":
print("latents shape: ", latents.shape)
print("latent_lengths: ", latent_lengths)
print("sr: ", sr)
torchaudio.save("test_reconstructed.wav", pred_wavs[0], sr)
sf.write("test_reconstructed.wav", pred_wavs[0].float().cpu().transpose(0, 1).numpy(), sr)
print("test_reconstructed.wav")