tindang commited on
Commit
54d02f8
·
verified ·
1 Parent(s): f981b0d

laya-code: fine-tuned Laya code-relevance re-ranker

Browse files
encoder/config.json ADDED
@@ -0,0 +1,84 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "architectures": [
3
+ "ModernBertForMaskedLM"
4
+ ],
5
+ "attention_bias": false,
6
+ "attention_dropout": 0.0,
7
+ "bos_token_id": 50281,
8
+ "classifier_activation": "gelu",
9
+ "classifier_bias": false,
10
+ "classifier_dropout": 0.0,
11
+ "classifier_pooling": "mean",
12
+ "cls_token_id": 50281,
13
+ "decoder_bias": true,
14
+ "deterministic_flash_attn": false,
15
+ "dtype": "float32",
16
+ "embedding_dropout": 0.0,
17
+ "eos_token_id": 50282,
18
+ "global_attn_every_n_layers": 3,
19
+ "gradient_checkpointing": false,
20
+ "hidden_activation": "gelu",
21
+ "hidden_size": 1024,
22
+ "initializer_cutoff_factor": 2.0,
23
+ "initializer_range": 0.02,
24
+ "intermediate_size": 2624,
25
+ "layer_norm_eps": 1e-05,
26
+ "layer_types": [
27
+ "full_attention",
28
+ "sliding_attention",
29
+ "sliding_attention",
30
+ "full_attention",
31
+ "sliding_attention",
32
+ "sliding_attention",
33
+ "full_attention",
34
+ "sliding_attention",
35
+ "sliding_attention",
36
+ "full_attention",
37
+ "sliding_attention",
38
+ "sliding_attention",
39
+ "full_attention",
40
+ "sliding_attention",
41
+ "sliding_attention",
42
+ "full_attention",
43
+ "sliding_attention",
44
+ "sliding_attention",
45
+ "full_attention",
46
+ "sliding_attention",
47
+ "sliding_attention",
48
+ "full_attention",
49
+ "sliding_attention",
50
+ "sliding_attention",
51
+ "full_attention",
52
+ "sliding_attention",
53
+ "sliding_attention",
54
+ "full_attention"
55
+ ],
56
+ "local_attention": 128,
57
+ "max_position_embeddings": 8192,
58
+ "mlp_bias": false,
59
+ "mlp_dropout": 0.0,
60
+ "model_type": "modernbert",
61
+ "norm_bias": false,
62
+ "norm_eps": 1e-05,
63
+ "num_attention_heads": 16,
64
+ "num_hidden_layers": 28,
65
+ "pad_token_id": 50283,
66
+ "position_embedding_type": "absolute",
67
+ "repad_logits_with_grad": false,
68
+ "rope_parameters": {
69
+ "full_attention": {
70
+ "rope_theta": 160000.0,
71
+ "rope_type": "default"
72
+ },
73
+ "sliding_attention": {
74
+ "rope_theta": 10000.0,
75
+ "rope_type": "default"
76
+ }
77
+ },
78
+ "sep_token_id": 50282,
79
+ "sparse_pred_ignore_index": -100,
80
+ "sparse_prediction": false,
81
+ "tie_word_embeddings": true,
82
+ "transformers_version": "5.0.0",
83
+ "vocab_size": 50368
84
+ }
model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:1cc1a3716fcd0121a1aafab184d844b6df9b54a8237c883d05780e101cbae5f8
3
+ size 842609210
rl_agent_api.py ADDED
@@ -0,0 +1,78 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Jev-compatible inference for a saved RL Agent model: system_one(state, questions) -> typed answers."""
2
+ import json
3
+ import math
4
+ import os
5
+
6
+ import numpy as np
7
+ import torch
8
+
9
+ from rl_common import (QTYPES, amp_dtype, build_model, build_sequence, collate_items, confidence_from_probs,
10
+ render_options, temp_bucket)
11
+
12
+
13
+ class RLAgent:
14
+ def __init__(self, model_dir, device=None):
15
+ from safetensors.torch import load_file
16
+ from transformers import AutoTokenizer
17
+ with open(os.path.join(model_dir, "rl_agent_config.json")) as f:
18
+ self.cfg = json.load(f)
19
+ self.device = torch.device(device or ("cuda" if torch.cuda.is_available() else "cpu"))
20
+ self.tok = AutoTokenizer.from_pretrained(os.path.join(model_dir, "tokenizer"))
21
+ self.model = build_model(self.cfg, encoder_dir=os.path.join(model_dir, "encoder"))
22
+ self.model.load_state_dict(load_file(os.path.join(model_dir, "model.safetensors")), strict=True)
23
+ self.model.to(self.device).eval()
24
+ self.model.encoder.config.reference_compile = False # torch.compile is a loss on small batches / few SMs (T4)
25
+ self.temperature = self.cfg.get("temperature", [1.0, 1.0, 1.0])
26
+ self.temperature_by_options = self.cfg.get("temperature_by_options", {})
27
+ self.dtype = amp_dtype(self.cfg.get("amp_dtype", "fp16"))
28
+ if self.device.type == "cuda" and torch.cuda.get_device_capability(self.device)[0] < 8:
29
+ self.dtype = torch.float16 # e.g. a bf16-trained model evaluated on a T4
30
+
31
+ @staticmethod
32
+ def _to_internal(qdef):
33
+ t = qdef["type"]
34
+ crit = qdef.get("criteria")
35
+ if t == "choice" and isinstance(crit, list):
36
+ crit = {c: None for c in crit}
37
+ return {"t": t, "ins": qdef["instructions"] if isinstance(qdef["instructions"], str) else json.dumps(qdef["instructions"]),
38
+ "crit": crit}
39
+
40
+ @torch.no_grad()
41
+ def system_one(self, state, questions):
42
+ """questions: {id: {"type": "choice"|"score"|"noul", "instructions": ..., "criteria": ...}} (Jev request shape)."""
43
+ ids, items = list(questions.keys()), []
44
+ for qid in ids:
45
+ q = self._to_internal(questions[qid])
46
+ seq, markers = build_sequence(self.tok, state, q, self.cfg["max_len"], self.cfg["head_max_len"])
47
+ if len(markers) != len(render_options(q)):
48
+ raise ValueError("question %r: options do not fit in head_max_len=%d tokens" % (qid, self.cfg["head_max_len"]))
49
+ items.append({"ids": seq, "markers": markers, "qtype": QTYPES[q["t"]], "target": [0.0] * len(markers), "label": -1,
50
+ "episode": 0, "ep_step": 0, "ep_len": 1, "src": "api"})
51
+ b = collate_items([items], self.tok.pad_token_id)
52
+ use_amp = self.device.type == "cuda"
53
+ with torch.autocast(device_type=self.device.type, dtype=self.dtype, enabled=use_amp):
54
+ logits, act = self.model(b["input_ids"].to(self.device), b["attention_mask"].to(self.device),
55
+ b["marker_pos"].to(self.device), b["marker_mask"].to(self.device), b["qtype"].to(self.device))
56
+ logits, act = logits.float().cpu().numpy(), torch.softmax(act.float(), -1).cpu().numpy()
57
+ answers, n_tokens = {}, int(b["attention_mask"].sum())
58
+ for r, qid in enumerate(ids):
59
+ q = self._to_internal(questions[qid])
60
+ k = len(items[r]["markers"])
61
+ qt = QTYPES[q["t"]]
62
+ z = logits[r, :k] / self.temperature_by_options.get(temp_bucket(qt, k), self.temperature[qt])
63
+ p = np.exp(z - z.max())
64
+ p = p / p.sum()
65
+ ext = {"act_probability": float(act[r, 0])}
66
+ if q["t"] == "choice":
67
+ keys = list(q["crit"].keys())
68
+ answers[qid] = {"type": "choice", "choice": keys[int(p.argmax())],
69
+ "probabilities": {kk: round(float(v), 4) for kk, v in zip(keys, p)},
70
+ "confidence": round(confidence_from_probs(p, k), 4), "rl_agent": ext}
71
+ elif q["t"] == "score":
72
+ answers[qid] = {"type": "score", "score": round(float((np.arange(k) * p).sum()), 4),
73
+ "legend": {str(i): c for i, c in enumerate(q["crit"])},
74
+ "probabilities": {str(i): round(float(v), 4) for i, v in enumerate(p)},
75
+ "confidence": round(confidence_from_probs(p, k), 4), "rl_agent": ext}
76
+ else:
77
+ answers[qid] = {"type": "noul", "noul": round(float(p[1]), 4), "rl_agent": ext}
78
+ return {"model": "rl-agent", "answers": answers, "usage": {"input_tokens": n_tokens, "output_tokens": 0}}
rl_agent_config.json ADDED
@@ -0,0 +1,60 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "encoder": "answerdotai/ModernBERT-large",
3
+ "head_layers": 2,
4
+ "max_len": 512,
5
+ "head_max_len": 192,
6
+ "max_prefixes": 6,
7
+ "act_costs": {
8
+ "escalate": 0.5
9
+ },
10
+ "cost_wrong_act": 3.0,
11
+ "amp_dtype": "bf16",
12
+ "model_name": "laya-code",
13
+ "temperature": [
14
+ 1.6369030475616455,
15
+ 1.2514300346374512,
16
+ 0.9409746926053619
17
+ ],
18
+ "temperature_by_options": {
19
+ "choice:3-5": 1.7601518630981445,
20
+ "choice:6-10": 1.0000158548355103,
21
+ "score:3-5": 1.2514300346374512,
22
+ "noul:2": 0.9409746926053619,
23
+ "choice:11+": 0.10058280825614929,
24
+ "choice:2": 1.9063563346862793
25
+ },
26
+ "training": {
27
+ "updates": 7313,
28
+ "epochs_completed": 1,
29
+ "hours": 1.96,
30
+ "world_size": 1,
31
+ "fine_tuned_from_checkpoint": true
32
+ },
33
+ "finetune": {
34
+ "base": "laya-base (convaiinnovations/laya)",
35
+ "task": "code relevance (noul)",
36
+ "noul_temperature_fit": {
37
+ "split": "val",
38
+ "before": {
39
+ "T": 1.9834,
40
+ "nll": 0.5774,
41
+ "ece_hard": 0.1673,
42
+ "auroc_pos_vs_neg": 0.7045,
43
+ "mean_p": 0.3733,
44
+ "base_rate": 0.2741,
45
+ "prec_at_p>=0.5_hard": 0.3333,
46
+ "n": 5136
47
+ },
48
+ "after": {
49
+ "T": 0.941,
50
+ "nll": 0.5421,
51
+ "ece_hard": 0.0578,
52
+ "auroc_pos_vs_neg": 0.7045,
53
+ "mean_p": 0.2681,
54
+ "base_rate": 0.2741,
55
+ "prec_at_p>=0.5_hard": 0.3333,
56
+ "n": 5136
57
+ }
58
+ }
59
+ }
60
+ }
rl_common.py ADDED
@@ -0,0 +1,408 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """RL Agent shared code: config, Jev-style question rendering, model, proper-scoring rewards, metrics.
2
+
3
+ Kept Python 3.9 compatible so the same file runs on Kaggle and on a laptop smoke test.
4
+ """
5
+ import json
6
+ import math
7
+ import os
8
+ import random
9
+ from typing import Dict, List, Optional
10
+
11
+ import numpy as np
12
+ import torch
13
+ import torch.nn as nn
14
+ import torch.nn.functional as F
15
+ import torch.utils.checkpoint
16
+
17
+ QTYPES = {"choice": 0, "score": 1, "noul": 2}
18
+ QTYPE_NAMES = {v: k for k, v in QTYPES.items()}
19
+
20
+
21
+ # ----------------------------------------------------------------------------- config
22
+ def load_cfg(path: Optional[str] = None) -> Dict:
23
+ path = path or os.environ.get("RL_AGENT_CFG", "rl_agent_config.json")
24
+ with open(path) as f:
25
+ return json.load(f)
26
+
27
+
28
+ # ----------------------------------------------------------------------------- rendering
29
+ def serialize_state(state) -> str:
30
+ if isinstance(state, str):
31
+ return state
32
+ return json.dumps(state, ensure_ascii=False)
33
+
34
+
35
+ def render_options(q: Dict) -> List[str]:
36
+ """Option texts in label-index order. Noul is always [false, true] so p[1] == noul."""
37
+ t, crit = q["t"], q.get("crit")
38
+ if t == "choice":
39
+ return [k if not v else "%s: %s" % (k, v) for k, v in crit.items()]
40
+ if t == "score":
41
+ return ["level %d: %s" % (i, c) for i, c in enumerate(crit)]
42
+ crit = crit or {}
43
+ return ["false: " + (crit.get("false") or "no, the statement does not hold"),
44
+ "true: " + (crit.get("true") or "yes, the statement holds")]
45
+
46
+
47
+ def build_sequence(tok, state, q: Dict, max_len: int, head_max_len: int,
48
+ option_order: Optional[List[int]] = None, truncate_left: bool = False):
49
+ """[CLS] <type> instructions [SEP] [MASK] opt0 [MASK] opt1 ... [SEP] state [SEP].
50
+
51
+ Returns input_ids and the positions of the per-option [MASK] markers (in the given option order).
52
+ """
53
+ mask_tok = tok.mask_token
54
+ opts = render_options(q)
55
+ order = option_order if option_order is not None else list(range(len(opts)))
56
+ ins = str(q["ins"]).replace(mask_tok, " ")
57
+ head_ids = tok("%s question: %s" % (q["t"], ins), add_special_tokens=False)["input_ids"]
58
+ opt_ids = []
59
+ for i in order:
60
+ opt_ids.append([tok.mask_token_id] + tok(" " + opts[i].replace(mask_tok, " "), add_special_tokens=False)["input_ids"][:48])
61
+ opt_budget = head_max_len - sum(len(o) for o in opt_ids)
62
+ if opt_budget < 16: # too many / too long options: shrink every option text evenly
63
+ per = max(4, (head_max_len - 16) // max(1, len(opt_ids)))
64
+ opt_ids = [o[:per] for o in opt_ids]
65
+ opt_budget = head_max_len - sum(len(o) for o in opt_ids)
66
+ head_ids = head_ids[:max(8, opt_budget)]
67
+ ids = [tok.cls_token_id] + head_ids + [tok.sep_token_id]
68
+ markers = []
69
+ for o in opt_ids:
70
+ markers.append(len(ids))
71
+ ids.extend(o)
72
+ ids.append(tok.sep_token_id)
73
+ room = max(0, max_len - len(ids) - 1)
74
+ st = tok(serialize_state(state).replace(mask_tok, " "), add_special_tokens=False)["input_ids"]
75
+ st = st[-room:] if truncate_left else st[:room]
76
+ ids = ids + st + [tok.sep_token_id]
77
+ return ids[:max_len], [m for m in markers if m < max_len]
78
+
79
+
80
+ # ----------------------------------------------------------------------------- model
81
+ class DecisionModel(nn.Module):
82
+ """Pretrained bidirectional encoder (no LLM, no LoRA) + from-scratch decision head.
83
+
84
+ Each option gets a [MASK] marker; the head scores markers -> softmax over the question's options.
85
+ """
86
+
87
+ def __init__(self, encoder: nn.Module, head_layers: int = 2, n_act: int = 2, dropout: float = 0.1):
88
+ super().__init__()
89
+ self.encoder = encoder
90
+ d = encoder.config.hidden_size
91
+ nhead = max(1, d // 64)
92
+ layer = nn.TransformerEncoderLayer(d, nhead, 4 * d, dropout, batch_first=True, norm_first=True)
93
+ self.head = nn.TransformerEncoder(layer, head_layers, enable_nested_tensor=False) if head_layers > 0 else None
94
+ self.type_emb = nn.Embedding(3, d)
95
+ self.scorer = nn.Sequential(nn.LayerNorm(d), nn.Linear(d, d), nn.GELU(), nn.Linear(d, 1))
96
+ self.act_head = nn.Sequential(nn.Linear(d + 4, 256), nn.GELU(), nn.Linear(256, n_act))
97
+ self.register_buffer("temperature", torch.ones(3)) # per qtype, fitted post-hoc in evaluate.py
98
+ self.head_checkpointing = False
99
+
100
+ def forward(self, input_ids, attention_mask, marker_pos, marker_mask, qtype, detach_encoder: bool = False):
101
+ h = self.encoder(input_ids=input_ids, attention_mask=attention_mask).last_hidden_state
102
+ if detach_encoder:
103
+ h = h.detach()
104
+ h = h + self.type_emb(qtype)[:, None, :]
105
+ if self.head is not None:
106
+ pad = ~attention_mask.bool()
107
+ for layer in self.head.layers:
108
+ if self.head_checkpointing and self.training and torch.is_grad_enabled():
109
+ h = torch.utils.checkpoint.checkpoint(layer, h, None, pad, use_reentrant=False)
110
+ else:
111
+ h = layer(h, src_key_padding_mask=pad)
112
+ idx = marker_pos.clamp(min=0)[:, :, None].expand(-1, -1, h.size(-1))
113
+ m = torch.gather(h, 1, idx)
114
+ logits = self.scorer(m).squeeze(-1).float()
115
+ logits = logits.masked_fill(~marker_mask, -1e4)
116
+ # act head sees the pooled sequence + detached summary of its own answer distribution
117
+ p = torch.softmax(logits.detach(), -1)
118
+ k = marker_mask.sum(-1).clamp(min=2).float()
119
+ ent = -(p * torch.log(p.clamp_min(1e-9))).sum(-1) / torch.log(k)
120
+ top2 = p.topk(2, -1).values
121
+ feats = torch.stack([top2[:, 0], top2[:, 0] - top2[:, 1], ent, k / 255.0], -1)
122
+ pooled = h[:, 0].float()
123
+ act_logits = self.act_head(torch.cat([pooled, feats], -1))
124
+ return logits, act_logits
125
+
126
+
127
+ def build_model(cfg: Dict, encoder_dir: Optional[str] = None) -> DecisionModel:
128
+ from transformers import AutoConfig, AutoModel
129
+ if encoder_dir: # offline: architecture only, weights come from the saved state dict
130
+ ecfg = AutoConfig.from_pretrained(encoder_dir)
131
+ enc = AutoModel.from_config(ecfg, attn_implementation="sdpa")
132
+ else:
133
+ enc = AutoModel.from_pretrained(cfg["encoder"], attn_implementation="sdpa")
134
+ return DecisionModel(enc, cfg["head_layers"], len(cfg["act_costs"]) + 1)
135
+
136
+
137
+ # ----------------------------------------------------------------------------- rewards (strictly proper)
138
+ def proper_reward(q: torch.Tensor, target: torch.Tensor, qtype: torch.Tensor, mask: torch.Tensor,
139
+ w_sph: float = 0.5, w_rps: float = 1.0, log_floor: float = -9.21) -> torch.Tensor:
140
+ """q: [..., N, K] reported distributions, target: [N, K] (one-hot or soft) -> reward [..., N].
141
+
142
+ log score + spherical score for all types, + ranked probability score for ordinal (score) questions.
143
+ All three are strictly proper, so the only way to maximize reward is to report honest probabilities.
144
+ """
145
+ q = q * mask
146
+ logq = torch.log(q.clamp_min(1e-12)).clamp_min(log_floor)
147
+ log_score = (target * logq).sum(-1)
148
+ sph = (target * q).sum(-1) / q.norm(dim=-1).clamp_min(1e-9)
149
+ r = log_score + w_sph * sph
150
+ is_score = (qtype == QTYPES["score"]).float()
151
+ if is_score.any():
152
+ k = mask.sum(-1).clamp(min=2).float()
153
+ cdf_q = torch.cumsum(q, -1)
154
+ cdf_t = torch.cumsum(target, -1)
155
+ rps = (((cdf_q - cdf_t) ** 2) * mask).sum(-1) / (k - 1)
156
+ r = r - w_rps * rps * is_score
157
+ return r
158
+
159
+
160
+ # ----------------------------------------------------------------------------- metrics (numpy, no sklearn)
161
+ def ece_score(conf: np.ndarray, correct: np.ndarray, bins: int = 15) -> float:
162
+ if len(conf) == 0:
163
+ return float("nan")
164
+ edges = np.linspace(0, 1, bins + 1)
165
+ e = 0.0
166
+ for lo, hi in zip(edges[:-1], edges[1:]):
167
+ sel = (conf > lo) & (conf <= hi)
168
+ if sel.any():
169
+ e += sel.mean() * abs(conf[sel].mean() - correct[sel].mean())
170
+ return float(e)
171
+
172
+
173
+ def auroc(scores: np.ndarray, labels: np.ndarray) -> float:
174
+ pos, neg = labels == 1, labels == 0
175
+ if pos.sum() == 0 or neg.sum() == 0:
176
+ return float("nan")
177
+ order = np.argsort(scores)
178
+ ranks = np.empty(len(scores))
179
+ ranks[order] = np.arange(1, len(scores) + 1)
180
+ # average ties
181
+ s_sorted = scores[order]
182
+ i = 0
183
+ while i < len(s_sorted):
184
+ j = i
185
+ while j + 1 < len(s_sorted) and s_sorted[j + 1] == s_sorted[i]:
186
+ j += 1
187
+ if j > i:
188
+ ranks[order[i:j + 1]] = (i + j + 2) / 2.0
189
+ i = j + 1
190
+ return float((ranks[pos].sum() - pos.sum() * (pos.sum() + 1) / 2) / (pos.sum() * neg.sum()))
191
+
192
+
193
+ def spearman(a: np.ndarray, b: np.ndarray) -> float:
194
+ if len(a) < 3:
195
+ return float("nan")
196
+ ra = np.argsort(np.argsort(a)).astype(float)
197
+ rb = np.argsort(np.argsort(b)).astype(float)
198
+ if ra.std() == 0 or rb.std() == 0:
199
+ return float("nan")
200
+ return float(np.corrcoef(ra, rb)[0, 1])
201
+
202
+
203
+ def aurc(conf: np.ndarray, correct: np.ndarray) -> float:
204
+ """Area under the risk-coverage curve (lower is better)."""
205
+ if len(conf) == 0:
206
+ return float("nan")
207
+ order = np.argsort(-conf)
208
+ err = 1 - correct[order]
209
+ return float((np.cumsum(err) / np.arange(1, len(err) + 1)).mean())
210
+
211
+
212
+ def confidence_from_probs(p: np.ndarray, k: int) -> float:
213
+ """Jev-style confidence: 1 - normalized entropy of the answer distribution."""
214
+ if k < 2:
215
+ return 1.0
216
+ p = p[:k]
217
+ ent = -(p * np.log(np.clip(p, 1e-12, 1))).sum()
218
+ return float(1 - ent / math.log(k))
219
+
220
+
221
+ def seed_all(seed: int):
222
+ random.seed(seed)
223
+ np.random.seed(seed)
224
+ torch.manual_seed(seed)
225
+
226
+
227
+ # ----------------------------------------------------------------------------- record -> model inputs
228
+ def episode_prefix_lengths(n_turns: int, max_prefixes: int) -> List[int]:
229
+ if n_turns <= max_prefixes:
230
+ return list(range(1, n_turns + 1))
231
+ return sorted(set(int(round(x)) for x in np.linspace(1, n_turns, max_prefixes)))
232
+
233
+
234
+ def encode_record(rec: Dict, tok, cfg: Dict, rng: Optional[random.Random], train: bool) -> List[Dict]:
235
+ """One stored record -> list of model sequences (one per question, or one per conversation prefix)."""
236
+ items = []
237
+ if rec.get("kind") == "episode":
238
+ ep, q = rec["ep"], rec["qs"][0]
239
+ lens = episode_prefix_lengths(len(ep["turns"]), cfg["max_prefixes"])
240
+ for step, t in enumerate(lens):
241
+ state = dict(ep["ctx"], conversation=ep["turns"][:t])
242
+ ids, markers = build_sequence(tok, state, q, cfg["max_len"], cfg["head_max_len"], truncate_left=True)
243
+ if len(markers) != 2:
244
+ continue
245
+ items.append({"ids": ids, "markers": markers, "qtype": QTYPES["noul"], "target": [1.0 - ep["y"], float(ep["y"])],
246
+ "label": int(ep["y"]), "episode": 1, "ep_step": step, "ep_len": len(lens), "src": rec.get("src", ""),
247
+ "prefix_frac": t / float(len(ep["turns"]))})
248
+ return items
249
+ for qi, q in enumerate(rec["qs"]):
250
+ k = len(render_options(q))
251
+ target = list(q["soft"]) if q.get("soft") else [1.0 if i == q["y"] else 0.0 for i in range(k)]
252
+ order = list(range(k))
253
+ if train and rng is not None and q["t"] != "score":
254
+ rng.shuffle(order)
255
+ ids, markers = build_sequence(tok, rec["state"], q, cfg["max_len"], cfg["head_max_len"], option_order=order)
256
+ if len(markers) != k:
257
+ continue # options did not fit; skip rather than train on a truncated answer space
258
+ target = [target[i] for i in order]
259
+ label = order.index(q["y"]) if q.get("y") is not None else -1
260
+ items.append({"ids": ids, "markers": markers, "qtype": QTYPES[q["t"]], "target": target, "label": label,
261
+ "episode": 0, "ep_step": 0, "ep_len": 1, "src": rec.get("src", ""), "q_index": qi, "order": order})
262
+ return items
263
+
264
+
265
+ def collate_items(batch, pad_id: int):
266
+ items = [it for group in batch for it in group]
267
+ if not items:
268
+ return None
269
+ n, L = len(items), max(len(it["ids"]) for it in items)
270
+ kmax = max(len(it["markers"]) for it in items)
271
+ ids = torch.full((n, L), pad_id, dtype=torch.long)
272
+ att = torch.zeros((n, L), dtype=torch.long)
273
+ mpos = torch.zeros((n, kmax), dtype=torch.long)
274
+ mmask = torch.zeros((n, kmax), dtype=torch.bool)
275
+ target = torch.zeros((n, kmax), dtype=torch.float32)
276
+ ep_group = torch.full((n,), -1, dtype=torch.long)
277
+ group_of = {}
278
+ for i, it in enumerate(items):
279
+ ids[i, :len(it["ids"])] = torch.tensor(it["ids"])
280
+ att[i, :len(it["ids"])] = 1
281
+ k = len(it["markers"])
282
+ mpos[i, :k] = torch.tensor(it["markers"])
283
+ mmask[i, :k] = True
284
+ target[i, :k] = torch.tensor(it["target"], dtype=torch.float32)
285
+ # episodes: all prefixes of the same record share a group id (used for TD(lambda) targets)
286
+ for i, it in enumerate(items):
287
+ if it["episode"]:
288
+ ep_group[i] = group_of.setdefault(it.get("rec_uid", -1 - i), len(group_of))
289
+ return {"input_ids": ids, "attention_mask": att, "marker_pos": mpos, "marker_mask": mmask, "target": target,
290
+ "qtype": torch.tensor([it["qtype"] for it in items]), "label": torch.tensor([it["label"] for it in items]),
291
+ "episode": torch.tensor([it["episode"] for it in items], dtype=torch.bool), "ep_group": ep_group,
292
+ "ep_step": torch.tensor([it["ep_step"] for it in items]), "meta": [{k: it[k] for k in it if k not in ("ids", "markers", "target")} for it in items],
293
+ "n_tokens": int(att.sum())}
294
+
295
+
296
+ def pack_groups(groups: List[List[Dict]], max_tokens: int, max_seqs: int) -> List[List[List[Dict]]]:
297
+ """Split one sampled batch into sub-batches using the *real* tokenized lengths, so padded tokens never exceed
298
+ max_tokens (the index only stores estimates). A record's items stay together (TD targets need all prefixes)."""
299
+ groups = sorted([g for g in groups if g], key=lambda g: max(len(it["ids"]) for it in g))
300
+ subs, cur, cur_max, cur_n = [], [], 0, 0
301
+ for g in groups:
302
+ g_max, g_n = max(len(it["ids"]) for it in g), len(g)
303
+ if g_max * g_n > max_tokens: # one record bigger than the budget (only if max_tokens < max_len * n_items)
304
+ step = max(1, max_tokens // g_max)
305
+ for s in range(0, g_n, step):
306
+ subs.append([g[s:s + step]])
307
+ continue
308
+ new_max, new_n = max(cur_max, g_max), cur_n + g_n
309
+ if cur and (new_max * new_n > max_tokens or new_n > max_seqs):
310
+ subs.append(cur)
311
+ cur, new_max, new_n = [], g_max, g_n
312
+ cur.append(g)
313
+ cur_max, cur_n = new_max, new_n
314
+ if cur:
315
+ subs.append(cur)
316
+ return subs
317
+
318
+
319
+ def td_lambda_targets(p_true: torch.Tensor, batch: Dict, lam: float) -> torch.Tensor:
320
+ """TD(lambda) soft targets for conversation prefixes: G_last = outcome, G_t = (1-lam) V_{t+1} + lam G_{t+1}."""
321
+ target = batch["target"].clone()
322
+ groups = batch["ep_group"]
323
+ for g in torch.unique(groups[groups >= 0]).tolist():
324
+ idx = (groups == g).nonzero(as_tuple=True)[0]
325
+ idx = idx[torch.argsort(batch["ep_step"][idx])]
326
+ y = batch["target"][idx[-1], 1]
327
+ G = y
328
+ for j in range(len(idx) - 1, -1, -1):
329
+ if j < len(idx) - 1:
330
+ G = (1 - lam) * p_true[idx[j + 1]] + lam * G
331
+ target[idx[j], 0], target[idx[j], 1] = 1 - G, G
332
+ return target
333
+
334
+
335
+ def make_token_batches(lengths: np.ndarray, nseq: np.ndarray, max_tokens: int, max_seqs: int, rng: np.random.RandomState,
336
+ chunk: int = 4096) -> List[List[int]]:
337
+ """Length-bucketed batches of record indices under a padded-token budget."""
338
+ order = rng.permutation(len(lengths))
339
+ batches = []
340
+ for s in range(0, len(order), chunk):
341
+ part = order[s:s + chunk]
342
+ part = part[np.argsort(lengths[part])]
343
+ cur, cur_max, cur_n = [], 0, 0
344
+ for i in part:
345
+ ln, ns = int(lengths[i]), int(nseq[i])
346
+ new_max, new_n = max(cur_max, ln), cur_n + ns
347
+ if cur and (new_max * new_n > max_tokens or new_n > max_seqs):
348
+ batches.append(cur)
349
+ cur, new_max, new_n = [], ln, ns
350
+ cur.append(int(i))
351
+ cur_max, cur_n = new_max, new_n
352
+ if cur:
353
+ batches.append(cur)
354
+ rng.shuffle(batches)
355
+ return batches
356
+
357
+
358
+ def temp_bucket(qtype: int, k: int) -> str:
359
+ """Key for per-cardinality temperature fitting: a 2-option noul and a 20-option choice need different scaling."""
360
+ size = "2" if k <= 2 else "3-5" if k <= 5 else "6-10" if k <= 10 else "11+"
361
+ return "%s:%s" % (QTYPE_NAMES[int(qtype)], size)
362
+
363
+
364
+ def amp_dtype(name: Optional[str]) -> torch.dtype:
365
+ """'bf16' on GPUs that support it (Ampere+, e.g. RTX 6000 Pro); 'fp16' on T4."""
366
+ return torch.bfloat16 if name == "bf16" else torch.float16
367
+
368
+
369
+ @torch.no_grad()
370
+ def predict_items(model, items: List[Dict], pad_id: int = 0, device=None, max_tokens: int = 16384, use_amp: bool = True,
371
+ dtype: torch.dtype = torch.float16, max_seqs: int = 256, progress: str = ""):
372
+ """Run the model over pre-encoded items; returns list of dicts with probs/logits (uncalibrated) and act probs."""
373
+ import sys
374
+ import time as _time
375
+ model.eval()
376
+ out = []
377
+ t0, done_tok = _time.time(), 0
378
+ order = sorted(range(len(items)), key=lambda i: len(items[i]["ids"]))
379
+ i = 0
380
+ while i < len(order):
381
+ j, L = i, 0
382
+ while j < len(order) and j - i < max_seqs and max(L, len(items[order[j]]["ids"])) * (j - i + 1) <= max_tokens:
383
+ L = max(L, len(items[order[j]]["ids"]))
384
+ j += 1
385
+ j = max(j, i + 1)
386
+ sel = [items[order[t]] for t in range(i, j)]
387
+ b = collate_items([sel], pad_id)
388
+ with torch.autocast(device_type=device.type, dtype=dtype, enabled=use_amp and device.type == "cuda"):
389
+ logits, act = model(b["input_ids"].to(device), b["attention_mask"].to(device), b["marker_pos"].to(device),
390
+ b["marker_mask"].to(device), b["qtype"].to(device))
391
+ logits, act = logits.float().cpu(), torch.softmax(act.float(), -1).cpu()
392
+ done_tok += int(b["attention_mask"].sum())
393
+ if progress and (j % max(1, len(order) // 2000) == 0 or j >= len(order)):
394
+ el = _time.time() - t0
395
+ eta = el * (len(order) - j) / max(1, j)
396
+ sys.stdout.write("\r [%s] %d/%d sequences | %.1fk tok/s | ETA %dm%02ds " %
397
+ (progress, j, len(order), done_tok / max(el, 1e-9) / 1000, int(eta // 60), int(eta % 60)))
398
+ sys.stdout.flush()
399
+ for r, it in enumerate(sel):
400
+ k = len(it["markers"])
401
+ out.append((order[i + r], {"logits": logits[r, :k].detach().numpy(), "act": act[r].detach().numpy()}))
402
+ i = j
403
+ if progress:
404
+ print("\r [%s] %d sequences in %.0fs (%.1fk tok/s)%s" % (progress, len(order), _time.time() - t0,
405
+ done_tok / max(_time.time() - t0, 1e-9) / 1000, " " * 20))
406
+ out.sort(key=lambda x: x[0])
407
+ model.train()
408
+ return [o for _, o in out]
tokenizer/tokenizer.json ADDED
The diff for this file is too large to render. See raw diff
 
tokenizer/tokenizer_config.json ADDED
@@ -0,0 +1,14 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "clean_up_tokenization_spaces": true,
3
+ "cls_token": "[CLS]",
4
+ "mask_token": "[MASK]",
5
+ "model_input_names": [
6
+ "input_ids",
7
+ "attention_mask"
8
+ ],
9
+ "model_max_length": 8192,
10
+ "pad_token": "[PAD]",
11
+ "sep_token": "[SEP]",
12
+ "tokenizer_class": "PreTrainedTokenizerFast",
13
+ "unk_token": "[UNK]"
14
+ }