File size: 10,105 Bytes
35700de
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
"""Routing policy: everything you can change without touching the weights.

Supra-Router-51M bakes its decision into the model. two tiers, one threshold,
one domain vocabulary, all frozen at fine-tune time. Changing "route legal work
to the frontier model regardless of complexity" means collecting data and
retraining a 51M model.

Here the model's job stops at *calibrated probabilities about the prompt*. What
to do with them is this file, driven by a JSON document the user owns:

  * N tiers, not 2, each with a cost and a capability
  * expected-cost arithmetic instead of a hard-coded boolean
  * an abstain band that escalates ties rather than coin-flipping them
  * hard overrides on any predicted field
  * user-defined categories matched by prototype, no retraining

Also here: the freshness / tool-use / multilingual signals. Those are keyword
and character-class rules. a learned head for them would only be memorising
the regex below, which costs parameters and teaches nothing. They live where a
rule belongs, and the user can edit them.
"""

from __future__ import annotations

import json
import re
from dataclasses import asdict, dataclass, field
from pathlib import Path

import torch
from torch import Tensor

from modeling import BINARY_FIELDS, DOMAINS, GlintRouter, complexity_from_cumulative

FRESHNESS = re.compile(
    r"\b(today|tonight|current|currently|latest|right now|this (week|month|year)|"
    r"recent|news|stock price|weather|who won|as of \d{4}|202\d|203\d)\b", re.I)
TOOL_USE = re.compile(
    r"\b(search the web|look ?up|browse|run this|execute|call the api|fetch|"
    r"my file|this (csv|pdf|spreadsheet|repo)|book a|send an email|schedule)\b", re.I)


def non_ascii_ratio(text: str) -> float:
    letters = [c for c in text if c.isalpha()]
    if not letters:
        return 0.0
    return sum(1 for c in letters if ord(c) > 127) / len(letters)


@dataclass
class Tier:
    name: str
    cost: float          # relative price of one call, any unit you like
    capability: float    # 0..1; the difficulty this tier handles before it starts failing
    tools: bool = False  # can this tier search / execute?


@dataclass
class Policy:
    tiers: list[Tier] = field(default_factory=lambda: [
        Tier("local", cost=0.0, capability=0.35),
        Tier("mid", cost=1.0, capability=0.65),
        Tier("frontier", cost=8.0, capability=0.95, tools=True),
    ])
    # A wrong answer must cost several times the most expensive call, or the
    # arithmetic degenerates: when every tier is likely to fail, "pay least"
    # beats "try hardest", and the router quietly sends its hardest prompts to
    # the weakest model. Penalty >> max(cost) is what makes escalation rational.
    failure_penalty: float = 50.0   # what a wrong answer costs, in the same units as `cost`
    # Sharpness sets the residual failure rate a tier carries on prompts well
    # inside its capability, and it has to be steep. At sharpness 12 the local
    # tier still shows a 2.6% failure rate on a difficulty-0.05 prompt, which
    # priced at `failure_penalty` (1.33) loses to the mid tier's flat cost of
    # 1.0, so a free tier that handles the prompt perfectly never gets used.
    # 20 puts that residual at 0.25% and the cheap tier wins what it should.
    sharpness: float = 20.0
    abstain_margin: float = 0.05    # ties inside this band escalate instead of coin-flipping
    difficulty_weights: tuple[float, float, float] = (0.6, 0.3, 0.1)  # route, complexity, code/math
    # {"when": {...}, "tier": "frontier"}. first match wins, checked before the arithmetic.
    overrides: list[dict] = field(default_factory=list)
    freshness_needs_tools: bool = True
    multilingual_threshold: float = 0.05
    prototype_threshold: float = 0.55

    @classmethod
    def load(cls, path: Path | None) -> Policy:
        if path is None:
            return cls()
        blob = json.loads(Path(path).read_text())
        tiers = [Tier(**t) for t in blob.pop("tiers", [])]
        policy = cls(**blob)
        if tiers:
            policy.tiers = tiers
        return policy

    def save(self, path: Path) -> None:
        Path(path).write_text(json.dumps(asdict(self), indent=2))


@dataclass
class Decision:
    tier: str
    domain: str
    complexity: int
    difficulty: float
    probabilities: dict[str, float]
    signals: dict[str, bool]
    reason: str
    category: str | None = None
    escalated: bool = False

    def as_line(self) -> str:
        """Supra's output format, so the two are directly comparable."""
        p = self.probabilities
        return (f"Domain: {self.domain} | Complexity: {self.complexity} | "
                f"Math: {'T' if p['math'] > 0.5 else 'F'} | "
                f"Code: {'T' if p['code'] > 0.5 else 'F'} | "
                f"Route: {self.tier} | Justification: {self.reason}")


class Prototypes:
    """User-defined categories, added from a handful of examples at runtime.

    The projection head is trained jointly with the classifier, so prompts that
    share a routing-relevant character land near each other. A new category is
    the mean of its examples' projections. adding one is a forward pass over
    5-10 prompts, not a training run.
    """

    def __init__(self, vectors: dict[str, list[float]] | None = None) -> None:
        self.vectors = {k: torch.tensor(v) for k, v in (vectors or {}).items()}

    def add(self, name: str, projections: Tensor) -> None:
        centroid = projections.mean(dim=0)
        self.vectors[name] = centroid / centroid.norm().clamp(min=1e-6)

    def match(self, projection: Tensor, threshold: float) -> tuple[str | None, float]:
        if not self.vectors:
            return None, 0.0
        names = list(self.vectors)
        bank = torch.stack([self.vectors[n] for n in names]).to(projection.device)
        scores = bank @ projection
        best = int(scores.argmax())
        score = float(scores[best])
        return (names[best], score) if score >= threshold else (None, score)

    def save(self, path: Path) -> None:
        Path(path).write_text(json.dumps({k: v.tolist() for k, v in self.vectors.items()}))

    @classmethod
    def load(cls, path: Path | None) -> Prototypes:
        if path is None or not Path(path).is_file():
            return cls()
        return cls(json.loads(Path(path).read_text()))


def difficulty_of(route_p: float, complexity: float, code_p: float, math_p: float,
                  weights: tuple[float, float, float]) -> float:
    w_route, w_complexity, w_technical = weights
    scaled_complexity = (complexity - 1.0) / 4.0
    technical = max(code_p, math_p)
    total = w_route + w_complexity + w_technical
    return (w_route * route_p + w_complexity * scaled_complexity
            + w_technical * technical) / max(total, 1e-6)


def _override_hit(rule: dict, domain: str, complexity: int, probs: dict[str, float],
                  signals: dict[str, bool], category: str | None) -> bool:
    when = rule.get("when", {})
    if "domain" in when and when["domain"] != domain:
        return False
    if "category" in when and when["category"] != category:
        return False
    if "complexity_min" in when and complexity < when["complexity_min"]:
        return False
    if "complexity_max" in when and complexity > when["complexity_max"]:
        return False
    for f in BINARY_FIELDS:
        if f in when and bool(when[f]) != (probs[f] > 0.5):
            return False
    for s in ("freshness", "tool_use", "multilingual"):
        if s in when and bool(when[s]) != signals[s]:
            return False
    return True


def decide(model: GlintRouter, tokens: Tensor, text: str, policy: Policy,
           prototypes: Prototypes | None = None) -> Decision:
    """One prompt -> one routing decision. `tokens` is a (1, seq) batch."""
    with torch.no_grad():
        out = model.calibrated(tokens)
    probs = {f: float(out["binary"][0, i]) for i, f in enumerate(BINARY_FIELDS)}
    probs["route_big"] = float(out["route"][0])
    domain = DOMAINS[int(out["domain"][0].argmax())]
    complexity_p = out["complexity"][0]
    complexity = int(complexity_from_cumulative(complexity_p))
    expected_complexity = 1.0 + float(complexity_p.sum())

    signals = {
        "freshness": bool(FRESHNESS.search(text)),
        "tool_use": bool(TOOL_USE.search(text)),
        "multilingual": non_ascii_ratio(text) > policy.multilingual_threshold,
    }
    category, _ = (prototypes.match(out["proj"][0], policy.prototype_threshold)
                   if prototypes else (None, 0.0))

    for rule in policy.overrides:
        if _override_hit(rule, domain, complexity, probs, signals, category):
            return Decision(rule["tier"], domain, complexity, 0.0, probs, signals,
                            f"override: {json.dumps(rule.get('when', {}))}", category)

    difficulty = difficulty_of(probs["route_big"], expected_complexity,
                               probs["code"], probs["math"], policy.difficulty_weights)

    tiers = policy.tiers
    if (signals["freshness"] or signals["tool_use"]) and policy.freshness_needs_tools:
        with_tools = [t for t in tiers if t.tools]
        if with_tools:
            tiers = with_tools

    costs = []
    for tier in tiers:
        fail = torch.sigmoid(torch.tensor(
            (difficulty - tier.capability) * policy.sharpness)).item()
        costs.append((tier.cost + policy.failure_penalty * fail, tier))
    costs.sort(key=lambda pair: pair[0])

    best_cost, best = costs[0]
    escalated = False
    if len(costs) > 1 and (costs[1][0] - best_cost) <= policy.abstain_margin * best_cost:
        runner_up = costs[1][1]
        if runner_up.capability > best.capability:
            best, escalated = runner_up, True

    reason = (f"difficulty={difficulty:.2f} complexity={complexity} "
              f"code={probs['code']:.2f} math={probs['math']:.2f}"
              + (" (escalated: tie inside abstain band)" if escalated else ""))
    return Decision(best.name, domain, complexity, difficulty, probs, signals,
                    reason, category, escalated)