Add isolated HeartMuLa 3B quality test

This commit is contained in:
Mikei386
2026-09-10 20:20:51 +02:00
parent 7c95dab324
commit fe2a93eeb1
5 changed files with 253 additions and 1 deletions
+148
View File
@@ -0,0 +1,148 @@
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)