diff options
Diffstat (limited to 'worldalign/manifold_gate.py')
| -rw-r--r-- | worldalign/manifold_gate.py | 677 |
1 files changed, 677 insertions, 0 deletions
diff --git a/worldalign/manifold_gate.py b/worldalign/manifold_gate.py new file mode 100644 index 0000000..30acceb --- /dev/null +++ b/worldalign/manifold_gate.py @@ -0,0 +1,677 @@ +"""On-manifold identifiability gate for assignment energies. + +The configuration space is restricted to permutations of real frozen text +states. On this space the language-only energy terms (sliced distribution, +prototype manifold) depend only on the set of states and are therefore +constant; the only varying terms are the cross-modal relation MSE and the +multiscale conditional KL from ``energy.relation_field_energy``. Hidden +pairs are used only to score orderings, never inside the energy. + +Gates, in increasing strictness: + +A. global ranking: energy of the true assignment against random and + structured permutations; +B. local identifiability: exact delta energy of every transposition of the + true assignment, via a closed form that one matrix product evaluates for + all N(N-1)/2 swaps; +C. basin audit: exact steepest 2-swap descent from the true assignment and + from random assignments, with a local-minimum certificate. Descent from + random assignments doubles as a blind transductive recovery baseline and + as a search for on-manifold counterfeits. +""" + +from __future__ import annotations + +import argparse +import json +from pathlib import Path + +import torch +import torch.nn.functional as F + +from .common import read_json, seed_everything, write_json +from .io import load_feature_pair, select_rows + +TEMPERATURES = (0.03, 0.07, 0.15) +M30_RELATION_WEIGHT = 2.0 +M30_CONDITIONAL_WEIGHT = 0.2 + + +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", + help="Flickr language node state definition.", + ) + 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", + help="Private pairing, loaded only to construct the evaluation order.", + ) + parser.add_argument( + "--vg-bundle-channels", + action="store_true", + help="Add std/q10/q90 view-pair relation channels to the scalar mean.", + ) + 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("--random-perms", type=int, default=1000) + parser.add_argument("--derangement-samples", type=int, default=200) + parser.add_argument("--descent-restarts", type=int, default=3) + parser.add_argument("--descent-max-steps", type=int, default=200000) + parser.add_argument( + "--descent-objective", + choices=["mse", "m30_total"], + default="m30_total", + help="m30_total preselects swaps by closed-form MSE and verifies the " + "exact weighted MSE+KL objective on the best candidates.", + ) + 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/gate.json") + parser.add_argument("--trajectory-output") + return parser.parse_args() + + +def cosine_relation(features: torch.Tensor) -> torch.Tensor: + features = F.normalize(features.double(), dim=-1) + return features @ features.T + + +def offdiag_mask(size: int, device: torch.device | str) -> torch.Tensor: + return ~torch.eye(size, dtype=torch.bool, device=device) + + +def standardize_relation(relation: torch.Tensor) -> tuple[torch.Tensor, float, float]: + """Standardized copy with zeroed diagonal. + + The mean and std are taken over off-diagonal values, which are a + permutation-invariant set, so the same constants apply to every + assignment of the same states. + """ + mask = offdiag_mask(len(relation), relation.device) + values = relation[mask] + mean = values.mean() + std = values.std().clamp_min(1e-6) + standardized = (relation - mean) / std + standardized = standardized.masked_fill(~mask, 0.0) + return standardized, float(mean), float(std) + + +def relation_mse( + text_standardized: torch.Tensor, visual_standardized: torch.Tensor +) -> torch.Tensor: + mask = offdiag_mask(len(visual_standardized), visual_standardized.device) + return (text_standardized[mask] - visual_standardized[mask]).square().mean() + + +def conditional_kl( + text_relation: torch.Tensor, visual_relation: torch.Tensor +) -> torch.Tensor: + """Multiscale conditional KL, identical to energy.relation_field_energy.""" + diagonal = torch.eye( + len(visual_relation), dtype=torch.bool, device=visual_relation.device + ) + total = visual_relation.new_zeros(()) + for temperature in TEMPERATURES: + visual_logits = (visual_relation / temperature).masked_fill(diagonal, -1e4) + text_logits = (text_relation / temperature).masked_fill(diagonal, -1e4) + visual_probability = F.softmax(visual_logits, dim=-1) + total = total + ( + visual_probability + * ( + F.log_softmax(visual_logits, dim=-1) + - F.log_softmax(text_logits, dim=-1) + ) + ).sum(-1).mean() + return total + + +def permuted(relation: torch.Tensor, permutation: torch.Tensor) -> torch.Tensor: + return relation[permutation][:, permutation] + + +def assignment_energy( + text_channels: torch.Tensor, + visual_channels: torch.Tensor, + text_relation: torch.Tensor, + visual_relation: torch.Tensor, + permutation: torch.Tensor, +) -> dict[str, float]: + """Exact energy of one assignment. Channel 0 is the scalar relation.""" + mse_channels = [ + float(relation_mse(permuted(text_channels[c], permutation), visual_channels[c])) + for c in range(len(text_channels)) + ] + kl = float(conditional_kl(permuted(text_relation, permutation), visual_relation)) + mse = mse_channels[0] + return { + "mse": mse, + "mse_channels": mse_channels, + "mse_channel_mean": sum(mse_channels) / len(mse_channels), + "conditional_kl": kl, + "m30_total": M30_RELATION_WEIGHT * mse + M30_CONDITIONAL_WEIGHT * kl, + } + + +def all_transposition_delta_mse( + text_standardized: torch.Tensor, visual_standardized: torch.Tensor +) -> torch.Tensor: + """Exact MSE change for every transposition of the current assignment. + + Swapping nodes p and q changes rows/columns p and q of the permuted text + relation. In the squared error the quadratic text terms cancel, leaving + delta(p, q) = (4 / M) * sum_{k not in {p, q}} + (T_pk - T_qk)(V_pk - V_qk), + with M the off-diagonal count and both matrices standardized with zeroed + diagonals. One matrix product evaluates the sum for all pairs. + """ + size = len(text_standardized) + count = size * (size - 1) + cross = text_standardized @ visual_standardized # (T V)_pq + self_terms = (text_standardized * visual_standardized).sum(-1) # s_i + corrections = 2.0 * text_standardized * visual_standardized # k in {p, q} + total = self_terms[:, None] + self_terms[None, :] - cross - cross.T - corrections + delta = (4.0 / count) * total + delta.fill_diagonal_(0.0) + return delta + + +def sum_channel_delta( + text_channels: torch.Tensor, visual_channels: torch.Tensor +) -> torch.Tensor: + delta = all_transposition_delta_mse(text_channels[0], visual_channels[0]) + for c in range(1, len(text_channels)): + delta = delta + all_transposition_delta_mse( + text_channels[c], visual_channels[c] + ) + return delta / len(text_channels) + + +def random_permutations( + count: int, size: int, generator: torch.Generator +) -> torch.Tensor: + return torch.argsort(torch.rand(count, size, generator=generator), dim=-1) + + +def k_derangement( + size: int, k: int, generator: torch.Generator +) -> torch.Tensor: + """Identity with a random cyclic derangement on k random positions.""" + permutation = torch.arange(size) + chosen = torch.randperm(size, generator=generator)[:k] + permutation[chosen] = chosen.roll(1) + return permutation + + +def gate_a_global_ranking( + text_channels: torch.Tensor, + visual_channels: torch.Tensor, + text_relation: torch.Tensor, + visual_relation: torch.Tensor, + args: argparse.Namespace, + generator: torch.Generator, +) -> dict: + size = len(visual_relation) + identity = torch.arange(size) + true_energy = assignment_energy( + text_channels, visual_channels, text_relation, visual_relation, identity + ) + keys = ("mse", "mse_channel_mean", "conditional_kl", "m30_total") + samples: dict[str, list[float]] = {key: [] for key in keys} + for index in range(args.random_perms): + permutation = random_permutations(1, size, generator)[0] + energy = assignment_energy( + text_channels, visual_channels, text_relation, visual_relation, permutation + ) + for key in keys: + samples[key].append(energy[key]) + # Structured negative: cyclic shift along the text-similarity order, a + # systematic misassignment that preserves neighborhood smoothness. + order = text_relation.sum(-1).argsort() + shift = torch.empty_like(order) + shift[order] = order.roll(1) + shifted_energy = assignment_energy( + text_channels, visual_channels, text_relation, visual_relation, shift + ) + report: dict = { + "true": true_energy, + "similarity_shift": shifted_energy, + "random": {}, + } + for key in keys: + values = torch.tensor(samples[key]) + z = (values.mean() - true_energy[key]) / values.std().clamp_min(1e-12) + rank = int((values <= true_energy[key]).sum()) + report["random"][key] = { + "mean": float(values.mean()), + "std": float(values.std()), + "min": float(values.min()), + "true_z": float(z), + "true_rank_among_random": rank, + "count": args.random_perms, + } + return report + + +def gate_b_transpositions( + text_channels: torch.Tensor, + visual_channels: torch.Tensor, + text_relation: torch.Tensor, + visual_relation: torch.Tensor, + captions: list[str] | None, +) -> dict: + size = len(visual_relation) + delta = sum_channel_delta(text_channels, visual_channels) + upper = torch.triu(torch.ones(size, size, dtype=torch.bool), diagonal=1) + values = delta[upper] + improving = values < 0 + report: dict = { + "pairs": int(values.numel()), + "improving_pairs": int(improving.sum()), + "improving_fraction": float(improving.double().mean()), + "delta_mean": float(values.mean()), + "delta_min": float(values.min()), + "identity_is_local_min_mse": bool(improving.sum() == 0), + } + if improving.any(): + flat = delta.masked_fill(~upper, float("inf")).flatten() + worst = flat.argsort()[:20] + offenders = [] + for index in worst.tolist(): + p, q = divmod(index, size) + if flat[index] == float("inf"): + break + exact = { + "pair": [p, q], + "delta_mse": float(delta[p, q]), + "text_cosine": float(text_relation[p, q]), + "visual_cosine": float(visual_relation[p, q]), + } + if captions is not None: + exact["captions"] = [captions[p][:90], captions[q][:90]] + offenders.append(exact) + report["worst_improving_swaps"] = offenders + return report + + +def derangement_curve( + text_channels: torch.Tensor, + visual_channels: torch.Tensor, + args: argparse.Namespace, + generator: torch.Generator, +) -> list[dict]: + size = len(text_channels[0]) + identity_mse = float( + relation_mse(text_channels[0], visual_channels[0]) + ) + curve = [] + k = 2 + while k <= size: + deltas = [] + for _ in range(args.derangement_samples): + permutation = k_derangement(size, k, generator) + mse = float( + relation_mse( + permuted(text_channels[0], permutation), visual_channels[0] + ) + ) + deltas.append(mse - identity_mse) + values = torch.tensor(deltas) + curve.append( + { + "k": k, + "delta_mean": float(values.mean()), + "delta_std": float(values.std()), + "improving_fraction": float((values < 0).double().mean()), + } + ) + k *= 2 + return curve + + +def steepest_descent( + text_channels: torch.Tensor, + visual_channels: torch.Tensor, + text_relation: torch.Tensor, + visual_relation: torch.Tensor, + start: torch.Tensor, + args: argparse.Namespace, +) -> dict: + """Exact steepest 2-swap descent with a local-minimum certificate. + + Every step evaluates the closed-form MSE delta of all transpositions of + the current assignment. With the m30_total objective the best candidates + by MSE delta are re-scored with the exact weighted MSE+KL objective, so + an accepted move always lowers the reported objective. + """ + permutation = start.clone() + identity = torch.arange(len(start)) + trajectory = [] + + def objective(perm: torch.Tensor) -> float: + energy = assignment_energy( + text_channels, visual_channels, text_relation, visual_relation, perm + ) + return energy["m30_total" if args.descent_objective == "m30_total" else "mse"] + + current = objective(permutation) + accepted_moves = 0 + for step in range(args.descent_max_steps): + perm_text = torch.stack( + [permuted(channel, permutation) for channel in text_channels] + ) + delta = sum_channel_delta(perm_text, visual_channels) + upper = torch.triu(torch.ones_like(delta, dtype=torch.bool), diagonal=1) + masked = delta.masked_fill(~upper, float("inf")) + if args.descent_objective == "mse": + best = masked.flatten().argmin() + p, q = divmod(int(best), len(permutation)) + if masked[p, q] >= 0: + break + permutation[[p, q]] = permutation[[q, p]] + current = objective(permutation) + accepted_moves += 1 + else: + candidates = masked.flatten().argsort()[: args.descent_verify_top] + accepted = False + for index in candidates.tolist(): + p, q = divmod(index, len(permutation)) + if masked[p, q] == float("inf"): + break + trial = permutation.clone() + trial[[p, q]] = trial[[q, p]] + value = objective(trial) + if value < current - 1e-12: + permutation = trial + current = value + accepted = True + accepted_moves += 1 + break + if not accepted: + break + if step % 50 == 0: + trajectory.append( + { + "step": step, + "objective": current, + "accuracy": float((permutation == identity).double().mean()), + } + ) + final_delta = sum_channel_delta( + torch.stack([permuted(channel, permutation) for channel in text_channels]), + visual_channels, + ) + upper = torch.triu(torch.ones_like(final_delta, dtype=torch.bool), diagonal=1) + certificate = bool((final_delta[upper] >= 0).all()) + return { + "start_accuracy": float((start == identity).double().mean()), + "final_accuracy": float((permutation == identity).double().mean()), + "final_objective": current, + "final_energy": assignment_energy( + text_channels, visual_channels, text_relation, visual_relation, permutation + ), + "accepted_moves": accepted_moves, + "moved_fraction": float((permutation != start).double().mean()), + "mse_local_min_certificate": certificate, + "trajectory": trajectory, + "final_permutation": permutation.tolist(), + } + + +def load_flickr(args: argparse.Namespace) -> dict: + manifest = read_json(args.manifest) + vision, text, vision_lookup, text_lookup = load_feature_pair( + args.vision, args.text + ) + rows = manifest[args.split][: args.samples] + visual_states = select_rows(vision["features"], vision_lookup, rows) + captions = None + if args.text_mode == "single": + text_states = select_rows(text["features"], text_lookup, rows) + captions = [ + text["captions"][text_lookup[int(row)]] for row in rows + ] + else: + state = torch.load(args.text_orbits, map_location="cpu", weights_only=False) + lookup = {int(row): index for index, row in enumerate(state["rows"])} + orbit_mean = F.normalize(state["features"].float().mean(1), dim=-1) + text_states = select_rows(orbit_mean, lookup, rows) + captions = [state["captions"][lookup[int(row)]][0] for row in rows] + return { + "visual_views": visual_states[:, None, :], + "text_views": text_states[:, None, :], + "captions": captions, + "meta": { + "dataset": "flickr30k", + "split": args.split, + "samples": len(rows), + "text_mode": args.text_mode, + "rows": rows, + }, + } + + +def load_vg(args: argparse.Namespace) -> dict: + vision = torch.load(args.vg_vision, map_location="cpu", weights_only=False) + text = torch.load(args.vg_text, map_location="cpu", weights_only=False) + pairs = [ + json.loads(line) + for line in Path(args.vg_ground_truth).read_text().splitlines() + if line.strip() + ] + vision_index = {node: i for i, node in enumerate(vision["node_ids"])} + text_index = {node: i for i, node in enumerate(text["node_ids"])} + vision_order = [vision_index[pair["vision_node_id"]] for pair in pairs] + text_order = [text_index[pair["text_node_id"]] for pair in pairs] + visual_views = F.normalize(vision["region_features"].float(), dim=-1)[vision_order] + text_views = F.normalize(text["region_features"].float(), dim=-1)[text_order] + if args.samples and args.samples < len(visual_views): + generator = torch.Generator().manual_seed(args.subset_seed) + subset = torch.randperm(len(visual_views), generator=generator)[: args.samples] + visual_views = visual_views[subset] + text_views = text_views[subset] + return { + "visual_views": visual_views, + "text_views": text_views, + "captions": None, + "meta": { + "dataset": "visual_genome_5k", + "tier": text.get("tier"), + "samples": len(visual_views), + "subset_seed": args.subset_seed, + "bundle_channels": bool(args.vg_bundle_channels), + }, + } + + +def view_bundle_channels(views: torch.Tensor) -> torch.Tensor: + """Distribution-valued relation field from per-node view sets. + + Channel order: mean, std, q10, q90 of the view-pair cosine distribution + between two nodes. The scalar mean channel equals the relation of the + (unnormalized) view-mean embeddings; the remaining channels carry + information a single pooled vector cannot. + """ + nodes, view_count, _ = views.shape + views = views.double() + pair_cosines = torch.einsum("aud,bvd->abuv", views, views).reshape( + nodes, nodes, view_count * view_count + ) + mean = pair_cosines.mean(-1) + std = pair_cosines.std(-1) + q10 = pair_cosines.quantile(0.10, dim=-1) + q90 = pair_cosines.quantile(0.90, dim=-1) + return torch.stack([mean, std, q10, q90]) + + +def build_channels( + views: torch.Tensor, bundle: bool +) -> tuple[torch.Tensor, torch.Tensor]: + """Standardized relation channels and the raw scalar relation.""" + node_states = F.normalize(views.double().mean(1), dim=-1) + scalar = node_states @ node_states.T + if bundle and views.shape[1] > 1: + raw = view_bundle_channels(views) + else: + raw = scalar[None] + channels = [] + for c in range(len(raw)): + standardized, _, _ = standardize_relation(raw[c]) + channels.append(standardized) + return torch.stack(channels), scalar + + +def main() -> None: + args = parse_args() + seed_everything(args.seed) + data = load_flickr(args) if args.dataset == "flickr" else load_vg(args) + device = torch.device(args.device) + bundle = args.dataset == "vg" and args.vg_bundle_channels + text_channels, text_relation = build_channels( + data["text_views"].to(device), bundle + ) + visual_channels, visual_relation = build_channels( + data["visual_views"].to(device), bundle + ) + generator = torch.Generator().manual_seed(args.seed) + + report: dict = { + "protocol": ( + "Assignments permute real frozen text states; hidden pairs are " + "used only to place the true assignment in the ranking. The " + "energy terms are the cross-modal relation MSE and conditional " + "KL; the language-only terms of the falsified free-particle " + "energy are permutation-invariant on this space." + ), + "meta": data["meta"], + "args": { + key: value + for key, value in vars(args).items() + if key not in ("manifest", "vision", "text") + }, + "channel_names": ( + ["mean", "std", "q10", "q90"] if bundle else ["mean"] + ), + } + + report["gate_a_global_ranking"] = gate_a_global_ranking( + text_channels, visual_channels, text_relation, visual_relation, args, generator + ) + print(json.dumps({"gate_a": report["gate_a_global_ranking"]["random"]})) + + report["gate_b_transpositions"] = gate_b_transpositions( + text_channels, visual_channels, text_relation, visual_relation, data["captions"] + ) + print( + json.dumps( + { + "gate_b": { + key: value + for key, value in report["gate_b_transpositions"].items() + if key != "worst_improving_swaps" + } + } + ) + ) + + report["derangement_curve"] = derangement_curve( + text_channels, visual_channels, args, generator + ) + + identity = torch.arange(len(visual_relation)) + report["gate_c_descent_from_true"] = steepest_descent( + text_channels, + visual_channels, + text_relation, + visual_relation, + identity, + args, + ) + print( + json.dumps( + { + "gate_c_from_true": { + key: value + for key, value in report["gate_c_descent_from_true"].items() + if key not in ("trajectory", "final_permutation") + } + } + ) + ) + + restarts = [] + for restart in range(args.descent_restarts): + start = random_permutations(1, len(visual_relation), generator)[0] + result = steepest_descent( + text_channels, + visual_channels, + text_relation, + visual_relation, + start, + args, + ) + result.pop("final_permutation") + restarts.append(result) + print( + json.dumps( + { + "gate_c_from_random": { + "restart": restart, + "final_objective": result["final_objective"], + "final_accuracy": result["final_accuracy"], + } + } + ) + ) + report["gate_c_descent_from_random"] = restarts + + true_total = report["gate_a_global_ranking"]["true"]["m30_total"] + counterfeit = [ + restart + for restart in restarts + if restart["final_energy"]["m30_total"] < true_total + and restart["final_accuracy"] < 0.5 + ] + report["verdict"] = { + "true_m30_total": true_total, + "identity_is_local_min_mse": report["gate_b_transpositions"][ + "identity_is_local_min_mse" + ], + "descent_from_true_stays": report["gate_c_descent_from_true"][ + "final_accuracy" + ], + "on_manifold_counterfeit_found": bool(counterfeit), + "best_random_descent_m30_total": min( + (restart["final_energy"]["m30_total"] for restart in restarts), + default=None, + ), + } + print(json.dumps({"verdict": report["verdict"]})) + + Path(args.output).parent.mkdir(parents=True, exist_ok=True) + write_json(args.output, report) + if args.trajectory_output: + torch.save( + { + "from_true": report["gate_c_descent_from_true"], + "meta": data["meta"], + }, + args.trajectory_output, + ) + print(f"Wrote {args.output}") + + +if __name__ == "__main__": + main() |
