rl_common.py
18.7 KB · 409 lines · python Raw
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]
409