from __future__ import annotations import os import threading import time import uuid from pathlib import Path import gradio as gr import soundfile as sf import torch import uvicorn from fastapi import FastAPI from heartlib import HeartMuLaGenPipeline MODEL_PATH = os.environ.get("HEARTMULA_MODEL_PATH", "/models/ckpt") OUTPUT_DIR = Path(os.environ.get("HEARTMULA_OUTPUT_DIR", "/output")) MULA_DEVICE = os.environ.get("HEARTMULA_MULA_DEVICE", "cuda:0") CODEC_DEVICE = os.environ.get("HEARTMULA_CODEC_DEVICE", "cuda:1") OUTPUT_DIR.mkdir(parents=True, exist_ok=True) pipeline: HeartMuLaGenPipeline | None = None load_error = "" generation_lock = threading.Lock() class AthenaHeartMuLaPipeline(HeartMuLaGenPipeline): """Write 24-bit WAV directly; torch 2.9 otherwise requires TorchCodec.""" def postprocess(self, model_outputs, save_path: str): frames = model_outputs["frames"].to(self.codec_device) wav = self.codec.detokenize(frames) self._unload() audio = wav.to(torch.float32).cpu().numpy().T sf.write(save_path, audio, 48_000, format="WAV", subtype="PCM_24") def load_pipeline() -> None: global pipeline, load_error try: pipeline = AthenaHeartMuLaPipeline.from_pretrained( MODEL_PATH, device={ "mula": torch.device(MULA_DEVICE), "codec": torch.device(CODEC_DEVICE), }, dtype={"mula": torch.bfloat16, "codec": torch.float32}, version="3B", lazy_load=False, ) except Exception as exc: load_error = f"{type(exc).__name__}: {exc}" raise def generate( tags: str, lyrics: str, instrumental: bool, duration_seconds: int, topk: int, temperature: float, cfg_scale: float, ): if pipeline is None: raise gr.Error(f"Modell ist nicht bereit. {load_error}".strip()) tags = ",".join(part.strip() for part in tags.split(",") if part.strip()) if not tags: raise gr.Error("Bitte mindestens ein Musik-Tag angeben.") if instrumental: effective_lyrics = "[Instrumental]" else: effective_lyrics = lyrics.strip() if not effective_lyrics: raise gr.Error("Für ein Lied mit Gesang fehlt der Liedtext.") target = OUTPUT_DIR / f"heartmula-{int(time.time())}-{uuid.uuid4().hex[:8]}.wav" with generation_lock, torch.inference_mode(): pipeline( {"lyrics": effective_lyrics, "tags": tags}, max_audio_length_ms=int(duration_seconds) * 1000, save_path=str(target), topk=int(topk), temperature=float(temperature), cfg_scale=float(cfg_scale), ) return str(target), ( f"Fertig: {target.name} · WAV 48 kHz. " "Instrumental ist bei der öffentlichen 3B-Version experimentell." ) with gr.Blocks(title="HeartMuLa 3B – Athena Test") as demo: gr.Markdown( "# HeartMuLa 3B – Athena Test\n" "Offizielles öffentliches 3B-Modell. HeartMuLa läuft auf der RTX 5080, " "HeartCodec verlustarm in FP32 auf der RTX 3060. " "Referenzaudio wird von dieser Version noch nicht unterstützt." ) tags = gr.Textbox( label="Stil und Instrumente (Komma-getrennte Tags)", value="synthwave,retrowave,1980s,melodic,nostalgic,arpeggiated synthesizer,drum machine,atmospheric,cinematic", lines=3, ) lyrics = gr.Textbox( label="Liedtext mit Abschnitten wie [Verse], [Chorus], [Bridge]", lines=14, placeholder="[Verse]\n...\n\n[Chorus]\n...", ) instrumental = gr.Checkbox(label="Instrumental (experimentell)", value=True) duration = gr.Slider(20, 240, value=60, step=5, label="Maximale Dauer in Sekunden") with gr.Accordion("Sampling – offizielle Standardwerte", open=False): topk = gr.Slider(1, 200, value=50, step=1, label="Top-k") temperature = gr.Slider(0.1, 2.0, value=1.0, step=0.05, label="Temperatur") cfg_scale = gr.Slider(1.0, 4.0, value=1.5, step=0.1, label="CFG") run = gr.Button("Musik erzeugen", variant="primary") audio = gr.Audio(label="Ergebnis", type="filepath") status = gr.Textbox(label="Status", interactive=False) run.click( generate, inputs=[tags, lyrics, instrumental, duration, topk, temperature, cfg_scale], outputs=[audio, status], ) demo.queue(default_concurrency_limit=1, max_size=4) api = FastAPI() @api.get("/healthz") def healthz(): return { "ready": pipeline is not None, "error": load_error or None, "model": "HeartMuLa-oss-3B-happy-new-year", "codec": "HeartCodec-oss-20260123", "mula_device": MULA_DEVICE, "codec_device": CODEC_DEVICE, } api = gr.mount_gradio_app(api, demo, path="/") if __name__ == "__main__": load_pipeline() uvicorn.run(api, host="0.0.0.0", port=7860)