Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
118 changes: 117 additions & 1 deletion acestep/training/dataset_builder_modules/preprocess_audio.py
Original file line number Diff line number Diff line change
@@ -1,9 +1,125 @@
import json
import shutil
import subprocess

import numpy as np
import torch
import torchaudio

# A long clip decodes in well under a minute. The cap stops a stuck ffmpeg
# from blocking preprocessing forever.
_FFMPEG_TIMEOUT_SECONDS = 300


def _run_checked(command: list[str], *, text: bool) -> subprocess.CompletedProcess:
"""Run a command and include its stderr when it fails."""
try:
return subprocess.run(
command,
check=True,
capture_output=True,
text=text,
timeout=_FFMPEG_TIMEOUT_SECONDS,
)
except subprocess.CalledProcessError as exc:
detail = exc.stderr or ""
if isinstance(detail, bytes):
detail = detail.decode("utf-8", errors="replace")
detail = detail.strip() or f"command exited {exc.returncode}"
raise RuntimeError(detail) from exc


def _probe_audio_stream(ffprobe: str, audio_path: str) -> tuple[int, int]:
"""Return the sample rate and channel count of the first audio stream."""
probe = _run_checked(
[
ffprobe,
"-v",
"error",
"-select_streams",
"a:0",
"-show_entries",
"stream=sample_rate,channels",
"-of",
"json",
"-i",
audio_path,
],
text=True,
)
try:
streams = json.loads(probe.stdout)["streams"]
stream = streams[0]
sample_rate = int(stream["sample_rate"])
channels = int(stream["channels"])
except (KeyError, IndexError, TypeError, ValueError, json.JSONDecodeError) as exc:
raise RuntimeError(
f"ffprobe returned no usable audio stream for {audio_path}"
) from exc
if sample_rate < 1:
raise RuntimeError(
f"ffprobe reported sample rate {sample_rate} for {audio_path}"
)
if channels < 1:
raise RuntimeError(f"ffprobe reported no audio channels for {audio_path}")
return sample_rate, channels


def _load_via_ffmpeg(audio_path: str) -> tuple[torch.Tensor, int]:
"""Decode with the ffmpeg binary when torchaudio's TorchCodec build cannot.

TorchCodec wheels only link FFmpeg 4-8. A newer system FFmpeg (for example
Homebrew FFmpeg 9, libavutil.61) makes torchaudio.load raise before any
samples are read. The ffmpeg CLI on PATH can still decode the file.
"""
ffmpeg = shutil.which("ffmpeg")
ffprobe = shutil.which("ffprobe")
if not ffmpeg or not ffprobe:
raise RuntimeError("ffmpeg and ffprobe are not on PATH")

sample_rate, channels = _probe_audio_stream(ffprobe, audio_path)
decoded = _run_checked(
[
ffmpeg,
"-v",
"error",
"-i",
audio_path,
"-map",
"0:a:0",
"-ac",
str(channels),
"-f",
"f32le",
"-acodec",
"pcm_f32le",
"-",
],
text=False,
)
pcm = np.frombuffer(decoded.stdout, dtype=np.float32).copy()
if pcm.size == 0:
raise RuntimeError(f"ffmpeg decoded no samples from {audio_path}")
if pcm.size % channels != 0:
raise RuntimeError(
f"ffmpeg output size is not divisible by {channels} channels"
)
audio = torch.from_numpy(pcm.reshape(-1, channels).T).contiguous()
Comment thread
coderabbitai[bot] marked this conversation as resolved.
return audio, sample_rate


def load_audio_stereo(audio_path: str, target_sample_rate: int, max_duration: float):
"""Load audio, resample, convert to stereo, and truncate."""
audio, sr = torchaudio.load(audio_path)
try:
audio, sr = torchaudio.load(audio_path)
except Exception as exc:
try:
audio, sr = _load_via_ffmpeg(audio_path)
except Exception as fallback_exc:
raise RuntimeError(
f"Could not decode {audio_path}. torchaudio failed ({exc}); "
f"ffmpeg fallback failed ({fallback_exc})."
) from fallback_exc

if sr != target_sample_rate:
resampler = torchaudio.transforms.Resample(sr, target_sample_rate)
Expand Down
176 changes: 176 additions & 0 deletions acestep/training/dataset_builder_modules/preprocess_audio_test.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,176 @@
"""Unit tests for the torchaudio / ffmpeg preprocess decoder."""

import json
import subprocess
import unittest
from unittest.mock import MagicMock, patch

import numpy as np
import torch

from acestep.training.dataset_builder_modules import preprocess_audio


def _completed(stdout: str | bytes, returncode: int = 0) -> MagicMock:
"""Build a subprocess result with the fields the decoder reads."""
result = MagicMock()
result.returncode = returncode
result.stdout = stdout
result.stderr = b""
return result


class LoadAudioStereoTests(unittest.TestCase):
"""load_audio_stereo prefers torchaudio and falls back to the ffmpeg CLI."""

def test_uses_torchaudio_when_it_decodes(self):
"""A successful torchaudio load does not call ffmpeg."""
audio = torch.zeros(2, 8)
with patch.object(
preprocess_audio.torchaudio, "load", return_value=(audio, 48000)
) as load, patch.object(preprocess_audio.subprocess, "run") as run:
out, sample_rate = preprocess_audio.load_audio_stereo(
"song.mp3", 48000, 240
)

load.assert_called_once_with("song.mp3")
run.assert_not_called()
self.assertEqual(sample_rate, 48000)
self.assertEqual(tuple(out.shape), (2, 8))

def test_ffmpeg_fallback_decodes_interleaved_f32(self):
"""The fallback probes with -i and decodes the probed channel count."""
pcm = np.array([0.0, 0.5, -0.25, 1.0], dtype=np.float32)
probe = _completed(
json.dumps({"streams": [{"sample_rate": "48000", "channels": 2}]})
)
decoded = _completed(pcm.tobytes())
with patch.object(
preprocess_audio.torchaudio,
"load",
side_effect=RuntimeError("libavutil 61 is not supported"),
), patch.object(
preprocess_audio.shutil, "which", return_value="/usr/bin/ffmpeg"
), patch.object(
preprocess_audio.subprocess, "run", side_effect=[probe, decoded]
) as run:
out, sample_rate = preprocess_audio.load_audio_stereo(
"song.mp3", 48000, 240
)

probe_cmd = run.call_args_list[0].args[0]
decode_cmd = run.call_args_list[1].args[0]
self.assertEqual(probe_cmd[-2:], ["-i", "song.mp3"])
self.assertEqual(run.call_args_list[0].kwargs["timeout"], 300)
self.assertIn("-map", decode_cmd)
self.assertEqual(decode_cmd[decode_cmd.index("-ac") + 1], "2")
self.assertEqual(run.call_args_list[1].kwargs["timeout"], 300)
self.assertEqual(sample_rate, 48000)
self.assertEqual(tuple(out.shape), (2, 2))
self.assertTrue(torch.allclose(out[0], torch.tensor([0.0, -0.25])))
self.assertTrue(torch.allclose(out[1], torch.tensor([0.5, 1.0])))

def test_reports_both_decoder_failures(self):
"""A missing ffmpeg binary is included with the torchaudio error."""
with patch.object(
preprocess_audio.torchaudio,
"load",
side_effect=RuntimeError("torchcodec mismatch"),
), patch.object(preprocess_audio.shutil, "which", return_value=None):
with self.assertRaises(RuntimeError) as caught:
preprocess_audio.load_audio_stereo("song.mp3", 48000, 240)

message = str(caught.exception)
self.assertIn("torchcodec mismatch", message)
self.assertIn("ffmpeg and ffprobe are not on PATH", message)

def test_ffmpeg_failure_includes_the_process_error(self):
"""A non-zero ffmpeg exit becomes part of the combined error."""
failure = subprocess.CalledProcessError(
1, ["ffmpeg"], stderr=b"Invalid data found when processing input"
)
with patch.object(
preprocess_audio.torchaudio,
"load",
side_effect=RuntimeError("torchcodec mismatch"),
), patch.object(
preprocess_audio.shutil, "which", return_value="/usr/bin/ffmpeg"
), patch.object(preprocess_audio.subprocess, "run", side_effect=failure):
with self.assertRaises(RuntimeError) as caught:
preprocess_audio.load_audio_stereo("song.mp3", 48000, 240)

self.assertIn("Invalid data found", str(caught.exception))

def test_ffmpeg_timeout_is_reported(self):
"""A hung ffmpeg is reported instead of blocking preprocessing."""
failure = subprocess.TimeoutExpired(["ffprobe"], 300)
with patch.object(
preprocess_audio.torchaudio,
"load",
side_effect=RuntimeError("torchcodec mismatch"),
), patch.object(
preprocess_audio.shutil, "which", return_value="/usr/bin/ffmpeg"
), patch.object(preprocess_audio.subprocess, "run", side_effect=failure):
with self.assertRaises(RuntimeError) as caught:
preprocess_audio.load_audio_stereo("song.mp3", 48000, 240)

self.assertIn("timed out", str(caught.exception))

def test_rejects_pcm_that_does_not_match_the_channel_count(self):
"""A short final frame is an error instead of a mis-shaped tensor."""
probe = _completed(
json.dumps({"streams": [{"sample_rate": "48000", "channels": 2}]})
)
decoded = _completed(np.array([0.1, 0.2, 0.3], dtype=np.float32).tobytes())
with patch.object(
preprocess_audio.torchaudio,
"load",
side_effect=RuntimeError("torchcodec mismatch"),
), patch.object(
preprocess_audio.shutil, "which", return_value="/usr/bin/ffmpeg"
), patch.object(
preprocess_audio.subprocess, "run", side_effect=[probe, decoded]
):
with self.assertRaises(RuntimeError) as caught:
preprocess_audio.load_audio_stereo("song.mp3", 48000, 240)

self.assertIn("not divisible by 2", str(caught.exception))

def test_rejects_empty_ffmpeg_output(self):
"""A successful ffmpeg process that writes nothing is an error."""
probe = _completed(
json.dumps({"streams": [{"sample_rate": "48000", "channels": 2}]})
)
decoded = _completed(b"")
with patch.object(
preprocess_audio.torchaudio,
"load",
side_effect=RuntimeError("torchcodec mismatch"),
), patch.object(
preprocess_audio.shutil, "which", return_value="/usr/bin/ffmpeg"
), patch.object(
preprocess_audio.subprocess, "run", side_effect=[probe, decoded]
):
with self.assertRaises(RuntimeError) as caught:
preprocess_audio.load_audio_stereo("song.mp3", 48000, 240)

self.assertIn("decoded no samples", str(caught.exception))

def test_rejects_probe_output_without_a_stream(self):
"""Probe JSON with no audio stream names the file in the error."""
probe = _completed(json.dumps({"streams": []}))
with patch.object(
preprocess_audio.torchaudio,
"load",
side_effect=RuntimeError("torchcodec mismatch"),
), patch.object(
preprocess_audio.shutil, "which", return_value="/usr/bin/ffmpeg"
), patch.object(preprocess_audio.subprocess, "run", return_value=probe):
with self.assertRaises(RuntimeError) as caught:
preprocess_audio.load_audio_stereo("song.mp3", 48000, 240)

self.assertIn("no usable audio stream", str(caught.exception))


if __name__ == "__main__":
unittest.main()