diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-08-01 16:23:26 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-08-01 16:23:26 -0500 |
| commit | 66af02f1e5c50e041ced3e8b85659709a1c1f4f3 (patch) | |
| tree | 6dd0de68e644be5e7a5546123f44ea3c4a63646f /worldalign/rank_gate.py | |
| parent | 735f9c7fd202d0eaed9183094d84d365e0e5404d (diff) | |
Record the gate correction in the concept document
The user-facing statement still carried the retired correlation
threshold. Adds the correction, the joint condition that replaces it,
the width diagnosis, and the closure of the shrink-N route.
Co-Authored-By: Claude <noreply@anthropic.com>
Diffstat (limited to 'worldalign/rank_gate.py')
| -rw-r--r-- | worldalign/rank_gate.py | 138 |
1 files changed, 138 insertions, 0 deletions
diff --git a/worldalign/rank_gate.py b/worldalign/rank_gate.py new file mode 100644 index 0000000..d55f12e --- /dev/null +++ b/worldalign/rank_gate.py @@ -0,0 +1,138 @@ +"""Is low-rank failure an information limit or a solver limit? + +The rank ladder is the result that retires the correlation gate, so it has to +survive the objection the project has fallen for three times before: a limit +that looks intrinsic and turns out to belong to the search operator. At rank +eight the composed solver reaches 13%, and that alone does not say whether the +truth is unreachable or merely unreached. + +The gate answers it. If the deepest state a strong searcher finds is deeper +than the truth, the truth is not the optimum and no solver of any cost +recovers it -- the width really is an information limit. If the truth is +deepest and the solver still misses it, the ladder measures search difficulty +instead and the conclusion has to be rewritten. +""" + +from __future__ import annotations + +import argparse +import json + +import numpy as np +import torch + +from .common import write_json +from .spectral_match import grampa +from .synth_fast_gate import ClosedFormEnergy, all_swaps, steepest_descent +from .synth_triangle_gate import standardized + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--fields", default="artifacts/synth_v1/fields_tier0_ws_256.pt") + parser.add_argument("--ranks", type=int, nargs="+", default=[4, 8, 16, 256]) + parser.add_argument("--restarts", type=int, default=60) + parser.add_argument("--trials", type=int, default=3) + parser.add_argument("--device", default="cuda:3") + parser.add_argument("--output", default="artifacts/synth_v1/rank_gate.json") + return parser.parse_args() + + +def offdiagonal(matrix: np.ndarray) -> np.ndarray: + return matrix[~np.eye(len(matrix), dtype=bool)] + + +def normalise(matrix: np.ndarray) -> np.ndarray: + values = offdiagonal(matrix) + out = (matrix - values.mean()) / values.std() + np.fill_diagonal(out, 0.0) + return out + + +def truncate(matrix: np.ndarray, rank: int) -> np.ndarray: + symmetric = (matrix + matrix.T) / 2.0 + values, vectors = np.linalg.eigh(symmetric) + order = np.argsort(np.abs(values))[::-1][:rank] + return (vectors[:, order] * values[order]) @ vectors[:, order].T + + +def main() -> None: + args = parse_args() + state = torch.load(args.fields, map_location="cpu", weights_only=False) + visual_full = state["visual_field"].double().numpy() + text_full = state["text_field"].double().numpy() + size = len(visual_full) + device = torch.device(args.device) + swaps = all_swaps(size, device) + + rows = [] + for rank in args.ranks: + visual = normalise(truncate(visual_full, rank) if rank < size else visual_full) + text = normalise(truncate(text_full, rank) if rank < size else text_full) + verdicts, gaps, accuracies = [], [], [] + for trial in range(args.trials): + generator = np.random.default_rng(trial) + hidden = generator.permutation(size) + shuffled = text[np.ix_(hidden, hidden)] + energy = ClosedFormEnergy( + standardized(torch.from_numpy(shuffled).to(device)).float(), + standardized(torch.from_numpy(visual).to(device)).float(), + 1.0, 1.0, 256, + ) + truth = torch.from_numpy(np.argsort(hidden).copy()).to(device) + truth_energy = float(energy.energy(truth[None])[0]) + + starts = [torch.from_numpy(grampa(visual, shuffled, 1.0).copy()).to(device)] + starts += [ + torch.from_numpy(generator.permutation(size).copy()).to(device) + for _ in range(args.restarts) + ] + best_energy, best_accuracy = np.inf, 0.0 + for start in starts: + final, _ = steepest_descent(energy, start, swaps, 4000) + value = float(energy.energy(final[None])[0]) + if value < best_energy: + best_energy = value + best_accuracy = float( + (hidden[final.cpu().numpy()] == np.arange(size)).mean() + ) + verdicts.append(truth_energy <= best_energy + 1e-6) + gaps.append(best_energy - truth_energy) + accuracies.append(best_accuracy) + + row = { + "rank": rank, + "truth_is_deepest": f"{sum(verdicts)}/{args.trials}", + "mean_energy_gap_best_minus_truth": float(np.mean(gaps)), + "best_accuracy": float(np.mean(accuracies)), + "reading": ( + "information limit" if sum(verdicts) == 0 else + "truth is optimal; failure is search" if sum(verdicts) == args.trials + else "mixed" + ), + } + rows.append(row) + print( + f"rank={rank:<5} truth deepest {row['truth_is_deepest']} " + f"gap(best-truth)={row['mean_energy_gap_best_minus_truth']:+.4f} " + f"best acc={row['best_accuracy']:.3f} -> {row['reading']}", + flush=True, + ) + + summary = { + "protocol": ( + "Spectral start plus random restarts, each run to a local optimum " + "under exact steepest descent. A negative gap means the searcher " + "found a state deeper than the truth, so the truth is not the " + "optimum and the limit is information rather than search." + ), + "size": size, + "restarts": args.restarts, + "rows": rows, + } + print(json.dumps({"done": True})) + write_json(args.output, summary) + + +if __name__ == "__main__": + main() |
