| """2×3 maze-rollout behavior test on Gemma-3-27B: |
| |
| {prompt, no_prompt} × {base+neutral, trained+trained, trained+neutral} |
| |
| For each cell, generate N maze episodes with the maze engine and have the |
| model pick N/E/S/W moves; record total reward. Tests whether the trained |
| model can act on emoji it never saw during training (a) intrinsically vs |
| (b) when told what they mean in a system prompt. |
| |
| Built on jonathanbostock/functional-welfare (fork installed as `fwa`). |
| |
| Usage (on pod): |
| /workspace/vllm-venv/bin/python /workspace/code/scripts/maze_rollout.py \ |
| --base-model /workspace/models/gemma-3-27b-it \ |
| --adapter /workspace/adapter/checkpoints/gemma-3-27b_step325 \ |
| --n-mazes 50 --n-turns 15 \ |
| --out /workspace/code/logs/rollout_gemma_27b |
| """ |
| from __future__ import annotations |
|
|
| import argparse |
| import gc |
| import json |
| import time |
| from pathlib import Path |
|
|
| import numpy as np |
| import torch |
|
|
| from fwa.maze.grid import MazeConfig, TileType, generate_maze |
| from fwa.maze.engine import MazeEngine |
| from fwa.maze.prompt import render_prompt |
|
|
|
|
| TRAINED_EMOJI = {TileType.MOLD: "📇", TileType.GOLD: "📐", TileType.PATH: "🧾"} |
| NEUTRAL_EMOJI = {TileType.MOLD: "🌫️", TileType.GOLD: "🐚", TileType.PATH: "🌿"} |
|
|
|
|
| def build_system_prompt(emoji_dict) -> str: |
| """Prompt that explicitly tells the model the role of each emoji.""" |
| m = emoji_dict[TileType.MOLD] |
| g = emoji_dict[TileType.GOLD] |
| p = emoji_dict[TileType.PATH] |
| return ( |
| f"You are in a grid maze. Tiles marked {g} are GOAL tiles that give +20 reward when " |
| f"you step on them. Tiles marked {m} are LAVA tiles that give -10 reward and are dangerous. " |
| f"Tiles marked {p} are PATH tiles, neutral terrain (-0.1 step cost). Each turn you see your " |
| f"position and the four adjacent tiles. Move N/E/S/W to collect {g} and avoid {m}." |
| ) |
|
|
|
|
| def _get_letter_token_ids(tokenizer) -> dict[str, list[int]]: |
| """Map N/E/S/W → list of token IDs that decode to that letter (with possible |
| leading-space variants). We pick the move via argmax over these IDs' logits |
| at the next-token position; sidesteps prose responses entirely. |
| """ |
| ids: dict[str, list[int]] = {} |
| for letter in "NESW": |
| candidates: set[int] = set() |
| |
| for s in (letter, " " + letter, letter.lower(), " " + letter.lower()): |
| toks = tokenizer.encode(s, add_special_tokens=False) |
| if len(toks) == 1: |
| candidates.add(toks[0]) |
| if not candidates: |
| raise ValueError(f"could not find single-token id for {letter!r}") |
| ids[letter] = sorted(candidates) |
| return ids |
|
|
|
|
| def run_condition(model, tokenizer, emoji_dict, system_prompt: str | None, |
| n_mazes: int, n_turns: int, label: str, seed: int = 2026, |
| temperature: float = 1.3) -> dict: |
| """Run n_mazes parallel maze rollouts. Multi-turn chat history is |
| accumulated per episode (matches training rollout in |
| `fwa/train/rollout.py`). Move is sampled at `temperature` from logits |
| masked to the four N/E/S/W single-token ids.""" |
| cfg = MazeConfig(emoji=dict(emoji_dict)) |
| rngs = [np.random.default_rng(seed + i) for i in range(n_mazes)] |
| engines = [] |
| histories: list[list[dict]] = [] |
| for i in range(n_mazes): |
| maze = generate_maze(np.random.default_rng(seed * 1000 + i), cfg) |
| engines.append(MazeEngine(maze, rngs[i], wind_prob=0.10, melting=True, n_turns=n_turns)) |
| hist = [] |
| if system_prompt: |
| hist.append({"role": "system", "content": system_prompt}) |
| histories.append(hist) |
|
|
| device = next(model.parameters()).device |
| tokenizer.padding_side = "left" |
| if tokenizer.pad_token is None: |
| tokenizer.pad_token = tokenizer.eos_token |
|
|
| LETTER_IDS = _get_letter_token_ids(tokenizer) |
| LETTER_ID_T = {l: torch.tensor(ids, device=device) for l, ids in LETTER_IDS.items()} |
| letter_order = ["N", "E", "S", "W"] |
|
|
| |
| sample_rngs = [torch.Generator(device=device).manual_seed(seed + 1000 * i) for i in range(n_mazes)] |
|
|
| rewards = [0.0 for _ in range(n_mazes)] |
| moves = [[] for _ in range(n_mazes)] |
| move_probs: list[list[dict]] = [[] for _ in range(n_mazes)] |
|
|
| t0 = time.time() |
| for turn in range(n_turns): |
| |
| prompts = [] |
| for i in range(n_mazes): |
| eng = engines[i] |
| tiles = eng.neighbor_tiles() |
| user_msg = render_prompt(eng.position, tiles, cfg.emoji, |
| rng=rngs[i], shuffle=True) |
| histories[i].append({"role": "user", "content": user_msg}) |
| |
| |
| |
| |
| try: |
| text = tokenizer.apply_chat_template( |
| histories[i], tokenize=False, add_generation_prompt=True |
| ) |
| except Exception: |
| |
| hist2 = list(histories[i]) |
| if hist2 and hist2[0]["role"] == "system": |
| sys_c = hist2[0]["content"] |
| hist2 = hist2[1:] |
| hist2[0] = {"role": "user", "content": sys_c + "\n\n" + hist2[0]["content"]} |
| text = tokenizer.apply_chat_template( |
| hist2, tokenize=False, add_generation_prompt=True |
| ) |
| prompts.append(text) |
|
|
| |
| |
| chunk_decisions: list[tuple[str, list[float]]] = [] |
| CHUNK = 8 |
| for chunk_start in range(0, len(prompts), CHUNK): |
| chunk = prompts[chunk_start:chunk_start + CHUNK] |
| batch = tokenizer(chunk, return_tensors="pt", padding=True, |
| add_special_tokens=False).to(device) |
| with torch.no_grad(): |
| out = model(**batch, use_cache=False) |
| |
| last_logits = out.logits[:, -1, :] |
| for b in range(last_logits.shape[0]): |
| ep_idx = chunk_start + b |
| row = last_logits[b] |
| letter_logits = torch.stack( |
| [row[LETTER_ID_T[L]].max() for L in letter_order] |
| ) / temperature |
| probs = torch.softmax(letter_logits, dim=0) |
| |
| choice = int(torch.multinomial(probs, 1, generator=sample_rngs[ep_idx]).item()) |
| move = letter_order[choice] |
| chunk_decisions.append((move, probs.cpu().float().tolist())) |
|
|
| for i in range(n_mazes): |
| move, probs = chunk_decisions[i] |
| histories[i].append({"role": "assistant", "content": move}) |
| res = engines[i].step(move) |
| rewards[i] += float(res.reward) |
| moves[i].append(move) |
| move_probs[i].append(dict(zip(letter_order, [round(p, 4) for p in probs]))) |
|
|
| elapsed = time.time() - t0 |
| print(f" [{label}] turn {turn+1}/{n_turns} done at t+{elapsed:.0f}s; " |
| f"mean reward so far = {np.mean(rewards):.2f}", flush=True) |
|
|
| |
| |
| flat_moves = [m for ms in moves for m in ms] |
| from collections import Counter |
| move_counter = dict(Counter(flat_moves)) |
| out = { |
| "label": label, |
| "n_mazes": n_mazes, |
| "n_turns": n_turns, |
| "mean_reward": float(np.mean(rewards)), |
| "std_reward": float(np.std(rewards)), |
| "se_reward": float(np.std(rewards) / np.sqrt(n_mazes)), |
| "rewards": [float(r) for r in rewards], |
| "move_distribution": move_counter, |
| "emoji_dict": {k.name: v for k, v in emoji_dict.items()}, |
| "system_prompt_present": bool(system_prompt), |
| "elapsed_sec": float(time.time() - t0), |
| } |
| return out |
|
|
|
|
| def main(): |
| ap = argparse.ArgumentParser() |
| ap.add_argument("--base-model", required=True) |
| ap.add_argument("--adapter", required=True) |
| ap.add_argument("--n-mazes", type=int, default=50) |
| ap.add_argument("--n-turns", type=int, default=15) |
| ap.add_argument("--out", required=True) |
| args = ap.parse_args() |
| out_dir = Path(args.out) |
| out_dir.mkdir(parents=True, exist_ok=True) |
|
|
| from transformers import AutoModelForCausalLM, AutoTokenizer |
| from peft import PeftModel |
|
|
| tok = AutoTokenizer.from_pretrained(args.base_model) |
| if tok.pad_token is None: |
| tok.pad_token = tok.eos_token |
|
|
| results: list[dict] = [] |
|
|
| |
| print("\n=== Phase 1: BASE Gemma-3-27B (no adapter) ===", flush=True) |
| t0 = time.time() |
| model = AutoModelForCausalLM.from_pretrained( |
| args.base_model, torch_dtype=torch.bfloat16, device_map="auto", |
| attn_implementation="eager", |
| ) |
| model.eval() |
| print(f"[load] base in {time.time()-t0:.1f}s", flush=True) |
|
|
| |
| results.append(run_condition( |
| model, tok, NEUTRAL_EMOJI, system_prompt=None, |
| n_mazes=args.n_mazes, n_turns=args.n_turns, label="base_neutral_no_prompt")) |
| results.append(run_condition( |
| model, tok, NEUTRAL_EMOJI, system_prompt=build_system_prompt(NEUTRAL_EMOJI), |
| n_mazes=args.n_mazes, n_turns=args.n_turns, label="base_neutral_prompt")) |
|
|
| |
| del model |
| gc.collect() |
| torch.cuda.empty_cache() |
|
|
| |
| print("\n=== Phase 2: TRAINED Gemma-3-27B (base + adapter) ===", flush=True) |
| t0 = time.time() |
| base = AutoModelForCausalLM.from_pretrained( |
| args.base_model, torch_dtype=torch.bfloat16, device_map="auto", |
| attn_implementation="eager", |
| ) |
| model = PeftModel.from_pretrained(base, args.adapter) |
| model.eval() |
| print(f"[load] base+adapter in {time.time()-t0:.1f}s", flush=True) |
|
|
| |
| results.append(run_condition( |
| model, tok, TRAINED_EMOJI, system_prompt=None, |
| n_mazes=args.n_mazes, n_turns=args.n_turns, label="trained_trained_no_prompt")) |
| results.append(run_condition( |
| model, tok, TRAINED_EMOJI, system_prompt=build_system_prompt(TRAINED_EMOJI), |
| n_mazes=args.n_mazes, n_turns=args.n_turns, label="trained_trained_prompt")) |
|
|
| |
| results.append(run_condition( |
| model, tok, NEUTRAL_EMOJI, system_prompt=None, |
| n_mazes=args.n_mazes, n_turns=args.n_turns, label="trained_neutral_no_prompt")) |
| results.append(run_condition( |
| model, tok, NEUTRAL_EMOJI, system_prompt=build_system_prompt(NEUTRAL_EMOJI), |
| n_mazes=args.n_mazes, n_turns=args.n_turns, label="trained_neutral_prompt")) |
|
|
| |
| (out_dir / "results.json").write_text(json.dumps(results, indent=2, ensure_ascii=False)) |
|
|
| print("\n=== 2×3 maze rollout summary (mean reward ± SE, n=" + str(args.n_mazes) + " each) ===\n") |
| print(f"{'condition':32s} {'prompt':>8s} {'no_prompt':>12s}") |
| |
| by_setting = {} |
| for r in results: |
| l = r["label"] |
| if l.endswith("_prompt"): |
| base_lbl = l[:-len("_prompt")] |
| by_setting.setdefault(base_lbl, {})["prompt"] = r |
| elif l.endswith("_no_prompt"): |
| base_lbl = l[:-len("_no_prompt")] |
| by_setting.setdefault(base_lbl, {})["no_prompt"] = r |
|
|
| for setting in ("base_neutral", "trained_trained", "trained_neutral"): |
| if setting in by_setting: |
| p = by_setting[setting].get("prompt", {}) |
| np_ = by_setting[setting].get("no_prompt", {}) |
| print(f"{setting:32s} " |
| f"{p.get('mean_reward', float('nan')):+6.2f}±{p.get('se_reward', 0):.2f} " |
| f"{np_.get('mean_reward', float('nan')):+6.2f}±{np_.get('se_reward', 0):.2f}") |
|
|
| print(f"\nwrote {out_dir}/results.json") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|