# -*- coding: utf-8 -*- """Lucy-Stimme v2: F5-TTS (deutsch) via ONNX Runtime + DirectML (9070 XT, nativ Windows, kein ROCm). Non-autoregressiv -> keine Kollaps-/Wiederhol-/Männerstimmen-Fehler wie pocket. Satzweises Streaming. Modelliert nach dem verifizierten DML-Benchmark (bench_dml.py) + Export_F5-Preprocessing + pocket_server-Struktur.""" import os, io, re, time, threading, logging import numpy as np, soundfile as sf, librosa, jieba, torch import onnxruntime as ort from pypinyin import lazy_pinyin, Style from fastapi import FastAPI from fastapi.responses import Response, JSONResponse, StreamingResponse from pydantic import BaseModel from contextlib import asynccontextmanager log = logging.getLogger("lucy-f5"); logging.basicConfig(level=logging.INFO) BASE = os.path.dirname(os.path.abspath(__file__)) ONNX_DIR = os.environ.get("LUCY_F5_ONNX", os.path.join(BASE, "onnx_de")) VOCAB = os.environ.get("LUCY_F5_VOCAB", os.path.join(BASE, "vocab.txt")) REF_WAV = os.environ.get("LUCY_F5_REF", os.path.join(BASE, "lucy_ref.wav")) REF_TXT = os.environ.get("LUCY_F5_REF_TXT", os.path.join(BASE, "lucy_ref.txt")) PROVIDER = os.environ.get("LUCY_F5_PROVIDER", "DmlExecutionProvider") NFE_STEP = int(os.environ.get("LUCY_F5_NFE", "32")) # MUSS zum Export passen (Zeitplan ist eingebacken) TARGET_RMS = float(os.environ.get("LUCY_TARGET_RMS", "0.09")) SR = 24000; HOP_LENGTH = 256 STATE, LOCK = {}, threading.Lock() # ---- Text-Preprocessing (aus Export_F5.py; für Deutsch laufen Nicht-CJK-Zeichen einfach durch) ---- def _load_vocab(path): m = {} with open(path, "r", encoding="utf-8") as f: for i, ch in enumerate(f): m[ch[:-1]] = i return m def convert_char_to_pinyin(text_list, polyphone=True): if jieba.dt.initialized is False: jieba.default_logger.setLevel(50); jieba.initialize() out, trans = [], str.maketrans({";": ",", "“": '"', "”": '"', "‘": "'", "’": "'"}) def is_zh(c): return "㄀" <= c <= "鿿" for text in text_list: cl = []; text = text.translate(trans) for seg in jieba.cut(text): blen = len(bytes(seg, "UTF-8")) if blen == len(seg): if cl and blen > 1 and cl[-1] not in " :'\"": cl.append(" ") cl.extend(seg) elif polyphone and blen == 3 * len(seg): pin = lazy_pinyin(seg, style=Style.TONE3, tone_sandhi=True) for i, c in enumerate(seg): if is_zh(c): cl.append(" ") cl.append(pin[i]) else: for c in seg: if ord(c) < 256: cl.extend(c) elif is_zh(c): cl.append(" "); cl.extend(lazy_pinyin(c, style=Style.TONE3, tone_sandhi=True)) else: cl.append(c) out.append(cl) return out def list_str_to_idx(text, vocab_map, padding_value=-1): get = vocab_map.get tensors = [torch.tensor([get(c, 0) for c in t], dtype=torch.int32) for t in text] return torch.nn.utils.rnn.pad_sequence(tensors, padding_value=padding_value, batch_first=True).numpy() _ZH_PUNC = r"。,、;:?!" def _text_len(s): return len(s.encode("utf-8")) + 3 * len(re.findall(_ZH_PUNC, s)) # ---- Satz-Splitter (wie pocket) ---- _SENT_RX = re.compile(r".+?(?:[.!?…]+(?:\s|$)|$)", re.S) def split_sentences(text, min_len=30): parts = [m.group(0).strip() for m in _SENT_RX.finditer(text.strip())] out = [] for p in parts: if not p: continue if out and len(out[-1]) < min_len: out[-1] = f"{out[-1]} {p}" else: out.append(p) return out or [text.strip()] def cleanup(a, sr): """F5-Output putzen: Stille-Trim hinten, RMS-Norm auf TARGET_RMS, Peak-Clamp, 80ms-Pads.""" a = np.asarray(a, dtype=np.float32).reshape(-1) if a.size == 0: return a rev, _ = librosa.effects.trim(a[::-1], top_db=40); a = rev[::-1] if rev.size else a yt, _ = librosa.effects.trim(a, top_db=40); a = yt if yt.size else a rms = float(np.sqrt(np.mean(a ** 2))) or 1e-9 a = a * (TARGET_RMS / rms) peak = float(np.max(np.abs(a))) if peak > 0.95: a = a * (0.95 / peak) fi = min(int(0.008 * sr), a.size // 2) if fi > 0: a[:fi] *= np.linspace(0., 1., fi, dtype=np.float32); a[-fi:] *= np.linspace(1., 0., fi, dtype=np.float32) pad = np.zeros(int(0.08 * sr), dtype=np.float32) return np.concatenate([pad, a, pad]) def _to_pcm16(a): a = np.asarray(a, dtype=np.float32).reshape(-1) np.clip(a, -0.95, 0.95, out=a) return (a * 32767.0).astype(" np.ndarray: s = STATE ref_text = s["ref_text"] rt_len = _text_len(ref_text); gt_len = max(_text_len(gen_text), 1) ref_audio_len = s["ref_audio"].shape[-1] // HOP_LENGTH + 1 max_duration = np.array([ref_audio_len + int(ref_audio_len / rt_len * gt_len)], dtype=np.int64) text = convert_char_to_pinyin([ref_text + gen_text]) text_ids = list_str_to_idx(text, s["vocab"]) A = s["A"].run(s["A_out"], {s["A_in"][0]: s["ref_audio"], s["A_in"][1]: text_ids, s["A_in"][2]: max_duration}) noise, rcq, rsq, rck, rsk, cmt, cmtd, ref_signal_len = A dev = s["dev"] if dev: # DirectML/CUDA: io_binding, Tensoren GPU-resident über die NFE-Schleife ts = np.array([0], dtype=np.int32) ins = [ort.OrtValue.ortvalue_from_numpy(x, dev, 0) for x in (noise, rcq, rsq, rck, rsk, cmt, cmtd, ts)] outs = [ins[0], ins[-1]] iob = s["B"].io_binding() for i in range(len(ins)): iob.bind_ortvalue_input(name=s["B_in"][i], ortvalue=ins[i]) for i in range(len(outs)): iob.bind_ortvalue_output(name=s["B_out"][i], ortvalue=outs[i]) for _ in range(0, NFE_STEP, 1): s["B"].run_with_iobinding(iob) noise = ort.OrtValue.numpy(iob.get_outputs()[0]) else: ts = np.array([0], dtype=np.int32) for _ in range(0, NFE_STEP - 1, 1): noise, ts = s["B"].run(s["B_out"], {s["B_in"][0]: noise, s["B_in"][1]: rcq, s["B_in"][2]: rsq, s["B_in"][3]: rck, s["B_in"][4]: rsk, s["B_in"][5]: cmt, s["B_in"][6]: cmtd, s["B_in"][7]: ts}) out = s["C"].run([s["C_out"]], {s["C_in"][0]: noise, s["C_in"][1]: ref_signal_len})[0] a = np.asarray(out).reshape(-1).astype(np.float32) if a.dtype != np.float32 or np.max(np.abs(a)) > 1.5: # int16-Decoder -> auf float a = a / 32768.0 return a @asynccontextmanager async def lifespan(app): t0 = time.time(); log.info("Lade F5 ONNX (%s) ...", PROVIDER) so = ort.SessionOptions(); so.log_severity_level = 4 so.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL A = ort.InferenceSession(os.path.join(ONNX_DIR, "F5_Preprocess.onnx"), so, providers=["CPUExecutionProvider"]) B = ort.InferenceSession(os.path.join(ONNX_DIR, "F5_Transformer.onnx"), so, providers=[PROVIDER]) C = ort.InferenceSession(os.path.join(ONNX_DIR, "F5_Decode.onnx"), so, providers=["CPUExecutionProvider"]) prov = B.get_providers()[0] dev = "dml" if "Dml" in prov else ("cuda" if "CUDA" in prov or "Tensorrt" in prov else None) # Lucy-Referenz als int16 laden (Decoder/Preprocess erwartet int16-Pfad) ref, _sr = sf.read(REF_WAV, dtype="float32", always_2d=False) ref = np.asarray(ref, dtype=np.float32).reshape(-1) if _sr != SR: ref = librosa.resample(ref, orig_sr=_sr, target_sr=SR) mx = np.max(np.abs(ref)) or 1.0 ref_i16 = (ref * (32767.0 / mx)).astype(np.int16).reshape(1, 1, -1) STATE.update( A=A, B=B, C=C, dev=dev, prov=prov, A_in=[i.name for i in A.get_inputs()], A_out=[o.name for o in A.get_outputs()], B_in=[i.name for i in B.get_inputs()], B_out=[o.name for o in B.get_outputs()], C_in=[i.name for i in C.get_inputs()], C_out=C.get_outputs()[0].name, vocab=_load_vocab(VOCAB), ref_audio=ref_i16, ref_text=open(REF_TXT, encoding="utf-8").read().strip(), ) log.info("Lucy-F5 bereit in %.1fs (Provider=%s, dev=%s, NFE=%d)", time.time() - t0, prov, dev, NFE_STEP) yield STATE.clear() app = FastAPI(title="Lucy TTS (F5/DirectML)", lifespan=lifespan) class Req(BaseModel): text: str @app.get("/health") def health(): return {"status": "ok" if "A" in STATE else "loading", "engine": "f5-tts", "provider": STATE.get("prov"), "nfe": NFE_STEP, "sr": SR} @app.post("/tts") def tts(req: Req): if "A" not in STATE: return JSONResponse({"error": "loading"}, status_code=503) t0 = time.time() parts = [] with LOCK: for sent in split_sentences(req.text): parts.append(cleanup(_infer(sent), SR)) a = np.concatenate(parts) if parts else np.zeros(0, np.float32) buf = io.BytesIO(); sf.write(buf, a, SR, format="WAV", subtype="PCM_16"); buf.seek(0) dur = a.size / SR; gen = time.time() - t0 log.info("/tts %dZ audio=%.1fs gen=%.1fs rtf=%.2f", len(req.text), dur, gen, gen / max(dur, 0.01)) return Response(buf.read(), media_type="audio/wav", headers={"X-Audio-Seconds": f"{dur:.2f}", "X-Gen-Seconds": f"{gen:.2f}"}) @app.post("/tts/stream") def tts_stream(req: Req): if "A" not in STATE: return JSONResponse({"error": "loading"}, status_code=503) sentences = split_sentences(req.text) def pcm(): t0 = time.time(); total = 0; first = True with LOCK: for sent in sentences: a = cleanup(_infer(sent), SR); total += a.size if first: log.info("/tts/stream TTFB=%.2fs (%d Sätze)", time.time() - t0, len(sentences)); first = False yield _to_pcm16(a) log.info("/tts/stream %dZ audio=%.1fs gen=%.1fs", len(req.text), total / SR, time.time() - t0) return StreamingResponse(pcm(), media_type="application/octet-stream", headers={"X-Sample-Rate": str(SR)})