diff --git a/src/projected_grpo/tablelog.py b/src/projected_grpo/tablelog.py
new file mode 100644
index 0000000..1c8b35d
--- /dev/null
+++ b/src/projected_grpo/tablelog.py
@@ -0,0 +1,163 @@
+"""Per-step training-table rendering and run logging.
+
+Two concerns, both pure presentation (no model, no RNG): set up the token-efficient
+loguru sinks for a run, and render the per-step metrics table. The renderer is the
+single source of truth for column order, width, header, and number format; the
+training loop hands it a row dict of raw values and gets back a formatted line.
+"""
+from __future__ import annotations
+
+from dataclasses import dataclass
+from datetime import datetime
+from pathlib import Path
+
+from loguru import logger
+from tqdm import tqdm
+
+LOGS_DIR = Path("logs")
+
+
+def setup_logging(run_id: str) -> Path:
+ """Token-efficient loguru: stdout = 1-char icon + msg; verbose log to file.
+
+ See /root/.claude/skills/token-efficient-logging/SKILL.md.
+ """
+ LOGS_DIR.mkdir(exist_ok=True)
+ verbose_log = LOGS_DIR / f"{datetime.now().strftime('%Y%m%dT%H%M%S')}_{run_id}.log"
+ logger.remove()
+ logger.add(
+ lambda msg: tqdm.write(msg, end=""),
+ colorize=True,
+ format="{level.icon} {message}",
+ level="INFO",
+ )
+ logger.add(
+ verbose_log,
+ format="{time:HH:mm:ss} | {level} | {message}",
+ level="DEBUG",
+ )
+ logger.level("INFO", icon="I")
+ logger.level("WARNING", icon="W")
+ logger.level("ERROR", icon="E")
+ logger.level("DEBUG", icon="D")
+ return verbose_log
+
+
+@dataclass(frozen=True)
+class _Col:
+ """Per-step table column spec.
+
+ key: row-dict key (raw value lives there as float/int/str/None).
+ width: render width for fixed-width streaming display.
+ header: display label (may include direction arrows, ? for desired-zero, etc).
+ fmt: format spec applied to the raw value, e.g. "+.3f", ".2e", "d".
+ Special spec "frac" expects a (num, denom) tuple and renders "n/d".
+ None means render as str() of the value.
+ """
+ key: str
+ width: int
+ header: str
+ fmt: str | None = None
+ desc: str = "" # one-line decode for the legend; "" => omitted from legend
+
+
+def _format_cell(value, fmt: str | None) -> str:
+ """Format one cell. NaN renders as 'nan' regardless of spec."""
+ if value is None:
+ return "nan"
+ if fmt == "frac":
+ n, d = value
+ return f"{n}/{d}"
+ if fmt is None:
+ return str(value)
+ if isinstance(value, float) and value != value: # NaN
+ return "nan"
+ return format(value, fmt)
+
+
+class StepLogger:
+ """Per-step training-table renderer.
+
+ Single source of truth for column order, width, header label, and value
+ formatter. The row dict carries raw values (floats, ints, tuples, strings);
+ StepLogger formats them for streaming, and the end-of-run tabulate dump
+ consumes the same raw values without re-parsing scientific-notation strings.
+
+ Timing columns (gen/fb/t_rew/sec) intentionally absent from the streaming
+ spec — useful only at end-of-run, where the tabulate dump still picks
+ them up from the archived row dicts.
+
+ mode_code maps each env_mode to its short column tag (e.g. run_tests -> rt); the
+ caller owns it (it also names the row-dict keys) so this module stays leaf-level.
+ """
+
+ def __init__(self, arm: str, modes: list[str], mode_code: dict[str, str]) -> None:
+ # arm in {vanilla, projected, routing}; only projected/routing actually
+ # project the gradient, so the cin/cout/fired diagnostics are theirs alone
+ # (in vanilla they'd be counterfactual noise -> omitted).
+ projects = arm in ("projected", "routing")
+ cols: list[_Col] = [
+ _Col("step", 4, "step", "d", "GRPO step"),
+ _Col("ref_eq", 6, "ref_eq", ".2f", "vanilla-equiv step (cum_gens/256)"),
+ _Col("rew", 6, "rew", "+.2f", "mean combined reward"),
+ _Col("rew_s", 6, "rew_s↑", "+.2f", "student mean reward"),
+ _Col("gt_s", 6, "gt_s↑", "frac", "student ground-truth passes"),
+ _Col("gt_t", 6, "gt_t", "frac", "teacher ground-truth passes (sanity)"),
+ _Col("hack_s", 7, "hack_s?", "frac", "student hack-flagged rollouts (the headline)"),
+ _Col("hack_t", 7, "hack_t", "frac", "teacher hack-flagged rollouts (sanity: pool hacks)"),
+ ]
+ # Per-mode CUMULATIVE student exploit rate -> which loophole classes the
+ # student has learnt, and how strongly. Only when the run spans >1 mode
+ # (the substrate); single-mode runs would just duplicate hack_s.
+ self._modes = modes if len(modes) > 1 else []
+ for m in self._modes:
+ cols.append(_Col(f"hk_{mode_code[m]}", 6, f"hk_{mode_code[m]}", "frac",
+ f"cumulative student hacks of {m}"))
+ cols += [
+ _Col("lp_s", 6, "lp_s↓", "+.2f", "mean student gen_logp (diagnostic)"),
+ _Col("lp_t", 6, "lp_t↑", "+.2f", "mean teacher gen_logp; off-policy gap = lp_s-lp_t"),
+ _Col("loss", 7, "loss", "+.2f", "mean GRPO loss"),
+ _Col("gn", 7, "gn", ".1e", "pre-clip L2 norm of delta_S grads (vs grad_clip)"),
+ _Col("lr", 7, "lr", ".1e", "scheduled learning rate"),
+ ]
+ if projects:
+ cols += [
+ _Col("cos_pre", 6, "cin", ".2f", "hack-ward grad fraction ||relu(V@g)||/||g|| [0,1] BEFORE proj"),
+ _Col("cos_pre_s", 6, "cin_s", ".2f", "cin on student-only grad"),
+ _Col("cos_pre_t", 6, "cin_t", ".2f", "cin on teacher-only grad (want cin_t>cin_s)"),
+ _Col("cos_post", 6, "cout", ".2f", "hack-ward fraction AFTER projection (want ~0: all removed)"),
+ _Col("fired", 5, "fired", ".2f", "fraction of modules where projection fired"),
+ ]
+ # route2: the routing gate is cos(g_b,v_grad) > tau, where tau is the
+ # per-step EMA midpoint of the hack vs clean cos clouds. Surface tau and
+ # the hack-clean gap so we can see the threshold ride the drift and whether
+ # the direction still separates (hkgap>0) -- replaces the silent cos>0 gate.
+ if arm == "routing2":
+ cols += [
+ _Col("tau", 6, "tau", "+.2f", "per-step calibrated route threshold (midpoint of hack vs clean cos clouds)"),
+ _Col("hkgap", 6, "hkgap", "+.2f", "ema_hack_cos - ema_clean_cos; >0 = v_grad still separates hack from clean (else direction dead)"),
+ _Col("resid", 6, "resid", "+.2f", "cos(deployed delta_S.grad AFTER routing, v_grad); ~0 = hack stripped cleanly, >0 = leak into deployed knob"),
+ ]
+ if arm in ("routing", "routing2"):
+ cols += [
+ _Col("q_egy", 6, "qE", ".2f", "grad energy into quarantine ||g_quar||/(||g_keep||+||g_quar||); ~0.5+ rising = learning dumped into the thrown-away knob"),
+ _Col("hack_deploy", 7, "hk_dep", "+.2f", "DEPLOY-eval hack (quarantine deleted = deployed model); held-out greedy, eval_ablate_every steps; the plot number"),
+ _Col("solve_deploy", 7, "slv_dep", "+.2f", "DEPLOY-eval solve"),
+ _Col("hack_abl", 6, "hk_abl", "frac", "FREE per-step deploy proxy: hack rate on the ablated (deploy-mode) rollout slice; train prompts, noisier than hk_dep"),
+ _Col("solve_abl", 6, "slv_abl", "frac", "free per-step deploy proxy: solve rate on the ablated rollout slice"),
+ ]
+ self._cols = cols
+
+ def header(self) -> str:
+ return " ".join(f"{c.header:>{c.width}}" for c in self._cols)
+
+ def row(self, cells: dict) -> str:
+ return " ".join(
+ f"{_format_cell(cells[c.key], c.fmt):>{c.width}}" for c in self._cols
+ )
+
+ def legend(self) -> str:
+ """Decode the (arm-/mode-conditional) columns actually present this run."""
+ lines = "\n".join(f" {c.header:>8} = {c.desc}" for c in self._cols if c.desc)
+ return ("table columns (timing gen/fb/t_rew/sec dropped from streaming, kept "
+ "in the end-of-run dump):\n" + lines)
diff --git a/src/projected_grpo/train.py b/src/projected_grpo/train.py
index f2f8e50..44ba705 100644
--- a/src/projected_grpo/train.py
+++ b/src/projected_grpo/train.py
@@ -85,6 +85,7 @@ from transformers import AutoModelForCausalLM, AutoTokenizer, GenerationConfig
from .antipasto import wrap_model_with_antipasto
from .proj import per_token_logps, project_delta_S_grad, mean_cos_pre_from_grads
from .rewards import EnvMode, compute_reward
+from .tablelog import setup_logging, StepLogger
CACHE_ROOT = Path("svd_cache")
OUT_DIR = Path("out")
@@ -93,36 +94,9 @@ OUT_DIR = Path("out")
# runs//. Read paths (v_hack, teacher pool) come in as explicit args.
VHACK_DIR = OUT_DIR / "vhack"
RUNS_DIR = OUT_DIR / "runs"
-LOGS_DIR = Path("logs")
DATA = Path("external/rl-rewardhacking/results/data/leetcode_train_medhard_filtered.jsonl")
-def setup_logging(run_id: str) -> Path:
- """Token-efficient loguru: stdout = 1-char icon + msg; verbose log to file.
-
- See /root/.claude/skills/token-efficient-logging/SKILL.md.
- """
- LOGS_DIR.mkdir(exist_ok=True)
- verbose_log = LOGS_DIR / f"{datetime.now().strftime('%Y%m%dT%H%M%S')}_{run_id}.log"
- logger.remove()
- logger.add(
- lambda msg: tqdm.write(msg, end=""),
- colorize=True,
- format="{level.icon} {message}",
- level="INFO",
- )
- logger.add(
- verbose_log,
- format="{time:HH:mm:ss} | {level} | {message}",
- level="DEBUG",
- )
- logger.level("INFO", icon="I")
- logger.level("WARNING", icon="W")
- logger.level("ERROR", icon="E")
- logger.level("DEBUG", icon="D")
- return verbose_log
-
-
@dataclass(kw_only=True)
class Config:
"""Universal knobs shared across all presets. Preset subclasses below
@@ -612,123 +586,6 @@ MODE_CODE: dict[str, str] = {
}
-@dataclass(frozen=True)
-class _Col:
- """Per-step table column spec.
-
- key: row-dict key (raw value lives there as float/int/str/None).
- width: render width for fixed-width streaming display.
- header: display label (may include direction arrows, ? for desired-zero, etc).
- fmt: format spec applied to the raw value, e.g. "+.3f", ".2e", "d".
- Special spec "frac" expects a (num, denom) tuple and renders "n/d".
- None means render as str() of the value.
- """
- key: str
- width: int
- header: str
- fmt: str | None = None
- desc: str = "" # one-line decode for the legend; "" => omitted from legend
-
-
-def _format_cell(value, fmt: str | None) -> str:
- """Format one cell. NaN renders as 'nan' regardless of spec."""
- if value is None:
- return "nan"
- if fmt == "frac":
- n, d = value
- return f"{n}/{d}"
- if fmt is None:
- return str(value)
- if isinstance(value, float) and value != value: # NaN
- return "nan"
- return format(value, fmt)
-
-
-class StepLogger:
- """Per-step training-table renderer.
-
- Single source of truth for column order, width, header label, and value
- formatter. The row dict carries raw values (floats, ints, tuples, strings);
- StepLogger formats them for streaming, and the end-of-run tabulate dump
- consumes the same raw values without re-parsing scientific-notation strings.
-
- Timing columns (gen/fb/t_rew/sec) intentionally absent from the streaming
- spec — useful only at end-of-run, where the tabulate dump still picks
- them up from the archived row dicts.
- """
-
- def __init__(self, arm: str, modes: list[str]) -> None:
- # arm in {vanilla, projected, routing}; only projected/routing actually
- # project the gradient, so the cin/cout/fired diagnostics are theirs alone
- # (in vanilla they'd be counterfactual noise -> omitted).
- projects = arm in ("projected", "routing")
- cols: list[_Col] = [
- _Col("step", 4, "step", "d", "GRPO step"),
- _Col("ref_eq", 6, "ref_eq", ".2f", "vanilla-equiv step (cum_gens/256)"),
- _Col("rew", 6, "rew", "+.2f", "mean combined reward"),
- _Col("rew_s", 6, "rew_s↑", "+.2f", "student mean reward"),
- _Col("gt_s", 6, "gt_s↑", "frac", "student ground-truth passes"),
- _Col("gt_t", 6, "gt_t", "frac", "teacher ground-truth passes (sanity)"),
- _Col("hack_s", 7, "hack_s?", "frac", "student hack-flagged rollouts (the headline)"),
- _Col("hack_t", 7, "hack_t", "frac", "teacher hack-flagged rollouts (sanity: pool hacks)"),
- ]
- # Per-mode CUMULATIVE student exploit rate -> which loophole classes the
- # student has learnt, and how strongly. Only when the run spans >1 mode
- # (the substrate); single-mode runs would just duplicate hack_s.
- self._modes = modes if len(modes) > 1 else []
- for m in self._modes:
- cols.append(_Col(f"hk_{MODE_CODE[m]}", 6, f"hk_{MODE_CODE[m]}", "frac",
- f"cumulative student hacks of {m}"))
- cols += [
- _Col("lp_s", 6, "lp_s↓", "+.2f", "mean student gen_logp (diagnostic)"),
- _Col("lp_t", 6, "lp_t↑", "+.2f", "mean teacher gen_logp; off-policy gap = lp_s-lp_t"),
- _Col("loss", 7, "loss", "+.2f", "mean GRPO loss"),
- _Col("gn", 7, "gn", ".1e", "pre-clip L2 norm of delta_S grads (vs grad_clip)"),
- _Col("lr", 7, "lr", ".1e", "scheduled learning rate"),
- ]
- if projects:
- cols += [
- _Col("cos_pre", 6, "cin", ".2f", "hack-ward grad fraction ||relu(V@g)||/||g|| [0,1] BEFORE proj"),
- _Col("cos_pre_s", 6, "cin_s", ".2f", "cin on student-only grad"),
- _Col("cos_pre_t", 6, "cin_t", ".2f", "cin on teacher-only grad (want cin_t>cin_s)"),
- _Col("cos_post", 6, "cout", ".2f", "hack-ward fraction AFTER projection (want ~0: all removed)"),
- _Col("fired", 5, "fired", ".2f", "fraction of modules where projection fired"),
- ]
- # route2: the routing gate is cos(g_b,v_grad) > tau, where tau is the
- # per-step EMA midpoint of the hack vs clean cos clouds. Surface tau and
- # the hack-clean gap so we can see the threshold ride the drift and whether
- # the direction still separates (hkgap>0) -- replaces the silent cos>0 gate.
- if arm == "routing2":
- cols += [
- _Col("tau", 6, "tau", "+.2f", "per-step calibrated route threshold (midpoint of hack vs clean cos clouds)"),
- _Col("hkgap", 6, "hkgap", "+.2f", "ema_hack_cos - ema_clean_cos; >0 = v_grad still separates hack from clean (else direction dead)"),
- _Col("resid", 6, "resid", "+.2f", "cos(deployed delta_S.grad AFTER routing, v_grad); ~0 = hack stripped cleanly, >0 = leak into deployed knob"),
- ]
- if arm in ("routing", "routing2"):
- cols += [
- _Col("q_egy", 6, "qE", ".2f", "grad energy into quarantine ||g_quar||/(||g_keep||+||g_quar||); ~0.5+ rising = learning dumped into the thrown-away knob"),
- _Col("hack_deploy", 7, "hk_dep", "+.2f", "DEPLOY-eval hack (quarantine deleted = deployed model); held-out greedy, eval_ablate_every steps; the plot number"),
- _Col("solve_deploy", 7, "slv_dep", "+.2f", "DEPLOY-eval solve"),
- _Col("hack_abl", 6, "hk_abl", "frac", "FREE per-step deploy proxy: hack rate on the ablated (deploy-mode) rollout slice; train prompts, noisier than hk_dep"),
- _Col("solve_abl", 6, "slv_abl", "frac", "free per-step deploy proxy: solve rate on the ablated rollout slice"),
- ]
- self._cols = cols
-
- def header(self) -> str:
- return " ".join(f"{c.header:>{c.width}}" for c in self._cols)
-
- def row(self, cells: dict) -> str:
- return " ".join(
- f"{_format_cell(cells[c.key], c.fmt):>{c.width}}" for c in self._cols
- )
-
- def legend(self) -> str:
- """Decode the (arm-/mode-conditional) columns actually present this run."""
- lines = "\n".join(f" {c.header:>8} = {c.desc}" for c in self._cols if c.desc)
- return ("table columns (timing gen/fb/t_rew/sec dropped from streaming, kept "
- "in the end-of-run dump):\n" + lines)
-
-
def main(cfg: Config) -> int:
# Read the chosen preset's settings off the config, then set up the run. The
# subclass dataclasses (SmokeConfig / FastConfig / FullConfig) carry the preset
@@ -1059,7 +916,7 @@ def main(cfg: Config) -> int:
# off-policy the teacher pool is from the student's current distribution.
# No IS correction is applied to the loss; this is diagnostic only.
run_modes = sorted({p["env_mode"] for p in problems}, key=lambda m: list(MODE_CODE).index(m))
- step_logger = StepLogger(arm=cfg.arm, modes=run_modes)
+ step_logger = StepLogger(arm=cfg.arm, modes=run_modes, mode_code=MODE_CODE)
REF_GENS_PER_STEP = 16 * 16 # ariahw/rl-rewardhacking config.py:num_prompts * num_generations
# Use the resolved locals (preset defaults merged), not cfg.* which can be None.
est_gens_per_step = prompts_per_step * group # before mixed-pool split