Spaces:
Running on Zero
Running on Zero
File size: 43,085 Bytes
a7c4f9e ff6e9ba a7c4f9e 7bbb06a a7c4f9e 7bbb06a a7c4f9e f289609 a7c4f9e f289609 a7c4f9e 7bbb06a a7c4f9e 7bbb06a a7c4f9e dfbe23d a7c4f9e 6a9e131 a7c4f9e 6a9e131 a7c4f9e dfbe23d a7c4f9e dfbe23d a7c4f9e 7bbb06a a7c4f9e 6a9e131 a7c4f9e 7bbb06a a7c4f9e dfbe23d | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 407 408 409 410 411 412 413 414 415 416 417 418 419 420 421 422 423 424 425 426 427 428 429 430 431 432 433 434 435 436 437 438 439 440 441 442 443 444 445 446 447 448 449 450 451 452 453 454 455 456 457 458 459 460 461 462 463 464 465 466 467 468 469 470 471 472 473 474 475 476 477 478 479 480 481 482 483 484 485 486 487 488 489 490 491 492 493 494 495 496 497 498 499 500 501 502 503 504 505 506 507 508 509 510 511 512 513 514 515 516 517 518 519 520 521 522 523 524 525 526 527 528 529 530 531 532 533 534 535 536 537 538 539 540 541 542 543 544 545 546 547 548 549 550 551 552 553 554 555 556 557 558 559 560 561 562 563 564 565 566 567 568 569 570 571 572 573 574 575 576 577 578 579 580 581 582 583 584 585 586 587 588 589 590 591 592 593 594 595 596 597 598 599 600 601 602 603 604 605 606 607 608 609 610 611 612 613 614 615 616 617 618 619 620 621 622 623 624 625 626 627 628 629 630 631 632 633 634 635 636 637 638 639 640 641 642 643 644 645 646 647 648 649 650 651 652 653 654 655 656 657 658 659 660 661 662 663 664 665 666 667 668 669 670 671 672 673 674 675 676 677 678 679 680 681 682 683 684 685 686 687 688 689 690 691 692 693 694 695 696 697 698 699 700 701 702 703 704 705 706 707 708 709 710 711 712 713 714 715 716 717 718 719 720 721 722 723 724 725 726 727 728 729 730 731 732 733 734 735 736 737 738 739 740 741 742 743 744 745 746 747 748 749 750 751 752 753 754 755 756 757 758 759 760 761 762 763 764 765 766 767 768 769 770 771 772 773 774 775 776 777 778 779 780 781 782 783 784 785 786 787 788 789 790 791 792 793 794 795 796 797 798 799 800 801 802 803 804 805 806 807 808 809 810 811 812 813 814 815 816 817 818 819 820 821 822 823 824 825 826 827 828 829 830 831 832 833 834 835 836 837 838 839 840 841 842 843 844 845 846 847 848 849 850 851 852 853 854 855 856 857 858 859 860 861 862 863 864 865 866 867 868 869 870 871 872 873 874 875 876 877 878 879 880 881 882 883 884 885 886 887 888 889 890 891 892 893 894 895 896 897 898 899 900 901 902 903 904 905 906 907 908 909 910 911 912 913 914 915 916 917 918 919 920 921 922 923 924 925 926 927 928 929 930 931 932 933 934 935 936 937 938 939 940 941 942 943 944 945 946 947 948 949 950 951 952 953 954 955 956 957 958 959 960 961 962 963 964 965 966 967 968 969 970 971 972 973 974 975 976 977 978 979 980 981 982 983 984 985 986 987 988 989 990 991 992 993 994 995 996 997 998 999 1000 1001 1002 1003 1004 1005 1006 1007 1008 | """Centauri Cognitive Simulator.
Play a psychology experiment trial by trial while small foundation models of human
cognition (Centauri / Qwentaur LoRA adapters over Qwen3-Base, Oh & Gobet 2026)
predict what a *human* would do next.
Everything the models see is plain Psych-101-formatted text, exactly as in training:
choices are wrapped in `<<...>>` and the model's choice distribution is read off the
next-token logits right after `You press <<`.
"""
import os
os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
import spaces # noqa: E402 (must precede torch / transformers)
import html
import json
import math
import random
import re
import time
from typing import Any, Dict, List, Optional, Tuple
import gradio as gr
import torch
from peft import PeftModel
from transformers import AutoModelForCausalLM, AutoTokenizer
# --------------------------------------------------------------------------------------
# Models
# --------------------------------------------------------------------------------------
PRIMARY = "socius/Qwentaur-8B-LoRA-r16"
SMALL = "socius/Qwentaur-0.6B-LoRA-r16"
BASE_8B = "unsloth/Qwen3-8B-Base"
BASE_06B = "unsloth/Qwen3-0.6B-Base"
# key -> display metadata. "big"/"small" are the fine-tuned adapters, "raw" is the
# 8B base model with the adapter switched off (free — same weights in VRAM).
PREDICTORS = [
("big", "Qwentaur-8B", "#4f46e5", "LoRA r=16 on Qwen3-8B-Base"),
("small", "Qwentaur-0.6B", "#0d9488", "LoRA r=16 on Qwen3-0.6B-Base"),
("raw", "Qwen3-8B-Base", "#a1a1aa", "no fine-tuning (adapter off)"),
]
PRED_LABEL = {k: label for k, label, _, _ in PREDICTORS}
PRED_COLOR = {k: color for k, _, color, _ in PREDICTORS}
def _load(base_id: str, adapter_id: str):
tok = AutoTokenizer.from_pretrained(adapter_id)
model = AutoModelForCausalLM.from_pretrained(base_id, dtype=torch.bfloat16)
# torch_device="cpu" is required on ZeroGPU: peft otherwise asks safetensors to
# materialise the adapter straight on CUDA at import time, when no GPU is attached.
try:
model = PeftModel.from_pretrained(model, adapter_id, torch_device="cpu")
except TypeError:
model = PeftModel.from_pretrained(model, adapter_id)
model.eval()
model = model.to("cuda")
return tok, model
print("Loading Qwentaur-8B (Qwen3-8B-Base + LoRA r=16) ...", flush=True)
TOK_BIG, MODEL_BIG = _load(BASE_8B, PRIMARY)
print("Loading Qwentaur-0.6B (Qwen3-0.6B-Base + LoRA r=16) ...", flush=True)
TOK_SMALL, MODEL_SMALL = _load(BASE_06B, SMALL)
print("Models ready.", flush=True)
# Single-token uppercase letters, shared by both (identical Qwen3) tokenizers.
LETTER_POOL = [
c
for c in "ABCDEFGHIJKLMNOPQRSTUVWXYZ"
if len(TOK_BIG(c, add_special_tokens=False).input_ids) == 1
and len(TOK_SMALL(c, add_special_tokens=False).input_ids) == 1
]
print(f"{len(LETTER_POOL)} single-token option letters available.", flush=True)
MAX_SCORE_TOKENS = 3072
@torch.no_grad()
def _next_letter_probs(model, tok, prefix: str, letters: List[str]) -> Dict[str, float]:
"""Renormalised P(letter) read off the next-token logits after `... You press <<`."""
ids = tok(prefix, return_tensors="pt", add_special_tokens=False).input_ids.to("cuda")
if ids.shape[1] > MAX_SCORE_TOKENS:
ids = ids[:, -MAX_SCORE_TOKENS:]
logits = model(ids).logits[0, -1].float()
logp = torch.log_softmax(logits, dim=-1)
raw = {}
for letter in letters:
tid = tok(letter, add_special_tokens=False).input_ids[0]
raw[letter] = float(logp[tid])
mx = max(raw.values())
ex = {k: math.exp(v - mx) for k, v in raw.items()}
z = sum(ex.values())
return {k: v / z for k, v in ex.items()}
def _predict_all(prefix: str, letters: List[str]) -> Dict[str, Dict[str, float]]:
"""Choice distributions from all three predictors for one upcoming choice."""
out = {"big": _next_letter_probs(MODEL_BIG, TOK_BIG, prefix, letters)}
with MODEL_BIG.disable_adapter():
out["raw"] = _next_letter_probs(MODEL_BIG, TOK_BIG, prefix, letters)
out["small"] = _next_letter_probs(MODEL_SMALL, TOK_SMALL, prefix, letters)
return out
# --------------------------------------------------------------------------------------
# Psych-101 task environments
#
# Wording is copied verbatim from the Psych-101 transcripts (Binz et al.) so the models
# see text drawn from exactly the distribution they were fine-tuned on. Only the option
# letters are re-randomised per session, as in the original dataset.
# --------------------------------------------------------------------------------------
TASK_ORDER = ["bandit", "igt", "itc", "cue"]
TASK_LABELS = {
"bandit": "🎰 Horizon task — explore or exploit (Wilson et al., 2014)",
"igt": "🃏 Iowa Gambling Task — learning under risk (Steingroever et al., 2015)",
"itc": "⏳ Intertemporal choice — patience (Ruggeri et al., 2022)",
"cue": "🔍 Multi-attribute inference — cue integration (Hilbig & Moshagen, 2014)",
}
LABEL_TO_TASK = {v: k for k, v in TASK_LABELS.items()}
IGT_DECKS = {
# (win, [loss schedule over a block of 10 cards]) — classic Bechara payoffs
"bad_freq": (100.0, [0, 0, 150, 0, 300, 0, 200, 0, 250, 350]),
"bad_rare": (100.0, [0, 0, 0, 0, 0, 0, 0, 0, 0, 1250]),
"good_freq": (50.0, [0, 0, 50, 0, 50, 0, 50, 0, 50, 50]),
"good_rare": (50.0, [0, 0, 0, 0, 0, 0, 0, 0, 0, 250]),
}
# The ten trials of the Ruggeri et al. intertemporal-choice block, verbatim.
ITC_TRIALS = [
("receiving 500$ immediately", "receiving 550$ in one year"),
("receiving 500$ immediately", "receiving 600$ in one year"),
("receiving 500$ immediately", "receiving 750$ in one year"),
("paying 500$ immediately", "paying 550$ in one year"),
("paying 500$ immediately", "paying 510$ in one year"),
("paying 500$ immediately", "paying 505$ in one year"),
("receiving 5000$ immediately", "receiving 5500$ in one year"),
("receiving 5000$ immediately", "receiving 6000$ in one year"),
("receiving 5000$ immediately", "receiving 7500$ in one year"),
("receiving 500$ in one year", "receiving 750$ in two years"),
]
CUE_VALIDITIES = [0.9, 0.8, 0.7, 0.6]
def new_session(task: str, seed: int, blind: bool) -> Dict[str, Any]:
rng = random.Random(seed)
st: Dict[str, Any] = {
"task": task,
"seed": int(seed),
"blind": bool(blind),
"trial": 0,
"log": [],
"pred": {},
"done": False,
}
if task == "bandit":
letters = rng.sample(LETTER_POOL, 2)
st["letters"] = letters
low = rng.randint(30, 58)
delta = rng.choice([4, 8, 12, 20, 30])
means = [low, low + delta]
rng.shuffle(means)
# Pre-generate the full (counterfactual) reward table for 4 instructed + 6 free trials.
rewards = [
[max(1, min(99, int(round(rng.gauss(m, 8))))) for _ in range(10)] for m in means
]
forced = ([0] * 3 + [1]) if rng.random() < 0.5 else [0, 0, 1, 1]
rng.shuffle(forced)
st["plan"] = {"rewards": rewards, "forced": forced, "means": means}
st["n_trials"] = 6
st["text"] = (
f"You are participating in multiple games involving two slot machines, labeled {letters[0]} and {letters[1]}.\n"
"The two slot machines are different across different games.\n"
"Each time you choose a slot machine, you get some points.\n"
"You choose a slot machine by pressing the corresponding key.\n"
"Each slot machine tends to pay out about the same amount of points on average.\n"
"Your goal is to choose the slot machines that will give you the most points across the experiment.\n"
"The first 4 trials in each game are instructed trials where you will be told which slot machine to choose.\n"
"After these instructed trials, you will have the freedom to choose for either 1 or 6 trials.\n"
"\nGame 1. There are 10 trials in this game.\n"
)
counts = [0, 0]
observed: List[List[int]] = [[], []]
for arm in forced:
r = rewards[arm][counts[arm]]
counts[arm] += 1
observed[arm].append(r)
st["text"] += f"You are instructed to press {letters[arm]} and get {r} points.\n"
st["counts"] = counts
st["observed"] = observed
st["total"] = sum(sum(o) for o in observed)
elif task == "igt":
letters = rng.sample(LETTER_POOL, 4)
st["letters"] = letters
kinds = list(IGT_DECKS.keys())
rng.shuffle(kinds)
schedule = {}
for letter, kind in zip(letters, kinds):
win, losses = IGT_DECKS[kind]
cards = []
for _ in range(4): # 40 cards per deck is plenty for 30 trials
block = list(losses)
rng.shuffle(block)
cards += [(win, float(x)) for x in block]
schedule[letter] = cards
st["plan"] = {"schedule": schedule, "kinds": dict(zip(letters, kinds))}
st["n_trials"] = 30
st["counts"] = {letter: 0 for letter in letters}
st["net"] = {letter: 0.0 for letter in letters}
st["balance"] = 2000.0
st["text"] = (
f"You see in front of you four decks of cards labeled {letters[0]}, {letters[1]}, {letters[2]}, and {letters[3]}.\n"
"You get a loan of 2000$ of play money.\n"
"You have to select one card at a time, from any of the four decks, for 100 trials.\n"
"You select a card from a deck by pressing the corresponding key.\n"
"After turning a card, you win some money, the amount varies with the deck.\n"
"You sometimes also have to pay a penalty, which also varies with the deck.\n"
"Your goal is to maximize profit on the loan of the play money.\n\n"
)
elif task == "itc":
letters = rng.sample(LETTER_POOL, 2)
st["letters"] = letters
st["plan"] = {"trials": ITC_TRIALS}
st["n_trials"] = len(ITC_TRIALS)
st["patient"] = 0
st["text"] = (
f"In the following you will be presented with multiple choices between two options {letters[0]} and {letters[1]}.\n"
"Please name which option you would prefer by pressing the corresponding key.\n\n"
)
elif task == "cue":
letters = rng.sample(LETTER_POOL, 2)
st["letters"] = letters
stimuli = []
while len(stimuli) < 20:
a = [rng.randint(0, 1) for _ in range(4)]
b = [rng.randint(0, 1) for _ in range(4)]
if a != b:
stimuli.append([a, b])
st["plan"] = {"stimuli": stimuli}
st["n_trials"] = 20
st["optimal"] = 0
st["text"] = (
f"You are repeatedly presented with two options, labeled {letters[0]} and {letters[1]}.\n"
"Each option represents a fictitious product and you have to infer which product is superior in terms of quality.\n"
"You select a product by pressing the corresponding key.\n"
"For each decision, you are provided with four expert ratings (with 1 representing a positive and 0 representing a negative rating).\n"
"The four experts differ in their validity.\n"
"The ratings of experts are given in descending order of their validity (having validities of 90%, 80%, 70%, and 60%).\n\n"
)
else:
raise ValueError(task)
return st
def pre_text(st: Dict[str, Any]) -> str:
"""Text of the current trial up to (not including) ` <<`."""
task, letters, t = st["task"], st["letters"], st["trial"]
if task in ("bandit", "igt"):
return "You press"
if task == "itc":
a, b = st["plan"]["trials"][t]
return (
f"You have the choice between {a} (press {letters[0]}) or "
f"{b} (press {letters[1]}). You press"
)
if task == "cue":
a, b = st["plan"]["stimuli"][t]
fmt = lambda v: "[" + " ".join(str(x) for x in v) + "]" # noqa: E731
return (
f"Product {letters[0]} ratings: {fmt(a)}. "
f"Product {letters[1]} ratings: {fmt(b)}. You press"
)
raise ValueError(task)
def score_prefix(st: Dict[str, Any]) -> str:
return st["text"] + pre_text(st) + " <<"
def apply_choice(st: Dict[str, Any], letter: str) -> str:
"""Append the completed trial line to the transcript; returns a plain-text outcome."""
task, letters, t = st["task"], st["letters"], st["trial"]
line = pre_text(st) + f" <<{letter}>>"
if task == "bandit":
arm = letters.index(letter)
r = st["plan"]["rewards"][arm][st["counts"][arm]]
st["counts"][arm] += 1
st["observed"][arm].append(r)
st["total"] += r
line += f" and get {r} points.\n"
outcome = f"Machine {letter} paid {r} points."
elif task == "igt":
win, loss = st["plan"]["schedule"][letter][st["counts"][letter]]
st["counts"][letter] += 1
st["net"][letter] += win - loss
st["balance"] += win - loss
line += f". You win {win:.1f}$ and lose {loss:.1f}$.\n"
outcome = f"Deck {letter}: won ${win:.0f}" + (f", lost ${loss:.0f}" if loss else ", no penalty")
elif task == "itc":
if letter == letters[1]:
st["patient"] += 1
line += ".\n"
outcome = f"You chose option {letter}."
elif task == "cue":
a, b = st["plan"]["stimuli"][t]
wa = sum(v * math.log(p / (1 - p)) for v, p in zip(a, CUE_VALIDITIES))
wb = sum(v * math.log(p / (1 - p)) for v, p in zip(b, CUE_VALIDITIES))
best = letters[0] if wa >= wb else letters[1]
if letter == best:
st["optimal"] += 1
line += ".\n"
outcome = (
f"Optimal cue integration would pick {best}."
if letter != best
else f"{letter} is also what optimal cue integration picks."
)
else:
raise ValueError(task)
st["text"] += line
st["trial"] += 1
return outcome
# --------------------------------------------------------------------------------------
# Rendering
# --------------------------------------------------------------------------------------
CSS_BLOCK = """
<style>
.cs-wrap{font-family:var(--font,ui-sans-serif,system-ui);}
.cs-head{font-size:13px;opacity:.75;margin:0 0 8px 2px;letter-spacing:.02em}
.cs-cards{display:flex;gap:10px;flex-wrap:wrap}
.cs-card{flex:1 1 130px;min-width:120px;border:1px solid var(--border-color-primary,#e5e7eb);
border-radius:12px;padding:10px 12px;background:var(--background-fill-secondary,#fafafa)}
.cs-card.sel{border-color:#f59e0b;box-shadow:0 0 0 2px rgba(245,158,11,.25)}
.cs-key{font-size:26px;font-weight:700;line-height:1.1}
.cs-sub{font-size:12px;opacity:.72;margin-top:3px;line-height:1.45}
.cs-out{margin-top:10px;font-size:13px;padding:7px 10px;border-radius:8px;
background:rgba(245,158,11,.12);border:1px solid rgba(245,158,11,.35)}
.cs-mrow{margin:0 0 12px 0}
.cs-mname{font-size:12px;font-weight:600;margin-bottom:4px;display:flex;
justify-content:space-between;align-items:baseline;gap:8px}
.cs-mnote{font-weight:400;opacity:.6;font-size:11px}
.cs-bar{display:flex;align-items:center;gap:6px;margin:2px 0;font-size:11px}
.cs-bl{width:16px;text-align:right;opacity:.8;font-weight:600}
.cs-btrack{flex:1;height:14px;border-radius:7px;background:var(--background-fill-secondary,#eee);overflow:hidden}
.cs-bfill{height:100%;border-radius:7px}
.cs-bv{width:42px;font-variant-numeric:tabular-nums;opacity:.8}
.cs-pick{color:#f59e0b;font-weight:700}
.cs-tbl{width:100%;border-collapse:collapse;font-size:12px;margin-top:4px}
.cs-tbl th,.cs-tbl td{padding:4px 6px;text-align:right;border-bottom:1px solid var(--border-color-primary,#eee)}
.cs-tbl th:first-child,.cs-tbl td:first-child{text-align:left}
.cs-note{font-size:11.5px;opacity:.65;margin-top:8px;line-height:1.5}
.cs-tx{font-family:ui-monospace,SFMono-Regular,Menlo,monospace;font-size:12px;line-height:1.75;
white-space:pre-wrap;max-height:420px;overflow:auto;padding:10px;border-radius:8px;
border:1px solid var(--border-color-primary,#e5e7eb)}
.cs-ch{border-radius:4px;padding:1px 3px;font-weight:700}
</style>
"""
def _pips(vals: List[int]) -> str:
out = []
for v, p in zip(vals, CUE_VALIDITIES):
col = "#16a34a" if v else "#d4d4d8"
out.append(
f"<span title='validity {int(p*100)}%' style='display:inline-block;width:13px;height:13px;"
f"border-radius:3px;background:{col};margin-right:3px'></span>"
)
return "".join(out)
def render_board(st: Optional[Dict[str, Any]], outcome: str = "") -> str:
if st is None:
return (
CSS_BLOCK
+ "<div class='cs-wrap'><div class='cs-head'>No session yet</div>"
"<div class='cs-sub'>Pick an experiment above and press <b>Start experiment</b>.</div></div>"
)
task, letters, t, n = st["task"], st["letters"], st["trial"], st["n_trials"]
done = st["trial"] >= n
cards = []
if task == "bandit":
head = (
f"Game 1 · free choice {min(t + 1, n)} of {n} · <b>{st['total']} points</b> so far"
if not done
else f"Game over · <b>{st['total']} points</b>"
)
for i, letter in enumerate(letters):
obs = st["observed"][i]
seen = ", ".join(str(x) for x in obs) if obs else "never played"
avg = f"mean {sum(obs)/len(obs):.1f}" if obs else "unknown"
cards.append((letter, f"slot machine {letter}", [f"played {st['counts'][i]}×", seen, avg]))
elif task == "igt":
head = (
f"Card {min(t + 1, n)} of {n} · balance <b>${st['balance']:.0f}</b>"
if not done
else f"Finished · balance <b>${st['balance']:.0f}</b>"
)
for letter in letters:
cards.append(
(
letter,
f"deck {letter}",
[f"played {st['counts'][letter]}×", f"net ${st['net'][letter]:+.0f}"],
)
)
elif task == "itc":
head = f"Choice {min(t + 1, n)} of {n} · you took the later option {st['patient']}×"
if done:
head = f"Finished · you took the later option {st['patient']}/{n} times"
trial = st["plan"]["trials"][min(t, n - 1)]
for j, (letter, opt) in enumerate(zip(letters, trial)):
cards.append((letter, opt, ["the sooner option" if j == 0 else "the later option"]))
else: # cue
head = f"Decision {min(t + 1, n)} of {n} · matched optimal integration {st['optimal']}×"
if done:
head = f"Finished · matched optimal integration {st['optimal']}/{n} times"
stim = st["plan"]["stimuli"][min(t, n - 1)]
for letter, vals in zip(letters, stim):
cards.append((letter, f"product {letter}", [_pips(vals), "experts 90/80/70/60%"]))
last_pick = st["log"][-1]["choice"] if st["log"] else None
body = []
for letter, title, lines in cards:
sel = " sel" if letter == last_pick else ""
subs = "<br>".join(lines)
body.append(
f"<div class='cs-card{sel}'><div class='cs-key'>{html.escape(letter)}</div>"
f"<div class='cs-sub'>{title}</div><div class='cs-sub'>{subs}</div></div>"
)
out = f"<div class='cs-out'>{html.escape(outcome)}</div>" if outcome else ""
return (
CSS_BLOCK
+ f"<div class='cs-wrap'><div class='cs-head'>{head}</div>"
+ f"<div class='cs-cards'>{''.join(body)}</div>{out}</div>"
)
def _bars(probs: Dict[str, float], color: str, letters: List[str], pick: Optional[str]) -> str:
rows = []
top = max(probs, key=probs.get)
for letter in letters:
p = probs[letter]
mark = " <span class='cs-pick'>← you</span>" if letter == pick else ""
star = "★" if letter == top else ""
rows.append(
f"<div class='cs-bar'><div class='cs-bl'>{html.escape(letter)}</div>"
f"<div class='cs-btrack'><div class='cs-bfill' style='width:{p*100:.1f}%;background:{color}'></div></div>"
f"<div class='cs-bv'>{p*100:.1f}%</div><div>{star}{mark}</div></div>"
)
return "".join(rows)
def render_predictions(st: Optional[Dict[str, Any]]) -> str:
if st is None:
return (
CSS_BLOCK
+ "<div class='cs-wrap'><div class='cs-head'>Model predictions</div>"
"<div class='cs-sub'>Start a session to see what each model expects a human to do.</div></div>"
)
letters, n_opt = st["letters"], len(st["letters"])
done = st["trial"] >= st["n_trials"]
chance = 1.0 / n_opt
blocks = []
if st["pred"] and not done and not st["blind"]:
title = f"Predicted next choice · trial {st['trial'] + 1}"
for key, label, color, note in PREDICTORS:
probs = st["pred"][key]
blocks.append(
f"<div class='cs-mrow'><div class='cs-mname'><span>{label}</span>"
f"<span class='cs-mnote'>{note}</span></div>{_bars(probs, color, letters, None)}</div>"
)
elif st["log"]:
title = "Last trial · what each model expected"
rec = st["log"][-1]
for key, label, color, note in PREDICTORS:
probs = rec["probs"][key]
blocks.append(
f"<div class='cs-mrow'><div class='cs-mname'><span>{label}</span>"
f"<span class='cs-mnote'>{note}</span></div>"
f"{_bars(probs, color, letters, rec['choice'])}</div>"
)
else:
title = "Model predictions"
blocks.append("<div class='cs-sub'>Waiting for the first trial…</div>")
# scoreboard
table = ""
if st["log"]:
rows = []
for key, label, _color, _note in PREDICTORS:
nll = [-math.log(max(r["probs"][key][r["choice"]], 1e-9)) for r in st["log"]]
hits = sum(1 for r in st["log"] if max(r["probs"][key], key=r["probs"][key].get) == r["choice"])
mean_nll = sum(nll) / len(nll)
pr2 = 1.0 - mean_nll / math.log(n_opt)
rows.append(
f"<tr><td>{label}</td><td>{mean_nll:.3f}</td><td>{pr2*100:.0f}%</td>"
f"<td>{hits}/{len(st['log'])}</td></tr>"
)
rows.append(
f"<tr><td>random guessing</td><td>{math.log(n_opt):.3f}</td><td>0%</td>"
f"<td>{len(st['log'])/n_opt:.1f}/{len(st['log'])}</td></tr>"
)
table = (
"<table class='cs-tbl'><tr><th>predictor</th><th>loss</th><th>pseudo-R²</th>"
f"<th>top-1 hits</th></tr>{''.join(rows)}</table>"
"<div class='cs-note'>loss = mean negative log-likelihood of <i>your</i> choices "
f"(nats/choice; {math.log(n_opt):.3f} = chance with {n_opt} options). "
"This is the metric the paper reports on Psych-101.</div>"
)
return (
CSS_BLOCK
+ f"<div class='cs-wrap'><div class='cs-head'>{title}</div>{''.join(blocks)}{table}"
+ (f"<div class='cs-note'>Chance level is {chance*100:.0f}% per option.</div>" if not st["log"] else "")
+ "</div>"
)
def render_chart(st: Optional[Dict[str, Any]]) -> str:
if st is None or not st["log"]:
return ""
n = len(st["log"])
n_opt = len(st["letters"])
W, H, PL, PR, PT, PB = 660, 190, 34, 12, 14, 26
iw, ih = W - PL - PR, H - PT - PB
def x(i):
return PL + (iw * (i / max(1, n - 1)) if n > 1 else iw / 2)
def y(p):
return PT + ih * (1 - p)
parts = [f"<svg viewBox='0 0 {W} {H}' width='100%' style='max-width:700px'>"]
for gy in (0.0, 0.25, 0.5, 0.75, 1.0):
parts.append(
f"<line x1='{PL}' y1='{y(gy):.1f}' x2='{W-PR}' y2='{y(gy):.1f}' "
f"stroke='currentColor' stroke-opacity='.12'/>"
f"<text x='{PL-6}' y='{y(gy)+3:.1f}' font-size='9' text-anchor='end' "
f"fill='currentColor' opacity='.5'>{int(gy*100)}%</text>"
)
ch = 1.0 / n_opt
parts.append(
f"<line x1='{PL}' y1='{y(ch):.1f}' x2='{W-PR}' y2='{y(ch):.1f}' stroke='#ef4444' "
f"stroke-dasharray='4 3' stroke-opacity='.7'/>"
f"<text x='{W-PR}' y='{y(ch)-4:.1f}' font-size='9' text-anchor='end' fill='#ef4444' "
f"opacity='.85'>chance</text>"
)
for key, label, color, _note in PREDICTORS:
pts = " ".join(
f"{x(i):.1f},{y(r['probs'][key][r['choice']]):.1f}" for i, r in enumerate(st["log"])
)
parts.append(
f"<polyline points='{pts}' fill='none' stroke='{color}' stroke-width='2' "
f"stroke-linejoin='round'/>"
)
for i, r in enumerate(st["log"]):
parts.append(
f"<circle cx='{x(i):.1f}' cy='{y(r['probs'][key][r['choice']]):.1f}' r='2.6' fill='{color}'/>"
)
parts.append(
f"<text x='{PL}' y='{H-8}' font-size='9.5' fill='currentColor' opacity='.6'>trial 1</text>"
f"<text x='{W-PR}' y='{H-8}' font-size='9.5' text-anchor='end' fill='currentColor' "
f"opacity='.6'>trial {n}</text></svg>"
)
legend = " ".join(
f"<span style='font-size:11px;margin-right:10px'>"
f"<span style='display:inline-block;width:9px;height:9px;border-radius:2px;"
f"background:{c};margin-right:4px'></span>{l}</span>"
for _k, l, c, _n in PREDICTORS
)
return (
CSS_BLOCK
+ "<div class='cs-wrap'><div class='cs-head'>Probability each model assigned to the choice "
"you actually made</div>" + "".join(parts) + f"<div>{legend}</div></div>"
)
def button_updates(st: Optional[Dict[str, Any]]) -> List[Any]:
ups = []
for i in range(4):
if st is None or st["trial"] >= st["n_trials"] or i >= len(st["letters"]):
ups.append(gr.update(visible=False))
else:
ups.append(gr.update(visible=True, value=f"Press {st['letters'][i]}"))
return ups
# --------------------------------------------------------------------------------------
# Interactive-session handlers (GPU)
# --------------------------------------------------------------------------------------
@spaces.GPU(duration=15)
def start_session(
task_label: str, seed: float = 0, randomize_seed: bool = True, blind: bool = False
) -> Tuple[Any, ...]:
"""Start a fresh cognitive-experiment session and predict the first choice.
Args:
task_label: which Psych-101 experiment to run.
seed: RNG seed for the experiment (stimuli, payoffs, option letters).
randomize_seed: draw a new random seed instead of using `seed`.
blind: hide the model predictions until after each trial is over.
"""
task = LABEL_TO_TASK.get(task_label, "bandit")
seed = random.randint(0, 2**31 - 1) if randomize_seed else int(seed)
st = new_session(task, seed, blind)
st["pred"] = _predict_all(score_prefix(st), st["letters"])
return (
st,
render_board(st),
render_predictions(st),
render_chart(st),
st["text"] + pre_text(st) + " <<",
gr.update(value=seed),
*button_updates(st),
)
@spaces.GPU(duration=15)
def choose(st: Optional[Dict[str, Any]], index: int = 0) -> Tuple[Any, ...]:
"""Register the human's choice for the current trial and predict the next one.
Args:
st: opaque session state.
index: which of the on-screen options was pressed.
"""
if st is None or st["trial"] >= st["n_trials"]:
return (
st,
render_board(st),
render_predictions(st),
render_chart(st),
"" if st is None else st["text"],
gr.update(),
*button_updates(st),
)
letter = st["letters"][int(index)]
st["log"].append({"trial": st["trial"], "choice": letter, "probs": st["pred"]})
outcome = apply_choice(st, letter)
if st["trial"] < st["n_trials"]:
st["pred"] = _predict_all(score_prefix(st), st["letters"])
transcript = st["text"] + pre_text(st) + " <<"
else:
st["pred"] = {}
st["done"] = True
transcript = st["text"]
outcome += " Session complete — start another to compare."
return (
st,
render_board(st, outcome),
render_predictions(st),
render_chart(st),
transcript,
gr.update(),
*button_updates(st),
)
# --------------------------------------------------------------------------------------
# Tab 2 — score a real Psych-101 transcript (the paper's own evaluation metric)
# --------------------------------------------------------------------------------------
with open("psych101_examples.json") as f:
EXAMPLES: Dict[str, Dict[str, str]] = json.load(f)
EXAMPLE_NAMES = list(EXAMPLES)
@torch.no_grad()
def _score_text(model, tok, text: str) -> Tuple[List[float], List[bool], int]:
"""Per-choice negative log-likelihood over every `<<...>>` response in `text`.
Mirrors the paper's evaluation: `DataCollatorForCompletionOnlyLM` with response
template ` <<` and instruction template `>>` keeps loss on exactly the tokens
*between* the markers, summed within a response.
"""
enc = tok(text, return_tensors="pt", add_special_tokens=False, return_offsets_mapping=True)
ids = enc.input_ids[:, :MAX_SCORE_TOKENS].to("cuda")
offsets = enc.offset_mapping[0][:MAX_SCORE_TOKENS].tolist()
logits = model(ids).logits[0] # [T, V], bf16
tgt = ids[0]
# Row i-1 predicts token i. Chunked so we never materialise a full [T, vocab] fp32
# tensor (vocab is ~152k, so an unchunked log_softmax is gigabytes).
rows, targets = logits[:-1], tgt[1:]
lp_parts, hit_parts = [], []
for s in range(0, rows.shape[0], 256):
chunk = rows[s : s + 256].float()
lp = torch.log_softmax(chunk, dim=-1)
lp_parts.append(lp.gather(1, targets[s : s + 256, None]).squeeze(1))
hit_parts.append(chunk.argmax(dim=-1) == targets[s : s + 256])
# tok_lp[i] / argmax_hit[i] describe token i (index 0 is unpredictable -> padded).
tok_lp = [0.0] + torch.cat(lp_parts).tolist()
argmax_hit = [False] + torch.cat(hit_parts).tolist()
nlls: List[float] = []
hits: List[bool] = []
for m in re.finditer(r"<<(.*?)>>", text, flags=re.S):
# Score the choice only, never the ` <<` / `>>` markers: the boundary tokens
# carry 4-16 nats of pure delimiter surprise and swamp the real signal.
start, end = m.start() + 2, m.end() - 2
idxs = [i for i, (a, b) in enumerate(offsets) if b > a and a >= start and b <= end]
if not idxs or idxs[0] == 0:
continue
nlls.append(-sum(tok_lp[i] for i in idxs))
hits.append(bool(argmax_hit[idxs[-1]]))
used = int(ids.shape[1])
return nlls, hits, used
def _score_impl(transcript: str, colour_by: str) -> Tuple[str, str]:
"""Shared, undecorated scoring implementation (never call GPU-decorated fns from GPU fns)."""
text = (transcript or "").strip()
if not text or "<<" not in text:
return (
CSS_BLOCK + "<div class='cs-wrap'><div class='cs-sub'>Paste a transcript whose human "
"choices are wrapped in <code><< >></code>.</div></div>",
"",
)
t0 = time.perf_counter()
res: Dict[str, Tuple[List[float], List[bool], int]] = {}
res["big"] = _score_text(MODEL_BIG, TOK_BIG, text)
with MODEL_BIG.disable_adapter():
res["raw"] = _score_text(MODEL_BIG, TOK_BIG, text)
res["small"] = _score_text(MODEL_SMALL, TOK_SMALL, text)
dt = time.perf_counter() - t0
key = {v: k for k, v in PRED_LABEL.items()}.get(colour_by, "big")
nlls = res[key][0]
# annotate
out, cursor, i = [], 0, 0
for m in re.finditer(r"<<(.*?)>>", text, flags=re.S):
out.append(html.escape(text[cursor : m.start()]))
if i < len(nlls):
p = math.exp(-nlls[i])
hue = 120 * min(max(p, 0.0), 1.0)
out.append(
f"<span class='cs-ch' title='p = {p*100:.1f}% · loss = {nlls[i]:.2f} nats' "
f"style='background:hsla({hue:.0f},70%,45%,.28)'><<{html.escape(m.group(1))}>></span>"
)
else:
out.append(f"<<{html.escape(m.group(1))}>>")
cursor = m.end()
i += 1
out.append(html.escape(text[cursor:]))
rows = []
for k, label, _c, note in PREDICTORS:
nl, hit, used = res[k]
if not nl:
continue
rows.append(
f"<tr><td>{label} <span class='cs-mnote'>{note}</span></td>"
f"<td>{sum(nl)/len(nl):.3f}</td>"
f"<td>{100*sum(hit)/len(hit):.0f}%</td><td>{len(nl)}</td></tr>"
)
truncated = res["big"][2] >= MAX_SCORE_TOKENS
table = (
CSS_BLOCK
+ "<div class='cs-wrap'><div class='cs-head'>Held-out choice prediction</div>"
"<table class='cs-tbl'><tr><th>model</th><th>loss (nats/choice)</th>"
f"<th>top-1</th><th>choices</th></tr>{''.join(rows)}</table>"
"<div class='cs-note'>Loss is the mean negative log-likelihood of the human's actual "
"choices — the same quantity reported in the paper (lower is more human-like). "
f"Scored {res['big'][2]} tokens in {dt:.1f}s"
+ (" (transcript truncated to fit the context budget)." if truncated else ".")
+ "</div></div>"
)
annotated = (
CSS_BLOCK
+ "<div class='cs-wrap'><div class='cs-head'>Green = the model expected that choice · "
f"red = it was surprised (colours from {html.escape(colour_by)})</div>"
f"<div class='cs-tx'>{''.join(out)}</div></div>"
)
return annotated, table
@spaces.GPU(duration=15)
def score_transcript(transcript: str, colour_by: str = "Qwentaur-8B") -> Tuple[str, str]:
"""Score every human choice in a Psych-101-formatted transcript.
Args:
transcript: text with human choices wrapped in `<<...>>`.
colour_by: which model's surprise colours the transcript.
Returns:
An annotated transcript and a per-model loss table.
"""
return _score_impl(transcript, colour_by)
def load_example(name: str) -> str:
"""Load one bundled Psych-101 human transcript into the editor."""
rec = EXAMPLES.get(name) or next(iter(EXAMPLES.values()))
return rec["text"]
@spaces.GPU(duration=15)
def run_example(name: str, colour_by: str = "Qwentaur-8B") -> Tuple[str, str, str]:
"""Load a bundled Psych-101 human session and score it.
Args:
name: which bundled human session to score.
colour_by: which model's surprise colours the transcript.
"""
text = load_example(name)
annotated, table = _score_impl(text, colour_by)
return text, annotated, table
# --------------------------------------------------------------------------------------
# UI
# --------------------------------------------------------------------------------------
CSS = """
#col-container { max-width: 1180px; margin: 0 auto; }
.dark .gradio-container { color: var(--body-text-color); }
"""
INTRO = """# 🧠 Centauri Cognitive Simulator
Play a real psychology experiment while three models predict, trial by trial, what a
**human** would do next — the fine-tuned **Qwentaur-8B**, the 13× smaller
**Qwentaur-0.6B**, and the **un-tuned Qwen3-8B-Base** for contrast.
The paper's headline claim is that sub-billion-parameter models already match a
70B Centaur at fitting human choices in-distribution — you can watch that happen here.
[Paper](https://huggingface.co/papers/2608.05224) ·
[Code](https://github.com/socius-org/Centauri) ·
[Adapters](https://huggingface.co/collections/socius/centauri-6a72e25a4669e413571fe4ac) ·
[Psych-101](https://huggingface.co/datasets/marcelbinz/Psych-101)
"""
with gr.Blocks(title="Centauri Cognitive Simulator") as demo:
with gr.Column(elem_id="col-container"):
gr.Markdown(INTRO)
with gr.Tabs():
with gr.Tab("Play an experiment"):
# A hidden gr.JSON rather than gr.State: it serialises, so the session
# round-trips through the REST / MCP API instead of resetting each call.
state = gr.JSON(value=None, visible=False, label="session")
with gr.Row():
task_radio = gr.Radio(
choices=[TASK_LABELS[k] for k in TASK_ORDER],
value=TASK_LABELS["bandit"],
label="Experiment",
scale=4,
)
start_btn = gr.Button("Start experiment", variant="primary", scale=1)
with gr.Row():
with gr.Column(scale=3):
board = gr.HTML(render_board(None))
with gr.Row():
btns = [
gr.Button("—", visible=False, variant="secondary") for _ in range(4)
]
# hidden constants so each button submits its own option index
# (and so `index` is a real API argument, unlike a gr.State)
idx_boxes = [
gr.Number(value=i, visible=False, precision=0) for i in range(4)
]
chart = gr.HTML()
with gr.Column(scale=2):
preds = gr.HTML(render_predictions(None))
with gr.Accordion("What the models actually see (Psych-101 prompt)", open=False):
transcript = gr.Textbox(
label="Prompt",
show_label=False,
lines=14,
max_lines=24,
buttons=["copy"],
interactive=False,
)
with gr.Accordion("Advanced", open=False):
with gr.Row():
seed_num = gr.Number(label="Seed", value=0, precision=0)
rand_seed = gr.Checkbox(label="Randomize seed", value=True)
blind_cb = gr.Checkbox(
label="Blind mode (reveal predictions only after each trial)",
value=False,
)
gr.Markdown(
"Choice probabilities are read straight off the next-token logits after "
"`You press <<`, renormalised over the options — no sampling, no prompt "
"engineering. Seeing the predictions before you act can bias you; turn on "
"blind mode for a cleaner test of yourself."
)
out_common = [
state,
board,
preds,
chart,
transcript,
seed_num,
*btns,
]
start_btn.click(
start_session,
inputs=[task_radio, seed_num, rand_seed, blind_cb],
outputs=out_common,
api_name="start_session",
)
for i, b in enumerate(btns):
# One documented /choose endpoint; the other three buttons are the
# same call with a different option index, so hide them. In Gradio 6
# `api_name=False` would name them "false", "false_1", ...
extra = (
{"api_name": "choose"}
if i == 0
else {"api_visibility": "private"}
)
b.click(
choose,
inputs=[state, idx_boxes[i]],
outputs=out_common,
**extra,
)
with gr.Tab("Score real human data"):
gr.Markdown(
"Score the models the way the paper does: how surprised is each one by the "
"choices a **real participant** actually made? These transcripts come "
"verbatim from [Psych-101]"
"(https://huggingface.co/datasets/marcelbinz/Psych-101) (Binz et al., "
"Apache-2.0), truncated to the first few trials."
)
with gr.Row():
preset = gr.Dropdown(
choices=EXAMPLE_NAMES,
value=EXAMPLE_NAMES[0],
label="Bundled human session",
scale=3,
)
colour_by = gr.Dropdown(
choices=[label for _k, label, _c, _n in PREDICTORS],
value="Qwentaur-8B",
label="Colour transcript by",
scale=2,
)
score_btn = gr.Button("Score choices", variant="primary", scale=1)
text_box = gr.Textbox(
label="Transcript (editable — paste your own Psych-101-style text)",
value=load_example(EXAMPLE_NAMES[0]),
lines=10,
max_lines=20,
)
score_table = gr.HTML()
annotated_out = gr.HTML()
preset.change(load_example, inputs=[preset], outputs=[text_box])
score_btn.click(
score_transcript,
inputs=[text_box, colour_by],
outputs=[annotated_out, score_table],
api_name="score_transcript",
)
gr.Examples(
examples=[[n] for n in EXAMPLE_NAMES],
inputs=[preset],
outputs=[text_box, annotated_out, score_table],
fn=run_example,
cache_examples=True,
cache_mode="lazy",
label="Real human sessions from Psych-101",
)
gr.Markdown(
"Models: [socius/Qwentaur-8B-LoRA-r16](https://huggingface.co/socius/Qwentaur-8B-LoRA-r16) "
"and [socius/Qwentaur-0.6B-LoRA-r16](https://huggingface.co/socius/Qwentaur-0.6B-LoRA-r16) "
"(LoRA adapters on Qwen3-Base, Apache-2.0) from *Small Foundation Models of Human "
"Cognition and Behaviour* (Oh & Gobet, 2026). Task wording and the bundled human "
"sessions come from Psych-101 (Binz et al., 2025)."
)
if __name__ == "__main__":
# Gradio 6 moved `theme` / `css` from the Blocks constructor to launch().
demo.launch(theme=gr.themes.Citrus(), css=CSS, mcp_server=True)
|