Spaces:
Running on Zero
Running on Zero
Download evaluate.py from Bayernator/HORST: direct link, hf CLI and curl.
- Browser
- Download file 7.94 kB
-
https://huggingface.co/spaces/Bayernator/HORST/resolve/main/evaluate.py
- Command line
-
hf download hf://spaces/Bayernator/HORST/evaluate.py
-
curl -L -o evaluate.py https://huggingface.co/spaces/Bayernator/HORST/resolve/main/evaluate.py
7.94 kB
| """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)") | |