Router / policy.py
CompactAI's picture
Upload 11 files
35700de verified
Raw
History Blame Contribute Delete
10.1 kB
"""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)