| """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 |
| |
| 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) |
| |
| 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)) |
| |
| from dataclasses import replace |
| return [replace(result, model_rows=model_rows, cache_hits=hits) for result in results] |
|
|