File size: 6,685 Bytes
7b67921
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Tiny-Jev — a small "System One" decision model: typed Choice / Score / Noul answers with calibrated probabilities.

    from transformers import AutoModel, AutoTokenizer
    tok = AutoTokenizer.from_pretrained("lostargon/Tiny-Jev")
    model = AutoModel.from_pretrained("lostargon/Tiny-Jev", trust_remote_code=True).eval()

    model.choice(tok, "My card was charged twice and nobody answers.", "Which team should handle this", ["billing", "technical", "sales"])
    # {'choice': 'billing', 'probabilities': {'billing': 0.97, 'technical': 0.02, 'sales': 0.01}, 'confidence': 0.97}
    model.noul(tok, "Customer: I want my money back now.", "The customer requests a refund")          # 0.98
    model.score(tok, "Arrived on time, works great.", "How satisfied is the reviewer", ["very unhappy", "neutral", "very happy"])
    model.decide(tok, state, [{"kind": "choice", "instructions": ..., "criteria": {...}}, {"kind": "noul", "instructions": ...}])

The whole model — decoder stack and decision head — lives in one safetensors file. No LoRA, no generation.
"""
from __future__ import annotations
import json
import torch
import torch.nn as nn
import torch.nn.functional as F
from transformers import AutoConfig, AutoModel
from transformers.models.qwen3.configuration_qwen3 import Qwen3Config
from transformers.models.qwen3.modeling_qwen3 import Qwen3Model, Qwen3PreTrainedModel


class TinyJevConfig(Qwen3Config):
    model_type = "tiny_jev"

    def __init__(self, temperature: float = 1.0, opt_token: str = "<opt>", cue: str = " - correct?", **kwargs):
        super().__init__(**kwargs)
        self.temperature = temperature
        self.opt_token = opt_token
        self.cue = cue


class TinyJevModel(Qwen3PreTrainedModel):
    config_class = TinyJevConfig

    def __init__(self, config: TinyJevConfig):
        super().__init__(config)
        self.model = Qwen3Model(config)
        self.head = nn.Linear(config.hidden_size, 1)
        self.post_init()

    # ---- core ---------------------------------------------------------------------------------------------
    def forward(self, input_ids, attention_mask, opt_positions, opt_mask=None):
        """opt_positions: [B, O] token indices of the option markers; returns logits [B, O] (masked with -inf)."""
        h = self.model(input_ids=input_ids, attention_mask=attention_mask).last_hidden_state
        g = torch.gather(h, 1, opt_positions.unsqueeze(-1).expand(-1, -1, h.shape[-1]))
        g = F.layer_norm(g.float(), (h.shape[-1],)).to(self.head.weight.dtype)
        logits = self.head(g).float().squeeze(-1) / self.config.temperature
        if opt_mask is not None:
            logits = logits.masked_fill(~opt_mask, float("-inf"))
        return logits

    # ---- prompt building ---------------------------------------------------------------------------------
    def encode(self, tok, state: str, question: str, options: list[str], max_len: int = 4096):
        ids_of = lambda t: tok(t, add_special_tokens=False)["input_ids"]
        opt_id = tok.convert_tokens_to_ids(self.config.opt_token)
        pre = ids_of("<state>\n")
        mid = ids_of("\n</state>\n<question>\n" + question + "\n</question>\n<options>\n")
        opts = [ids_of(o + self.config.cue) + [opt_id] + ids_of("\n") for o in options]
        post = ids_of("</options>")
        budget = max_len - (len(pre) + len(mid) + sum(map(len, opts)) + len(post))
        st = ids_of(state)
        if len(st) > budget:                       # truncate the state only, keeping head and tail
            head = max(0, int(budget * 0.7))
            st = st[:head] + st[-(budget - head):] if budget > 0 else []
        ids, pos = pre + st + mid, []
        for o in opts:
            pos.append(len(ids) + len(o) - 3)      # the cue's last token, right before <opt>
            ids += o
        return ids + post, pos

    @torch.no_grad()
    def probabilities(self, tok, state, questions: list[tuple[str, list[str]]]) -> list[list[float]]:
        state = state if isinstance(state, str) else json.dumps(state, ensure_ascii=False)
        encs = [self.encode(tok, state, q, o) for q, o in questions]
        L, O = max(len(e[0]) for e in encs), max(len(e[1]) for e in encs)
        pad = tok.pad_token_id if tok.pad_token_id is not None else 0
        dev = next(self.parameters()).device
        ids = torch.full((len(encs), L), pad, dtype=torch.long); att = torch.zeros_like(ids)
        pos = torch.zeros((len(encs), O), dtype=torch.long); mask = torch.zeros((len(encs), O), dtype=torch.bool)
        for i, (e, p) in enumerate(encs):
            ids[i, :len(e)] = torch.tensor(e); att[i, :len(e)] = 1
            pos[i, :len(p)] = torch.tensor(p); mask[i, :len(p)] = True
        logits = self(ids.to(dev), att.to(dev), pos.to(dev), mask.to(dev))
        probs = F.softmax(logits, -1).cpu()
        return [probs[i, :len(p)].tolist() for i, (_, p) in enumerate(encs)]

    # ---- typed API -----------------------------------------------------------------------------------------
    def decide(self, tok, state, questions: list[dict]) -> list[dict]:
        specs, keys = [], []
        for q in questions:
            crit = q.get("criteria")
            if q["kind"] == "noul":
                specs.append((q["instructions"], ["no", "yes"])); keys.append(None)
            elif isinstance(crit, dict):
                specs.append((q["instructions"], [f"{k}: {v}" for k, v in crit.items()])); keys.append(list(crit))
            else:
                specs.append((q["instructions"], list(crit))); keys.append(list(crit))
        out = []
        for q, k, p in zip(questions, keys, self.probabilities(tok, state, specs)):
            if q["kind"] == "noul":
                out.append({"noul": p[1]})
            elif q["kind"] == "choice":
                i = max(range(len(p)), key=p.__getitem__)
                out.append({"choice": k[i], "probabilities": dict(zip(k, p)), "confidence": p[i]})
            else:
                out.append({"score": sum(i * x for i, x in enumerate(p)), "probabilities": dict(zip(k, p)), "confidence": max(p)})
        return out

    def choice(self, tok, state, instructions, criteria): return self.decide(tok, state, [{"kind": "choice", "instructions": instructions, "criteria": criteria}])[0]
    def score(self, tok, state, instructions, levels): return self.decide(tok, state, [{"kind": "score", "instructions": instructions, "criteria": levels}])[0]
    def noul(self, tok, state, instructions): return self.decide(tok, state, [{"kind": "noul", "instructions": instructions}])[0]["noul"]


AutoConfig.register("tiny_jev", TinyJevConfig)
AutoModel.register(TinyJevConfig, TinyJevModel)