"""F0.5 auf Falko-MERLIN (dev/test) und dem eigenen Hard-Benchmark (bench-dev: 1000, bench: 4000) + Korrektur-CLI. python evaluate.py --data data --split dev # Scorer-Check: Oracle und Identität python evaluate.py --data data --split bench --ckpt ckpts/ckpt-*.pt # mehrere Checkpoints -> gemittelt python evaluate.py --data data --ckpt best.pt --correct "ich weis nicht das es so ist" """ import argparse import difflib import json import sentencepiece as spm import torch from model import EOS, GEC from prepare import _pair_split, detokenize, split_sentences, tokenize def load_model(paths, device): states = [torch.load(p, map_location="cpu", weights_only=False) for p in paths] avg = {k: sum(s["model"][k].float() for s in states) / len(states) for k in states[0]["model"]} model = GEC(**states[0]["cfg"]) model.load_state_dict(avg) return model.to(device).eval() def correct_tokens(model, sp, sents, beam=5): """sents: tokenisierte Sätze -> korrigierte tokenisierte Sätze.""" dev = next(model.parameters()).device out = [] with torch.autocast(dev.type, dtype=torch.float16, enabled=dev.type == "cuda"): for s in sents: ids = sp.encode(s)[: model.cfg["max_len"] - 1] + [EOS] out.append(sp.decode(model.beam_search(torch.tensor(ids, device=dev), beam=beam))) return out def parse_m2(path): """-> [(Quelltokens, {(start, ende, korrektur)})], nur Annotator 0.""" data = [] for block in open(path, encoding="utf-8").read().strip().split("\n\n"): lines = block.split("\n") edits = set() for a in lines[1:]: span, typ, cor, *_, annot = a[2:].split("|||") if typ != "noop" and annot == "0": s, e = map(int, span.split()) edits.add((s, e, cor)) data.append((lines[0][2:], edits)) return data def _token_edits(a, b, off): """Levenshtein über Tokens -> Einzel-Token-Edits wie in den Gold-M2 (ähnliche Wörter werden bevorzugt ersetzt).""" D = [[i + j if not i * j else 0 for j in range(len(b) + 1)] for i in range(len(a) + 1)] sub = lambda x, y: 2 - difflib.SequenceMatcher(None, x.lower(), y.lower()).ratio() for i in range(1, len(a) + 1): for j in range(1, len(b) + 1): D[i][j] = min(D[i - 1][j] + 1, D[i][j - 1] + 1, D[i - 1][j - 1] + sub(a[i - 1], b[j - 1])) edits, i, j = set(), len(a), len(b) while i or j: if i and j and D[i][j] == D[i - 1][j - 1] + sub(a[i - 1], b[j - 1]): edits.add((off + i - 1, off + i, b[j - 1])); i -= 1; j -= 1 elif j and D[i][j] == D[i][j - 1] + 1: edits.add((off + i, off + i, b[j - 1])); j -= 1 else: edits.add((off + i - 1, off + i, "")); i -= 1 return edits def hyp_edits(src, hyp, gold=frozenset()): """Wie der M2-Scorer (MaxMatch): pro Änderungsblock die Zerlegung wählen, die am besten zu Gold passt.""" # ponytail: nur zwei Zerlegungen (ganzer Block / pro Token) statt voller MaxMatch-Lattice -> Werte annähernd literaturvergleichbar a, b = src.split(), hyp.split() edits = set() for tag, i1, i2, j1, j2 in difflib.SequenceMatcher(None, a, b, autojunk=False).get_opcodes(): if tag != "equal": merged = {(i1, i2, " ".join(b[j1:j2]))} split = _token_edits(a[i1:i2], b[j1:j2], i1) edits |= max((merged, split), key=lambda e: (len(e & gold), -len(e))) return edits def ref_edits(src, trg): """Gold-Edits aus einem Paar ohne M2-Annotation (Hard-Benchmark), pro Token geschnitten.""" a, b = src.split(), trg.split() ops = difflib.SequenceMatcher(None, a, b, autojunk=False).get_opcodes() return set().union(*(_token_edits(a[i1:i2], b[j1:j2], i1) for tag, i1, i2, j1, j2 in ops if tag != "equal")) def load_split(data, split): """-> [(Quelle tokenisiert, Gold-Edits, zu korrigierende Sätze)], [Zieltexte fürs Oracle].""" if split in ("dev", "test"): m2 = parse_m2(f"{data}/real/fm-{split}.m2") trg = open(f"{data}/real/fm-{split}.trg", encoding="utf-8").read().split("\n") return [(s, g, [s]) for s, g in m2], trg[: len(m2)] rows = [json.loads(l) for l in open(f"{data}/real/bench.jsonl", encoding="utf-8")] rows = rows[:1000] if split == "bench-dev" else rows[1000:] items = [] for r in rows: src = tokenize(r["input"]) items.append((src, ref_edits(src, tokenize(r["target"])), [tokenize(x) for x in _pair_split(r["input"])])) return items, [tokenize(r["target"]) for r in rows] def group_sentences(sents, sp, limit): """Sätze zu Gruppen bis `limit` Tokens bündeln (Kontext!). Ein Satz, der allein nicht passt, wird an Wortgrenzen geteilt: correct_tokens würde ihn sonst nach `limit` Tokens abschneiden und den Rest still verlieren.""" n = lambda words: len(sp.encode(" ".join(words))) pieces = [] for s in sents: if n([s]) <= limit: pieces.append(s) continue cur = [] for w in s.split(): if cur and n(cur + [w]) > limit: pieces.append(" ".join(cur)) cur = [] cur.append(w) pieces.append(" ".join(cur)) groups, cur = [], [] for p in pieces: if cur and n(cur + [p]) > limit: groups.append(" ".join(cur)) cur = [] cur.append(p) return groups + [" ".join(cur)] def correct_text(model, sp, sents, beam=5): """Tokenisierte Sätze -> in Gruppen korrigiert (so viel Kontext, wie ins Modell passt) und zusammengesetzt.""" return " ".join(correct_tokens(model, sp, group_sentences(sents, sp, model.cfg["max_len"] - 1), beam)) def correct_items(model, sp, items, beam=5): return [correct_text(model, sp, sents, beam) for _, _, sents in items] def f05(items, hyps): tp = fp = fn = 0 for (src, gold, *_), hyp in zip(items, hyps): h = hyp_edits(src, hyp, gold) tp, fp, fn = tp + len(h & gold), fp + len(h - gold), fn + len(gold - h) p = tp / (tp + fp) if tp + fp else 1.0 r = tp / (tp + fn) if tp + fn else 1.0 f = 1.25 * p * r / (0.25 * p + r) if p + r else 0.0 return p, r, f if __name__ == "__main__": ap = argparse.ArgumentParser() ap.add_argument("--data", default="data") ap.add_argument("--split", default="dev", choices=["dev", "test", "bench-dev", "bench"]) ap.add_argument("--ckpt", nargs="*") ap.add_argument("--beam", type=int, default=5) ap.add_argument("--n", type=int, help="nur die ersten n Beispiele") ap.add_argument("--correct", help="beliebigen Text korrigieren") a = ap.parse_args() items, trg = load_split(a.data, a.split) items, trg = items[: a.n], trg[: a.n] if not a.ckpt: # Scorer-Selbsttest: Zieltexte müssen nahe 1.0 liegen, Identität bei 0 p, r, f = f05(items, trg) print(f"Oracle P={p:.3f} R={r:.3f} F0.5={f:.3f}") assert f > 0.85, "Scorer weicht zu stark von den Gold-Edits ab" assert all(" ".join(sents) == src for src, _, sents in items), "Satzzerlegung verändert Tokens" p, r, f = f05(items, [s for s, *_ in items]) print(f"Identität P={p:.3f} R={r:.3f} F0.5={f:.3f}") assert r == 0 raise SystemExit device = "cuda" if torch.cuda.is_available() else "cpu" model = load_model(a.ckpt, device) sp = spm.SentencePieceProcessor(model_file=f"{a.data}/spm.model") if a.correct: sents = [tokenize(s) for s in split_sentences(a.correct) if s] print(detokenize(correct_text(model, sp, sents, a.beam))) raise SystemExit hyps = correct_items(model, sp, items, a.beam) with open(f"hyp-{a.split}.txt", "w", encoding="utf-8") as fh: fh.write("\n".join(hyps) + "\n") p, r, f = f05(items, hyps) print(f"{a.split}: P={p:.3f} R={r:.3f} F0.5={f:.3f} (Hypothesen in hyp-{a.split}.txt)")