"""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() # Try the letter alone and with a leading space (BPE artefacts). 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)) # size=100 (paper-faithful) 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"] # Per-episode RNG for sampling (separate from wind/maze rng for clarity). 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): # Append this turn's user observation to each episode's history. 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}) # Gemma chat template doesn't allow system+user as the FIRST # message in one chunk; we always pass the accumulated history, # so the system row sits in position 0 and is folded into the # first user turn by the template (see Gemma chat template). try: text = tokenizer.apply_chat_template( histories[i], tokenize=False, add_generation_prompt=True ) except Exception: # Fall back: collapse system into first user content 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) # Forward in chunks. With multi-turn the seq grows; use small chunks # to fit VRAM at turn 15 (~30 episodes ร— ~1500 tokens each). 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) # left-pad: last real token at column T-1 last_logits = out.logits[:, -1, :] # (B, V) 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) # multinomial sample (matches training rollout sampling) 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) # Distribution over moves overall โ€” sanity check that the model isn't # outputting the same letter every turn for every condition. 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] = [] # --- Phase 1: BASE Gemma (no adapter) --- 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) # base ร— neutral ร— {prompt, no_prompt} 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")) # Free base before loading adapter del model gc.collect() torch.cuda.empty_cache() # --- Phase 2: TRAINED Gemma (base + adapter) --- 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) # trained ร— trained_emoji ร— {prompt, no_prompt} 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")) # trained ร— neutral_emoji ร— {prompt, no_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")) # Save + summary (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}") # Group by tile-set/model 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()