diff --git a/acestep/audio_utils.py b/acestep/audio_utils.py index 45d9a7ae9..f5c689142 100644 --- a/acestep/audio_utils.py +++ b/acestep/audio_utils.py @@ -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 @@ -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): + """ + 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, @@ -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}") def save_audio( self, @@ -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 diff --git a/acestep/audio_utils_test.py b/acestep/audio_utils_test.py index 6751c07c1..4389d6554 100644 --- a/acestep/audio_utils_test.py +++ b/acestep/audio_utils_test.py @@ -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.""" @@ -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 .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.""" diff --git a/acestep/ui/gradio/events/results/generation_progress.py b/acestep/ui/gradio/events/results/generation_progress.py index d8935417e..ffa47f813 100644 --- a/acestep/ui/gradio/events/results/generation_progress.py +++ b/acestep/ui/gradio/events/results/generation_progress.py @@ -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, @@ -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 @@ -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("\\", "/") @@ -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,