rl_agent_api.py
| 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}} |
| 79 | |