"""Fit per-question-type temperatures on held-out cases (NLL), write them into rl_agent_config.json. python finetune/calibrate.py out/pages.jsonl out/eval_cases.jsonl """ import json, os, sys import numpy as np, torch sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))); sys.path.insert(0, "/home/ckl/projects/S/laya-upstream") from common_ft import build_request, gold_for import laya from laya.common import QTYPES, build_sequence, collate_items def main(): pages = [json.loads(l) for l in open(sys.argv[1])]; cases = [c for c in (json.loads(l) for l in open(sys.argv[2])) if not c.get("noul_only")]; ck = sys.argv[3] agent = laya.load(ck); agent.cfg["max_len"], agent.cfg["head_max_len"] = int(os.environ.get("LAYA_MAXLEN", "1024")), int(os.environ.get("LAYA_HEAD", agent.cfg.get("head_max_len_train", 512))); agent.accelerate() Z, T, K = [], [], [] for c in cases[::2]: state, questions, targets, controls = build_request(c.get("page_obj") or pages[c["page"]], c["goal"], c.get("history", [])) gop, gidx = gold_for(c, targets, controls) golds = {"operation": gop} if gidx is not None: golds[gop.lower() + "_target"] = gidx items, keys = [], [] for qid, gold in golds.items(): q = questions[qid]; qq = {"t": "choice", "ins": json.dumps(q["instructions"]), "crit": q["criteria"]} seq, markers = build_sequence(agent.tok, state, qq, 1024, agent.cfg['head_max_len']) items.append({"ids": seq, "markers": markers, "qtype": QTYPES["choice"]}); keys.append((list(q["criteria"]), gold)) b = collate_items([items], agent.tok.pad_token_id) with torch.no_grad(), torch.autocast("cuda", dtype=agent.dtype): logits, _ = agent.model(b["input_ids"].cuda(), b["attention_mask"].cuda(), b["marker_pos"].cuda(), b["marker_mask"].cuda(), b["qtype"].cuda()) for r, (ks, gold) in enumerate(keys): z = logits[r, :len(ks)].float().cpu(); Z.append(z); T.append(ks.index(gold)) kmax = max(len(z) for z in Z); M = torch.full((len(Z), kmax), -1e4) for i, z in enumerate(Z): M[i, :len(z)] = z y = torch.tensor(T) def nll(t): return torch.nn.functional.cross_entropy(M / t, y).item() ts = np.exp(np.linspace(np.log(0.2), np.log(10), 200)); best = min(ts, key=nll) acc = (M.argmax(-1) == y).float().mean().item() conf0 = torch.softmax(M, -1).max(-1).values.mean().item(); conf1 = torch.softmax(M / best, -1).max(-1).values.mean().item() print(f"n={len(Z)} acc={acc:.3f} T=1: nll {nll(1.0):.3f} mean conf {conf0:.3f} | T={best:.2f}: nll {nll(best):.3f} mean conf {conf1:.3f}") cfgp = os.path.join(ck, "rl_agent_config.json"); cfg = json.load(open(cfgp)); cfg["temperature"] = [float(best), 1.0, 1.0] json.dump(cfg, open(cfgp, "w"), indent=2); print("wrote temperature", best, "->", cfgp) if __name__ == "__main__": main()