initial_prompt personale + glossario: POST /transcribe accetta prompt e glossary (normalizzazione alias->canonical post-ASR, raw_text/segments[].raw), GET/PUT /glossary con GLOSSARY_FILE, glossary.example.json, --prompt/--carry-initial-prompt in transcribe.py

This commit is contained in:
Matteo Benedetto
2026-08-21 12:55:53 +02:00
parent 87bcbd49b0
commit 0b583ab39d
5 changed files with 200 additions and 25 deletions
+1
View File
@@ -61,6 +61,7 @@ ENV WHISPER_CLI=/app/whisper-cli
ENV MODELS_DIR=/models ENV MODELS_DIR=/models
ENV STT_PYTHON=python3 ENV STT_PYTHON=python3
ENV HF_HOME=/hf-cache ENV HF_HOME=/hf-cache
ENV GLOSSARY_FILE=/glossary/glossary.json
EXPOSE 8883 EXPOSE 8883
ENTRYPOINT ["./entrypoint.sh"] ENTRYPOINT ["./entrypoint.sh"]
+2
View File
@@ -33,9 +33,11 @@ services:
HF_HOME: /hf-cache HF_HOME: /hf-cache
HF_TOKEN: ${HF_TOKEN:-} HF_TOKEN: ${HF_TOKEN:-}
PORT: "8883" PORT: "8883"
GLOSSARY_FILE: /glossary/glossary.json
volumes: volumes:
- "${MODELS_DIR:-/home/enne2/dev/whisper.cpp/models}:/models:ro" - "${MODELS_DIR:-/home/enne2/dev/whisper.cpp/models}:/models:ro"
- "${HF_CACHE:-/home/enne2/.cache/huggingface}:/hf-cache:ro" - "${HF_CACHE:-/home/enne2/.cache/huggingface}:/hf-cache:ro"
- "${GLOSSARY_FILE:-/home/enne2/dev/stt/glossary.json}:/glossary/glossary.json:ro"
ports: ports:
- "${PORT:-8883}:8883" - "${PORT:-8883}:8883"
restart: unless-stopped restart: unless-stopped
+39
View File
@@ -0,0 +1,39 @@
{
"people": [
{
"canonical": "Giulia Bianchi",
"aliases": ["Giulia Bianci", "Julia Bianchi", "Giulia Bianche"],
"domain": "persone"
},
{
"canonical": "Pietrolli",
"aliases": ["Pietro Lli", "Pietrolli"],
"domain": "persone"
}
],
"organizations": [
{
"canonical": "Perplexity AI",
"aliases": ["Perplexiti", "Perplexity"],
"domain": "aziende"
}
],
"terms": [
{
"canonical": "retrieval-augmented generation",
"aliases": ["rag", "retriaval augmented generation"],
"domain": "AI"
},
{
"canonical": "fine-tuning",
"aliases": ["fine tuning", "fain tuning"],
"domain": "AI"
},
{
"canonical": "stt-server",
"aliases": ["stt server", "esteti server", "siti server"],
"domain": "infrastruttura"
}
],
"custom": []
}
+51 -13
View File
@@ -4,13 +4,18 @@ ASR: whisper.cpp (Vulkan) | Diarization opzionale: pyannote Community-1 (CPU)
Endpoints: Endpoints:
GET /health → {"status": "ok", "model": ...} GET /health → {"status": "ok", "model": ...}
POST /transcribe → multipart: file=<audio> + form: lang, model, diarize GET /glossary → glossario personale (se GLOSSARY_FILE configurato)
diarize=false (default): ASR only → {"text": "...", "segments": [...]} PUT /glossary → salva glossario (se GLOSSARY_FILE configurato)
POST /transcribe → multipart: file=<audio> + form:
lang, model, diarize, prompt (initial_prompt personale),
glossary (JSON string con persone/termini → normalizzazione
deterministica alias→canonical; il testo originale resta in 'raw')
diarize=false (default): ASR only → {"text": ..., "raw_text": ...}
diarize=true: ASR + diarization → segments con speaker diarize=true: ASR + diarization → segments con speaker
Uso: uvicorn server:app --host 0.0.0.0 --port 8883 Uso: uvicorn server:app --host 0.0.0.0 --port 8883
""" """
import asyncio, os, tempfile, time import asyncio, json, os, tempfile, time
from pathlib import Path from pathlib import Path
from fastapi import FastAPI, File, Form, UploadFile from fastapi import FastAPI, File, Form, UploadFile
@@ -18,7 +23,7 @@ from fastapi.responses import JSONResponse
import transcribe import transcribe
app = FastAPI(title="STT Server", version="1.0") app = FastAPI(title="STT Server", version="1.1")
# Semaforo: whisper-cli Vulkan e pyannote non sono thread-safe per uso concorrente # Semaforo: whisper-cli Vulkan e pyannote non sono thread-safe per uso concorrente
_lock = asyncio.Lock() _lock = asyncio.Lock()
@@ -26,18 +31,21 @@ _lock = asyncio.Lock()
# Cache dei modelli già caricati (pyannote pipeline è costosa da caricare) # Cache dei modelli già caricati (pyannote pipeline è costosa da caricare)
_pipeline_cache = {} _pipeline_cache = {}
# Glossario server-side condiviso (opzionale): env GLOSSARY_FILE
_glossary_file = os.environ.get("GLOSSARY_FILE") or ""
def _asr_only(audio_path: str, lang: str, model: str) -> list[dict]:
def _asr_only(audio_path: str, lang: str, model: str, prompt: str | None) -> list[dict]:
"""Path veloce: solo ASR (whisper.cpp Vulkan), niente diarization.""" """Path veloce: solo ASR (whisper.cpp Vulkan), niente diarization."""
with tempfile.TemporaryDirectory() as tmp: with tempfile.TemporaryDirectory() as tmp:
segs = transcribe.asr(audio_path, model, lang, tmp) segs = transcribe.asr(audio_path, model, lang, tmp, prompt=prompt)
return segs return segs
def _full_pipeline(audio_path: str, lang: str, model: str) -> list[dict]: def _full_pipeline(audio_path: str, lang: str, model: str, prompt: str | None) -> list[dict]:
"""Path completo: ASR + diarization pyannote.""" """Path completo: ASR + diarization pyannote."""
with tempfile.TemporaryDirectory() as tmp: with tempfile.TemporaryDirectory() as tmp:
segs = transcribe.asr(audio_path, model, lang, tmp) segs = transcribe.asr(audio_path, model, lang, tmp, prompt=prompt)
turns = transcribe.diarize(audio_path, tmp) turns = transcribe.diarize(audio_path, tmp)
segs = transcribe.merge(segs, turns) segs = transcribe.merge(segs, turns)
return segs return segs
@@ -45,7 +53,27 @@ def _full_pipeline(audio_path: str, lang: str, model: str) -> list[dict]:
@app.get("/health") @app.get("/health")
async def health(): async def health():
return {"status": "ok", "service": "stt", "model": "large-v3-turbo"} return {"status": "ok", "service": "stt", "model": "large-v3-turbo",
"version": "1.1"}
@app.get("/glossary")
async def get_glossary():
if not _glossary_file or not os.path.exists(_glossary_file):
return JSONResponse({"error": "glossario non configurato"}, status_code=404)
with open(_glossary_file, encoding="utf-8") as f:
return JSONResponse(json.load(f))
@app.put("/glossary")
async def put_glossary(body: dict):
if not _glossary_file:
return JSONResponse({"error": "glossario non configurato (GLOSSARY_FILE)"},
status_code=404)
with open(_glossary_file, "w", encoding="utf-8") as f:
json.dump(body, f, ensure_ascii=False, indent=2)
return {"status": "ok", "entries": sum(len(body.get(g) or []) for g in
("people", "organizations", "terms", "custom"))}
@app.post("/transcribe") @app.post("/transcribe")
@@ -54,6 +82,8 @@ async def transcribe_endpoint(
lang: str = Form("it"), lang: str = Form("it"),
model: str = Form("large-v3-turbo"), model: str = Form("large-v3-turbo"),
diarize: bool = Form(False), diarize: bool = Form(False),
prompt: str | None = Form(None),
glossary: str | None = Form(None),
): ):
t0 = time.time() t0 = time.time()
# Salva l'upload in un file temporaneo # Salva l'upload in un file temporaneo
@@ -61,26 +91,34 @@ async def transcribe_endpoint(
with tempfile.NamedTemporaryFile(suffix=suffix, delete=False) as tmp: with tempfile.NamedTemporaryFile(suffix=suffix, delete=False) as tmp:
tmp.write(await file.read()) tmp.write(await file.read())
audio_path = tmp.name audio_path = tmp.name
flat_map = transcribe.load_glossary(glossary) if glossary else {}
try: try:
async with _lock: async with _lock:
loop = asyncio.get_running_loop() loop = asyncio.get_running_loop()
if diarize: if diarize:
segs = await loop.run_in_executor(None, _full_pipeline, audio_path, lang, model) segs = await loop.run_in_executor(None, _full_pipeline, audio_path, lang, model, prompt)
else: else:
segs = await loop.run_in_executor(None, _asr_only, audio_path, lang, model) segs = await loop.run_in_executor(None, _asr_only, audio_path, lang, model, prompt)
finally: finally:
os.unlink(audio_path) os.unlink(audio_path)
raw_text = " ".join(s["text"] for s in segs).strip()
if flat_map:
segs = transcribe.apply_glossary(segs, flat_map)
text = " ".join(s["text"] for s in segs).strip() text = " ".join(s["text"] for s in segs).strip()
elapsed = round(time.time() - t0, 2) elapsed = round(time.time() - t0, 2)
return JSONResponse({ resp = {
"text": text, "text": text,
"segments": segs, "segments": segs,
"diarize": diarize, "diarize": diarize,
"lang": lang, "lang": lang,
"model": model, "model": model,
"elapsed_s": elapsed, "elapsed_s": elapsed,
}) }
# raw_text solo quando c'è stato un intervento (prompt/glossario)
if flat_map or prompt:
resp["raw_text"] = raw_text
return JSONResponse(resp)
if __name__ == "__main__": if __name__ == "__main__":
+104 -9
View File
@@ -2,26 +2,101 @@
"""STT + Speaker Diarization pipeline (frigate.vpn). """STT + Speaker Diarization pipeline (frigate.vpn).
ASR: whisper.cpp (Vulkan) | Diarization: pyannote Community-1 (CPU) ASR: whisper.cpp (Vulkan) | Diarization: pyannote Community-1 (CPU)
Uso: transcribe.py <audio> [--lang it] [--model large-v3-turbo] [--out base] Uso: transcribe.py <audio> [--lang it] [--model large-v3-turbo] [--out base]
[--prompt \"frase personale\"] [--glossary glossario.json]
""" """
import argparse, json, os, subprocess, sys, tempfile import argparse, json, os, re, subprocess, sys, tempfile
WHISPER_CLI = os.environ.get('WHISPER_CLI', os.path.expanduser('~/dev/whisper.cpp/build/bin/whisper-cli')) WHISPER_CLI = os.environ.get('WHISPER_CLI', os.path.expanduser('~/dev/whisper.cpp/build/bin/whisper-cli'))
MODELS_DIR = os.environ.get('MODELS_DIR', os.path.expanduser('~/dev/whisper.cpp/models')) MODELS_DIR = os.environ.get('MODELS_DIR', os.path.expanduser('~/dev/whisper.cpp/models'))
PYTHON = os.environ.get('STT_PYTHON', os.path.expanduser('~/dev/stt-venv/bin/python')) PYTHON = os.environ.get('STT_PYTHON', os.path.expanduser('~/dev/stt-venv/bin/python'))
def asr(audio, model, lang, tmp): # ---------------------------------------------------------------------------
# Glossario personale: estensione vocabolario + normalizzazione deterministica
# ---------------------------------------------------------------------------
# Formato (JSON):
# {
# "people": [{"canonical": "Giulia Bianchi", "aliases": ["Giulia Bianci", ...]}],
# "organizations": [{"canonical": "...", "aliases": [...]}],
# "terms": [{"canonical": "fine-tuning", "aliases": ["fine tuning", "fain tuning"]}]
# }
# La normalizzazione sostituisce gli alias con la forma canonica SOLO su match
# esatto (case-insensitive, confini di parola). Le sostituzioni sono annotate
# nel campo "raw" del segmento per audit/rollback.
def load_glossary(path_or_json):
"""Accetta un path JSON o una stringa JSON; restituisce la mappa
alias->canonical (alias in lowercase). Invalido/vuoto -> {}."""
if not path_or_json:
return {}
try:
data = path_or_json
if isinstance(data, str):
stripped = data.strip()
if stripped.startswith('{'):
data = json.loads(stripped) # stringa JSON
else:
with open(stripped, encoding='utf-8') as f: # path
data = json.load(f)
flat = {}
for group in ('people', 'organizations', 'terms', 'custom'):
for entry in (data.get(group) or []):
canon = str(entry.get('canonical') or '').strip()
if not canon:
continue
for alias in (entry.get('aliases') or []):
alias = str(alias).strip()
if alias and alias.lower() != canon.lower():
flat[alias.lower()] = canon
# Formato alternativo semplice: {"alias": "canonical"}
if not flat:
for alias, canon in data.items():
if isinstance(canon, str) and isinstance(alias, str):
flat[alias.lower()] = canon
return flat
except Exception:
return {}
def normalize_text(text, flat_map):
"""Sostituisce alias (case-insensitive, confini parola) con la forma
canonica. Le sostituzioni avvengono su segnaposto univoci: un alias non
può ri-matchare testo generato da un'altra sostituzione (no cascate).
Restituisce (nuovo_testo, sostituzioni_fatte)."""
if not flat_map or not text:
return text, 0
out, count = text, 0
placeholders = []
# Passo 1: alias -> segnaposto (alias più lunghi prima, evita match parziali)
for alias in sorted(flat_map, key=len, reverse=True):
canon = flat_map[alias]
pattern = r'(?<!\w)' + re.escape(alias) + r'(?!\w)'
ph = f'\x00PH{len(placeholders)}\x00'
new_out, n = re.subn(pattern, ph, out, flags=re.IGNORECASE)
if n:
placeholders.append((ph, canon))
count += n
out = new_out
# Passo 2: segnaposto -> forma canonica
for ph, canon in placeholders:
out = out.replace(ph, canon)
return out, count
def asr(audio, model, lang, tmp, prompt=None):
model_path = os.path.join(MODELS_DIR, f'ggml-{model}.bin') model_path = os.path.join(MODELS_DIR, f'ggml-{model}.bin')
out_json = os.path.join(tmp, 'asr.json') out_json = os.path.join(tmp, 'asr.json')
vad_model = os.path.join(MODELS_DIR, 'ggml-silero-v6.2.0.bin') vad_model = os.path.join(MODELS_DIR, 'ggml-silero-v6.2.0.bin')
# VAD (silero): salta il silenzio/rumore → elimina le allucinazioni cmd = [WHISPER_CLI, '-m', model_path, '-f', audio, '-l', lang,
# tipo "Grazie a tutti" su audio senza parlato. --suppress-regex come
# rete di sicurezza per frasi allucinate residue.
subprocess.run([WHISPER_CLI, '-m', model_path, '-f', audio, '-l', lang,
'-oj', '-of', os.path.join(tmp, 'asr'), '--no-prints', '-oj', '-of', os.path.join(tmp, 'asr'), '--no-prints',
'-vm', vad_model, '--vad', '-vm', vad_model, '--vad',
'--suppress-regex', '--suppress-regex',
r'(Grazie a tutti|Grazie per l.attenzione|Thank you|Thanks for watching|Sottotitoli creati|Sottotitoli)'], r'(Grazie a tutti|Grazie per l.attenzione|Thank you|Thanks for watching|Sottotitoli creati|Sottotitoli)']
check=True, capture_output=True) # Initial prompt (bias vocabolario personale, nomi, gergo): boost
# zero-training. --carry-initial-prompt lo riapplica su ogni finestra
# (necessario per audio lunghi; verificare con whisper-cli -h).
if prompt:
cmd += ['--prompt', prompt, '--carry-initial-prompt']
subprocess.run(cmd, check=True, capture_output=True)
data = json.load(open(out_json)) data = json.load(open(out_json))
segs = [] segs = []
for s in data['transcription']: for s in data['transcription']:
@@ -37,6 +112,18 @@ def asr(audio, model, lang, tmp):
'text': text}) 'text': text})
return segs return segs
def apply_glossary(segs, flat_map):
"""Normalizza i testi dei segmenti con il glossario. I segmenti modificati
ricevono 'raw' = testo ASR originale (audit/rollback)."""
for seg in segs:
new_text, count = normalize_text(seg['text'], flat_map)
if count:
seg['raw'] = seg['text']
seg['text'] = new_text
return segs
def diarize(audio, tmp): def diarize(audio, tmp):
script = f''' script = f'''
import json, sys import json, sys
@@ -85,12 +172,20 @@ def main():
ap.add_argument('--lang', default='it') ap.add_argument('--lang', default='it')
ap.add_argument('--model', default='large-v3-turbo') ap.add_argument('--model', default='large-v3-turbo')
ap.add_argument('--out', default=None) ap.add_argument('--out', default=None)
ap.add_argument('--prompt', default=None,
help='initial prompt personale (nomi, termini, gergo)')
ap.add_argument('--glossary', default=None,
help='path a glossario JSON (alias -> forma canonica)')
args = ap.parse_args() args = ap.parse_args()
base = args.out or os.path.splitext(args.audio)[0] base = args.out or os.path.splitext(args.audio)[0]
flat_map = load_glossary(args.glossary)
with tempfile.TemporaryDirectory() as tmp: with tempfile.TemporaryDirectory() as tmp:
print(f'[1/3] ASR whisper.cpp ({args.model}, {args.lang})...', file=sys.stderr) print(f'[1/3] ASR whisper.cpp ({args.model}, {args.lang})...', file=sys.stderr)
segs = asr(args.audio, args.model, args.lang, tmp) segs = asr(args.audio, args.model, args.lang, tmp, prompt=args.prompt)
print(f' {len(segs)} segmenti', file=sys.stderr) print(f' {len(segs)} segmenti', file=sys.stderr)
if flat_map:
segs = apply_glossary(segs, flat_map)
print(f' glossario: {len(flat_map)} alias applicati', file=sys.stderr)
print('[2/3] Diarization pyannote Community-1...', file=sys.stderr) print('[2/3] Diarization pyannote Community-1...', file=sys.stderr)
turns = diarize(args.audio, tmp) turns = diarize(args.audio, tmp)
print(f' {len(turns)} turni parlante', file=sys.stderr) print(f' {len(turns)} turni parlante', file=sys.stderr)