diff --git a/src/projected_grpo/probe_distill.py b/src/projected_grpo/probe_distill.py index cfb1c06..0bd8478 100644 --- a/src/projected_grpo/probe_distill.py +++ b/src/projected_grpo/probe_distill.py @@ -213,7 +213,7 @@ def main(cfg: Config) -> int: student, wrappers, tok = load_student(device) delta_params = [info["delta_S"] for info in wrappers.values()] logger.info(f"student delta_S params: {sum(p.numel() for p in delta_params):,}") - v_hack_cpu, _v_sv = load_v_hack(cfg.v_hack_path, STUDENT_MODEL, wrappers) + v_hack_cpu = load_v_hack(cfg.v_hack_path, STUDENT_MODEL, wrappers) v_hack = {n: v.to(device) for n, v in v_hack_cpu.items()} opt = torch.optim.AdamW(delta_params, lr=cfg.lr) diff --git a/src/projected_grpo/train.py b/src/projected_grpo/train.py index cb4d3fc..f2d6b15 100644 --- a/src/projected_grpo/train.py +++ b/src/projected_grpo/train.py @@ -223,18 +223,17 @@ def load_problems(n: int) -> list[dict]: def load_v_hack( path: Path, model_name: str, wrappers: dict, k_use: int | None = None, -) -> tuple[dict[str, torch.Tensor], dict[str, torch.Tensor]]: - """Load v_hack (top-k directions) + sv (singular values) for this wrapped model. +) -> dict[str, torch.Tensor]: + """Load v_hack (top-k directions) for this wrapped model. File schema (v2): bare `{name}` keys hold V[k_max, r]; `_sv/{name}` keys hold - S[k_max]. v_hack is model-specific because module names and per-module SVD - ranks depend on the exact checkpoint; a smoke (Qwen3.5-0.8B) v_hack must - not be reused for a full (Qwen3-4B) run. + S[k_max] (read but not returned — no caller uses them yet). v_hack is + model-specific because module names and per-module SVD ranks depend on the + exact checkpoint; a smoke (Qwen3.5-0.8B) v_hack must not be reused for a + full (Qwen3-4B) run. - If `k_use` is given, slices V and S to top-k_use rows. Errors if k_use > k_max + If `k_use` is given, slices V to top-k_use rows. Errors if k_use > k_max saved (re-extract with a higher top_k). - - Returns (v_hack, v_sv). v_sv may be empty for v1 files without _sv/ keys. """ with safe_open(str(path), framework="pt", device="cpu") as f: meta = f.metadata() or {} @@ -253,10 +252,9 @@ def load_v_hack( f"v_hack dtype/SVD-basis mismatch: {path} was extracted with dtype={saved_dtype}; " "train.py loads models in bf16. Re-extract with `--dtype=bf16`." ) - # Split V keys (bare module names) from S keys (prefixed _sv/). - all_tensors = {k: f.get_tensor(k) for k in f.keys()} - v_hack = {k: v for k, v in all_tensors.items() if not k.startswith("_sv/")} - v_sv = {k[len("_sv/"):]: v for k, v in all_tensors.items() if k.startswith("_sv/")} + # Read only V keys (bare module names); _sv/{name} keys are saved + # alongside but no runtime path consumes them currently. + v_hack = {k: f.get_tensor(k) for k in f.keys() if not k.startswith("_sv/")} wrapper_keys = set(wrappers) vhack_keys = set(v_hack) @@ -286,12 +284,10 @@ def load_v_hack( f"Re-extract with `--top-k={k_use}`." ) v_hack = {n: v[:k_use].contiguous() for n, v in v_hack.items()} - v_sv = {n: s[:k_use].contiguous() for n, s in v_sv.items()} logger.info( - f"loaded v_hack from {path}: modules={len(v_hack)}; k_saved={k_max}, k_use={k_use or k_max}; " - f"sv={'yes' if v_sv else 'no'}" + f"loaded v_hack from {path}: modules={len(v_hack)}; k_saved={k_max}, k_use={k_use or k_max}" ) - return v_hack, v_sv + return v_hack @torch.no_grad() @@ -390,7 +386,7 @@ def main(cfg: Config) -> int: "tau_axis": "0.0", "schema": "v2_with_sv"}) # extract zeros grads at exit; opt is built below so no opt-state taint. model.train() # restore train mode; eval was set only for the extract pass - v_hack_cpu, _v_sv_cpu = load_v_hack(v_hack_path, model_name, wrappers, k_use=cfg.v_hack_k) + v_hack_cpu = load_v_hack(v_hack_path, model_name, wrappers, k_use=cfg.v_hack_k) v_hack = {name: v.to(device) for name, v in v_hack_cpu.items()} # Teacher pool: pre-generated rollouts on disk keyed by problem_id. Each step's # G_t teacher rollouts come from a uniform random sample of that prompt's cache, @@ -806,19 +802,15 @@ def main(cfg: Config) -> int: # ratio is bounded in [0,1]). Discriminator: cin_t > cin_s on a clean # base means v_hack lights up for hack grads more than non-hack — the # extraction direction is doing real work. - cin_s = mean_cin_from_grads(step_grad_s, v_hack) if v_hack is not None else float("nan") - cin_t = mean_cin_from_grads(step_grad_t, v_hack) if v_hack is not None else float("nan") + cin_s = mean_cin_from_grads(step_grad_s, v_hack) + cin_t = mean_cin_from_grads(step_grad_t, v_hack) # Diagnostic cos_in for both arms; projection only mutates grad if arm=projected. - if v_hack is not None: - diag = project_delta_S_grad( - wrappers, v_hack, cfg.preserve_magnitude, - measure_only=(cfg.arm != "projected"), - gate_mode=cfg.gate_mode, - ) - else: - diag = {"mean_cos_in": float("nan"), "mean_cos_out": float("nan"), - "frac_fired": float("nan")} + diag = project_delta_S_grad( + wrappers, v_hack, cfg.preserve_magnitude, + measure_only=(cfg.arm != "projected"), + gate_mode=cfg.gate_mode, + ) diag["mean_cin_s"] = cin_s diag["mean_cin_t"] = cin_t diff --git a/src/projected_grpo/verify_vhack_heldout.py b/src/projected_grpo/verify_vhack_heldout.py index c5190bc..20d79c0 100644 --- a/src/projected_grpo/verify_vhack_heldout.py +++ b/src/projected_grpo/verify_vhack_heldout.py @@ -61,7 +61,7 @@ def main(cfg: Config) -> int: wrappers = wrap_model_with_antipasto( model, model_name=cfg.model, cache_root=CACHE_ROOT, svd_device=device, ) - v_hack, _v_sv = load_v_hack(cfg.v_hack_path, cfg.model, wrappers) + v_hack = load_v_hack(cfg.v_hack_path, cfg.model, wrappers) logger.info(f"loaded v_hack: {len(v_hack)} modules") grads_hack: dict[str, list[torch.Tensor]] = defaultdict(list)