diff --git a/acestep/music_dcae/music_dcae_pipeline.py b/acestep/music_dcae/music_dcae_pipeline.py index 8112638..af23ed8 100644 --- a/acestep/music_dcae/music_dcae_pipeline.py +++ b/acestep/music_dcae/music_dcae_pipeline.py @@ -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") diff --git a/acestep/pipeline_ace_step.py b/acestep/pipeline_ace_step.py index a734d1a..552a2a1 100644 --- a/acestep/pipeline_ace_step.py +++ b/acestep/pipeline_ace_step.py @@ -12,6 +12,7 @@ import os import re import torch +import soundfile as sf from loguru import logger from tqdm import tqdm import json @@ -46,7 +47,6 @@ from acestep.apg_guidance import ( cfg_zero_star, cfg_double_condition_forward, ) -import torchaudio from .cpu_offload import cpu_offload @@ -1405,13 +1405,17 @@ class ACEStepPipeline: else: output_path_wav = save_path - target_wav = target_wav.float() - backend = "soundfile" - if format == "ogg": - backend = "sox" - logger.info(f"Saving audio to {output_path_wav} using backend {backend}") - torchaudio.save( - output_path_wav, target_wav, sample_rate=sample_rate, format=format, backend=backend + target_wav = target_wav.float().cpu() + logger.info(f"Saving audio to {output_path_wav}") + # Write with soundfile rather than torchaudio.save(): since torchaudio + # 2.11 the latter ignores the `backend` argument and routes everything + # through TorchCodec, an extra native dependency we do not require. + # soundfile expects (samples, channels), torch tensors are (channels, samples). + sf.write( + output_path_wav, + target_wav.transpose(0, 1).numpy(), + sample_rate, + format=format.upper(), ) return output_path_wav diff --git a/acestep/text2music_dataset.py b/acestep/text2music_dataset.py index 0a5b307..0a964c4 100644 --- a/acestep/text2music_dataset.py +++ b/acestep/text2music_dataset.py @@ -7,6 +7,7 @@ from loguru import logger import time import traceback import torchaudio +import soundfile as sf from pathlib import Path import re from acestep.language_segmentation import LangSegment @@ -398,7 +399,10 @@ class Text2MusicDataset(Dataset): filename = item["filename"] sr = 48000 try: - audio, sr = torchaudio.load(filename) + # soundfile instead of torchaudio.load(): torchaudio 2.11 routes + # I/O through TorchCodec, an extra native dependency. + _data, sr = sf.read(filename, dtype="float32", always_2d=True) + audio = torch.from_numpy(_data.T) except Exception as e: logger.error(f"Failed to load audio {item}: {e}") return None