#!/usr/bin/env python3 """ BART-MNLI backend: implements Jev's three primitives on facebook/bart-large-mnli. WHY, AND WHAT WE ARE TESTING The strongest "Jev is not new" argument (Medium, "What Jev Probably Is, And Why You Already Have One") is that the three primitives are ~20 lines on top of a free zero-shot classifier. It gives that code, using facebook/bart-large-mnli, and poses the experiment nobody has published: do the primitives work, and are the probabilities honest? We already have Needle 3 answering the same question at 87.5% held out on a 28-ticket labelled set. This backend runs the SAME questions over the SAME set so the two can be compared directly. HOW IT MAPS ONTO THE PRIMITIVES (following both the article and Jev's docs) choice -> zero-shot over the option NAMES. Jev's `choice` returns probabilities across options, so we keep the classifier's scores as a distribution. score -> zero-shot over the level descriptions, then the probability-weighted MEAN over the ordered levels. This is exactly what Jev's docs describe ("a position along your levels... can fall between two of them") and exactly what the article's `score()` does. noul -> zero-shot over (yes_text, no_text). The article flags the real weakness here: "writing the no side of a yes-or-no question by hand is clumsy, and the wording changes the answer." We take the yes probability. TWO CONFIGURATION CHOICES THAT MATTER, AND WHY 1. `multi_label`: - False -> scores are SOFTMAXED across the candidates and sum to 1. The options compete. This is what the article's code uses. - True -> each candidate is scored independently (sigmoid), no normalisation. Closer to Jev's documented "each level is judged on its own against the state", but then the numbers are not a distribution until you normalise. We expose both (`multi_label` parameter) and use False by default to match the article, because that yields a real distribution and a comparable `confidence`. 2. HYPOTHESIS TEMPLATE. The zero-shot pipeline defaults to "This example is {}." -- for our option names ("billing", "frustrated but civil") that is weak. The standard MNLI template for this family is "This example is {}." but a topical one reads far better for our labels, so we allow a template and default to "This text is about {}." This is a real tuning knob: the same paper the article cites warns the wording of the menu drives the result, so it should be set explicitly rather than left implicit. NOTE ON HOSTING bart-large-mnli is 1.63 GB and ships a flax_model.msgpack, so it can load in this Gradio Space. It is loaded LAZILY and only when a BART request arrives, so the Needle path keeps its current cold-start cost. """ from __future__ import annotations import os from typing import Any BART_ID = "facebook/bart-large-mnli" # Module-level cache: {"pipe": ..., "template": ...} _BART: dict[str, Any] = {} def bart_available() -> bool: try: import transformers # noqa: F401 import torch # noqa: F401 return True except Exception: return False def load_bart(template: str | None = None, use_flax: bool = False): """Lazily load the zero-shot pipeline. 1.6 GB, so only on first BART call.""" template = template or "This text is about {}." if _BART.get("pipe") is not None and _BART.get("template") == template: return _BART["pipe"] from transformers import pipeline if use_flax: pipe = pipeline("zero-shot-classification", model=BART_ID, framework="flax", hypothesis_template=template) else: pipe = pipeline("zero-shot-classification", model=BART_ID, device=-1, hypothesis_template=template) _BART["pipe"] = pipe _BART["template"] = template return pipe def _rank(state: str, labels: list[str], template: str | None, multi_label: bool) -> dict[str, float]: """Score labels against state; return {label: probability}.""" pipe = load_bart(template) out = pipe(state, labels, multi_label=multi_label) return dict(zip(out["labels"], [float(s) for s in out["scores"]])) def bart_choice(state: str, instructions: str, criteria: dict, template: str | None = None, multi_label: bool = False) -> dict: """Jev `choice`: pick one option, with probabilities across the options. `instructions` is used as the hypothesis when the caller gives no per-option description, because the classifier needs a sentence to test entailment against. """ options = list(criteria.keys()) labels = [] for opt in options: d = criteria[opt] desc = d.get("description") if isinstance(d, dict) else d labels.append(f"{instructions} {desc}".strip() if desc else opt) probs = _rank(state, labels, template, multi_label) if multi_label: tot = sum(probs.values()) or 1.0 probs = {k: v / tot for k, v in probs.items()} by_opt = {opt: probs.get(lbl, 0.0) for opt, lbl in zip(options, labels)} best = max(by_opt, key=by_opt.get) if by_opt else None return {"type": "choice", "choice": best, "probabilities": by_opt, "confidence": by_opt.get(best, 0.0) if best else 0.0} def bart_score(state: str, instructions: str, levels: list[str], template: str | None = None, multi_label: bool = False) -> dict: """Jev `score`: position on an ordered scale, can land BETWEEN levels. Probability-weighted mean over ordered levels, as Jev's docs describe and as the article's score() does. """ labels = [f"{instructions} {lv}".strip() for lv in levels] probs = _rank(state, labels, template, multi_label) if multi_label: tot = sum(probs.values()) or 1.0 probs = {k: v / tot for k, v in probs.items()} dist = {str(i): probs.get(lbl, 0.0) for i, lbl in enumerate(levels)} score = sum(i * dist[str(i)] for i in range(len(levels))) return {"type": "score", "score": score, "legend": {str(i): lv for i, lv in enumerate(levels)}, "probabilities": dist, "confidence": max(dist.values()) if dist else 0.0} def bart_noul(state: str, instructions: str, yes_text: str | None = None, no_text: str | None = None, template: str | None = None, multi_label: bool = False) -> dict: """Jev `noul`: probability the answer is yes. Jev returns no separate confidence here, and neither do we: with two options the probability IS the confidence. The article flags the hand-written negative side as this approach's real weakness, so both sides are caller-supplied where possible. """ yes_text = yes_text or instructions no_text = no_text or f"the opposite of: {instructions}" probs = _rank(state, [yes_text, no_text], template, multi_label) if multi_label: tot = sum(probs.values()) or 1.0 probs = {k: v / tot for k, v in probs.items()} p_yes = probs.get(yes_text, 0.0) return {"type": "noul", "noul": p_yes}