File size: 7,200 Bytes
f37e3f0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
#!/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}