summaryrefslogtreecommitdiff
path: root/worldalign/rank_gate.py
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-08-01 16:23:26 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-08-01 16:23:26 -0500
commit66af02f1e5c50e041ced3e8b85659709a1c1f4f3 (patch)
tree6dd0de68e644be5e7a5546123f44ea3c4a63646f /worldalign/rank_gate.py
parent735f9c7fd202d0eaed9183094d84d365e0e5404d (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.py138
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()