"""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"])