needle3-gpu / bart_backend.py
dkappe's picture
Upload bart_backend.py with huggingface_hub
f37e3f0 verified
Raw History Blame Contribute Delete
7.2 kB
#!/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}