| |
| """Analyse des erreurs residuelles sur devhard (config optimale) pour trouver des |
| corrections SYSTEMATIQUES exploitables. Cherche : substitutions les plus frequentes, |
| insertions/suppressions recurrentes, erreurs de casse, mots hors-vocabulaire du train. |
| """ |
| import collections |
| import json |
| import pickle |
| import re |
| import unicodedata |
|
|
| import jiwer |
| from multiprocessing import Pool |
| from pyctcdecode import build_ctcdecoder |
| from transformers import AutoProcessor |
|
|
| MDL = "/root/models/joint_cont_best" |
| ARPA = "/scratch/lm/lin_5g.arpa" |
|
|
|
|
| def comb(refs, hyps): |
| pr = [(r, h) for r, h in zip(refs, hyps) if r.strip()] |
| a = [x for x, _ in pr] |
| b = [y for _, y in pr] |
| return 0.5 * jiwer.wer(a, b) + 0.5 * jiwer.cer(a, b) |
|
|
|
|
| def main(): |
| rows = [json.loads(l) for l in open("/root/devhard/devhard_linsna.jsonl", encoding="utf-8")] |
| sub = [r for r in rows if r["lang"] == "lin"] |
| refs = [r["text"] for r in sub] |
| log = pickle.load(open("/scratch/lm/logits_lin.pkl", "rb")) |
| tok = AutoProcessor.from_pretrained(MDL).tokenizer |
|
|
| v = tok.get_vocab() |
| lab = [None] * len(v) |
| for t, i in v.items(): |
| lab[i] = t |
| lab[tok.word_delimiter_token_id] = " " |
| lab[tok.unk_token_id] = "⁇" |
| lab[tok.pad_token_id] = "" |
| greedy = [" ".join(tok.decode(l.argmax(-1)).replace("|", " ").split()) for l in log] |
| dec = build_ctcdecoder(lab, kenlm_model_path=ARPA, alpha=0.5, beta=0.5, |
| lm_score_boundary=False) |
| with Pool(8) as p: |
| hyps = [" ".join(x.split()) for x in dec.decode_batch(p, log, beam_width=64)] |
| hyps = [(g[:1] + h[1:] if h and g else h) for h, g in zip(hyps, greedy)] |
| print("combine actuel: %.4f" % comb(refs, hyps), flush=True) |
|
|
| |
| subs = collections.Counter() |
| dels = collections.Counter() |
| ins = collections.Counter() |
| out = jiwer.process_words(refs, hyps) |
| for ref_w, hyp_w, chunks in zip(out.references, out.hypotheses, out.alignments): |
| for ch in chunks: |
| if ch.type == "substitute": |
| for a, b in zip(ref_w[ch.ref_start_idx:ch.ref_end_idx], |
| hyp_w[ch.hyp_start_idx:ch.hyp_end_idx]): |
| subs[(a, b)] += 1 |
| elif ch.type == "delete": |
| for a in ref_w[ch.ref_start_idx:ch.ref_end_idx]: |
| dels[a] += 1 |
| elif ch.type == "insert": |
| for b in hyp_w[ch.hyp_start_idx:ch.hyp_end_idx]: |
| ins[b] += 1 |
| print("\n=== 25 substitutions les plus frequentes (ref -> hyp) ===", flush=True) |
| for (a, b), n in subs.most_common(25): |
| flag = "" |
| if a.lower() == b.lower(): |
| flag = " [CASSE SEULE]" |
| elif re.sub(r"[^\w]", "", a.lower()) == re.sub(r"[^\w]", "", b.lower()): |
| flag = " [PONCTUATION SEULE]" |
| print(" %5d %-22s -> %-22s%s" % (n, a, b, flag), flush=True) |
| print("\n=== 15 mots le plus souvent OMIS (dans ref, absents hyp) ===", flush=True) |
| for a, n in dels.most_common(15): |
| print(" %5d %s" % (n, a), flush=True) |
| print("\n=== 15 mots le plus souvent AJOUTES a tort ===", flush=True) |
| for b, n in ins.most_common(15): |
| print(" %5d %s" % (n, b), flush=True) |
|
|
| |
| tot = sum(subs.values()) |
| only_case = sum(n for (a, b), n in subs.items() if a.lower() == b.lower()) |
| only_punct = sum(n for (a, b), n in subs.items() |
| if a.lower() != b.lower() |
| and re.sub(r"[^\w]", "", a.lower()) == re.sub(r"[^\w]", "", b.lower())) |
| print("\n=== nature des substitutions ===", flush=True) |
| print(" total substitutions : %d" % tot, flush=True) |
| print(" casse seule : %d (%.1f%%)" % (only_case, 100.0 * only_case / max(tot, 1)), flush=True) |
| print(" ponctuation seule : %d (%.1f%%)" % (only_punct, 100.0 * only_punct / max(tot, 1)), flush=True) |
| print(" vrai contenu : %d (%.1f%%)" % (tot - only_case - only_punct, |
| 100.0 * (tot - only_case - only_punct) / max(tot, 1)), flush=True) |
|
|
| |
| def norm(t): |
| return unicodedata.normalize("NFC", t) |
| vocab = set() |
| for line in open("/scratch/lm/corpus_lin.txt", encoding="utf-8"): |
| vocab.update(re.findall(r"[\w']+", norm(line).lower())) |
| hw = [w for h in hyps for w in re.findall(r"[\w']+", norm(h).lower())] |
| oov = [w for w in hw if w not in vocab] |
| print("\n=== mots des hypotheses hors vocabulaire du LM ===", flush=True) |
| print(" %d / %d tokens (%.1f%%) | %d formes distinctes" |
| % (len(oov), len(hw), 100.0 * len(oov) / max(len(hw), 1), len(set(oov))), flush=True) |
| print(" top:", collections.Counter(oov).most_common(15), flush=True) |
| print("\nERRANALYSIS_DONE", flush=True) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|