hfvladkon's picture
Upload main defect tagger (weights, tokenizer, inference code, card)
22ca93a
Raw History Blame Contribute Delete
15.2 kB
"""Windowed inference over long answers and defect-level metrics.
Scores: p_def(token) = P(BAD) + P(TAIL). A token is predicted defective when p_def >= threshold.
With the grouped label scheme (O, BAD/TAIL-script, BAD/TAIL-grammar) the threshold is a pair
(t_script, t_grammar) applied to P(script) and P(grammar) separately, since script leaks get
p close to 1 while grammar errors score much lower.
Predicted spans are maximal runs of predicted tokens.
Gold events are the BAD/TAIL tokens of one annotated error (event_ids). Positions with label -100
inside the answer (DEPENDENT / EXCLUDED / UNCERTAIN) are neutral: a prediction there is neither
a hit nor a false alarm.
event_recall share of gold events with at least one predicted token on their BAD/TAIL tokens
head_recall share of BAD heads whose own token is predicted
span_precision share of predicted spans (not fully neutral) overlapping a gold event token
event_f1 harmonic mean of span_precision and event_recall
row_* answer-level: "has at least one defect" vs "at least one predicted span"
corrected_fp_per_1k predicted spans per 1000 tokens on repaired answers (should be ~0)
ap_token average precision of p_def on labelled tokens (O vs BAD/TAIL), threshold-free
*_tol2 the same with ±2 token tolerance: a prediction next to the annotated head
(the annotation anchors e.g. an agreement error on one word of the phrase)
counts as a hit and not as a false alarm
"""
from __future__ import annotations
from collections import defaultdict
import numpy as np
import torch
from modeling import model_inputs
def windows(n_answer: int, a0: int, max_len: int, stride: int):
"""Split answer positions [a0, a0+n) into windows that fit max_len together with the prompt."""
span = max_len - a0
if span <= stride:
raise ValueError("prompt too long for max_len")
starts, s = [], 0
while True:
starts.append(s)
if s + span >= n_answer:
break
s += span - stride
return [(s, min(n_answer, s + span)) for s in starts]
@torch.no_grad()
def predict_row(model, input_ids, a0: int, max_len: int, stride: int, device, batch: int = 8):
"""Return p(class) for every answer token, merging overlapping windows by centrality."""
ids = torch.as_tensor(input_ids, dtype=torch.long)
prompt, ans = ids[:a0], ids[a0:]
n = len(ans)
if n == 0:
return np.zeros((0, model.config.num_labels), dtype=np.float32)
ws = windows(n, a0, max_len, stride)
probs = np.zeros((n, model.config.num_labels), dtype=np.float32)
best = np.full(n, -1.0)
for b in range(0, len(ws), batch):
chunk = ws[b:b + batch]
seqs = [torch.cat([prompt, ans[s:e]]) for s, e in chunk]
L = max(len(x) for x in seqs)
pad = getattr(model.config, "qwen_pad_token_id", None) or model.config.pad_token_id or 0
x = torch.full((len(seqs), L), pad, dtype=torch.long)
m = torch.zeros((len(seqs), L), dtype=torch.long)
for i, sq in enumerate(seqs):
x[i, :len(sq)] = sq
m[i, :len(sq)] = 1
logits = model(**model_inputs(model, x.to(device), m.to(device))).logits.float()
p = torch.softmax(logits, -1).cpu().numpy()
for i, (s, e) in enumerate(chunk):
pos = np.arange(s, e)
centr = np.minimum(pos - s, e - 1 - pos).astype(float)
upd = centr > best[pos]
probs[pos[upd]] = p[i, a0 + (pos[upd] - s)]
best[pos[upd]] = centr[upd]
return probs
def average_precision(scores: np.ndarray, y: np.ndarray) -> float:
if y.sum() == 0:
return float("nan")
order = np.argsort(-scores, kind="stable")
y = y[order]
tp = np.cumsum(y)
prec = tp / np.arange(1, len(y) + 1)
return float((prec * y).sum() / y.sum())
def spans_of(mask: np.ndarray):
out, i, n = [], 0, len(mask)
while i < n:
if mask[i]:
j = i
while j + 1 < n and mask[j + 1]:
j += 1
out.append((i, j + 1))
i = j + 1
else:
i += 1
return out
SCRIPT_TYPES = {"CJK", "LATIN_INSERT", "MIXED_SCRIPT", "TRANSLITERATION"}
def p_def(p: np.ndarray) -> np.ndarray:
return p[:, 1:].sum(1)
def defect_mask(p: np.ndarray, threshold) -> np.ndarray:
if isinstance(threshold, (list, tuple)):
return (p[:, 1:3].sum(1) >= threshold[0]) | (p[:, 3:5].sum(1) >= threshold[1])
return p_def(p) >= threshold
def dilate(mask: np.ndarray, k: int) -> np.ndarray:
if k <= 0:
return mask
out = mask.copy()
for d in range(1, k + 1):
out[d:] |= mask[:-d]
out[:-d] |= mask[d:]
return out
def score_rows(rows, probs_list, threshold, tol: int = 2):
"""rows: dicts with labels (answer part, >0 = defect), event_ids, defect_type, variant."""
ev_total = ev_hit = heads = head_hit = 0
sp_total = sp_good = 0
ev_hit_t = sp_good_t = 0
by_type_t = defaultdict(int)
row_tp = row_fp = row_fn = row_tn = 0
corr_spans = corr_tokens = 0
by_type = defaultdict(lambda: [0, 0])
for r, probs in zip(rows, probs_list):
lab = np.asarray(r["labels"])
pred = defect_mask(probs, threshold) & (lab != -100)
spans = spans_of(pred)
if r["variant"] == "corrected":
corr_spans += len(spans)
corr_tokens += int((lab != -100).sum())
continue
gold_pos = lab > 0
pred_near = dilate(pred, tol)
gold_near = dilate(gold_pos, tol)
events = defaultdict(list)
for i, (l, ev) in enumerate(zip(lab, r["event_ids"])):
if l > 0:
events[ev].append(i)
for ev, pos in events.items():
hit = bool(pred[pos].any())
ev_total += 1
ev_hit += hit
t = next((r["defect_type"][i] for i in pos if lab[i] == 1), r["defect_type"][pos[0]])
by_type[t][0] += 1
by_type[t][1] += hit
hit_t = bool(pred_near[pos].any())
ev_hit_t += hit_t
by_type_t[t] += hit_t
for i in np.flatnonzero(lab == 1):
heads += 1
head_hit += bool(pred[i])
for s, e in spans:
sp_total += 1
sp_good += bool(gold_pos[s:e].any())
sp_good_t += bool(gold_near[s:e].any())
has_gold, has_pred = bool(events), bool(spans)
row_tp += has_gold and has_pred
row_fp += (not has_gold) and has_pred
row_fn += has_gold and not has_pred
row_tn += (not has_gold) and not has_pred
rec = ev_hit / ev_total if ev_total else float("nan")
prec = sp_good / sp_total if sp_total else float("nan")
f1 = 2 * prec * rec / (prec + rec) if sp_total and ev_total and prec + rec > 0 else 0.0
rec_t = ev_hit_t / ev_total if ev_total else float("nan")
prec_t = sp_good_t / sp_total if sp_total else float("nan")
f1_t = 2 * prec_t * rec_t / (prec_t + rec_t) if sp_total and ev_total and prec_t + rec_t > 0 else 0.0
rp = row_tp / (row_tp + row_fp) if row_tp + row_fp else float("nan")
rr = row_tp / (row_tp + row_fn) if row_tp + row_fn else float("nan")
return {
"threshold": threshold, "event_recall": rec, "span_precision": prec, "event_f1": f1,
"head_recall": head_hit / heads if heads else float("nan"), "events": ev_total, "pred_spans": sp_total,
"row_precision": rp, "row_recall": rr, "rows_tp_fp_fn_tn": [row_tp, row_fp, row_fn, row_tn],
"corrected_fp_per_1k": 1000 * corr_spans / corr_tokens if corr_tokens else float("nan"),
"event_recall_tol2": rec_t, "span_precision_tol2": prec_t, "event_f1_tol2": f1_t,
"recall_by_type": {k: {"events": v[0], "recall": v[1] / v[0], "recall_tol2": by_type_t[k] / v[0]}
for k, v in sorted(by_type.items())},
}
def full_report(rows, probs_list, thresholds=None, fixed_threshold=None, tune_on: str = "event_f1_tol2"):
lab = np.concatenate([np.asarray(r["labels"]) for r in rows])
pd = np.concatenate([p_def(p) for p in probs_list])
keep = lab != -100
ap = average_precision(pd[keep], (lab[keep] > 0).astype(int))
if fixed_threshold is not None:
best = score_rows(rows, probs_list, fixed_threshold)
else:
grid = thresholds or [0.02, 0.05, 0.1, 0.15, 0.2, 0.25, 0.3, 0.35, 0.4, 0.45, 0.5, 0.55, 0.6, 0.65,
0.7, 0.75, 0.8, 0.85, 0.9, 0.93, 0.95, 0.97, 0.98, 0.99, 0.995, 0.998]
if probs_list[0].shape[1] == 5:
coarse = [0.02, 0.05, 0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9, 0.95, 0.98, 0.99, 0.995]
grid = [(a, b) for a in coarse for b in coarse]
best = max((score_rows(rows, probs_list, t) for t in grid), key=lambda m: m[tune_on])
best["ap_token"] = ap
half = (0.5, 0.5) if probs_list[0].shape[1] == 5 else 0.5
best["at_0.5"] = {k: v for k, v in score_rows(rows, probs_list, half).items() if k != "recall_by_type"}
return best
# ---- standard span-level evaluation: seqeval (via HF evaluate) over BIO tags ------------------------
_SEQEVAL = None
def seqeval_metric():
global _SEQEVAL
if _SEQEVAL is None:
import evaluate
_SEQEVAL = evaluate.load("seqeval")
return _SEQEVAL
def gold_bio(r, grouped: bool):
"""labels 0/1/2 (O/BAD/TAIL) -> BIO over reviewed positions only (label -100 dropped)."""
out = []
for l, t in zip(r["labels"], r["defect_type"]):
if l < 0:
continue
if l == 0:
out.append("O")
continue
ent = ("SCRIPT" if t in SCRIPT_TYPES else "GRAM") if grouped else "DEF"
out.append(("B-" if l == 1 else "I-") + ent)
return out
def pred_bio(p: np.ndarray, labels, threshold=None):
"""Predicted BIO: argmax by default; with a threshold, a group fires when its summed p >= threshold."""
k = p.shape[1]
groups = [("DEF", 1)] if k == 3 else [("SCRIPT", 1), ("GRAM", 3)]
out = []
for i, l in enumerate(labels):
if l < 0:
continue
if threshold is None:
c = int(p[i].argmax())
if c == 0:
out.append("O")
continue
name, base = groups[(c - 1) // 2]
out.append(("B-" if c == base else "I-") + name)
continue
thr = threshold if isinstance(threshold, (list, tuple)) else [threshold] * len(groups)
best = None
for (name, base), t in zip(groups, thr):
s = p[i, base] + p[i, base + 1]
if s >= t and (best is None or s > best[0]):
best = (s, name, base)
if best is None:
out.append("O")
else:
_, name, base = best
out.append(("B-" if p[i, base] >= p[i, base + 1] else "I-") + name)
return out
def seqeval_report(rows, probs_list, threshold=None, grouped=None):
"""seqeval on source answers; false alarms on repaired answers reported separately."""
grouped = probs_list[0].shape[1] == 5 if grouped is None else grouped
gold, pred, fp_ent, corr_tok = [], [], 0, 0
for r, p in zip(rows, probs_list):
pb = pred_bio(p, r["labels"], threshold)
if r["variant"] == "corrected":
fp_ent += sum(1 for t in pb if t.startswith("B-")) + sum(
1 for a, b in zip(["O"] + pb, pb) if b.startswith("I-") and a == "O")
corr_tok += len(pb)
continue
gold.append(gold_bio(r, grouped))
pred.append(pb)
res = seqeval_metric().compute(predictions=pred, references=gold, zero_division=0)
out = {"precision": res["overall_precision"], "recall": res["overall_recall"], "f1": res["overall_f1"],
"threshold": threshold, "corrected_fp_per_1k": 1000 * fp_ent / max(corr_tok, 1)}
for k, v in res.items():
if isinstance(v, dict):
out[k] = {m: float(v[m]) for m in ("precision", "recall", "f1", "number")}
return out
def seqeval_sweep(rows, probs_list):
"""Best threshold on these rows by overall seqeval F1: a single threshold for 3 classes, a full
2-D grid (script, grammar) for grouped models, so a group that never helps gets a high threshold."""
grid = [0.02, 0.05, 0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9, 0.95, 0.98, 0.99, 0.999]
cands = [(a, b) for a in grid for b in grid] if probs_list[0].shape[1] == 5 else grid
return max((seqeval_report(rows, probs_list, t) for t in cands), key=lambda r: r["f1"])
# ---- span-overlap evaluation (same seqeval entities, overlap matching instead of exact) -------------
def _entities(tags):
from seqeval.metrics.sequence_labeling import get_entities
return [(t, s, e + 1) for t, s, e in get_entities(tags)] # [start, end)
def overlap_report(rows, probs_list, threshold=None, grouped=None):
"""Entities are extracted from the BIO sequences exactly as seqeval does (get_entities);
a gold entity counts as found if any predicted entity of the same group overlaps it, and a
predicted entity counts as correct if it overlaps any gold entity of the same group.
'overall' ignores the group. False alarms on repaired answers are reported separately."""
grouped = probs_list[0].shape[1] == 5 if grouped is None else grouped
stats = defaultdict(lambda: [0, 0, 0, 0]) # gold, gold_found, pred, pred_correct
fp_corr = corr_tok = 0
for r, p in zip(rows, probs_list):
pb = pred_bio(p, r["labels"], threshold)
pe = _entities(pb)
if r["variant"] == "corrected":
fp_corr += len(pe)
corr_tok += len(pb)
continue
ge = _entities(gold_bio(r, grouped))
for key, same in (("overall", False), (None, True)):
for t, s, e in ge:
k = key or t
stats[k][0] += 1
stats[k][1] += any(ps < e and s < pe_ and (not same or pt == t) for pt, ps, pe_ in pe)
for t, s, e in pe:
k = key or t
stats[k][2] += 1
stats[k][3] += any(gs < e and s < ge_ and (not same or gt == t) for gt, gs, ge_ in ge)
def prf(v):
p_ = v[3] / v[2] if v[2] else 0.0
r_ = v[1] / v[0] if v[0] else 0.0
return {"precision": p_, "recall": r_, "f1": 2 * p_ * r_ / (p_ + r_) if p_ + r_ else 0.0, "number": v[0]}
out = prf(stats["overall"])
out |= {"threshold": threshold, "corrected_fp_per_1k": 1000 * fp_corr / max(corr_tok, 1),
"_counts": {k: list(v) for k, v in stats.items()} | {"_corrected": [fp_corr, corr_tok]}}
for k, v in stats.items():
if k != "overall":
out[k] = prf(v)
return out
def overlap_sweep(rows, probs_list):
"""Best threshold by overlap F1 (single for 3 classes, 2-D grid for grouped models)."""
grid = [0.02, 0.05, 0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9, 0.95, 0.98, 0.99, 0.999]
cands = [(a, b) for a in grid for b in grid] if probs_list[0].shape[1] == 5 else grid
return max((overlap_report(rows, probs_list, t) for t in cands), key=lambda r: r["f1"])