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)
|