Kids-Rhyme-Studio / audio_engine.py
CodeAntidote's picture
Fix ACE-Step ZeroGPU singing pipeline
52eeb43 verified
Raw History Blame Contribute Delete
16.2 kB
"""Generate a short song with sung vocals from the edited rhyme."""
from __future__ import annotations
import logging
import os
import re
import tempfile
import time
import traceback
from typing import Any
# IMPORTANT:
# Import spaces before torch/diffusers so Hugging Face ZeroGPU can install
# its CUDA shim before PyTorch is imported.
try:
import spaces
except ModuleNotFoundError:
if os.environ.get("SPACE_ID"):
raise
class _LocalSpaces:
@staticmethod
def GPU(**_kwargs):
return lambda fn: fn
spaces = _LocalSpaces()
import numpy as np
import soundfile as sf
from rhyme_engine import LANGUAGES, validate_lyric_for_audio
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger("kids_rhyme.audio")
MODEL_ID = "ACE-Step/acestep-v15-xl-turbo-diffusers"
VOCAL_LANGUAGE = {
key: value["code"]
for key, value in LANGUAGES.items()
}
FALLBACK_SAMPLE_RATE = 48000
DEBUG_ERRORS = os.environ.get(
"KIDS_DEBUG",
"1",
) == "1"
# ---------------------------------------------------------------------
# DO NOT create/move the pipeline to CUDA at module import time.
#
# On ZeroGPU there may be no CUDA device allocated while the module is
# importing. CUDA becomes available inside the @spaces.GPU function.
# ---------------------------------------------------------------------
_pipe = None
def _debug_detail(exc: BaseException) -> str:
if not DEBUG_ERRORS:
return ""
message = str(exc).replace("\n", " ").strip()
return (
f" [debug: {type(exc).__name__}: "
f"{message[:1000]}]"
)
def _load_pipeline():
"""
Load ACE-Step while inside the ZeroGPU allocation.
The pipeline is cached between calls when possible, but CUDA placement
is never attempted during module import.
"""
global _pipe
import torch
from diffusers import AceStepPipeline
if not torch.cuda.is_available():
raise RuntimeError(
"CUDA is unavailable inside the ZeroGPU function. "
"Check that the Space is using ZeroGPU hardware and that "
"this function is running through spaces.GPU."
)
if _pipe is None:
logger.info(
"Loading ACE-Step pipeline: %s",
MODEL_ID,
)
try:
_pipe = AceStepPipeline.from_pretrained(
MODEL_ID,
torch_dtype=torch.bfloat16,
)
except TypeError:
# Newer Diffusers versions prefer dtype.
_pipe = AceStepPipeline.from_pretrained(
MODEL_ID,
dtype=torch.bfloat16,
)
logger.info(
"ACE-Step pipeline loaded on CPU."
)
try:
_pipe.vae.enable_tiling()
logger.info("ACE-Step VAE tiling enabled.")
except Exception:
logger.info(
"VAE tiling unavailable; continuing."
)
logger.info(
"Moving ACE-Step pipeline to ZeroGPU CUDA device."
)
_pipe.to("cuda")
return _pipe
def _find_sample_rate(
pipe: Any,
result: Any = None,
) -> int:
"""
Try known locations for the model/output sample rate instead of blindly
assuming 48 kHz.
"""
candidates = []
if result is not None:
for attr in (
"sample_rate",
"sampling_rate",
"audio_sample_rate",
):
candidates.append(
getattr(result, attr, None)
)
if isinstance(result, dict):
for key in (
"sample_rate",
"sampling_rate",
"audio_sample_rate",
):
candidates.append(result.get(key))
for attr in (
"sample_rate",
"sampling_rate",
"audio_sample_rate",
):
candidates.append(
getattr(pipe, attr, None)
)
vae = getattr(pipe, "vae", None)
vae_config = getattr(vae, "config", None)
if vae_config is not None:
for attr in (
"sample_rate",
"sampling_rate",
"audio_sample_rate",
):
candidates.append(
getattr(vae_config, attr, None)
)
config = getattr(pipe, "config", None)
if config is not None:
for attr in (
"sample_rate",
"sampling_rate",
"audio_sample_rate",
):
candidates.append(
getattr(config, attr, None)
)
for rate in candidates:
if (
not isinstance(rate, (bool, np.bool_))
and isinstance(rate, (int, np.integer))
and int(rate) > 0
):
logger.info(
"Detected ACE-Step sample rate: %s Hz",
rate,
)
return int(rate)
logger.warning(
"ACE-Step did not expose a sample rate; "
"falling back to %s Hz.",
FALLBACK_SAMPLE_RATE,
)
return FALLBACK_SAMPLE_RATE
def _song_lyrics(text: str) -> str:
text = (
text.replace("\r\n", "\n")
.replace("\r", "\n")
)
stanzas = [
part.strip()
for part in re.split(r"\n\s*\n", text)
if part.strip()
]
if len(stanzas) >= 2:
first = stanzas[0]
second = "\n".join(stanzas[1:])
else:
lines = [
line.strip()
for line in text.splitlines()
if line.strip()
]
if len(lines) < 2:
raise ValueError(
"Add at least two short lines to sing."
)
midpoint = (len(lines) + 1) // 2
first = "\n".join(lines[:midpoint])
second = "\n".join(lines[midpoint:])
return (
f"[verse]\n{first}\n"
f"[chorus]\n{second}"
)
def _song_settings(
text: str,
language: str,
mood: str,
theme: str = "",
theme_prompt: str = "",
) -> dict:
text = validate_lyric_for_audio(text)
if (
language not in VOCAL_LANGUAGE
or mood not in ("Bouncy", "Calm")
):
raise ValueError(
"Write a rhyme first to select its "
"language and music mood."
)
if mood == "Calm":
prompt = (
"Gentle original children's lullaby, "
"a clear warm voice SINGING a simple "
"memorable melody in the language of the lyrics. "
"Soft piano, glockenspiel, light acoustic guitar, "
"slow swaying rhythm. Vocal-forward mix. "
"Sing the supplied lyrics; "
"no spoken words or narration."
)
bpm = 82
else:
prompt = (
"Playful original children's sing-along, "
"a clear cheerful voice SINGING "
"a simple catchy melody in the language of the lyrics. "
"Ukulele, handclaps, toy piano, bright steady beat. "
"Vocal-forward mix. "
"Sing the supplied lyrics; "
"no spoken words or narration."
)
bpm = 112
language_name = LANGUAGES[language]["name"]
prompt = (
f"Sing all vocals in {language_name}. "
+ prompt
)
if theme:
prompt += f" Song theme: {theme}."
if theme_prompt:
prompt += f" Topic: {theme_prompt}."
line_count = sum(
bool(line.strip())
for line in text.splitlines()
)
return {
"prompt": prompt,
"lyrics": _song_lyrics(text),
"vocal_language": VOCAL_LANGUAGE[language],
"audio_duration": min(
56.0,
40.0 + 4.0 * max(0, line_count - 8),
),
"num_inference_steps": 8,
"bpm": bpm,
"task_type": "text2music",
}
def _extract_audio(result: Any) -> np.ndarray:
"""
Normalize ACE-Step output into float32 [channels, samples].
Handles tensors, numpy arrays, lists/batches and common Diffusers
pipeline output containers.
"""
value = None
if hasattr(result, "audios"):
value = result.audios
elif hasattr(result, "audio"):
value = result.audio
elif isinstance(result, dict):
for key in (
"audios",
"audio",
"waveform",
"sample",
):
if key in result:
value = result[key]
break
elif isinstance(result, (tuple, list)) and result:
value = result[0]
if value is None:
raise RuntimeError(
"ACE-Step returned no audio. "
f"Result type: {type(result).__name__}"
)
# result.audios may itself be a batch/list.
if isinstance(value, (list, tuple)):
if not value:
raise RuntimeError(
"ACE-Step returned an empty audio list."
)
value = value[0]
if hasattr(value, "detach"):
value = (
value.detach()
.float()
.cpu()
.numpy()
)
arr = np.asarray(value)
logger.info(
"Raw ACE-Step audio: type=%s shape=%s dtype=%s",
type(value).__name__,
getattr(arr, "shape", None),
getattr(arr, "dtype", None),
)
arr = arr.astype(
np.float32,
copy=False,
)
# Typical batched forms:
# [batch, channels, samples]
# [batch, samples]
while arr.ndim > 2:
arr = arr[0]
if arr.ndim == 1:
arr = arr[np.newaxis, :]
if arr.ndim != 2:
raise RuntimeError(
"Unexpected ACE-Step audio shape: "
f"{arr.shape}"
)
# Normalize to [channels, samples].
#
# If first dimension clearly looks like samples and the second
# dimension is mono/stereo, transpose it.
if (
arr.shape[0] > 2
and arr.shape[1] in (1, 2)
):
arr = arr.T
if arr.shape[0] not in (1, 2):
raise RuntimeError(
"Unexpected ACE-Step channel layout: "
f"{arr.shape}"
)
if not np.isfinite(arr).all():
raise RuntimeError(
"ACE-Step produced NaN or infinite audio values."
)
peak = float(np.max(np.abs(arr)))
if peak < 0.001:
raise RuntimeError(
"ACE-Step returned silent audio."
)
return np.ascontiguousarray(arr)
def _polish(
wave: np.ndarray,
sr: int,
) -> np.ndarray:
wave = np.array(
wave,
dtype=np.float32,
copy=True,
)
peak = float(np.max(np.abs(wave)))
if peak > 0:
wave *= 0.89 / peak
fade_samples = min(
int(0.6 * sr),
wave.shape[1] // 4,
)
if fade_samples > 0:
wave[:, -fade_samples:] *= np.linspace(
1.0,
0.0,
fade_samples,
dtype=np.float32,
)
return wave
def _save_wav(
waveform: np.ndarray,
sample_rate: int,
) -> str:
"""
Save an actual WAV file and return its path.
Returning a filepath is reliable for Gradio Audio outputs and gives
the user a downloadable WAV.
"""
if waveform.ndim != 2:
raise RuntimeError(
f"Invalid waveform shape before WAV save: "
f"{waveform.shape}"
)
# soundfile expects:
# mono -> [samples]
# stereo -> [samples, channels]
if waveform.shape[0] == 1:
output = waveform[0]
else:
output = waveform.T
output = np.ascontiguousarray(
np.clip(
output,
-1.0,
1.0,
),
dtype=np.float32,
)
temp = tempfile.NamedTemporaryFile(
suffix=".wav",
delete=False,
)
path = temp.name
temp.close()
sf.write(
path,
output,
samplerate=sample_rate,
subtype="PCM_16",
format="WAV",
)
if (
not os.path.isfile(path)
or os.path.getsize(path) <= 44
):
raise RuntimeError(
"WAV file creation failed."
)
logger.info(
"Song WAV saved: %s (%d bytes)",
path,
os.path.getsize(path),
)
return path
@spaces.GPU(duration=120)
def _generate_on_gpu(
settings: dict,
) -> str:
"""
EVERYTHING requiring CUDA happens after ZeroGPU allocation.
"""
import torch
started = time.monotonic()
logger.info(
"ZeroGPU allocation entered. "
"cuda_available=%s torch=%s",
torch.cuda.is_available(),
torch.__version__,
)
if not torch.cuda.is_available():
raise RuntimeError(
"ZeroGPU allocation did not expose CUDA."
)
try:
logger.info(
"CUDA device: %s",
torch.cuda.get_device_name(0),
)
except Exception:
logger.info(
"CUDA device name unavailable."
)
pipe = _load_pipeline()
logger.info(
"Starting ACE-Step inference with settings: %r",
settings,
)
try:
with torch.inference_mode():
result = pipe(**settings)
logger.info(
"ACE-Step result type: %s",
type(result).__name__,
)
waveform = _extract_audio(result)
sample_rate = _find_sample_rate(
pipe,
result,
)
logger.info(
"ACE-Step normalized audio shape=%s "
"sample_rate=%s",
waveform.shape,
sample_rate,
)
if waveform.shape[1] < sample_rate:
raise RuntimeError(
"ACE-Step returned less than one second "
f"of audio: shape={waveform.shape}, "
f"sample_rate={sample_rate}"
)
waveform = _polish(
waveform,
sample_rate,
)
wav_path = _save_wav(
waveform,
sample_rate,
)
logger.info(
"Generated %.2f seconds of audio in %.2f seconds.",
waveform.shape[1] / sample_rate,
time.monotonic() - started,
)
return wav_path
except Exception:
# This is deliberately logger.exception rather than a generic
# "Singing failed" message. Hugging Face runtime logs will contain
# the complete traceback and original exception.
logger.exception(
"ACE-Step inference failed."
)
raise
finally:
# Do not delete _pipe here. Keeping the CPU-side object cached can
# avoid re-downloading/reconstructing it. Move it back off the
# leased ZeroGPU CUDA device before leaving the GPU scope.
if _pipe is not None:
try:
_pipe.to("cpu")
logger.info(
"ACE-Step pipeline moved back to CPU."
)
except Exception:
logger.exception(
"Could not move ACE-Step pipeline back to CPU."
)
try:
torch.cuda.empty_cache()
except Exception:
pass
def make_sung_song(
text: str,
language: str,
mood: str,
theme: str = "",
theme_prompt: str = "",
):
"""
Called by app.py.
Keep this exact five-argument signature.
"""
try:
settings = _song_settings(
text=text,
language=language,
mood=mood,
theme=theme,
theme_prompt=theme_prompt,
)
wav_path = _generate_on_gpu(
settings
)
return (
wav_path,
"Your sung song is ready to listen to and download.",
)
except Exception as exc:
# Full original traceback in Hugging Face logs.
logger.error(
"Singing generation failed with full traceback:\n%s",
traceback.format_exc(),
)
# DEBUG_ERRORS=1 also exposes the underlying exception in Gradio,
# which is useful while fixing the Space.
raise RuntimeError(
"Singing failed."
+ _debug_detail(exc)
) from exc