AudioDecisionModel / server_runtime /learned_modality_router.py
gojiteji's picture
Audio Decision Model demo
148af80
Raw History Blame Contribute Delete
6.18 kB
"""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