Spaces:
Running on Zero
Running on Zero
| """Evaluation service (spec §6): seeded sampling with fixed item IDs, CIs on | |
| every metric, paired significance tests for baseline-vs-finetuned, item-level | |
| persistence for resume/reproduction, and rule-based diagnostics. | |
| Pure-Python metric implementations (LCS ROUGE-L, corpus-free BLEU-4 per item) | |
| keep the Space light; sacrebleu/bertscore hook in when enabled in limits.yaml. | |
| """ | |
| import math | |
| import random | |
| import re | |
| import statistics | |
| import time | |
| # ---------- sampling ---------- | |
| def sample_items(records: list[dict], n: int, seed: int) -> list[dict]: | |
| """Stable, seeded sample with item ids; same (records, n, seed) => same items.""" | |
| rng = random.Random(seed) | |
| idx = list(range(len(records))) | |
| rng.shuffle(idx) | |
| chosen = sorted(idx[: min(n, len(records))]) | |
| items = [] | |
| for i in chosen: | |
| msgs = records[i]["messages"] | |
| ref = next((m["content"] for m in reversed(msgs) if m["role"] == "assistant"), "") | |
| prompt = [m for m in msgs if m["role"] != "assistant"] | |
| items.append({"item_id": i, "prompt": prompt, "reference": ref}) | |
| return items | |
| # ---------- metrics ---------- | |
| def _norm(s): | |
| return re.sub(r"\s+", " ", re.sub(r"[^\w\s]", "", s.lower())).strip() | |
| def exact_match(pred, ref): | |
| return float(_norm(pred) == _norm(ref)) | |
| def token_f1(pred, ref): | |
| p, r = _norm(pred).split(), _norm(ref).split() | |
| if not p or not r: | |
| return 0.0 | |
| common = {} | |
| for t in p: | |
| common[t] = common.get(t, 0) | |
| overlap = 0 | |
| rc = {} | |
| for t in r: | |
| rc[t] = rc.get(t, 0) + 1 | |
| pc = {} | |
| for t in p: | |
| pc[t] = pc.get(t, 0) + 1 | |
| for t, c in pc.items(): | |
| overlap += min(c, rc.get(t, 0)) | |
| if overlap == 0: | |
| return 0.0 | |
| prec, rec = overlap / len(p), overlap / len(r) | |
| return 2 * prec * rec / (prec + rec) | |
| def rouge_l(pred, ref): | |
| a, b = _norm(pred).split(), _norm(ref).split() | |
| if not a or not b: | |
| return 0.0 | |
| dp = [0] * (len(b) + 1) | |
| for x in a: | |
| prev = 0 | |
| for j, y in enumerate(b, 1): | |
| cur = dp[j] | |
| dp[j] = prev + 1 if x == y else max(dp[j], dp[j - 1]) | |
| prev = cur | |
| lcs = dp[-1] | |
| prec, rec = lcs / len(a), lcs / len(b) | |
| return 0.0 if prec + rec == 0 else 2 * prec * rec / (prec + rec) | |
| def bleu4(pred, ref): | |
| p, r = _norm(pred).split(), _norm(ref).split() | |
| if len(p) == 0: | |
| return 0.0 | |
| logs = [] | |
| for n in range(1, 5): | |
| pn = [tuple(p[i:i + n]) for i in range(len(p) - n + 1)] | |
| rn = [tuple(r[i:i + n]) for i in range(len(r) - n + 1)] | |
| if not pn: | |
| return 0.0 | |
| rc = {} | |
| for g in rn: | |
| rc[g] = rc.get(g, 0) + 1 | |
| hit = 0 | |
| pc = {} | |
| for g in pn: | |
| pc[g] = pc.get(g, 0) + 1 | |
| for g, c in pc.items(): | |
| hit += min(c, rc.get(g, 0)) | |
| logs.append(math.log((hit + 1e-9) / len(pn))) | |
| bp = 1.0 if len(p) > len(r) else math.exp(1 - len(r) / max(len(p), 1)) | |
| return bp * math.exp(sum(logs) / 4) | |
| def unsupported_claim_rate(pred, ref): | |
| """Fraction of predicted sentences with <30% token overlap vs reference — | |
| a component of the hallucination ESTIMATE, not a truth measurement.""" | |
| sents = [s for s in re.split(r"(?<=[.!?])\s+", pred) if len(s.split()) >= 4] | |
| if not sents: | |
| return 0.0 | |
| ref_toks = set(_norm(ref).split()) | |
| bad = sum(1 for s in sents | |
| if len(set(_norm(s).split()) & ref_toks) / max(len(set(_norm(s).split())), 1) < 0.3) | |
| return bad / len(sents) | |
| # ---------- statistics ---------- | |
| def mean_ci(values, iters=2000, seed=0): | |
| """Bootstrap mean + 95% CI.""" | |
| if not values: | |
| return {"mean": 0.0, "ci_low": 0.0, "ci_high": 0.0, "n": 0} | |
| rng = random.Random(seed) | |
| n = len(values) | |
| means = sorted(statistics.fmean(rng.choices(values, k=n)) for _ in range(iters)) | |
| return {"mean": round(statistics.fmean(values), 4), | |
| "ci_low": round(means[int(0.025 * iters)], 4), | |
| "ci_high": round(means[int(0.975 * iters)], 4), "n": n} | |
| def paired_pvalue(base, post, iters=2000, seed=0): | |
| """Paired permutation test on mean difference (sign-flip).""" | |
| diffs = [b - a for a, b in zip(base, post)] | |
| if not diffs or all(d == 0 for d in diffs): | |
| return 1.0 | |
| obs = abs(statistics.fmean(diffs)) | |
| rng = random.Random(seed) | |
| hits = sum(1 for _ in range(iters) | |
| if abs(statistics.fmean([d if rng.random() < 0.5 else -d for d in diffs])) >= obs) | |
| return round(hits / iters, 4) | |
| # ---------- evaluation run ---------- | |
| METRICS = { | |
| "accuracy": exact_match, | |
| "token_f1": token_f1, | |
| "bleu": bleu4, | |
| "rougeL": rouge_l, | |
| "unsupported_claim_rate": unsupported_claim_rate, | |
| } | |
| def evaluate_items(generate_fn, items, existing: dict | None = None, | |
| max_new_tokens=192, progress=None): | |
| """generate_fn(prompt_messages) -> text. Resumable: pass previously | |
| completed per-item results as `existing` (item_id -> record) (spec P8).""" | |
| done = dict(existing or {}) | |
| for k, item in enumerate(items): | |
| iid = str(item["item_id"]) | |
| if iid in done: | |
| continue | |
| t0 = time.time() | |
| try: | |
| pred = generate_fn(item["prompt"]) | |
| except Exception as e: # noqa: BLE001 | |
| pred = f"[generation error: {type(e).__name__}]" | |
| latency = time.time() - t0 | |
| rec = {"item_id": item["item_id"], "prediction": pred, | |
| "reference": item["reference"], "latency_s": round(latency, 3), | |
| "pred_tokens": len(pred.split())} | |
| for name, fn in METRICS.items(): | |
| rec[name] = round(fn(pred, item["reference"]), 4) | |
| done[iid] = rec | |
| if progress: | |
| progress((k + 1) / len(items)) | |
| return done | |
| def summarize(item_results: dict, seed: int, level: str) -> dict: | |
| recs = list(item_results.values()) | |
| out = {"level": level, "seed": seed, "n_items": len(recs), | |
| "full_benchmark_executed": False, "metrics": {}} | |
| for name in METRICS: | |
| out["metrics"][name] = mean_ci([r[name] for r in recs], seed=seed) | |
| lat = sorted(r["latency_s"] for r in recs) | |
| out["metrics"]["latency_s"] = mean_ci([r["latency_s"] for r in recs], seed=seed) | |
| out["latency_p95_s"] = round(lat[int(0.95 * (len(lat) - 1))], 3) if lat else 0 | |
| out["avg_response_tokens"] = round(statistics.fmean([r["pred_tokens"] for r in recs]), 1) if recs else 0 | |
| # composite hallucination estimate (spec §6.2) — labeled estimate everywhere | |
| ucr = out["metrics"]["unsupported_claim_rate"]["mean"] | |
| fact = out["metrics"]["token_f1"]["mean"] | |
| out["hallucination_estimate"] = { | |
| "factual_consistency": round(fact, 3), | |
| "unsupported_claim_rate": round(ucr, 3), | |
| "composite_pct": round(100 * (0.6 * ucr + 0.4 * (1 - fact)), 1), | |
| "label": "Estimated hallucination risk — not a direct measurement of truthfulness", | |
| "judge_model": None, "human_verified": False, | |
| } | |
| return out | |
| def compare(baseline: dict, post: dict, base_items: dict, post_items: dict) -> dict: | |
| """Paired comparison (identical item ids) with significance (spec §6.3).""" | |
| rows, verdicts = [], [] | |
| shared = sorted(set(base_items) & set(post_items), key=int) | |
| higher_better = {"accuracy": True, "token_f1": True, "bleu": True, "rougeL": True, | |
| "unsupported_claim_rate": False, "latency_s": False} | |
| primary = {"accuracy", "token_f1", "rougeL", "bleu"} | |
| for name, hb in higher_better.items(): | |
| b = [base_items[i][name] for i in shared] | |
| p = [post_items[i][name] for i in shared] | |
| delta = round(statistics.fmean(p) - statistics.fmean(b), 4) if shared else 0.0 | |
| pval = paired_pvalue(b, p) | |
| sig = pval < 0.05 | |
| improved = (delta > 0) == hb and delta != 0 | |
| rows.append({"metric": name, "baseline": round(statistics.fmean(b), 4) if b else 0, | |
| "finetuned": round(statistics.fmean(p), 4) if p else 0, | |
| "change": delta, "p_value": pval, "significant": sig, | |
| "direction": "improved" if improved else ("degraded" if delta != 0 else "unchanged")}) | |
| if sig and name in primary: | |
| verdicts.append("improved" if improved else "degraded") | |
| elif sig and not improved and name in ("latency_s",): | |
| verdicts.append("minor-regression") | |
| if "degraded" in verdicts: | |
| overall = "Degraded" | |
| elif "improved" in verdicts: | |
| overall = "Improved" | |
| else: | |
| overall = "Neutral" | |
| return {"rows": rows, "overall": overall, "n_paired_items": len(shared), | |
| "method": "paired permutation test (sign-flip), alpha=0.05"} | |
| def diagnostics(summary_ds: dict, training_log: dict, comparison: dict) -> list[dict]: | |
| """Trial & Error panel: evidence-backed possible reasons (spec §6.3).""" | |
| out = [] | |
| n = summary_ds.get("samples", 0) | |
| losses = training_log.get("losses", []) | |
| if comparison["overall"] != "Improved": | |
| if n and n < 500: | |
| out.append({"reason": "Dataset too small", | |
| "evidence": f"only {n} training samples; <500 rarely shifts a pretrained model"}) | |
| if losses and len(losses) > 4 and losses[-1] > losses[0] * 0.9: | |
| out.append({"reason": "Learning rate too low or too few steps", | |
| "evidence": f"loss only moved {losses[0]:.2f}→{losses[-1]:.2f}"}) | |
| if losses and min(losses) < 0.5 and n < 2000: | |
| out.append({"reason": "Overfitting risk", | |
| "evidence": f"train loss reached {min(losses):.2f} on a small dataset"}) | |
| deg = [r for r in comparison["rows"] if r["direction"] == "degraded" and r["significant"]] | |
| if deg: | |
| out.append({"reason": "Catastrophic forgetting possible", | |
| "evidence": f"significant regressions: {', '.join(r['metric'] for r in deg)}"}) | |
| if not out: | |
| out.append({"reason": "Neutral result", | |
| "evidence": "no significant movement — consider more epochs or higher-quality data"}) | |
| return out | |