laya-code: fine-tuned Laya code-relevance re-ranker
Browse files- encoder/config.json +84 -0
- model.safetensors +3 -0
- rl_agent_api.py +78 -0
- rl_agent_config.json +60 -0
- rl_common.py +408 -0
- tokenizer/tokenizer.json +0 -0
- tokenizer/tokenizer_config.json +14 -0
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 |
+
}
|