summaryrefslogtreecommitdiff
path: root/worldalign/ricci_control.py
diff options
context:
space:
mode:
Diffstat (limited to 'worldalign/ricci_control.py')
-rw-r--r--worldalign/ricci_control.py258
1 files changed, 258 insertions, 0 deletions
diff --git a/worldalign/ricci_control.py b/worldalign/ricci_control.py
new file mode 100644
index 0000000..3091b22
--- /dev/null
+++ b/worldalign/ricci_control.py
@@ -0,0 +1,258 @@
+"""Ricci-flow control for the on-manifold assignment gate.
+
+Hypothesis under test: a discrete Ricci flow that smooths each modality's
+relational geometry before matching could repair the local ordering that
+static relation fields fail. The flow is run independently per modality
+with identical hyperparameters; hidden pairs never touch the flow or the
+energy and only score orderings, as in the base gate.
+
+Two geometries are gated per condition: heat-kernel channels on the raw
+kNN graph (diffusion without flow) and the same construction after
+Ollivier-Ricci weight evolution. Differences between them are attributable
+to the flow itself rather than to the diffusion representation.
+"""
+
+from __future__ import annotations
+
+import argparse
+import json
+from pathlib import Path
+
+import numpy as np
+import ot
+import torch
+from scipy.sparse import csr_matrix
+from scipy.sparse.csgraph import shortest_path
+
+from .common import seed_everything, write_json
+from .manifold_gate import (
+ gate_a_global_ranking,
+ gate_b_transpositions,
+ load_flickr,
+ load_vg,
+ standardize_relation,
+ steepest_descent,
+)
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--dataset", choices=["flickr", "vg"], default="flickr")
+ parser.add_argument("--manifest", default="artifacts/manifest.json")
+ parser.add_argument("--vision", default="artifacts/vision.pt")
+ parser.add_argument("--text", default="artifacts/text.pt")
+ parser.add_argument("--text-orbits", default="artifacts/text_orbits_qwen0p5b.pt")
+ parser.add_argument("--text-mode", choices=["single", "orbit_mean"], default="orbit_mean")
+ parser.add_argument("--vg-vision", default="artifacts/vg_5k/vision_features.pt")
+ parser.add_argument("--vg-text", default="artifacts/vg_5k/text_features.pt")
+ parser.add_argument(
+ "--vg-ground-truth", default="artifacts/vg_5k/ground_truth.private.jsonl"
+ )
+ parser.add_argument("--vg-bundle-channels", action="store_true")
+ parser.add_argument("--split", choices=["val", "test"], default="test")
+ parser.add_argument("--samples", type=int, default=512)
+ parser.add_argument("--subset-seed", type=int, default=0)
+ parser.add_argument("--neighbors", type=int, default=16)
+ parser.add_argument("--flow-iterations", type=int, default=10)
+ parser.add_argument("--flow-step", type=float, default=0.4)
+ parser.add_argument("--lazy-alpha", type=float, default=0.5)
+ parser.add_argument("--heat-times", default="1.0,4.0")
+ parser.add_argument("--random-perms", type=int, default=300)
+ parser.add_argument("--derangement-samples", type=int, default=50)
+ parser.add_argument("--descent-restarts", type=int, default=2)
+ parser.add_argument("--descent-max-steps", type=int, default=200000)
+ parser.add_argument("--descent-objective", default="mse", choices=["mse", "m30_total"])
+ parser.add_argument("--descent-verify-top", type=int, default=64)
+ parser.add_argument("--device", default="cpu")
+ parser.add_argument("--seed", type=int, default=20260729)
+ parser.add_argument("--output", default="artifacts/manifold_gate/ricci_control.json")
+ parser.add_argument("--trajectory-output")
+ return parser.parse_args()
+
+
+def knn_edges(distance: np.ndarray, k: int) -> list[tuple[int, int]]:
+ order = distance.argsort(-1)[:, 1 : k + 1]
+ edges = {
+ (min(i, int(j)), max(i, int(j)))
+ for i in range(len(distance))
+ for j in order[i]
+ }
+ return sorted(edges)
+
+
+def graph_apsp(size: int, weights: dict[tuple[int, int], float]) -> np.ndarray:
+ rows, cols, vals = [], [], []
+ for (i, j), w in weights.items():
+ rows.extend((i, j))
+ cols.extend((j, i))
+ vals.extend((w, w))
+ graph = csr_matrix((vals, (rows, cols)), shape=(size, size))
+ return shortest_path(graph, method="D", directed=False)
+
+
+def ollivier_ricci_apsp(
+ distance: np.ndarray, args: argparse.Namespace, iterations: int
+) -> np.ndarray:
+ """All-pairs geodesics after Ollivier-Ricci weight evolution.
+
+ Lazy uniform neighbor measures, W1 ground costs from current geodesics,
+ multiplicative weight update w <- w * (1 - step * kappa), total edge
+ mass renormalized each iteration. iterations=0 gives the un-flowed
+ graph geometry for the diffusion-only control.
+ """
+ edges = knn_edges(distance, args.neighbors)
+ weights = {edge: max(float(distance[edge]), 1e-9) for edge in edges}
+ total = sum(weights.values())
+ neighbor_map: dict[int, list[int]] = {}
+ for i, j in edges:
+ neighbor_map.setdefault(i, []).append(j)
+ neighbor_map.setdefault(j, []).append(i)
+ apsp = graph_apsp(len(distance), weights)
+ for _ in range(iterations):
+ updated: dict[tuple[int, int], float] = {}
+ for i, j in edges:
+ support_i = [i] + neighbor_map[i]
+ support_j = [j] + neighbor_map[j]
+ mass_i = np.full(len(support_i), (1 - args.lazy_alpha) / len(neighbor_map[i]))
+ mass_i[0] = args.lazy_alpha
+ mass_j = np.full(len(support_j), (1 - args.lazy_alpha) / len(neighbor_map[j]))
+ mass_j[0] = args.lazy_alpha
+ ground = apsp[np.ix_(support_i, support_j)]
+ if not np.isfinite(ground).all():
+ finite_max = apsp[np.isfinite(apsp)].max()
+ ground = np.where(np.isfinite(ground), ground, 2.0 * finite_max)
+ wasserstein = ot.emd2(mass_i, mass_j, ground)
+ geodesic = max(float(apsp[i, j]), 1e-9)
+ curvature = 1.0 - wasserstein / geodesic
+ updated[(i, j)] = max(1e-9, weights[(i, j)] * (1.0 - args.flow_step * curvature))
+ scale = total / sum(updated.values())
+ weights = {edge: w * scale for edge, w in updated.items()}
+ apsp = graph_apsp(len(distance), weights)
+ return apsp
+
+
+def geometry_channels(
+ apsp: np.ndarray, heat_times: tuple[float, ...]
+) -> torch.Tensor:
+ """Standardized relation channels of a flowed geometry.
+
+ Channel 0 is the negative geodesic field; the rest are heat kernels of
+ the normalized Laplacian of a geodesic-scale affinity.
+ """
+ finite = apsp[np.isfinite(apsp) & (apsp > 0)]
+ scale = np.median(finite)
+ capped = np.where(np.isfinite(apsp), apsp, finite.max() * 2.0)
+ affinity = np.exp(-capped / scale)
+ np.fill_diagonal(affinity, 1.0)
+ degree = affinity.sum(-1)
+ normalized = affinity / np.sqrt(degree[:, None] * degree[None, :])
+ values, vectors = np.linalg.eigh((normalized + normalized.T) / 2.0)
+ laplacian_eigen = 1.0 - values
+ channels = [torch.from_numpy(-capped).double()]
+ for t in heat_times:
+ heat = (vectors * np.exp(-t * laplacian_eigen)) @ vectors.T
+ channels.append(torch.from_numpy(heat).double())
+ return torch.stack(
+ [standardize_relation(channel)[0] for channel in channels]
+ )
+
+
+def run_gates(
+ text_channels: torch.Tensor,
+ visual_channels: torch.Tensor,
+ args: argparse.Namespace,
+ generator: torch.Generator,
+) -> dict:
+ text_relation = text_channels[0]
+ visual_relation = visual_channels[0]
+ report = {
+ "gate_a": gate_a_global_ranking(
+ text_channels, visual_channels, text_relation, visual_relation, args, generator
+ ),
+ "gate_b": gate_b_transpositions(
+ text_channels, visual_channels, text_relation, visual_relation, None
+ ),
+ "descent_from_true": steepest_descent(
+ text_channels,
+ visual_channels,
+ text_relation,
+ visual_relation,
+ torch.arange(len(visual_relation)),
+ args,
+ ),
+ }
+ report["descent_from_true"].pop("final_permutation", None)
+ report["descent_from_random"] = []
+ for _ in range(args.descent_restarts):
+ start = torch.argsort(torch.rand(len(visual_relation), generator=generator))
+ result = steepest_descent(
+ text_channels, visual_channels, text_relation, visual_relation, start, args
+ )
+ result.pop("final_permutation", None)
+ result.pop("trajectory", None)
+ report["descent_from_random"].append(result)
+ return report
+
+
+def main() -> None:
+ args = parse_args()
+ seed_everything(args.seed)
+ data = load_flickr(args) if args.dataset == "flickr" else load_vg(args)
+ heat_times = tuple(float(t) for t in args.heat_times.split(","))
+ states = {
+ "visual": torch.nn.functional.normalize(
+ data["visual_views"].double().mean(1), dim=-1
+ ),
+ "text": torch.nn.functional.normalize(
+ data["text_views"].double().mean(1), dim=-1
+ ),
+ }
+ distances = {
+ side: (1.0 - state @ state.T).clamp_min(0.0).numpy()
+ for side, state in states.items()
+ }
+ report: dict = {
+ "protocol": (
+ "Each modality's kNN geometry evolves independently under "
+ "Ollivier-Ricci flow; the identical heat-kernel channels are "
+ "gated with and without the flow. Hidden pairs score orderings "
+ "only."
+ ),
+ "meta": {**data["meta"], "flow": vars(args)},
+ "conditions": {},
+ }
+ for label, iterations in (
+ ("diffusion_no_flow", 0),
+ ("ricci_flow", args.flow_iterations),
+ ):
+ channels = {}
+ for side in ("visual", "text"):
+ apsp = ollivier_ricci_apsp(distances[side], args, iterations)
+ channels[side] = geometry_channels(apsp, heat_times)
+ generator = torch.Generator().manual_seed(args.seed)
+ result = run_gates(channels["text"], channels["visual"], args, generator)
+ report["conditions"][label] = result
+ print(
+ json.dumps(
+ {
+ label: {
+ "true_z_mse": result["gate_a"]["random"]["mse"]["true_z"],
+ "improving_fraction": result["gate_b"]["improving_fraction"],
+ "descent_keeps": result["descent_from_true"]["final_accuracy"],
+ "counterfeit_found": any(
+ r["final_objective"]
+ < result["gate_a"]["true"]["mse"]
+ and r["final_accuracy"] < 0.5
+ for r in result["descent_from_random"]
+ ),
+ }
+ }
+ )
+ )
+ Path(args.output).parent.mkdir(parents=True, exist_ok=True)
+ write_json(args.output, report)
+ print(f"Wrote {args.output}")
+
+
+if __name__ == "__main__":
+ main()