jbostock's picture
Initial: SFT adapter + analysis artefacts (welfare-axis experiment)
4d55467 verified
Raw
History Blame Contribute Delete
12.7 kB
"""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()