149 lines
4.8 KiB
Python
149 lines
4.8 KiB
Python
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)
|