File size: 6,181 Bytes
148af80
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
"""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