Spaces:
Running on Zero
Running on Zero
File size: 7,438 Bytes
f1ef7e2 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 | """
Phase 5 (prototype): conservative LLM cleanup pass over Phase 1 transcripts.
Reframed by the error analysis: the remaining errors are dominated by
function-word/filler deletions an LLM cannot restore, so this pass is NOT a
WER-to-99 play. Its job is to fix the small, genuinely addressable slice --
misheard homophones, garbled proper nouns/names, obviously wrong words -- to
improve DOWNSTREAM transcript quality for compliance. The main risk is
OVER-correction (the LLM paraphrasing or "fixing" correct words), so:
- the prompt is strict / minimal-edit, temperature low
- text is corrected in small sentence-aware chunks (edits stay local)
- a guardrail rejects a chunk's correction if its length deviates too much
- WER is measured BEFORE vs AFTER as the guardrail metric: if WER regresses,
the pass is doing more harm than good.
Model-agnostic (--model). Default qwen2.5:7b-instruct (best quality, needs
~5 GB free RAM); fall back to llama3.2:3b on low-RAM machines.
Usage:
python llm_cleanup.py # all 12 probe calls
python llm_cleanup.py --model llama3.2:3b --limit 3
"""
import os
import re
import json
import time
import argparse
import urllib.request
import difflib
import jiwer
from eval_common import normalise
DATA = r"d:\Desktop\ai-ml-capstone\data\na_testset"
RES = "results_channels" # Phase 1 input
OLLAMA = "http://localhost:11434/api/generate"
PROMPT = """You are correcting errors in an automatic speech-recognition transcript of a {domain} phone call (the {speaker} is speaking). Fix ONLY clear recognition errors:
- misheard words and wrong homophones (e.g. "segway" -> "segue")
- misspelled or garbled proper nouns, names, and places
- obviously wrong words that don't fit the sentence
STRICT RULES:
- Do NOT paraphrase, rephrase, reorder, summarize, or improve style.
- Do NOT add or remove words. Keep ALL filler and disfluencies (um, uh, okay, yeah, repeated words).
- Keep all numbers exactly as written.
- Preserve punctuation and capitalization as-is unless clearly wrong.
- If the text is already correct, return it EXACTLY unchanged.
- Output ONLY the corrected transcript text. No preamble, no quotes, no commentary.
Transcript:
{text}"""
def ollama(prompt, model, temperature=0.1, timeout=180):
body = json.dumps({
"model": model, "prompt": prompt, "stream": False,
"options": {"temperature": temperature, "num_predict": 1024},
}).encode()
req = urllib.request.Request(OLLAMA, data=body, headers={"Content-Type": "application/json"})
with urllib.request.urlopen(req, timeout=timeout) as r:
return json.loads(r.read())["response"].strip()
def sentence_chunks(text, max_words=130):
"""Sentence-aware chunks of <= max_words, so corrections stay local and the
model keeps sentence context."""
sents = re.split(r"(?<=[.?!])\s+", text)
chunk, n, out = [], 0, []
for s in sents:
w = len(s.split())
if n + w > max_words and chunk:
out.append(" ".join(chunk)); chunk, n = [], 0
chunk.append(s); n += w
if chunk:
out.append(" ".join(chunk))
return out
def strip_wrapper(s):
s = s.strip()
s = re.sub(r'^(here is|here\'s|corrected( transcript)?:?)\s*', '', s, flags=re.IGNORECASE).strip()
if len(s) >= 2 and s[0] in '"“' and s[-1] in '"”':
s = s[1:-1].strip()
return s
def clean_text(text, domain, speaker, model, temperature, max_words):
out, n_changed = [], 0
for chunk in sentence_chunks(text, max_words):
try:
corrected = strip_wrapper(ollama(
PROMPT.format(domain=domain, speaker=speaker, text=chunk), model, temperature))
except Exception:
corrected = chunk
# guardrail: reject implausible rewrites (length blow-up/collapse or empty)
if (not corrected) or abs(len(corrected.split()) - len(chunk.split())) > 0.25 * len(chunk.split()) + 3:
corrected = chunk
if corrected != chunk:
n_changed += 1
out.append(corrected)
return " ".join(out), n_changed
def wer_norm(ref, hyp):
r, h = normalise(ref), normalise(hyp)
n = len(r.split())
return (jiwer.wer(r, h) if n else 0.0), n
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--model", default="qwen2.5:7b-instruct")
ap.add_argument("--out", default="results_channels_llm")
ap.add_argument("--chunk-words", type=int, default=130)
ap.add_argument("--temperature", type=float, default=0.1)
ap.add_argument("--limit", type=int, default=None)
args = ap.parse_args()
with open(os.path.join(DATA, "manifest.json"), encoding="utf-8") as f:
manifest = {m["call_id"]: m for m in json.load(f)}
with open(os.path.join(DATA, "probe_set.json"), encoding="utf-8") as f:
probe = json.load(f)["calls"]
if args.limit:
probe = probe[: args.limit]
out_root = os.path.join(DATA, args.out)
print(f"LLM cleanup | model={args.model} | {len(probe)} calls\n")
print(f" {'call_id':<32} {'before':>7} {'after':>7} {'delta':>6} {'chunks_chg':>11}")
print(" " + "-" * 70)
agg = {"before": [0.0, 0], "after": [0.0, 0]}
t_start = time.time()
for p in probe:
cid, accent = p["call_id"], p["accent"]
m = manifest[cid]
with open(os.path.join(DATA, RES, accent, cid + ".json"), encoding="utf-8") as f:
hyp = json.load(f)
before_a = " ".join(w["word"] for w in hyp["agent"])
before_c = " ".join(w["word"] for w in hyp["customer"])
after_a, chg_a = clean_text(before_a, m["domain"], "agent", args.model, args.temperature, args.chunk_words)
after_c, chg_c = clean_text(before_c, m["domain"], "customer", args.model, args.temperature, args.chunk_words)
# word-weighted WER before/after
wba, na = wer_norm(m["agent_transcript"], before_a)
wbc, nc = wer_norm(m["customer_transcript"], before_c)
waa, _ = wer_norm(m["agent_transcript"], after_a)
wac, _ = wer_norm(m["customer_transcript"], after_c)
before = (wba * na + wbc * nc) / (na + nc)
after = (waa * na + wac * nc) / (na + nc)
agg["before"][0] += before * (na + nc); agg["before"][1] += na + nc
agg["after"][0] += after * (na + nc); agg["after"][1] += na + nc
os.makedirs(os.path.join(out_root, accent), exist_ok=True)
with open(os.path.join(out_root, accent, cid + ".json"), "w", encoding="utf-8") as f:
json.dump({"call_id": cid, "accent": accent, "domain": m["domain"], "model": args.model,
"agent_text": after_a, "customer_text": after_c,
"agent_text_before": before_a, "customer_text_before": before_c}, f, indent=2)
d = (after - before) * 100
print(f" {cid:<32} {(1-before)*100:>6.1f}% {(1-after)*100:>6.1f}% {-d:>+5.1f} {chg_a+chg_c:>11}")
ba = (1 - agg["before"][0] / agg["before"][1]) * 100
aa = (1 - agg["after"][0] / agg["after"][1]) * 100
print(" " + "-" * 70)
print(f" {'OVERALL':<32} {ba:>6.1f}% {aa:>6.1f}% {aa-ba:>+5.1f}")
print(f"\n {'IMPROVED' if aa > ba + 0.1 else 'NO GAIN / REGRESSED'} "
f"| {len(probe)} calls in {(time.time()-t_start)/60:.1f} min")
print(" (positive delta = WER reduced = good; negative = over-correction, pass is harmful)")
if __name__ == "__main__":
main()
|