Instructions to use hfvladkon/bert_token_classification_detector with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use hfvladkon/bert_token_classification_detector with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("token-classification", model="hfvladkon/bert_token_classification_detector")# Load model directly from transformers import AutoTokenizer, AutoModelForTokenClassification tokenizer = AutoTokenizer.from_pretrained("hfvladkon/bert_token_classification_detector") model = AutoModelForTokenClassification.from_pretrained("hfvladkon/bert_token_classification_detector", device_map="auto") - Notebooks
- Google Colab
- Kaggle
Download code/metrics.py from hfvladkon/bert_token_classification_detector: direct link, hf CLI and curl.
- Browser
- Download file 15.2 kB
-
https://huggingface.co/hfvladkon/bert_token_classification_detector/resolve/main/code/metrics.py
- Command line
-
hf download hf://hfvladkon/bert_token_classification_detector/code/metrics.py
-
curl -L -o metrics.py https://huggingface.co/hfvladkon/bert_token_classification_detector/resolve/main/code/metrics.py
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] | |
| 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"]) | |