rl_common.py
| 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 | |