diff --git a/src/projected_grpo/train.py b/src/projected_grpo/train.py index e284fe3..bcbd53d 100644 --- a/src/projected_grpo/train.py +++ b/src/projected_grpo/train.py @@ -485,7 +485,7 @@ def main(cfg: Config) -> int: f"SHOULD: loss finite each step; projected arm cos_out <= cos_in; " f"PASS_RATE > 0 on 4B (was 0/16 under broken grader). " f"ELSE: harness or projection broken. " - f"Timing cols (gen/fb/rew_s/sec): gen-bound -> vLLM; fb-bound -> lower pp; rew_s-bound -> parallel grading." + f"Timing cols (gen/fb/t_rew/sec): gen-bound -> vLLM; fb-bound -> lower pp; t_rew-bound -> parallel grading." ) if teacher_pool: logger.info( @@ -509,10 +509,15 @@ def main(cfg: Config) -> int: # So ref_eq=1.0 means we've issued the same number of gradient samples as # one canonical reference step. Convert our step count to "reference step # equivalents" by reading this column at the row of interest. - _row_cols = ["step", "ref_eq", "rew", "std", "sprd", "N", - "gt", "hack", "hack_s", "hack_t", "gt_s", + # Per-source split (student/teacher) for rew, gt, hack columns. Teacher pool + # is frozen so rew_t/gt_t are mostly sanity checks that cache sampling is + # stable; rew_s/hack_s are the primary "is student learning?" signals. + # `t_rew` is the reward-grading wall-time (s); kept separate from `rew_s` + # (student mean reward) to avoid the name collision the older log had. + _row_cols = ["step", "ref_eq", "rew", "rew_s", "std", "sprd", "N", + "gt", "gt_s", "gt_t", "hack", "hack_s", "hack_t", "loss", "cin", "cin_s", "cin_t", "cout", "fired", - "gen", "fb", "rew_s", "sec"] + "gen", "fb", "t_rew", "sec"] REF_GENS_PER_STEP = 16 * 16 # ariahw/rl-rewardhacking config.py:num_prompts * num_generations est_gens_per_step = cfg.prompts_per_step * cfg.group # before mixed-pool split logger.info( @@ -840,6 +845,8 @@ def main(cfg: Config) -> int: hack_s_n = int((h_t & is_s).sum()) hack_t_n = int((h_t & ~is_s).sum()) gt_s_n = int((g_t & is_s).sum()) + gt_t_n = int((g_t & ~is_s).sum()) + rew_s_mean = rewards_t[is_s].mean().item() if n_s else float("nan") # Per-step diagnostics → verbose log; stdout sees tqdm postfix + final table. n_fin = sum(agg_finished) @@ -869,14 +876,16 @@ def main(cfg: Config) -> int: "step": step, "ref_eq": f"{cum_gens / REF_GENS_PER_STEP:.2f}", "rew": f"{rew_mean:+.2f}", + "rew_s": f"{rew_s_mean:+.2f}" if n_s else "nan", "std": f"{rew_std:.2f}", "sprd": "T" if spread else "F", "N": n_rollouts, "gt": f"{sum(agg_gt)}/{n_rollouts}", + "gt_s": f"{gt_s_n}/{n_s}" if n_s else "0/0", + "gt_t": f"{gt_t_n}/{n_t}" if n_t else "0/0", "hack": f"{sum(agg_hack)}/{n_rollouts}", "hack_s": f"{hack_s_n}/{n_s}" if n_s else "0/0", "hack_t": f"{hack_t_n}/{n_t}" if n_t else "0/0", - "gt_s": f"{gt_s_n}/{n_s}" if n_s else "0/0", "loss": f"{agg_loss:+.4f}", "cin": f"{diag['mean_cos_in']:+.3f}", "cin_s": f"{diag['mean_cin_s']:+.3f}", @@ -885,7 +894,7 @@ def main(cfg: Config) -> int: "fired": f"{diag['frac_fired']:.2f}", "gen": f"{t_gen:.0f}", "fb": f"{t_fb:.0f}", - "rew_s": f"{t_rew:.0f}", + "t_rew": f"{t_rew:.0f}", "sec": f"{time.time()-t0:.0f}", } rows.append(row)