"""Learned sparse routing and option-level fusion for Audio Jev. The router consumes only ModernBERT embeddings of context, instructions, and criteria. It does not inspect aliases, task names, or schema keys. Its soft weights decide whether the speech expert, the non-verbal expert, or both run. """ from __future__ import annotations from collections import OrderedDict from dataclasses import dataclass import math from pathlib import Path from typing import Mapping import numpy as np import onnxruntime as ort from tokenizers import Tokenizer from .typed_audio import question_options, typed_answer ROUTES = ("speech", "nonverbal", "joint") ENTROPY_THRESHOLD = 0.40 JOINT_THRESHOLD = 0.10 NATURAL_ROUTER_WEIGHT = 0.60 def _softmax(values: np.ndarray) -> np.ndarray: values = np.asarray(values, dtype=np.float64) shifted = values - values.max(axis=-1, keepdims=True) result = np.exp(shifted) return result / result.sum(axis=-1, keepdims=True) @dataclass(frozen=True) class LearnedQuestionRoute: mode: str weights: tuple[float, float, float] option_weights: tuple[tuple[float, float, float], ...] entropy: float @property def experts(self) -> tuple[str, ...]: if self.mode == "language": return ("speech",) if self.mode == "sound": return ("nonverbal",) return ("speech", "nonverbal") def public(self) -> dict[str, object]: value: dict[str, object] = { "mode": self.mode, "experts": list(self.experts), "weights": {name: round(value, 4) for name, value in zip(ROUTES, self.weights)}, "entropy": round(self.entropy, 4), "method": "learned_sparse_router_ensemble_v1", } if self.mode == "joint": value["fusion"] = "learned_option_weighted_logit_fusion_v1" return value class LearnedModalityRouter: """Frozen text encoder and learned sparse-router ensemble.""" def __init__(self, model_dir: str | Path) -> None: root = Path(model_dir) options = ort.SessionOptions() options.intra_op_num_threads = 2 self.text = ort.InferenceSession( str(root / "text_q8.onnx"), sess_options=options, providers=["CPUExecutionProvider"], ) self.routers = tuple( ort.InferenceSession(str(root / name), sess_options=options, providers=["CPUExecutionProvider"]) for name in ("modality_router_core_fp32.onnx", "modality_router_natural_fp32.onnx") ) self.tokenizer = Tokenizer.from_file(str(root / "tokenizer" / "tokenizer.json")) self.cache: OrderedDict[str, np.ndarray] = OrderedDict() def _encode(self, text: str) -> np.ndarray: cached = self.cache.get(text) if cached is not None: self.cache.move_to_end(text) return cached ids = self.tokenizer.encode(text, add_special_tokens=True).ids if len(ids) > 128: ids = [*ids[:127], 2] input_ids = np.full((1, 128), 3, dtype=np.int64) attention_mask = np.zeros((1, 128), dtype=np.int64) input_ids[0, :len(ids)] = ids attention_mask[0, :len(ids)] = 1 vector = self.text.run( ["features"], {"input_ids": input_ids, "attention_mask": attention_mask} )[0][0].astype(np.float32, copy=False) self.cache[text] = vector if len(self.cache) > 1024: self.cache.popitem(last=False) return vector def route(self, context: str, question: Mapping) -> LearnedQuestionRoute: pairs = question_options(dict(question)) context_vector = self._encode(context) question_vector = self._encode(str(question["instructions"])) option_vectors = np.stack([self._encode(description) for _, description in pairs]) inputs = { "context": context_vector[None], "question": question_vector[None], "options": option_vectors[None], } core, natural = (runtime.run(["question_logits", "option_logits"], inputs) for runtime in self.routers) question_logits = (1 - NATURAL_ROUTER_WEIGHT) * core[0] + NATURAL_ROUTER_WEIGHT * natural[0] option_logits = (1 - NATURAL_ROUTER_WEIGHT) * core[1] + NATURAL_ROUTER_WEIGHT * natural[1] weights = _softmax(question_logits[0]) option_weights = _softmax(option_logits[0]) entropy = float(-(weights * np.log(np.clip(weights, 1e-12, 1))).sum() / math.log(len(ROUTES))) both = bool(weights[2] >= JOINT_THRESHOLD or entropy >= ENTROPY_THRESHOLD) mode = "joint" if both else ("language" if weights[0] >= weights[1] else "sound") return LearnedQuestionRoute( mode=mode, weights=tuple(float(value) for value in weights), option_weights=tuple(tuple(float(value) for value in row) for row in option_weights), entropy=entropy, ) def route_all(self, context: str, questions: Mapping[str, Mapping]) -> dict[str, LearnedQuestionRoute]: return {name: self.route(context, question) for name, question in questions.items()} def learned_fuse_answers( question: Mapping, language_answer: Mapping, sound_answer: Mapping, route: LearnedQuestionRoute, ) -> dict: """Fuse expert distributions with learned per-option modality weights.""" pairs = question_options(dict(question)) language = np.asarray([language_answer["probabilities"][key] for key, _ in pairs], dtype=np.float64) sound = np.asarray([sound_answer["probabilities"][key] for key, _ in pairs], dtype=np.float64) option_weights = np.asarray(route.option_weights, dtype=np.float64) language_weight = option_weights[:, 0] + 0.5 * option_weights[:, 2] sound_weight = option_weights[:, 1] + 0.5 * option_weights[:, 2] logits = language_weight * np.log(np.clip(language, 1e-8, 1)) logits += sound_weight * np.log(np.clip(sound, 1e-8, 1)) probabilities = _softmax(logits) answer = typed_answer(dict(question), probabilities.tolist()) answer["fusion"] = {"method": "learned_option_weighted_logit_fusion_v1"} return answer