Add isolated HeartMuLa 3B quality test
This commit is contained in:
@@ -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)
|
||||
Reference in New Issue
Block a user