File size: 9,280 Bytes
dfba102
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
"""OpenThai-SystemOne decision model.

text tower (Qwen3.5-0.8B, LM head removed)  ->  hidden state at every <|ts_answer|>  ->  SlotHead (256 logits)
                                                 mask slots >= k  ->  softmax  ->  probabilities over the k options
"""
from __future__ import annotations

import math
from dataclasses import dataclass
from typing import Optional

import torch
import torch.nn as nn
import torch.nn.functional as F
from transformers import AutoModel, AutoModelForCausalLM, AutoTokenizer, PreTrainedModel
from transformers.utils import ModelOutput

from .configuration import OpenThaiSystemOneConfig
from .formatting import SPECIAL_TOKENS, TOK_ANSWER, add_special_tokens

QTYPE_INDEX = {"choice": 0, "score": 1, "noul": 2}


@dataclass
class DecisionOutput(ModelOutput):
    loss: Optional[torch.Tensor] = None
    logits: Optional[torch.Tensor] = None  # (B, Q, n_slots), masked with -inf
    probs: Optional[torch.Tensor] = None  # (B, Q, n_slots)
    hidden_states: Optional[torch.Tensor] = None  # (B, Q, H) at answer positions


class OpenThaiSystemOneForDecision(PreTrainedModel):
    config_class = OpenThaiSystemOneConfig
    base_model_prefix = "model"
    supports_gradient_checkpointing = True
    _supports_flash_attn = True
    _supports_sdpa = True

    def __init__(self, config: OpenThaiSystemOneConfig):
        super().__init__(config)
        self.model = AutoModel.from_config(config.text_config)
        self.slot_head = nn.Linear(config.hidden_size, config.n_slots, bias=config.head_bias)
        # log-temperatures per question type (choice/score/noul); learned in the calibration stage
        self.log_temperature = nn.Parameter(torch.zeros(config.n_temperatures))
        self.post_init()

    # ------------------------------------------------------------------ construction
    @classmethod
    def from_causal_lm(
        cls,
        path: str,
        *,
        tokenizer=None,
        n_slots: int = 256,
        torch_dtype=torch.bfloat16,
        **kwargs,
    ):
        """Build a decision model from a (text-only) causal-LM checkpoint: drop lm_head, add tokens + head."""
        tok = tokenizer or AutoTokenizer.from_pretrained(path)
        added = add_special_tokens(tok)
        lm = AutoModelForCausalLM.from_pretrained(path, dtype=torch_dtype, **kwargs)
        base = lm.model if hasattr(lm, "model") else lm.base_model
        text_cfg = base.config
        if added:
            lm.resize_token_embeddings(len(tok), mean_resizing=False)
            text_cfg.vocab_size = lm.get_input_embeddings().weight.shape[0]
            _init_new_token_embeddings(lm.get_input_embeddings().weight, tok, added)
        cfg = OpenThaiSystemOneConfig(
            text_config=text_cfg,
            n_slots=n_slots,
            answer_token_id=tok.convert_tokens_to_ids(TOK_ANSWER),
            pad_token_id=tok.pad_token_id,
        )
        cfg.text_config.tie_word_embeddings = False  # there is no LM head any more
        model = cls(cfg).to(torch_dtype)
        missing, unexpected = model.model.load_state_dict(base.state_dict(), strict=False)
        assert not unexpected, unexpected
        _init_slot_head(model.slot_head)
        model.model.config = cfg.text_config
        return model, tok

    # ------------------------------------------------------------------ forward
    def gather_answer_states(self, hidden: torch.Tensor, answer_positions: torch.Tensor) -> torch.Tensor:
        idx = answer_positions.clamp(min=0).unsqueeze(-1).expand(-1, -1, hidden.shape[-1])
        return torch.gather(hidden, 1, idx)  # (B, Q, H)

    def slot_logits(
        self,
        answer_hidden: torch.Tensor,
        option_counts: torch.Tensor,
        *,
        include_abstain: bool = True,
        qtypes: Optional[torch.Tensor] = None,
        apply_temperature: bool = True,
    ) -> torch.Tensor:
        logits = self.slot_head(answer_hidden.to(self.slot_head.weight.dtype)).float()
        if apply_temperature:
            if qtypes is None:
                t = self.log_temperature[0].exp()
            else:
                t = self.log_temperature.exp()[qtypes.clamp(min=0)]  # (B, Q)
                t = t.unsqueeze(-1)
            logits = logits / t
        ar = torch.arange(logits.shape[-1], device=logits.device)
        valid = ar[None, None, :] < option_counts.unsqueeze(-1)
        if include_abstain:
            valid = valid.clone()
            valid[..., self.config.abstain_slot] = True
        # questions that are padding (option_counts == 0) keep slot 0 valid to avoid NaNs
        valid[..., 0] |= option_counts.unsqueeze(-1).squeeze(-1) == 0
        return logits.masked_fill(~valid, float("-inf"))

    def forward(
        self,
        input_ids: torch.Tensor,
        attention_mask: Optional[torch.Tensor] = None,
        answer_positions: Optional[torch.Tensor] = None,
        option_counts: Optional[torch.Tensor] = None,
        labels: Optional[torch.Tensor] = None,
        soft_labels: Optional[torch.Tensor] = None,
        qtypes: Optional[torch.Tensor] = None,
        include_abstain: bool = True,
        label_smoothing: float = 0.0,
        brier_weight: float = 0.0,
        apply_temperature: bool = True,
        **kwargs,
    ) -> DecisionOutput:
        out = self.model(input_ids=input_ids, attention_mask=attention_mask, **kwargs)
        hidden = out.last_hidden_state
        if answer_positions is None:
            answer_positions = (input_ids == self.config.answer_token_id).nonzero()[:, 1].unsqueeze(0)
            if option_counts is None:
                raise ValueError("option_counts required")
        h = self.gather_answer_states(hidden, answer_positions)
        logits = self.slot_logits(h, option_counts, include_abstain=include_abstain, qtypes=qtypes, apply_temperature=apply_temperature)
        probs = logits.softmax(-1)

        loss = None
        if labels is not None or soft_labels is not None:
            logp = logits.log_softmax(-1)
            if soft_labels is not None:
                valid = (option_counts > 0)
                tgt = soft_labels.float()
                nll = -(tgt * logp.masked_fill(torch.isinf(logp), 0.0)).sum(-1)
                loss = (nll * valid).sum() / valid.sum().clamp(min=1)
            else:
                flat_logp = logp.reshape(-1, logp.shape[-1])
                flat_lab = labels.reshape(-1)
                keep = flat_lab != -100
                if keep.any():
                    lp = flat_logp[keep]
                    lb = flat_lab[keep]
                    nll = -lp.gather(1, lb[:, None]).squeeze(1)
                    if label_smoothing > 0:
                        n_valid = torch.isfinite(lp).sum(-1).clamp(min=1).float()
                        smooth = -(lp.masked_fill(torch.isinf(lp), 0.0)).sum(-1) / n_valid
                        nll = (1 - label_smoothing) * nll + label_smoothing * smooth
                    loss = nll.mean()
                    if brier_weight > 0:
                        p = lp.exp()
                        onehot = F.one_hot(lb, p.shape[-1]).float()
                        loss = loss + brier_weight * ((p - onehot) ** 2).sum(-1).mean()
                else:
                    loss = logits.sum() * 0.0
        return DecisionOutput(loss=loss, logits=logits, probs=probs, hidden_states=h)


def use_reference_kernels():
    """Force the pure-PyTorch Gated-DeltaNet / causal-conv paths.

    transformers routes `chunk_gated_delta_rule` & co. to the Triton kernels (flash-linear-attention, causal-conv1d)
    whenever those packages are importable, without checking the tensor device, which crashes on CPU/MPS.
    Call this before running on a non-CUDA device.
    """
    try:
        from transformers.models.qwen3_5 import modeling_qwen3_5 as m
    except Exception:  # pragma: no cover
        return
    for name in ("torch_chunk_gated_delta_rule", "torch_recurrent_gated_delta_rule", "chunk_gated_delta_rule",
                 "fused_recurrent_gated_delta_rule", "causal_conv1d_fn", "causal_conv1d_update"):
        fn = getattr(m, name, None)
        if fn is not None and hasattr(fn, "__wrapped__"):
            setattr(m, name, fn.__wrapped__)


def _init_slot_head(head: nn.Linear):
    nn.init.normal_(head.weight, std=0.02)
    if head.bias is not None:
        nn.init.zeros_(head.bias)


@torch.no_grad()
def _init_new_token_embeddings(weight: torch.Tensor, tok, n_added: int):
    """New control tokens start near the mean of digit-token embeddings + small noise."""
    digit_ids = [tok.convert_tokens_to_ids(d) for d in "0123456789"]
    digit_ids = [i for i in digit_ids if i is not None and i != tok.unk_token_id]
    mean = weight[digit_ids].float().mean(0) if digit_ids else weight[: weight.shape[0] - n_added].float().mean(0)
    std = weight[: weight.shape[0] - n_added].float().std()
    new = mean[None, :] + torch.randn(n_added, weight.shape[1]) * std * 0.1
    weight[-n_added:] = new.to(weight.dtype)


def confidence_from_probs(p: torch.Tensor, k: int) -> float:
    """1 - normalised entropy over the k valid options."""
    if k <= 1:
        return 1.0
    p = p[:k].clamp(min=1e-12)
    p = p / p.sum()
    h = -(p * p.log()).sum().item()
    return float(max(0.0, min(1.0, 1.0 - h / math.log(k))))