CodeAntidote commited on
Commit
52eeb43
·
verified ·
1 Parent(s): c69e7af

Fix ACE-Step ZeroGPU singing pipeline

Browse files
Files changed (1) hide show
  1. audio_engine.py +554 -127
audio_engine.py CHANGED
@@ -5,17 +5,20 @@ from __future__ import annotations
5
  import logging
6
  import os
7
  import re
 
8
  import time
9
  import traceback
 
10
 
11
- # Import spaces before torch/diffusers so ZeroGPU can install its CUDA shim.
 
 
12
  try:
13
  import spaces
14
  except ModuleNotFoundError:
15
  if os.environ.get("SPACE_ID"):
16
- raise # A Space needs the ZeroGPU runtime and its spaces module.
17
 
18
- # Local CPU machines can still open the form and write rhymes.
19
  class _LocalSpaces:
20
  @staticmethod
21
  def GPU(**_kwargs):
@@ -23,101 +26,233 @@ except ModuleNotFoundError:
23
 
24
  spaces = _LocalSpaces()
25
 
 
26
  import numpy as np
 
 
27
  from rhyme_engine import LANGUAGES, validate_lyric_for_audio
28
 
 
29
  logging.basicConfig(level=logging.INFO)
30
  logger = logging.getLogger("kids_rhyme.audio")
31
 
 
32
  MODEL_ID = "ACE-Step/acestep-v15-xl-turbo-diffusers"
33
- VOCAL_LANGUAGE = {key: value["code"] for key, value in LANGUAGES.items()}
34
- FALLBACK_SAMPLE_RATE = 48000
35
 
36
- # Turn this on temporarily in Hugging Face Settings -> Variables and secrets:
37
- # KIDS_DEBUG=1
38
- # Leave it unset (or set it to 0) when the public app is working.
39
- DEBUG_ERRORS = os.environ.get("KIDS_DEBUG", "0") == "1"
40
 
41
- # ZeroGPU expects the model to be placed on CUDA at module import time. Its
42
- # CUDA shim makes this possible before the real GPU is allocated to a request.
43
- _pipe = None
44
- _load_error = None
45
- try:
46
- import torch
47
- except ModuleNotFoundError:
48
- torch = None
49
 
50
- if torch is not None and torch.cuda.is_available():
51
- try:
52
- from diffusers import AceStepPipeline
 
53
 
54
- # diffusers versions have used both names; support either without
55
- # changing the rest of the app.
56
- try:
57
- _pipe = AceStepPipeline.from_pretrained(MODEL_ID, dtype=torch.bfloat16)
58
- except TypeError:
59
- _pipe = AceStepPipeline.from_pretrained(MODEL_ID, torch_dtype=torch.bfloat16)
60
 
61
- _pipe = _pipe.to("cuda")
62
- logger.info("Singing model loaded: %s", MODEL_ID)
63
- except Exception as exc:
64
- _load_error = exc
65
- logger.error("Singing model failed to load:\n%s", traceback.format_exc())
 
 
 
66
 
67
 
68
  def _debug_detail(exc: BaseException) -> str:
69
- """Return a short browser-safe debug suffix only while KIDS_DEBUG=1."""
70
  if not DEBUG_ERRORS:
71
  return ""
 
72
  message = str(exc).replace("\n", " ").strip()
73
- return f" [debug: {type(exc).__name__}: {message[:500]}]"
74
 
 
 
 
 
75
 
76
- def _require_pipeline():
77
- if _pipe is None:
78
- if _load_error is not None:
79
- raise RuntimeError(
80
- "The singing model could not load. Check the Space logs."
81
- + _debug_detail(_load_error)
82
- ) from _load_error
 
 
 
 
 
 
 
83
  raise RuntimeError(
84
- "Singing needs GPU hardware. Set the Space hardware to ZeroGPU."
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
85
  )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
86
  return _pipe
87
 
88
 
89
- def _sample_rate() -> int:
90
- pipe = _require_pipeline()
91
- rate = getattr(pipe, "sample_rate", None)
92
- if rate is None:
93
- vae_cfg = getattr(getattr(pipe, "vae", None), "config", None)
94
- rate = getattr(vae_cfg, "sampling_rate", None)
95
- if rate is None:
96
- logger.warning(
97
- "No sample rate found on pipeline; using %s Hz for %s",
98
- FALLBACK_SAMPLE_RATE,
99
- MODEL_ID,
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
100
  )
101
- rate = FALLBACK_SAMPLE_RATE
102
- if isinstance(rate, (bool, np.bool_)) or not isinstance(rate, (int, np.integer)) or rate <= 0:
103
- raise ValueError(f"The pipeline returned an invalid sample rate: {rate!r}")
104
- return int(rate)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
105
 
106
 
107
  def _song_lyrics(text: str) -> str:
108
- """Preserve every edited line and mark its two parts for the song model."""
109
- text = text.replace("\r\n", "\n").replace("\r", "\n")
110
- stanzas = [part.strip() for part in re.split(r"\n\s*\n", text) if part.strip()]
 
 
 
 
 
 
 
 
111
  if len(stanzas) >= 2:
112
  first = stanzas[0]
113
  second = "\n".join(stanzas[1:])
 
114
  else:
115
- lines = [line.strip() for line in text.splitlines() if line.strip()]
 
 
 
 
 
116
  if len(lines) < 2:
117
- raise ValueError("Add at least two short lines to sing.")
 
 
 
118
  midpoint = (len(lines) + 1) // 2
119
- first, second = "\n".join(lines[:midpoint]), "\n".join(lines[midpoint:])
120
- return f"[verse]\n{first}\n[chorus]\n{second}"
 
 
 
 
 
 
121
 
122
 
123
  def _song_settings(
@@ -128,130 +263,398 @@ def _song_settings(
128
  theme_prompt: str = "",
129
  ) -> dict:
130
  text = validate_lyric_for_audio(text)
131
- if language not in VOCAL_LANGUAGE or mood not in ("Bouncy", "Calm"):
132
- raise ValueError("Write a rhyme first to select its language and music mood.")
 
 
 
 
 
 
 
133
 
134
  if mood == "Calm":
135
  prompt = (
136
- "Gentle original children's lullaby, a clear warm voice SINGING a simple "
137
- "memorable melody in the language of the lyrics. Soft piano, glockenspiel, "
138
- "light acoustic guitar, slow swaying rhythm. Vocal-forward mix. "
139
- "Sing the supplied lyrics; no spoken words or narration."
 
 
 
140
  )
141
  bpm = 82
 
142
  else:
143
  prompt = (
144
- "Playful original children's sing-along, a clear cheerful voice SINGING "
145
- "a simple catchy melody in the language of the lyrics. Ukulele, handclaps, "
146
- "toy piano, bright steady beat. Vocal-forward mix. "
147
- "Sing the supplied lyrics; no spoken words or narration."
 
 
 
148
  )
149
  bpm = 112
150
 
151
  language_name = LANGUAGES[language]["name"]
152
- prompt = f"Sing all vocals in {language_name}. " + prompt
 
 
 
 
 
153
  if theme:
154
  prompt += f" Song theme: {theme}."
 
155
  if theme_prompt:
156
  prompt += f" Topic: {theme_prompt}."
157
 
158
- line_count = sum(bool(line.strip()) for line in text.splitlines())
 
 
 
 
159
  return {
160
  "prompt": prompt,
161
  "lyrics": _song_lyrics(text),
162
  "vocal_language": VOCAL_LANGUAGE[language],
163
- # 40 seconds for a normal rhyme, with a little extra room for longer text.
164
- "audio_duration": min(56.0, 40.0 + 4.0 * max(0, line_count - 8)),
 
 
165
  "num_inference_steps": 8,
166
  "bpm": bpm,
167
  "task_type": "text2music",
168
- "output_type": "np",
169
  }
170
 
171
 
172
- def _extract_audio(result) -> np.ndarray:
173
- """Extract ACE-Step audio and normalize it to [channels, samples]."""
 
 
 
 
 
 
174
  value = None
175
 
176
- # Current diffusers ACE-Step output uses `audios`.
177
  if hasattr(result, "audios"):
178
  value = result.audios
 
179
  elif hasattr(result, "audio"):
180
  value = result.audio
 
181
  elif isinstance(result, dict):
182
- for key in ("audios", "audio", "waveform", "sample"):
 
 
 
 
 
183
  if key in result:
184
  value = result[key]
185
  break
 
186
  elif isinstance(result, (tuple, list)) and result:
187
  value = result[0]
188
 
189
  if value is None:
190
- raise ValueError(f"No audio was returned by {type(result).__name__}.")
 
 
 
 
 
 
 
 
 
 
 
191
 
192
  if hasattr(value, "detach"):
193
- value = value.detach().float().cpu().numpy()
 
 
 
 
 
 
 
 
 
 
 
 
 
 
194
 
195
- arr = np.asarray(value, dtype=np.float32)
 
 
 
196
 
197
- # Typical result: [batch, channels, samples]. Take first batch.
 
 
198
  while arr.ndim > 2:
199
  arr = arr[0]
 
200
  if arr.ndim == 1:
201
- arr = arr[None, :]
 
202
  if arr.ndim != 2:
203
- raise ValueError(f"Unexpected audio rank/shape {arr.shape}.")
 
 
 
204
 
205
- # If it came back [samples, channels], transpose it.
206
- if arr.shape[0] > 2 and arr.shape[1] in (1, 2):
 
 
 
 
 
 
207
  arr = arr.T
208
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
209
  return np.ascontiguousarray(arr)
210
 
211
 
212
- def _polish(wave: np.ndarray, sr: int) -> np.ndarray:
213
- """Peak-normalize and fade the tail so playback ends cleanly."""
214
- wave = np.array(wave, dtype=np.float32, copy=True)
 
 
 
 
 
 
 
215
  peak = float(np.max(np.abs(wave)))
 
216
  if peak > 0:
217
  wave *= 0.89 / peak
218
- n = min(int(0.6 * sr), wave.shape[1] // 4)
219
- if n > 0:
220
- wave[:, -n:] *= np.linspace(1.0, 0.0, n, dtype=np.float32)
 
 
 
 
 
 
 
 
 
 
 
221
  return wave
222
 
223
 
224
- def generate_song(settings: dict) -> tuple[int, np.ndarray]:
225
- """Inference core. The caller must provide GPU access."""
226
- pipe = _require_pipeline()
227
- started = time.monotonic()
 
 
228
 
229
- result = pipe(**settings)
230
- waveform = _extract_audio(result)
231
- sr = _sample_rate()
232
 
233
- logger.info("ACE-Step returned audio shape %s at %s Hz", waveform.shape, sr)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
234
 
235
- if waveform.shape[0] not in (1, 2) or waveform.shape[1] < sr:
236
- raise ValueError(f"Unexpected audio shape {waveform.shape} at {sr} Hz")
237
- if not np.isfinite(waveform).all() or float(np.max(np.abs(waveform))) < 0.005:
238
- raise ValueError("Model returned empty or invalid audio")
 
 
 
239
 
240
- waveform = _polish(waveform, sr)
241
- pcm = (np.clip(waveform.T, -1.0, 1.0) * 32767).astype(np.int16)
 
 
 
 
 
 
 
 
 
 
 
 
 
242
 
243
  logger.info(
244
- "Generated %.1fs of audio at %s Hz in %.1fs",
245
- len(pcm) / sr,
246
- sr,
247
- time.monotonic() - started,
248
  )
249
- return sr, pcm
 
250
 
251
 
252
  @spaces.GPU(duration=120)
253
- def _generate_on_gpu(settings: dict) -> tuple[int, np.ndarray]:
254
- return generate_song(settings)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
255
 
256
 
257
  def make_sung_song(
@@ -260,17 +663,41 @@ def make_sung_song(
260
  mood: str,
261
  theme: str = "",
262
  theme_prompt: str = "",
263
- ) -> tuple[tuple[int, np.ndarray], str]:
264
- """Validate first, then generate. Keep the app's five-argument API."""
265
- settings = _song_settings(text, language, mood, theme, theme_prompt)
266
- _require_pipeline()
 
 
267
 
268
  try:
269
- audio = _generate_on_gpu(settings)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
270
  except Exception as exc:
271
- logger.error("Song generation failed:\n%s", traceback.format_exc())
272
- raise RuntimeError(
273
- "Singing failed this time. Please try again." + _debug_detail(exc)
274
- ) from exc
 
275
 
276
- return audio, "Your sung song is ready to listen to and download."
 
 
 
 
 
 
5
  import logging
6
  import os
7
  import re
8
+ import tempfile
9
  import time
10
  import traceback
11
+ from typing import Any
12
 
13
+ # IMPORTANT:
14
+ # Import spaces before torch/diffusers so Hugging Face ZeroGPU can install
15
+ # its CUDA shim before PyTorch is imported.
16
  try:
17
  import spaces
18
  except ModuleNotFoundError:
19
  if os.environ.get("SPACE_ID"):
20
+ raise
21
 
 
22
  class _LocalSpaces:
23
  @staticmethod
24
  def GPU(**_kwargs):
 
26
 
27
  spaces = _LocalSpaces()
28
 
29
+
30
  import numpy as np
31
+ import soundfile as sf
32
+
33
  from rhyme_engine import LANGUAGES, validate_lyric_for_audio
34
 
35
+
36
  logging.basicConfig(level=logging.INFO)
37
  logger = logging.getLogger("kids_rhyme.audio")
38
 
39
+
40
  MODEL_ID = "ACE-Step/acestep-v15-xl-turbo-diffusers"
 
 
41
 
42
+ VOCAL_LANGUAGE = {
43
+ key: value["code"]
44
+ for key, value in LANGUAGES.items()
45
+ }
46
 
47
+ FALLBACK_SAMPLE_RATE = 48000
 
 
 
 
 
 
 
48
 
49
+ DEBUG_ERRORS = os.environ.get(
50
+ "KIDS_DEBUG",
51
+ "1",
52
+ ) == "1"
53
 
 
 
 
 
 
 
54
 
55
+ # ---------------------------------------------------------------------
56
+ # DO NOT create/move the pipeline to CUDA at module import time.
57
+ #
58
+ # On ZeroGPU there may be no CUDA device allocated while the module is
59
+ # importing. CUDA becomes available inside the @spaces.GPU function.
60
+ # ---------------------------------------------------------------------
61
+
62
+ _pipe = None
63
 
64
 
65
  def _debug_detail(exc: BaseException) -> str:
 
66
  if not DEBUG_ERRORS:
67
  return ""
68
+
69
  message = str(exc).replace("\n", " ").strip()
 
70
 
71
+ return (
72
+ f" [debug: {type(exc).__name__}: "
73
+ f"{message[:1000]}]"
74
+ )
75
 
76
+
77
+ def _load_pipeline():
78
+ """
79
+ Load ACE-Step while inside the ZeroGPU allocation.
80
+
81
+ The pipeline is cached between calls when possible, but CUDA placement
82
+ is never attempted during module import.
83
+ """
84
+ global _pipe
85
+
86
+ import torch
87
+ from diffusers import AceStepPipeline
88
+
89
+ if not torch.cuda.is_available():
90
  raise RuntimeError(
91
+ "CUDA is unavailable inside the ZeroGPU function. "
92
+ "Check that the Space is using ZeroGPU hardware and that "
93
+ "this function is running through spaces.GPU."
94
+ )
95
+
96
+ if _pipe is None:
97
+ logger.info(
98
+ "Loading ACE-Step pipeline: %s",
99
+ MODEL_ID,
100
+ )
101
+
102
+ try:
103
+ _pipe = AceStepPipeline.from_pretrained(
104
+ MODEL_ID,
105
+ torch_dtype=torch.bfloat16,
106
+ )
107
+ except TypeError:
108
+ # Newer Diffusers versions prefer dtype.
109
+ _pipe = AceStepPipeline.from_pretrained(
110
+ MODEL_ID,
111
+ dtype=torch.bfloat16,
112
+ )
113
+
114
+ logger.info(
115
+ "ACE-Step pipeline loaded on CPU."
116
  )
117
+
118
+ try:
119
+ _pipe.vae.enable_tiling()
120
+ logger.info("ACE-Step VAE tiling enabled.")
121
+ except Exception:
122
+ logger.info(
123
+ "VAE tiling unavailable; continuing."
124
+ )
125
+
126
+ logger.info(
127
+ "Moving ACE-Step pipeline to ZeroGPU CUDA device."
128
+ )
129
+
130
+ _pipe.to("cuda")
131
+
132
  return _pipe
133
 
134
 
135
+ def _find_sample_rate(
136
+ pipe: Any,
137
+ result: Any = None,
138
+ ) -> int:
139
+ """
140
+ Try known locations for the model/output sample rate instead of blindly
141
+ assuming 48 kHz.
142
+ """
143
+
144
+ candidates = []
145
+
146
+ if result is not None:
147
+ for attr in (
148
+ "sample_rate",
149
+ "sampling_rate",
150
+ "audio_sample_rate",
151
+ ):
152
+ candidates.append(
153
+ getattr(result, attr, None)
154
+ )
155
+
156
+ if isinstance(result, dict):
157
+ for key in (
158
+ "sample_rate",
159
+ "sampling_rate",
160
+ "audio_sample_rate",
161
+ ):
162
+ candidates.append(result.get(key))
163
+
164
+ for attr in (
165
+ "sample_rate",
166
+ "sampling_rate",
167
+ "audio_sample_rate",
168
+ ):
169
+ candidates.append(
170
+ getattr(pipe, attr, None)
171
  )
172
+
173
+ vae = getattr(pipe, "vae", None)
174
+ vae_config = getattr(vae, "config", None)
175
+
176
+ if vae_config is not None:
177
+ for attr in (
178
+ "sample_rate",
179
+ "sampling_rate",
180
+ "audio_sample_rate",
181
+ ):
182
+ candidates.append(
183
+ getattr(vae_config, attr, None)
184
+ )
185
+
186
+ config = getattr(pipe, "config", None)
187
+
188
+ if config is not None:
189
+ for attr in (
190
+ "sample_rate",
191
+ "sampling_rate",
192
+ "audio_sample_rate",
193
+ ):
194
+ candidates.append(
195
+ getattr(config, attr, None)
196
+ )
197
+
198
+ for rate in candidates:
199
+ if (
200
+ not isinstance(rate, (bool, np.bool_))
201
+ and isinstance(rate, (int, np.integer))
202
+ and int(rate) > 0
203
+ ):
204
+ logger.info(
205
+ "Detected ACE-Step sample rate: %s Hz",
206
+ rate,
207
+ )
208
+ return int(rate)
209
+
210
+ logger.warning(
211
+ "ACE-Step did not expose a sample rate; "
212
+ "falling back to %s Hz.",
213
+ FALLBACK_SAMPLE_RATE,
214
+ )
215
+
216
+ return FALLBACK_SAMPLE_RATE
217
 
218
 
219
  def _song_lyrics(text: str) -> str:
220
+ text = (
221
+ text.replace("\r\n", "\n")
222
+ .replace("\r", "\n")
223
+ )
224
+
225
+ stanzas = [
226
+ part.strip()
227
+ for part in re.split(r"\n\s*\n", text)
228
+ if part.strip()
229
+ ]
230
+
231
  if len(stanzas) >= 2:
232
  first = stanzas[0]
233
  second = "\n".join(stanzas[1:])
234
+
235
  else:
236
+ lines = [
237
+ line.strip()
238
+ for line in text.splitlines()
239
+ if line.strip()
240
+ ]
241
+
242
  if len(lines) < 2:
243
+ raise ValueError(
244
+ "Add at least two short lines to sing."
245
+ )
246
+
247
  midpoint = (len(lines) + 1) // 2
248
+
249
+ first = "\n".join(lines[:midpoint])
250
+ second = "\n".join(lines[midpoint:])
251
+
252
+ return (
253
+ f"[verse]\n{first}\n"
254
+ f"[chorus]\n{second}"
255
+ )
256
 
257
 
258
  def _song_settings(
 
263
  theme_prompt: str = "",
264
  ) -> dict:
265
  text = validate_lyric_for_audio(text)
266
+
267
+ if (
268
+ language not in VOCAL_LANGUAGE
269
+ or mood not in ("Bouncy", "Calm")
270
+ ):
271
+ raise ValueError(
272
+ "Write a rhyme first to select its "
273
+ "language and music mood."
274
+ )
275
 
276
  if mood == "Calm":
277
  prompt = (
278
+ "Gentle original children's lullaby, "
279
+ "a clear warm voice SINGING a simple "
280
+ "memorable melody in the language of the lyrics. "
281
+ "Soft piano, glockenspiel, light acoustic guitar, "
282
+ "slow swaying rhythm. Vocal-forward mix. "
283
+ "Sing the supplied lyrics; "
284
+ "no spoken words or narration."
285
  )
286
  bpm = 82
287
+
288
  else:
289
  prompt = (
290
+ "Playful original children's sing-along, "
291
+ "a clear cheerful voice SINGING "
292
+ "a simple catchy melody in the language of the lyrics. "
293
+ "Ukulele, handclaps, toy piano, bright steady beat. "
294
+ "Vocal-forward mix. "
295
+ "Sing the supplied lyrics; "
296
+ "no spoken words or narration."
297
  )
298
  bpm = 112
299
 
300
  language_name = LANGUAGES[language]["name"]
301
+
302
+ prompt = (
303
+ f"Sing all vocals in {language_name}. "
304
+ + prompt
305
+ )
306
+
307
  if theme:
308
  prompt += f" Song theme: {theme}."
309
+
310
  if theme_prompt:
311
  prompt += f" Topic: {theme_prompt}."
312
 
313
+ line_count = sum(
314
+ bool(line.strip())
315
+ for line in text.splitlines()
316
+ )
317
+
318
  return {
319
  "prompt": prompt,
320
  "lyrics": _song_lyrics(text),
321
  "vocal_language": VOCAL_LANGUAGE[language],
322
+ "audio_duration": min(
323
+ 56.0,
324
+ 40.0 + 4.0 * max(0, line_count - 8),
325
+ ),
326
  "num_inference_steps": 8,
327
  "bpm": bpm,
328
  "task_type": "text2music",
 
329
  }
330
 
331
 
332
+ def _extract_audio(result: Any) -> np.ndarray:
333
+ """
334
+ Normalize ACE-Step output into float32 [channels, samples].
335
+
336
+ Handles tensors, numpy arrays, lists/batches and common Diffusers
337
+ pipeline output containers.
338
+ """
339
+
340
  value = None
341
 
 
342
  if hasattr(result, "audios"):
343
  value = result.audios
344
+
345
  elif hasattr(result, "audio"):
346
  value = result.audio
347
+
348
  elif isinstance(result, dict):
349
+ for key in (
350
+ "audios",
351
+ "audio",
352
+ "waveform",
353
+ "sample",
354
+ ):
355
  if key in result:
356
  value = result[key]
357
  break
358
+
359
  elif isinstance(result, (tuple, list)) and result:
360
  value = result[0]
361
 
362
  if value is None:
363
+ raise RuntimeError(
364
+ "ACE-Step returned no audio. "
365
+ f"Result type: {type(result).__name__}"
366
+ )
367
+
368
+ # result.audios may itself be a batch/list.
369
+ if isinstance(value, (list, tuple)):
370
+ if not value:
371
+ raise RuntimeError(
372
+ "ACE-Step returned an empty audio list."
373
+ )
374
+ value = value[0]
375
 
376
  if hasattr(value, "detach"):
377
+ value = (
378
+ value.detach()
379
+ .float()
380
+ .cpu()
381
+ .numpy()
382
+ )
383
+
384
+ arr = np.asarray(value)
385
+
386
+ logger.info(
387
+ "Raw ACE-Step audio: type=%s shape=%s dtype=%s",
388
+ type(value).__name__,
389
+ getattr(arr, "shape", None),
390
+ getattr(arr, "dtype", None),
391
+ )
392
 
393
+ arr = arr.astype(
394
+ np.float32,
395
+ copy=False,
396
+ )
397
 
398
+ # Typical batched forms:
399
+ # [batch, channels, samples]
400
+ # [batch, samples]
401
  while arr.ndim > 2:
402
  arr = arr[0]
403
+
404
  if arr.ndim == 1:
405
+ arr = arr[np.newaxis, :]
406
+
407
  if arr.ndim != 2:
408
+ raise RuntimeError(
409
+ "Unexpected ACE-Step audio shape: "
410
+ f"{arr.shape}"
411
+ )
412
 
413
+ # Normalize to [channels, samples].
414
+ #
415
+ # If first dimension clearly looks like samples and the second
416
+ # dimension is mono/stereo, transpose it.
417
+ if (
418
+ arr.shape[0] > 2
419
+ and arr.shape[1] in (1, 2)
420
+ ):
421
  arr = arr.T
422
 
423
+ if arr.shape[0] not in (1, 2):
424
+ raise RuntimeError(
425
+ "Unexpected ACE-Step channel layout: "
426
+ f"{arr.shape}"
427
+ )
428
+
429
+ if not np.isfinite(arr).all():
430
+ raise RuntimeError(
431
+ "ACE-Step produced NaN or infinite audio values."
432
+ )
433
+
434
+ peak = float(np.max(np.abs(arr)))
435
+
436
+ if peak < 0.001:
437
+ raise RuntimeError(
438
+ "ACE-Step returned silent audio."
439
+ )
440
+
441
  return np.ascontiguousarray(arr)
442
 
443
 
444
+ def _polish(
445
+ wave: np.ndarray,
446
+ sr: int,
447
+ ) -> np.ndarray:
448
+ wave = np.array(
449
+ wave,
450
+ dtype=np.float32,
451
+ copy=True,
452
+ )
453
+
454
  peak = float(np.max(np.abs(wave)))
455
+
456
  if peak > 0:
457
  wave *= 0.89 / peak
458
+
459
+ fade_samples = min(
460
+ int(0.6 * sr),
461
+ wave.shape[1] // 4,
462
+ )
463
+
464
+ if fade_samples > 0:
465
+ wave[:, -fade_samples:] *= np.linspace(
466
+ 1.0,
467
+ 0.0,
468
+ fade_samples,
469
+ dtype=np.float32,
470
+ )
471
+
472
  return wave
473
 
474
 
475
+ def _save_wav(
476
+ waveform: np.ndarray,
477
+ sample_rate: int,
478
+ ) -> str:
479
+ """
480
+ Save an actual WAV file and return its path.
481
 
482
+ Returning a filepath is reliable for Gradio Audio outputs and gives
483
+ the user a downloadable WAV.
484
+ """
485
 
486
+ if waveform.ndim != 2:
487
+ raise RuntimeError(
488
+ f"Invalid waveform shape before WAV save: "
489
+ f"{waveform.shape}"
490
+ )
491
+
492
+ # soundfile expects:
493
+ # mono -> [samples]
494
+ # stereo -> [samples, channels]
495
+ if waveform.shape[0] == 1:
496
+ output = waveform[0]
497
+ else:
498
+ output = waveform.T
499
+
500
+ output = np.ascontiguousarray(
501
+ np.clip(
502
+ output,
503
+ -1.0,
504
+ 1.0,
505
+ ),
506
+ dtype=np.float32,
507
+ )
508
 
509
+ temp = tempfile.NamedTemporaryFile(
510
+ suffix=".wav",
511
+ delete=False,
512
+ )
513
+
514
+ path = temp.name
515
+ temp.close()
516
 
517
+ sf.write(
518
+ path,
519
+ output,
520
+ samplerate=sample_rate,
521
+ subtype="PCM_16",
522
+ format="WAV",
523
+ )
524
+
525
+ if (
526
+ not os.path.isfile(path)
527
+ or os.path.getsize(path) <= 44
528
+ ):
529
+ raise RuntimeError(
530
+ "WAV file creation failed."
531
+ )
532
 
533
  logger.info(
534
+ "Song WAV saved: %s (%d bytes)",
535
+ path,
536
+ os.path.getsize(path),
 
537
  )
538
+
539
+ return path
540
 
541
 
542
  @spaces.GPU(duration=120)
543
+ def _generate_on_gpu(
544
+ settings: dict,
545
+ ) -> str:
546
+ """
547
+ EVERYTHING requiring CUDA happens after ZeroGPU allocation.
548
+ """
549
+
550
+ import torch
551
+
552
+ started = time.monotonic()
553
+
554
+ logger.info(
555
+ "ZeroGPU allocation entered. "
556
+ "cuda_available=%s torch=%s",
557
+ torch.cuda.is_available(),
558
+ torch.__version__,
559
+ )
560
+
561
+ if not torch.cuda.is_available():
562
+ raise RuntimeError(
563
+ "ZeroGPU allocation did not expose CUDA."
564
+ )
565
+
566
+ try:
567
+ logger.info(
568
+ "CUDA device: %s",
569
+ torch.cuda.get_device_name(0),
570
+ )
571
+ except Exception:
572
+ logger.info(
573
+ "CUDA device name unavailable."
574
+ )
575
+
576
+ pipe = _load_pipeline()
577
+
578
+ logger.info(
579
+ "Starting ACE-Step inference with settings: %r",
580
+ settings,
581
+ )
582
+
583
+ try:
584
+ with torch.inference_mode():
585
+ result = pipe(**settings)
586
+
587
+ logger.info(
588
+ "ACE-Step result type: %s",
589
+ type(result).__name__,
590
+ )
591
+
592
+ waveform = _extract_audio(result)
593
+ sample_rate = _find_sample_rate(
594
+ pipe,
595
+ result,
596
+ )
597
+
598
+ logger.info(
599
+ "ACE-Step normalized audio shape=%s "
600
+ "sample_rate=%s",
601
+ waveform.shape,
602
+ sample_rate,
603
+ )
604
+
605
+ if waveform.shape[1] < sample_rate:
606
+ raise RuntimeError(
607
+ "ACE-Step returned less than one second "
608
+ f"of audio: shape={waveform.shape}, "
609
+ f"sample_rate={sample_rate}"
610
+ )
611
+
612
+ waveform = _polish(
613
+ waveform,
614
+ sample_rate,
615
+ )
616
+
617
+ wav_path = _save_wav(
618
+ waveform,
619
+ sample_rate,
620
+ )
621
+
622
+ logger.info(
623
+ "Generated %.2f seconds of audio in %.2f seconds.",
624
+ waveform.shape[1] / sample_rate,
625
+ time.monotonic() - started,
626
+ )
627
+
628
+ return wav_path
629
+
630
+ except Exception:
631
+ # This is deliberately logger.exception rather than a generic
632
+ # "Singing failed" message. Hugging Face runtime logs will contain
633
+ # the complete traceback and original exception.
634
+ logger.exception(
635
+ "ACE-Step inference failed."
636
+ )
637
+ raise
638
+
639
+ finally:
640
+ # Do not delete _pipe here. Keeping the CPU-side object cached can
641
+ # avoid re-downloading/reconstructing it. Move it back off the
642
+ # leased ZeroGPU CUDA device before leaving the GPU scope.
643
+ if _pipe is not None:
644
+ try:
645
+ _pipe.to("cpu")
646
+ logger.info(
647
+ "ACE-Step pipeline moved back to CPU."
648
+ )
649
+ except Exception:
650
+ logger.exception(
651
+ "Could not move ACE-Step pipeline back to CPU."
652
+ )
653
+
654
+ try:
655
+ torch.cuda.empty_cache()
656
+ except Exception:
657
+ pass
658
 
659
 
660
  def make_sung_song(
 
663
  mood: str,
664
  theme: str = "",
665
  theme_prompt: str = "",
666
+ ):
667
+ """
668
+ Called by app.py.
669
+
670
+ Keep this exact five-argument signature.
671
+ """
672
 
673
  try:
674
+ settings = _song_settings(
675
+ text=text,
676
+ language=language,
677
+ mood=mood,
678
+ theme=theme,
679
+ theme_prompt=theme_prompt,
680
+ )
681
+
682
+ wav_path = _generate_on_gpu(
683
+ settings
684
+ )
685
+
686
+ return (
687
+ wav_path,
688
+ "Your sung song is ready to listen to and download.",
689
+ )
690
+
691
  except Exception as exc:
692
+ # Full original traceback in Hugging Face logs.
693
+ logger.error(
694
+ "Singing generation failed with full traceback:\n%s",
695
+ traceback.format_exc(),
696
+ )
697
 
698
+ # DEBUG_ERRORS=1 also exposes the underlying exception in Gradio,
699
+ # which is useful while fixing the Space.
700
+ raise RuntimeError(
701
+ "Singing failed."
702
+ + _debug_detail(exc)
703
+ ) from exc