Julia-1 / julia /router /router.py
kleeedolinux
Publish Julia 1 model and Python runtime
5278d6b
Raw
History Blame Contribute Delete
8.49 kB
"""Bounded Julia inference with a Bend candidate-tree reducer."""
from collections import OrderedDict
from dataclasses import dataclass
import json
import math
import threading
from ..probabilities import display_probabilities
from .native import BendReducer
@dataclass(frozen=True)
class RouteResult:
index: int
candidates: tuple[int, ...]
probabilities: tuple[float, ...]
rounds: int
model_rows: int
cache_hits: int
hierarchical: bool
# Probabilities are conditional on candidates, never global for a tournament.
probability_scope: str
class Router:
"""Wrap an Engine-compatible logits(rows) implementation.
Up to width options: one unchanged model request. Larger choice requests:
retain survivors per group and rerank until one final group remains.
This increases supported option count, not the model's trained capacity.
"""
def __init__(self, engine, *, library=None, width=20, survivors=2,
batch_size=16, cache_size=0, max_options=4096):
for name, value in [('width', width), ('survivors', survivors),
('batch_size', batch_size), ('cache_size', cache_size),
('max_options', max_options)]:
if type(value) is not int:
raise ValueError(f'{name} must be an integer')
if not 2 <= width <= 20 or not 1 <= survivors < width:
raise ValueError('Require 2 <= width <= 20 and 1 <= survivors < width')
if batch_size < 1 or cache_size < 0 or max_options < width:
raise ValueError('Invalid batch/cache/capacity limit')
self.engine, self.reducer = engine, BendReducer(library)
self.width, self.survivors = width, survivors
self.batch_size, self.cache_size, self.max_options = batch_size, cache_size, max_options
self._cache = OrderedDict()
self._lock = threading.RLock()
def clear_cache(self):
"""Call after modifying model weights or inference settings."""
with self._lock:
self._cache.clear()
def _validate(self, row):
if not isinstance(row, dict):
raise ValueError('A request must be a dictionary')
if not isinstance(row.get('state'), (str, dict, list)) or not isinstance(row.get('question'), str):
raise ValueError('state must be text/JSON and question must be text')
options = row.get('options')
if not isinstance(options, list) or not 2 <= len(options) <= self.max_options:
raise ValueError(f'Expected 2–{self.max_options} options')
if not all(isinstance(x, str) and x for x in options):
raise ValueError('Options must be nonempty strings')
kind = row.get('type', 'choice')
if kind not in ('choice', 'score', 'noul'):
raise ValueError('Unknown decision type')
if kind == 'noul' and len(options) != 2:
raise ValueError('noul requires [false, true]')
if kind != 'choice' and len(options) > self.width:
raise ValueError('Hierarchical routing supports choice decisions only')
clean = dict(state=row['state'], question=row['question'], options=options, type=kind)
# Snapshot mutable state; preserve dict order used by Julia serialization.
return json.loads(json.dumps(clean, ensure_ascii=False, allow_nan=False))
def _score(self, jobs):
values = [None] * len(jobs)
missing = OrderedDict()
hits = 0
for i, row in enumerate(jobs):
key = json.dumps(row, ensure_ascii=False, separators=(',', ':'), allow_nan=False)
if key in self._cache:
values[i] = self._cache[key]
self._cache.move_to_end(key)
hits += 1
elif key in missing:
missing[key][1].append(i)
hits += 1
else:
missing[key] = (row, [i])
pending = list(missing.items())
for offset in range(0, len(pending), self.batch_size):
chunk = pending[offset:offset + self.batch_size]
output = list(self.engine.logits([entry[1][0] for entry in chunk]))
if len(output) != len(chunk):
raise ValueError('Engine returned the wrong number of rows')
for (key, (row, indices)), scores in zip(chunk, output):
scores = tuple(float(x) for x in scores)
if len(scores) != len(row['options']) or not all(math.isfinite(x) for x in scores):
raise ValueError('Engine logits must be finite and match option count')
for i in indices:
values[i] = scores
if self.cache_size:
self._cache[key] = scores
self._cache.move_to_end(key)
while len(self._cache) > self.cache_size:
self._cache.popitem(last=False)
return values, len(pending), hits
@staticmethod
def _confident_winner(scores, best):
"""Keep one candidate when its local softmax dominates the group."""
runner_up = max(score for i, score in enumerate(scores) if i != best)
if runner_up - scores[best] >= math.log(0.045 / 0.95):
return False
total = sum(math.exp(score - scores[best]) for score in scores)
return 1 / total > 0.95 and math.exp(runner_up - scores[best]) / total < 0.045
def route(self, row):
return self.route_many([row])[0]
def route_many(self, rows):
"""Batch independent groups across requests; preserve request/option order.
Calls on one Router serialize because most model engines and the LRU are
mutable. The C bridge also serializes access to Bend's global runtime.
"""
with self._lock:
requests = [self._validate(row) for row in rows]
candidates = [list(range(len(row['options']))) for row in requests]
results = [None] * len(requests)
rounds = [0] * len(requests)
model_rows = hits = 0
while any(result is None for result in results):
jobs, layout = [], []
for i, row in enumerate(requests):
if results[i] is not None:
continue
rounds[i] += 1
current = candidates[i]
final = len(current) <= self.width
groups = [current[j:j + self.width] for j in range(0, len(current), self.width)]
candidates[i] = []
for group in groups:
if len(group) == 1:
candidates[i].extend(group)
continue
jobs.append(dict(row, options=[row['options'][k] for k in group]))
layout.append((i, group, final))
scored, used, cached = self._score(jobs)
model_rows += used
hits += cached
for (i, group, final), scores in zip(layout, scored):
best = self.reducer.argmax(scores)
if final:
maximum = max(scores)
weights = [math.exp(x - maximum) for x in scores]
total = sum(weights)
results[i] = RouteResult(group[best], tuple(group),
tuple(display_probabilities([x / total for x in weights])), rounds[i], 0, 0,
len(requests[i]['options']) > self.width,
'final_candidates' if len(requests[i]['options']) > self.width else 'all_options')
elif self._confident_winner(scores, best):
candidates[i].append(group[best])
else:
remaining = list(range(len(group)))
chosen = []
for _ in range(min(self.survivors, len(group))):
position = self.reducer.argmax([scores[k] for k in remaining])
chosen.append(group[remaining.pop(position)])
candidates[i].extend(sorted(chosen))
# These counters are shared batch totals, not per-request attribution.
from dataclasses import replace
return [replace(result, model_rows=model_rows, cache_hits=hits) for result in results]