Lucy-TTS/F5: Skripte + Batches versionieren, schwere Assets ignoriert
- pocket_server.py (Produktions-TTS mit Stimmen-Waechter), text_norm, Bench-/Diag-Skripte - lucy-f5: f5_server/f5_test/bench_dml (DirectML-Experiment, Phase C/D offen) - .gitignore: venvs/Modelle/Audio/Logs der beiden Ordner + box_recon/gemma_swap-Scratch Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
@@ -0,0 +1,196 @@
|
||||
# -*- 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("<i2").tobytes()
|
||||
|
||||
# ---- ONNX-Inferenz (A=Preprocess CPU, B=Transformer DML+io_binding, C=Decode CPU) ----
|
||||
def _infer(gen_text: str) -> 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)})
|
||||
Reference in New Issue
Block a user