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
61 changes: 51 additions & 10 deletions acestep/audio_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,9 +11,11 @@
import io
import json
import os
import shutil
import subprocess
import hashlib
import tempfile
import uuid
from pathlib import Path
from typing import Union, Optional, List, Tuple
import torch
Expand All @@ -22,6 +24,29 @@
from loguru import logger


class AudioExportDegradedError(RuntimeError):
"""Raised when audio generation succeeded but the requested export format failed.

The underlying audio was already synthesized successfully; only the
format conversion (e.g. WAV -> MP3 via ffmpeg) failed. ``wav_fallback_path``
points to a valid WAV file already saved to disk, so callers can present
this as a degraded-but-successful result instead of a hard failure.
"""

def __init__(self, message: str, wav_fallback_path: str, requested_format: str):
Comment thread
coderabbitai[bot] marked this conversation as resolved.
"""
Initialize degraded export error.

Args:
message: Error message describing the failure.
wav_fallback_path: Path to the preserved WAV fallback file.
requested_format: The audio format that failed to export.
"""
super().__init__(message)
self.wav_fallback_path = wav_fallback_path
self.requested_format = requested_format


def apply_fade(
audio_data: Union[torch.Tensor, np.ndarray],
fade_in_samples: int = 0,
Expand Down Expand Up @@ -181,18 +206,33 @@ def _save_mp3(
]
subprocess.run(cmd, check=True, capture_output=True, timeout=120)
logger.debug(f"[AudioSaver] Saved audio to {output_path} (mp3, {target_sample_rate}Hz, {bitrate})")
except FileNotFoundError as e:
raise RuntimeError("ffmpeg executable not found. Install ffmpeg or add it to PATH to export MP3 files.") from e
except subprocess.TimeoutExpired as e:
raise RuntimeError("ffmpeg MP3 export timed out after 120 seconds.") from e
except subprocess.CalledProcessError as e:
stderr = e.stderr.decode('utf-8', errors='ignore') if e.stderr else str(e)
raise RuntimeError(f"ffmpeg MP3 export failed: {stderr}") from e
except (FileNotFoundError, subprocess.TimeoutExpired, subprocess.CalledProcessError) as e:
if isinstance(e, FileNotFoundError):
reason = "ffmpeg executable not found. Install ffmpeg or add it to PATH to export MP3 files."
elif isinstance(e, subprocess.TimeoutExpired):
reason = "ffmpeg MP3 export timed out after 120 seconds."
else:
stderr = e.stderr.decode("utf-8", errors="ignore") if e.stderr else str(e)
reason = f"ffmpeg MP3 export failed: {stderr}"

# The WAV was already synthesized successfully before the ffmpeg
# step -- preserve it instead of discarding a valid result.
wav_fallback_path = output_path.with_suffix(".wav")
if wav_fallback_path.exists():
wav_fallback_path = output_path.with_name(f"{output_path.stem}_fallback_{uuid.uuid4().hex[:8]}.wav")
try:
shutil.move(str(temp_wav_path), str(wav_fallback_path))
except Exception:
logger.error(f"[AudioSaver] {reason} Additionally failed to preserve WAV fallback.")
raise RuntimeError(reason) from e

logger.warning(f"[AudioSaver] {reason} Saved WAV fallback to {wav_fallback_path} instead.")
raise AudioExportDegradedError(reason, str(wav_fallback_path), "mp3") from e
finally:
try:
temp_wav_path.unlink(missing_ok=True)
except Exception:
logger.warning(f"[AudioSaver] Failed to remove temporary WAV file: {temp_wav_path}")
except OSError as exc:
logger.warning(f"[AudioSaver] Failed to remove temporary WAV file {temp_wav_path}: {exc}")
Comment thread
coderabbitai[bot] marked this conversation as resolved.

def save_audio(
self,
Expand Down Expand Up @@ -321,7 +361,8 @@ def save_audio(

except Exception as e:
if format == "mp3":
logger.error(f"[AudioSaver] MP3 export failed without fallback: {e}")
if not isinstance(e, AudioExportDegradedError):
logger.error(f"[AudioSaver] MP3 export failed without fallback: {e}")
raise
try:
import soundfile as sf
Expand Down
102 changes: 101 additions & 1 deletion acestep/audio_utils_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,7 @@
import torch
import numpy as np

from acestep.audio_utils import AudioSaver, apply_fade, save_audio
from acestep.audio_utils import AudioSaver, AudioExportDegradedError, apply_fade, save_audio

class AudioSaverFormatTests(unittest.TestCase):
"""Tests for AudioSaver format support, especially new Opus and AAC formats."""
Expand Down Expand Up @@ -335,6 +335,106 @@ def test_save_audio_mp3_does_not_fallback_to_soundfile_on_failure(self):
format="mp3",
)

def test_save_mp3_missing_ffmpeg_preserves_wav_fallback(self):
"""When ffmpeg is missing, the already-synthesized WAV must be preserved.

Generation itself already succeeded before the ffmpeg step runs; a
missing ffmpeg binary should not discard that valid audio.
"""
saver = AudioSaver()
output_path = Path(self.temp_dir) / "test.mp3"

with patch(
"acestep.audio_utils.subprocess.run",
side_effect=FileNotFoundError("ffmpeg not found"),
):
with self.assertRaises(AudioExportDegradedError) as ctx:
saver._save_mp3(self.sample_audio, output_path, self.sample_rate)

exc = ctx.exception
self.assertEqual(exc.requested_format, "mp3")
self.assertTrue(exc.wav_fallback_path.endswith(".wav"))
self.assertTrue(os.path.exists(exc.wav_fallback_path))
self.assertIn("ffmpeg", str(exc).lower())

def test_save_mp3_fallback_does_not_overwrite_existing_wav(self):
"""Fallback WAV should use a unique name if the default <stem>.wav already exists."""
saver = AudioSaver()
output_path = Path(self.temp_dir) / "test.mp3"
existing_wav = output_path.with_suffix(".wav")
existing_data = b"original data"
existing_wav.write_bytes(existing_data)

with patch(
"acestep.audio_utils.subprocess.run",
side_effect=FileNotFoundError("ffmpeg not found"),
):
with self.assertRaises(AudioExportDegradedError) as ctx:
saver._save_mp3(self.sample_audio, output_path, self.sample_rate)

exc = ctx.exception
# Ensure the original file is untouched
self.assertEqual(existing_wav.read_bytes(), existing_data)
# Ensure a different fallback file was created and returned
self.assertNotEqual(exc.wav_fallback_path, str(existing_wav))
self.assertTrue(os.path.exists(exc.wav_fallback_path))
self.assertTrue(Path(exc.wav_fallback_path).name.startswith("test_fallback_"))

def test_save_mp3_timeout_preserves_wav_fallback(self):
"""An ffmpeg timeout also preserves the already-synthesized WAV."""
import subprocess as subprocess_module

saver = AudioSaver()
output_path = Path(self.temp_dir) / "test.mp3"

with patch(
"acestep.audio_utils.subprocess.run",
side_effect=subprocess_module.TimeoutExpired(cmd="ffmpeg", timeout=120),
):
with self.assertRaises(AudioExportDegradedError) as ctx:
saver._save_mp3(self.sample_audio, output_path, self.sample_rate)

self.assertTrue(os.path.exists(ctx.exception.wav_fallback_path))

def test_save_audio_mp3_missing_ffmpeg_raises_degraded_error_with_valid_wav(self):
"""save_audio() propagates AudioExportDegradedError with a playable WAV fallback."""
saver = AudioSaver()
output_path = Path(self.temp_dir) / "test.mp3"

with patch(
"acestep.audio_utils.subprocess.run",
side_effect=FileNotFoundError("ffmpeg not found"),
):
with self.assertRaises(AudioExportDegradedError) as ctx:
saver.save_audio(
self.sample_audio,
output_path,
sample_rate=self.sample_rate,
format="mp3",
)

wav_path = ctx.exception.wav_fallback_path
self.assertTrue(os.path.exists(wav_path))
self.assertGreater(os.path.getsize(wav_path), 0)

def test_save_mp3_temp_file_cleaned_up_when_fallback_move_fails(self):
"""If even the WAV fallback move fails, the original ffmpeg reason still surfaces."""
saver = AudioSaver()
output_path = Path(self.temp_dir) / "test.mp3"

with (
patch(
"acestep.audio_utils.subprocess.run",
side_effect=FileNotFoundError("ffmpeg not found"),
),
patch("acestep.audio_utils.shutil.move", side_effect=OSError("disk full")),
):
with self.assertRaises(RuntimeError) as ctx:
saver._save_mp3(self.sample_audio, output_path, self.sample_rate)

self.assertNotIsInstance(ctx.exception, AudioExportDegradedError)
self.assertIn("ffmpeg", str(ctx.exception).lower())

class ApplyFadeTests(unittest.TestCase):
"""Tests for apply_fade function."""

Expand Down
28 changes: 21 additions & 7 deletions acestep/ui/gradio/events/results/generation_progress.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,7 @@
from loguru import logger

from acestep.inference import generate_music, GenerationParams, GenerationConfig
from acestep.audio_utils import save_audio
from acestep.audio_utils import save_audio, AudioExportDegradedError
from acestep.gpu_config import (
get_global_gpu_config,
check_duration_limit,
Expand Down Expand Up @@ -273,6 +273,7 @@ def generate_with_progress(
)
time_module.sleep(0.1)

export_warnings = []
for i in range(8):
if i >= len(audios):
continue
Expand All @@ -291,11 +292,19 @@ def generate_with_progress(
ext = "wav" if audio_format == "wav32" else audio_format
audio_path = os.path.join(temp_dir, f"{key}.{ext}").replace("\\", "/")

saved_path = save_audio(
audio_data=audio_tensor, output_path=audio_path,
sample_rate=sample_rate, format=audio_format, channels_first=True,
mp3_bitrate=mp3_bitrate, mp3_sample_rate=mp3_sample_rate,
)
try:
saved_path = save_audio(
audio_data=audio_tensor, output_path=audio_path,
sample_rate=sample_rate, format=audio_format, channels_first=True,
mp3_bitrate=mp3_bitrate, mp3_sample_rate=mp3_sample_rate,
)
except AudioExportDegradedError as exc:
# Generation itself already succeeded (audio_tensor exists); only
# the requested export format failed. Use the WAV ACE-Step already
# saved instead of discarding a valid result as a hard error.
saved_path = exc.wav_fallback_path
logger.warning(f"[generate_with_progress] Sample {key}: {exc}")
export_warnings.append(f"{key}: {exc.requested_format} export unavailable ({exc}), saved as WAV")
if saved_path:
audio_path = saved_path.replace("\\", "/")

Expand Down Expand Up @@ -413,9 +422,14 @@ def generate_with_progress(
{**result.extra_outputs, "lrcs": final_lrcs_list, "subtitles": final_subtitles_list}
)

if export_warnings:
final_status = "Generation Complete (WAV). " + "; ".join(export_warnings)
else:
final_status = "Generation Complete"

yield (
*audio_playback_updates,
all_audio_paths, generation_info, "Generation Complete", seed_value_for_ui,
all_audio_paths, generation_info, final_status, seed_value_for_ui,
*final_scores_list, *final_codes_display, *final_accordions, *final_lrcs_list,
lm_generated_metadata, is_format_caption,
extra_to_store,
Expand Down