HORST / evaluate.py
Bayernator's picture
HORST: eigenes Korrekturmodell auf ZeroGPU
6424eee verified
Raw History Blame Contribute Delete
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)")