diff options
| author | Yuren Hao <blackhao0426@gmail.com> | 2026-08-01 14:10:03 -0500 |
|---|---|---|
| committer | Yuren Hao <blackhao0426@gmail.com> | 2026-08-01 14:10:03 -0500 |
| commit | a62cf4d2a99b4a7985c61b2a7feb92a82a8218b7 (patch) | |
| tree | ee2248078db7edf3812a07f195afa3d9bd6f10c6 /worldalign | |
World Alignment: unpaired cross-modal correspondence by relational identifiability
Method: scene states are sets of part states; relation fields are built
within each modality and are invariant to how each side labels its own
features; the cross-modal bridge is a coupling searched under an energy
that is a closed-form functional of one matrix; solving is spectral
initialisation followed by exact local refinement.
Evidence: in a procedurally generated closed world, blind recovery of a
hidden image-caption correspondence reaches 95.3% at 256 scenes against
0.39% chance, and the recovered pairs transfer to 200 held-out scenes at
93.0% exact retrieval with random-pair and shuffled-image controls at or
near chance. Cross-modal value correspondence is derived from disjoint
corpora rather than declared. On Visual Genome the field correlation
reaches 0.656 against the 0.9 that polynomial recovery needs, with the
deficit attributed away from segmentation and discretisation.
Protocol: no image-text pair enters any objective, optimiser,
initialisation, or model selection; hidden pairs score orderings only.
Co-Authored-By: Claude <noreply@anthropic.com>
Diffstat (limited to 'worldalign')
61 files changed, 13364 insertions, 0 deletions
diff --git a/worldalign/__init__.py b/worldalign/__init__.py new file mode 100644 index 0000000..4013580 --- /dev/null +++ b/worldalign/__init__.py @@ -0,0 +1,4 @@ +"""Unpaired vision-language representation alignment experiments.""" + +__version__ = "0.1.0" + diff --git a/worldalign/anchor_battery.py b/worldalign/anchor_battery.py new file mode 100644 index 0000000..9594168 --- /dev/null +++ b/worldalign/anchor_battery.py @@ -0,0 +1,274 @@ +"""C5 battery: gauge-free unary anchors from raw observations. + +Two declared anchor classes, per `STRUCTURE_DESIGN.md`: + +1. pure invariants -- quantities whose cross-modal identity needs no + specification at all (numeral content of phrases); +2. minimally specified physical invariants -- fixed world-knowledge maps + (color lexicon to hue bands, size words to box areas, luminance words + to pixel luminance), frozen before evaluation and never tuned against + pairs. + +Anchors are unary and computed independently per modality: text from the +released phrase bundles, vision from cached image pixels and the region +boxes of the private preprocessing file (vision-side observables). Hidden +pairs score anchor quality; the anchor matrix itself is built without +them and feeds the tempering search as side information. +""" + +from __future__ import annotations + +import argparse +import json +import re +from concurrent.futures import ThreadPoolExecutor +from pathlib import Path + +import numpy as np +import torch +from PIL import Image + +from .common import write_json + +NUMBER_WORDS = { + "one": 1, "two": 2, "three": 3, "four": 4, "five": 5, "six": 6, + "seven": 7, "eight": 8, "nine": 9, "ten": 10, "eleven": 11, + "twelve": 12, "dozen": 12, +} +# Hue band centers in degrees on the HSV wheel; frozen world knowledge. +COLOR_BANDS = { + "red": 0.0, "orange": 30.0, "yellow": 60.0, "green": 120.0, + "cyan": 180.0, "blue": 220.0, "purple": 275.0, "pink": 330.0, +} +ACHROMATIC = ("white", "black", "gray", "grey", "brown", "tan", "beige") +SIZE_WORDS = { + "huge": 2.0, "large": 1.5, "big": 1.5, "tall": 1.0, "long": 1.0, + "small": -1.0, "little": -1.0, "tiny": -2.0, "short": -0.5, +} +LIGHT_WORDS = { + "white": 1.0, "bright": 1.0, "light": 0.5, "sunny": 1.0, + "black": -1.0, "dark": -1.0, "shadow": -0.5, "night": -1.0, +} + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--vg-dir", default="artifacts/vg_5k") + parser.add_argument("--image-cache", default="/tmp/yurenh2-worldalign-vg-images") + parser.add_argument("--image-size", type=int, default=224) + parser.add_argument("--workers", type=int, default=16) + parser.add_argument( + "--output", default="artifacts/manifold_gate/anchor_battery.json" + ) + parser.add_argument( + "--matrix-output", default="artifacts/manifold_gate/anchor_unary.pt" + ) + return parser.parse_args() + + +def read_jsonl(path: Path) -> list[dict]: + return [ + json.loads(line) + for line in path.read_text(encoding="utf-8").splitlines() + if line.strip() + ] + + +def text_anchors(phrases: list[str]) -> dict[str, np.ndarray | float]: + tokens = [ + token + for phrase in phrases + for token in re.findall(r"[a-z]+|\d+", phrase.lower()) + ] + numeral_total = 0.0 + for token in tokens: + if token.isdigit() and int(token) <= 50: + numeral_total += int(token) + elif token in NUMBER_WORDS: + numeral_total += NUMBER_WORDS[token] + color_histogram = np.zeros(len(COLOR_BANDS)) + for index, color in enumerate(COLOR_BANDS): + color_histogram[index] = sum(token == color for token in tokens) + achromatic = sum(token in ACHROMATIC for token in tokens) + chroma_score = color_histogram.sum() / max(color_histogram.sum() + achromatic, 1) + size_score = float(sum(SIZE_WORDS.get(token, 0.0) for token in tokens)) + light_score = float(sum(LIGHT_WORDS.get(token, 0.0) for token in tokens)) + return { + "numeral_total": float(numeral_total), + "color_histogram": color_histogram, + "chroma_score": float(chroma_score), + "size_score": size_score, + "light_score": light_score, + } + + +def vision_anchors(record: dict, args: argparse.Namespace) -> dict: + path = Path(args.image_cache, f"{record['source_image_id']}.jpg") + with Image.open(path) as image: + hsv = image.convert("HSV").resize( + (args.image_size, args.image_size), Image.BILINEAR + ) + values = np.asarray(hsv, dtype=np.float64) + hue = values[..., 0] * (360.0 / 255.0) + saturation = values[..., 1] / 255.0 + value = values[..., 2] / 255.0 + chromatic = (saturation > 0.25) & (value > 0.2) + hue_histogram = np.zeros(len(COLOR_BANDS)) + for index, center in enumerate(COLOR_BANDS.values()): + distance = np.minimum(np.abs(hue - center), 360.0 - np.abs(hue - center)) + hue_histogram[index] = float(((distance < 25.0) & chromatic).mean()) + width, height = record["width"], record["height"] + areas = [ + (region["width"] * region["height"]) / max(width * height, 1) + for region in record["regions"] + ] + return { + "hue_histogram": hue_histogram, + "chroma_fraction": float(chromatic.mean()), + "luminance_mean": float(value.mean()), + "mean_box_area": float(np.mean(areas)), + "log_mean_box_area": float(np.log(np.mean(areas) + 1e-6)), + } + + +def rank_z(x: np.ndarray) -> np.ndarray: + order = x.argsort().argsort().astype(np.float64) + order = (order - order.mean()) / max(order.std(), 1e-9) + return order + + +def spearman(x: np.ndarray, y: np.ndarray) -> float: + return float(np.corrcoef(rank_z(x), rank_z(y))[0, 1]) + + +def main() -> None: + args = parse_args() + text_nodes = read_jsonl(Path(args.vg_dir, "text_nodes.jsonl")) + vision_nodes = read_jsonl(Path(args.vg_dir, "vision_nodes.private.jsonl")) + truth = read_jsonl(Path(args.vg_dir, "ground_truth.private.jsonl")) + + text_features = { + record["node_id"]: text_anchors(record["region_closed"]) + for record in text_nodes + } + with ThreadPoolExecutor(max_workers=args.workers) as pool: + vision_results = list( + pool.map(lambda record: vision_anchors(record, args), vision_nodes) + ) + vision_features = { + record["node_id"]: result + for record, result in zip(vision_nodes, vision_results) + } + + # Aligned arrays in ground-truth order: index i pairs vision i, text i. + vision_ids = [pair["vision_node_id"] for pair in truth] + text_ids = [pair["text_node_id"] for pair in truth] + n = len(truth) + + def text_array(key: str) -> np.ndarray: + return np.array([text_features[node][key] for node in text_ids]) + + def vision_array(key: str) -> np.ndarray: + return np.array([vision_features[node][key] for node in vision_ids]) + + declared_pairs = { + "color": ("color_histogram", "hue_histogram"), + "chroma": ("chroma_score", "chroma_fraction"), + "light": ("light_score", "luminance_mean"), + "size": ("size_score", "log_mean_box_area"), + "numerals_vs_area": ("numeral_total", "log_mean_box_area"), + } + + report: dict = { + "protocol": ( + "Unary anchors, computed independently per modality; hidden " + "pairs score anchor quality only. Lexicon-to-physics maps are " + "frozen world knowledge, declared before evaluation." + ), + "nodes": n, + "anchors": {}, + } + channels: list[np.ndarray] = [] + channel_names: list[str] = [] + + # Color: cosine between lexicon histogram and hue-band histogram under + # the frozen band map, evaluated for every cross pair. + text_color = np.stack([text_features[node]["color_histogram"] for node in text_ids]) + vision_color = np.stack( + [vision_features[node]["hue_histogram"] for node in vision_ids] + ) + text_color_norm = text_color / np.linalg.norm(text_color, axis=1, keepdims=True).clip( + min=1e-9 + ) + vision_color_norm = vision_color / np.linalg.norm( + vision_color, axis=1, keepdims=True + ).clip(min=1e-9) + color_similarity = vision_color_norm @ text_color_norm.T # [vision, text] + has_signal = (text_color.sum(1) > 0)[None, :] * np.ones((n, 1)) + matched_color = np.diag(color_similarity) + informative = text_color.sum(1) > 0 + shuffle = np.random.default_rng(0).permutation(n) + report["anchors"]["color_histogram"] = { + "informative_text_fraction": float(informative.mean()), + "matched_mean": float(matched_color[informative].mean()), + "shuffled_mean": float(color_similarity[shuffle, np.arange(n)][informative].mean()), + "matched_minus_shuffled_z": float( + (matched_color[informative] - color_similarity[shuffle, np.arange(n)][informative]).mean() + / (matched_color[informative] - color_similarity[shuffle, np.arange(n)][informative]).std() + * np.sqrt(informative.sum()) + ), + } + channels.append(color_similarity * has_signal) + channel_names.append("color") + + for name, (text_key, vision_key) in declared_pairs.items(): + if name == "color": + continue + t = text_array(text_key) + v = vision_array(vision_key) + rho = spearman(t, v) + report["anchors"][name] = {"true_pairing_spearman": rho} + similarity = -np.abs(rank_z(v)[:, None] - rank_z(t)[None, :]) + channels.append(similarity) + channel_names.append(name) + + # Combined unary similarity: mean of per-channel rank-standardized maps. + stacked = [] + for channel in channels: + flat = channel.flatten() + standardized = (channel - flat.mean()) / max(flat.std(), 1e-9) + stacked.append(standardized) + combined = np.mean(stacked, axis=0) + matched = np.diag(combined) + ranks = (combined >= matched[:, None]).sum(1) + report["combined_unary"] = { + "channels": channel_names, + "matched_mean": float(matched.mean()), + "grand_mean": float(combined.mean()), + "true_z": float( + (matched.mean() - combined.mean()) + / combined.std() + * np.sqrt(n) + ), + "retrieval_r@10": float((ranks <= 10).mean()), + "retrieval_r@100": float((ranks <= 100).mean()), + "median_rank": float(np.median(ranks)), + "chance_r@10": 10.0 / n, + } + torch.save( + { + "vision_node_ids": vision_ids, + "text_node_ids": text_ids, + "order_note": "row i = vision node truth[i]; column j = text node truth[j]", + "channels": {name: torch.from_numpy(np.asarray(channel)) for name, channel in zip(channel_names, channels)}, + "combined": torch.from_numpy(combined), + }, + args.matrix_output, + ) + write_json(args.output, report) + print(json.dumps(report, indent=2)[:2200]) + print(f"Wrote {args.output} and {args.matrix_output}") + + +if __name__ == "__main__": + main() diff --git a/worldalign/blind_recovery.py b/worldalign/blind_recovery.py new file mode 100644 index 0000000..625f407 --- /dev/null +++ b/worldalign/blind_recovery.py @@ -0,0 +1,476 @@ +"""Blind unpaired recovery on gate-passing states. + +The basin audit licenses this experiment: on content-projected VG states +the true assignment sits below every quench, so the remaining question is +search. Three arms attack the measured landscape shape (deep true basin, +glassy surroundings, smooth coarse spectrum): + +A. entropic Sinkhorn annealing: soft coupling, temperature and entropy + schedules, barycentric language states; +B. parallel tempering: replica-exchange Metropolis over permutations with + the exact closed-form swap deltas, imported from spin-glass practice; +C. spectral-band homotopy: solve the coupling in a low-dimensional + spectral band first, then warm-start progressively finer bands -- + a band-limited continuation in the sense of low-order Fourier terms on + the symmetric group, instantiated through the content spectrum. + +The text side enters through a hidden shuffle, so the optimizer's identity +carries no information. The hidden truth scores results and is used +nowhere else. +""" + +from __future__ import annotations + +import argparse +import json +from pathlib import Path + +import torch +import torch.nn.functional as F +from scipy.optimize import linear_sum_assignment + +from .common import seed_everything, write_json +from .content_gate import flickr_states, vg_states +from .energy import log_sinkhorn +from .manifold_gate import standardize_relation + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--dataset", choices=["flickr", "vg"], default="vg") + 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("--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("--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("--dims", type=int, default=256) + parser.add_argument("--shrinkage", type=float, default=0.05) + parser.add_argument("--arms", default="tempering,sinkhorn,homotopy") + parser.add_argument("--replicas", type=int, default=8) + parser.add_argument("--tempering-rounds", type=int, default=40000) + parser.add_argument("--temp-high", type=float, default=3e-3) + parser.add_argument("--temp-low", type=float, default=1e-5) + parser.add_argument("--exchange-every", type=int, default=20) + parser.add_argument("--sinkhorn-steps", type=int, default=2500) + parser.add_argument("--sinkhorn-restarts", type=int, default=4) + parser.add_argument("--homotopy-bands", default="4,8,16,32,64,full") + parser.add_argument( + "--unary-matrix", + default="", + help="Optional anchor similarity matrix (.pt) in ground-truth node " + "order; injected as a unary field on the tempering energy.", + ) + parser.add_argument("--unary-weight", type=float, default=2.0) + parser.add_argument("--init", choices=["random", "unary"], default="random") + parser.add_argument("--device", default="cuda:3") + parser.add_argument("--seed", type=int, default=20260730) + parser.add_argument("--output", required=True) + return parser.parse_args() + + +def relation_mse_soft( + language_states: torch.Tensor, visual_standardized: torch.Tensor +) -> torch.Tensor: + """Differentiable standardized-relation MSE for soft states.""" + normalized = F.normalize(language_states, dim=-1) + relation = normalized @ normalized.T + mask = ~torch.eye(len(relation), dtype=torch.bool, device=relation.device) + values = relation[mask] + standardized = (values - values.mean()) / values.std().clamp_min(1e-6) + return (standardized - visual_standardized[mask]).square().mean() + + +def permutation_energy_batch( + text_standardized: torch.Tensor, + visual_standardized: torch.Tensor, + permutations: torch.Tensor, +) -> torch.Tensor: + fields = text_standardized[permutations[:, :, None], permutations[:, None, :]] + mask = ~torch.eye( + text_standardized.shape[-1], dtype=torch.bool, device=fields.device + ) + return ((fields - visual_standardized) ** 2)[:, mask].mean(-1) + + +def proposal_swap_deltas( + permuted_fields: torch.Tensor, + visual_standardized: torch.Tensor, + pairs_p: torch.Tensor, + pairs_q: torch.Tensor, +) -> torch.Tensor: + """Exact deltas for proposed swaps only, O(N) per proposal. + + permuted_fields: [R, N, N] text fields under each replica's current + permutation; pairs_p/pairs_q: [R, P] proposal endpoints. + """ + size = permuted_fields.shape[-1] + count = size * (size - 1) + replica_index = torch.arange(len(permuted_fields), device=permuted_fields.device) + rows_p = permuted_fields[replica_index[:, None], pairs_p] # [R, P, N] + rows_q = permuted_fields[replica_index[:, None], pairs_q] + visual_p = visual_standardized[pairs_p] + visual_q = visual_standardized[pairs_q] + self_p = (rows_p * visual_p).sum(-1) + self_q = (rows_q * visual_q).sum(-1) + cross_pq = (rows_p * visual_q).sum(-1) + cross_qp = (rows_q * visual_p).sum(-1) + direct = permuted_fields[replica_index[:, None], pairs_p, pairs_q] + direct_visual = visual_standardized[pairs_p, pairs_q] + total = self_p + self_q - cross_pq - cross_qp - 2.0 * direct * direct_visual + delta = (4.0 / count) * total + return delta.masked_fill(pairs_p == pairs_q, float("inf")) + + +def score( + permutation: torch.Tensor, truth: torch.Tensor, energy: float, true_energy: float +) -> dict: + return { + "accuracy": float((permutation.cpu() == truth.cpu()).double().mean()), + "energy": energy, + "energy_over_true": energy / true_energy - 1.0, + } + + +def coupling_metrics(coupling: torch.Tensor, truth: torch.Tensor) -> dict: + n = len(coupling) + truth = truth.to(coupling.device) + mass_at_truth = float(coupling[torch.arange(n, device=coupling.device), truth].mean()) + rows, cols = linear_sum_assignment(-coupling.detach().cpu().numpy()) + permutation = torch.from_numpy(cols) + forward = coupling.argmax(-1).cpu() + backward = coupling.argmax(0).cpu() + mutual = backward[forward] == torch.arange(n) + sorted_mass = coupling.sort(-1, descending=True).values + margin = (sorted_mass[:, 0] - sorted_mass[:, 1]).cpu() + top = margin.argsort(descending=True)[: max(1, n // 10)] + return { + "mass_at_truth": mass_at_truth, + "hungarian_accuracy": float((permutation == truth.cpu()).double().mean()), + "mutual_nn_count": int(mutual.sum()), + "mutual_nn_precision": float( + (forward[mutual] == truth.cpu()[mutual]).double().mean() + ) + if mutual.any() + else None, + "top_margin_decile_precision": float( + (forward[top] == truth.cpu()[top]).double().mean() + ), + } + + +def arm_tempering( + text_standardized: torch.Tensor, + visual_standardized: torch.Tensor, + truth: torch.Tensor, + true_energy: float, + args: argparse.Namespace, + generator: torch.Generator, + unary: torch.Tensor | None = None, +) -> dict: + device = text_standardized.device + size = text_standardized.shape[-1] + replicas = args.replicas + weight = args.unary_weight if unary is not None else 0.0 + + def total_energy(permutations: torch.Tensor) -> torch.Tensor: + energy = permutation_energy_batch( + text_standardized, visual_standardized, permutations + ) + if unary is not None: + index = torch.arange(size, device=device) + energy = energy - weight * unary[index, permutations].mean(-1) + return energy + + temperatures = torch.logspace( + torch.log10(torch.tensor(args.temp_low)), + torch.log10(torch.tensor(args.temp_high)), + replicas, + ).to(device) + if args.init == "unary" and unary is not None: + rows, cols = linear_sum_assignment(-unary.cpu().numpy()) + start = torch.from_numpy(cols).to(device) + permutations = torch.stack([start.clone() for _ in range(replicas)]) + else: + permutations = torch.stack( + [ + torch.randperm(size, generator=generator).to(device) + for _ in range(replicas) + ] + ) + energies = total_energy(permutations) + best = {"energy": float("inf"), "permutation": permutations[0].clone()} + accepted = 0 + for round_index in range(args.tempering_rounds): + fields = text_standardized[ + permutations[:, :, None], permutations[:, None, :] + ] + pairs_p = torch.randint(0, size, (replicas, 24), generator=generator).to( + device + ) + pairs_q = torch.randint(0, size, (replicas, 24), generator=generator).to( + device + ) + proposal_deltas = proposal_swap_deltas( + fields, visual_standardized, pairs_p, pairs_q + ) + if unary is not None: + replica_index = torch.arange(replicas, device=device)[:, None] + assigned_p = permutations[replica_index, pairs_p] + assigned_q = permutations[replica_index, pairs_q] + unary_delta = ( + unary[pairs_p, assigned_q] + + unary[pairs_q, assigned_p] + - unary[pairs_p, assigned_p] + - unary[pairs_q, assigned_q] + ) + proposal_deltas = proposal_deltas - (weight / size) * unary_delta + noise = torch.rand(replicas, 24, generator=generator).to(device) + acceptable = ( + proposal_deltas < -temperatures[:, None] * noise.clamp_min(1e-12).log() + ) & torch.isfinite(proposal_deltas) + for replica in range(replicas): + hits = torch.nonzero(acceptable[replica]) + if not len(hits): + continue + first = int(hits[0, 0]) + p = int(pairs_p[replica, first]) + q = int(pairs_q[replica, first]) + permutations[replica][[p, q]] = permutations[replica][[q, p]] + energies[replica] = energies[replica] + proposal_deltas[replica, first] + accepted += 1 + if round_index % args.exchange_every == 0: + for replica in range(replicas - 1): + gap = (energies[replica] - energies[replica + 1]) * ( + 1.0 / temperatures[replica] - 1.0 / temperatures[replica + 1] + ) + if gap > 0 or torch.rand(1, generator=generator).item() < float( + gap.exp() + ): + permutations[[replica, replica + 1]] = permutations[ + [replica + 1, replica] + ] + energies[[replica, replica + 1]] = energies[ + [replica + 1, replica] + ] + cold = int(energies.argmin()) + if float(energies[cold]) < best["energy"]: + best = { + "energy": float(energies[cold]), + "permutation": permutations[cold].clone(), + } + if round_index % 5000 == 0: + energies = total_energy(permutations) # refresh against drift + print( + json.dumps( + { + "tempering_round": round_index, + "cold_energy": float(energies.min()), + "cold_accuracy": float( + (permutations[int(energies.argmin())].cpu() == truth) + .double() + .mean() + ), + "accepted": accepted, + } + ) + ) + final = [ + score( + permutations[r], + truth, + float( + permutation_energy_batch( + text_standardized, visual_standardized, permutations[r : r + 1] + )[0] + ), + true_energy, + ) + for r in range(replicas) + ] + best_score = score( + best["permutation"], + truth, + float( + permutation_energy_batch( + text_standardized, visual_standardized, best["permutation"][None] + )[0] + ), + true_energy, + ) + return {"replicas": final, "best": best_score, "accepted_moves": accepted} + + +def arm_sinkhorn( + text_states: torch.Tensor, + visual_standardized: torch.Tensor, + truth: torch.Tensor, + true_energy: float, + args: argparse.Namespace, + generator: torch.Generator, + bands: list[int | None] | None = None, +) -> dict: + device = text_states.device + size = len(text_states) + restarts = [] + for restart in range(args.sinkhorn_restarts if bands is None else 1): + logits = torch.nn.Parameter( + 0.01 + * torch.randn(size, size, generator=generator).to(device) + ) + optimizer = torch.optim.Adam([logits], lr=0.08) + stage_states = text_states + stages = bands if bands is not None else [None] + for stage_index, band in enumerate(stages): + if band is not None: + keep = min(band, text_states.shape[-1]) + stage_states = text_states[:, :keep] + steps = args.sinkhorn_steps // len(stages) + for step in range(steps): + progress = step / max(steps - 1, 1) + temperature = 0.5 * (0.05 / 0.5) ** progress + coupling = log_sinkhorn(logits, temperature, iterations=12) + barycenter = coupling @ stage_states + loss = relation_mse_soft(barycenter, visual_standardized) + entropy = -(coupling * coupling.clamp_min(1e-12).log()).sum(-1).mean() + loss = loss + (0.002 + 0.02 * progress) * entropy + optimizer.zero_grad(set_to_none=True) + loss.backward() + optimizer.step() + with torch.no_grad(): + coupling = log_sinkhorn(logits, 0.05, iterations=30) + metrics = coupling_metrics(coupling, truth) + rows, cols = linear_sum_assignment(-coupling.detach().cpu().numpy()) + permutation = torch.from_numpy(cols).to(device) + rounded_energy = float( + permutation_energy_batch( + standardize_relation( + F.normalize(text_states, dim=-1) + @ F.normalize(text_states, dim=-1).T + )[0][None].squeeze(0), + visual_standardized, + permutation[None], + )[0] + ) + metrics.update( + score(permutation, truth, rounded_energy, true_energy) + ) + restarts.append(metrics) + print(json.dumps({"sinkhorn_restart": restart, **metrics})) + return {"restarts": restarts} + + +def main() -> None: + args = parse_args() + seed_everything(args.seed) + if args.dataset == "flickr": + text_states, visual_states = flickr_states(args) + else: + text_states, visual_states = vg_states(args) + device = torch.device(args.device) + generator = torch.Generator().manual_seed(args.seed) + + size = len(text_states) + hidden = torch.randperm(size, generator=generator) + truth = torch.argsort(hidden) # input row j holds true counterpart hidden[j] + text_input = text_states[hidden].float().to(device) + visual_states = visual_states.float().to(device) + + unary = None + if args.unary_matrix: + anchor_state = torch.load( + args.unary_matrix, map_location="cpu", weights_only=False + ) + combined = anchor_state["combined"].double() + subset_generator = torch.Generator().manual_seed(args.subset_seed) + subset = torch.randperm(len(combined), generator=subset_generator)[ + : args.samples + ] + unary_subset = combined[subset][:, subset] + unary_subset = (unary_subset - unary_subset.mean()) / unary_subset.std() + unary = unary_subset[:, hidden].float().to(device) + + text_standardized = standardize_relation( + F.normalize(text_input, dim=-1) @ F.normalize(text_input, dim=-1).T + )[0] + visual_standardized = standardize_relation( + F.normalize(visual_states, dim=-1) @ F.normalize(visual_states, dim=-1).T + )[0] + true_energy = float( + permutation_energy_batch( + text_standardized, visual_standardized, truth[None].to(device) + )[0] + ) + report: dict = { + "protocol": ( + "The text side enters through a hidden shuffle; the optimizer " + "never sees pair, order, or truth information. Hidden truth " + "scores outcomes only." + ), + "dataset": args.dataset, + "dims": args.dims, + "samples": size, + "true_energy": true_energy, + "chance_accuracy": 1.0 / size, + "arms": {}, + } + arms = [arm.strip() for arm in args.arms.split(",")] + if "tempering" in arms: + report["arms"]["tempering"] = arm_tempering( + text_standardized, + visual_standardized, + truth, + true_energy, + args, + generator, + unary=unary, + ) + if unary is not None: + report["unary"] = { + "weight": args.unary_weight, + "matched_mean": float( + unary[torch.arange(size, device=unary.device), truth.to(unary.device)].mean() + ), + "grand_mean": float(unary.mean()), + "init": args.init, + } + if "sinkhorn" in arms: + report["arms"]["sinkhorn"] = arm_sinkhorn( + text_input, visual_standardized, truth, true_energy, args, generator + ) + if "homotopy" in arms: + bands: list[int | None] = [ + None if band == "full" else int(band) + for band in args.homotopy_bands.split(",") + ] + report["arms"]["homotopy"] = arm_sinkhorn( + text_input, + visual_standardized, + truth, + true_energy, + args, + generator, + bands=bands, + ) + summary = {} + for arm, result in report["arms"].items(): + candidates = result.get("restarts") or result.get("replicas") or [] + if "best" in result: + candidates = candidates + [result["best"]] + best = max(candidates, key=lambda c: c["accuracy"], default=None) + summary[arm] = best + report["summary"] = summary + print(json.dumps({"summary": summary})) + Path(args.output).parent.mkdir(parents=True, exist_ok=True) + write_json(args.output, report) + print(f"Wrote {args.output}") + + +if __name__ == "__main__": + main() diff --git a/worldalign/common.py b/worldalign/common.py new file mode 100644 index 0000000..aa7b174 --- /dev/null +++ b/worldalign/common.py @@ -0,0 +1,142 @@ +from __future__ import annotations + +import json +import math +import os +import random +from pathlib import Path +from typing import Any + +import numpy as np +import torch +import torch.nn.functional as F + + +DATASET_NAME = "nlphuji/flickr30k" +DATASET_SPLIT = "test" + + +def seed_everything(seed: int) -> None: + random.seed(seed) + np.random.seed(seed) + torch.manual_seed(seed) + if torch.cuda.is_available(): + torch.cuda.manual_seed_all(seed) + + +def read_json(path: str | os.PathLike[str]) -> dict[str, Any]: + with open(path, encoding="utf-8") as f: + return json.load(f) + + +def write_json(path: str | os.PathLike[str], value: Any) -> None: + path = Path(path) + path.parent.mkdir(parents=True, exist_ok=True) + with open(path, "w", encoding="utf-8") as f: + json.dump(value, f, indent=2, ensure_ascii=False) + + +def normalized(x: torch.Tensor, eps: float = 1e-8) -> torch.Tensor: + return F.normalize(x.float(), dim=-1, eps=eps) + + +def pairwise_cosine_distance(x: np.ndarray) -> np.ndarray: + x = x.astype(np.float64, copy=False) + x /= np.linalg.norm(x, axis=1, keepdims=True).clip(min=1e-12) + d = 1.0 - x @ x.T + np.fill_diagonal(d, 0.0) + scale = np.median(d[d > 0]) + return d / max(float(scale), 1e-12) + + +def linear_cka(x: torch.Tensor, y: torch.Tensor) -> float: + x = x.float() - x.float().mean(0, keepdim=True) + y = y.float() - y.float().mean(0, keepdim=True) + xty = x.T @ y + numerator = (xty * xty).sum() + xx = x.T @ x + yy = y.T @ y + denominator = torch.sqrt((xx * xx).sum() * (yy * yy).sum()) + return float((numerator / denominator.clamp_min(1e-12)).item()) + + +def retrieval_metrics( + image_features: torch.Tensor, + text_features: torch.Tensor, + ks: tuple[int, ...] = (1, 5, 10), +) -> dict[str, float]: + image_features = normalized(image_features) + text_features = normalized(text_features) + similarities = image_features @ text_features.T + n = similarities.shape[0] + truth = torch.arange(n, device=similarities.device) + + i2t_order = similarities.argsort(dim=1, descending=True) + t2i_order = similarities.T.argsort(dim=1, descending=True) + i2t_rank = (i2t_order == truth[:, None]).nonzero()[:, 1] + t2i_rank = (t2i_order == truth[:, None]).nonzero()[:, 1] + + result: dict[str, float] = {} + for k in ks: + result[f"i2t_r@{k}"] = float((i2t_rank < k).float().mean().item()) + result[f"t2i_r@{k}"] = float((t2i_rank < k).float().mean().item()) + result["i2t_median_rank"] = float(i2t_rank.float().median().item() + 1) + result["t2i_median_rank"] = float(t2i_rank.float().median().item() + 1) + result["chance_r@1"] = 1.0 / max(n, 1) + return result + + +def batch_indices(n: int, batch_size: int, shuffle: bool = False, seed: int = 0): + order = np.arange(n) + if shuffle: + rng = np.random.default_rng(seed) + rng.shuffle(order) + for start in range(0, n, batch_size): + yield order[start : start + batch_size] + + +def sliced_wasserstein( + x: torch.Tensor, + y: torch.Tensor, + num_projections: int = 64, +) -> torch.Tensor: + """Differentiable empirical sliced W2 for equal-sized minibatches.""" + n = min(x.shape[0], y.shape[0]) + x = x[:n] + y = y[:n] + directions = torch.randn( + x.shape[-1], num_projections, device=x.device, dtype=x.dtype + ) + directions = F.normalize(directions, dim=0) + x_proj = (x @ directions).sort(dim=0).values + y_proj = (y @ directions).sort(dim=0).values + return (x_proj - y_proj).square().mean() + + +def cosine_isometry_loss(source: torch.Tensor, mapped: torch.Tensor) -> torch.Tensor: + source = normalized(source) + mapped = normalized(mapped) + source_gram = source @ source.T + mapped_gram = mapped @ mapped.T + mask = ~torch.eye(source.shape[0], dtype=torch.bool, device=source.device) + return (source_gram[mask] - mapped_gram[mask]).square().mean() + + +def cosine_loss(x: torch.Tensor, y: torch.Tensor) -> torch.Tensor: + return 1.0 - F.cosine_similarity(x.float(), y.float(), dim=-1).mean() + + +def dtype_for_device(device: str) -> torch.dtype: + return torch.bfloat16 if device.startswith("cuda") else torch.float32 + + +def parameter_count(module: torch.nn.Module) -> int: + return sum(p.numel() for p in module.parameters()) + + +def cosine_schedule(step: int, steps: int, warmup: int) -> float: + if step < warmup: + return (step + 1) / max(1, warmup) + progress = (step - warmup) / max(1, steps - warmup) + return 0.5 * (1.0 + math.cos(math.pi * progress)) + diff --git a/worldalign/content_gate.py b/worldalign/content_gate.py new file mode 100644 index 0000000..4e1383a --- /dev/null +++ b/worldalign/content_gate.py @@ -0,0 +1,164 @@ +"""Full assignment gate on content-projected states. + +The R6 battery collapsed the improving-swap fraction by one to two orders +of magnitude. This runs the complete gate -- global ranking, exact +transposition enumeration, descent from the truth, and counterfeit search +from random starts -- on the projected states, which the fraction alone +cannot decide. + +Projection hygiene: Flickr directions are fitted on the text-only +training orbits and applied to held-out test states. VG directions are +fitted only on nodes outside the evaluated subset. No pairs anywhere in +the fit; hidden pairs score orderings only. +""" + +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 .content_projection import content_directions, project +from .io import load_feature_pair, select_rows +from .manifold_gate import standardize_relation +from .ricci_control import run_gates + + +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("--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("--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("--dims", type=int, default=32) + parser.add_argument("--shrinkage", type=float, default=0.05) + parser.add_argument("--random-perms", type=int, default=1000) + parser.add_argument("--descent-restarts", type=int, default=3) + 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("--seed", type=int, default=20260729) + parser.add_argument("--output", required=True) + return parser.parse_args() + + +def flickr_states(args: argparse.Namespace) -> tuple[torch.Tensor, torch.Tensor]: + manifest = read_json(args.manifest) + vision, _, vision_lookup, _ = load_feature_pair(args.vision, args.text) + orbits = torch.load(args.text_orbits, map_location="cpu", weights_only=False) + lookup = {int(row): i for i, row in enumerate(orbits["rows"])} + features = F.normalize(orbits["features"].double(), dim=-1) + train_views = features[[lookup[int(r)] for r in manifest["text_only_train"]]] + mean, directions = content_directions(train_views, args.shrinkage) + rows = manifest[args.split][: args.samples] + text_states = project( + features[[lookup[int(r)] for r in rows]].mean(1), mean, directions, args.dims + ) + visual_states = F.normalize( + select_rows(vision["features"], vision_lookup, rows).double(), dim=-1 + ) + return text_states, visual_states + + +def vg_states(args: argparse.Namespace) -> tuple[torch.Tensor, torch.Tensor]: + vision = torch.load(args.vg_vision, map_location="cpu", weights_only=False) + text = torch.load(args.vg_text, map_location="cpu", weights_only=False) + vision_key = "context_states" if "context_states" in vision else "region_features" + text_key = "context_states" if "context_states" in text else "region_features" + pairs = [ + json.loads(line) + for line in open(args.vg_ground_truth, encoding="utf-8") + 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[p["vision_node_id"]] for p in pairs] + text_order = [text_index[p["text_node_id"]] for p in pairs] + visual_views = F.normalize(vision[vision_key].double(), dim=-1)[vision_order] + text_views = F.normalize(text[text_key].double(), dim=-1)[text_order] + generator = torch.Generator().manual_seed(args.subset_seed) + order = torch.randperm(len(visual_views), generator=generator) + subset = order[: args.samples] + holdout = order[args.samples :] + visual_mean, visual_directions = content_directions( + visual_views[holdout], args.shrinkage + ) + text_mean, text_directions = content_directions( + text_views[holdout], args.shrinkage + ) + text_states = project( + text_views[subset].mean(1), text_mean, text_directions, args.dims + ) + visual_states = project( + visual_views[subset].mean(1), visual_mean, visual_directions, args.dims + ) + return text_states, visual_states + + +def main() -> None: + args = parse_args() + seed_everything(args.seed) + if args.dataset == "flickr": + text_states, visual_states = flickr_states(args) + else: + text_states, visual_states = vg_states(args) + text_channels = standardize_relation(text_states @ text_states.T)[0][None] + visual_channels = standardize_relation(visual_states @ visual_states.T)[0][None] + generator = torch.Generator().manual_seed(args.seed) + report = { + "protocol": ( + "Content-projected states, directions fitted without pairs on " + "held-out scenes; the complete assignment gate is scored with " + "hidden pairs." + ), + "dataset": args.dataset, + "dims": args.dims, + "samples": args.samples, + **run_gates( + text_channels, visual_channels, args, generator + ), + } + verdict = { + "true_z_mse": report["gate_a"]["random"]["mse"]["true_z"], + "improving_fraction": report["gate_b"]["improving_fraction"], + "identity_strict_local_min": report["gate_b"]["identity_is_local_min_mse"], + "descent_keeps": report["descent_from_true"]["final_accuracy"], + "true_mse": report["gate_a"]["true"]["mse"], + "best_random_descent": min( + (r["final_objective"] for r in report["descent_from_random"]), + default=None, + ), + "best_random_accuracy": max( + (r["final_accuracy"] for r in report["descent_from_random"]), + default=None, + ), + } + verdict["counterfeit_found"] = bool( + verdict["best_random_descent"] is not None + and verdict["best_random_descent"] < verdict["true_mse"] + and verdict["best_random_accuracy"] < 0.5 + ) + report["verdict"] = verdict + print(json.dumps({"verdict": verdict})) + Path(args.output).parent.mkdir(parents=True, exist_ok=True) + write_json(args.output, report) + print(f"Wrote {args.output}") + + +if __name__ == "__main__": + main() diff --git a/worldalign/content_projection.py b/worldalign/content_projection.py new file mode 100644 index 0000000..b3e7422 --- /dev/null +++ b/worldalign/content_projection.py @@ -0,0 +1,214 @@ +"""R6 battery: cross-view predictable (content) subspace projection. + +Views of the same scene share content and differ in style. Directions that +maximize between-scene over within-scene variance are estimated from +within-modality orbit structure alone (no pairs anywhere), then states are +projected onto the top content directions. Hidden pairs are used only to +score the cross-modal effect. + +Note the contrast with population whitening, which was destructive: the +generalized eigenproblem whitens the within-scene (view-noise) covariance, +not the total covariance, so directions where redescriptions of the same +scene agree are amplified rather than equalized away. +""" + +from __future__ import annotations + +import argparse +import json + +import numpy as np +import torch +import torch.nn.functional as F +from scipy.linalg import eigh + +from .common import read_json, write_json +from .io import load_feature_pair, select_rows +from .manifold_gate import all_transposition_delta_mse, standardize_relation + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + 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("--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("--samples", type=int, default=512) + parser.add_argument("--dims", default="8,16,32,64,128,256") + parser.add_argument("--shrinkage", type=float, default=0.05) + parser.add_argument( + "--output", default="artifacts/manifold_gate/content_projection.json" + ) + return parser.parse_args() + + +def content_directions( + views: torch.Tensor, shrinkage: float +) -> tuple[np.ndarray, np.ndarray]: + """Generalized eigenvectors of between-scene vs within-scene covariance. + + views: [scenes, views_per_scene, dim], any consistent preprocessing. + Returns (mean, directions) with directions sorted by decreasing ratio. + """ + flat = views.reshape(-1, views.shape[-1]).double().numpy() + mean = flat.mean(0) + scene_means = views.double().mean(1).numpy() + within = views.double().numpy() - scene_means[:, None, :] + within = within.reshape(-1, views.shape[-1]) + sigma_within = within.T @ within / max(len(within) - 1, 1) + centered_means = scene_means - mean + sigma_between = ( + centered_means.T @ centered_means / max(len(centered_means) - 1, 1) + ) + trace_scale = np.trace(sigma_within) / len(sigma_within) + regularized = sigma_within + shrinkage * trace_scale * np.eye(len(sigma_within)) + values, vectors = eigh(sigma_between, regularized) + order = np.argsort(values)[::-1] + return mean, vectors[:, order] + + +def project( + states: torch.Tensor, mean: np.ndarray, directions: np.ndarray, dims: int +) -> torch.Tensor: + basis = torch.from_numpy(directions[:, :dims]).double() + centered = states.double() - torch.from_numpy(mean).double() + return F.normalize(centered @ basis, dim=-1) + + +def relation_spearman(text_states: torch.Tensor, visual_states: torch.Tensor) -> float: + text_relation = text_states @ text_states.T + visual_relation = visual_states @ visual_states.T + mask = ~torch.eye(len(text_relation), dtype=torch.bool) + t, v = text_relation[mask], visual_relation[mask] + ranks = torch.stack( + [t.argsort().argsort().double(), v.argsort().argsort().double()] + ) + return float(torch.corrcoef(ranks)[0, 1]) + + +def improving_fraction( + text_states: torch.Tensor, visual_states: torch.Tensor +) -> float: + text_field, _, _ = standardize_relation(text_states @ text_states.T) + visual_field, _, _ = standardize_relation(visual_states @ visual_states.T) + delta = all_transposition_delta_mse(text_field, visual_field) + upper = torch.triu(torch.ones_like(delta, dtype=torch.bool), diagonal=1) + return float((delta[upper] < 0).double().mean()) + + +def flickr_battery(args: argparse.Namespace, dims: list[int]) -> dict: + manifest = read_json(args.manifest) + vision, _, vision_lookup, _ = load_feature_pair(args.vision, args.text) + orbits = torch.load(args.text_orbits, map_location="cpu", weights_only=False) + lookup = {int(row): i for i, row in enumerate(orbits["rows"])} + features = F.normalize(orbits["features"].double(), dim=-1) + + train_rows = [int(r) for r in manifest["text_only_train"]] + test_rows = manifest["test"][: args.samples] + train_views = features[[lookup[r] for r in train_rows]] + mean, directions = content_directions(train_views, args.shrinkage) + + visual = select_rows(vision["features"], vision_lookup, test_rows).double() + visual = F.normalize(visual, dim=-1) + test_views = features[[lookup[int(r)] for r in test_rows]] + orbit_mean = F.normalize(test_views.mean(1), dim=-1) + single = test_views[:, 0] + + report = { + "baseline_orbit_mean": { + "spearman": relation_spearman(orbit_mean, visual), + "improving_fraction": improving_fraction(orbit_mean, visual), + }, + "baseline_single": { + "spearman": relation_spearman(F.normalize(single, dim=-1), visual), + "improving_fraction": improving_fraction( + F.normalize(single, dim=-1), visual + ), + }, + "projected": {}, + } + for k in dims: + projected_mean = project(test_views.mean(1), mean, directions, k) + projected_single = project(single, mean, directions, k) + report["projected"][k] = { + "orbit_mean_spearman": relation_spearman(projected_mean, visual), + "orbit_mean_improving_fraction": improving_fraction( + projected_mean, visual + ), + "single_spearman": relation_spearman(projected_single, visual), + } + return report + + +def vg_battery(args: argparse.Namespace, dims: list[int]) -> 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 open(args.vg_ground_truth, encoding="utf-8") + 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[p["vision_node_id"]] for p in pairs] + text_order = [text_index[p["text_node_id"]] for p in pairs] + visual_views = F.normalize(vision["region_features"].double(), dim=-1)[ + vision_order + ] + text_views = F.normalize(text["region_features"].double(), dim=-1)[text_order] + + visual_mean, visual_directions = content_directions(visual_views, args.shrinkage) + text_mean, text_directions = content_directions(text_views, args.shrinkage) + + generator = torch.Generator().manual_seed(0) + subset = torch.randperm(len(visual_views), generator=generator)[: args.samples] + visual_node = F.normalize(visual_views[subset].mean(1), dim=-1) + text_node = F.normalize(text_views[subset].mean(1), dim=-1) + + report = { + "baseline": { + "spearman": relation_spearman(text_node, visual_node), + "improving_fraction": improving_fraction(text_node, visual_node), + }, + "projected": {}, + } + for k in dims: + projected_text = project( + text_views[subset].mean(1), text_mean, text_directions, k + ) + projected_visual = project( + visual_views[subset].mean(1), visual_mean, visual_directions, k + ) + report["projected"][k] = { + "both_sides_spearman": relation_spearman(projected_text, projected_visual), + "both_sides_improving_fraction": improving_fraction( + projected_text, projected_visual + ), + "text_only_spearman": relation_spearman(projected_text, visual_node), + } + return report + + +def main() -> None: + args = parse_args() + dims = [int(d) for d in args.dims.split(",")] + report = { + "protocol": ( + "Content directions maximize between-scene over within-scene " + "variance of view states, fitted per modality on unpaired " + "orbit structure only. Hidden pairs score the effect." + ), + "flickr": flickr_battery(args, dims), + "vg_region_closed": vg_battery(args, dims), + } + write_json(args.output, report) + print(json.dumps(report, indent=2)) + + +if __name__ == "__main__": + main() diff --git a/worldalign/cooccurrence_dictionary.py b/worldalign/cooccurrence_dictionary.py new file mode 100644 index 0000000..4c338e4 --- /dev/null +++ b/worldalign/cooccurrence_dictionary.py @@ -0,0 +1,237 @@ +"""Derive the cross-modal value correspondence from joint structure. + +Marginal frequency pairs values only when the two corpora rank them the +same way, which reporting bias erodes: text mentions the salient, pixels +count the common. Joint structure is sturdier. Which colours co-occur +with which shapes is a property of the world that both modalities +observe, and reporting bias distorts the marginals long before it +scrambles that pattern. + +The correspondence is therefore recovered by matching two small +value-level graphs -- colour-by-shape co-occurrence, estimated per +modality on its own corpus -- with the same spectral-plus-refinement +solver used at scene level. Both corpora stay disjoint; no pair is read. +""" + +from __future__ import annotations + +import argparse +import json +from collections import Counter +from pathlib import Path + +import numpy as np +import torch +from scipy.cluster.hierarchy import fcluster, linkage +from scipy.optimize import linear_sum_assignment +from sklearn.cluster import KMeans +from tqdm import tqdm + +from .common import read_json, seed_everything, write_json +from .synth_set_battery import parse_group_phrases +from .synth_towers import load_image +from .tier0_pipeline import ( + appearance, + extract_objects, + group_objects, +) +from .tier0_dictionary import partition_text_vocabulary, text_token_statistics + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--data-dir", default="artifacts/synth_v3") + parser.add_argument("--fit-scenes", type=int, default=3000) + parser.add_argument("--peak-distance", type=int, default=5) + parser.add_argument("--group-threshold", type=float, default=0.35) + parser.add_argument("--shape-classes", type=int, default=6) + parser.add_argument("--restarts", type=int, default=40) + parser.add_argument("--seed", type=int, default=0) + parser.add_argument( + "--output", default="artifacts/synth_v3/cooccurrence_dictionary.json" + ) + return parser.parse_args() + + +def text_joint( + captions: list[list[str]], + rows: list[int], + colour_words: list[str], + shape_words: list[str], +) -> np.ndarray: + joint = np.zeros((len(colour_words), len(shape_words))) + colour_index = {word: i for i, word in enumerate(colour_words)} + shape_index = {word: i for i, word in enumerate(shape_words)} + for row in rows: + for phrase in parse_group_phrases(captions[row][0]): + tokens = phrase.split() + colour = next((colour_index[t] for t in tokens if t in colour_index), None) + shape = next( + ( + shape_index[t.rstrip("es") if t.rstrip("es") in shape_index else t] + for t in tokens + if t in shape_index or t.rstrip("es") in shape_index + ), + None, + ) + if colour is not None and shape is not None: + joint[colour, shape] += 1.0 + return joint + + +def vision_joint( + rows: list[int], + image_dir: Path, + args: argparse.Namespace, + colour_classes: int, +) -> tuple[np.ndarray, KMeans, KMeans, dict[int, int]]: + groups, shapes = [], [] + for row in tqdm(rows, desc="vision joint"): + objects = extract_objects( + load_image(image_dir / f"scene{row:06d}_v0.png"), args.peak_distance + ) + members = group_objects(objects, args.group_threshold) + if not members: + continue + features = np.stack([appearance(item) for item in objects]) + if len(objects) == 1: + labels = np.array([0]) + else: + labels = fcluster( + linkage(features, "complete"), args.group_threshold, "distance" + ) + buckets: dict[int, list[dict]] = {} + for item, label in zip(objects, labels): + buckets.setdefault(int(label), []).append(item) + for group, bucket in zip(members, buckets.values()): + groups.append(group) + shapes.append(np.mean([item["shape"] for item in bucket], axis=0)) + colour_model = KMeans(colour_classes, n_init=10, random_state=args.seed).fit( + np.stack([group["rgb"] for group in groups]).astype(np.float64) + ) + shape_model = KMeans(args.shape_classes, n_init=10, random_state=args.seed).fit( + np.stack(shapes).astype(np.float64) + ) + colour_labels = colour_model.labels_ + shape_labels = shape_model.labels_ + frequency = Counter(colour_labels.tolist()) + rank = {label: index for index, (label, _) in enumerate(frequency.most_common())} + joint = np.zeros((colour_classes, args.shape_classes)) + for colour, shape in zip(colour_labels, shape_labels): + joint[rank[int(colour)], int(shape)] += 1.0 + return joint, colour_model, shape_model, rank + + +def normalise_joint(joint: np.ndarray) -> np.ndarray: + """Row-normalised joint: the shape profile of each colour.""" + return joint / joint.sum(axis=1, keepdims=True).clip(1e-9) + + +def match_values( + text: np.ndarray, vision: np.ndarray, restarts: int, seed: int +) -> tuple[np.ndarray, np.ndarray, float]: + """Align colour rows and shape columns of two joint tables. + + Alternating assignment: given a column correspondence, rows are + matched by profile similarity; given rows, columns are rematched. + Iterating from many random column starts and keeping the best + agreement avoids the trivial fixed point. Row indices are returned as + vision-colour to text-colour. + """ + generator = np.random.default_rng(seed) + text_profile = normalise_joint(text) + vision_profile = normalise_joint(vision) + shapes = text.shape[1] + best_row, best_column, best_value = None, None, -np.inf + for restart in range(restarts): + column = ( + np.arange(shapes) if restart == 0 else generator.permutation(shapes) + ) + row = None + for _ in range(20): + # Rows: vision colour i against text colour j under current columns. + score = vision_profile @ text_profile[:, column].T + _, row = linear_sum_assignment(-score) + # Columns: vision shape a against text shape b under current rows. + column_score = vision_profile.T @ text_profile[row] + _, new_column = linear_sum_assignment(-column_score) + if np.array_equal(new_column, column): + break + column = new_column + value = float((vision_profile * text_profile[row][:, column]).sum()) + if value > best_value: + best_row, best_column, best_value = row, column, value + return best_row, best_column, best_value + + +def main() -> None: + args = parse_args() + seed_everything(args.seed) + manifest = read_json(Path(args.data_dir, "manifest.json")) + captions = read_json(Path(args.data_dir, "captions.json"))["captions"] + scenes = read_json(Path(args.data_dir, "scenes.private.json"))["scenes"] + image_dir = Path(manifest["image_dir"]) + + families = partition_text_vocabulary( + text_token_statistics(captions, manifest["text_only_train"]) + ) + colour_words = families["colour_words"] + from .synth_world import SHAPES + + shape_words = list(SHAPES) + text_table = text_joint( + captions, manifest["text_only_train"], colour_words, shape_words + ) + vision_table, colour_model, _, rank = vision_joint( + manifest["vision_only_train"][: args.fit_scenes], + image_dir, + args, + len(colour_words), + ) + + row, column, agreement = match_values( + text_table, vision_table, args.restarts, args.seed + ) + + # Evaluation only: names attached to vision clusters by nearest true RGB. + from .synth_world import COLORS + + names = list(COLORS) + reference = np.asarray([COLORS[name] for name in names], dtype=np.float64) / 255.0 + inverse_rank = {index: label for label, index in rank.items()} + cluster_name = { + index: names[ + int(np.argmin(((colour_model.cluster_centers_[inverse_rank[index]] - reference) ** 2).sum(1))) + ] + for index in range(len(colour_words)) + } + marginal_correct = sum( + 1 + for index, word in enumerate(colour_words) + if index < len(cluster_name) and word == cluster_name[index] + ) + joint_correct = sum( + 1 + for vision_index, text_index in enumerate(row) + if colour_words[text_index] == cluster_name[vision_index] + ) + report = { + "protocol": ( + "Colour-by-shape joint tables are estimated per modality on " + "disjoint corpora and aligned by alternating assignment; " + "truth is read only to count correct entries." + ), + "colour_words": colour_words, + "agreement": agreement, + "marginal_rank_correct": marginal_correct, + "joint_structure_correct": joint_correct, + "classes": len(colour_words), + } + print(json.dumps({k: report[k] for k in + ("marginal_rank_correct", "joint_structure_correct", "classes")})) + write_json(args.output, report) + print(f"Wrote {args.output}") + + +if __name__ == "__main__": + main() diff --git a/worldalign/diagnose.py b/worldalign/diagnose.py new file mode 100644 index 0000000..2d05fd9 --- /dev/null +++ b/worldalign/diagnose.py @@ -0,0 +1,76 @@ +from __future__ import annotations + +import argparse +from pathlib import Path + +import numpy as np +import torch +from scipy.stats import spearmanr + +from .common import linear_cka, normalized, read_json, write_json +from .io import load_feature_pair, select_rows + + +def parse_args() -> argparse.Namespace: + p = argparse.ArgumentParser() + p.add_argument("--manifest", default="artifacts/manifest.json") + p.add_argument("--vision", default="artifacts/vision.pt") + p.add_argument("--text", default="artifacts/text.pt") + p.add_argument("--split", choices=["val", "test"], default="val") + p.add_argument("--max-samples", type=int, default=1_000) + p.add_argument("--permutations", type=int, default=20) + p.add_argument("--seed", type=int, default=20260728) + p.add_argument("--output", default="artifacts/diagnostics.json") + return p.parse_args() + + +def upper_triangle(x: torch.Tensor) -> np.ndarray: + n = x.shape[0] + i, j = torch.triu_indices(n, n, offset=1) + return x[i, j].cpu().numpy() + + +def main() -> None: + args = parse_args() + manifest = read_json(args.manifest) + vision, text, vlookup, tlookup = load_feature_pair(args.vision, args.text) + rows = manifest[args.split][: args.max_samples] + x = select_rows(vision["features"], vlookup, rows) + y = select_rows(text["features"], tlookup, rows) + + gx = normalized(x) @ normalized(x).T + gy = normalized(y) @ normalized(y).T + gx_upper = upper_triangle(gx) + gy_upper = upper_triangle(gy) + rho = spearmanr(gx_upper, gy_upper).statistic + generator = torch.Generator().manual_seed(args.seed) + shuffled_rhos = [] + for _ in range(args.permutations): + permutation = torch.randperm(len(y), generator=generator) + shuffled_gy = gy[permutation][:, permutation] + shuffled_rhos.append( + float(spearmanr(gx_upper, upper_triangle(shuffled_gy)).statistic) + ) + shuffled_mean = float(np.mean(shuffled_rhos)) + shuffled_std = float(np.std(shuffled_rhos)) + result = { + "split": args.split, + "samples": len(rows), + "linear_cka": linear_cka(x, y), + "pairwise_cosine_spearman": float(rho), + "shuffled_spearman_mean": shuffled_mean, + "shuffled_spearman_std": shuffled_std, + "spearman_shuffle_z": float( + (rho - shuffled_mean) / max(shuffled_std, 1e-12) + ), + "permutations": args.permutations, + "vision_dim": x.shape[-1], + "text_dim": y.shape[-1], + } + Path(args.output).parent.mkdir(parents=True, exist_ok=True) + write_json(args.output, result) + print(result) + + +if __name__ == "__main__": + main() diff --git a/worldalign/energy.py b/worldalign/energy.py new file mode 100644 index 0000000..796c41a --- /dev/null +++ b/worldalign/energy.py @@ -0,0 +1,137 @@ +from __future__ import annotations + +import torch +import torch.nn.functional as F + + +def off_diagonal_mask(size: int, device: torch.device | str) -> torch.Tensor: + return ~torch.eye(size, dtype=torch.bool, device=device) + + +def standardized_relation( + features: torch.Tensor, mask: torch.Tensor | None = None +) -> tuple[torch.Tensor, torch.Tensor]: + """Cosine relation field and its standardized off-diagonal values.""" + features = F.normalize(features.float(), dim=-1) + relation = features @ features.T + if mask is None: + mask = off_diagonal_mask(len(features), features.device) + values = relation[mask] + standardized = (values - values.mean()) / values.std().clamp_min(1e-6) + return relation, standardized + + +def relation_field_energy( + visual_relation: torch.Tensor, + visual_standardized: torch.Tensor, + language_particles: torch.Tensor, + temperatures: tuple[float, ...] = (0.03, 0.07, 0.15), +) -> tuple[torch.Tensor, torch.Tensor]: + """Second-order and multiscale conditional relation energies.""" + language_relation, language_standardized = standardized_relation( + language_particles + ) + mse = F.mse_loss(language_standardized, visual_standardized) + diagonal = torch.eye( + len(language_particles), + dtype=torch.bool, + device=language_particles.device, + ) + conditional_kl = language_particles.new_zeros(()) + for temperature in temperatures: + visual_logits = (visual_relation / temperature).masked_fill( + diagonal, -1e4 + ) + language_logits = (language_relation / temperature).masked_fill( + diagonal, -1e4 + ) + visual_probability = F.softmax(visual_logits, dim=-1) + conditional_kl = conditional_kl + ( + visual_probability + * ( + F.log_softmax(visual_logits, dim=-1) + - F.log_softmax(language_logits, dim=-1) + ) + ).sum(-1).mean() + return mse, conditional_kl + + +def projection_quantile_target( + text_features: torch.Tensor, + particles: int, + projections: int, + generator: torch.Generator, +) -> tuple[torch.Tensor, torch.Tensor]: + """Fixed sliced-distribution target from an unpaired text population.""" + directions = F.normalize( + torch.randn( + text_features.shape[-1], + projections, + generator=generator, + device=text_features.device, + ), + dim=0, + ) + projected = (text_features @ directions).sort(dim=0).values + quantile_indices = ( + torch.linspace( + 0, + len(projected) - 1, + particles, + device=text_features.device, + ) + .round() + .long() + ) + return directions, projected[quantile_indices] + + +def sliced_distribution_energy( + particles: torch.Tensor, + directions: torch.Tensor, + target_quantiles: torch.Tensor, +) -> torch.Tensor: + projected = (F.normalize(particles, dim=-1) @ directions).sort( + dim=0 + ).values + return F.mse_loss(projected, target_quantiles) + + +def prototype_manifold_energy( + particles: torch.Tensor, prototypes: torch.Tensor +) -> torch.Tensor: + particles = F.normalize(particles, dim=-1) + prototypes = F.normalize(prototypes, dim=-1) + return (1 - (particles @ prototypes.T).max(dim=-1).values).mean() + + +def log_sinkhorn( + logits: torch.Tensor, temperature: float, iterations: int = 12 +) -> torch.Tensor: + """Doubly stochastic coupling with differentiable log-domain updates.""" + log_coupling = logits / temperature + for _ in range(iterations): + log_coupling = log_coupling - torch.logsumexp( + log_coupling, dim=1, keepdim=True + ) + log_coupling = log_coupling - torch.logsumexp( + log_coupling, dim=0, keepdim=True + ) + return log_coupling.exp() + + +def retrieval_metrics( + particles: torch.Tensor, paired_text: torch.Tensor +) -> dict[str, float]: + particles = F.normalize(particles.float(), dim=-1) + paired_text = F.normalize(paired_text.float(), dim=-1) + similarity = particles @ paired_text.T + target = similarity.diagonal() + ranks = (similarity > target[:, None]).sum(-1) + 1 + return { + "r@1": float((ranks <= 1).float().mean()), + "r@5": float((ranks <= 5).float().mean()), + "r@10": float((ranks <= 10).float().mean()), + "median_rank": float(ranks.float().median()), + "paired_cosine": float(target.mean()), + } diff --git a/worldalign/energy_coupling.py b/worldalign/energy_coupling.py new file mode 100644 index 0000000..cf5b859 --- /dev/null +++ b/worldalign/energy_coupling.py @@ -0,0 +1,205 @@ +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 .energy import ( + log_sinkhorn, + relation_field_energy, + retrieval_metrics, + standardized_relation, +) +from .io import load_feature_pair, select_rows + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + 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") + parser.add_argument("--split", choices=["val", "test"], default="test") + parser.add_argument("--samples", type=int, default=256) + parser.add_argument( + "--reservoir", + choices=["same_set_shuffled", "unpaired"], + default="same_set_shuffled", + ) + parser.add_argument("--steps", type=int, default=1_200) + parser.add_argument("--lr", type=float, default=0.12) + parser.add_argument("--temperature-start", type=float, default=0.8) + parser.add_argument("--temperature-end", type=float, default=0.12) + parser.add_argument("--conditional-weight", type=float, default=0.08) + parser.add_argument("--entropy-weight-start", type=float, default=0.005) + parser.add_argument("--entropy-weight-end", type=float, default=0.055) + parser.add_argument("--device", default="cuda:1") + parser.add_argument("--seed", type=int, default=20260729) + parser.add_argument( + "--output", default="artifacts/energy_coupling_test.pt" + ) + parser.add_argument( + "--metrics-output", default="artifacts/energy_coupling_test.json" + ) + return parser.parse_args() + + +def main() -> None: + args = parse_args() + seed_everything(args.seed) + 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 = select_rows( + vision["features"], vision_lookup, rows + ).to(args.device) + paired_text = select_rows( + text["features"], text_lookup, rows + ).to(args.device) + text_population = select_rows( + text["features"], text_lookup, manifest["text_only_train"] + ).to(args.device) + if args.text_orbits: + 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 + ) + paired_text = select_rows( + orbit_mean, lookup, rows + ).to(args.device) + text_population = select_rows( + orbit_mean, lookup, manifest["text_only_train"] + ).to(args.device) + + generator = torch.Generator(device=args.device).manual_seed(args.seed) + target: torch.Tensor | None = None + if args.reservoir == "same_set_shuffled": + permutation = torch.randperm( + len(paired_text), generator=generator, device=args.device + ) + anchors = paired_text[permutation] + target = torch.argsort(permutation) + else: + anchor_indices = torch.randperm( + len(text_population), + generator=generator, + device=args.device, + )[: len(visual)] + anchors = text_population[anchor_indices] + + initial_permutation = torch.randperm( + len(visual), generator=generator, device=args.device + ) + logits = torch.nn.Parameter( + 0.02 + * torch.randn( + len(visual), + len(visual), + generator=generator, + device=args.device, + ) + ) + with torch.no_grad(): + logits[torch.arange(len(visual), device=args.device), initial_permutation] += 3 + visual_relation, visual_standardized = standardized_relation(visual) + optimizer = torch.optim.Adam([logits], lr=args.lr) + history: list[dict] = [] + + for step in range(args.steps + 1): + progress = min(step / max(args.steps, 1), 1.0) + temperature = max( + args.temperature_end, + args.temperature_start + + progress + * (args.temperature_end - args.temperature_start), + ) + coupling = log_sinkhorn(logits, temperature, iterations=15) + particles = F.normalize(coupling @ anchors, dim=-1) + relation, conditional = relation_field_energy( + visual_relation, visual_standardized, particles + ) + entropy = -( + coupling * coupling.clamp_min(1e-12).log() + ).sum(-1).mean() + entropy_weight = ( + args.entropy_weight_start + + progress + * (args.entropy_weight_end - args.entropy_weight_start) + ) + loss = ( + relation + + args.conditional_weight * conditional + + entropy_weight * entropy + ) + if step % 100 == 0 or step == args.steps: + record = { + "step": step, + "total": float(loss.detach()), + "relation": float(relation.detach()), + "conditional": float(conditional.detach()), + "entropy": float(entropy.detach()), + "temperature": temperature, + "paired_evaluation_only": retrieval_metrics( + particles.detach(), paired_text + ), + } + if target is not None: + prediction = coupling.argmax(-1) + record["exact_coupling_evaluation_only"] = { + "accuracy": float((prediction == target).float().mean()), + "true_mass": float( + coupling[ + torch.arange( + len(coupling), device=args.device + ), + target, + ].mean() + ), + } + history.append(record) + print(json.dumps(record)) + if step == args.steps: + break + optimizer.zero_grad(set_to_none=True) + loss.backward() + torch.nn.utils.clip_grad_norm_([logits], 5.0) + optimizer.step() + + result = { + "protocol": ( + "The optimizer sees only frozen within-modality states and " + "energy. Pair identity and paired text are used only for " + "evaluation diagnostics." + ), + "mode": "mass_conserving_language_state_coupling", + "reservoir": args.reservoir, + "split": args.split, + "rows": rows, + "args": vars(args), + "history": history, + } + state = { + **result, + "anchors": anchors.detach().cpu(), + "final_coupling": coupling.detach().cpu(), + "final_particles": particles.detach().cpu(), + } + Path(args.output).parent.mkdir(parents=True, exist_ok=True) + torch.save(state, args.output) + write_json(args.metrics_output, result) + print(f"Wrote {args.output} and {args.metrics_output}") + + +if __name__ == "__main__": + main() diff --git a/worldalign/energy_functional_infer.py b/worldalign/energy_functional_infer.py new file mode 100644 index 0000000..6864011 --- /dev/null +++ b/worldalign/energy_functional_infer.py @@ -0,0 +1,262 @@ +from __future__ import annotations + +import argparse +import json +from pathlib import Path + +import torch +import torch.nn.functional as F +from transformers import AutoModelForCausalLM, AutoTokenizer + +from .common import dtype_for_device, read_json, seed_everything, write_json +from .energy import ( + projection_quantile_target, + relation_field_energy, + retrieval_metrics, + sliced_distribution_energy, + standardized_relation, +) +from .extract_text import mean_pool +from .io import load_feature_pair, select_rows +from .models import load_prefix + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + 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("--prefix", default="artifacts/prefix.pt") + parser.add_argument("--split", choices=["val", "test"], default="val") + parser.add_argument("--samples", type=int, default=128) + parser.add_argument("--steps", type=int, default=60) + parser.add_argument("--refresh-steps", type=int, default=20) + parser.add_argument("--max-new-tokens", type=int, default=24) + parser.add_argument("--lr", type=float, default=0.03) + parser.add_argument("--projections", type=int, default=128) + parser.add_argument("--relation-weight", type=float, default=1.0) + parser.add_argument("--conditional-weight", type=float, default=0.08) + parser.add_argument("--distribution-weight", type=float, default=80.0) + parser.add_argument("--functional-weight", type=float, default=10.0) + parser.add_argument("--device", default="cuda:1") + parser.add_argument("--seed", type=int, default=20260729) + parser.add_argument( + "--output", default="artifacts/energy_functional_val.pt" + ) + parser.add_argument( + "--metrics-output", default="artifacts/energy_functional_val.json" + ) + return parser.parse_args() + + +@torch.no_grad() +def refresh_functional_anchor( + particles: torch.Tensor, + prefix: torch.nn.Module, + lm: AutoModelForCausalLM, + tokenizer: AutoTokenizer, + dtype: torch.dtype, + max_new_tokens: int, +) -> tuple[torch.Tensor, list[str]]: + captions: list[str] = [] + for chunk in particles.split(16): + prefix_embedding = prefix(chunk).to(dtype) + attention = torch.ones( + prefix_embedding.shape[:2], + dtype=torch.long, + device=particles.device, + ) + output = lm.generate( + inputs_embeds=prefix_embedding, + attention_mask=attention, + max_new_tokens=max_new_tokens, + do_sample=False, + eos_token_id=tokenizer.eos_token_id, + pad_token_id=tokenizer.pad_token_id, + ) + captions.extend( + tokenizer.batch_decode(output, skip_special_tokens=True) + ) + anchor_parts: list[torch.Tensor] = [] + for start in range(0, len(captions), 32): + tokens = tokenizer( + captions[start : start + 32], + padding=True, + truncation=True, + max_length=64, + return_tensors="pt", + ) + tokens = { + key: value.to(particles.device) + for key, value in tokens.items() + } + result = lm( + **tokens, output_hidden_states=True, return_dict=True + ) + anchor_parts.append( + mean_pool( + result.hidden_states[-1], tokens["attention_mask"] + ).float() + ) + return F.normalize(torch.cat(anchor_parts), dim=-1), captions + + +def main() -> None: + args = parse_args() + seed_everything(args.seed) + manifest = read_json(args.manifest) + vision, text, vision_lookup, _ = load_feature_pair( + args.vision, args.text + ) + orbit_state = torch.load( + args.text_orbits, map_location="cpu", weights_only=False + ) + orbit_lookup = { + int(row): index for index, row in enumerate(orbit_state["rows"]) + } + orbit_mean = F.normalize( + orbit_state["features"].float().mean(1), dim=-1 + ) + rows = manifest[args.split][: args.samples] + visual = select_rows( + vision["features"], vision_lookup, rows + ).to(args.device) + paired_text = select_rows( + orbit_mean, orbit_lookup, rows + ).to(args.device) + text_population = select_rows( + orbit_mean, orbit_lookup, manifest["text_only_train"] + ).to(args.device) + + prefix, prefix_state = load_prefix(args.prefix, args.device) + if prefix.semantic_dim != text_population.shape[-1]: + raise ValueError("Prefix and text-orbit dimensions do not match") + tokenizer = AutoTokenizer.from_pretrained(text["model"]) + if tokenizer.pad_token_id is None: + tokenizer.pad_token = tokenizer.eos_token + dtype = dtype_for_device(args.device) + lm = AutoModelForCausalLM.from_pretrained( + text["model"], torch_dtype=dtype + ).to(args.device) + lm.eval() + for parameter in lm.parameters(): + parameter.requires_grad_(False) + for parameter in prefix.parameters(): + parameter.requires_grad_(False) + + generator = torch.Generator(device=args.device).manual_seed(args.seed) + initial_indices = torch.randperm( + len(text_population), + generator=generator, + device=args.device, + )[: len(visual)] + particles = torch.nn.Parameter(text_population[initial_indices].clone()) + initial_particles = particles.detach().cpu() + directions, target_quantiles = projection_quantile_target( + text_population, + len(particles), + args.projections, + generator, + ) + visual_relation, visual_standardized = standardized_relation(visual) + optimizer = torch.optim.Adam([particles], lr=args.lr) + functional_anchor, functional_captions = refresh_functional_anchor( + particles.detach(), + prefix, + lm, + tokenizer, + dtype, + args.max_new_tokens, + ) + + history: list[dict] = [] + refresh_captions: dict[int, list[str]] = { + 0: functional_captions[:25] + } + for step in range(args.steps + 1): + if step > 0 and step % args.refresh_steps == 0: + functional_anchor, functional_captions = ( + refresh_functional_anchor( + particles.detach(), + prefix, + lm, + tokenizer, + dtype, + args.max_new_tokens, + ) + ) + refresh_captions[step] = functional_captions[:25] + relation, conditional = relation_field_energy( + visual_relation, visual_standardized, particles + ) + distribution = sliced_distribution_energy( + particles, directions, target_quantiles + ) + functional = ( + 1 + - ( + F.normalize(particles, dim=-1) + * functional_anchor.detach() + ).sum(-1) + ).mean() + loss = ( + args.relation_weight * relation + + args.conditional_weight * conditional + + args.distribution_weight * distribution + + args.functional_weight * functional + ) + if step % 10 == 0 or step == args.steps: + record = { + "step": step, + "total": float(loss.detach()), + "relation": float(relation.detach()), + "conditional": float(conditional.detach()), + "distribution": float(distribution.detach()), + "functional": float(functional.detach()), + "paired_evaluation_only": retrieval_metrics( + particles.detach(), paired_text + ), + } + history.append(record) + print(json.dumps(record)) + if step == args.steps: + break + optimizer.zero_grad(set_to_none=True) + loss.backward() + torch.nn.utils.clip_grad_norm_([particles], 2.0) + optimizer.step() + with torch.no_grad(): + particles.copy_(F.normalize(particles, dim=-1)) + + result = { + "protocol": ( + "No image-text pair and no cross-modal parameter is used. The " + "functional language energy is a frozen decode-reencode cycle " + "through a text-only prefix interface and frozen Qwen." + ), + "mode": "functional_cycle_language_latent_particles", + "split": args.split, + "rows": rows, + "args": vars(args), + "history": history, + "prefix_training": prefix_state["training"], + "refresh_caption_examples": refresh_captions, + } + state = { + **result, + "initial_particles": initial_particles, + "final_particles": F.normalize( + particles.detach(), dim=-1 + ).cpu(), + } + Path(args.output).parent.mkdir(parents=True, exist_ok=True) + torch.save(state, args.output) + write_json(args.metrics_output, result) + print(f"Wrote {args.output} and {args.metrics_output}") + + +if __name__ == "__main__": + main() diff --git a/worldalign/energy_infer.py b/worldalign/energy_infer.py new file mode 100644 index 0000000..d5a67b6 --- /dev/null +++ b/worldalign/energy_infer.py @@ -0,0 +1,249 @@ +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 .energy import ( + projection_quantile_target, + prototype_manifold_energy, + relation_field_energy, + retrieval_metrics, + sliced_distribution_energy, + standardized_relation, +) +from .io import load_feature_pair, select_rows + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + 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", + help=( + "Optional multi-description feature cache. When supplied, each " + "language particle is the normalized mean of its observation " + "orbit." + ), + ) + parser.add_argument("--prototypes", default="artifacts/gw.pt") + parser.add_argument("--split", choices=["val", "test"], default="test") + parser.add_argument("--samples", type=int, default=256) + parser.add_argument("--steps", type=int, default=100) + parser.add_argument("--lr", type=float, default=0.03) + parser.add_argument("--projections", type=int, default=128) + parser.add_argument("--relation-weight-start", type=float, default=0.4) + parser.add_argument("--relation-weight-end", type=float, default=2.0) + parser.add_argument("--conditional-weight-start", type=float, default=0.05) + parser.add_argument("--conditional-weight-end", type=float, default=0.2) + parser.add_argument("--distribution-weight", type=float, default=80.0) + parser.add_argument("--manifold-weight", type=float, default=1.0) + parser.add_argument("--noise", type=float, default=0.01) + parser.add_argument( + "--shuffle-visual-energy", + action="store_true", + help=( + "Evaluation control: permute image particles before constructing " + "the energy while leaving evaluation rows unchanged." + ), + ) + parser.add_argument("--device", default="cuda:1") + parser.add_argument("--seed", type=int, default=20260729) + parser.add_argument( + "--output", default="artifacts/energy_free_test.pt" + ) + parser.add_argument( + "--metrics-output", default="artifacts/energy_free_test.json" + ) + return parser.parse_args() + + +def main() -> None: + args = parse_args() + seed_everything(args.seed) + 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 = select_rows( + vision["features"], vision_lookup, rows + ).to(args.device) + paired_text = select_rows( + text["features"], text_lookup, rows + ).to(args.device) + text_population = select_rows( + text["features"], text_lookup, manifest["text_only_train"] + ).to(args.device) + if args.text_orbits: + orbit_state = torch.load( + args.text_orbits, map_location="cpu", weights_only=False + ) + orbit_lookup = { + int(row): index + for index, row in enumerate(orbit_state["rows"]) + } + orbit_mean = F.normalize( + orbit_state["features"].float().mean(1), dim=-1 + ) + paired_text = select_rows( + orbit_mean, orbit_lookup, rows + ).to(args.device) + text_population = select_rows( + orbit_mean, orbit_lookup, manifest["text_only_train"] + ).to(args.device) + if args.shuffle_visual_energy: + control_generator = torch.Generator( + device=args.device + ).manual_seed(args.seed + 10_000) + visual = visual[ + torch.randperm( + len(visual), + generator=control_generator, + device=args.device, + ) + ] + prototype_state = torch.load( + args.prototypes, map_location="cpu", weights_only=False + ) + prototypes = prototype_state["text_centers"].to(args.device) + if prototypes.shape[-1] != text_population.shape[-1]: + raise ValueError( + "Prototype dimension does not match text features; use the GW " + "cache built from the selected text backbone." + ) + + generator = torch.Generator(device=args.device).manual_seed(args.seed) + initial_indices = torch.randperm( + len(text_population), + generator=generator, + device=args.device, + )[: len(visual)] + initial = text_population[initial_indices].clone() + initial = initial + args.noise * torch.randn( + initial.shape, generator=generator, device=args.device + ) + particles = torch.nn.Parameter(F.normalize(initial, dim=-1)) + directions, target_quantiles = projection_quantile_target( + text_population, + particles=len(particles), + projections=args.projections, + generator=generator, + ) + visual_relation, visual_standardized = standardized_relation(visual) + optimizer = torch.optim.Adam([particles], lr=args.lr) + + def energy_values() -> tuple[torch.Tensor, ...]: + relation, conditional = relation_field_energy( + visual_relation, visual_standardized, particles + ) + distribution = sliced_distribution_energy( + particles, directions, target_quantiles + ) + manifold = prototype_manifold_energy(particles, prototypes) + return relation, conditional, distribution, manifold + + history: list[dict] = [] + initial_cpu = F.normalize(particles.detach(), dim=-1).cpu() + for step in range(args.steps + 1): + relation, conditional, distribution, manifold = energy_values() + progress = min(step / max(args.steps, 1), 1.0) + relation_weight = ( + args.relation_weight_start + + progress + * (args.relation_weight_end - args.relation_weight_start) + ) + conditional_weight = ( + args.conditional_weight_start + + progress + * ( + args.conditional_weight_end + - args.conditional_weight_start + ) + ) + loss = ( + relation_weight * relation + + conditional_weight * conditional + + args.distribution_weight * distribution + + args.manifold_weight * manifold + ) + if step % 20 == 0 or step == args.steps: + history.append( + { + "step": step, + "total": float(loss.detach()), + "relation": float(relation.detach()), + "conditional": float(conditional.detach()), + "distribution": float(distribution.detach()), + "manifold": float(manifold.detach()), + "paired_evaluation_only": retrieval_metrics( + particles.detach(), paired_text + ), + } + ) + print(json.dumps(history[-1])) + if step == args.steps: + break + optimizer.zero_grad(set_to_none=True) + loss.backward() + torch.nn.utils.clip_grad_norm_([particles], 2.0) + optimizer.step() + with torch.no_grad(): + particles.copy_(F.normalize(particles, dim=-1)) + + with torch.no_grad(): + oracle_relation, oracle_conditional = relation_field_energy( + visual_relation, visual_standardized, paired_text + ) + oracle_distribution = sliced_distribution_energy( + paired_text, directions, target_quantiles + ) + oracle_manifold = prototype_manifold_energy( + paired_text, prototypes + ) + result = { + "protocol": ( + "No cross-modal map and no image-text pair is used by the energy " + "or optimizer. Paired text is loaded only for trajectory and " + "oracle diagnostics; the final step is fixed by CLI arguments." + ), + "mode": "free_language_latent_particles", + "split": args.split, + "rows": rows, + "vision_model": vision["model"], + "text_model": text["model"], + "text_observation": ( + "multi-description orbit mean" + if args.text_orbits + else "single description" + ), + "args": vars(args), + "history": history, + "oracle_energy_components": { + "relation": float(oracle_relation), + "conditional": float(oracle_conditional), + "distribution": float(oracle_distribution), + "manifold": float(oracle_manifold), + }, + } + state = { + **result, + "initial_particles": initial_cpu, + "final_particles": F.normalize( + particles.detach(), dim=-1 + ).cpu(), + } + Path(args.output).parent.mkdir(parents=True, exist_ok=True) + torch.save(state, args.output) + write_json(args.metrics_output, result) + print(f"Wrote {args.output} and {args.metrics_output}") + + +if __name__ == "__main__": + main() diff --git a/worldalign/evaluate.py b/worldalign/evaluate.py new file mode 100644 index 0000000..6f97873 --- /dev/null +++ b/worldalign/evaluate.py @@ -0,0 +1,149 @@ +from __future__ import annotations + +import argparse +import json +from pathlib import Path +import re +from collections import Counter + +import torch +from transformers import AutoModelForCausalLM, AutoTokenizer + +from .common import ( + dtype_for_device, + read_json, + retrieval_metrics, + write_json, +) +from .io import load_feature_pair, select_rows +from .models import load_bridge, load_prefix + + +TOKEN_RE = re.compile(r"[a-z0-9]+") + + +def parse_args() -> argparse.Namespace: + p = argparse.ArgumentParser() + p.add_argument("--manifest", default="artifacts/manifest.json") + p.add_argument("--vision", default="artifacts/vision.pt") + p.add_argument("--text", default="artifacts/text.pt") + p.add_argument("--bridge", required=True) + p.add_argument("--prefix") + p.add_argument("--split", choices=["val", "test"], default="test") + p.add_argument("--device", default="cuda:1") + p.add_argument("--generation-samples", type=int, default=100) + p.add_argument("--max-new-tokens", type=int, default=32) + p.add_argument( + "--shuffle-mapped", + action="store_true", + help="Permute image-conditioned latents before retrieval/generation as a null control.", + ) + p.add_argument("--seed", type=int, default=20260728) + p.add_argument("--output", default="artifacts/evaluation.json") + return p.parse_args() + + +def unigram_f1(candidate: str, references: list[str]) -> float: + candidate_tokens = TOKEN_RE.findall(candidate.lower()) + if not candidate_tokens: + return 0.0 + candidate_count = Counter(candidate_tokens) + best = 0.0 + for reference in references: + reference_count = Counter(TOKEN_RE.findall(reference.lower())) + overlap = sum((candidate_count & reference_count).values()) + precision = overlap / max(sum(candidate_count.values()), 1) + recall = overlap / max(sum(reference_count.values()), 1) + f1 = 2 * precision * recall / max(precision + recall, 1e-12) + best = max(best, f1) + return best + + +def main() -> None: + args = parse_args() + manifest = read_json(args.manifest) + vision, text, vlookup, tlookup = load_feature_pair(args.vision, args.text) + rows = manifest[args.split] + x = select_rows(vision["features"], vlookup, rows) + y = select_rows(text["features"], tlookup, rows) + + bridge, bridge_state = load_bridge(args.bridge, args.device) + mapped = [] + with torch.inference_mode(): + for chunk in x.split(512): + mapped.append(bridge(chunk.to(args.device)).cpu()) + mapped = torch.cat(mapped) + if args.shuffle_mapped: + generator = torch.Generator().manual_seed(args.seed) + mapped = mapped[torch.randperm(len(mapped), generator=generator)] + result: dict = { + "split": args.split, + "samples": len(rows), + "bridge_mode": bridge_state["mode"], + "shuffle_mapped": args.shuffle_mapped, + "retrieval": retrieval_metrics(mapped, y), + } + + if args.prefix: + prefix, prefix_state = load_prefix(args.prefix, args.device) + if prefix_state["text_model"] != text["model"]: + raise ValueError("Prefix adapter and text feature model differ") + tokenizer = AutoTokenizer.from_pretrained(text["model"]) + if tokenizer.pad_token_id is None: + tokenizer.pad_token = tokenizer.eos_token + dtype = dtype_for_device(args.device) + lm = AutoModelForCausalLM.from_pretrained( + text["model"], torch_dtype=dtype + ).to(args.device) + lm.eval() + generated: list[str] = [] + n = min(args.generation_samples, len(rows)) + with torch.inference_mode(): + for chunk in mapped[:n].split(16): + prefix_embeds = prefix(chunk.to(args.device)).to(dtype) + attention = torch.ones( + prefix_embeds.shape[:2], + dtype=torch.long, + device=args.device, + ) + output = lm.generate( + inputs_embeds=prefix_embeds, + attention_mask=attention, + max_new_tokens=args.max_new_tokens, + do_sample=False, + eos_token_id=tokenizer.eos_token_id, + pad_token_id=tokenizer.pad_token_id, + ) + generated.extend(tokenizer.batch_decode(output, skip_special_tokens=True)) + + caption_lookup = { + int(row): caps + for row, caps in zip(text["rows"], text["all_captions"]) + } + records = [] + scores = [] + for row, caption in zip(rows[:n], generated): + refs = caption_lookup[int(row)] + score = unigram_f1(caption, refs) + scores.append(score) + records.append( + { + "row": int(row), + "generated": caption, + "references": refs, + "unigram_f1": score, + } + ) + result["generation"] = { + "samples": n, + "mean_best_reference_unigram_f1": sum(scores) / max(len(scores), 1), + "examples": records[:25], + } + + Path(args.output).parent.mkdir(parents=True, exist_ok=True) + write_json(args.output, result) + print(json.dumps(result, indent=2, ensure_ascii=False)) + + +if __name__ == "__main__": + main() diff --git a/worldalign/evaluate_energy_prefix.py b/worldalign/evaluate_energy_prefix.py new file mode 100644 index 0000000..2f3af7a --- /dev/null +++ b/worldalign/evaluate_energy_prefix.py @@ -0,0 +1,151 @@ +from __future__ import annotations + +import argparse +import json +from pathlib import Path + +import numpy as np +import torch +from transformers import AutoModelForCausalLM, AutoTokenizer + +from .common import dtype_for_device, write_json +from .evaluate import unigram_f1 +from .io import load_feature_pair, select_rows +from .models import load_prefix + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--energy", default="artifacts/energy_free_test.pt") + parser.add_argument("--vision", default="artifacts/vision.pt") + parser.add_argument("--text", default="artifacts/text.pt") + parser.add_argument("--text-orbits") + parser.add_argument("--prefix", default="artifacts/prefix.pt") + parser.add_argument("--samples", type=int, default=100) + parser.add_argument("--max-new-tokens", type=int, default=32) + parser.add_argument("--device", default="cuda:3") + parser.add_argument( + "--output", default="artifacts/energy_prefix_evaluation.json" + ) + return parser.parse_args() + + +def main() -> None: + args = parse_args() + energy = torch.load(args.energy, map_location="cpu", weights_only=False) + _, text, _, text_lookup = load_feature_pair(args.vision, args.text) + rows = energy["rows"][: args.samples] + initial = energy["initial_particles"][: args.samples] + final = energy["final_particles"][: args.samples] + oracle = select_rows(text["features"], text_lookup, rows) + if args.text_orbits: + orbit_state = torch.load( + args.text_orbits, map_location="cpu", weights_only=False + ) + orbit_lookup = { + int(row): index + for index, row in enumerate(orbit_state["rows"]) + } + orbit_mean = torch.nn.functional.normalize( + orbit_state["features"].float().mean(1), dim=-1 + ) + oracle = select_rows(orbit_mean, orbit_lookup, rows) + reference_lookup = { + int(row): captions + for row, captions in zip(text["rows"], text["all_captions"]) + } + + prefix, prefix_state = load_prefix(args.prefix, args.device) + if prefix.semantic_dim != initial.shape[-1]: + raise ValueError("Energy latent and text-only prefix dimensions differ") + tokenizer = AutoTokenizer.from_pretrained(text["model"]) + if tokenizer.pad_token_id is None: + tokenizer.pad_token = tokenizer.eos_token + dtype = dtype_for_device(args.device) + lm = AutoModelForCausalLM.from_pretrained( + text["model"], torch_dtype=dtype + ).to(args.device) + lm.eval() + + conditions = {"initial": initial, "final": final, "oracle": oracle} + decoded: dict[str, list[str]] = {key: [] for key in conditions} + with torch.inference_mode(): + for name, semantic in conditions.items(): + for chunk in semantic.split(16): + embeds = prefix(chunk.to(args.device)).to(dtype) + attention = torch.ones( + embeds.shape[:2], dtype=torch.long, device=args.device + ) + output = lm.generate( + inputs_embeds=embeds, + attention_mask=attention, + max_new_tokens=args.max_new_tokens, + do_sample=False, + eos_token_id=tokenizer.eos_token_id, + pad_token_id=tokenizer.pad_token_id, + ) + decoded[name].extend( + tokenizer.batch_decode( + output, skip_special_tokens=True + ) + ) + + scores = {} + per_condition: dict[str, list[float]] = {} + for name, generations in decoded.items(): + values = [ + unigram_f1(generation, reference_lookup[int(row)]) + for row, generation in zip(rows, generations) + ] + per_condition[name] = values + scores[name] = sum(values) / max(len(values), 1) + difference = np.asarray(per_condition["final"]) - np.asarray( + per_condition["initial"] + ) + bootstrap_generator = np.random.default_rng(20260729) + bootstrap = difference[ + bootstrap_generator.integers( + 0, len(difference), size=(10_000, len(difference)) + ) + ].mean(1) + result = { + "samples": len(rows), + "mean_best_reference_unigram_f1": scores, + "paired_final_minus_initial": { + "mean": float(difference.mean()), + "bootstrap_95_percentile_interval": [ + float(np.quantile(bootstrap, 0.025)), + float(np.quantile(bootstrap, 0.975)), + ], + "improved": int((difference > 0).sum()), + "tied": int((difference == 0).sum()), + "worsened": int((difference < 0).sum()), + }, + "energy_protocol": energy["protocol"], + "prefix_training": prefix_state["training"], + "per_sample_unigram_f1": [ + { + "row": int(row), + **{ + name: per_condition[name][index] + for name in per_condition + }, + } + for index, row in enumerate(rows) + ], + "examples": [ + { + "row": int(row), + "references": reference_lookup[int(row)], + **{name: decoded[name][i] for name in decoded}, + } + for i, row in enumerate(rows[:25]) + ], + } + Path(args.output).parent.mkdir(parents=True, exist_ok=True) + write_json(args.output, result) + print(json.dumps(result, indent=2, ensure_ascii=False)) + + +if __name__ == "__main__": + main() diff --git a/worldalign/evaluate_prefix.py b/worldalign/evaluate_prefix.py new file mode 100644 index 0000000..2d53a51 --- /dev/null +++ b/worldalign/evaluate_prefix.py @@ -0,0 +1,93 @@ +from __future__ import annotations + +import argparse +import json +from pathlib import Path + +import torch +from transformers import AutoModelForCausalLM, AutoTokenizer + +from .common import dtype_for_device, read_json, write_json +from .evaluate import unigram_f1 +from .io import load_feature_pair, select_rows +from .models import load_prefix + + +def parse_args() -> argparse.Namespace: + p = argparse.ArgumentParser() + p.add_argument("--manifest", default="artifacts/manifest.json") + p.add_argument("--vision", default="artifacts/vision.pt") + p.add_argument("--text", default="artifacts/text.pt") + p.add_argument("--prefix", default="artifacts/prefix.pt") + p.add_argument("--split", choices=["val", "test"], default="val") + p.add_argument("--device", default="cuda:1") + p.add_argument("--samples", type=int, default=100) + p.add_argument("--max-new-tokens", type=int, default=32) + p.add_argument("--output", default="artifacts/prefix_evaluation.json") + return p.parse_args() + + +def main() -> None: + args = parse_args() + manifest = read_json(args.manifest) + _, text, _, tlookup = load_feature_pair(args.vision, args.text) + rows = manifest[args.split][: args.samples] + semantic = select_rows(text["features"], tlookup, rows) + caption_lookup = { + int(row): caps for row, caps in zip(text["rows"], text["all_captions"]) + } + prefix, prefix_state = load_prefix(args.prefix, args.device) + tokenizer = AutoTokenizer.from_pretrained(text["model"]) + if tokenizer.pad_token_id is None: + tokenizer.pad_token = tokenizer.eos_token + dtype = dtype_for_device(args.device) + lm = AutoModelForCausalLM.from_pretrained( + text["model"], torch_dtype=dtype + ).to(args.device) + lm.eval() + + generated = [] + with torch.inference_mode(): + for chunk in semantic.split(16): + embeds = prefix(chunk.to(args.device)).to(dtype) + attention = torch.ones( + embeds.shape[:2], dtype=torch.long, device=args.device + ) + output = lm.generate( + inputs_embeds=embeds, + attention_mask=attention, + max_new_tokens=args.max_new_tokens, + do_sample=False, + eos_token_id=tokenizer.eos_token_id, + pad_token_id=tokenizer.pad_token_id, + ) + generated.extend(tokenizer.batch_decode(output, skip_special_tokens=True)) + + records = [] + for row, caption in zip(rows, generated): + references = caption_lookup[int(row)] + records.append( + { + "row": int(row), + "generated": caption, + "references": references, + "unigram_f1": unigram_f1(caption, references), + } + ) + result = { + "split": args.split, + "samples": len(records), + "mean_best_reference_unigram_f1": sum( + r["unigram_f1"] for r in records + ) + / max(len(records), 1), + "examples": records[:25], + "training": prefix_state["training"], + } + Path(args.output).parent.mkdir(parents=True, exist_ok=True) + write_json(args.output, result) + print(json.dumps(result, indent=2, ensure_ascii=False)) + + +if __name__ == "__main__": + main() diff --git a/worldalign/extract_text.py b/worldalign/extract_text.py new file mode 100644 index 0000000..8586e21 --- /dev/null +++ b/worldalign/extract_text.py @@ -0,0 +1,95 @@ +from __future__ import annotations + +import argparse +from pathlib import Path + +import torch +from datasets import load_dataset +from tqdm import tqdm +from transformers import AutoModel, AutoTokenizer + +from .common import batch_indices, dtype_for_device, read_json + + +def parse_args() -> argparse.Namespace: + p = argparse.ArgumentParser() + p.add_argument("--manifest", default="artifacts/manifest.json") + p.add_argument("--output", default="artifacts/text.pt") + p.add_argument("--model", default="Qwen/Qwen2.5-0.5B") + p.add_argument("--device", default="cuda:3") + p.add_argument("--batch-size", type=int, default=96) + p.add_argument("--max-length", type=int, default=64) + p.add_argument( + "--layer", + type=int, + default=-1, + help="Hidden-state index; -1 is the final transformer output.", + ) + p.add_argument("--limit", type=int) + return p.parse_args() + + +def mean_pool(hidden: torch.Tensor, mask: torch.Tensor) -> torch.Tensor: + weights = mask.to(hidden.dtype).unsqueeze(-1) + return (hidden * weights).sum(1) / weights.sum(1).clamp_min(1) + + +@torch.inference_mode() +def main() -> None: + args = parse_args() + manifest = read_json(args.manifest) + rows = manifest["all_rows"] + if args.limit: + rows = rows[: args.limit] + dataset = load_dataset(manifest["dataset"], split=manifest["dataset_split"]) + # Avoid decoding the 4.3GB image column when only captions are requested. + dataset = dataset.remove_columns("image") + tokenizer = AutoTokenizer.from_pretrained(args.model) + if tokenizer.pad_token_id is None: + tokenizer.pad_token = tokenizer.eos_token + tokenizer.padding_side = "right" + dtype = dtype_for_device(args.device) + # The backbone is sufficient for hidden states. Using a causal-LM wrapper + # would also materialize vocabulary logits that are discarded here. + model = AutoModel.from_pretrained(args.model, torch_dtype=dtype).to(args.device) + model.eval() + + outputs: list[torch.Tensor] = [] + captions: list[str] = [] + all_captions: list[list[str]] = [] + for ids in tqdm( + batch_indices(len(rows), args.batch_size), desc="Qwen text features" + ): + batch_captions = [dataset[int(rows[i])]["caption"] for i in ids] + texts = [caps[0] for caps in batch_captions] + tokens = tokenizer( + texts, + padding=True, + truncation=True, + max_length=args.max_length, + return_tensors="pt", + ) + tokens = {k: v.to(args.device) for k, v in tokens.items()} + result = model(**tokens, output_hidden_states=True, return_dict=True) + feature = mean_pool( + result.hidden_states[args.layer], tokens["attention_mask"] + ) + outputs.append(feature.float().cpu()) + captions.extend(texts) + all_captions.extend(batch_captions) + + value = { + "model": args.model, + "layer": args.layer, + "rows": rows, + "features": torch.cat(outputs), + "captions": captions, + "all_captions": all_captions, + } + Path(args.output).parent.mkdir(parents=True, exist_ok=True) + torch.save(value, args.output) + print(f"Wrote {args.output}: {tuple(value['features'].shape)}") + + +if __name__ == "__main__": + main() diff --git a/worldalign/extract_text_orbits.py b/worldalign/extract_text_orbits.py new file mode 100644 index 0000000..43b9020 --- /dev/null +++ b/worldalign/extract_text_orbits.py @@ -0,0 +1,108 @@ +from __future__ import annotations + +import argparse +from pathlib import Path + +from datasets import load_dataset +import torch +from tqdm import tqdm +from transformers import AutoModel, AutoTokenizer + +from .common import batch_indices, dtype_for_device, read_json +from .extract_text import mean_pool + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--manifest", default="artifacts/manifest.json") + parser.add_argument( + "--output", default="artifacts/text_orbits_qwen0p5b.pt" + ) + parser.add_argument("--model", default="Qwen/Qwen2.5-0.5B") + parser.add_argument("--device", default="cuda:3") + parser.add_argument("--batch-size", type=int, default=128) + parser.add_argument("--max-length", type=int, default=64) + parser.add_argument("--layer", type=int, default=-1) + parser.add_argument( + "--row-groups", + default="text_only_train,val,test", + help="Comma-separated manifest row lists to encode.", + ) + parser.add_argument("--limit", type=int) + return parser.parse_args() + + +@torch.inference_mode() +def main() -> None: + args = parse_args() + manifest = read_json(args.manifest) + groups = [group.strip() for group in args.row_groups.split(",")] + rows = list( + dict.fromkeys( + int(row) + for group in groups + for row in manifest[group] + ) + ) + if args.limit: + rows = rows[: args.limit] + dataset = load_dataset( + manifest["dataset"], split=manifest["dataset_split"] + ).remove_columns("image") + captions = [dataset[row]["caption"] for row in rows] + views = len(captions[0]) + if any(len(items) != views for items in captions): + raise ValueError("Every text orbit must have the same view count") + flat_text = [text for items in captions for text in items] + + tokenizer = AutoTokenizer.from_pretrained(args.model) + if tokenizer.pad_token_id is None: + tokenizer.pad_token = tokenizer.eos_token + tokenizer.padding_side = "right" + dtype = dtype_for_device(args.device) + model = AutoModel.from_pretrained( + args.model, torch_dtype=dtype + ).to(args.device) + model.eval() + + output: list[torch.Tensor] = [] + for indices in tqdm( + batch_indices(len(flat_text), args.batch_size), + desc="Qwen text orbits", + ): + tokens = tokenizer( + [flat_text[index] for index in indices], + padding=True, + truncation=True, + max_length=args.max_length, + return_tensors="pt", + ) + tokens = {key: value.to(args.device) for key, value in tokens.items()} + result = model( + **tokens, output_hidden_states=True, return_dict=True + ) + output.append( + mean_pool( + result.hidden_states[args.layer], + tokens["attention_mask"], + ) + .float() + .cpu() + ) + features = torch.cat(output).reshape(len(rows), views, -1) + state = { + "model": args.model, + "layer": args.layer, + "rows": rows, + "features": features, + "captions": captions, + "views_per_orbit": views, + "row_groups": groups, + } + Path(args.output).parent.mkdir(parents=True, exist_ok=True) + torch.save(state, args.output) + print(f"Wrote {args.output}: {tuple(features.shape)}") + + +if __name__ == "__main__": + main() diff --git a/worldalign/extract_vision.py b/worldalign/extract_vision.py new file mode 100644 index 0000000..348c6a4 --- /dev/null +++ b/worldalign/extract_vision.py @@ -0,0 +1,61 @@ +from __future__ import annotations + +import argparse +from pathlib import Path + +import torch +from datasets import load_dataset +from tqdm import tqdm +from transformers import AutoImageProcessor, AutoModel + +from .common import batch_indices, dtype_for_device, read_json + + +def parse_args() -> argparse.Namespace: + p = argparse.ArgumentParser() + p.add_argument("--manifest", default="artifacts/manifest.json") + p.add_argument("--output", default="artifacts/vision.pt") + p.add_argument("--model", default="facebook/dinov2-small") + p.add_argument("--device", default="cuda:1") + p.add_argument("--batch-size", type=int, default=96) + p.add_argument("--limit", type=int) + return p.parse_args() + + +@torch.inference_mode() +def main() -> None: + args = parse_args() + manifest = read_json(args.manifest) + rows = manifest["all_rows"] + if args.limit: + rows = rows[: args.limit] + dataset = load_dataset(manifest["dataset"], split=manifest["dataset_split"]) + processor = AutoImageProcessor.from_pretrained(args.model) + dtype = dtype_for_device(args.device) + model = AutoModel.from_pretrained(args.model, torch_dtype=dtype).to(args.device) + model.eval() + + outputs: list[torch.Tensor] = [] + for ids in tqdm( + batch_indices(len(rows), args.batch_size), desc="DINO image features" + ): + images = [dataset[int(rows[i])]["image"].convert("RGB") for i in ids] + batch = processor(images=images, return_tensors="pt") + batch = {k: v.to(args.device) for k, v in batch.items()} + result = model(**batch) + feature = result.last_hidden_state[:, 0] + outputs.append(feature.float().cpu()) + + value = { + "model": args.model, + "rows": rows, + "features": torch.cat(outputs), + } + Path(args.output).parent.mkdir(parents=True, exist_ok=True) + torch.save(value, args.output) + print(f"Wrote {args.output}: {tuple(value['features'].shape)}") + + +if __name__ == "__main__": + main() + diff --git a/worldalign/gw.py b/worldalign/gw.py new file mode 100644 index 0000000..62c3f5c --- /dev/null +++ b/worldalign/gw.py @@ -0,0 +1,79 @@ +from __future__ import annotations + +import numpy as np +import ot +import torch +from sklearn.cluster import MiniBatchKMeans + +from .common import pairwise_cosine_distance + + +def gw_pseudo_targets( + vision: torch.Tensor, + text: torch.Tensor, + clusters: int, + seed: int, + max_iter: int = 100, +) -> dict[str, torch.Tensor | float | int]: + """Build structure-only pseudo-correspondences between modality prototypes.""" + x = vision.float().numpy() + y = text.float().numpy() + k = min(clusters, len(x), len(y)) + vx = MiniBatchKMeans( + n_clusters=k, + random_state=seed, + batch_size=min(4096, len(x)), + n_init=3, + max_iter=200, + ).fit(x) + ty = MiniBatchKMeans( + n_clusters=k, + random_state=seed + 1, + batch_size=min(4096, len(y)), + n_init=3, + max_iter=200, + ).fit(y) + + cx = pairwise_cosine_distance(vx.cluster_centers_) + cy = pairwise_cosine_distance(ty.cluster_centers_) + p = np.bincount(vx.labels_, minlength=k).astype(np.float64) + q = np.bincount(ty.labels_, minlength=k).astype(np.float64) + p /= p.sum() + q /= q.sum() + + coupling, log = ot.gromov.gromov_wasserstein( + cx, + cy, + p, + q, + loss_fun="square_loss", + armijo=False, + log=True, + max_iter=max_iter, + tol_rel=1e-8, + tol_abs=1e-8, + ) + target_centers = coupling @ ty.cluster_centers_ + target_centers /= p[:, None].clip(min=1e-12) + target_centers /= np.linalg.norm(target_centers, axis=1, keepdims=True).clip( + min=1e-12 + ) + row_entropy = -np.sum( + (coupling / p[:, None].clip(min=1e-12)) + * np.log( + (coupling / p[:, None].clip(min=1e-12)).clip(min=1e-12) + ), + axis=1, + ).mean() + + return { + "vision_centers": torch.from_numpy(vx.cluster_centers_).float(), + "text_centers": torch.from_numpy(ty.cluster_centers_).float(), + "target_centers": torch.from_numpy(target_centers).float(), + "vision_assignments": torch.from_numpy(vx.labels_).long(), + "coupling": torch.from_numpy(coupling).float(), + "gw_distance": float(log["gw_dist"]), + "coupling_row_entropy": float(row_entropy), + "clusters": k, + } + diff --git a/worldalign/io.py b/worldalign/io.py new file mode 100644 index 0000000..6cac007 --- /dev/null +++ b/worldalign/io.py @@ -0,0 +1,30 @@ +from __future__ import annotations + +import torch + +from .common import normalized + + +def load_feature_pair( + vision_path: str, text_path: str +) -> tuple[dict, dict, dict[int, int], dict[int, int]]: + vision = torch.load(vision_path, map_location="cpu", weights_only=False) + text = torch.load(text_path, map_location="cpu", weights_only=False) + vision["features"] = normalized(vision["features"]) + text["features"] = normalized(text["features"]) + vision_lookup = {int(row): i for i, row in enumerate(vision["rows"])} + text_lookup = {int(row): i for i, row in enumerate(text["rows"])} + return vision, text, vision_lookup, text_lookup + + +def select_rows( + feature: torch.Tensor, lookup: dict[int, int], rows: list[int] +) -> torch.Tensor: + missing = [int(row) for row in rows if int(row) not in lookup] + if missing: + raise KeyError( + f"{len(missing)} requested rows are absent from feature cache; " + f"first missing rows: {missing[:5]}" + ) + return feature[[lookup[int(row)] for row in rows]] + 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() diff --git a/worldalign/models.py b/worldalign/models.py new file mode 100644 index 0000000..5f98605 --- /dev/null +++ b/worldalign/models.py @@ -0,0 +1,90 @@ +from __future__ import annotations + +import torch +from torch import nn +import torch.nn.functional as F + + +class Bridge(nn.Module): + def __init__( + self, + vision_dim: int, + text_dim: int, + hidden_dim: int = 1536, + linear: bool = False, + ): + super().__init__() + self.vision_dim = vision_dim + self.text_dim = text_dim + self.hidden_dim = hidden_dim + self.linear = linear + if linear: + self.network = nn.Linear(vision_dim, text_dim, bias=False) + else: + self.network = nn.Sequential( + nn.LayerNorm(vision_dim), + nn.Linear(vision_dim, hidden_dim), + nn.GELU(), + nn.Linear(hidden_dim, text_dim), + ) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + return F.normalize(self.network(x).float(), dim=-1) + + def config(self) -> dict: + return { + "vision_dim": self.vision_dim, + "text_dim": self.text_dim, + "hidden_dim": self.hidden_dim, + "linear": self.linear, + } + + +class PrefixAdapter(nn.Module): + def __init__( + self, + semantic_dim: int, + lm_dim: int, + prefix_length: int = 8, + hidden_dim: int = 2048, + ): + super().__init__() + self.semantic_dim = semantic_dim + self.lm_dim = lm_dim + self.prefix_length = prefix_length + self.hidden_dim = hidden_dim + self.network = nn.Sequential( + nn.LayerNorm(semantic_dim), + nn.Linear(semantic_dim, hidden_dim), + nn.GELU(), + nn.Linear(hidden_dim, prefix_length * lm_dim), + ) + + def forward(self, semantic: torch.Tensor) -> torch.Tensor: + prefix = self.network(semantic.float()) + return prefix.view(-1, self.prefix_length, self.lm_dim) + + def config(self) -> dict: + return { + "semantic_dim": self.semantic_dim, + "lm_dim": self.lm_dim, + "prefix_length": self.prefix_length, + "hidden_dim": self.hidden_dim, + } + + +def load_bridge(path: str, device: str = "cpu") -> tuple[Bridge, dict]: + state = torch.load(path, map_location="cpu", weights_only=False) + model = Bridge(**state["config"]) + model.load_state_dict(state["state_dict"]) + model.to(device).eval() + return model, state + + +def load_prefix(path: str, device: str = "cpu") -> tuple[PrefixAdapter, dict]: + state = torch.load(path, map_location="cpu", weights_only=False) + model = PrefixAdapter(**state["config"]) + model.load_state_dict(state["state_dict"]) + model.to(device).eval() + return model, state + diff --git a/worldalign/natural_families.py b/worldalign/natural_families.py new file mode 100644 index 0000000..a13a1d9 --- /dev/null +++ b/worldalign/natural_families.py @@ -0,0 +1,264 @@ +"""Do factor families survive real language? + +The synthetic derivation found colour, count, and size families because +templated phrases place exactly one member of each family in a fixed +slot, so family members never co-occur. Real descriptions break every +part of that: free word order, stacked adjectives, synonyms, and phrases +that mention no attribute at all. This measures how much of the +mutual-exclusivity signal survives, on Visual Genome region descriptions, +using no lexicon and no labels. + +Families are recovered as low-co-occurrence, high-context-similarity +groups: two words of one factor rarely modify the same head, and when +they do appear they appear in the same distributional company. That is +the paradigmatic relation of distributional semantics, and the greedy +exclusivity pass of the synthetic pipeline is its degenerate case. +""" + +from __future__ import annotations + +import argparse +import json +import re +from collections import Counter, defaultdict +from pathlib import Path + +import numpy as np + +from .common import read_json, seed_everything, write_json + +# Reference sets, used only to score the recovered families, never to find +# them. Membership is checked after the fact. +REFERENCE = { + "colour": { + "black", "white", "red", "blue", "green", "yellow", "brown", "gray", + "grey", "orange", "purple", "pink", "tan", "beige", "silver", "gold", + "golden", "dark", "light", + }, + "number": { + "one", "two", "three", "four", "five", "six", "seven", "eight", + "nine", "ten", "a", "an", "the", "some", "many", + }, + "size": {"small", "large", "big", "little", "tiny", "huge", "tall", "short", "long"}, + "material": {"wooden", "metal", "plastic", "glass", "brick", "stone", "leather"}, +} + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--nodes", default="artifacts/vg_5k/text_nodes.jsonl") + parser.add_argument("--field", default="region_closed") + parser.add_argument("--min-count", type=int, default=200) + parser.add_argument("--max-vocabulary", type=int, default=300) + parser.add_argument("--context-window", type=int, default=2) + parser.add_argument("--exclusivity", type=float, default=0.25, + help="Maximum observed-over-expected co-occurrence for " + "two words to count as mutually exclusive.") + parser.add_argument( + "--method", + choices=["exclusivity", "distributional"], + default="distributional", + help="Mutual exclusivity is a template artifact; real paradigms are " + "found by distributional similarity, of which it is a special case.", + ) + parser.add_argument("--components", type=int, default=64) + parser.add_argument("--clusters", type=int, default=30) + parser.add_argument("--min-family", type=int, default=3) + parser.add_argument("--seed", type=int, default=0) + parser.add_argument( + "--output", default="artifacts/vg_5k/natural_families.json" + ) + return parser.parse_args() + + +def load_phrases(path: Path, field: str) -> list[list[str]]: + phrases = [] + for line in path.read_text(encoding="utf-8").splitlines(): + if not line.strip(): + continue + record = json.loads(line) + for phrase in record[field]: + tokens = re.findall(r"[a-z]+", phrase.lower()) + if tokens: + phrases.append(tokens) + return phrases + + +def statistics( + phrases: list[list[str]], vocabulary: set[str], window: int +) -> tuple[Counter, Counter, dict[str, Counter]]: + unigram: Counter = Counter() + pair: Counter = Counter() + context: dict[str, Counter] = defaultdict(Counter) + for tokens in phrases: + present = [token for token in tokens if token in vocabulary] + for token in set(present): + unigram[token] += 1 + for first in set(present): + for second in set(present): + if first < second: + pair[(first, second)] += 1 + for index, token in enumerate(tokens): + if token not in vocabulary: + continue + for offset in range(1, window + 1): + for neighbour_index in (index - offset, index + offset): + if 0 <= neighbour_index < len(tokens): + context[token][tokens[neighbour_index]] += 1 + return unigram, pair, context + + +def exclusivity_ratio( + first: str, second: str, unigram: Counter, pair: Counter, total: int +) -> float: + """Observed co-occurrence over the independent expectation.""" + expected = unigram[first] * unigram[second] / max(total, 1) + key = (first, second) if first < second else (second, first) + return pair[key] / max(expected, 1e-9) + + +def context_similarity(first: Counter, second: Counter) -> float: + keys = set(first) | set(second) + a = np.array([first[k] for k in keys], dtype=float) + b = np.array([second[k] for k in keys], dtype=float) + a /= max(np.linalg.norm(a), 1e-9) + b /= max(np.linalg.norm(b), 1e-9) + return float(a @ b) + + +def build_families( + words: list[str], + unigram: Counter, + pair: Counter, + context: dict[str, Counter], + total: int, + args: argparse.Namespace, +) -> list[list[str]]: + """Greedy paradigmatic grouping: exclusive and distributionally alike.""" + families: list[list[str]] = [] + for word in words: + best_family, best_score = None, 0.0 + for family in families: + ratios = [ + exclusivity_ratio(word, member, unigram, pair, total) + for member in family + ] + if max(ratios) > args.exclusivity: + continue + similarity = float( + np.mean( + [context_similarity(context[word], context[member]) + for member in family] + ) + ) + if similarity > best_score: + best_family, best_score = family, similarity + if best_family is not None and best_score > 0.15: + best_family.append(word) + else: + families.append([word]) + return [family for family in families if len(family) >= args.min_family] + + +def distributional_families( + words: list[str], context: dict[str, Counter], args: argparse.Namespace +) -> list[list[str]]: + """Word classes from context vectors: positive PMI, truncated SVD, k-means. + + The standard recipe of distributional word-class induction. Words of + one factor modify the same heads and so share company, whether or not + they exclude each other -- real colour terms co-occur freely. + """ + from sklearn.cluster import KMeans + + features = sorted({key for word in words for key in context[word]}) + index = {key: position for position, key in enumerate(features)} + matrix = np.zeros((len(words), len(features))) + for row, word in enumerate(words): + for key, value in context[word].items(): + matrix[row, index[key]] = value + total = matrix.sum() + row_sum = matrix.sum(1, keepdims=True) + column_sum = matrix.sum(0, keepdims=True) + expected = row_sum * column_sum / max(total, 1e-9) + pmi = np.log(np.maximum(matrix, 1e-9) / np.maximum(expected, 1e-9)) + pmi[matrix == 0] = 0.0 + pmi = np.maximum(pmi, 0.0) + left, values, _ = np.linalg.svd(pmi, full_matrices=False) + embedding = left[:, : args.components] * values[: args.components] + embedding /= np.linalg.norm(embedding, axis=1, keepdims=True).clip(1e-9) + labels = KMeans(args.clusters, n_init=10, random_state=args.seed).fit_predict( + embedding + ) + families: dict[int, list[str]] = defaultdict(list) + for word, label in zip(words, labels): + families[int(label)].append(word) + return [family for family in families.values() if len(family) >= args.min_family] + + +def label_family(family: list[str]) -> tuple[str, float]: + best_name, best_purity = "unlabelled", 0.0 + for name, reference in REFERENCE.items(): + purity = sum(1 for word in family if word in reference) / len(family) + if purity > best_purity: + best_name, best_purity = name, purity + return best_name, best_purity + + +def main() -> None: + args = parse_args() + seed_everything(args.seed) + phrases = load_phrases(Path(args.nodes), args.field) + counts = Counter(token for tokens in phrases for token in set(tokens)) + vocabulary = [ + word + for word, count in counts.most_common(args.max_vocabulary) + if count >= args.min_count + ] + unigram, pair, context = statistics(phrases, set(vocabulary), args.context_window) + if args.method == "distributional": + families = distributional_families(vocabulary, context, args) + else: + families = build_families( + vocabulary, unigram, pair, context, len(phrases), args + ) + + described = [] + for family in sorted(families, key=len, reverse=True): + name, purity = label_family(family) + described.append( + { + "size": len(family), + "label": name, + "purity": purity, + "words": sorted(family, key=lambda w: -unigram[w])[:14], + } + ) + report = { + "protocol": ( + "Families are recovered from region descriptions alone by " + "mutual exclusivity plus distributional similarity. Reference " + "word sets are read only to label the result." + ), + "phrases": len(phrases), + "vocabulary": len(vocabulary), + "families": described, + "recovered_labels": Counter(item["label"] for item in described), + } + for item in described[:12]: + print( + json.dumps( + { + "label": item["label"], + "purity": round(item["purity"], 2), + "size": item["size"], + "words": item["words"][:10], + } + ) + ) + write_json(args.output, report) + print(f"Wrote {args.output}") + + +if __name__ == "__main__": + main() diff --git a/worldalign/natural_fields.py b/worldalign/natural_fields.py new file mode 100644 index 0000000..7221a3a --- /dev/null +++ b/worldalign/natural_fields.py @@ -0,0 +1,224 @@ +"""Natural-data relation fields under the Tier 0 recipe. + +Each side is encoded into its own discovered factor coordinates -- text +word classes from distributional induction, vision segment classes from +feature clustering -- and the two coordinate systems are paired by joint +structure, never by a declared lexicon. Scene states are the resulting +sets of segment or phrase codes; relations are moment-kernel similarities +between scenes. + +The field correlation at the hidden pairing is the go/no-go statistic: +polynomial recovery needs roughly 0.9, and the synthetic world showed +nothing works below it. Hidden pairs are read only to compute it. +""" + +from __future__ import annotations + +import argparse +import json +import re +from collections import Counter, defaultdict +from pathlib import Path + +import numpy as np +import torch +import torch.nn.functional as F +from scipy.optimize import linear_sum_assignment +from sklearn.cluster import KMeans + +from .common import read_json, seed_everything, write_json +from .natural_families import distributional_families, load_phrases, statistics +from .synth_cc_battery import moment_field + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--vg-dir", default="artifacts/vg_5k") + parser.add_argument("--objects", default="artifacts/vg_5k/natural_objects.pt") + parser.add_argument("--samples", type=int, default=256) + parser.add_argument("--fit-nodes", type=int, default=2000) + parser.add_argument("--vision-classes", type=int, default=24) + parser.add_argument("--text-classes", type=int, default=24) + parser.add_argument("--clusters", type=int, default=24) + parser.add_argument("--min-count", type=int, default=200) + parser.add_argument("--max-vocabulary", type=int, default=300) + parser.add_argument("--context-window", type=int, default=2) + parser.add_argument("--min-family", type=int, default=3) + parser.add_argument("--components", type=int, default=64) + parser.add_argument("--seed", type=int, default=0) + parser.add_argument("--output", default="artifacts/vg_5k/natural_fields.pt") + return parser.parse_args() + + +def text_scene_codes( + records: list[dict], word_class: dict[str, int], classes: int +) -> dict[str, torch.Tensor]: + """One code vector per region description: its word-class profile.""" + codes: dict[str, list[torch.Tensor]] = {} + for record in records: + vectors = [] + for phrase in record["region_closed"]: + vector = torch.zeros(classes) + for token in re.findall(r"[a-z]+", phrase.lower()): + if token in word_class: + vector[word_class[token]] += 1.0 + if vector.sum() > 0: + vectors.append(vector) + if vectors: + codes[record["node_id"]] = F.normalize(torch.stack(vectors), dim=-1) + return codes + + +def vision_scene_codes( + state: dict, model: KMeans, classes: int +) -> dict[str, torch.Tensor]: + """One code vector per segment: its feature-class assignment.""" + codes: dict[str, torch.Tensor] = {} + for node, segments in zip(state["node_ids"], state["segments"]): + if not segments: + continue + labels = model.predict( + np.stack([segment["feature"] for segment in segments]).astype(np.float64) + ) + vectors = torch.zeros(len(segments), classes) + for row, label in enumerate(labels): + vectors[row, int(label)] = 1.0 + codes[node] = F.normalize(vectors, dim=-1) + return codes + + +def align_classes( + text_codes: dict[str, torch.Tensor], + vision_codes: dict[str, torch.Tensor], + classes: int, +) -> np.ndarray: + """Pair vision classes to text classes by marginal frequency rank. + + Within-scene co-occurrence is the stronger signal but needs a second + factor to condition on; frequency is the available unimodal statistic + at this stage and is reported as the first pass. + """ + text_mass = torch.zeros(classes) + for code in text_codes.values(): + text_mass += code.sum(0) + vision_mass = torch.zeros(classes) + for code in vision_codes.values(): + vision_mass += code.sum(0) + text_order = torch.argsort(text_mass, descending=True).numpy() + vision_order = torch.argsort(vision_mass, descending=True).numpy() + mapping = np.empty(classes, dtype=int) + mapping[vision_order] = text_order + return mapping + + +def moment_states(codes: torch.Tensor) -> torch.Tensor: + first = codes.mean(0) + second = (codes[:, :, None] * codes[:, None, :]).mean(0).flatten() + return torch.cat([first, second]) + + +def main() -> None: + args = parse_args() + seed_everything(args.seed) + vg_dir = Path(args.vg_dir) + text_records = [ + json.loads(line) + for line in (vg_dir / "text_nodes.jsonl").read_text(encoding="utf-8").splitlines() + if line.strip() + ] + truth = [ + json.loads(line) + for line in (vg_dir / "ground_truth.private.jsonl") + .read_text(encoding="utf-8") + .splitlines() + if line.strip() + ] + state = torch.load(args.objects, map_location="cpu", weights_only=False) + + phrases = load_phrases(vg_dir / "text_nodes.jsonl", "region_closed") + counts = Counter(token for tokens in phrases for token in set(tokens)) + vocabulary = [ + word + for word, count in counts.most_common(args.max_vocabulary) + if count >= args.min_count + ] + _, _, context = statistics(phrases, set(vocabulary), args.context_window) + args.clusters = args.text_classes + families = distributional_families(vocabulary, context, args) + word_class = { + word: index for index, family in enumerate(families) for word in family + } + text_classes = len(families) + + features = np.stack( + [ + segment["feature"] + for segments in state["segments"][: args.fit_nodes] + for segment in segments + ] + ).astype(np.float64) + vision_model = KMeans( + args.vision_classes, n_init=10, random_state=args.seed + ).fit(features) + + text_codes = text_scene_codes(text_records, word_class, text_classes) + vision_codes = vision_scene_codes(state, vision_model, args.vision_classes) + + shared = min(text_classes, args.vision_classes) + mapping = align_classes( + {k: v[:, :shared] for k, v in text_codes.items()}, + {k: v[:, :shared] for k, v in vision_codes.items()}, + shared, + ) + + pairs = [ + pair + for pair in truth + if pair["vision_node_id"] in vision_codes and pair["text_node_id"] in text_codes + ][: args.samples] + vision_sets, text_sets = [], [] + for pair in pairs: + vision = vision_codes[pair["vision_node_id"]][:, :shared] + remapped = torch.zeros_like(vision) + for source in range(shared): + remapped[:, mapping[source]] = vision[:, source] + vision_sets.append(F.normalize(remapped, dim=-1)) + text_sets.append( + F.normalize(text_codes[pair["text_node_id"]][:, :shared], dim=-1) + ) + + visual_field = moment_field(vision_sets) + text_field = moment_field(text_sets) + size = len(pairs) + mask = ~np.eye(size, dtype=bool) + correlation = float( + np.corrcoef( + visual_field.double().numpy()[mask], text_field.double().numpy()[mask] + )[0, 1] + ) + torch.save( + { + "visual_field": visual_field, + "text_field": text_field, + "pairs": [p["vision_node_id"] for p in pairs], + }, + args.output, + ) + summary = { + "protocol": ( + "Text classes from distributional induction, vision classes " + "from segment-feature clustering, paired by marginal frequency " + "rank. Hidden pairs are read only for the correlation." + ), + "samples": size, + "text_classes": text_classes, + "vision_classes": args.vision_classes, + "field_correlation_at_truth": correlation, + "go_no_go": "recovery needs about 0.9", + } + print(json.dumps(summary)) + write_json(str(args.output).replace(".pt", ".json"), summary) + + +if __name__ == "__main__": + main() diff --git a/worldalign/natural_objects.py b/worldalign/natural_objects.py new file mode 100644 index 0000000..65e035b --- /dev/null +++ b/worldalign/natural_objects.py @@ -0,0 +1,262 @@ +"""Unsupervised object states for natural images. + +The synthetic world's objects came from watershed on a flat background, +which real photographs do not offer. Self-supervised vision features do: +DINOv2 patch descriptors segment objects without labels, which is what +the deep-spectral line of work exploits. Each image is cut into segments +by spectral clustering of the patch affinity graph, and each segment is +described by observables that a text corpus also states -- mean colour, +relative area, position, and its own feature centroid for later category +clustering. + +Nothing here reads text, pairs, or region annotations. Boxes from the +preprocessing file are not used; segmentation is derived from pixels. +""" + +from __future__ import annotations + +import argparse +import json +from pathlib import Path + +import numpy as np +import torch +import torch.nn.functional as F +from PIL import Image +from tqdm import tqdm +from transformers import AutoImageProcessor, AutoModel + +from .common import batch_indices, read_json, seed_everything, write_json + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--vg-dir", default="artifacts/vg_5k") + parser.add_argument("--image-cache", default="/tmp/yurenh2-worldalign-vg-images") + parser.add_argument("--model", default="facebook/dinov2-base") + parser.add_argument("--image-size", type=int, default=448) + parser.add_argument("--segments", type=int, default=6) + parser.add_argument("--min-patches", type=int, default=8) + parser.add_argument( + "--oracle-regions", + action="store_true", + help="Diagnostic upper bound: take segments from the annotated " + "region boxes instead of unsupervised segmentation, isolating " + "how much of the field deficit segmentation accounts for.", + ) + parser.add_argument( + "--views", + type=int, + default=1, + help="Augmentation orbit size. Each view is a random resized crop, " + "so content is fixed and framing varies -- the natural-image " + "analogue of the synthetic world's re-renders.", + ) + parser.add_argument("--crop-low", type=float, default=0.75) + parser.add_argument("--nodes", type=int, default=0) + parser.add_argument("--batch-size", type=int, default=16) + parser.add_argument("--device", default="cuda:3") + parser.add_argument("--seed", type=int, default=0) + parser.add_argument("--output", default="artifacts/vg_5k/natural_objects.pt") + return parser.parse_args() + + +def spectral_segments( + features: torch.Tensor, grid: int, segments: int, seed: int +) -> np.ndarray: + """Segment a patch grid by clustering the normalised affinity spectrum. + + Feature cosine affinity is combined with a spatial prior so segments + stay connected, then the leading eigenvectors of the normalised + Laplacian are clustered. This is the standard unsupervised recipe over + self-supervised features. + """ + from sklearn.cluster import KMeans + + normalised = F.normalize(features.double(), dim=-1) + affinity = (normalised @ normalised.T).clamp_min(0).numpy() + coordinates = np.stack( + np.meshgrid(np.arange(grid), np.arange(grid), indexing="ij"), -1 + ).reshape(-1, 2).astype(np.float64) + distance = ((coordinates[:, None, :] - coordinates[None, :, :]) ** 2).sum(-1) + affinity = affinity * np.exp(-distance / (2 * (grid / 4.0) ** 2)) + degree = affinity.sum(1) + laplacian = affinity / np.sqrt(np.outer(degree, degree) + 1e-9) + values, vectors = np.linalg.eigh(laplacian) + embedding = vectors[:, -segments:] + embedding /= np.linalg.norm(embedding, axis=1, keepdims=True).clip(1e-9) + return KMeans(segments, n_init=10, random_state=seed).fit_predict(embedding) + + +def describe_segments( + labels: np.ndarray, + features: torch.Tensor, + pixels: torch.Tensor, + grid: int, + min_patches: int, +) -> list[dict]: + """Observables of each segment: colour, area, position, feature centre.""" + image = pixels.permute(1, 2, 0).numpy() + size = image.shape[0] + scale = size // grid + described = [] + for label in np.unique(labels): + mask = labels == label + if mask.sum() < min_patches: + continue + rows, cols = np.nonzero(mask.reshape(grid, grid)) + pixel_mask = np.zeros((size, size), dtype=bool) + for row, col in zip(rows, cols): + pixel_mask[ + row * scale : (row + 1) * scale, col * scale : (col + 1) * scale + ] = True + described.append( + { + "rgb": image[pixel_mask].mean(0), + "area": float(mask.mean()), + "centre": np.array([rows.mean() / grid, cols.mean() / grid]), + "feature": features[mask].mean(0).numpy(), + } + ) + return described + + +def region_segments( + record: dict, features: torch.Tensor, pixels: torch.Tensor, grid: int +) -> list[dict]: + """Segments taken from annotated boxes: an oracle for segmentation only.""" + image = pixels.permute(1, 2, 0).numpy() + size = image.shape[0] + width, height = record["width"], record["height"] + described = [] + for region in record["regions"]: + x0 = max(0.0, min(region["x"] / width, 1.0)) * grid + x1 = max(0.0, min((region["x"] + region["width"]) / width, 1.0)) * grid + y0 = max(0.0, min(region["y"] / height, 1.0)) * grid + y1 = max(0.0, min((region["y"] + region["height"]) / height, 1.0)) * grid + columns = np.arange(grid) + 0.5 + inside_x = (columns >= x0) & (columns <= x1) + inside_y = (columns >= y0) & (columns <= y1) + mask = (inside_y[:, None] & inside_x[None, :]).flatten() + if not mask.any(): + centre_x = min(grid - 1, max(0, int((x0 + x1) / 2))) + centre_y = min(grid - 1, max(0, int((y0 + y1) / 2))) + mask = np.zeros(grid * grid, dtype=bool) + mask[centre_y * grid + centre_x] = True + rows, cols = np.nonzero(mask.reshape(grid, grid)) + scale = size // grid + pixel_mask = np.zeros((size, size), dtype=bool) + for row, col in zip(rows, cols): + pixel_mask[ + row * scale : (row + 1) * scale, col * scale : (col + 1) * scale + ] = True + described.append( + { + "rgb": image[pixel_mask].mean(0), + "area": float(mask.mean()), + "centre": np.array([rows.mean() / grid, cols.mean() / grid]), + "feature": features[mask].mean(0).numpy(), + } + ) + return described + + +@torch.inference_mode() +def main() -> None: + args = parse_args() + seed_everything(args.seed) + records = [ + json.loads(line) + for line in Path(args.vg_dir, "vision_nodes.private.jsonl") + .read_text(encoding="utf-8") + .splitlines() + if line.strip() + ] + if args.nodes: + records = records[: args.nodes] + processor = AutoImageProcessor.from_pretrained(args.model) + model = AutoModel.from_pretrained(args.model, torch_dtype=torch.float32).to( + args.device + ) + model.eval() + patch = model.config.patch_size + grid = args.image_size // patch + mean = torch.tensor(processor.image_mean).view(3, 1, 1) + std = torch.tensor(processor.image_std).view(3, 1, 1) + + generator = np.random.default_rng(args.seed) + + def load(record: dict, view: int = 0) -> torch.Tensor: + path = Path(args.image_cache, f"{record['source_image_id']}.jpg") + with Image.open(path) as image: + picture = image.convert("RGB") + if view > 0: + width, height = picture.size + scale = generator.uniform(args.crop_low, 1.0) + box_w, box_h = int(width * scale), int(height * scale) + left = int(generator.integers(0, max(width - box_w, 1))) + top = int(generator.integers(0, max(height - box_h, 1))) + picture = picture.crop((left, top, left + box_w, top + box_h)) + resized = picture.resize( + (args.image_size, args.image_size), Image.BILINEAR + ) + return torch.from_numpy(np.asarray(resized).copy()).permute(2, 0, 1).float() / 255.0 + + node_ids, all_segments = [], [] + view_segments: list[list] = [[] for _ in range(args.views)] + for indices in tqdm( + list(batch_indices(len(records), args.batch_size)), desc="segment" + ): + for view in range(args.views): + batch = [records[index] for index in indices] + raw = torch.stack([load(record, view) for record in batch]) + normalised = ((raw - mean) / std).to(args.device) + hidden = model(pixel_values=normalised, return_dict=True).last_hidden_state + patches = hidden[:, 1:].float().cpu() + for position, record in enumerate(batch): + if args.oracle_regions: + described = region_segments( + record, patches[position], raw[position], grid + ) + else: + labels = spectral_segments( + patches[position], grid, args.segments, args.seed + ) + described = describe_segments( + labels, patches[position], raw[position], grid, args.min_patches + ) + if view == 0: + node_ids.append(record["node_id"]) + all_segments.append(described) + view_segments[view].append(described) + + torch.save( + { + "model": args.model, + "node_ids": node_ids, + "segments": all_segments, + "view_segments": view_segments if args.views > 1 else None, + "views": args.views, + "grid": grid, + "protocol": ( + "Segments come from spectral clustering of self-supervised " + "patch features; no boxes, no text, no pairs." + ), + }, + args.output, + ) + counts = [len(item) for item in all_segments] + print( + json.dumps( + { + "nodes": len(node_ids), + "mean_segments": float(np.mean(counts)), + "min_segments": int(np.min(counts)), + } + ) + ) + print(f"Wrote {args.output}") + + +if __name__ == "__main__": + main() diff --git a/worldalign/natural_pipeline.py b/worldalign/natural_pipeline.py new file mode 100644 index 0000000..c30af4d --- /dev/null +++ b/worldalign/natural_pipeline.py @@ -0,0 +1,204 @@ +"""Natural-data field pipeline: continuous states and content projection. + +Produces the field correlation reported for Visual Genome. Two choices +carry most of it, both measured. States stay continuous -- quantising +segments and phrases into class codes costs more than half the signal, +because relation fields need scene similarity and nothing else. And each +side is projected onto the directions that separate scenes rather than +parts, fitted per modality on scenes outside the evaluated set. + +Vision states come from `natural_objects`; language states are +positive-PMI context vectors of the corpus, averaged per phrase. Hidden +pairs are read only to report the correlation. +""" + +from __future__ import annotations + +import argparse +import json +import re +from collections import Counter +from pathlib import Path + +import numpy as np +import torch +import torch.nn.functional as F +from scipy.linalg import eigh +from sklearn.decomposition import TruncatedSVD + +from .common import read_json, seed_everything, write_json +from .natural_families import load_phrases, statistics +from .synth_cc_battery import moment_field + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--vg-dir", default="artifacts/vg_5k") + parser.add_argument("--objects", default="artifacts/vg_5k/natural_objects.pt") + parser.add_argument("--samples", type=int, default=256) + parser.add_argument("--fit-scenes", type=int, default=1200) + parser.add_argument("--word-vectors", type=int, default=48) + parser.add_argument("--min-count", type=int, default=60) + parser.add_argument("--max-vocabulary", type=int, default=600) + parser.add_argument("--context-window", type=int, default=2) + parser.add_argument("--keep", type=int, default=8) + parser.add_argument("--shrinkage", type=float, default=0.05) + parser.add_argument("--views", type=int, default=1) + parser.add_argument("--seed", type=int, default=0) + parser.add_argument("--output", default="artifacts/vg_5k/natural_pipeline.pt") + return parser.parse_args() + + +def word_vectors( + phrases: list[list[str]], args: argparse.Namespace +) -> dict[str, np.ndarray]: + """Positive-PMI context vectors, reduced by truncated SVD.""" + counts = Counter(token for tokens in phrases for token in set(tokens)) + vocabulary = [ + word + for word, count in counts.most_common(args.max_vocabulary) + if count >= args.min_count + ] + _, _, context = statistics(phrases, set(vocabulary), args.context_window) + features = sorted({key for word in vocabulary for key in context[word]}) + index = {key: position for position, key in enumerate(features)} + matrix = np.zeros((len(vocabulary), len(features))) + for row, word in enumerate(vocabulary): + for key, value in context[word].items(): + matrix[row, index[key]] = value + total = matrix.sum() + expected = matrix.sum(1, keepdims=True) * matrix.sum(0, keepdims=True) / total + pmi = np.log(np.maximum(matrix, 1e-9) / np.maximum(expected, 1e-9)) + pmi[matrix == 0] = 0.0 + embedding = TruncatedSVD(args.word_vectors, random_state=args.seed).fit_transform( + np.maximum(pmi, 0.0) + ) + embedding /= np.linalg.norm(embedding, axis=1, keepdims=True).clip(1e-9) + return {word: embedding[row] for row, word in enumerate(vocabulary)} + + +def content_directions( + sets: list[np.ndarray], shrinkage: float +) -> tuple[np.ndarray, np.ndarray]: + """Directions separating scenes more than they separate parts.""" + means = np.stack([item.mean(0) for item in sets]) + centre = means.mean(0) + within = np.concatenate([item - item.mean(0, keepdims=True) for item in sets]) + scatter_within = within.T @ within / max(len(within) - 1, 1) + centred = means - centre + scatter_between = centred.T @ centred / max(len(centred) - 1, 1) + trace = np.trace(scatter_within) / len(scatter_within) + values, vectors = eigh( + scatter_between, scatter_within + shrinkage * trace * np.eye(len(scatter_within)) + ) + return centre, vectors[:, np.argsort(values)[::-1]] + + +def main() -> None: + args = parse_args() + seed_everything(args.seed) + vg_dir = Path(args.vg_dir) + truth = [ + json.loads(line) + for line in (vg_dir / "ground_truth.private.jsonl") + .read_text(encoding="utf-8") + .splitlines() + if line.strip() + ] + records = { + json.loads(line)["node_id"]: json.loads(line) + for line in (vg_dir / "text_nodes.jsonl").read_text(encoding="utf-8").splitlines() + if line.strip() + } + state = torch.load(args.objects, map_location="cpu", weights_only=False) + per_view = state.get("view_segments") or [state["segments"]] + views = min(args.views, len(per_view)) + index = {node: position for position, node in enumerate(state["node_ids"])} + + vectors = word_vectors(load_phrases(vg_dir / "text_nodes.jsonl", "region_closed"), args) + + def language_state(node: str) -> np.ndarray | None: + rows = [] + for phrase in records[node]["region_closed"]: + hits = [vectors[token] for token in re.findall(r"[a-z]+", phrase.lower()) + if token in vectors] + if hits: + rows.append(np.mean(hits, axis=0)) + return np.stack(rows) if rows else None + + def vision_state(node: str, view: int) -> np.ndarray | None: + segments = per_view[view][index[node]] + return ( + np.stack([item["feature"] for item in segments]).astype(np.float64) + if segments + else None + ) + + pairs = [ + pair + for pair in truth + if pair["vision_node_id"] in index and pair["text_node_id"] in records + ] + fit = pairs[args.samples : args.samples + args.fit_scenes] + language_fit = [language_state(p["text_node_id"]) for p in fit] + language_fit = [item for item in language_fit if item is not None and len(item) > 1] + vision_fit = [vision_state(p["vision_node_id"], 0) for p in fit] + vision_fit = [item for item in vision_fit if item is not None and len(item) > 1] + language_centre, language_basis = content_directions(language_fit, args.shrinkage) + vision_centre, vision_basis = content_directions(vision_fit, args.shrinkage) + + evaluated = [ + pair for pair in pairs[: args.samples] + if language_state(pair["text_node_id"]) is not None + ] + keep = args.keep + + def project(raw: np.ndarray, centre: np.ndarray, basis: np.ndarray) -> torch.Tensor: + width = min(keep, basis.shape[1]) + return F.normalize( + torch.tensor((raw - centre) @ basis[:, :width], dtype=torch.float32), dim=-1 + ) + + text_sets = [ + project(language_state(p["text_node_id"]), language_centre, language_basis) + for p in evaluated + ] + view_fields = [] + for view in range(views): + sets = [ + project(vision_state(p["vision_node_id"], view), vision_centre, vision_basis) + for p in evaluated + ] + view_fields.append(moment_field(sets)) + visual_field = torch.stack(view_fields).mean(0) + text_field = moment_field(text_sets) + + mask = ~np.eye(len(evaluated), dtype=bool) + correlation = float( + np.corrcoef( + visual_field.double().numpy()[mask], text_field.double().numpy()[mask] + )[0, 1] + ) + torch.save( + {"visual_field": visual_field, "text_field": text_field, + "nodes": [p["vision_node_id"] for p in evaluated]}, + args.output, + ) + summary = { + "protocol": ( + "Continuous states, content projection fitted per modality on " + "scenes outside the evaluated set. Hidden pairs are read only " + "for the correlation." + ), + "samples": len(evaluated), + "views": views, + "kept_directions": keep, + "field_correlation_at_truth": correlation, + "recovery_threshold": 0.9, + } + print(json.dumps(summary)) + write_json(str(args.output).replace(".pt", ".json"), summary) + + +if __name__ == "__main__": + main() diff --git a/worldalign/precompute_gw.py b/worldalign/precompute_gw.py new file mode 100644 index 0000000..8eac056 --- /dev/null +++ b/worldalign/precompute_gw.py @@ -0,0 +1,51 @@ +from __future__ import annotations + +import argparse +from pathlib import Path + +import torch + +from .common import read_json, seed_everything +from .gw import gw_pseudo_targets +from .io import load_feature_pair, select_rows + + +def parse_args() -> argparse.Namespace: + p = argparse.ArgumentParser() + p.add_argument("--manifest", default="artifacts/manifest.json") + p.add_argument("--vision", default="artifacts/vision.pt") + p.add_argument("--text", default="artifacts/text.pt") + p.add_argument("--output", default="artifacts/gw.pt") + p.add_argument("--clusters", type=int, default=128) + p.add_argument("--seed", type=int, default=20260728) + p.add_argument("--max-iter", type=int, default=100) + return p.parse_args() + + +def main() -> None: + args = parse_args() + seed_everything(args.seed) + manifest = read_json(args.manifest) + vision, text, vlookup, tlookup = load_feature_pair(args.vision, args.text) + x = select_rows( + vision["features"], vlookup, manifest["vision_only_train"] + ) + y = select_rows(text["features"], tlookup, manifest["text_only_train"]) + result = gw_pseudo_targets( + x, + y, + clusters=args.clusters, + seed=args.seed, + max_iter=args.max_iter, + ) + Path(args.output).parent.mkdir(parents=True, exist_ok=True) + torch.save(result, args.output) + print( + f"Wrote {args.output}: K={result['clusters']}, " + f"GW={result['gw_distance']:.6f}, " + f"row_entropy={result['coupling_row_entropy']:.6f}" + ) + + +if __name__ == "__main__": + main() diff --git a/worldalign/prepare.py b/worldalign/prepare.py new file mode 100644 index 0000000..dbdd101 --- /dev/null +++ b/worldalign/prepare.py @@ -0,0 +1,80 @@ +from __future__ import annotations + +import argparse +from collections import Counter + +import numpy as np +from datasets import load_dataset + +from .common import DATASET_NAME, DATASET_SPLIT, write_json + + +def parse_args() -> argparse.Namespace: + p = argparse.ArgumentParser() + p.add_argument("--output", default="artifacts/manifest.json") + p.add_argument("--dataset", default=DATASET_NAME) + p.add_argument("--seed", type=int, default=20260728) + p.add_argument("--unpaired-per-modality", type=int, default=12_000) + p.add_argument("--paired-train", type=int, default=12_000) + p.add_argument("--max-eval", type=int, default=1_000) + return p.parse_args() + + +def main() -> None: + args = parse_args() + dataset = load_dataset(args.dataset, split=DATASET_SPLIT) + by_split: dict[str, list[int]] = {} + for i, split in enumerate(dataset["split"]): + by_split.setdefault(split, []).append(i) + + counts = Counter(dataset["split"]) + print(f"Internal Flickr splits: {dict(counts)}") + train = np.asarray(by_split["train"], dtype=np.int64) + rng = np.random.default_rng(args.seed) + rng.shuffle(train) + + n_unpaired = min(args.unpaired_per_modality, len(train) // 2) + vision_only = train[:n_unpaired] + text_only = train[n_unpaired : 2 * n_unpaired] + assert not set(vision_only.tolist()) & set(text_only.tolist()) + + paired_n = min(args.paired_train, len(train)) + paired_train = train[:paired_n] + + val_key = "val" if "val" in by_split else "validation" + val = np.asarray(by_split[val_key], dtype=np.int64)[: args.max_eval] + test = np.asarray(by_split["test"], dtype=np.int64)[: args.max_eval] + + all_rows = sorted( + set(vision_only.tolist()) + | set(text_only.tolist()) + | set(paired_train.tolist()) + | set(val.tolist()) + | set(test.tolist()) + ) + manifest = { + "dataset": args.dataset, + "dataset_split": DATASET_SPLIT, + "seed": args.seed, + "vision_only_train": vision_only.tolist(), + "text_only_train": text_only.tolist(), + "paired_train": paired_train.tolist(), + "val": val.tolist(), + "test": test.tolist(), + "all_rows": all_rows, + "protocol": ( + "vision_only_train and text_only_train contain disjoint image IDs; " + "val/test pairs are held out from all bridge training" + ), + } + write_json(args.output, manifest) + print( + f"Wrote {args.output}: unpaired={n_unpaired}+{n_unpaired}, " + f"paired_upper_bound={paired_n}, val={len(val)}, test={len(test)}, " + f"features={len(all_rows)} rows" + ) + + +if __name__ == "__main__": + main() + 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() diff --git a/worldalign/spectral_match.py b/worldalign/spectral_match.py new file mode 100644 index 0000000..9e1a0fe --- /dev/null +++ b/worldalign/spectral_match.py @@ -0,0 +1,172 @@ +"""Spectral matching of relation fields: Umeyama and GRAMPA. + +Relation fields are N x N regardless of embedding dimension, so the two +modalities need no common representation dimension. What they do need is +comparable spectra: eigenvector matching degrades when effective ranks +differ or when eigenvalues cluster. GRAMPA is built for that regime -- it +weights every pair of eigenvectors by 1 / ((lambda_i - mu_j)^2 + eta^2) +instead of pairing them one to one -- so both solvers are provided along +with the spectral compatibility diagnostic that predicts whether either +can work. + +No local search: these are polynomial-time solvers that sidestep the +glassy landscape entirely. Hidden pairs score the output only. +""" + +from __future__ import annotations + +import argparse +import json +from pathlib import Path + +import numpy as np +import torch +from scipy.optimize import linear_sum_assignment + +from .common import read_json, seed_everything, write_json + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--fields", help="Saved .pt with visual/text fields.") + parser.add_argument("--data-dir", default="artifacts/synth_v0") + parser.add_argument("--split", choices=["val", "test"], default="test") + parser.add_argument("--samples", type=int, default=256) + parser.add_argument("--merge-distance", type=float, default=30.0) + parser.add_argument("--vision-views", type=int, default=4) + parser.add_argument("--eta", default="0.05,0.1,0.2,0.5") + parser.add_argument("--seed", type=int, default=20260731) + parser.add_argument( + "--output", default="artifacts/synth_v0/spectral_match.json" + ) + parser.add_argument("--fields-output", default="") + return parser.parse_args() + + +def spectral_profile(matrix: np.ndarray) -> dict: + values = np.linalg.eigvalsh(matrix)[::-1] + magnitude = np.abs(values) + total = magnitude.sum() + cumulative = np.cumsum(magnitude) / max(total, 1e-12) + participation = (magnitude.sum() ** 2) / max((magnitude**2).sum(), 1e-12) + gaps = np.abs(np.diff(values)) + return { + "top_eigenvalues": values[:12].tolist(), + "effective_rank_participation": float(participation), + "rank_for_90_percent": int(np.searchsorted(cumulative, 0.90) + 1), + "rank_for_99_percent": int(np.searchsorted(cumulative, 0.99) + 1), + "median_relative_gap": float( + np.median(gaps) / max(magnitude.max(), 1e-12) + ), + "min_relative_gap_top20": float( + gaps[:20].min() / max(magnitude.max(), 1e-12) + ), + } + + +def umeyama(visual: np.ndarray, text: np.ndarray) -> np.ndarray: + """Classic eigenvector-magnitude matching, sign ambiguity absorbed.""" + _, u = np.linalg.eigh(visual) + _, v = np.linalg.eigh(text) + score = np.abs(u) @ np.abs(v).T + rows, cols = linear_sum_assignment(-score) + return cols + + +def grampa(visual: np.ndarray, text: np.ndarray, eta: float) -> np.ndarray: + """Pairwise eigen-alignment similarity, robust to clustered spectra.""" + lam, u = np.linalg.eigh(visual) + mu, v = np.linalg.eigh(text) + ones = np.ones(len(visual)) + left = u.T @ ones # [N] + right = v.T @ ones + weight = np.outer(left, right) / ((lam[:, None] - mu[None, :]) ** 2 + eta**2) + similarity = u @ weight @ v.T + rows, cols = linear_sum_assignment(-similarity) + return cols + + +def score_assignment( + assignment: np.ndarray, hidden: np.ndarray, size: int +) -> dict: + """assignment[i] is the shuffled-space index matched to visual row i. + + Shuffled index j denotes original node hidden[j], and visual row i + denotes original node i, so the match is correct when + hidden[assignment[i]] == i. + """ + return { + "accuracy": float((hidden[assignment] == np.arange(size)).mean()), + "chance": 1.0 / size, + } + + +def main() -> None: + args = parse_args() + seed_everything(args.seed) + if args.fields: + state = torch.load(args.fields, map_location="cpu", weights_only=False) + visual_field = state["visual_field"] + text_field = state["text_field"] + else: + from .synth_triangle_gate import build_fields + + manifest = read_json(Path(args.data_dir, "manifest.json")) + rows = manifest[args.split][: args.samples] + visual_field, text_field = build_fields(args, rows, manifest) + if args.fields_output: + torch.save( + {"visual_field": visual_field, "text_field": text_field}, + args.fields_output, + ) + + size = len(visual_field) + visual = visual_field.double().numpy() + text = text_field.double().numpy() + # Center and scale: spectral methods compare shapes, not offsets. + mask = ~np.eye(size, dtype=bool) + for matrix in (visual, text): + values = matrix[mask] + matrix -= values.mean() + matrix /= values.std() + np.fill_diagonal(matrix, 0.0) + + generator = np.random.default_rng(args.seed) + hidden = generator.permutation(size) + text_shuffled = text[np.ix_(hidden, hidden)] + + report = { + "protocol": ( + "Polynomial-time spectral solvers on N x N relation fields; " + "embedding dimensions are irrelevant by construction. The " + "hidden shuffle is applied to the text field and used only to " + "score the returned assignment." + ), + "samples": size, + "spectra": { + "visual": spectral_profile(visual), + "text": spectral_profile(text_shuffled), + }, + "solvers": {}, + } + # Both fields are built over the same row list, so they are aligned at + # the truth without any permutation. + correlation = float(np.corrcoef(visual[mask], text[mask])[0, 1]) + report["field_correlation_at_truth"] = correlation + + assignment = umeyama(visual, text_shuffled) + report["solvers"]["umeyama"] = score_assignment(assignment, hidden, size) + print(json.dumps({"umeyama": report["solvers"]["umeyama"]})) + + for eta in (float(value) for value in args.eta.split(",")): + assignment = grampa(visual, text_shuffled, eta) + result = score_assignment(assignment, hidden, size) + report["solvers"][f"grampa_eta{eta}"] = result + print(json.dumps({f"grampa_eta{eta}": result})) + + write_json(args.output, report) + print(f"Wrote {args.output}") + + +if __name__ == "__main__": + main() diff --git a/worldalign/synth_cc_battery.py b/worldalign/synth_cc_battery.py new file mode 100644 index 0000000..be28a01 --- /dev/null +++ b/worldalign/synth_cc_battery.py @@ -0,0 +1,377 @@ +"""Upper-bound set battery: connected-component sprites as vision sets. + +The learned towers have not yet produced object states, which leaves two +hypotheses entangled: the set-kernel machinery could be wrong, or only +the towers could be short. This battery separates them. On this world +the background is flat, so connected bright components ARE the objects; +per-component centered sprites are model-free object states of the +minimal-world-knowledge class (like the color anchors on natural data: +declared, fixed, no learning). If set-kernel relation fields built from +these pass the gate and support recovery, the machinery is validated and +the remaining gap is exactly "an SSL objective that discovers objects". + +Text sets are the per-group phrase states of the set battery. Hidden +pairs score orderings only. +""" + +from __future__ import annotations + +import argparse +import json +from pathlib import Path + +import numpy as np +import torch +import torch.nn.functional as F +from scipy import ndimage +from tqdm import tqdm + +from .common import read_json, seed_everything, write_json +from .manifold_gate import standardize_relation +from .ricci_control import run_gates +from .synth_set_battery import ( + parse_group_phrases, + set_similarity_field, + text_group_sets, +) +from .synth_towers import load_image + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--data-dir", default="artifacts/synth_v0") + parser.add_argument("--text-tower", default="artifacts/synth_v0/text_tower.pt") + parser.add_argument("--split", choices=["val", "test"], default="test") + parser.add_argument("--samples", type=int, default=512) + parser.add_argument("--sprite-window", type=int, default=56) + parser.add_argument("--merge-distance", type=float, default=30.0) + parser.add_argument( + "--features", choices=["pixels", "descriptors", "onehot"], default="descriptors" + ) + parser.add_argument("--text-mode", choices=["lm", "bow"], default="bow") + parser.add_argument( + "--kernel", choices=["matching", "moment"], default="matching", + help="moment: symmetric-tensor (Fock) set kernel, no matching step", + ) + parser.add_argument("--vision-views", type=int, default=4) + parser.add_argument("--random-perms", type=int, default=300) + parser.add_argument("--descent-restarts", type=int, default=5) + 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="cuda:3") + parser.add_argument("--seed", type=int, default=20260731) + parser.add_argument( + "--output", default="artifacts/synth_v0/cc_battery_gate.json" + ) + return parser.parse_args() + + +def component_sprites( + image: torch.Tensor, window: int, merge_distance: float +) -> tuple[torch.Tensor, torch.Tensor]: + """Centered sprites of bright connected components, ring-merged. + + Same-group objects sit on a small ring; merging nearby components + reassembles groups so a sprite carries multiplicity as pattern. + """ + array = image.permute(1, 2, 0).numpy() + background = np.median(array.reshape(-1, 3), axis=0) + foreground = (np.abs(array - background).sum(-1) > 0.12) + labels, count = ndimage.label(foreground) + if count == 0: + return torch.zeros(1, 3 * window * window), torch.ones(1) + centers = np.array(ndimage.center_of_mass(foreground, labels, range(1, count + 1))) + sizes = ndimage.sum(foreground, labels, range(1, count + 1)) + # Merge components whose centers are close (ring members). + parent = list(range(count)) + + def find(a: int) -> int: + while parent[a] != a: + parent[a] = parent[parent[a]] + a = parent[a] + return a + + for a in range(count): + for b in range(a + 1, count): + if np.linalg.norm(centers[a] - centers[b]) < merge_distance: + parent[find(a)] = find(b) + groups: dict[int, list[int]] = {} + for a in range(count): + groups.setdefault(find(a), []).append(a) + + height, width = foreground.shape + half = window // 2 + padded = np.pad(array, ((half, half), (half, half), (0, 0))) + padded_mask = np.pad(foreground, half) + sprites, weights = [], [] + for members in groups.values(): + member_mask = np.isin(labels, [m + 1 for m in members]) + mass = float(member_mask.sum()) + ys, xs = np.nonzero(member_mask) + cy, cx = int(ys.mean()), int(xs.mean()) + patch = padded[cy : cy + window, cx : cx + window].copy() + mask_patch = padded_mask[cy : cy + window, cx : cx + window] + patch[~mask_patch] = 0.0 + sprites.append(torch.from_numpy(patch).float().flatten()) + weights.append(mass) + weights = torch.tensor(weights) + return torch.stack(sprites), weights / weights.sum().clamp_min(1e-8) + + +HUE_CENTERS = { + "red": 0.0, "orange": 30.0, "yellow": 60.0, "green": 120.0, + "cyan": 180.0, "blue": 220.0, "purple": 275.0, "pink": 330.0, +} + + +def component_descriptors( + image: torch.Tensor, merge_distance: float +) -> tuple[torch.Tensor, torch.Tensor]: + """Rotation-invariant group descriptors: color classes, log-area, + member multiplicity, and normalized single-member area.""" + import colorsys + + array = image.permute(1, 2, 0).numpy() + background = np.median(array.reshape(-1, 3), axis=0) + foreground = np.abs(array - background).sum(-1) > 0.12 + labels, count = ndimage.label(foreground) + if count == 0: + return torch.zeros(1, 19), torch.ones(1) + centers = np.array( + ndimage.center_of_mass(foreground, labels, range(1, count + 1)) + ) + parent = list(range(count)) + + def find(a: int) -> int: + while parent[a] != a: + parent[a] = parent[parent[a]] + a = parent[a] + return a + + for a in range(count): + for b in range(a + 1, count): + if np.linalg.norm(centers[a] - centers[b]) < merge_distance: + parent[find(a)] = find(b) + groups: dict[int, list[int]] = {} + for a in range(count): + groups.setdefault(find(a), []).append(a) + descriptors, weights = [], [] + total = foreground.size + for members in groups.values(): + member_mask = np.isin(labels, [m + 1 for m in members]) + pixels = array[member_mask] + mean_rgb = pixels.mean(0) + h, s_, v = colorsys.rgb_to_hsv(*mean_rgb.tolist()) + hue = h * 360.0 + vector = np.zeros(16) + if v > 0.75 and s_ < 0.25: + vector[8] = 1.0 # white + elif v < 0.25: + vector[9] = 1.0 # black + elif s_ < 0.25: + vector[10] = 1.0 # gray + elif abs(hue - 30.0) < 25.0 and v < 0.55: + vector[11] = 1.0 # brown + else: + for index, center in enumerate(HUE_CENTERS.values()): + distance = min(abs(hue - center), 360.0 - abs(hue - center)) + if distance < 25.0: + vector[index] = 1.0 + break + member_count = len(members) + area = float(member_mask.sum()) / total + vector[12] = np.log(area + 1e-6) / 6.0 + vector[13] = (member_count - 1) / 3.0 + vector[14] = np.log(area / member_count + 1e-6) / 6.0 + vector[15] = 1.0 + # Rotation-invariant shape features of the largest single member: + # compactness, convexity, and inertia eccentricity separate the + # six shapes without orientation. + largest = max(members, key=lambda m: (labels == m + 1).sum()) + single = labels == largest + 1 + area_px = float(single.sum()) + eroded = ndimage.binary_erosion(single) + perimeter = float((single & ~eroded).sum()) + compactness = 4.0 * np.pi * area_px / max(perimeter, 1.0) ** 2 + ys, xs = np.nonzero(single) + ys = ys - ys.mean(); xs = xs - xs.mean() + cov = np.cov(np.stack([xs, ys])) + 1e-6 * np.eye(2) + eigenvalues = np.linalg.eigvalsh(cov) + eccentricity = float(1.0 - eigenvalues[0] / eigenvalues[1]) + hull_span = (xs.max() - xs.min() + 1) * (ys.max() - ys.min() + 1) + boxfill = area_px / max(hull_span, 1.0) + shape_vector = np.array([compactness, eccentricity, boxfill]) + vector = np.concatenate([vector, shape_vector]) + descriptors.append(torch.tensor(vector, dtype=torch.float32)) + weights.append(float(member_mask.sum())) + weights = torch.tensor(weights) + return torch.stack(descriptors), weights / weights.sum().clamp_min(1e-8) + + +def onehot_descriptors( + raw_sets: list[torch.Tensor], +) -> list[torch.Tensor]: + """Factor-mirrored one-hot recoding of descriptor sets. + + Sizes are binned by corpus terciles of single-member log-area with the + middle bin unmarked, mirroring the text side where medium size has no + word; shapes are k-means clusters of the rotation-invariant shape + features. Both statistics come from the evaluated corpus itself, + unimodally. Output channels mirror the bag-of-words support: color + (12), small/large (2), count (4), shape cluster (6). + """ + from sklearn.cluster import KMeans + + all_groups = torch.cat(raw_sets) + # The renderer shrinks radii with member count, so raw single-member + # area confounds size class with multiplicity; residualize log-area on + # member count before binning (unimodal statistics). + counts_all = (all_groups[:, 13] * 3.0).round().clamp(0, 3) + area_all = all_groups[:, 14] + count_means = {} + for value in (0.0, 1.0, 2.0, 3.0): + chosen = counts_all == value + count_means[value] = float(area_all[chosen].mean()) if chosen.any() else 0.0 + adjusted_all = area_all - torch.tensor( + [count_means[float(v)] for v in counts_all] + ) + low, high = adjusted_all.quantile(1.0 / 3.0), adjusted_all.quantile(2.0 / 3.0) + shape_features = all_groups[:, 16:19].numpy() + clusters = KMeans(n_clusters=6, n_init=10, random_state=0).fit(shape_features) + recoded = [] + for groups in raw_sets: + vectors = torch.zeros(len(groups), 24) + vectors[:, :12] = groups[:, :12] + counts_here = (groups[:, 13] * 3.0).round().clamp(0, 3) + adjusted = groups[:, 14] - torch.tensor( + [count_means[float(v)] for v in counts_here] + ) + vectors[:, 12] = (adjusted <= low).float() # small + vectors[:, 13] = (adjusted >= high).float() # large + counts = (groups[:, 13] * 3.0).round().long().clamp(0, 3) + vectors[torch.arange(len(groups)), 14 + counts] = 1.0 + labels = clusters.predict(groups[:, 16:19].numpy()) + vectors[torch.arange(len(groups)), 18 + labels] = 1.0 + recoded.append(vectors) + return recoded + + +def moment_field(sets: list[torch.Tensor]) -> torch.Tensor: + """Symmetric-tensor (second-quantized) set kernel, matching-free. + + phi(S) concatenates the degree-1 and degree-2 moments of the set; the + field is the Gram matrix of normalized phi. No assignment step, so no + matching-value or set-size bias can enter. + """ + phis = [] + for members in sets: + m1 = members.mean(0) + m2 = (members[:, :, None] * members[:, None, :]).mean(0).flatten() + phi = torch.cat([m1, m2]) + phis.append(phi / phi.norm().clamp_min(1e-9)) + stacked = torch.stack(phis) + return stacked @ stacked.T + + +def phrase_bow_sets( + rows: list[int], captions: list[list[str]], vocabulary: list[str] +) -> list[torch.Tensor]: + index = {word: i for i, word in enumerate(vocabulary)} + sets = [] + for row in rows: + phrases = parse_group_phrases(captions[row][0]) + vectors = torch.zeros(len(phrases), len(index)) + for p, phrase in enumerate(phrases): + for token in phrase.split(): + if token in index: + vectors[p, index[token]] += 1.0 + sets.append(F.normalize(vectors, dim=-1)) + return sets + + +def main() -> None: + args = parse_args() + seed_everything(args.seed) + manifest = read_json(Path(args.data_dir, "manifest.json")) + captions = read_json(Path(args.data_dir, "captions.json"))["captions"] + rows = manifest[args.split][: args.samples] + image_dir = Path(manifest["image_dir"]) + + per_view_fields = [] + set_sizes = [] + for view in range(args.vision_views): + vision_sets = [] + raw_sets = [] + for row in tqdm(rows, desc=f"cc sprites v{view}"): + image = load_image(image_dir / f"scene{row:06d}_v{view}.png") + if args.features in ("descriptors", "onehot"): + sprites, _ = component_descriptors(image, args.merge_distance) + else: + sprites, _ = component_sprites( + image, args.sprite_window, args.merge_distance + ) + raw_sets.append(sprites) + if view == 0: + set_sizes.append(len(sprites)) + if args.features == "onehot": + vision_sets = [ + F.normalize(v, dim=-1) for v in onehot_descriptors(raw_sets) + ] + else: + vision_sets = [F.normalize(v, dim=-1) for v in raw_sets] + if args.kernel == "moment": + per_view_fields.append(moment_field(vision_sets)) + else: + per_view_fields.append(set_similarity_field(vision_sets)) + + if args.text_mode == "bow": + text_sets = phrase_bow_sets(rows, captions, manifest["vocabulary"]) + else: + text_sets = text_group_sets(rows, captions, args) + if args.kernel == "moment": + text_field = moment_field(text_sets) + else: + text_field = set_similarity_field(text_sets) + # Mass weighting inside the matching corrupts similarity grading + # (0.16 vs 0.45 against soft truth); match unweighted. Averaging the + # per-view fields cancels segmentation errors across resampled layouts. + visual_field = torch.stack(per_view_fields).mean(0) + visual_channels = standardize_relation(visual_field.double())[0][None] + text_channels = standardize_relation(text_field.double())[0][None] + generator = torch.Generator().manual_seed(args.seed) + report = { + "protocol": ( + "Vision sets are model-free connected-component sprites " + "(declared minimal world knowledge); text sets are per-group " + "phrase states. Hidden pairs score orderings only." + ), + "split": args.split, + "samples": len(rows), + "mean_vision_set_size": float(np.mean(set_sizes)), + **run_gates(text_channels, visual_channels, args, generator), + } + verdict = { + "true_z_mse": report["gate_a"]["random"]["mse"]["true_z"], + "improving_fraction": report["gate_b"]["improving_fraction"], + "descent_keeps": report["descent_from_true"]["final_accuracy"], + "true_mse": report["gate_a"]["true"]["mse"], + "best_random_descent": min( + (r["final_objective"] for r in report["descent_from_random"]), + default=None, + ), + } + verdict["counterfeit_found"] = bool( + verdict["best_random_descent"] is not None + and verdict["best_random_descent"] < verdict["true_mse"] + ) + report["verdict"] = verdict + print(json.dumps({"verdict": verdict})) + write_json(args.output, report) + print(f"Wrote {args.output}") + + +if __name__ == "__main__": + main() diff --git a/worldalign/synth_decode.py b/worldalign/synth_decode.py new file mode 100644 index 0000000..183e4a9 --- /dev/null +++ b/worldalign/synth_decode.py @@ -0,0 +1,150 @@ +"""Close the loop: recovered pairs to a map to generated descriptions. + +Matching is transductive -- it aligns one fixed population. A usable +system needs a map, so the recovered assignment is treated as pseudo-pair +supervision for a small vision-to-text-state regressor, and the map is +then applied to scenes that took no part in the matching. Descriptions +are read out through the frozen text tower. + +Two controls decide whether the output is image-conditioned: the same +pipeline with the recovered assignment replaced by a random one, and the +same map applied to a shuffled image. Hidden pairs score; nothing here +trains on them. +""" + +from __future__ import annotations + +import argparse +import json +from pathlib import Path + +import numpy as np +import torch +import torch.nn.functional as F + +from .common import read_json, seed_everything, write_json + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--data-dir", default="artifacts/synth_v1") + parser.add_argument("--states", required=True, + help=".pt with vision_states, text_states, rows for " + "the matched population and the held-out split.") + parser.add_argument("--assignment", required=True, + help=".pt with the recovered permutation and truth.") + parser.add_argument("--ridge", type=float, default=1.0) + parser.add_argument("--seed", type=int, default=0) + parser.add_argument("--output", default="artifacts/synth_v1/decode.json") + return parser.parse_args() + + +def fit_map(source: np.ndarray, target: np.ndarray, ridge: float) -> np.ndarray: + source = np.concatenate([source, np.ones((len(source), 1))], axis=1) + gram = source.T @ source + ridge * np.eye(source.shape[1]) + return np.linalg.solve(gram, source.T @ target) + + +def apply_map(weights: np.ndarray, source: np.ndarray) -> np.ndarray: + source = np.concatenate([source, np.ones((len(source), 1))], axis=1) + return source @ weights + + +def nearest_caption( + predicted: np.ndarray, bank: np.ndarray +) -> np.ndarray: + predicted = predicted / np.linalg.norm(predicted, axis=1, keepdims=True).clip(1e-9) + bank = bank / np.linalg.norm(bank, axis=1, keepdims=True).clip(1e-9) + return (predicted @ bank.T).argmax(1) + + +def factor_scores( + predicted_rows: list[int], truth_rows: list[int], scenes: list[dict] +) -> dict: + """Do the retrieved descriptions state the right world facts?""" + colour_f1, count_ok, group_ok = [], [], [] + for predicted, truth in zip(predicted_rows, truth_rows): + p, t = scenes[predicted], scenes[truth] + pc = {g["color"] for g in p["groups"]} + tc = {g["color"] for g in t["groups"]} + overlap = len(pc & tc) + precision = overlap / max(len(pc), 1) + recall = overlap / max(len(tc), 1) + colour_f1.append( + 0.0 if precision + recall == 0 else 2 * precision * recall / (precision + recall) + ) + count_ok.append( + sorted(g["count"] for g in p["groups"]) + == sorted(g["count"] for g in t["groups"]) + ) + group_ok.append(len(p["groups"]) == len(t["groups"])) + return { + "colour_set_f1": float(np.mean(colour_f1)), + "count_multiset_exact": float(np.mean(count_ok)), + "group_count_exact": float(np.mean(group_ok)), + } + + +def main() -> None: + args = parse_args() + seed_everything(args.seed) + scenes = read_json(Path(args.data_dir, "scenes.private.json"))["scenes"] + states = torch.load(args.states, map_location="cpu", weights_only=False) + recovered = torch.load(args.assignment, map_location="cpu", weights_only=False) + + match_vision = states["match_vision"].double().numpy() + match_text = states["match_text"].double().numpy() + held_vision = states["held_vision"].double().numpy() + held_text = states["held_text"].double().numpy() + held_rows = states["held_rows"] + + assignment = np.asarray(recovered["assignment"]) + generator = np.random.default_rng(args.seed) + random_assignment = generator.permutation(len(assignment)) + + report = { + "protocol": ( + "The recovered assignment supplies pseudo-pairs for a ridge " + "map from vision states to text states; the map is applied to " + "held-out scenes that took no part in matching. Random " + "assignment and shuffled-image controls bound the claim." + ), + "matched_population": len(assignment), + "held_out": len(held_rows), + "recovery_accuracy": float( + (assignment == np.arange(len(assignment))).mean() + ), + "conditions": {}, + } + + for label, permutation in ( + ("recovered_pairs", assignment), + ("random_pairs", random_assignment), + ): + weights = fit_map(match_vision, match_text[permutation], args.ridge) + predicted = apply_map(weights, held_vision) + retrieved = nearest_caption(predicted, held_text) + exact = float((retrieved == np.arange(len(held_rows))).mean()) + entry = { + "held_out_retrieval_exact": exact, + "chance": 1.0 / len(held_rows), + **factor_scores( + [held_rows[i] for i in retrieved], held_rows, scenes + ), + } + if label == "recovered_pairs": + shuffled = generator.permutation(len(held_vision)) + predicted_shuffled = apply_map(weights, held_vision[shuffled]) + retrieved_shuffled = nearest_caption(predicted_shuffled, held_text) + entry["shuffled_image_control"] = factor_scores( + [held_rows[i] for i in retrieved_shuffled], held_rows, scenes + ) + report["conditions"][label] = entry + print(json.dumps({label: entry})) + + write_json(args.output, report) + print(f"Wrote {args.output}") + + +if __name__ == "__main__": + main() diff --git a/worldalign/synth_deep_gate.py b/worldalign/synth_deep_gate.py new file mode 100644 index 0000000..74bfea9 --- /dev/null +++ b/worldalign/synth_deep_gate.py @@ -0,0 +1,203 @@ +"""Corrected gate: the deepest reachable minimum, not descent retention. + +Today's lesson: descent retention measures the search operator, not the +energy. Sampled-proposal descent keeps the truth for every candidate +energy, while long tempering on the same energy reaches states well below +it. The only decision-relevant question is therefore + + E(truth) <= E(deepest state a strong searcher reaches) ? + +This module answers it uniformly for the candidate energies (pairwise +moment kernel, triangle-only, and their sum) with one strong searcher: +long parallel tempering with large proposal batches from random starts, +plus a truth-initialized tempering arm that reports whether the truth +itself survives thermal agitation. Hidden pairs score only. +""" + +from __future__ import annotations + +import argparse +import json +from pathlib import Path + +import torch + +from .common import read_json, seed_everything, write_json +from .synth_triangle_gate import TriangleEnergy, build_fields, standardized + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--data-dir", default="artifacts/synth_v0") + parser.add_argument("--split", choices=["val", "test"], default="test") + parser.add_argument("--samples", type=int, default=256) + parser.add_argument("--merge-distance", type=float, default=30.0) + parser.add_argument("--vision-views", type=int, default=4) + parser.add_argument("--triples", type=int, default=200000) + parser.add_argument( + "--energies", + default="pair,triangle,both", + help="Comma list from pair, triangle, both.", + ) + parser.add_argument("--replicas", type=int, default=6) + parser.add_argument("--rounds", type=int, default=1500) + parser.add_argument("--proposals", type=int, default=64) + parser.add_argument("--temp-high", type=float, default=3e-2) + parser.add_argument("--temp-low", type=float, default=1e-4) + parser.add_argument("--exchange-every", type=int, default=20) + parser.add_argument("--device", default="cuda:3") + parser.add_argument("--seed", type=int, default=20260731) + parser.add_argument("--output", default="artifacts/synth_v0/deep_gate.json") + return parser.parse_args() + + +def temper( + energy: TriangleEnergy, + starts: list[torch.Tensor], + truth: torch.Tensor, + args: argparse.Namespace, + generator: torch.Generator, + label: str, +) -> dict: + size = energy.size + temperatures = torch.logspace( + torch.log10(torch.tensor(args.temp_low)), + torch.log10(torch.tensor(args.temp_high)), + len(starts), + ) + states = [start.clone() for start in starts] + energies = [energy.total(state) for state in states] + best = {"energy": min(energies), "accuracy": 0.0} + for round_index in range(args.rounds): + for replica in range(len(states)): + temperature = float(temperatures[replica]) + for _ in range(args.proposals): + p = int(torch.randint(0, size, (1,), generator=generator)) + q = int(torch.randint(0, size, (1,), generator=generator)) + if p == q: + continue + delta = energy.swap_delta(states[replica], p, q) + threshold = -temperature * float( + torch.rand(1, generator=generator).clamp_min(1e-12).log() + ) + if delta < threshold: + states[replica][[p, q]] = states[replica][[q, p]] + energies[replica] += delta + if round_index % args.exchange_every == 0: + for replica in range(len(states) - 1): + gap = (energies[replica] - energies[replica + 1]) * ( + 1.0 / float(temperatures[replica]) + - 1.0 / float(temperatures[replica + 1]) + ) + accept = gap > 0 or float( + torch.rand(1, generator=generator) + ) < min(1.0, float(torch.tensor(gap).exp())) + if accept: + states[replica], states[replica + 1] = ( + states[replica + 1], + states[replica], + ) + energies[replica], energies[replica + 1] = ( + energies[replica + 1], + energies[replica], + ) + cold = min(range(len(states)), key=lambda r: energies[r]) + if energies[cold] < best["energy"]: + best = { + "energy": energies[cold], + "accuracy": float( + (states[cold].cpu() == truth.cpu()).float().mean() + ), + "round": round_index, + } + exact = [energy.total(state) for state in states] + cold = min(range(len(states)), key=lambda r: exact[r]) + return { + "arm": label, + "best_seen": best, + "final_cold_energy": exact[cold], + "final_cold_accuracy": float( + (states[cold].cpu() == truth.cpu()).float().mean() + ), + "final_accuracies": [ + float((state.cpu() == truth.cpu()).float().mean()) for state in states + ], + } + + +def main() -> None: + args = parse_args() + seed_everything(args.seed) + manifest = read_json(Path(args.data_dir, "manifest.json")) + rows = manifest[args.split][: args.samples] + visual_field, text_field = build_fields(args, rows, manifest) + + device = torch.device(args.device) + size = len(rows) + generator = torch.Generator().manual_seed(args.seed) + hidden = torch.randperm(size, generator=generator) + truth = torch.argsort(hidden).to(device) + text = standardized(text_field[hidden][:, hidden].double().to(device)).float() + visual = standardized(visual_field.double().to(device)).float() + + triples = torch.randint(0, size, (args.triples, 3), generator=generator) + triples = triples[ + (triples[:, 0] != triples[:, 1]) + & (triples[:, 1] != triples[:, 2]) + & (triples[:, 0] != triples[:, 2]) + ].to(device) + + weights = { + "pair": (1.0, 0.0), + "triangle": (0.0, 1.0), + "both": (1.0, 1.0), + } + report = { + "protocol": ( + "The decision statistic is the deepest energy a strong " + "searcher reaches versus the energy of the truth. Descent " + "retention is reported but not used: it measures the search " + "operator. Hidden pairs score only." + ), + "samples": size, + "triples": len(triples), + "energies": {}, + } + for name in (item.strip() for item in args.energies.split(",")): + pair_weight, triangle_weight = weights[name] + energy = TriangleEnergy(text, visual, triples, pair_weight, triangle_weight) + true_energy = energy.total(truth) + random_starts = [ + torch.argsort(torch.rand(size, generator=generator)).to(device) + for _ in range(args.replicas) + ] + from_random = temper(energy, random_starts, truth, args, generator, "random") + from_truth = temper( + energy, + [truth.clone() for _ in range(args.replicas)], + truth, + args, + generator, + "truth", + ) + deepest = min(from_random["best_seen"]["energy"], from_truth["best_seen"]["energy"]) + entry = { + "true_energy": true_energy, + "from_random": from_random, + "from_truth": from_truth, + "deepest_seen": deepest, + "margin_over_true": deepest / true_energy - 1.0, + "passes": bool(deepest >= true_energy - 1e-9), + "recovery_accuracy": max( + from_random["best_seen"]["accuracy"], + from_random["final_cold_accuracy"], + ), + } + report["energies"][name] = entry + print(json.dumps({name: {k: entry[k] for k in ("true_energy", "deepest_seen", "margin_over_true", "passes", "recovery_accuracy")}})) + write_json(args.output, report) + print(f"Wrote {args.output}") + + +if __name__ == "__main__": + main() diff --git a/worldalign/synth_extract.py b/worldalign/synth_extract.py new file mode 100644 index 0000000..3186c3e --- /dev/null +++ b/worldalign/synth_extract.py @@ -0,0 +1,144 @@ +"""Feature extraction for the synthetic world, in main-pipeline schema. + +Emits vision.pt, text.pt, and text_orbits.pt files with the same fields +the Flickr loaders read, so the diagnostic, gate, projection, and recovery +stack runs on the synthetic world unchanged. +""" + +from __future__ import annotations + +import argparse +from pathlib import Path + +import torch +import torch.nn.functional as F +from tqdm import tqdm + +from .common import batch_indices, read_json +from .synth_towers import TextTower, VisionTower, load_image, tokenize + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--data-dir", default="artifacts/synth_v0") + parser.add_argument("--vision-tower", default="artifacts/synth_v0/vision_tower.pt") + parser.add_argument("--text-tower", default="artifacts/synth_v0/text_tower.pt") + parser.add_argument("--batch-size", type=int, default=512) + parser.add_argument("--device", default="cuda:3") + parser.add_argument("--vision-output", default="artifacts/synth_v0/vision.pt") + parser.add_argument("--text-output", default="artifacts/synth_v0/text.pt") + parser.add_argument( + "--orbits-output", default="artifacts/synth_v0/text_orbits.pt" + ) + return parser.parse_args() + + +@torch.inference_mode() +def main() -> None: + args = parse_args() + manifest = read_json(Path(args.data_dir, "manifest.json")) + captions = read_json(Path(args.data_dir, "captions.json"))["captions"] + image_dir = Path(manifest["image_dir"]) + views = manifest["visual_views"] + + vision_state = torch.load(args.vision_tower, map_location="cpu", weights_only=False) + vision_args = vision_state["args"] + vision = VisionTower( + manifest["image_size"], + vision_args["patch"], + vision_args["dim"], + vision_args["depth"], + vision_args["heads"], + ).to(args.device) + vision.load_state_dict(vision_state["model"]) + vision.eval() + + rendered_rows = sorted( + set(manifest["vision_only_train"]) | set(manifest["val"]) | set(manifest["test"]) + ) + jobs = [(row, view) for row in rendered_rows for view in range(views)] + features = [] + objective = vision_args.get("objective", "infonce") + for indices in tqdm(list(batch_indices(len(jobs), args.batch_size)), desc="vision"): + pixels = torch.stack( + [ + load_image(image_dir / f"scene{jobs[i][0]:06d}_v{jobs[i][1]}.png") + for i in indices + ] + ).to(args.device) + tokens = vision.encode(pixels) + state = tokens[:, 1:].mean(1) if objective in ("simmim", "data2vec") else tokens[:, 0] + features.append(state.float().cpu()) + view_features = torch.cat(features).reshape(len(rendered_rows), views, -1) + torch.save( + { + "model": "synth_vision_tower", + "rows": rendered_rows, + "features": F.normalize(view_features.mean(1), dim=-1), + "view_features": view_features, + "views_per_scene": views, + }, + args.vision_output, + ) + + text_state = torch.load(args.text_tower, map_location="cpu", weights_only=False) + text_args = text_state["args"] + vocab = text_state["vocab"] + text = TextTower( + len(vocab), + text_args["text_dim"], + text_args["depth"], + text_args.get("text_heads", 4), + text_args["context"], + ).to(args.device) + text.load_state_dict(text_state["model"]) + text.eval() + + rows = list(range(manifest["all_rows"])) + orbit_size = len(captions[0]) + jobs = [(row, k) for row in rows for k in range(orbit_size)] + pooled = [] + for indices in tqdm(list(batch_indices(len(jobs), args.batch_size)), desc="text"): + batch = [tokenize(captions[jobs[i][0]][jobs[i][1]], vocab) for i in indices] + longest = min(text_args["context"], max(len(s) for s in batch)) + tokens = torch.zeros(len(batch), longest, dtype=torch.long) + for index, sentence in enumerate(batch): + clipped = sentence[:longest] + tokens[index, : len(clipped)] = torch.tensor(clipped) + tokens = tokens.to(args.device) + hidden = text(tokens) + mask = (tokens != 0).float()[..., None] + pooled.append( + ((hidden * mask).sum(1) / mask.sum(1).clamp_min(1.0)).float().cpu() + ) + orbit_features = torch.cat(pooled).reshape(len(rows), orbit_size, -1) + torch.save( + { + "model": "synth_text_tower", + "layer": -1, + "rows": rows, + "features": orbit_features[:, 0], + "captions": [captions[row][0] for row in rows], + "all_captions": captions, + }, + args.text_output, + ) + torch.save( + { + "model": "synth_text_tower", + "layer": -1, + "rows": rows, + "features": orbit_features, + "captions": captions, + "views_per_orbit": orbit_size, + "row_groups": ["all"], + }, + args.orbits_output, + ) + print( + f"Wrote {args.vision_output}, {args.text_output}, {args.orbits_output}" + ) + + +if __name__ == "__main__": + main() diff --git a/worldalign/synth_fast_gate.py b/worldalign/synth_fast_gate.py new file mode 100644 index 0000000..54b6891 --- /dev/null +++ b/worldalign/synth_fast_gate.py @@ -0,0 +1,345 @@ +"""Closed-form batched gate: exact all-triple triangle energy. + +The sampled triangle energy cost 5.7 ms per proposal because each +evaluation gathered a variable-length triple set in Python. It has a +closed form. Writing M(sigma) for the elementwise product of the +sigma-permuted text field with the visual field, + + pairwise = const - 2 * sum(M) + all-triple = const - 2 * trace(M^3) / 6 + +because trace of a power is invariant under simultaneous row-column +permutation, so the text-only and vision-only terms do not move. One +batched matrix product evaluates trace(M^3) = sum(M * (M @ M)) for +hundreds of candidate permutations at once, over every C(N,3) triple +rather than a sample. + +This buys a genuinely stronger searcher: exact steepest descent over all +N(N-1)/2 transpositions per step, plus batched-proposal tempering. +Hidden pairs score orderings only. +""" + +from __future__ import annotations + +import argparse +import json +from pathlib import Path + +import torch + +from .common import read_json, seed_everything, write_json +from .synth_triangle_gate import build_fields, standardized + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--fields", default="", help="Saved .pt with visual_field/text_field.") + parser.add_argument("--data-dir", default="artifacts/synth_v0") + parser.add_argument("--split", choices=["val", "test"], default="test") + parser.add_argument("--samples", type=int, default=512) + parser.add_argument("--merge-distance", type=float, default=30.0) + parser.add_argument("--vision-views", type=int, default=4) + parser.add_argument("--energies", default="pair,triangle,both") + parser.add_argument("--chunk", type=int, default=512) + parser.add_argument("--descent-restarts", type=int, default=5) + parser.add_argument("--descent-max-steps", type=int, default=4000) + parser.add_argument("--replicas", type=int, default=8) + parser.add_argument("--rounds", type=int, default=3000) + parser.add_argument("--proposals", type=int, default=64) + parser.add_argument("--temp-high", type=float, default=3e-2) + parser.add_argument("--temp-low", type=float, default=1e-4) + parser.add_argument("--exchange-every", type=int, default=20) + parser.add_argument("--polish-steps", type=int, default=2000) + parser.add_argument("--device", default="cuda:3") + parser.add_argument("--seed", type=int, default=20260731) + parser.add_argument("--output", default="artifacts/synth_v0/fast_gate.json") + return parser.parse_args() + + +class ClosedFormEnergy: + """Energy as a function of M = permuted-text * visual, batched.""" + + def __init__( + self, + text: torch.Tensor, + visual: torch.Tensor, + pair_weight: float, + triangle_weight: float, + chunk: int, + ) -> None: + self.text = text + self.visual = visual + self.pair_weight = pair_weight + self.triangle_weight = triangle_weight + self.chunk = chunk + self.size = len(visual) + self.pair_count = self.size * (self.size - 1) + self.triple_count = self.size * (self.size - 1) * (self.size - 2) + # Permutation-invariant constants, kept so reported values match + # the direct definitions of the two energies. + square_text = text * text + square_visual = visual * visual + self.pair_constant = float(square_text.sum() + square_visual.sum()) + self.triangle_constant = float( + self._trace_cube(square_text[None])[0] + + self._trace_cube(square_visual[None])[0] + ) + + @staticmethod + def _trace_cube(matrices: torch.Tensor) -> torch.Tensor: + return (matrices * torch.bmm(matrices, matrices)).sum((-2, -1)) + + def energy(self, permutations: torch.Tensor) -> torch.Tensor: + """Exact energy for a batch of permutations [B, N].""" + values = [] + for start in range(0, len(permutations), self.chunk): + block = permutations[start : start + self.chunk] + permuted = self.text[block[:, :, None], block[:, None, :]] + product = permuted * self.visual + total = torch.zeros(len(block), device=permuted.device) + if self.pair_weight: + pair = (self.pair_constant - 2.0 * product.sum((-2, -1))) / ( + self.pair_count + ) + total = total + self.pair_weight * pair + if self.triangle_weight: + triangle = ( + self.triangle_constant - 2.0 * self._trace_cube(product) + ) / self.triple_count + total = total + self.triangle_weight * triangle + values.append(total) + return torch.cat(values) + + +def all_swaps(size: int, device: torch.device) -> torch.Tensor: + rows, cols = torch.triu_indices(size, size, offset=1, device=device) + return torch.stack([rows, cols], dim=1) + + +def apply_swaps(permutation: torch.Tensor, swaps: torch.Tensor) -> torch.Tensor: + batch = permutation[None].repeat(len(swaps), 1) + index = torch.arange(len(swaps), device=permutation.device) + p, q = swaps[:, 0], swaps[:, 1] + values_p = batch[index, p].clone() + batch[index, p] = batch[index, q] + batch[index, q] = values_p + return batch + + +def steepest_descent( + energy: ClosedFormEnergy, + start: torch.Tensor, + swaps: torch.Tensor, + max_steps: int, +) -> tuple[torch.Tensor, float]: + """Exact steepest descent over every transposition each step.""" + current = start.clone() + value = float(energy.energy(current[None])[0]) + for _ in range(max_steps): + candidates = apply_swaps(current, swaps) + values = energy.energy(candidates) + best = int(values.argmin()) + if float(values[best]) >= value - 1e-12: + break + current = candidates[best] + value = float(values[best]) + return current, value + + +def temper( + energy: ClosedFormEnergy, + starts: torch.Tensor, + args: argparse.Namespace, + generator: torch.Generator, + device: torch.device, +) -> tuple[torch.Tensor, torch.Tensor]: + """Batched-proposal parallel tempering; returns states and energies.""" + size = energy.size + replicas = len(starts) + temperatures = torch.logspace( + torch.log10(torch.tensor(args.temp_low)), + torch.log10(torch.tensor(args.temp_high)), + replicas, + ).to(device) + states = starts.clone() + values = energy.energy(states) + for round_index in range(args.rounds): + p = torch.randint( + 0, size, (replicas, args.proposals), generator=generator + ).to(device) + q = torch.randint( + 0, size, (replicas, args.proposals), generator=generator + ).to(device) + valid = p != q + batch = states[:, None, :].repeat(1, args.proposals, 1) + index_r = torch.arange(replicas, device=device)[:, None] + original_p = batch.gather(2, p[..., None]).squeeze(-1) + original_q = batch.gather(2, q[..., None]).squeeze(-1) + batch.scatter_(2, p[..., None], original_q[..., None]) + batch.scatter_(2, q[..., None], original_p[..., None]) + flat = batch.reshape(replicas * args.proposals, size) + proposal_values = energy.energy(flat).reshape(replicas, args.proposals) + deltas = proposal_values - values[:, None] + noise = torch.rand( + replicas, args.proposals, generator=generator + ).to(device) + threshold = -temperatures[:, None] * noise.clamp_min(1e-12).log() + accept = (deltas < threshold) & valid + first = torch.where( + accept.any(-1), + accept.float().argmax(-1), + torch.zeros(replicas, dtype=torch.long, device=device), + ) + taken = accept.any(-1) + chosen = batch[index_r.squeeze(-1), first] + states = torch.where(taken[:, None], chosen, states) + values = torch.where( + taken, proposal_values[index_r.squeeze(-1), first], values + ) + if round_index % args.exchange_every == 0: + for replica in range(replicas - 1): + gap = (values[replica] - values[replica + 1]) * ( + 1.0 / temperatures[replica] - 1.0 / temperatures[replica + 1] + ) + accept_swap = gap > 0 or float( + torch.rand(1, generator=generator) + ) < float(gap.exp().clamp(max=1.0)) + if accept_swap: + states[[replica, replica + 1]] = states[ + [replica + 1, replica] + ] + values[[replica, replica + 1]] = values[ + [replica + 1, replica] + ] + return states, values + + +def main() -> None: + args = parse_args() + seed_everything(args.seed) + if args.fields: + state = torch.load(args.fields, map_location="cpu", weights_only=False) + visual_field, text_field = state["visual_field"], state["text_field"] + else: + manifest = read_json(Path(args.data_dir, "manifest.json")) + rows = manifest[args.split][: args.samples] + visual_field, text_field = build_fields(args, rows, manifest) + + device = torch.device(args.device) + size = len(visual_field) + generator = torch.Generator().manual_seed(args.seed) + hidden = torch.randperm(size, generator=generator) + truth = torch.argsort(hidden).to(device) # scored, never optimized against + text = standardized(text_field[hidden][:, hidden].double().to(device)).float() + visual = standardized(visual_field.double().to(device)).float() + swaps = all_swaps(size, device) + + weights = {"pair": (1.0, 0.0), "triangle": (0.0, 1.0), "both": (1.0, 1.0)} + report = { + "protocol": ( + "Closed-form energies over all triples; the decision statistic " + "is the energy of the truth against the deepest state reached " + "by exact steepest descent and batched tempering. Hidden pairs " + "score only." + ), + "samples": size, + "energies": {}, + } + for name in (item.strip() for item in args.energies.split(",")): + pair_weight, triangle_weight = weights[name] + energy = ClosedFormEnergy( + text, visual, pair_weight, triangle_weight, args.chunk + ) + true_energy = float(energy.energy(truth[None])[0]) + random_batch = torch.stack( + [ + torch.argsort(torch.rand(size, generator=generator)) + for _ in range(200) + ] + ).to(device) + random_values = energy.energy(random_batch) + + kept, kept_value = steepest_descent( + energy, truth, swaps, args.descent_max_steps + ) + retention = float((kept.cpu() == truth.cpu()).float().mean()) + + quenches = [] + for restart in range(args.descent_restarts): + start = torch.argsort(torch.rand(size, generator=generator)).to(device) + final, value = steepest_descent( + energy, start, swaps, args.descent_max_steps + ) + quenches.append( + { + "energy": value, + "accuracy": float((final.cpu() == truth.cpu()).float().mean()), + } + ) + + starts = torch.stack( + [ + torch.argsort(torch.rand(size, generator=generator)) + for _ in range(args.replicas) + ] + ).to(device) + states, values = temper(energy, starts, args, generator, device) + cold = int(values.argmin()) + polished, polished_value = steepest_descent( + energy, states[cold], swaps, args.polish_steps + ) + tempering = { + "cold_energy": float(values[cold]), + "polished_energy": polished_value, + "polished_accuracy": float( + (polished.cpu() == truth.cpu()).float().mean() + ), + "best_replica_accuracy": max( + float((state.cpu() == truth.cpu()).float().mean()) + for state in states + ), + } + deepest = min( + [polished_value, float(values.min())] + + [item["energy"] for item in quenches] + ) + entry = { + "true_energy": true_energy, + "random_mean": float(random_values.mean()), + "true_z": float( + (random_values.mean() - true_energy) + / random_values.std().clamp_min(1e-12) + ), + "descent_retention_diagnostic": retention, + "descent_from_truth_energy": kept_value, + "quenches": quenches, + "tempering": tempering, + "deepest_seen": deepest, + "margin_over_true": deepest / abs(true_energy) - true_energy / abs(true_energy), + "passes": bool(deepest >= true_energy - 1e-9), + "recovery_accuracy": tempering["polished_accuracy"], + } + report["energies"][name] = entry + print( + json.dumps( + { + name: { + key: entry[key] + for key in ( + "true_energy", + "deepest_seen", + "passes", + "recovery_accuracy", + "descent_retention_diagnostic", + "true_z", + ) + } + } + ) + ) + write_json(args.output, report) + print(f"Wrote {args.output}") + + +if __name__ == "__main__": + main() diff --git a/worldalign/synth_precision_gate.py b/worldalign/synth_precision_gate.py new file mode 100644 index 0000000..46887fa --- /dev/null +++ b/worldalign/synth_precision_gate.py @@ -0,0 +1,279 @@ +"""Channel-matched (precision-weighted) relational energy for the synth world. + +Planted-problem theory: search is glassy when the energy mismatches the +generative channel. The orbit provides the channel unimodally -- the +variance of each vision relation entry across re-rendered views measures +exactly how corrupted that entry is (segmentation errors are the dominant +noise and are view-dependent), so entries are weighted by their orbit +precision. Text fields are parse-deterministic and enter unweighted. + +The weighted energy loses the all-pairs closed form (the quadratic terms +no longer cancel), but each swap delta is still O(N) and vectorizes over +sampled pairs; gates and tempering below use that form. Hidden pairs +score orderings only. +""" + +from __future__ import annotations + +import argparse +import json +from pathlib import Path + +import torch +import torch.nn.functional as F +from tqdm import tqdm + +from .common import read_json, seed_everything, write_json +from .synth_cc_battery import ( + component_descriptors, + moment_field, + onehot_descriptors, + phrase_bow_sets, +) + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--data-dir", default="artifacts/synth_v0") + parser.add_argument("--split", choices=["val", "test"], default="test") + parser.add_argument("--samples", type=int, default=512) + parser.add_argument("--merge-distance", type=float, default=30.0) + parser.add_argument("--vision-views", type=int, default=4) + parser.add_argument("--precision-floor", type=float, default=1e-4) + parser.add_argument("--random-perms", type=int, default=300) + parser.add_argument("--transposition-samples", type=int, default=100000) + parser.add_argument("--descent-restarts", type=int, default=5) + parser.add_argument("--descent-max-steps", type=int, default=5000) + parser.add_argument("--replicas", type=int, default=8) + parser.add_argument("--tempering-rounds", type=int, default=60000) + parser.add_argument("--temp-high", type=float, default=3e-3) + parser.add_argument("--temp-low", type=float, default=1e-5) + parser.add_argument("--exchange-every", type=int, default=20) + parser.add_argument("--device", default="cuda:3") + parser.add_argument("--seed", type=int, default=20260731) + parser.add_argument( + "--output", default="artifacts/synth_v0/precision_gate.json" + ) + return parser.parse_args() + + +def standardized(field: torch.Tensor) -> torch.Tensor: + mask = ~torch.eye(len(field), dtype=torch.bool) + values = field[mask] + out = (field - values.mean()) / values.std().clamp_min(1e-9) + return out.masked_fill(~mask, 0.0) + + +def weighted_energy( + text: torch.Tensor, visual: torch.Tensor, weights: torch.Tensor, + permutations: torch.Tensor, +) -> torch.Tensor: + fields = text[permutations[:, :, None], permutations[:, None, :]] + difference = (fields - visual) ** 2 * weights + mask = ~torch.eye(text.shape[-1], dtype=torch.bool, device=text.device) + return difference[:, mask].sum(-1) / weights[mask].sum() + + +def swap_deltas( + permuted_text: torch.Tensor, + visual: torch.Tensor, + weights: torch.Tensor, + pairs_p: torch.Tensor, + pairs_q: torch.Tensor, +) -> torch.Tensor: + """Weighted swap deltas, O(N) per proposal, batched over replicas. + + permuted_text: [R, N, N]; pairs: [R, P]. Row and column contributions + are equal by symmetry of all three matrices; the k in {p, q} terms are + excluded because the (p, q) entry itself is unchanged by the swap. + """ + replica_index = torch.arange(len(permuted_text), device=permuted_text.device) + rows_p = permuted_text[replica_index[:, None], pairs_p] + rows_q = permuted_text[replica_index[:, None], pairs_q] + visual_p, visual_q = visual[pairs_p], visual[pairs_q] + weight_p, weight_q = weights[pairs_p], weights[pairs_q] + new_p = (rows_q - visual_p) ** 2 * weight_p + old_p = (rows_p - visual_p) ** 2 * weight_p + new_q = (rows_p - visual_q) ** 2 * weight_q + old_q = (rows_q - visual_q) ** 2 * weight_q + total = (new_p - old_p + new_q - old_q).sum(-1) + columns = torch.stack([pairs_p, pairs_q], -1) + correction = torch.zeros_like(total) + for slot in range(2): + chosen = columns[..., slot] + correction = correction + ( + (rows_q.gather(-1, chosen[..., None]) - visual_p.gather(-1, chosen[..., None])) ** 2 + - (rows_p.gather(-1, chosen[..., None]) - visual_p.gather(-1, chosen[..., None])) ** 2 + ).squeeze(-1) * weight_p.gather(-1, chosen[..., None]).squeeze(-1) + correction = correction + ( + (rows_p.gather(-1, chosen[..., None]) - visual_q.gather(-1, chosen[..., None])) ** 2 + - (rows_q.gather(-1, chosen[..., None]) - visual_q.gather(-1, chosen[..., None])) ** 2 + ).squeeze(-1) * weight_q.gather(-1, chosen[..., None]).squeeze(-1) + mask_count = weights[~torch.eye(len(visual), dtype=torch.bool, device=visual.device)].sum() + return 2.0 * (total - correction) / mask_count + + +def main() -> None: + args = parse_args() + seed_everything(args.seed) + manifest = read_json(Path(args.data_dir, "manifest.json")) + captions = read_json(Path(args.data_dir, "captions.json"))["captions"] + rows = manifest[args.split][: args.samples] + image_dir = Path(manifest["image_dir"]) + from .synth_towers import load_image + + per_view = [] + for view in range(args.vision_views): + raw = [] + for row in tqdm(rows, desc=f"cc v{view}"): + sprites, _ = component_descriptors( + load_image(image_dir / f"scene{row:06d}_v{view}.png"), + args.merge_distance, + ) + raw.append(sprites) + sets = [F.normalize(v, dim=-1) for v in onehot_descriptors(raw)] + per_view.append(moment_field(sets)) + stack = torch.stack(per_view) + visual_field = stack.mean(0) + variance = stack.var(0) + weights = 1.0 / (variance + args.precision_floor) + weights = weights / weights.mean() + + text_sets = phrase_bow_sets(rows, captions, manifest["vocabulary"]) + text_field = moment_field(text_sets) + + device = torch.device(args.device) + size = len(rows) + generator = torch.Generator().manual_seed(args.seed) + hidden = torch.randperm(size, generator=generator) + truth = torch.argsort(hidden) + text_input = standardized(text_field[hidden][:, hidden].double()).float().to(device) + visual = standardized(visual_field.double()).float().to(device) + weight_map = weights.float().to(device) + + identity = torch.arange(size, device=device) + true_energy = float( + weighted_energy(text_input, visual, weight_map, truth[None].to(device))[0] + ) + random_perms = torch.stack( + [torch.argsort(torch.rand(size, generator=generator)) for _ in range(args.random_perms)] + ).to(device) + random_energies = weighted_energy(text_input, visual, weight_map, random_perms) + gate_a = { + "true": true_energy, + "random_mean": float(random_energies.mean()), + "random_std": float(random_energies.std()), + "true_z": float((random_energies.mean() - true_energy) / random_energies.std().clamp_min(1e-12)), + } + + samples = min(args.transposition_samples, size * (size - 1) // 2) + pairs_p = torch.randint(0, size, (1, samples), generator=generator).to(device) + pairs_q = torch.randint(0, size, (1, samples), generator=generator).to(device) + valid = (pairs_p != pairs_q).squeeze(0) + fields_true = text_input[truth.to(device)][:, truth.to(device)][None] + deltas = swap_deltas(fields_true, visual, weight_map, pairs_p, pairs_q).squeeze(0)[valid] + gate_b = { + "sampled": int(valid.sum()), + "improving_fraction": float((deltas < 0).float().mean()), + } + + def descent(start: torch.Tensor) -> dict: + current = start.clone() + energy = float(weighted_energy(text_input, visual, weight_map, current[None])[0]) + for _ in range(args.descent_max_steps): + fields = text_input[current][:, current][None] + cp = torch.randint(0, size, (1, 4096), generator=generator).to(device) + cq = torch.randint(0, size, (1, 4096), generator=generator).to(device) + dd = swap_deltas(fields, visual, weight_map, cp, cq).squeeze(0) + best = int(dd.argmin()) + if float(dd[best]) >= -1e-12: + break + p, q = int(cp[0, best]), int(cq[0, best]) + current[[p, q]] = current[[q, p]] + energy += float(dd[best]) + return { + "accuracy": float((current.cpu() == truth).float().mean()), + "energy": float(weighted_energy(text_input, visual, weight_map, current[None])[0]), + } + + from_true = descent(truth.to(device)) + restarts = [ + descent(torch.argsort(torch.rand(size, generator=generator)).to(device)) + for _ in range(args.descent_restarts) + ] + best_random = min(r["energy"] for r in restarts) + + # Tempering recovery with the weighted deltas. + temperatures = torch.logspace( + torch.log10(torch.tensor(args.temp_low)), + torch.log10(torch.tensor(args.temp_high)), + args.replicas, + ).to(device) + permutations = torch.stack( + [torch.randperm(size, generator=generator).to(device) for _ in range(args.replicas)] + ) + energies = weighted_energy(text_input, visual, weight_map, permutations) + for round_index in range(args.tempering_rounds): + fields = text_input[permutations[:, :, None], permutations[:, None, :]] + cp = torch.randint(0, size, (args.replicas, 24), generator=generator).to(device) + cq = torch.randint(0, size, (args.replicas, 24), generator=generator).to(device) + dd = swap_deltas(fields, visual, weight_map, cp, cq) + noise = torch.rand(args.replicas, 24, generator=generator).to(device) + ok = (dd < -temperatures[:, None] * noise.clamp_min(1e-12).log()) & (cp != cq) + for replica in range(args.replicas): + hits = torch.nonzero(ok[replica]) + if not len(hits): + continue + first = int(hits[0, 0]) + p, q = int(cp[replica, first]), int(cq[replica, first]) + permutations[replica][[p, q]] = permutations[replica][[q, p]] + energies[replica] = energies[replica] + dd[replica, first] + if round_index % args.exchange_every == 0: + for replica in range(args.replicas - 1): + gap = (energies[replica] - energies[replica + 1]) * ( + 1.0 / temperatures[replica] - 1.0 / temperatures[replica + 1] + ) + if gap > 0 or torch.rand(1, generator=generator).item() < float(gap.exp()): + permutations[[replica, replica + 1]] = permutations[[replica + 1, replica]] + energies[[replica, replica + 1]] = energies[[replica + 1, replica]] + if round_index % 10000 == 0: + energies = weighted_energy(text_input, visual, weight_map, permutations) + energies = weighted_energy(text_input, visual, weight_map, permutations) + accuracies = (permutations.cpu() == truth[None]).float().mean(-1) + cold = int(energies.argmin()) + + report = { + "protocol": ( + "Relation entries are weighted by orbit-derived precision " + "(variance of the vision field across re-rendered views); " + "weights are unimodal statistics. Hidden pairs score only." + ), + "samples": size, + "weight_stats": { + "min": float(weights.min()), + "median": float(weights.median()), + "max": float(weights.max()), + }, + "gate_a": gate_a, + "gate_b": gate_b, + "descent_from_true": from_true, + "descent_from_random": restarts, + "tempering": { + "cold_energy": float(energies[cold]), + "cold_accuracy": float(accuracies[cold]), + "best_accuracy": float(accuracies.max()), + }, + "verdict": { + "true_energy": true_energy, + "best_random_descent": best_random, + "counterfeit_found": bool(best_random < true_energy and min(r["accuracy"] for r in restarts) < 0.5), + "recovery_accuracy": float(accuracies.max()), + }, + } + print(json.dumps({"verdict": report["verdict"], "gate_b": gate_b, "from_true": from_true})) + write_json(args.output, report) + print(f"Wrote {args.output}") + + +if __name__ == "__main__": + main() diff --git a/worldalign/synth_probes.py b/worldalign/synth_probes.py new file mode 100644 index 0000000..180b96a --- /dev/null +++ b/worldalign/synth_probes.py @@ -0,0 +1,129 @@ +"""Ground-truth factor probes for synthetic-world representations. + +Linear probes from a representation to the discrete scene factors measure +which world variables survive the encoder and its pooling. Probes are +fitted per modality on that modality's own training split and evaluated +on held-out scenes; scene truth is generator metadata, so this is a +within-modality diagnostic, not a cross-modal signal. +""" + +from __future__ import annotations + +import argparse +import json +from pathlib import Path + +import numpy as np +import torch + +from .common import read_json, write_json +from .synth_world import COLORS, SHAPES + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--data-dir", default="artifacts/synth_v0") + parser.add_argument("--features", required=True) + parser.add_argument( + "--side", choices=["vision", "text"], required=True, + help="Selects the training split whose scenes fit the probes.", + ) + parser.add_argument("--ridge", type=float, default=1.0) + parser.add_argument("--output", required=True) + return parser.parse_args() + + +def scene_targets(scene: dict) -> dict[str, np.ndarray | float]: + color_presence = np.zeros(len(COLORS)) + shape_presence = np.zeros(len(SHAPES)) + for group in scene["groups"]: + color_presence[list(COLORS).index(group["color"])] = 1.0 + shape_presence[SHAPES.index(group["shape"])] = 1.0 + return { + "color_presence": color_presence, + "shape_presence": shape_presence, + "group_count": float(len(scene["groups"])), + "object_total": float(sum(g["count"] for g in scene["groups"])), + "relation_count": float(len(scene["relations"])), + } + + +def ridge_fit( + x: np.ndarray, y: np.ndarray, ridge: float +) -> tuple[np.ndarray, np.ndarray]: + x = np.concatenate([x, np.ones((len(x), 1))], axis=1) + gram = x.T @ x + ridge * np.eye(x.shape[1]) + weights = np.linalg.solve(gram, x.T @ y) + return weights, x @ weights + + +def evaluate( + features_train: np.ndarray, + features_test: np.ndarray, + train_targets: np.ndarray, + test_targets: np.ndarray, + ridge: float, + binary: bool, +) -> dict: + weights, _ = ridge_fit(features_train, train_targets, ridge) + prediction = ( + np.concatenate([features_test, np.ones((len(features_test), 1))], axis=1) + @ weights + ) + if binary: + accuracy = float(((prediction > 0.5) == (test_targets > 0.5)).mean()) + balanced = [] + for column in range(test_targets.shape[1]): + truth = test_targets[:, column] > 0.5 + if truth.any() and (~truth).any(): + hit = (prediction[:, column] > 0.5) == truth + balanced.append( + (hit[truth].mean() + hit[~truth].mean()) / 2.0 + ) + return { + "accuracy": accuracy, + "balanced_accuracy": float(np.mean(balanced)), + } + residual = prediction[:, 0] - test_targets[:, 0] + variance = test_targets[:, 0].var() + return { + "r2": float(1.0 - residual.var() / max(variance, 1e-9)), + "mae": float(np.abs(residual).mean()), + } + + +def main() -> None: + args = parse_args() + manifest = read_json(Path(args.data_dir, "manifest.json")) + scenes = read_json(Path(args.data_dir, "scenes.private.json"))["scenes"] + state = torch.load(args.features, map_location="cpu", weights_only=False) + lookup = {int(row): i for i, row in enumerate(state["rows"])} + features = state["features"].float().numpy() + + split_key = "vision_only_train" if args.side == "vision" else "text_only_train" + train_rows = [row for row in manifest[split_key] if row in lookup][:8000] + test_rows = [row for row in manifest["test"] if row in lookup] + + x_train = features[[lookup[row] for row in train_rows]] + x_test = features[[lookup[row] for row in test_rows]] + report = {"features": args.features, "side": args.side, "targets": {}} + for name in ("color_presence", "shape_presence"): + y_train = np.stack([scene_targets(scenes[row])[name] for row in train_rows]) + y_test = np.stack([scene_targets(scenes[row])[name] for row in test_rows]) + report["targets"][name] = evaluate( + x_train, x_test, y_train, y_test, args.ridge, binary=True + ) + for name in ("group_count", "object_total", "relation_count"): + y_train = np.array( + [[scene_targets(scenes[row])[name]] for row in train_rows] + ) + y_test = np.array([[scene_targets(scenes[row])[name]] for row in test_rows]) + report["targets"][name] = evaluate( + x_train, x_test, y_train, y_test, args.ridge, binary=False + ) + write_json(args.output, report) + print(json.dumps(report, indent=2)) + + +if __name__ == "__main__": + main() diff --git a/worldalign/synth_recovery.py b/worldalign/synth_recovery.py new file mode 100644 index 0000000..e045a72 --- /dev/null +++ b/worldalign/synth_recovery.py @@ -0,0 +1,184 @@ +"""Blind recovery on synthetic set-kernel fields: the end-to-end test. + +Builds the connected-component descriptor field and the phrase +bag-of-words field, hides the text order behind a shuffle, and runs +parallel tempering on the relational energy alone. Recovery accuracy +against the hidden truth is the first end-to-end measurement of world +matching in the closed world. +""" + +from __future__ import annotations + +import argparse +import json +from pathlib import Path + +import torch +from tqdm import tqdm + +from .blind_recovery import arm_tempering, permutation_energy_batch +from .common import read_json, seed_everything, write_json +from .manifold_gate import standardize_relation +from .synth_cc_battery import ( + component_descriptors, + moment_field, + onehot_descriptors, + phrase_bow_sets, +) +from .synth_set_battery import set_similarity_field +from .synth_towers import load_image + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--data-dir", default="artifacts/synth_v0") + parser.add_argument("--split", choices=["val", "test"], default="test") + parser.add_argument("--samples", type=int, default=512) + parser.add_argument("--merge-distance", type=float, default=30.0) + parser.add_argument("--vision-views", type=int, default=4) + parser.add_argument("--replicas", type=int, default=8) + parser.add_argument("--tempering-rounds", type=int, default=60000) + parser.add_argument("--temp-high", type=float, default=3e-3) + parser.add_argument("--temp-low", type=float, default=1e-5) + parser.add_argument("--exchange-every", type=int, default=20) + parser.add_argument("--unary-weight", type=float, default=0.0) + parser.add_argument("--init", default="random") + parser.add_argument("--residualize-size", action="store_true", default=False) + parser.add_argument( + "--features", choices=["descriptors", "onehot"], default="onehot" + ) + parser.add_argument("--kernel", choices=["matching", "moment"], default="moment") + parser.add_argument("--device", default="cuda:3") + parser.add_argument("--seed", type=int, default=20260731) + parser.add_argument( + "--output", default="artifacts/synth_v0/recovery_end_to_end.json" + ) + return parser.parse_args() + + +def residualize(field: torch.Tensor, sizes: torch.Tensor) -> torch.Tensor: + """Regress the set-size nuisance out of a matching-value field. + + Set sizes are unimodal observables; their sum, difference, and product + explain a size-driven component that differs between modalities and + is exploitable by counterfeit assignments. + """ + n = len(field) + mask = ~torch.eye(n, dtype=torch.bool) + features = torch.stack( + [ + (sizes[:, None] + sizes[None, :])[mask], + (sizes[:, None] - sizes[None, :]).abs()[mask], + (sizes[:, None] * sizes[None, :])[mask], + torch.ones(int(mask.sum()), dtype=torch.float64), + ], + dim=1, + ) + values = field.double()[mask] + solution = torch.linalg.lstsq(features, values[:, None]).solution + residual = values - (features @ solution).squeeze(1) + output = field.double().clone() + output[mask] = residual + output.fill_diagonal_(0.0) + return output.float() + + +def main() -> None: + args = parse_args() + seed_everything(args.seed) + manifest = read_json(Path(args.data_dir, "manifest.json")) + captions = read_json(Path(args.data_dir, "captions.json"))["captions"] + rows = manifest[args.split][: args.samples] + image_dir = Path(manifest["image_dir"]) + + import torch.nn.functional as F + + per_view_fields = [] + for view in range(args.vision_views): + raw_sets = [] + for row in tqdm(rows, desc=f"cc v{view}"): + image = load_image(image_dir / f"scene{row:06d}_v{view}.png") + sprites, _ = component_descriptors(image, args.merge_distance) + raw_sets.append(sprites) + if args.features == "onehot": + vision_sets = [ + F.normalize(v, dim=-1) for v in onehot_descriptors(raw_sets) + ] + else: + vision_sets = [F.normalize(v, dim=-1) for v in raw_sets] + if args.kernel == "moment": + per_view_fields.append(moment_field(vision_sets)) + else: + per_view_fields.append(set_similarity_field(vision_sets)) + visual_field = torch.stack(per_view_fields).mean(0) + text_sets = phrase_bow_sets(rows, captions, manifest["vocabulary"]) + if args.kernel == "moment": + text_field = moment_field(text_sets) + else: + text_field = set_similarity_field(text_sets) + + if args.residualize_size: + vision_sizes = torch.tensor( + [len(s) for s in vision_sets], dtype=torch.float64 + ) + text_sizes = torch.tensor([len(s) for s in text_sets], dtype=torch.float64) + visual_field = residualize(visual_field, vision_sizes) + text_field = residualize(text_field, text_sizes) + + device = torch.device(args.device) + size = len(rows) + generator = torch.Generator().manual_seed(args.seed) + hidden = torch.randperm(size, generator=generator) + truth = torch.argsort(hidden) + text_input = text_field[hidden][:, hidden] + + text_standardized = standardize_relation(text_input.double())[0].float().to(device) + visual_standardized = ( + standardize_relation(visual_field.double())[0].float().to(device) + ) + true_energy = float( + permutation_energy_batch( + text_standardized, visual_standardized, truth[None].to(device) + )[0] + ) + report = { + "protocol": ( + "Set-kernel fields from released renders and captions; text " + "order hidden behind a shuffle; tempering sees no truth. " + "Hidden truth scores the outcome only." + ), + "split": args.split, + "samples": size, + "true_energy": true_energy, + "chance_accuracy": 1.0 / size, + "tempering": arm_tempering( + text_standardized, + visual_standardized, + truth, + true_energy, + args, + generator, + unary=None, + ), + } + best = max( + report["tempering"]["replicas"] + [report["tempering"]["best"]], + key=lambda c: c["accuracy"], + ) + by_energy = min( + report["tempering"]["replicas"] + [report["tempering"]["best"]], + key=lambda c: c["energy"], + ) + report["summary"] = { + "best_accuracy": best["accuracy"], + "best_accuracy_energy_over_true": best["energy_over_true"], + "lowest_energy_accuracy": by_energy["accuracy"], + "lowest_energy_over_true": by_energy["energy_over_true"], + } + print(json.dumps({"summary": report["summary"]})) + write_json(args.output, report) + print(f"Wrote {args.output}") + + +if __name__ == "__main__": + main() diff --git a/worldalign/synth_set_battery.py b/worldalign/synth_set_battery.py new file mode 100644 index 0000000..9ed66ec --- /dev/null +++ b/worldalign/synth_set_battery.py @@ -0,0 +1,283 @@ +"""Set-kernel relation fields for the synthetic world: the structure battery. + +The pooled readout of a set representation destroys it. Here scene states +stay sets -- vision: slot vectors with alpha masses; text: per-group +phrase states parsed from the caption's enumeration sentence and encoded +individually -- and scene-to-scene relations are computed within each +modality as set-matching similarities. The cross-modal gate then runs on +these set-kernel relation fields exactly as on any relation channel. + +Phrase parsing reads only released captions; slot sets read only renders. +Hidden pairs score orderings, as always. +""" + +from __future__ import annotations + +import argparse +import json +import re +from pathlib import Path + +import torch +import torch.nn.functional as F +from scipy.optimize import linear_sum_assignment +from tqdm import tqdm + +from .common import batch_indices, read_json, seed_everything, write_json +from .manifold_gate import standardize_relation +from .ricci_control import run_gates +from .synth_slots import SlotAutoencoder +from .synth_towers import TextTower, load_image, tokenize + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--data-dir", default="artifacts/synth_v0") + parser.add_argument( + "--slots", default="artifacts/synth_v0/vision_slots.pt" + ) + parser.add_argument("--text-tower", default="artifacts/synth_v0/text_tower.pt") + parser.add_argument("--split", choices=["val", "test"], default="test") + parser.add_argument("--samples", type=int, default=512) + parser.add_argument("--mass-floor", type=float, default=0.02) + parser.add_argument( + "--vision-mode", choices=["slot_vectors", "sprites"], default="sprites" + ) + parser.add_argument( + "--slot-tower", default="artifacts/synth_v0/slot_tower.pt" + ) + parser.add_argument("--sprite-window", type=int, default=48) + parser.add_argument("--random-perms", type=int, default=300) + parser.add_argument("--descent-restarts", type=int, default=5) + 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="cuda:3") + parser.add_argument("--seed", type=int, default=20260731) + parser.add_argument( + "--output", default="artifacts/synth_v0/set_battery_gate.json" + ) + return parser.parse_args() + + +def parse_group_phrases(caption: str) -> list[str]: + """Group phrases from the enumeration sentence of a released caption.""" + first = caption.split(".")[0] + for opener in ("there are ", "the picture shows ", "you can see "): + if first.startswith(opener): + first = first[len(opener):] + break + first = first.replace(" and ", ", ") + return [phrase.strip() for phrase in first.split(",") if phrase.strip()] + + +@torch.inference_mode() +def text_group_sets( + rows: list[int], captions: list[list[str]], args: argparse.Namespace +) -> list[torch.Tensor]: + state = torch.load(args.text_tower, map_location="cpu", weights_only=False) + saved = state["args"] + vocab = state["vocab"] + model = TextTower( + len(vocab), + saved["text_dim"], + saved["depth"], + saved.get("text_heads", 4), + saved["context"], + ).to(args.device) + model.load_state_dict(state["model"]) + model.eval() + phrases_per_row = [parse_group_phrases(captions[row][0]) for row in rows] + flat = [ + (index, phrase) + for index, phrases in enumerate(phrases_per_row) + for phrase in phrases + ] + states: list[list[torch.Tensor]] = [[] for _ in rows] + for indices in tqdm(list(batch_indices(len(flat), 256)), desc="text sets"): + batch = [flat[i] for i in indices] + sequences = [tokenize(phrase, vocab) for _, phrase in batch] + longest = max(len(s) for s in sequences) + tokens = torch.zeros(len(batch), longest, dtype=torch.long) + for row, sequence in enumerate(sequences): + tokens[row, : len(sequence)] = torch.tensor(sequence) + tokens = tokens.to(args.device) + hidden = model(tokens) + mask = (tokens != 0).float()[..., None] + pooled = (hidden * mask).sum(1) / mask.sum(1).clamp_min(1.0) + for (index, _), vector in zip(batch, pooled.float().cpu()): + states[index].append(vector) + return [F.normalize(torch.stack(s), dim=-1) for s in states] + + +def set_similarity_field( + sets: list[torch.Tensor], weights: list[torch.Tensor] | None = None +) -> torch.Tensor: + """Symmetric matching-value similarity between all set pairs.""" + n = len(sets) + field = torch.zeros(n, n) + for a in range(n): + for b in range(a, n): + similarity = sets[a] @ sets[b].T + if weights is not None: + similarity = similarity * torch.sqrt( + weights[a][:, None] * weights[b][None, :] + ) + rows, cols = linear_sum_assignment(-similarity.numpy()) + value = float(similarity[rows, cols].sum()) / max( + min(similarity.shape), 1 + ) + field[a, b] = field[b, a] = value + return field + + +@torch.inference_mode() +def decode_sprites( + rows: list[int], + lookup: dict[int, int], + slot_sets_all: torch.Tensor, + args: argparse.Namespace, +) -> list[torch.Tensor]: + """Centered per-slot appearance sprites from the tower's own decoder. + + The joint decode gives each slot an rgb map and a competitive alpha + mask; centering the masked appearance at the alpha centroid removes + layout, leaving color, shape, size, and multiplicity pattern. + """ + manifest = read_json(Path(args.data_dir, "manifest.json")) + state = torch.load(args.slot_tower, map_location="cpu", weights_only=False) + saved = state["args"] + model = SlotAutoencoder( + manifest["image_size"], saved["slots"], saved["slot_dim"], saved["iterations"] + ).to(args.device) + model.load_state_dict(state["model"]) + model.eval() + size = manifest["image_size"] + window = args.sprite_window + axis = torch.arange(size, dtype=torch.float32, device=args.device) + sprites: list[torch.Tensor] = [] + for start in tqdm(range(0, len(rows), 64), desc="sprites"): + batch_rows = rows[start : start + 64] + slots = torch.stack( + [slot_sets_all[lookup[int(row)]][0] for row in batch_rows] + ).to(args.device) + rgb_alpha_rgb, alpha = model.decode(slots) + del rgb_alpha_rgb + # Re-decode retaining per-slot rgb: replicate decode internals. + batch, count, dim = slots.shape + x = slots.reshape(batch * count, dim, 1, 1).expand( + -1, -1, model.broadcast, model.broadcast + ) + from .synth_slots import coordinate_grid + + grid = coordinate_grid(model.broadcast, slots.device).reshape(1, -1, 4) + position = model.position_decoder(grid).transpose(1, 2).reshape( + 1, dim, model.broadcast, model.broadcast + ) + decoded = model.decoder(x + position) + decoded = F.interpolate( + decoded, size=size, mode="bilinear", align_corners=False + ).reshape(batch, count, 4, size, size) + rgb = decoded[:, :, :3] + masked = rgb * alpha # [B, K, 3, H, W] + weight_y = alpha.squeeze(2).sum(-1) # [B, K, H] + weight_x = alpha.squeeze(2).sum(-2) # [B, K, W] + cy = (weight_y * axis).sum(-1) / weight_y.sum(-1).clamp_min(1e-6) + cx = (weight_x * axis).sum(-1) / weight_x.sum(-1).clamp_min(1e-6) + half = window // 2 + batch_sprites = torch.zeros(batch, count, 3, window, window) + padded = F.pad(masked, (half, half, half, half)) + for b in range(batch): + for k in range(count): + y0 = int(cy[b, k].round()) + x0 = int(cx[b, k].round()) + batch_sprites[b, k] = padded[ + b, k, :, y0 : y0 + window, x0 : x0 + window + ].cpu() + sprites.extend(batch_sprites.flatten(2).unbind(0)) + return sprites + + +def main() -> None: + args = parse_args() + seed_everything(args.seed) + manifest = read_json(Path(args.data_dir, "manifest.json")) + captions = read_json(Path(args.data_dir, "captions.json"))["captions"] + rows = manifest[args.split][: args.samples] + + slot_state = torch.load(args.slots, map_location="cpu", weights_only=False) + lookup = {int(row): i for i, row in enumerate(slot_state["rows"])} + slot_sets_all = slot_state["slot_sets"] + masses_all = slot_state["slot_masses"] + sprites_all = None + if args.vision_mode == "sprites": + sprites_all = decode_sprites(rows, lookup, slot_sets_all, args) + vision_sets, vision_weights = [], [] + for position, row in enumerate(rows): + index = lookup[int(row)] + # Slot identity is not stable across forwards, so views cannot be + # averaged slot-wise; one view keeps object-slot binding intact. + slots = slot_sets_all[index][0] # [K, D] + mass = masses_all[index][0] + dominant = mass.argmax() + keep = torch.ones(len(mass), dtype=torch.bool) + keep[dominant] = False + keep &= mass > args.mass_floor + if not keep.any(): + keep = torch.ones(len(mass), dtype=torch.bool) + if sprites_all is not None: + vision_sets.append(F.normalize(sprites_all[position][keep], dim=-1)) + else: + vision_sets.append(F.normalize(slots[keep], dim=-1)) + weight = mass[keep] + vision_weights.append(weight / weight.sum().clamp_min(1e-8)) + + text_sets = text_group_sets(rows, captions, args) + + print(json.dumps({"building": "set similarity fields"})) + visual_field = set_similarity_field(vision_sets, vision_weights) + text_field = set_similarity_field(text_sets) + + visual_channels = standardize_relation(visual_field.double())[0][None] + text_channels = standardize_relation(text_field.double())[0][None] + generator = torch.Generator().manual_seed(args.seed) + report = { + "protocol": ( + "Scene states are sets (slot vectors; per-group phrase " + "states); within-modality relations are set-matching values; " + "hidden pairs score orderings only." + ), + "split": args.split, + "samples": len(rows), + "mean_vision_set_size": float( + torch.tensor([len(s) for s in vision_sets]).float().mean() + ), + "mean_text_set_size": float( + torch.tensor([len(s) for s in text_sets]).float().mean() + ), + **run_gates(text_channels, visual_channels, args, generator), + } + verdict = { + "true_z_mse": report["gate_a"]["random"]["mse"]["true_z"], + "improving_fraction": report["gate_b"]["improving_fraction"], + "descent_keeps": report["descent_from_true"]["final_accuracy"], + "true_mse": report["gate_a"]["true"]["mse"], + "best_random_descent": min( + (r["final_objective"] for r in report["descent_from_random"]), + default=None, + ), + } + verdict["counterfeit_found"] = bool( + verdict["best_random_descent"] is not None + and verdict["best_random_descent"] < verdict["true_mse"] + ) + report["verdict"] = verdict + print(json.dumps({"verdict": verdict})) + write_json(args.output, report) + print(f"Wrote {args.output}") + + +if __name__ == "__main__": + main() diff --git a/worldalign/synth_slots.py b/worldalign/synth_slots.py new file mode 100644 index 0000000..72178d8 --- /dev/null +++ b/worldalign/synth_slots.py @@ -0,0 +1,264 @@ +"""Object-centric vision tower for the synthetic world: slot attention. + +A slot-attention autoencoder decomposes each render into K competing +slots that jointly reconstruct the image through per-slot alpha masks. +Reconstruction cannot shortcut (every object must be painted by some +slot), and the state is a SET of object vectors rather than a pooled +summary -- the literal form of the node-as-object-set doctrine. Training +is plain unimodal autoencoding on the vision-only split. + +Extraction emits a flickr-schema vision file (pooled states for the +existing gate stack) plus the slot sets and alpha masses for the +structure battery. +""" + +from __future__ import annotations + +import argparse +import json +from pathlib import Path + +import torch +import torch.nn.functional as F +from torch import nn +from tqdm import tqdm + +from .common import batch_indices, read_json, seed_everything +from .synth_towers import load_image + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--mode", choices=["train", "extract"], required=True) + parser.add_argument("--data-dir", default="artifacts/synth_v0") + parser.add_argument("--slots", type=int, default=6) + parser.add_argument("--slot-dim", type=int, default=64) + parser.add_argument("--iterations", type=int, default=3) + parser.add_argument("--epochs", type=int, default=40) + parser.add_argument("--batch-size", type=int, default=64) + parser.add_argument("--lr", type=float, default=4e-4) + parser.add_argument("--warmup-steps", type=int, default=1500) + parser.add_argument("--device", default="cuda:3") + parser.add_argument("--seed", type=int, default=20260730) + parser.add_argument("--checkpoint", default="artifacts/synth_v0/slot_tower.pt") + parser.add_argument( + "--vision-output", default="artifacts/synth_v0/vision_slots.pt" + ) + return parser.parse_args() + + +def coordinate_grid(size: int, device: torch.device) -> torch.Tensor: + axis = torch.linspace(0.0, 1.0, size, device=device) + y, x = torch.meshgrid(axis, axis, indexing="ij") + return torch.stack([x, y, 1 - x, 1 - y], dim=-1) + + +class SlotAttention(nn.Module): + def __init__(self, slots: int, dim: int, iterations: int) -> None: + super().__init__() + self.slots = slots + self.iterations = iterations + self.scale = dim**-0.5 + self.mu = nn.Parameter(torch.randn(1, 1, dim) * 0.02) + self.log_sigma = nn.Parameter(torch.zeros(1, 1, dim)) + self.norm_input = nn.LayerNorm(dim) + self.norm_slots = nn.LayerNorm(dim) + self.norm_mlp = nn.LayerNorm(dim) + self.project_q = nn.Linear(dim, dim, bias=False) + self.project_k = nn.Linear(dim, dim, bias=False) + self.project_v = nn.Linear(dim, dim, bias=False) + self.gru = nn.GRUCell(dim, dim) + self.mlp = nn.Sequential(nn.Linear(dim, dim * 2), nn.ReLU(), nn.Linear(dim * 2, dim)) + + def forward(self, inputs: torch.Tensor) -> torch.Tensor: + batch, _, dim = inputs.shape + inputs = self.norm_input(inputs) + k = self.project_k(inputs) + v = self.project_v(inputs) + slots = self.mu + self.log_sigma.exp() * torch.randn( + batch, self.slots, dim, device=inputs.device + ) + for _ in range(self.iterations): + previous = slots + q = self.project_q(self.norm_slots(slots)) + attention = F.softmax( + torch.einsum("bkd,bnd->bkn", q, k) * self.scale, dim=1 + ) + attention = attention / attention.sum(-1, keepdim=True).clamp_min(1e-8) + updates = torch.einsum("bkn,bnd->bkd", attention, v) + slots = self.gru( + updates.reshape(-1, dim), previous.reshape(-1, dim) + ).reshape(batch, self.slots, dim) + slots = slots + self.mlp(self.norm_mlp(slots)) + return slots + + +class SlotAutoencoder(nn.Module): + def __init__(self, image_size: int, slots: int, dim: int, iterations: int) -> None: + super().__init__() + self.image_size = image_size + self.encoder = nn.Sequential( + nn.Conv2d(3, dim, 5, 2, 2), nn.ReLU(), + nn.Conv2d(dim, dim, 5, 2, 2), nn.ReLU(), + nn.Conv2d(dim, dim, 5, 1, 2), nn.ReLU(), + nn.Conv2d(dim, dim, 5, 1, 2), nn.ReLU(), + ) + self.grid_size = image_size // 4 + self.position_encoder = nn.Linear(4, dim) + self.norm = nn.LayerNorm(dim) + self.pre_mlp = nn.Sequential(nn.Linear(dim, dim), nn.ReLU(), nn.Linear(dim, dim)) + self.slot_attention = SlotAttention(slots, dim, iterations) + self.broadcast = 8 + self.position_decoder = nn.Linear(4, dim) + self.decoder = nn.Sequential( + nn.ConvTranspose2d(dim, dim, 4, 2, 1), nn.ReLU(), + nn.ConvTranspose2d(dim, dim, 4, 2, 1), nn.ReLU(), + nn.ConvTranspose2d(dim, dim, 4, 2, 1), nn.ReLU(), + nn.ConvTranspose2d(dim, dim, 4, 2, 1), nn.ReLU(), + nn.Conv2d(dim, 4, 3, 1, 1), + ) + + def encode(self, pixels: torch.Tensor) -> torch.Tensor: + features = self.encoder(pixels) + batch, dim = features.shape[:2] + features = features.flatten(2).transpose(1, 2) + grid = coordinate_grid(self.grid_size, pixels.device).reshape(1, -1, 4) + features = features + self.position_encoder(grid) + return self.slot_attention(self.pre_mlp(self.norm(features))) + + def decode(self, slots: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + batch, count, dim = slots.shape + x = slots.reshape(batch * count, dim, 1, 1).expand( + -1, -1, self.broadcast, self.broadcast + ) + grid = coordinate_grid(self.broadcast, slots.device).reshape(1, -1, 4) + position = self.position_decoder(grid).transpose(1, 2).reshape( + 1, dim, self.broadcast, self.broadcast + ) + decoded = self.decoder(x + position) + decoded = F.interpolate( + decoded, size=self.image_size, mode="bilinear", align_corners=False + ) + decoded = decoded.reshape(batch, count, 4, self.image_size, self.image_size) + rgb, alpha = decoded[:, :, :3], decoded[:, :, 3:] + alpha = F.softmax(alpha, dim=1) + return (rgb * alpha).sum(1), alpha + + def forward(self, pixels: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + slots = self.encode(pixels) + reconstruction, alpha = self.decode(slots) + return reconstruction, slots, alpha + + +def train(args: argparse.Namespace) -> None: + manifest = read_json(Path(args.data_dir, "manifest.json")) + rows = manifest["vision_only_train"] + views = manifest["visual_views"] + image_dir = Path(manifest["image_dir"]) + model = SlotAutoencoder( + manifest["image_size"], args.slots, args.slot_dim, args.iterations + ).to(args.device) + optimizer = torch.optim.Adam(model.parameters(), lr=args.lr) + jobs = [(row, view) for row in rows for view in range(views)] + steps_total = args.epochs * (len(jobs) // args.batch_size) + schedule = torch.optim.lr_scheduler.LambdaLR( + optimizer, + lambda step: min(1.0, step / max(args.warmup_steps, 1)) + * 0.5 + * (1 + torch.cos(torch.tensor(step / max(steps_total, 1) * 3.14159)).item()), + ) + import random + + rng = random.Random(args.seed) + step = 0 + for epoch in range(args.epochs): + rng.shuffle(jobs) + total, count = 0.0, 0 + for start in range(0, len(jobs) - args.batch_size + 1, args.batch_size): + batch = jobs[start : start + args.batch_size] + pixels = torch.stack( + [ + load_image(image_dir / f"scene{row:06d}_v{view}.png") + for row, view in batch + ] + ).to(args.device) + reconstruction, _, _ = model(pixels) + loss = F.mse_loss(reconstruction, pixels) + optimizer.zero_grad(set_to_none=True) + loss.backward() + torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) + optimizer.step() + schedule.step() + step += 1 + total += float(loss) + count += 1 + if epoch % 5 == 0 or epoch == args.epochs - 1: + print(json.dumps({"epoch": epoch, "loss": total / max(count, 1)})) + torch.save({"model": model.state_dict(), "args": vars(args)}, args.checkpoint) + print(f"Wrote {args.checkpoint}") + + +@torch.inference_mode() +def extract(args: argparse.Namespace) -> None: + manifest = read_json(Path(args.data_dir, "manifest.json")) + state = torch.load(args.checkpoint, map_location="cpu", weights_only=False) + saved = state["args"] + model = SlotAutoencoder( + manifest["image_size"], saved["slots"], saved["slot_dim"], saved["iterations"] + ).to(args.device) + model.load_state_dict(state["model"]) + model.eval() + views = manifest["visual_views"] + image_dir = Path(manifest["image_dir"]) + rendered_rows = sorted( + set(manifest["vision_only_train"]) | set(manifest["val"]) | set(manifest["test"]) + ) + jobs = [(row, view) for row in rendered_rows for view in range(views)] + slot_sets, masses = [], [] + for indices in tqdm(list(batch_indices(len(jobs), args.batch_size)), desc="slots"): + pixels = torch.stack( + [ + load_image(image_dir / f"scene{jobs[i][0]:06d}_v{jobs[i][1]}.png") + for i in indices + ] + ).to(args.device) + _, slots, alpha = model(pixels) + slot_sets.append(slots.float().cpu()) + masses.append(alpha.mean(dim=(2, 3, 4)).float().cpu()) + slot_sets = torch.cat(slot_sets).reshape( + len(rendered_rows), views, saved["slots"], saved["slot_dim"] + ) + masses = torch.cat(masses).reshape(len(rendered_rows), views, saved["slots"]) + # Pooled state: mass-weighted mean of the non-dominant slots. The + # largest-mass slot is the background in this world (the scene is + # mostly background) and would otherwise dominate the pool. + dominant = masses.argmax(-1, keepdim=True) + keep = torch.ones_like(masses).scatter(-1, dominant, 0.0) + weights = (masses * keep).clamp_min(1e-8) + weights = weights / weights.sum(-1, keepdim=True) + pooled = (slot_sets * weights[..., None]).sum(2).mean(1) + torch.save( + { + "model": "synth_slot_tower", + "rows": rendered_rows, + "features": F.normalize(pooled, dim=-1), + "slot_sets": slot_sets, + "slot_masses": masses, + "views_per_scene": views, + }, + args.vision_output, + ) + print(f"Wrote {args.vision_output}") + + +def main() -> None: + args = parse_args() + seed_everything(args.seed) + if args.mode == "train": + train(args) + else: + extract(args) + + +if __name__ == "__main__": + main() diff --git a/worldalign/synth_towers.py b/worldalign/synth_towers.py new file mode 100644 index 0000000..15221ac --- /dev/null +++ b/worldalign/synth_towers.py @@ -0,0 +1,348 @@ +"""From-scratch unimodal towers for the synthetic closed world. + +Vision: a compact ViT trained with InfoNCE over natural orbit positives -- +two renders of the same scene, which differ exactly by the world's +continuous nuisance. No flips (they would erase left/right semantics) and +no color jitter (color is world content). Scene identity within the +vision-only split is unimodal metadata. + +Text: a small word-level causal LM trained on the captions of the +text-only split. Both towers therefore learn the same world from disjoint +scenes and separate modalities, with dialable capacity. +""" + +from __future__ import annotations + +import argparse +import json +import math +import random +from pathlib import Path + +import numpy as np +import torch +import torch.nn.functional as F +from PIL import Image +from torch import nn + +from .common import read_json, seed_everything + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--side", choices=["vision", "text"], required=True) + parser.add_argument( + "--objective", choices=["infonce", "simmim", "hybrid", "data2vec"], default="infonce" + ) + parser.add_argument("--mask-ratio", type=float, default=0.6) + parser.add_argument("--recon-weight", type=float, default=25.0) + parser.add_argument("--ema-decay", type=float, default=0.999) + parser.add_argument("--data-dir", default="artifacts/synth_v0") + parser.add_argument("--dim", type=int, default=192) + parser.add_argument("--depth", type=int, default=6) + parser.add_argument("--heads", type=int, default=3) + parser.add_argument("--text-dim", type=int, default=256) + parser.add_argument("--text-heads", type=int, default=4) + parser.add_argument("--patch", type=int, default=16) + parser.add_argument("--context", type=int, default=80) + parser.add_argument("--epochs", type=int, default=80) + parser.add_argument("--batch-size", type=int, default=256) + parser.add_argument("--lr", type=float, default=1e-3) + parser.add_argument("--temperature", type=float, default=0.2) + parser.add_argument("--device", default="cuda:3") + parser.add_argument("--seed", type=int, default=20260730) + parser.add_argument("--output", required=True) + return parser.parse_args() + + +class Block(nn.Module): + def __init__(self, dim: int, heads: int) -> None: + super().__init__() + self.norm1 = nn.LayerNorm(dim) + self.attention = nn.MultiheadAttention(dim, heads, batch_first=True) + self.norm2 = nn.LayerNorm(dim) + self.mlp = nn.Sequential( + nn.Linear(dim, dim * 4), nn.GELU(), nn.Linear(dim * 4, dim) + ) + + def forward( + self, x: torch.Tensor, mask: torch.Tensor | None = None + ) -> torch.Tensor: + normed = self.norm1(x) + attended, _ = self.attention( + normed, normed, normed, attn_mask=mask, need_weights=False + ) + x = x + attended + return x + self.mlp(self.norm2(x)) + + +class VisionTower(nn.Module): + def __init__(self, image_size: int, patch: int, dim: int, depth: int, heads: int) -> None: + super().__init__() + self.patch = patch + self.patch_embed = nn.Conv2d(3, dim, kernel_size=patch, stride=patch) + tokens = (image_size // patch) ** 2 + self.cls = nn.Parameter(torch.zeros(1, 1, dim)) + self.mask_token = nn.Parameter(torch.zeros(1, 1, dim)) + self.positions = nn.Parameter(torch.zeros(1, tokens + 1, dim)) + nn.init.trunc_normal_(self.positions, std=0.02) + nn.init.trunc_normal_(self.cls, std=0.02) + nn.init.trunc_normal_(self.mask_token, std=0.02) + self.blocks = nn.ModuleList(Block(dim, heads) for _ in range(depth)) + self.norm = nn.LayerNorm(dim) + self.head = nn.Sequential( + nn.Linear(dim, dim), nn.GELU(), nn.Linear(dim, 128) + ) + self.reconstruction = nn.Linear(dim, patch * patch * 3) + self.feature_head = nn.Linear(dim, dim) + + def encode( + self, pixels: torch.Tensor, mask: torch.Tensor | None = None + ) -> torch.Tensor: + x = self.patch_embed(pixels).flatten(2).transpose(1, 2) + if mask is not None: + x = torch.where(mask[..., None], self.mask_token.expand_as(x), x) + x = torch.cat([self.cls.expand(len(x), -1, -1), x], dim=1) + x = x + self.positions + for block in self.blocks: + x = block(x) + return self.norm(x) + + def forward(self, pixels: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + tokens = self.encode(pixels) + return tokens[:, 0], self.head(tokens[:, 0]) + + +class TextTower(nn.Module): + def __init__(self, vocab: int, dim: int, depth: int, heads: int, context: int) -> None: + super().__init__() + self.embed = nn.Embedding(vocab, dim) + self.positions = nn.Parameter(torch.zeros(1, context, dim)) + nn.init.trunc_normal_(self.positions, std=0.02) + self.blocks = nn.ModuleList(Block(dim, heads) for _ in range(depth)) + self.norm = nn.LayerNorm(dim) + self.context = context + + def forward(self, tokens: torch.Tensor) -> torch.Tensor: + length = tokens.shape[1] + x = self.embed(tokens) + self.positions[:, :length] + mask = torch.triu( + torch.full((length, length), float("-inf"), device=tokens.device), 1 + ) + for block in self.blocks: + x = block(x, mask) + return self.norm(x) + + def logits(self, hidden: torch.Tensor) -> torch.Tensor: + return hidden @ self.embed.weight.T + + +def load_image(path: Path) -> torch.Tensor: + with Image.open(path) as image: + array = np.asarray(image.convert("RGB"), dtype=np.float32) / 255.0 + return torch.from_numpy(array).permute(2, 0, 1) + + +def patchify(pixels: torch.Tensor, patch: int) -> torch.Tensor: + batch, channels, height, width = pixels.shape + grid = height // patch + x = pixels.reshape(batch, channels, grid, patch, grid, patch) + return x.permute(0, 2, 4, 3, 5, 1).reshape(batch, grid * grid, -1) + + +def train_vision(args: argparse.Namespace) -> None: + manifest = read_json(Path(args.data_dir, "manifest.json")) + rows = manifest["vision_only_train"] + views = manifest["visual_views"] + image_dir = Path(manifest["image_dir"]) + model = VisionTower( + manifest["image_size"], args.patch, args.dim, args.depth, args.heads + ).to(args.device) + optimizer = torch.optim.AdamW(model.parameters(), lr=args.lr, weight_decay=0.05) + tokens = (manifest["image_size"] // args.patch) ** 2 + if args.objective in ("infonce", "hybrid"): + samples_per_epoch = len(rows) + else: + samples_per_epoch = len(rows) * views + teacher = None + if args.objective == "data2vec": + import copy + + teacher = copy.deepcopy(model) + for parameter in teacher.parameters(): + parameter.requires_grad_(False) + steps_per_epoch = max(1, samples_per_epoch // args.batch_size) + schedule = torch.optim.lr_scheduler.CosineAnnealingLR( + optimizer, T_max=args.epochs * steps_per_epoch + ) + rng = random.Random(args.seed) + for epoch in range(args.epochs): + total, count = 0.0, 0 + if args.objective in ("infonce", "hybrid"): + order = rows.copy() + rng.shuffle(order) + for start in range(0, len(order) - args.batch_size + 1, args.batch_size): + batch_rows = order[start : start + args.batch_size] + pairs = [] + for row in batch_rows: + first, second = rng.sample(range(views), 2) + pairs.append( + load_image(image_dir / f"scene{row:06d}_v{first}.png") + ) + pairs.append( + load_image(image_dir / f"scene{row:06d}_v{second}.png") + ) + pixels = torch.stack(pairs).to(args.device) + _, projected = model(pixels) + projected = F.normalize(projected, dim=-1) + logits = projected @ projected.T / args.temperature + logits.fill_diagonal_(float("-inf")) + targets = torch.arange(len(projected), device=args.device) ^ 1 + loss = F.cross_entropy(logits, targets) + if args.objective == "hybrid": + mask = ( + torch.rand(len(pixels), tokens, device=args.device) + < args.mask_ratio + ) + encoded = model.encode(pixels, mask=mask) + predicted = model.reconstruction(encoded[:, 1:][mask]) + target_patches = patchify(pixels, args.patch)[mask] + loss = loss + args.recon_weight * F.mse_loss( + predicted, target_patches + ) + optimizer.zero_grad(set_to_none=True) + loss.backward() + optimizer.step() + schedule.step() + total += float(loss) + count += 1 + else: + jobs = [(row, view) for row in rows for view in range(views)] + rng.shuffle(jobs) + for start in range(0, len(jobs) - args.batch_size + 1, args.batch_size): + batch = jobs[start : start + args.batch_size] + pixels = torch.stack( + [ + load_image(image_dir / f"scene{row:06d}_v{view}.png") + for row, view in batch + ] + ).to(args.device) + mask = ( + torch.rand(len(batch), tokens, device=args.device) + < args.mask_ratio + ) + encoded = model.encode(pixels, mask=mask) + if args.objective == "simmim": + predicted = model.reconstruction(encoded[:, 1:][mask]) + target = patchify(pixels, args.patch)[mask] + loss = F.mse_loss(predicted, target) + else: + with torch.no_grad(): + reference = teacher.encode(pixels)[:, 1:] + reference = F.layer_norm( + reference, reference.shape[-1:] + ) + predicted = model.feature_head(encoded[:, 1:][mask]) + loss = F.smooth_l1_loss(predicted, reference[mask]) + optimizer.zero_grad(set_to_none=True) + loss.backward() + optimizer.step() + schedule.step() + if teacher is not None: + with torch.no_grad(): + for student_p, teacher_p in zip( + model.parameters(), teacher.parameters() + ): + teacher_p.lerp_(student_p, 1.0 - args.ema_decay) + total += float(loss) + count += 1 + if epoch % 5 == 0 or epoch == args.epochs - 1: + print(json.dumps({"epoch": epoch, "loss": total / max(count, 1)})) + torch.save( + {"model": model.state_dict(), "args": vars(args), "side": "vision"}, + args.output, + ) + print(f"Wrote {args.output}") + + +def tokenize(caption: str, vocab: dict[str, int]) -> list[int]: + tokens = caption.replace(",", " ").replace(".", " .").split() + return [vocab["<bos>"]] + [vocab[token] for token in tokens] + + +def build_vocab(manifest: dict) -> dict[str, int]: + words = list(manifest["vocabulary"]) + ["."] + vocab = {"<pad>": 0, "<bos>": 1} + for word in sorted(set(words)): + vocab.setdefault(word, len(vocab)) + return vocab + + +def train_text(args: argparse.Namespace) -> None: + manifest = read_json(Path(args.data_dir, "manifest.json")) + captions = read_json(Path(args.data_dir, "captions.json"))["captions"] + vocab = build_vocab(manifest) + sentences = [ + tokenize(caption, vocab) + for row in manifest["text_only_train"] + for caption in captions[row] + ] + model = TextTower( + len(vocab), args.text_dim, args.depth, args.text_heads, args.context + ).to(args.device) + optimizer = torch.optim.AdamW(model.parameters(), lr=args.lr, weight_decay=0.01) + rng = random.Random(args.seed) + steps_per_epoch = max(1, len(sentences) // args.batch_size) + schedule = torch.optim.lr_scheduler.CosineAnnealingLR( + optimizer, T_max=args.epochs * steps_per_epoch + ) + for epoch in range(args.epochs): + rng.shuffle(sentences) + total, count = 0.0, 0 + for start in range(0, len(sentences) - args.batch_size + 1, args.batch_size): + batch = sentences[start : start + args.batch_size] + longest = min(args.context, max(len(s) for s in batch)) + tokens = torch.zeros(len(batch), longest, dtype=torch.long) + for index, sentence in enumerate(batch): + clipped = sentence[:longest] + tokens[index, : len(clipped)] = torch.tensor(clipped) + tokens = tokens.to(args.device) + hidden = model(tokens) + logits = model.logits(hidden[:, :-1]) + targets = tokens[:, 1:] + loss = F.cross_entropy( + logits.reshape(-1, logits.shape[-1]), + targets.reshape(-1), + ignore_index=0, + ) + optimizer.zero_grad(set_to_none=True) + loss.backward() + optimizer.step() + schedule.step() + total += float(loss) + count += 1 + if epoch % 5 == 0 or epoch == args.epochs - 1: + print(json.dumps({"epoch": epoch, "loss": total / max(count, 1)})) + torch.save( + { + "model": model.state_dict(), + "vocab": vocab, + "args": vars(args), + "side": "text", + }, + args.output, + ) + print(f"Wrote {args.output}") + + +def main() -> None: + args = parse_args() + seed_everything(args.seed) + if args.side == "vision": + train_vision(args) + else: + train_text(args) + + +if __name__ == "__main__": + main() diff --git a/worldalign/synth_transfer_energy.py b/worldalign/synth_transfer_energy.py new file mode 100644 index 0000000..67a081f --- /dev/null +++ b/worldalign/synth_transfer_energy.py @@ -0,0 +1,276 @@ +"""Prediction-transfer alignment energy for the synthetic world. + +The third rung of the energy ladder: no hand-designed geometry and no +pair-trained EBM. Each modality's own predictor induces a substitution +kernel over scenes -- which other scenes the world model finds compatible +-- and the alignment energy is the conjugacy defect of the two kernels +under a candidate coupling. + +- Vision kernel: the masked-reconstruction tower inpaints a half-masked + render; pixel distance between the inpainting and other scenes' renders + gives K_V[i, i']. +- Text kernel: the causal LM scores scene j's caption sentences as a + continuation of scene i's caption prefix; referential consistency makes + same-content continuations likely, giving K_T. + +Both kernels are unimodal-predictor functionals. Hidden pairs enter only +through the gate that scores orderings of the energy. +""" + +from __future__ import annotations + +import argparse +import json +from pathlib import Path + +import torch +import torch.nn.functional as F +from tqdm import tqdm + +from .common import batch_indices, read_json, seed_everything, write_json +from .manifold_gate import standardize_relation +from .ricci_control import run_gates +from .synth_towers import TextTower, VisionTower, load_image, tokenize + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--data-dir", default="artifacts/synth_v0") + parser.add_argument( + "--vision-tower", default="artifacts/synth_v0/vision_tower_simmim.pt" + ) + parser.add_argument("--text-tower", default="artifacts/synth_v0/text_tower.pt") + parser.add_argument("--split", choices=["val", "test"], default="test") + parser.add_argument("--samples", type=int, default=512) + parser.add_argument("--mask-ratio", type=float, default=0.5) + parser.add_argument("--kernel-temperature", type=float, default=0.03) + parser.add_argument("--batch-size", type=int, default=256) + parser.add_argument("--random-perms", type=int, default=300) + parser.add_argument("--descent-restarts", type=int, default=5) + 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="cuda:3") + parser.add_argument("--seed", type=int, default=20260731) + parser.add_argument( + "--output", default="artifacts/synth_v0/transfer_energy_gate.json" + ) + parser.add_argument( + "--kernels-output", default="artifacts/synth_v0/transfer_kernels.pt" + ) + return parser.parse_args() + + +@torch.inference_mode() +def vision_kernel( + rows: list[int], manifest: dict, args: argparse.Namespace +) -> torch.Tensor: + state = torch.load(args.vision_tower, map_location="cpu", weights_only=False) + saved = state["args"] + model = VisionTower( + manifest["image_size"], saved["patch"], saved["dim"], saved["depth"], saved["heads"] + ).to(args.device) + model.load_state_dict(state["model"]) + model.eval() + image_dir = Path(manifest["image_dir"]) + tokens = (manifest["image_size"] // saved["patch"]) ** 2 + generator = torch.Generator(device=args.device).manual_seed(args.seed) + + references = [] + for indices in batch_indices(len(rows), args.batch_size): + pixels = torch.stack( + [load_image(image_dir / f"scene{rows[i]:06d}_v0.png") for i in indices] + ) + references.append(pixels) + references = torch.cat(references) # [N, 3, H, W] view-0 renders on cpu + + inpainted = [] + for indices in tqdm( + list(batch_indices(len(rows), args.batch_size)), desc="vision kernel" + ): + # Inpaint view 1 under a random half mask; compare to view-0 renders. + pixels = torch.stack( + [load_image(image_dir / f"scene{rows[i]:06d}_v1.png") for i in indices] + ).to(args.device) + mask = ( + torch.rand(len(pixels), tokens, generator=generator, device=args.device) + < args.mask_ratio + ) + encoded = model.encode(pixels, mask=mask) + patches = model.reconstruction(encoded[:, 1:]) # [B, T, p*p*3] + grid = manifest["image_size"] // saved["patch"] + patch = saved["patch"] + image = patches.reshape(len(pixels), grid, grid, patch, patch, 3) + image = image.permute(0, 5, 1, 3, 2, 4).reshape( + len(pixels), 3, manifest["image_size"], manifest["image_size"] + ) + blended = torch.where( + mask.reshape(len(pixels), 1, grid, 1, grid, 1) + .expand(-1, 3, -1, patch, -1, patch) + .reshape_as(image), + image, + pixels, + ) + inpainted.append(blended.float().cpu()) + inpainted = torch.cat(inpainted) + + # Layout-invariant content distance: Chamfer between foreground patch + # sets in pixel space. Renders of the same scene share patch CONTENT + # (colors, local shape fragments) while positions are resampled, so + # per-pixel comparison would measure layout instead. + patch = saved["patch"] + grid = manifest["image_size"] // patch + + def patch_sets(images: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + pieces = images.reshape(len(images), 3, grid, patch, grid, patch) + pieces = pieces.permute(0, 2, 4, 1, 3, 5).reshape(len(images), grid * grid, -1) + background = pieces.mean(-1) + foreground = background > background.min(1, keepdim=True).values + 0.02 + return pieces, foreground + + inpaint_pieces, inpaint_fg = patch_sets(inpainted) + reference_pieces, reference_fg = patch_sets(references) + distances = torch.zeros(len(rows), len(rows)) + device = args.device + reference_pieces = reference_pieces.to(device) + reference_fg = reference_fg.to(device) + for start in range(0, len(rows), 32): + stop = min(start + 32, len(rows)) + block = inpaint_pieces[start:stop].to(device) + block_fg = inpaint_fg[start:stop].to(device) + cross = torch.cdist(block, reference_pieces.reshape(-1, block.shape[-1])) + cross = cross.reshape(stop - start, block.shape[1], len(rows), -1) + masked_ab = cross.masked_fill( + ~reference_fg[None, None, :, :], float("inf") + ).amin(-1) + forward = ( + (masked_ab * block_fg[:, :, None]).sum(1) + / block_fg.sum(1).clamp_min(1)[:, None] + ) + masked_ba = cross.masked_fill( + ~block_fg[:, :, None, None].expand_as(cross), float("inf") + ).amin(1) + backward = ( + (masked_ba * reference_fg[None, :, :]).sum(-1) + / reference_fg.sum(-1).clamp_min(1)[None, :] + ) + distances[start:stop] = (0.5 * (forward + backward)).float().cpu() + distances[torch.isinf(distances) | torch.isnan(distances)] = distances[ + torch.isfinite(distances) + ].max() + return distances + + +@torch.inference_mode() +def text_kernel( + rows: list[int], captions: list[list[str]], args: argparse.Namespace +) -> torch.Tensor: + state = torch.load(args.text_tower, map_location="cpu", weights_only=False) + saved = state["args"] + vocab = state["vocab"] + model = TextTower( + len(vocab), saved["text_dim"], saved["depth"], saved.get("text_heads", 4), saved["context"] + ).to(args.device) + model.load_state_dict(state["model"]) + model.eval() + + prefixes = [tokenize(captions[row][0], vocab) for row in rows] + # Continuations are the RELATION sentences only: they refer back to + # groups the prefix must have introduced, so their likelihood grades + # shared content instead of rewarding verbatim repetition. + continuations = [] + for row in rows: + sentences = captions[row][1].split(". ") + relational = ". ".join(sentences[1:]) if len(sentences) > 1 else sentences[0] + continuations.append(tokenize(relational, vocab)[1:]) # drop bos + scores = torch.zeros(len(rows), len(rows)) + jobs = [(i, j) for i in range(len(rows)) for j in range(len(rows))] + for indices in tqdm( + list(batch_indices(len(jobs), args.batch_size)), desc="text kernel" + ): + batch = [jobs[k] for k in indices] + sequences = [ + (prefixes[i] + continuations[j])[: saved["context"]] for i, j in batch + ] + longest = max(len(s) for s in sequences) + tokens = torch.zeros(len(batch), longest, dtype=torch.long) + for row, sequence in enumerate(sequences): + tokens[row, : len(sequence)] = torch.tensor(sequence) + tokens = tokens.to(args.device) + hidden = model(tokens) + logits = model.logits(hidden[:, :-1]) + log_probs = F.log_softmax(logits, dim=-1) + for row, (i, j) in enumerate(batch): + begin = len(prefixes[i]) - 1 + end = min(len(prefixes[i]) + len(continuations[j]), longest) - 1 + if end <= begin: + scores[i, j] = 0.0 + continue + targets = tokens[row, begin + 1 : end + 1] + picked = log_probs[row, begin:end].gather(1, targets[:, None]) + scores[i, j] = float(picked.mean()) + return -scores # negative mean log-likelihood as a distance + + +def main() -> None: + args = parse_args() + seed_everything(args.seed) + manifest = read_json(Path(args.data_dir, "manifest.json")) + captions = read_json(Path(args.data_dir, "captions.json"))["captions"] + rows = manifest[args.split][: args.samples] + + visual_distance = vision_kernel(rows, manifest, args) + text_distance = text_kernel(rows, captions, args) + torch.save( + { + "rows": rows, + "visual_distance": visual_distance, + "text_distance": text_distance, + }, + args.kernels_output, + ) + + # Substitution kernels as standardized relation channels: negative + # distances, symmetrized, standardized off-diagonal. The gate machinery + # then scores orderings exactly as for any relation field. + visual_similarity = -visual_distance + text_similarity = -0.5 * (text_distance + text_distance.T) + visual_channels = standardize_relation( + 0.5 * (visual_similarity + visual_similarity.T).double() + )[0][None] + text_channels = standardize_relation(text_similarity.double())[0][None] + generator = torch.Generator().manual_seed(args.seed) + report = { + "protocol": ( + "Substitution kernels are unimodal-predictor functionals " + "(masked inpainting distance; caption continuation " + "likelihood). Hidden pairs score orderings only." + ), + "split": args.split, + "samples": len(rows), + **run_gates(text_channels, visual_channels, args, generator), + } + verdict = { + "true_z_mse": report["gate_a"]["random"]["mse"]["true_z"], + "improving_fraction": report["gate_b"]["improving_fraction"], + "descent_keeps": report["descent_from_true"]["final_accuracy"], + "true_mse": report["gate_a"]["true"]["mse"], + "best_random_descent": min( + (r["final_objective"] for r in report["descent_from_random"]), + default=None, + ), + } + verdict["counterfeit_found"] = bool( + verdict["best_random_descent"] is not None + and verdict["best_random_descent"] < verdict["true_mse"] + ) + report["verdict"] = verdict + print(json.dumps({"verdict": verdict})) + write_json(args.output, report) + print(f"Wrote {args.output}") + + +if __name__ == "__main__": + main() diff --git a/worldalign/synth_triangle_gate.py b/worldalign/synth_triangle_gate.py new file mode 100644 index 0000000..35bc2b2 --- /dev/null +++ b/worldalign/synth_triangle_gate.py @@ -0,0 +1,296 @@ +"""Third-order (triangle) relational invariants for the synthetic world. + +Counterfeits so far satisfy pairwise statistics: they are wrong sections +that look flat when probed along edges. The holonomy analogue on a +relation field is a triangle: for a node triple the product-like +statistic T[a,b,c] combines three edges at once, so preserving pairwise +marginals no longer suffices. Each node participates in O(N^2) triangles +rather than O(N) edges, which multiplies the constraint count without +touching the states. + +The triangle energy sums squared cross-modal differences of standardized +triangle tensors over a fixed random triple sample; the sample is drawn +once from node indices only. Swap deltas are evaluated exactly on the +triples touching the swapped nodes. Hidden pairs score orderings only. +""" + +from __future__ import annotations + +import argparse +import json +from pathlib import Path + +import torch +import torch.nn.functional as F +from tqdm import tqdm + +from .common import read_json, seed_everything, write_json +from .synth_cc_battery import ( + component_descriptors, + moment_field, + onehot_descriptors, + phrase_bow_sets, +) +from .synth_towers import load_image + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--data-dir", default="artifacts/synth_v0") + parser.add_argument("--split", choices=["val", "test"], default="test") + parser.add_argument("--samples", type=int, default=512) + parser.add_argument("--merge-distance", type=float, default=30.0) + parser.add_argument("--vision-views", type=int, default=4) + parser.add_argument("--triples", type=int, default=200000) + parser.add_argument("--pair-weight", type=float, default=1.0) + parser.add_argument("--triangle-weight", type=float, default=1.0) + parser.add_argument("--random-perms", type=int, default=200) + parser.add_argument("--transposition-samples", type=int, default=20000) + parser.add_argument("--descent-restarts", type=int, default=3) + parser.add_argument("--descent-max-steps", type=int, default=3000) + parser.add_argument("--descent-proposals", type=int, default=2048) + parser.add_argument("--replicas", type=int, default=8) + parser.add_argument("--tempering-rounds", type=int, default=20000) + parser.add_argument("--temp-high", type=float, default=3e-3) + parser.add_argument("--temp-low", type=float, default=1e-5) + parser.add_argument("--exchange-every", type=int, default=20) + parser.add_argument("--device", default="cuda:3") + parser.add_argument("--seed", type=int, default=20260731) + parser.add_argument( + "--output", default="artifacts/synth_v0/triangle_gate.json" + ) + return parser.parse_args() + + +def standardized(field: torch.Tensor) -> torch.Tensor: + mask = ~torch.eye(len(field), dtype=torch.bool, device=field.device) + values = field[mask] + out = (field - values.mean()) / values.std().clamp_min(1e-9) + return out.masked_fill(~mask, 0.0) + + +class TriangleEnergy: + """Pairwise-plus-triangle energy over a fixed triple sample.""" + + def __init__( + self, + text: torch.Tensor, + visual: torch.Tensor, + triples: torch.Tensor, + pair_weight: float, + triangle_weight: float, + ) -> None: + self.text = text + self.visual = visual + self.triples = triples + self.pair_weight = pair_weight + self.triangle_weight = triangle_weight + self.size = len(visual) + self.mask = ~torch.eye(self.size, dtype=torch.bool, device=visual.device) + self.pair_count = int(self.mask.sum()) + a, b, c = triples[:, 0], triples[:, 1], triples[:, 2] + self.visual_triangle = ( + visual[a, b] * visual[b, c] * visual[a, c] + ) + # Triples touching each node, for local delta evaluation. + self.touch = [[] for _ in range(self.size)] + for index, (x, y, z) in enumerate(triples.tolist()): + self.touch[x].append(index) + self.touch[y].append(index) + self.touch[z].append(index) + self.touch = [ + torch.tensor(items, device=visual.device, dtype=torch.long) + for items in self.touch + ] + + def pairwise(self, permutation: torch.Tensor) -> torch.Tensor: + fields = self.text[permutation][:, permutation] + return ((fields - self.visual) ** 2)[self.mask].sum() / self.pair_count + + def triangle(self, permutation: torch.Tensor) -> torch.Tensor: + fields = self.text[permutation][:, permutation] + a, b, c = self.triples[:, 0], self.triples[:, 1], self.triples[:, 2] + text_triangle = fields[a, b] * fields[b, c] * fields[a, c] + return ((text_triangle - self.visual_triangle) ** 2).mean() + + def total(self, permutation: torch.Tensor) -> float: + return float( + self.pair_weight * self.pairwise(permutation) + + self.triangle_weight * self.triangle(permutation) + ) + + def swap_delta(self, permutation: torch.Tensor, p: int, q: int) -> float: + """Exact delta for swapping positions p and q.""" + before_pair = self._local_pair(permutation, p, q) + indices = torch.cat([self.touch[p], self.touch[q]]).unique() + before_triangle = self._local_triangle(permutation, indices) + trial = permutation.clone() + trial[[p, q]] = trial[[q, p]] + after_pair = self._local_pair(trial, p, q) + after_triangle = self._local_triangle(trial, indices) + return float( + self.pair_weight * (after_pair - before_pair) / self.pair_count + + self.triangle_weight + * (after_triangle - before_triangle) + / len(self.triples) + ) + + def _local_pair( + self, permutation: torch.Tensor, p: int, q: int + ) -> torch.Tensor: + rows = permutation[[p, q]] + fields = self.text[rows][:, permutation] + visual_rows = self.visual[[p, q]] + difference = (fields - visual_rows) ** 2 + difference[0, p] = 0.0 + difference[1, q] = 0.0 + # Rows and columns are symmetric; count row contributions twice and + # subtract the doubly counted (p, q) pair once. + total = 2.0 * difference.sum() + return total - 2.0 * difference[0, q] + + def _local_triangle( + self, permutation: torch.Tensor, indices: torch.Tensor + ) -> torch.Tensor: + triples = self.triples[indices] + a, b, c = triples[:, 0], triples[:, 1], triples[:, 2] + pa, pb, pc = permutation[a], permutation[b], permutation[c] + text_triangle = ( + self.text[pa, pb] * self.text[pb, pc] * self.text[pa, pc] + ) + return ((text_triangle - self.visual_triangle[indices]) ** 2).sum() + + +def build_fields(args: argparse.Namespace, rows: list[int], manifest: dict) -> tuple: + image_dir = Path(manifest["image_dir"]) + per_view = [] + for view in range(args.vision_views): + raw = [] + for row in tqdm(rows, desc=f"cc v{view}"): + sprites, _ = component_descriptors( + load_image(image_dir / f"scene{row:06d}_v{view}.png"), + args.merge_distance, + ) + raw.append(sprites) + sets = [F.normalize(v, dim=-1) for v in onehot_descriptors(raw)] + per_view.append(moment_field(sets)) + visual_field = torch.stack(per_view).mean(0) + captions = read_json(Path(args.data_dir, "captions.json"))["captions"] + text_sets = phrase_bow_sets(rows, captions, manifest["vocabulary"]) + return visual_field, moment_field(text_sets) + + +def main() -> None: + args = parse_args() + seed_everything(args.seed) + manifest = read_json(Path(args.data_dir, "manifest.json")) + rows = manifest[args.split][: args.samples] + visual_field, text_field = build_fields(args, rows, manifest) + + device = torch.device(args.device) + size = len(rows) + generator = torch.Generator().manual_seed(args.seed) + hidden = torch.randperm(size, generator=generator) + truth = torch.argsort(hidden).to(device) + text = standardized(text_field[hidden][:, hidden].double().to(device)).float() + visual = standardized(visual_field.double().to(device)).float() + + triples = torch.randint(0, size, (args.triples, 3), generator=generator) + triples = triples[ + (triples[:, 0] != triples[:, 1]) + & (triples[:, 1] != triples[:, 2]) + & (triples[:, 0] != triples[:, 2]) + ].to(device) + energy = TriangleEnergy( + text, visual, triples, args.pair_weight, args.triangle_weight + ) + + true_total = energy.total(truth) + randoms = [] + for _ in range(args.random_perms): + permutation = torch.argsort(torch.rand(size, generator=generator)).to(device) + randoms.append(energy.total(permutation)) + randoms = torch.tensor(randoms) + gate_a = { + "true": true_total, + "random_mean": float(randoms.mean()), + "true_z": float((randoms.mean() - true_total) / randoms.std().clamp_min(1e-12)), + } + + improving = 0 + checked = 0 + for _ in range(min(args.transposition_samples, 4000)): + p = int(torch.randint(0, size, (1,), generator=generator)) + q = int(torch.randint(0, size, (1,), generator=generator)) + if p == q: + continue + checked += 1 + improving += energy.swap_delta(truth, p, q) < 0 + gate_b = { + "sampled": checked, + "improving_fraction": improving / max(checked, 1), + } + + def descent(start: torch.Tensor) -> dict: + current = start.clone() + for _ in range(args.descent_max_steps): + best_delta, best_pair = 0.0, None + for _ in range(64): + p = int(torch.randint(0, size, (1,), generator=generator)) + q = int(torch.randint(0, size, (1,), generator=generator)) + if p == q: + continue + delta = energy.swap_delta(current, p, q) + if delta < best_delta: + best_delta, best_pair = delta, (p, q) + if best_pair is None: + break + p, q = best_pair + current[[p, q]] = current[[q, p]] + return { + "accuracy": float((current.cpu() == truth.cpu()).float().mean()), + "energy": energy.total(current), + } + + from_true = descent(truth) + restarts = [ + descent(torch.argsort(torch.rand(size, generator=generator)).to(device)) + for _ in range(args.descent_restarts) + ] + best_random = min(r["energy"] for r in restarts) + + report = { + "protocol": ( + "Pairwise-plus-triangle energy on moment set-kernel fields; " + "triples sampled from node indices only. Hidden pairs score " + "orderings." + ), + "samples": size, + "triples": len(triples), + "weights": { + "pair": args.pair_weight, + "triangle": args.triangle_weight, + }, + "gate_a": gate_a, + "gate_b": gate_b, + "descent_from_true": from_true, + "descent_from_random": restarts, + "verdict": { + "true_energy": true_total, + "best_random_descent": best_random, + "counterfeit_found": bool( + best_random < true_total + and min(r["accuracy"] for r in restarts) < 0.5 + ), + "descent_keeps": from_true["accuracy"], + "true_z": gate_a["true_z"], + "improving_fraction": gate_b["improving_fraction"], + }, + } + print(json.dumps({"verdict": report["verdict"]})) + write_json(args.output, report) + print(f"Wrote {args.output}") + + +if __name__ == "__main__": + main() diff --git a/worldalign/synth_triangle_recovery.py b/worldalign/synth_triangle_recovery.py new file mode 100644 index 0000000..a9bca61 --- /dev/null +++ b/worldalign/synth_triangle_recovery.py @@ -0,0 +1,173 @@ +"""Blind recovery under the triangle energy. + +The triangle gate holds the true assignment at N=512 with no counterfeit +below it. The search question is separate and now worth asking again: +tempering with third-order swap deltas, from random starts, with the +hidden order behind a shuffle. Recovery accuracy is the end-to-end world +matching measurement. +""" + +from __future__ import annotations + +import argparse +import json +from pathlib import Path + +import torch + +from .common import read_json, seed_everything, write_json +from .synth_triangle_gate import TriangleEnergy, build_fields, standardized + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--data-dir", default="artifacts/synth_v0") + parser.add_argument("--split", choices=["val", "test"], default="test") + parser.add_argument("--samples", type=int, default=512) + parser.add_argument("--merge-distance", type=float, default=30.0) + parser.add_argument("--vision-views", type=int, default=4) + parser.add_argument("--triples", type=int, default=400000) + parser.add_argument("--pair-weight", type=float, default=1.0) + parser.add_argument("--triangle-weight", type=float, default=1.0) + parser.add_argument("--replicas", type=int, default=6) + parser.add_argument("--rounds", type=int, default=4000) + parser.add_argument("--proposals", type=int, default=48) + parser.add_argument("--temp-high", type=float, default=2e-2) + parser.add_argument("--temp-low", type=float, default=1e-4) + parser.add_argument("--exchange-every", type=int, default=20) + parser.add_argument("--greedy-rounds", type=int, default=3000) + parser.add_argument("--device", default="cuda:3") + parser.add_argument("--seed", type=int, default=20260731) + parser.add_argument( + "--output", default="artifacts/synth_v0/triangle_recovery.json" + ) + return parser.parse_args() + + +def main() -> None: + args = parse_args() + seed_everything(args.seed) + manifest = read_json(Path(args.data_dir, "manifest.json")) + rows = manifest[args.split][: args.samples] + visual_field, text_field = build_fields(args, rows, manifest) + + device = torch.device(args.device) + size = len(rows) + generator = torch.Generator().manual_seed(args.seed) + hidden = torch.randperm(size, generator=generator) + truth = torch.argsort(hidden).to(device) + text = standardized(text_field[hidden][:, hidden].double().to(device)).float() + visual = standardized(visual_field.double().to(device)).float() + + triples = torch.randint(0, size, (args.triples, 3), generator=generator) + triples = triples[ + (triples[:, 0] != triples[:, 1]) + & (triples[:, 1] != triples[:, 2]) + & (triples[:, 0] != triples[:, 2]) + ].to(device) + energy = TriangleEnergy( + text, visual, triples, args.pair_weight, args.triangle_weight + ) + true_energy = energy.total(truth) + + temperatures = torch.logspace( + torch.log10(torch.tensor(args.temp_low)), + torch.log10(torch.tensor(args.temp_high)), + args.replicas, + ) + states = [ + torch.argsort(torch.rand(size, generator=generator)).to(device) + for _ in range(args.replicas) + ] + energies = [energy.total(state) for state in states] + history = [] + for round_index in range(args.rounds): + for replica in range(args.replicas): + temperature = float(temperatures[replica]) + for _ in range(args.proposals): + p = int(torch.randint(0, size, (1,), generator=generator)) + q = int(torch.randint(0, size, (1,), generator=generator)) + if p == q: + continue + delta = energy.swap_delta(states[replica], p, q) + threshold = -temperature * float( + torch.rand(1, generator=generator).clamp_min(1e-12).log() + ) + if delta < threshold: + states[replica][[p, q]] = states[replica][[q, p]] + energies[replica] += delta + if round_index % args.exchange_every == 0: + for replica in range(args.replicas - 1): + gap = (energies[replica] - energies[replica + 1]) * ( + 1.0 / float(temperatures[replica]) + - 1.0 / float(temperatures[replica + 1]) + ) + if gap > 0 or float(torch.rand(1, generator=generator)) < min( + 1.0, float(torch.tensor(gap).exp()) + ): + states[replica], states[replica + 1] = ( + states[replica + 1], + states[replica], + ) + energies[replica], energies[replica + 1] = ( + energies[replica + 1], + energies[replica], + ) + if round_index % 200 == 0: + cold = min(range(args.replicas), key=lambda r: energies[r]) + record = { + "round": round_index, + "cold_energy": energies[cold], + "cold_accuracy": float( + (states[cold].cpu() == truth.cpu()).float().mean() + ), + "energy_over_true": energies[cold] / true_energy - 1.0, + } + history.append(record) + print(json.dumps(record)) + + # Greedy polish of the coldest replica. + cold = min(range(args.replicas), key=lambda r: energies[r]) + current = states[cold].clone() + for _ in range(args.greedy_rounds): + best_delta, best_pair = 0.0, None + for _ in range(96): + p = int(torch.randint(0, size, (1,), generator=generator)) + q = int(torch.randint(0, size, (1,), generator=generator)) + if p == q: + continue + delta = energy.swap_delta(current, p, q) + if delta < best_delta: + best_delta, best_pair = delta, (p, q) + if best_pair is None: + break + p, q = best_pair + current[[p, q]] = current[[q, p]] + + final = { + "accuracy": float((current.cpu() == truth.cpu()).float().mean()), + "energy": energy.total(current), + } + final["energy_over_true"] = final["energy"] / true_energy - 1.0 + report = { + "protocol": ( + "Tempering on the pairwise-plus-triangle energy from random " + "starts; text order hidden behind a shuffle; hidden truth " + "scores the outcome only." + ), + "samples": size, + "true_energy": true_energy, + "chance_accuracy": 1.0 / size, + "history": history, + "final_polished": final, + "replica_accuracies": [ + float((state.cpu() == truth.cpu()).float().mean()) for state in states + ], + } + print(json.dumps({"final": final})) + write_json(args.output, report) + print(f"Wrote {args.output}") + + +if __name__ == "__main__": + main() diff --git a/worldalign/synth_world.py b/worldalign/synth_world.py new file mode 100644 index 0000000..8875fde --- /dev/null +++ b/worldalign/synth_world.py @@ -0,0 +1,506 @@ +"""Synthetic closed world v0: procedural scenes with exact closure. + +A scene's discrete state is a set of object groups (count, size, color, +shape) plus a set of spatial relation constraints between groups. Captions +mention exactly the discrete state; renders realize it with continuous +nuisance (positions, jitter) resampled per view. Shared content and +modality-private variation are therefore separated by construction, and +every dial -- vocabulary, ontology size, groups per scene, scene count, +orbit multiplicity, intervention density -- is a generator argument. + +Outputs mirror the Flickr pipeline layout: a manifest with disjoint +vision-only and text-only scene rows plus held-out val/test, an image +directory with `sceneNNNNNN_vK.png` views, and per-scene caption lists. +Intervention variants with edit metadata are stored for the response +battery; nothing downstream reads them yet. +""" + +from __future__ import annotations + +import argparse +import json +import math +import random +from pathlib import Path + +from PIL import Image, ImageDraw + +from .common import write_json + +COLORS = { + "red": (205, 49, 49), + "orange": (224, 133, 44), + "yellow": (229, 213, 74), + "green": (64, 168, 75), + "blue": (59, 104, 214), + "purple": (139, 72, 190), + "pink": (228, 136, 179), + "white": (238, 238, 238), + "gray": (140, 140, 140), + "brown": (125, 84, 48), +} +SHAPES = ("circle", "square", "triangle", "star", "diamond", "cross") +SIZES = {"small": (7, 10), "medium": (13, 17), "large": (21, 26)} +COUNT_WORDS = {1: "one", 2: "two", 3: "three", 4: "four"} +RELATIONS = ("left of", "right of", "above", "below") +BACKGROUND = (24, 24, 28) + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--output-dir", default="artifacts/synth_v0") + parser.add_argument("--image-size", type=int, default=128) + parser.add_argument("--vision-only", type=int, default=12000) + parser.add_argument("--text-only", type=int, default=12000) + parser.add_argument("--val", type=int, default=1000) + parser.add_argument("--test", type=int, default=1000) + parser.add_argument("--min-groups", type=int, default=2) + parser.add_argument("--max-groups", type=int, default=4) + parser.add_argument("--visual-views", type=int, default=4) + parser.add_argument("--captions", type=int, default=4) + parser.add_argument("--interventions", type=int, default=2, + help="Edited variants per eval scene.") + parser.add_argument("--seed", type=int, default=20260730) + parser.add_argument( + "--skew", + type=float, + default=0.0, + help="Zipf exponent for colour and shape sampling. Non-uniform " + "marginals make the cross-modal value correspondence recoverable " + "from frequency alone, with no declared dictionary.", + ) + parser.add_argument( + "--correlate", + type=float, + default=0.0, + help="Colour-shape coupling strength. Real worlds pair values " + "(bananas are yellow), which makes co-occurrence structure " + "informative where bare marginals are not.", + ) + parser.add_argument( + "--report-bias", + type=float, + default=0.0, + help="Reporting bias strength. Captions mention rarer colours " + "preferentially, so text frequency stops tracking pixel " + "frequency -- the situation in real corpora, where nobody " + "describes the grey road.", + ) + parser.add_argument( + "--texture", + action="store_true", + help="Dot-textured object fills: within-object structure that " + "gives binding objectives a purchase, without changing the " + "discrete state or the captions.", + ) + return parser.parse_args() + + +_PROFILE_CACHE: dict[tuple[str, float], list[float]] = {} + + +def colour_shape_profile(colour: str, concentration: float) -> list[float]: + """Fixed per-colour shape distribution, deterministic in the colour.""" + key = (colour, concentration) + if key not in _PROFILE_CACHE: + local = random.Random(hash(key) & 0xFFFFFFFF) + raw = [local.gammavariate(1.0 / concentration, 1.0) for _ in SHAPES] + total = sum(raw) or 1.0 + _PROFILE_CACHE[key] = [value / total for value in raw] + return _PROFILE_CACHE[key] + + +def zipf_weights(count: int, exponent: float) -> list[float]: + raw = [1.0 / (rank + 1) ** exponent for rank in range(count)] + total = sum(raw) + return [value / total for value in raw] + + +def sample_group( + rng: random.Random, skew: float = 0.0, correlate: float = 0.0 +) -> dict: + if skew > 0.0: + colors = rng.choices( + list(COLORS), weights=zipf_weights(len(COLORS), skew) + )[0] + else: + colors = rng.choice(list(COLORS)) + if correlate > 0.0: + # Each colour carries its own shape distribution, drawn once from a + # Dirichlet with concentration 1/correlate. Distinct profiles make + # the joint table identifying, as in a real world where objects of + # a kind take characteristic forms. + weights = colour_shape_profile(colors, correlate) + if skew > 0.0: + base = zipf_weights(len(SHAPES), skew) + weights = [w * b for w, b in zip(weights, base)] + shapes = rng.choices(SHAPES, weights=weights)[0] + elif skew > 0.0: + shapes = rng.choices( + SHAPES, weights=zipf_weights(len(SHAPES), skew) + )[0] + else: + shapes = rng.choice(SHAPES) + return { + "count": rng.randint(1, 4), + "size": rng.choice(list(SIZES)), + "color": colors, + "shape": shapes, + } + + +def sample_scene(rng: random.Random, args: argparse.Namespace) -> dict: + skew = getattr(args, "skew", 0.0) + correlate = getattr(args, "correlate", 0.0) + groups = [ + sample_group(rng, skew, correlate) + for _ in range(rng.randint(args.min_groups, args.max_groups)) + ] + # Reject referential ambiguity: color-shape pairs are unique per scene, + # so every relational mention has exactly one referent. + signatures = [(group["color"], group["shape"]) for group in groups] + while len(set(signatures)) < len(signatures): + groups = [sample_group(rng, skew, correlate) for _ in range(len(groups))] + signatures = [(group["color"], group["shape"]) for group in groups] + relation_count = rng.randint(1, min(3, len(groups) * (len(groups) - 1) // 2)) + pairs = [(a, b) for a in range(len(groups)) for b in range(len(groups)) if a < b] + rng.shuffle(pairs) + relations = [ + {"a": a, "b": b, "relation": rng.choice(RELATIONS)} + for a, b in pairs[:relation_count] + ] + return {"groups": groups, "relations": relations} + + +def relation_holds(relation: str, pa: tuple[float, float], pb: tuple[float, float], margin: float) -> bool: + if relation == "left of": + return pa[0] < pb[0] - margin + if relation == "right of": + return pa[0] > pb[0] + margin + if relation == "above": + return pa[1] < pb[1] - margin + return pa[1] > pb[1] + margin + + +def group_geometry(group: dict, rng: random.Random, size: int) -> dict: + low, high = SIZES[group["size"]] + shrink = (1.0, 0.95, 0.85, 0.75)[group["count"] - 1] + radius = rng.uniform(low, high) * size / 128.0 * shrink + count = group["count"] + if count == 1: + offsets = [(0.0, 0.0)] + extent = radius + else: + # Ring placement guarantees exact visible multiplicity: adjacent + # spacing 2 s sin(pi/k) stays above 2.15 r by construction. + ring = 1.10 * 1.075 * radius / math.sin(math.pi / count) + phase = rng.uniform(0.0, 2.0 * math.pi) + offsets = [ + ( + ring * math.cos(phase + 2.0 * math.pi * k / count), + ring * math.sin(phase + 2.0 * math.pi * k / count), + ) + for k in range(count) + ] + extent = ring + radius + return {"radius": radius, "offsets": offsets, "extent": extent} + + +def worst_case_extent(group: dict, size: int) -> float: + high = SIZES[group["size"]][1] * size / 128.0 + shrink = (1.0, 0.95, 0.85, 0.75)[group["count"] - 1] + radius = high * shrink + if group["count"] == 1: + return radius + return 1.10 * 1.075 * radius / math.sin(math.pi / group["count"]) + radius + + +def place_groups( + scene: dict, + rng: random.Random, + size: int, + extents: list[float] | None = None, +) -> list[tuple[float, float]] | None: + margin = size * 0.08 + if extents is None: + extents = [worst_case_extent(group, size) for group in scene["groups"]] + for _ in range(300): + centers = [] + feasible = True + for extent in extents: + low, high = extent + 2.0, size - extent - 2.0 + if low >= high: + feasible = False + break + centers.append((rng.uniform(low, high), rng.uniform(low, high))) + if not feasible: + return None + if any( + math.dist(centers[a], centers[b]) < extents[a] + extents[b] + 4.0 + for a in range(len(centers)) + for b in range(a + 1, len(centers)) + ): + continue + if all( + relation_holds(r["relation"], centers[r["a"]], centers[r["b"]], margin) + for r in scene["relations"] + ): + return centers + return None + + +def draw_shape(draw: ImageDraw.ImageDraw, shape: str, x: float, y: float, radius: float, fill: tuple) -> None: + if shape == "circle": + draw.ellipse([x - radius, y - radius, x + radius, y + radius], fill=fill) + elif shape == "square": + draw.rectangle([x - radius, y - radius, x + radius, y + radius], fill=fill) + elif shape == "triangle": + draw.polygon( + [(x, y - radius), (x - radius, y + radius), (x + radius, y + radius)], + fill=fill, + ) + elif shape == "diamond": + draw.polygon( + [(x, y - radius), (x + radius, y), (x, y + radius), (x - radius, y)], + fill=fill, + ) + elif shape == "cross": + arm = radius * 0.42 + draw.rectangle([x - arm, y - radius, x + arm, y + radius], fill=fill) + draw.rectangle([x - radius, y - arm, x + radius, y + arm], fill=fill) + else: # star + points = [] + for k in range(10): + r = radius if k % 2 == 0 else radius * 0.45 + angle = -math.pi / 2 + k * math.pi / 5 + points.append((x + r * math.cos(angle), y + r * math.sin(angle))) + draw.polygon(points, fill=fill) + + +def render_scene(scene: dict, rng: random.Random, size: int, texture: bool = False) -> Image.Image | None: + geometries = [group_geometry(group, rng, size) for group in scene["groups"]] + centers = place_groups( + scene, rng, size, extents=[g["extent"] for g in geometries] + ) + if centers is None: + return None + shade = rng.randint(-6, 6) + image = Image.new("RGB", (size, size), tuple(c + shade for c in BACKGROUND)) + draw = ImageDraw.Draw(image) + for group, geometry, center in zip(scene["groups"], geometries, centers): + base = tuple( + min(255, max(0, channel + rng.randint(-10, 10))) + for channel in COLORS[group["color"]] + ) + for off in geometry["offsets"]: + draw_shape( + draw, + group["shape"], + center[0] + off[0], + center[1] + off[1], + geometry["radius"], + base, + ) + return image + + +def group_phrase(group: dict, rng: random.Random) -> str: + size_word = "" if group["size"] == "medium" else group["size"] + " " + plural = "es" if group["shape"] == "cross" else "s" + noun = group["shape"] + (plural if group["count"] > 1 else "") + count_word = COUNT_WORDS[group["count"]] if group["count"] > 1 else ( + "a" if size_word == "" or size_word[0] not in "aeiou" else "an" + ) + return f"{count_word} {size_word}{group['color']} {noun}" + + +def caption_scene( + scene: dict, rng: random.Random, report_bias: float = 0.0 +) -> str: + order = list(range(len(scene["groups"]))) + rng.shuffle(order) + if report_bias > 0.0 and len(order) > 1: + # Mention probability falls with the colour's population frequency, + # so the text marginal inverts the pixel marginal. + ranks = {name: index for index, name in enumerate(COLORS)} + keep = [] + for index in order: + rank = ranks[scene["groups"][index]["color"]] + salience = ((rank + 1) / len(COLORS)) ** report_bias + if rng.random() < 0.25 + 0.75 * salience: + keep.append(index) + order = keep or order[:1] + phrases = [group_phrase(scene["groups"][index], rng) for index in order] + if len(phrases) > 1: + listed = ", ".join(phrases[:-1]) + " and " + phrases[-1] + else: + listed = phrases[0] + opener = rng.choice(["there are", "the picture shows", "you can see"]) + sentences = [f"{opener} {listed}."] + relations = list(scene["relations"]) + rng.shuffle(relations) + mentioned = set(order) + relations = [ + relation + for relation in relations + if relation["a"] in mentioned and relation["b"] in mentioned + ] + for relation in relations: + subject = group_phrase(scene["groups"][relation["a"]], rng) + target = group_phrase(scene["groups"][relation["b"]], rng) + sentences.append(f"the {subject.split(' ', 1)[1]} {'are' if scene['groups'][relation['a']]['count'] > 1 else 'is'} {relation['relation']} the {target.split(' ', 1)[1]}.") + return " ".join(sentences) + + +def intervene(scene: dict, rng: random.Random) -> tuple[dict, dict]: + edited = json.loads(json.dumps(scene)) + kinds = ["recolor", "count", "remove", "relation"] + if len(edited["groups"]) <= 2: + kinds.remove("remove") + if not edited["relations"]: + kinds.remove("relation") + kind = rng.choice(kinds) + if kind == "recolor": + index = rng.randrange(len(edited["groups"])) + old = edited["groups"][index]["color"] + shape = edited["groups"][index]["shape"] + taken = { + g["color"] + for i, g in enumerate(edited["groups"]) + if i != index and g["shape"] == shape + } + edited["groups"][index]["color"] = rng.choice( + [c for c in COLORS if c != old and c not in taken] + ) + detail = {"kind": kind, "group": index, "from": old, "to": edited["groups"][index]["color"]} + elif kind == "count": + index = rng.randrange(len(edited["groups"])) + old = edited["groups"][index]["count"] + edited["groups"][index]["count"] = old % 4 + 1 + detail = {"kind": kind, "group": index, "from": old, "to": edited["groups"][index]["count"]} + elif kind == "remove": + index = rng.randrange(len(edited["groups"])) + edited["groups"].pop(index) + edited["relations"] = [ + r for r in edited["relations"] if r["a"] != index and r["b"] != index + ] + for relation in edited["relations"]: + relation["a"] -= relation["a"] > index + relation["b"] -= relation["b"] > index + detail = {"kind": kind, "group": index} + else: + index = rng.randrange(len(edited["relations"])) + old = edited["relations"][index]["relation"] + flip = {"left of": "right of", "right of": "left of", "above": "below", "below": "above"} + edited["relations"][index]["relation"] = flip[old] + detail = {"kind": kind, "relation_index": index, "from": old, "to": flip[old]} + return edited, detail + + +def main() -> None: + args = parse_args() + rng = random.Random(args.seed) + output = Path(args.output_dir) + (output / "images").mkdir(parents=True, exist_ok=True) + + total = args.vision_only + args.text_only + args.val + args.test + scenes, captions = [], [] + while len(scenes) < total: + scene = sample_scene(rng, args) + if place_groups(scene, rng, args.image_size) is None: + continue + scenes.append(scene) + captions.append( + [ + caption_scene(scene, rng, args.report_bias) + for _ in range(args.captions) + ] + ) + + rows = list(range(total)) + vision_only = rows[: args.vision_only] + text_only = rows[args.vision_only : args.vision_only + args.text_only] + val = rows[args.vision_only + args.text_only : args.vision_only + args.text_only + args.val] + test = rows[-args.test :] + + needs_render = set(vision_only) | set(val) | set(test) + rendered = 0 + for row in sorted(needs_render): + for view in range(args.visual_views): + image = None + while image is None: + image = render_scene( + scenes[row], rng, args.image_size, texture=args.texture + ) + image.save(output / "images" / f"scene{row:06d}_v{view}.png") + rendered += 1 + + interventions = [] + for row in val + test: + for _ in range(args.interventions): + edited, detail = intervene(scenes[row], rng) + if place_groups(edited, rng, args.image_size) is None: + continue + index = len(interventions) + image = None + while image is None: + image = render_scene( + edited, rng, args.image_size, texture=args.texture + ) + image.save(output / "images" / f"edit{index:06d}.png") + interventions.append( + { + "index": index, + "row": row, + "detail": detail, + "caption": caption_scene(edited, rng, args.report_bias), + } + ) + + vocabulary = sorted( + { + token + for caption_list in captions + for caption in caption_list + for token in caption.replace(",", " ").replace(".", " ").split() + } + ) + manifest = { + "dataset": "synth_v0", + "image_dir": str(output / "images"), + "image_size": args.image_size, + "seed": args.seed, + "visual_views": args.visual_views, + "vision_only_train": vision_only, + "text_only_train": text_only, + "paired_train": [], + "val": val, + "test": test, + "all_rows": total, + "vocabulary_size": len(vocabulary), + "vocabulary": vocabulary, + "protocol": ( + "vision_only_train and text_only_train are disjoint scene rows; " + "val/test pairs are held out for evaluation only. Captions " + "mention exactly the discrete scene state." + ), + } + write_json(output / "manifest.json", manifest) + write_json(output / "scenes.private.json", {"scenes": scenes}) + write_json(output / "captions.json", {"captions": captions}) + write_json(output / "interventions.private.json", {"interventions": interventions}) + print( + json.dumps( + { + "scenes": total, + "rendered_images": rendered, + "interventions": len(interventions), + "vocabulary_size": len(vocabulary), + "example_caption": captions[test[0]][0], + } + ) + ) + + +if __name__ == "__main__": + main() diff --git a/worldalign/tier0_dictionary.py b/worldalign/tier0_dictionary.py new file mode 100644 index 0000000..96fcbac --- /dev/null +++ b/worldalign/tier0_dictionary.py @@ -0,0 +1,324 @@ +"""Tier 0: derive the cross-modal value correspondence, never declare it. + +A declared lexicon ("red" means hue 0) is a hand-supplied cross-modal +prior. Tier 0 forbids it. Every factor value correspondence is instead +recovered from unimodal statistics of the two disjoint training splits: + +- ordered factors (count, size) match by their intrinsic order, with the + small residual ambiguity enumerated and settled by the alignment + criterion rather than by assertion; +- unordered factors (colour, shape) match by marginal frequency rank, + which is a unimodal observable on both sides. This works exactly when + the world's factor marginals are non-uniform -- true of real corpora + and of the skewed synthetic world, false of the uniform one, where the + correspondence is information-theoretically unrecoverable. + +The recovered dictionary is then applied to build comparable object +descriptors. Hidden pairs are used only to report how many entries the +derivation got right; nothing here reads them. +""" + +from __future__ import annotations + +import argparse +import json +import re +from collections import Counter +from pathlib import Path + +import numpy as np +import torch +import torch.nn.functional as F +from tqdm import tqdm + +from .common import read_json, seed_everything, write_json +from .synth_cc_battery import component_descriptors +from .synth_set_battery import parse_group_phrases +from .synth_towers import load_image + +SINGULAR_ARTICLES = ("a", "an") +STOPWORDS = { + "there", "are", "the", "picture", "shows", "you", "can", "see", + "is", "of", "left", "right", "above", "below", "and", +} + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--data-dir", default="artifacts/synth_v1") + parser.add_argument("--text-scenes", type=int, default=6000) + parser.add_argument("--vision-scenes", type=int, default=3000) + parser.add_argument("--merge-distance", type=float, default=30.0) + parser.add_argument("--colour-classes", type=int, default=10) + parser.add_argument("--shape-classes", type=int, default=6) + parser.add_argument("--seed", type=int, default=20260731) + parser.add_argument( + "--output", default="artifacts/synth_v1/tier0_dictionary.json" + ) + return parser.parse_args() + + +def text_token_statistics( + captions: list[list[str]], rows: list[int] +) -> dict: + """Group-slot token frequencies and singular/plural association. + + Everything here reads captions only. Number words are identified by + their association with singular noun forms and their mutual exclusion + inside a group phrase; the remaining modifier vocabulary splits into + frequency-ranked classes. + """ + phrase_total = 0 + slot_counts: Counter = Counter() + with_singular: Counter = Counter() + total_singular = 0 + cooccurrence: Counter = Counter() + for row in rows: + for phrase in parse_group_phrases(captions[row][0]): + tokens = [ + token + for token in re.findall(r"[a-z]+|\d+", phrase.lower()) + if token not in STOPWORDS + ] + if not tokens: + continue + phrase_total += 1 + head = tokens[-1] + singular = not head.endswith("s") + total_singular += singular + modifiers = tokens[:-1] + for token in modifiers: + slot_counts[token] += 1 + with_singular[token] += singular + for first in modifiers: + for second in modifiers: + if first != second: + cooccurrence[(first, second)] += 1 + return { + "slot_counts": slot_counts, + "with_singular": with_singular, + "cooccurrence": cooccurrence, + "total_singular": total_singular, + "phrase_total": phrase_total, + } + + +def partition_text_vocabulary(stats: dict) -> dict: + """Discover modifier families by mutual exclusivity, then name them. + + Tokens of one factor never co-occur inside a group phrase, so families + are maximal mutually exclusive sets: greedily place each token in the + first family none of whose members it ever co-occurs with. Families are + then identified by two unimodal signals -- coverage (how many phrases + carry a member) and morphological association (whether the choice + predicts the head noun's plural suffix). No token list is declared. + """ + counts = stats["slot_counts"] + cooccurrence = stats["cooccurrence"] + phrases = stats["phrase_total"] + ordered = sorted(counts, key=lambda token: -counts[token]) + families: list[list[str]] = [] + for token in ordered: + for family in families: + if all( + cooccurrence[(token, member)] == 0 + and cooccurrence[(member, token)] == 0 + for member in family + ): + family.append(token) + break + else: + families.append([token]) + described = [] + for family in families: + coverage = sum(counts[token] for token in family) / max(phrases, 1) + rates = [ + stats["with_singular"][token] / max(counts[token], 1) + for token in family + ] + described.append( + { + "tokens": sorted(family, key=lambda token: -counts[token]), + "coverage": coverage, + "morphology_spread": float(max(rates) - min(rates)), + } + ) + described.sort(key=lambda item: -item["coverage"]) + # The count family is the near-complete family whose choice predicts the + # plural suffix; the other near-complete family is the dominant + # unordered attribute; partial families are optional modifiers. + complete = [item for item in described if item["coverage"] > 0.8] + partial = [item for item in described if item["coverage"] <= 0.8] + complete.sort(key=lambda item: -item["morphology_spread"]) + count_family = complete[0]["tokens"] if complete else [] + attribute_families = [item["tokens"] for item in complete[1:]] + return { + "count_words": count_family, + "colour_words": attribute_families[0] if attribute_families else [], + "other_attribute_words": attribute_families[1:], + "size_words": [item["tokens"] for item in partial], + "families_detail": described, + } + + +def component_raw(image: torch.Tensor, merge_distance: float) -> dict: + """Connected-component groups with raw appearance, no colour rules. + + Returns mean RGB, area fraction, and member count per group. Nothing + here quantises colour, so no declared hue boundary enters Tier 0. + """ + from scipy import ndimage + + array = image.permute(1, 2, 0).numpy() + background = np.median(array.reshape(-1, 3), axis=0) + foreground = np.abs(array - background).sum(-1) > 0.12 + labels, count = ndimage.label(foreground) + if count == 0: + return {"rgb": np.zeros((0, 3)), "area": np.zeros(0), "members": np.zeros(0)} + centers = np.array( + ndimage.center_of_mass(foreground, labels, range(1, count + 1)) + ) + parent = list(range(count)) + + def find(a: int) -> int: + while parent[a] != a: + parent[a] = parent[parent[a]] + a = parent[a] + return a + + for a in range(count): + for b in range(a + 1, count): + if np.linalg.norm(centers[a] - centers[b]) < merge_distance: + parent[find(a)] = find(b) + groups: dict[int, list[int]] = {} + for a in range(count): + groups.setdefault(find(a), []).append(a) + rgb, area, members = [], [], [] + total = foreground.size + for group in groups.values(): + mask = np.isin(labels, [m + 1 for m in group]) + rgb.append(array[mask].mean(0)) + area.append(float(mask.sum()) / total) + members.append(len(group)) + return { + "rgb": np.stack(rgb), + "area": np.array(area), + "members": np.array(members, dtype=np.int64), + } + + +def vision_value_statistics( + rows: list[int], manifest: dict, args: argparse.Namespace +) -> dict: + """Colour-class and size-class frequencies from pixels alone. + + Object colours are clustered in hue-saturation-value space with the + requested number of classes; class identity is arbitrary, only the + frequency ranking is used downstream. + """ + from sklearn.cluster import KMeans + + image_dir = Path(manifest["image_dir"]) + raw = [ + component_raw( + load_image(image_dir / f"scene{row:06d}_v0.png"), args.merge_distance + ) + for row in tqdm(rows, desc="vision values") + ] + rgb = np.concatenate([item["rgb"] for item in raw]) + # Cluster raw appearance: class boundaries come from the data, not from + # a declared hue table. + clusters = KMeans( + n_clusters=args.colour_classes, n_init=10, random_state=args.seed + ).fit(rgb) + labels = clusters.labels_ + return { + "colour_frequency": Counter(labels.tolist()), + "colour_labels": labels, + "cluster_centres": clusters.cluster_centers_, + "raw": raw, + "clusters": clusters, + } + + +def frequency_rank_map( + text_words: list[str], vision_frequency: Counter, classes: int +) -> dict: + """Match unordered values by descending marginal frequency.""" + vision_ranked = [ + label for label, _ in vision_frequency.most_common(classes) + ] + return { + word: vision_ranked[index] + for index, word in enumerate(text_words[:classes]) + } + + +def main() -> None: + args = parse_args() + seed_everything(args.seed) + manifest = read_json(Path(args.data_dir, "manifest.json")) + captions = read_json(Path(args.data_dir, "captions.json"))["captions"] + + text_rows = manifest["text_only_train"][: args.text_scenes] + stats = text_token_statistics(captions, text_rows) + families = partition_text_vocabulary(stats) + + vision_rows = manifest["vision_only_train"][: args.vision_scenes] + vision = vision_value_statistics(vision_rows, manifest, args) + + dictionary = frequency_rank_map( + families["colour_words"], vision["colour_frequency"], args.colour_classes + ) + + report = { + "protocol": ( + "Factor families and their value correspondence are derived " + "from unimodal statistics of disjoint splits: noun-number " + "association separates count words, coverage separates colour " + "from size words, and marginal frequency rank pairs colour " + "values across modalities. No declared lexicon." + ), + "text_families": { + key: families[key] + for key in ("count_words", "colour_words", "size_words") + }, + "families_detail": families["families_detail"], + "vision_colour_frequency": [ + [int(label), int(count)] + for label, count in vision["colour_frequency"].most_common() + ], + "derived_colour_map": { + word: int(label) for word, label in dictionary.items() + }, + } + + # Evaluation only: how many derived entries are semantically right? + scenes = read_json(Path(args.data_dir, "scenes.private.json"))["scenes"] + truth_frequency = Counter( + group["color"] for row in vision_rows for group in scenes[row]["groups"] + ) + truth_rank = [name for name, _ in truth_frequency.most_common()] + text_frequency = Counter( + group["color"] for row in text_rows for group in scenes[row]["groups"] + ) + text_rank = [name for name, _ in text_frequency.most_common()] + report["evaluation_only"] = { + "vision_side_truth_frequency_rank": truth_rank, + "text_side_truth_frequency_rank": text_rank, + "rank_agreement": float( + np.mean([a == b for a, b in zip(truth_rank, text_rank)]) + ), + "text_colour_words_recovered": sorted( + set(families["colour_words"][: args.colour_classes]) + & set(truth_rank) + ), + } + print(json.dumps(report["text_families"], indent=1)) + print(json.dumps(report["evaluation_only"], indent=1)) + write_json(args.output, report) + print(f"Wrote {args.output}") + + +if __name__ == "__main__": + main() diff --git a/worldalign/tier0_pipeline.py b/worldalign/tier0_pipeline.py new file mode 100644 index 0000000..778e20f --- /dev/null +++ b/worldalign/tier0_pipeline.py @@ -0,0 +1,310 @@ +"""Tier 0 pipeline: unpaired corpora to aligned relation fields. + +Consolidates the components that produced the synthetic world's +end-to-end result, each of which was developed and measured separately: + +- watershed object extraction, which splits touching ring members that + connected components merge (exact object count 74.2% to 100%); +- Gestalt appearance grouping, which assembles members into groups by + shared colour, size, and shape rather than by a distance threshold + (exact group count 71.9% to 90.2%); +- size classes from radial extent under one-dimensional k-means per + member count, because area confounds size with shape and the classes + are gap-separated rather than equally populated (52.7% to 87.9%); +- text factor families from mutual exclusivity within a group phrase; +- the cross-modal value correspondence from marginal frequency rank. + +Nothing crosses modalities except the frequency ranking, and the two +corpora it reads are disjoint: no instance appears on both sides. +""" + +from __future__ import annotations + +import argparse +import json +from collections import Counter, defaultdict +from pathlib import Path + +import numpy as np +import torch +import torch.nn.functional as F +from scipy import ndimage +from scipy.cluster.hierarchy import fcluster, linkage +from sklearn.cluster import KMeans +from tqdm import tqdm + +from .common import read_json, seed_everything, write_json +from .synth_cc_battery import moment_field +from .synth_set_battery import parse_group_phrases +from .synth_towers import load_image +from .tier0_dictionary import partition_text_vocabulary, text_token_statistics + +NUMBER_TO_COUNT = {"two": 2, "three": 3, "four": 4} +SINGULAR = ("a", "an") + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--data-dir", default="artifacts/synth_v1") + parser.add_argument("--split", choices=["val", "test"], default="test") + parser.add_argument("--samples", type=int, default=256) + parser.add_argument("--offset", type=int, default=0) + parser.add_argument("--fit-scenes", type=int, default=1500) + parser.add_argument("--peak-distance", type=int, default=5) + parser.add_argument("--group-threshold", type=float, default=0.35) + parser.add_argument("--views", type=int, default=0, help="0 uses manifest.") + parser.add_argument("--seed", type=int, default=0) + parser.add_argument("--output", required=True) + parser.add_argument("--states-output", default="") + return parser.parse_args() + + +def extract_objects(image: torch.Tensor, peak_distance: int) -> list[dict]: + """Foreground objects, splitting touching ones by distance watershed.""" + array = image.permute(1, 2, 0).numpy() + background = np.median(array.reshape(-1, 3), axis=0) + foreground = np.abs(array - background).sum(-1) > 0.12 + distance = ndimage.distance_transform_edt(foreground) + window = 2 * peak_distance + 1 + peaks = ( + distance >= ndimage.maximum_filter(distance, size=window) - 1e-9 + ) & (distance > 1.0) + labels, count = ndimage.label(peaks) + if count == 0: + labels, count = ndimage.label(foreground) + else: + while True: + grown = ndimage.grey_dilation(labels, size=3) + take = (labels == 0) & foreground & (grown > 0) + if not take.any(): + break + labels = np.where(take, grown, labels) + count = labels.max() + objects = [] + for index in range(1, count + 1): + mask = labels == index + area = float(mask.sum()) + if area < 6: + continue + ys, xs = np.nonzero(mask) + centred_y, centred_x = ys - ys.mean(), xs - xs.mean() + covariance = np.cov(np.stack([centred_x, centred_y])) + 1e-6 * np.eye(2) + eigenvalues = np.linalg.eigvalsh(covariance) + eroded = ndimage.binary_erosion(mask) + perimeter = max(float((mask & ~eroded).sum()), 1.0) + box = (xs.max() - xs.min() + 1) * (ys.max() - ys.min() + 1) + objects.append( + { + "rgb": array[mask].mean(0), + "area": area, + "extent": float(np.hypot(centred_y, centred_x).max()), + "shape": np.array( + [ + 4 * np.pi * area / perimeter**2, + 1 - eigenvalues[0] / eigenvalues[1], + area / max(box, 1), + ] + ), + } + ) + return objects + + +def appearance(item: dict) -> np.ndarray: + return np.concatenate( + [item["rgb"] * 3.0, [np.log(item["area"] + 1e-6) * 0.6], item["shape"]] + ) + + +def group_objects(objects: list[dict], threshold: float) -> list[dict]: + """Members of one group share appearance; group by similarity.""" + if not objects: + return [] + if len(objects) == 1: + labels = np.array([0]) + else: + features = np.stack([appearance(item) for item in objects]) + labels = fcluster(linkage(features, "complete"), threshold, "distance") + buckets: dict[int, list[dict]] = {} + for item, label in zip(objects, labels): + buckets.setdefault(int(label), []).append(item) + return [ + { + "rgb": np.mean([item["rgb"] for item in members], axis=0), + "members": len(members), + "extent": float(np.mean([item["extent"] for item in members])), + } + for members in buckets.values() + ] + + +class VisionCoder: + """Colour classes and size classes fitted on the vision corpus alone.""" + + def __init__(self, groups: list[dict], classes: int, seed: int) -> None: + colours = np.stack([group["rgb"] for group in groups]).astype(np.float64) + self.colour_model = KMeans(classes, n_init=10, random_state=seed).fit(colours) + frequency = Counter(self.colour_model.labels_.tolist()) + self.colour_rank = { + label: rank for rank, (label, _) in enumerate(frequency.most_common()) + } + by_count: dict[int, list[float]] = defaultdict(list) + for group in groups: + by_count[min(group["members"], 4)].append(group["extent"]) + self.size_models = {} + for count, extents in by_count.items(): + model = KMeans(3, n_init=10, random_state=seed).fit( + np.asarray(extents, dtype=np.float64)[:, None] + ) + order = np.argsort(model.cluster_centers_[:, 0]) + self.size_models[count] = ( + model, + {int(label): rank for rank, label in enumerate(order)}, + ) + + def encode(self, groups: list[dict], classes: int) -> torch.Tensor: + if not groups: + groups = [{"rgb": np.zeros(3), "members": 1, "extent": 1.0}] + colours = self.colour_model.predict( + np.stack([group["rgb"] for group in groups]).astype(np.float64) + ) + vectors = [] + for group, colour in zip(groups, colours): + model, order = self.size_models[min(group["members"], 4)] + size = order[ + int(model.predict(np.array([[group["extent"]]], dtype=np.float64))[0]) + ] + vectors.append( + factor_vector(self.colour_rank[int(colour)], group["members"], size, classes) + ) + return F.normalize(torch.stack(vectors), dim=-1) + + +def factor_vector(colour: int, count: int, size: int, classes: int) -> torch.Tensor: + vector = torch.zeros(classes + 4 + 3) + vector[colour] = 1.0 + vector[classes + min(count - 1, 3)] = 1.0 + vector[classes + 4 + size] = 1.0 + return vector + + +def encode_caption(caption: str, colour_rank: dict[str, int], classes: int) -> torch.Tensor: + vectors = [] + for phrase in parse_group_phrases(caption): + tokens = phrase.split() + colour = next((colour_rank[t] for t in tokens if t in colour_rank), 0) + count = ( + 1 + if any(token in SINGULAR for token in tokens) + else next( + (NUMBER_TO_COUNT[t] for t in tokens if t in NUMBER_TO_COUNT), 1 + ) + ) + size = 0 if "small" in tokens else (2 if "large" in tokens else 1) + vectors.append(factor_vector(colour, count, size, classes)) + if not vectors: + vectors = [factor_vector(0, 1, 1, classes)] + return F.normalize(torch.stack(vectors), dim=-1) + + +def moment_state(states: torch.Tensor) -> torch.Tensor: + first = states.mean(0) + second = (states[:, :, None] * states[:, None, :]).mean(0).flatten() + return torch.cat([first, second]) + + +def main() -> None: + args = parse_args() + seed_everything(args.seed) + manifest = read_json(Path(args.data_dir, "manifest.json")) + captions = read_json(Path(args.data_dir, "captions.json"))["captions"] + image_dir = Path(manifest["image_dir"]) + views = args.views or manifest["visual_views"] + rows = manifest[args.split][args.offset : args.offset + args.samples] + + families = partition_text_vocabulary( + text_token_statistics(captions, manifest["text_only_train"]) + ) + colour_words = families["colour_words"] + classes = len(colour_words) + colour_rank = {word: rank for rank, word in enumerate(colour_words)} + + fit_groups = [ + group + for row in tqdm( + manifest["vision_only_train"][: args.fit_scenes], desc="fit vision" + ) + for group in group_objects( + extract_objects( + load_image(image_dir / f"scene{row:06d}_v0.png"), args.peak_distance + ), + args.group_threshold, + ) + ] + coder = VisionCoder(fit_groups, classes, args.seed) + + per_view_fields = [] + view_states = [] + for view in range(views): + sets = [ + coder.encode( + group_objects( + extract_objects( + load_image(image_dir / f"scene{row:06d}_v{view}.png"), + args.peak_distance, + ), + args.group_threshold, + ), + classes, + ) + for row in tqdm(rows, desc=f"encode view {view}") + ] + per_view_fields.append(moment_field(sets)) + if view == 0: + view_states = [moment_state(item) for item in sets] + visual_field = torch.stack(per_view_fields).mean(0) + + text_sets = [encode_caption(captions[row][0], colour_rank, classes) for row in rows] + text_field = moment_field(text_sets) + + mask = ~np.eye(len(rows), dtype=bool) + correlation = float( + np.corrcoef( + visual_field.double().numpy()[mask], text_field.double().numpy()[mask] + )[0, 1] + ) + torch.save( + {"visual_field": visual_field, "text_field": text_field, "rows": rows}, + args.output, + ) + if args.states_output: + torch.save( + { + "vision_states": torch.stack(view_states), + "text_states": torch.stack([moment_state(item) for item in text_sets]), + "rows": rows, + }, + args.states_output, + ) + summary = { + "data_dir": args.data_dir, + "split": args.split, + "samples": len(rows), + "colour_classes": classes, + "text_families": { + key: families[key] for key in ("count_words", "colour_words", "size_words") + }, + "field_correlation_at_truth": correlation, + "note": ( + "The dictionary is derived from disjoint corpora; the " + "correlation is a diagnostic computed with hidden pairs and " + "never used by the pipeline." + ), + } + print(json.dumps({"field_correlation_at_truth": correlation})) + write_json(args.output.replace(".pt", ".json"), summary) + print(f"Wrote {args.output}") + + +if __name__ == "__main__": + main() diff --git a/worldalign/train_bridge.py b/worldalign/train_bridge.py new file mode 100644 index 0000000..da29194 --- /dev/null +++ b/worldalign/train_bridge.py @@ -0,0 +1,255 @@ +from __future__ import annotations + +import argparse +from pathlib import Path + +import numpy as np +import torch +from torch.optim import AdamW +from tqdm import tqdm +from scipy.optimize import linear_sum_assignment + +from .common import ( + cosine_isometry_loss, + cosine_loss, + cosine_schedule, + parameter_count, + read_json, + retrieval_metrics, + seed_everything, + sliced_wasserstein, +) +from .gw import gw_pseudo_targets +from .io import load_feature_pair, select_rows +from .models import Bridge + + +def parse_args() -> argparse.Namespace: + p = argparse.ArgumentParser() + p.add_argument("--mode", choices=["paired", "unpaired_swd", "unpaired_gw"]) + p.add_argument("--manifest", default="artifacts/manifest.json") + p.add_argument("--vision", default="artifacts/vision.pt") + p.add_argument("--text", default="artifacts/text.pt") + p.add_argument("--output", required=True) + p.add_argument("--device", default="cuda:1") + p.add_argument("--steps", type=int, default=4_000) + p.add_argument("--batch-size", type=int, default=256) + p.add_argument("--hidden-dim", type=int, default=1536) + p.add_argument("--linear", action="store_true") + p.add_argument("--lr", type=float, default=3e-4) + p.add_argument("--warmup", type=int, default=200) + p.add_argument("--clusters", type=int, default=128) + p.add_argument( + "--gw-cache", + help="Optional path for reusable GW prototype coupling and assignments.", + ) + p.add_argument("--swd-weight", type=float, default=10.0) + p.add_argument("--isometry-weight", type=float, default=1.0) + p.add_argument("--gw-weight", type=float, default=1.0) + p.add_argument("--seed", type=int, default=20260728) + p.add_argument("--eval-every", type=int, default=200) + return p.parse_args() + + +def evaluate( + bridge: Bridge, + x: torch.Tensor, + y: torch.Tensor, + device: str, +) -> dict[str, float]: + bridge.eval() + mapped = [] + with torch.inference_mode(): + for chunk in x.split(512): + mapped.append(bridge(chunk.to(device)).cpu()) + bridge.train() + return retrieval_metrics(torch.cat(mapped), y) + + +def evaluate_gw_cluster_mapping( + gw: dict, + val_x: torch.Tensor, + val_y: torch.Tensor, +) -> dict[str, float]: + vx = torch.nn.functional.normalize(gw["vision_centers"], dim=-1) + ty = torch.nn.functional.normalize(gw["text_centers"], dim=-1) + x_labels = (torch.nn.functional.normalize(val_x, dim=-1) @ vx.T).argmax(1) + y_labels = (torch.nn.functional.normalize(val_y, dim=-1) @ ty.T).argmax(1) + predicted_map = gw["coupling"].argmax(1) + predicted = predicted_map[x_labels] + accuracy = (predicted == y_labels).float().mean().item() + + k = vx.shape[0] + contingency = torch.zeros(k, k, dtype=torch.float64) + for i, j in zip(x_labels.tolist(), y_labels.tolist()): + contingency[i, j] += 1 + row, col = linear_sum_assignment(-contingency.numpy()) + oracle_correct = contingency[row, col].sum().item() + return { + "paired_val_cluster_accuracy": float(accuracy), + "paired_val_cluster_chance": 1.0 / k, + "paired_val_cluster_oracle_permutation": float( + oracle_correct / max(len(val_x), 1) + ), + } + + +def main() -> None: + args = parse_args() + seed_everything(args.seed) + manifest = read_json(args.manifest) + vision, text, vlookup, tlookup = load_feature_pair(args.vision, args.text) + + if args.mode == "paired": + rows = manifest["paired_train"] + train_x = select_rows(vision["features"], vlookup, rows) + train_y = select_rows(text["features"], tlookup, rows) + target_by_cluster = None + assignments = None + gw_meta = {} + gw_result = None + else: + train_x = select_rows( + vision["features"], vlookup, manifest["vision_only_train"] + ) + train_y = select_rows( + text["features"], tlookup, manifest["text_only_train"] + ) + target_by_cluster = None + assignments = None + gw_meta = {} + gw_result = None + if args.mode == "unpaired_gw": + if args.gw_cache and Path(args.gw_cache).exists(): + print(f"Loading GW coupling from {args.gw_cache}") + gw = torch.load( + args.gw_cache, map_location="cpu", weights_only=False + ) + else: + print("Computing unpaired GW prototype coupling...") + gw = gw_pseudo_targets( + train_x, train_y, args.clusters, args.seed + ) + if args.gw_cache: + Path(args.gw_cache).parent.mkdir(parents=True, exist_ok=True) + torch.save(gw, args.gw_cache) + print(f"Wrote GW coupling to {args.gw_cache}") + gw_result = gw + target_by_cluster = gw["target_centers"] + assignments = gw["vision_assignments"] + gw_meta = { + "gw_distance": gw["gw_distance"], + "coupling_row_entropy": gw["coupling_row_entropy"], + "clusters": gw["clusters"], + } + print(f"GW diagnostics: {gw_meta}") + + val_rows = manifest["val"] + val_x = select_rows(vision["features"], vlookup, val_rows) + val_y = select_rows(text["features"], tlookup, val_rows) + if gw_result is not None: + gw_meta.update(evaluate_gw_cluster_mapping(gw_result, val_x, val_y)) + print(f"GW paired-eval diagnostics (not used for training): {gw_meta}") + + bridge = Bridge( + train_x.shape[-1], + train_y.shape[-1], + hidden_dim=args.hidden_dim, + linear=args.linear, + ).to(args.device) + print(f"Bridge parameters: {parameter_count(bridge):,}") + optimizer = AdamW(bridge.parameters(), lr=args.lr, weight_decay=1e-4) + generator = torch.Generator().manual_seed(args.seed) + + history = [] + best_score = -1.0 + best_state = None + progress = tqdm(range(args.steps), desc=f"bridge:{args.mode}") + for step in progress: + ix = torch.randint( + len(train_x), (args.batch_size,), generator=generator + ) + if args.mode == "paired": + iy = ix + else: + iy = torch.randint( + len(train_y), (args.batch_size,), generator=generator + ) + x = train_x[ix].to(args.device) + y = train_y[iy].to(args.device) + mapped = bridge(x) + + if args.mode == "paired": + alignment = cosine_loss(mapped, y) + swd = mapped.new_zeros(()) + gw_loss = mapped.new_zeros(()) + else: + alignment = mapped.new_zeros(()) + swd = sliced_wasserstein(mapped, y, num_projections=64) + if target_by_cluster is not None and assignments is not None: + pseudo = target_by_cluster[assignments[ix]].to(args.device) + gw_loss = cosine_loss(mapped, pseudo) + else: + gw_loss = mapped.new_zeros(()) + isometry = cosine_isometry_loss(x, mapped) + loss = ( + alignment + + args.swd_weight * swd + + args.gw_weight * gw_loss + + args.isometry_weight * isometry + ) + + optimizer.zero_grad(set_to_none=True) + loss.backward() + torch.nn.utils.clip_grad_norm_(bridge.parameters(), 1.0) + optimizer.step() + scale = cosine_schedule(step, args.steps, args.warmup) + for group in optimizer.param_groups: + group["lr"] = args.lr * scale + + if step % 20 == 0: + progress.set_postfix( + loss=f"{loss.item():.3f}", + align=f"{alignment.item():.3f}", + swd=f"{swd.item():.3g}", + gw=f"{gw_loss.item():.3f}", + iso=f"{isometry.item():.3f}", + ) + if step % args.eval_every == 0 or step == args.steps - 1: + metrics = evaluate(bridge, val_x, val_y, args.device) + record = {"step": step, "loss": float(loss.item()), **metrics} + history.append(record) + # Paired validation is recorded for scientific evaluation, not used to + # select the unsupervised checkpoint. Save final for unpaired modes. + if args.mode == "paired" and metrics["i2t_r@1"] > best_score: + best_score = metrics["i2t_r@1"] + best_state = { + k: v.detach().cpu().clone() + for k, v in bridge.state_dict().items() + } + print(record) + + if args.mode == "paired" and best_state is not None: + bridge.load_state_dict(best_state) + checkpoint_selection = "best paired validation R@1 (upper bound only)" + else: + checkpoint_selection = "final step; no paired validation selection" + + state = { + "config": bridge.config(), + "state_dict": bridge.state_dict(), + "mode": args.mode, + "args": vars(args), + "vision_model": vision["model"], + "text_model": text["model"], + "history": history, + "gw": gw_meta, + "checkpoint_selection": checkpoint_selection, + } + Path(args.output).parent.mkdir(parents=True, exist_ok=True) + torch.save(state, args.output) + print(f"Wrote {args.output}") + + +if __name__ == "__main__": + main() diff --git a/worldalign/train_prefix.py b/worldalign/train_prefix.py new file mode 100644 index 0000000..76607e6 --- /dev/null +++ b/worldalign/train_prefix.py @@ -0,0 +1,141 @@ +from __future__ import annotations + +import argparse +from pathlib import Path + +import torch +from torch.optim import AdamW +from tqdm import tqdm +from transformers import AutoModelForCausalLM, AutoTokenizer + +from .common import ( + cosine_schedule, + dtype_for_device, + parameter_count, + read_json, + seed_everything, +) +from .io import load_feature_pair, select_rows +from .models import PrefixAdapter + + +def parse_args() -> argparse.Namespace: + p = argparse.ArgumentParser() + p.add_argument("--manifest", default="artifacts/manifest.json") + p.add_argument("--vision", default="artifacts/vision.pt") + p.add_argument("--text", default="artifacts/text.pt") + p.add_argument("--output", default="artifacts/prefix.pt") + p.add_argument("--device", default="cuda:1") + p.add_argument("--steps", type=int, default=3_000) + p.add_argument("--batch-size", type=int, default=32) + p.add_argument("--prefix-length", type=int, default=8) + p.add_argument("--hidden-dim", type=int, default=2048) + p.add_argument("--max-length", type=int, default=48) + p.add_argument("--lr", type=float, default=3e-4) + p.add_argument("--warmup", type=int, default=200) + p.add_argument("--seed", type=int, default=20260728) + return p.parse_args() + + +def main() -> None: + args = parse_args() + seed_everything(args.seed) + manifest = read_json(args.manifest) + _, text, _, tlookup = load_feature_pair(args.vision, args.text) + train_rows = manifest["text_only_train"] + semantic = select_rows(text["features"], tlookup, train_rows) + captions_by_row = { + int(row): caption for row, caption in zip(text["rows"], text["captions"]) + } + captions = [captions_by_row[int(row)] for row in train_rows] + + tokenizer = AutoTokenizer.from_pretrained(text["model"]) + if tokenizer.pad_token_id is None: + tokenizer.pad_token = tokenizer.eos_token + tokenizer.padding_side = "right" + dtype = dtype_for_device(args.device) + lm = AutoModelForCausalLM.from_pretrained( + text["model"], torch_dtype=dtype + ).to(args.device) + lm.eval() + for parameter in lm.parameters(): + parameter.requires_grad_(False) + lm_dim = lm.get_input_embeddings().embedding_dim + + adapter = PrefixAdapter( + semantic_dim=semantic.shape[-1], + lm_dim=lm_dim, + prefix_length=args.prefix_length, + hidden_dim=args.hidden_dim, + ).to(args.device) + print(f"Prefix adapter parameters: {parameter_count(adapter):,}") + optimizer = AdamW(adapter.parameters(), lr=args.lr, weight_decay=1e-4) + generator = torch.Generator().manual_seed(args.seed) + + history = [] + progress = tqdm(range(args.steps), desc="text-only prefix") + for step in progress: + ids = torch.randint( + len(semantic), (args.batch_size,), generator=generator + ) + batch_captions = [captions[int(i)] for i in ids] + tokens = tokenizer( + batch_captions, + padding=True, + truncation=True, + max_length=args.max_length, + return_tensors="pt", + ) + input_ids = tokens["input_ids"].to(args.device) + attention = tokens["attention_mask"].to(args.device) + prefix = adapter(semantic[ids].to(args.device)).to(dtype) + token_embeddings = lm.get_input_embeddings()(input_ids) + inputs_embeds = torch.cat([prefix, token_embeddings], dim=1) + prefix_attention = torch.ones( + prefix.shape[:2], dtype=attention.dtype, device=args.device + ) + full_attention = torch.cat([prefix_attention, attention], dim=1) + labels = input_ids.clone() + labels[attention == 0] = -100 + prefix_labels = torch.full( + prefix.shape[:2], -100, dtype=labels.dtype, device=args.device + ) + full_labels = torch.cat([prefix_labels, labels], dim=1) + result = lm( + inputs_embeds=inputs_embeds, + attention_mask=full_attention, + labels=full_labels, + use_cache=False, + return_dict=True, + ) + loss = result.loss + optimizer.zero_grad(set_to_none=True) + loss.backward() + torch.nn.utils.clip_grad_norm_(adapter.parameters(), 1.0) + optimizer.step() + scale = cosine_schedule(step, args.steps, args.warmup) + for group in optimizer.param_groups: + group["lr"] = args.lr * scale + + if step % 20 == 0: + progress.set_postfix(loss=f"{loss.item():.3f}") + if step % 100 == 0 or step == args.steps - 1: + history.append({"step": step, "loss": float(loss.item())}) + + state = { + "config": adapter.config(), + "state_dict": adapter.state_dict(), + "text_model": text["model"], + "text_layer": text["layer"], + "args": vars(args), + "history": history, + "training": "text only; no image features or image-text pairs", + } + Path(args.output).parent.mkdir(parents=True, exist_ok=True) + torch.save(state, args.output) + print(f"Wrote {args.output}") + + +if __name__ == "__main__": + main() + diff --git a/worldalign/vg_attention_extract.py b/worldalign/vg_attention_extract.py new file mode 100644 index 0000000..0ceb4ce --- /dev/null +++ b/worldalign/vg_attention_extract.py @@ -0,0 +1,269 @@ +"""R3 battery extraction: model-layer relational readout for VG nodes. + +The hypothesis under test is that relations live in the model's +computation rather than in output embedding geometry. For each node this +extracts, per modality, a stack of view-pair relation channels read from +inside the frozen models: + +- text: the sixteen region phrases are encoded jointly in one context; + cross-phrase attention mass, pooled over layer groups and averaged over + several phrase orders (position and causality artifacts cancel), plus + in-context phrase states from the final layer. +- vision: the full image is encoded once; patch-patch attention pooled + over region boxes per layer group, plus in-context region states pooled + from patch tokens. + +Block means are computed as P A P^T with span-indicator matrices P, and +attention layers are reduced into group accumulators one layer at a time, +so neither the full layer stack nor per-pair loops materialize. + +No pairs, node identities, or view correspondences are used. Outputs are +keyed by released node IDs in released view order. +""" + +from __future__ import annotations + +import argparse +import json +from concurrent.futures import ThreadPoolExecutor +from pathlib import Path + +import numpy as np +import torch +from PIL import Image +from tqdm import tqdm +from transformers import AutoImageProcessor, AutoModel, AutoTokenizer + +from .common import batch_indices + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--side", choices=["text", "vision"], required=True) + parser.add_argument("--vg-dir", default="artifacts/vg_5k") + parser.add_argument("--image-cache", default="/tmp/yurenh2-worldalign-vg-images") + parser.add_argument("--text-model", default="Qwen/Qwen2.5-0.5B") + parser.add_argument("--vision-model", default="facebook/dinov2-small") + parser.add_argument("--device", default="cuda:3") + parser.add_argument("--orders", type=int, default=4, help="Text phrase orders.") + parser.add_argument("--image-size", type=int, default=224) + parser.add_argument( + "--batch-size", type=int, default=16, help="Sequences or images per forward." + ) + parser.add_argument("--limit", type=int) + parser.add_argument("--seed", type=int, default=20260729) + parser.add_argument("--output", required=True) + return parser.parse_args() + + +def read_jsonl(path: Path) -> list[dict]: + return [ + json.loads(line) + for line in path.read_text(encoding="utf-8").splitlines() + if line.strip() + ] + + +def layer_groups(count: int, groups: int) -> list[list[int]]: + bounds = torch.linspace(0, count, groups + 1).long().tolist() + return [list(range(a, b)) for a, b in zip(bounds[:-1], bounds[1:])] + + +def grouped_attention( + attentions: tuple[torch.Tensor, ...], groups: list[list[int]] +) -> torch.Tensor: + """Mean over heads and over each layer group, one layer at a time.""" + batch, _, length, _ = attentions[0].shape + result = torch.zeros( + len(groups), batch, length, length, device=attentions[0].device + ) + for g, layer_list in enumerate(groups): + for layer in layer_list: + result[g] += attentions[layer].float().mean(1) + result[g] /= len(layer_list) + return result + + +@torch.inference_mode() +def extract_text(args: argparse.Namespace) -> None: + records = read_jsonl(Path(args.vg_dir, "text_nodes.jsonl")) + if args.limit: + records = records[: args.limit] + tokenizer = AutoTokenizer.from_pretrained(args.text_model) + model = AutoModel.from_pretrained( + args.text_model, torch_dtype=torch.bfloat16, attn_implementation="eager" + ).to(args.device) + model.eval() + groups = layer_groups(model.config.num_hidden_layers, 4) + separator = tokenizer("\n", add_special_tokens=False)["input_ids"] + generator = torch.Generator().manual_seed(args.seed) + + jobs: list[tuple[int, list[int], list[tuple[int, int]]]] = [] + for index, record in enumerate(records): + phrases = record["region_closed"] + for _ in range(args.orders): + order = torch.randperm(len(phrases), generator=generator) + ids: list[int] = [] + span: list[tuple[int, int]] = [(0, 0)] * len(phrases) + for position in order.tolist(): + tokens = tokenizer( + " " + phrases[position].strip(), add_special_tokens=False + )["input_ids"] + span[position] = (len(ids), len(ids) + len(tokens)) + ids.extend(tokens + separator) + jobs.append((index, ids, span)) + + views = len(records[0]["region_closed"]) + attention_sum = torch.zeros(len(records), len(groups), views, views) + state_sum = torch.zeros(len(records), views, model.config.hidden_size) + + for start in tqdm(range(0, len(jobs), args.batch_size), desc="text attention"): + batch = jobs[start : start + args.batch_size] + longest = max(len(ids) for _, ids, _ in batch) + pad_id = tokenizer.pad_token_id or tokenizer.eos_token_id + input_ids = torch.full((len(batch), longest), pad_id, dtype=torch.long) + attention_mask = torch.zeros_like(input_ids) + indicator = torch.zeros(len(batch), views, longest) + for row, (_, ids, span) in enumerate(batch): + input_ids[row, : len(ids)] = torch.tensor(ids) + attention_mask[row, : len(ids)] = 1 + for view, (a0, a1) in enumerate(span): + indicator[row, view, a0:a1] = 1.0 / max(a1 - a0, 1) + result = model( + input_ids=input_ids.to(args.device), + attention_mask=attention_mask.to(args.device), + output_attentions=True, + return_dict=True, + ) + grouped = grouped_attention(result.attentions, groups) # [G, B, S, S] + indicator_device = indicator.to(args.device) + pooled = torch.einsum( + "bvs,gbst,bwt->gbvw", indicator_device, grouped, indicator_device + ).cpu() + hidden = result.last_hidden_state.float() + states = torch.bmm(indicator_device, hidden).cpu() + for row, (index, _, _) in enumerate(batch): + attention_sum[index] += pooled[:, row] + state_sum[index] += states[row] + + attention_channels = attention_sum / args.orders + attention_channels = 0.5 * ( + attention_channels + attention_channels.transpose(-2, -1) + ) + torch.save( + { + "side": "text", + "model": args.text_model, + "node_ids": [record["node_id"] for record in records], + "attention_channels": attention_channels, + "context_states": state_sum / args.orders, + "orders": args.orders, + "layer_groups": [len(g) for g in groups], + }, + args.output, + ) + print(f"Wrote {args.output}") + + +@torch.inference_mode() +def extract_vision(args: argparse.Namespace) -> None: + records = read_jsonl(Path(args.vg_dir, "vision_nodes.private.jsonl")) + if args.limit: + records = records[: args.limit] + processor = AutoImageProcessor.from_pretrained(args.vision_model) + model = AutoModel.from_pretrained( + args.vision_model, torch_dtype=torch.float32, attn_implementation="eager" + ).to(args.device) + model.eval() + groups = layer_groups(model.config.num_hidden_layers, 3) + patch = model.config.patch_size + grid = args.image_size // patch + tokens = grid * grid + mean = torch.tensor(processor.image_mean).view(3, 1, 1) + std = torch.tensor(processor.image_std).view(3, 1, 1) + + def load_one(record: dict) -> torch.Tensor: + path = Path(args.image_cache, f"{record['source_image_id']}.jpg") + with Image.open(path) as image: + resized = image.convert("RGB").resize( + (args.image_size, args.image_size), Image.BILINEAR + ) + pixels = torch.from_numpy(np.asarray(resized).copy()).permute(2, 0, 1) + return (pixels.float() / 255.0 - mean) / std + + def region_indicator(record: dict) -> torch.Tensor: + width, height = record["width"], record["height"] + rows = [] + centers = torch.arange(grid) + 0.5 + for region in record["regions"]: + x0 = region["x"] / width * grid + x1 = (region["x"] + region["width"]) / width * grid + y0 = region["y"] / height * grid + y1 = (region["y"] + region["height"]) / height * grid + in_x = (centers >= x0) & (centers <= x1) + in_y = (centers >= y0) & (centers <= y1) + mask = (in_y[:, None] & in_x[None, :]).flatten().double() + if mask.sum() == 0: + cx = min(grid - 1, max(0, int((x0 + x1) / 2))) + cy = min(grid - 1, max(0, int((y0 + y1) / 2))) + mask[cy * grid + cx] = 1.0 + rows.append(mask / mask.sum()) + return torch.stack(rows).float() + + views = len(records[0]["regions"]) + node_ids: list[str] = [] + attention_channels: list[torch.Tensor] = [] + context_states: list[torch.Tensor] = [] + with ThreadPoolExecutor(max_workers=8) as pool: + for indices in tqdm( + list(batch_indices(len(records), args.batch_size)), + desc="vision attention", + ): + batch = [records[i] for i in indices] + pixels = torch.stack(list(pool.map(load_one, batch))) + result = model( + pixel_values=pixels.to(args.device), + output_attentions=True, + return_dict=True, + ) + grouped = grouped_attention(result.attentions, groups) + special = grouped.shape[-1] - tokens # CLS and any registers + grouped = grouped[:, :, special:, special:] + indicator = torch.stack( + [region_indicator(record) for record in batch] + ).to(args.device) + pooled = torch.einsum( + "bvs,gbst,bwt->gbvw", indicator, grouped, indicator + ).cpu() + pooled = 0.5 * (pooled + pooled.transpose(-2, -1)) + patches = result.last_hidden_state[:, special:].float() + states = torch.bmm(indicator, patches).cpu() + for row, record in enumerate(batch): + node_ids.append(record["node_id"]) + attention_channels.append(pooled[:, row]) + context_states.append(states[row]) + torch.save( + { + "side": "vision", + "model": args.vision_model, + "node_ids": node_ids, + "attention_channels": torch.stack(attention_channels), + "context_states": torch.stack(context_states), + "image_size": args.image_size, + "layer_groups": [len(g) for g in groups], + }, + args.output, + ) + print(f"Wrote {args.output}") + + +def main() -> None: + args = parse_args() + if args.side == "text": + extract_text(args) + else: + extract_vision(args) + + +if __name__ == "__main__": + main() diff --git a/worldalign/vg_attention_probe.py b/worldalign/vg_attention_probe.py new file mode 100644 index 0000000..3dac7a3 --- /dev/null +++ b/worldalign/vg_attention_probe.py @@ -0,0 +1,170 @@ +"""R3 battery evaluation: do model-internal relations align across modalities? + +Compares, at the replayed true view correspondence, the cross-modal +alignment of relation fields read from inside the frozen models +(cross-phrase attention per layer group, in-context state cosines) against +the isolated-encoding cosine baseline (matched rho 0.128, z 46.6). +Spearman is rank-based, so monotone attention transforms are immaterial. +View truth is evaluation-only. +""" + +from __future__ import annotations + +import argparse +import json + +import torch +import torch.nn.functional as F + +from .common import write_json +from .vg_view_probe import offdiag, replay_view_permutations, spearman + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--vg-dir", default="artifacts/vg_5k") + parser.add_argument("--cache-dir", default="/tmp/yurenh2-worldalign-vg-hf") + parser.add_argument("--seed", type=int, default=20260728) + parser.add_argument("--views", type=int, default=16) + parser.add_argument("--text-attn", default="artifacts/manifold_gate/attn_text_full.pt") + parser.add_argument( + "--vision-attn", default="artifacts/manifold_gate/attn_vision_full.pt" + ) + parser.add_argument("--nodes", type=int, default=0) + parser.add_argument( + "--output", default="artifacts/manifold_gate/attention_probe.json" + ) + return parser.parse_args() + + +def channel_fields(state: dict, side: str) -> dict[str, torch.Tensor]: + """Named [N, V, V] relation fields for one side.""" + attention = state["attention_channels"].double() + fields = { + f"{side}_attn_g{group}": attention[:, group] + for group in range(attention.shape[1]) + } + context = F.normalize(state["context_states"].double(), dim=-1) + fields[f"{side}_ctx_cos"] = context @ context.transpose(-2, -1) + return fields + + +def main() -> None: + args = parse_args() + mappings = replay_view_permutations(args) + + text_state = torch.load(args.text_attn, map_location="cpu", weights_only=False) + vision_state = torch.load(args.vision_attn, map_location="cpu", weights_only=False) + text_index = {node: i for i, node in enumerate(text_state["node_ids"])} + vision_index = {node: i for i, node in enumerate(vision_state["node_ids"])} + + baseline_vision = torch.load( + f"{args.vg_dir}/vision_features.pt", map_location="cpu", weights_only=False + ) + baseline_text = torch.load( + f"{args.vg_dir}/text_features.pt", map_location="cpu", weights_only=False + ) + baseline_vision_index = { + node: i for i, node in enumerate(baseline_vision["node_ids"]) + } + baseline_text_index = {node: i for i, node in enumerate(baseline_text["node_ids"])} + + text_fields = channel_fields(text_state, "text") + vision_fields = channel_fields(vision_state, "vision") + + node_ids = sorted(mappings) + if args.nodes: + node_ids = node_ids[: args.nodes] + + pairs = [ + (t_name, v_name) for t_name in text_fields for v_name in vision_fields + ] + matched: dict[tuple[str, str], list[float]] = {pair: [] for pair in pairs} + shuffled: dict[tuple[str, str], list[float]] = {pair: [] for pair in pairs} + baseline_matched: list[float] = [] + baseline_shuffled: list[float] = [] + generator = torch.Generator().manual_seed(args.seed) + + for node in node_ids: + mapping = mappings[node] + text_row = text_index[mapping["text_node_id"]] + vision_row = vision_index[node] + text_to_vision = torch.tensor(mapping["text_to_vision"]) + vision_to_text = torch.empty_like(text_to_vision) + vision_to_text[text_to_vision] = torch.arange(len(text_to_vision)) + shuffle = torch.randperm(args.views, generator=generator) + + for t_name, v_name in pairs: + text_field = text_fields[t_name][text_row] + aligned = text_field[vision_to_text][:, vision_to_text] + visual_field = vision_fields[v_name][vision_row] + matched[(t_name, v_name)].append( + spearman(offdiag(visual_field), offdiag(aligned)) + ) + scrambled = text_field[shuffle][:, shuffle] + shuffled[(t_name, v_name)].append( + spearman(offdiag(visual_field), offdiag(scrambled)) + ) + + crops = F.normalize( + baseline_vision["region_features"][ + baseline_vision_index[node] + ].double(), + dim=-1, + ) + phrases = F.normalize( + baseline_text["region_features"][ + baseline_text_index[mapping["text_node_id"]] + ].double(), + dim=-1, + ) + crop_field = crops @ crops.T + phrase_field = (phrases @ phrases.T)[vision_to_text][:, vision_to_text] + baseline_matched.append(spearman(offdiag(crop_field), offdiag(phrase_field))) + scrambled = (phrases @ phrases.T)[shuffle][:, shuffle] + baseline_shuffled.append(spearman(offdiag(crop_field), offdiag(scrambled))) + + def summarize(matched_values: list[float], shuffled_values: list[float]) -> dict: + m = torch.tensor(matched_values) + s = torch.tensor(shuffled_values) + return { + "matched_mean": float(m.mean()), + "shuffled_mean": float(s.mean()), + "gap_z": float( + (m.mean() - s.mean()) + / (m - s).std().clamp_min(1e-12) + * len(m) ** 0.5 + ), + } + + report = { + "protocol": ( + "Replayed view truth, evaluation-only. Each cell is the " + "within-node cross-modal Spearman of relation fields at the " + "true view correspondence, against a shuffled-view control." + ), + "nodes": len(node_ids), + "baseline_isolated_cosine": summarize(baseline_matched, baseline_shuffled), + "channels": { + f"{t} x {v}": summarize(matched[(t, v)], shuffled[(t, v)]) + for t, v in pairs + }, + } + write_json(args.output, report) + best = max( + report["channels"].items(), key=lambda item: item[1]["matched_mean"] + ) + print( + json.dumps( + { + "baseline": report["baseline_isolated_cosine"], + "best_channel": {best[0]: best[1]}, + "nodes": len(node_ids), + } + ) + ) + print(f"Wrote {args.output}") + + +if __name__ == "__main__": + main() diff --git a/worldalign/vg_diagnose.py b/worldalign/vg_diagnose.py new file mode 100644 index 0000000..4c6d4a8 --- /dev/null +++ b/worldalign/vg_diagnose.py @@ -0,0 +1,200 @@ +from __future__ import annotations + +import argparse +import json +from pathlib import Path + +import numpy as np +from scipy.stats import spearmanr +import torch + +from .common import normalized, write_json + + +def parse_args() -> argparse.Namespace: + p = argparse.ArgumentParser() + p.add_argument("--vision", default="artifacts/vg/vision_features.pt") + p.add_argument("--text", default="artifacts/vg/text_features.pt") + p.add_argument( + "--truth", default="artifacts/vg/ground_truth.private.jsonl" + ) + p.add_argument("--max-nodes", type=int, default=5_000) + p.add_argument("--quantiles", type=int, default=33) + p.add_argument("--output", default="artifacts/vg/diagnostics.json") + return p.parse_args() + + +def read_jsonl(path: str) -> list[dict]: + with open(path, encoding="utf-8") as handle: + return [json.loads(line) for line in handle if line.strip()] + + +def bundle_signature( + features: torch.Tensor, quantiles: int = 33 +) -> torch.Tensor: + """Permutation- and rotation-invariant signature of a bag of views.""" + x = normalized(features.float()) + gram = x @ x.transpose(1, 2) + views = x.shape[1] + i, j = torch.triu_indices(views, views, offset=1) + pairwise = gram[:, i, j] + q = torch.linspace(0, 1, quantiles) + distribution = torch.quantile(pairwise, q, dim=1).T + eigenvalues = torch.linalg.eigvalsh(gram).flip(-1) / max(views, 1) + return torch.cat([distribution, eigenvalues], dim=-1) + + +def rank_normalize_columns(x: torch.Tensor) -> torch.Tensor: + order = torch.argsort(x, dim=0) + ranks = torch.argsort(order, dim=0).float() + ranks = ranks / max(x.shape[0] - 1, 1) + return (ranks - 0.5) * 2 + + +def aligned_text_indices( + vision_ids: list[str], text_ids: list[str], truth_path: str +) -> torch.Tensor: + truth = read_jsonl(truth_path) + mapping = { + item["vision_node_id"]: item["text_node_id"] for item in truth + } + text_lookup = {node_id: idx for idx, node_id in enumerate(text_ids)} + return torch.tensor([text_lookup[mapping[node_id]] for node_id in vision_ids]) + + +def arbitrary_target_retrieval( + queries: torch.Tensor, candidates: torch.Tensor, targets: torch.Tensor +) -> dict: + similarity = normalized(queries) @ normalized(candidates).T + target_score = similarity[ + torch.arange(len(queries)), targets.to(similarity.device) + ] + ranks = (similarity > target_score[:, None]).sum(-1) + 1 + top_values, top_indices = similarity.topk( + min(2, similarity.shape[1]), dim=1 + ) + prediction = top_indices[:, 0] + correct = prediction == targets.to(prediction.device) + if top_values.shape[1] == 2: + margin = top_values[:, 0] - top_values[:, 1] + else: + margin = top_values[:, 0] + confidence_order = torch.argsort(margin, descending=True) + confidence_precision = {} + for count in (10, 50, 100, 500, 1_000): + if count <= len(queries): + selected = confidence_order[:count] + confidence_precision[str(count)] = { + "correct": int(correct[selected].sum()), + "precision": float(correct[selected].float().mean()), + } + + reverse_prediction = similarity.argmax(dim=0) + mutual = ( + reverse_prediction[prediction] + == torch.arange(len(queries), device=prediction.device) + ) + mutual_count = int(mutual.sum()) + result = { + "r@1": float((ranks <= 1).float().mean()), + "r@5": float((ranks <= 5).float().mean()), + "r@10": float((ranks <= 10).float().mean()), + "mean_reciprocal_rank": float((1.0 / ranks.float()).mean()), + "median_rank": float(ranks.float().median()), + "chance_r@1": 1.0 / len(candidates), + "chance_r@5": min(5.0 / len(candidates), 1.0), + "chance_r@10": min(10.0 / len(candidates), 1.0), + "confidence_margin_precision": confidence_precision, + "mutual_nearest": { + "selected": mutual_count, + "correct": int(correct[mutual].sum()), + "precision": ( + float(correct[mutual].float().mean()) + if mutual_count + else 0.0 + ), + }, + } + return result + + +def upper_triangle(x: torch.Tensor) -> np.ndarray: + i, j = torch.triu_indices(len(x), len(x), offset=1) + return x[i, j].cpu().numpy() + + +def main() -> None: + args = parse_args() + vision = torch.load(args.vision, map_location="cpu", weights_only=False) + text = torch.load(args.text, map_location="cpu", weights_only=False) + n = min(args.max_nodes, len(vision["node_ids"]), len(text["node_ids"])) + vision_ids = vision["node_ids"][:n] + target_full = aligned_text_indices( + vision_ids, text["node_ids"], args.truth + ) + candidate_indices = torch.unique(target_full, sorted=False) + if len(candidate_indices) != n: + raise ValueError("Ground truth is not a one-to-one permutation") + candidate_lookup = { + int(old): new for new, old in enumerate(candidate_indices.tolist()) + } + targets = torch.tensor( + [candidate_lookup[int(old)] for old in target_full.tolist()] + ) + + v_views = vision["region_features"][:n] + t_views = text["region_features"][candidate_indices] + v_signature = rank_normalize_columns( + bundle_signature(v_views, args.quantiles) + ) + t_signature = rank_normalize_columns( + bundle_signature(t_views, args.quantiles) + ) + signature_retrieval = arbitrary_target_retrieval( + v_signature, t_signature, targets + ) + + paired_t_signature = t_signature[targets] + signature_cosine = ( + normalized(v_signature) * normalized(paired_t_signature) + ).sum(-1) + generator = torch.Generator().manual_seed(20260728) + shuffled = paired_t_signature[torch.randperm(n, generator=generator)] + shuffled_cosine = ( + normalized(v_signature) * normalized(shuffled) + ).sum(-1) + + v_scene = vision["global_features"][:n] + t_scene = text["region_features"][candidate_indices].mean(1)[targets] + gv = normalized(v_scene) @ normalized(v_scene).T + gt = normalized(t_scene) @ normalized(t_scene).T + scene_rho = spearmanr( + upper_triangle(gv), upper_triangle(gt) + ).statistic + + result = { + "nodes": n, + "views_per_node": int(v_views.shape[1]), + "vision_model": vision["model"], + "text_model": text["model"], + "text_tier": text["tier"], + "bundle_signature_retrieval": signature_retrieval, + "paired_bundle_signature_cosine_mean": float( + signature_cosine.mean() + ), + "shuffled_bundle_signature_cosine_mean": float( + shuffled_cosine.mean() + ), + "between_scene_pairwise_cosine_spearman": float(scene_rho), + "evaluation_note": ( + "Ground-truth permutation is used only to score signatures and " + "relation geometry, never to fit them." + ), + } + Path(args.output).parent.mkdir(parents=True, exist_ok=True) + write_json(args.output, result) + print(json.dumps(result, indent=2)) + + +if __name__ == "__main__": + main() diff --git a/worldalign/vg_extract_text.py b/worldalign/vg_extract_text.py new file mode 100644 index 0000000..affff63 --- /dev/null +++ b/worldalign/vg_extract_text.py @@ -0,0 +1,111 @@ +from __future__ import annotations + +import argparse +import json +from pathlib import Path + +import torch +from tqdm import tqdm +from transformers import AutoModel, AutoTokenizer + +from .common import dtype_for_device +from .extract_text import mean_pool + + +def parse_args() -> argparse.Namespace: + p = argparse.ArgumentParser() + p.add_argument("--nodes", default="artifacts/vg/text_nodes.jsonl") + p.add_argument("--output", default="artifacts/vg/text_features.pt") + p.add_argument( + "--tier", + choices=["region_closed", "visible_relations", "qa_expanded"], + default="region_closed", + ) + p.add_argument("--model", default="Qwen/Qwen2.5-1.5B") + p.add_argument("--device", default="cuda:3") + p.add_argument("--batch-size", type=int, default=96) + p.add_argument("--max-length", type=int, default=48) + p.add_argument("--layer", type=int, default=-1) + p.add_argument("--views", type=int, default=32) + p.add_argument("--limit", type=int) + return p.parse_args() + + +def read_jsonl(path: str) -> list[dict]: + with open(path, encoding="utf-8") as handle: + return [json.loads(line) for line in handle if line.strip()] + + +@torch.inference_mode() +def main() -> None: + args = parse_args() + records = read_jsonl(args.nodes) + if args.limit: + records = records[: args.limit] + if not records: + raise ValueError("No text nodes found") + for record in records: + if args.tier not in record: + raise ValueError( + f"Tier {args.tier!r} is absent; regenerate bundles with " + "--with-extra-tiers if needed" + ) + if len(record[args.tier]) < args.views: + raise ValueError( + f"Node {record['node_id']} has fewer than {args.views} views" + ) + + tokenizer = AutoTokenizer.from_pretrained(args.model) + if tokenizer.pad_token_id is None: + tokenizer.pad_token = tokenizer.eos_token + tokenizer.padding_side = "right" + dtype = dtype_for_device(args.device) + model = AutoModel.from_pretrained(args.model, torch_dtype=dtype).to( + args.device + ) + model.eval() + + flat_texts = [ + text + for record in records + for text in record[args.tier][: args.views] + ] + outputs: list[torch.Tensor] = [] + for start in tqdm( + range(0, len(flat_texts), args.batch_size), + desc=f"Qwen VG text:{args.tier}", + ): + tokens = tokenizer( + flat_texts[start : start + args.batch_size], + padding=True, + truncation=True, + max_length=args.max_length, + return_tensors="pt", + ) + tokens = {key: value.to(args.device) for key, value in tokens.items()} + result = model( + **tokens, output_hidden_states=True, return_dict=True + ) + outputs.append( + mean_pool( + result.hidden_states[args.layer], tokens["attention_mask"] + ) + .float() + .cpu() + ) + features = torch.cat(outputs).reshape(len(records), args.views, -1) + state = { + "model": args.model, + "layer": args.layer, + "tier": args.tier, + "node_ids": [record["node_id"] for record in records], + "region_features": features, + "views_per_node": args.views, + } + Path(args.output).parent.mkdir(parents=True, exist_ok=True) + torch.save(state, args.output) + print(f"Wrote {args.output}: {tuple(features.shape)}") + + +if __name__ == "__main__": + main() diff --git a/worldalign/vg_extract_vision.py b/worldalign/vg_extract_vision.py new file mode 100644 index 0000000..504840a --- /dev/null +++ b/worldalign/vg_extract_vision.py @@ -0,0 +1,150 @@ +from __future__ import annotations + +import argparse +from concurrent.futures import ThreadPoolExecutor, as_completed +import json +from pathlib import Path +import time + +from PIL import Image +import requests +import torch +from tqdm import tqdm +from transformers import AutoImageProcessor, AutoModel + +from .common import dtype_for_device + + +def parse_args() -> argparse.Namespace: + p = argparse.ArgumentParser() + p.add_argument( + "--nodes", default="artifacts/vg/vision_nodes.private.jsonl" + ) + p.add_argument("--output", default="artifacts/vg/vision_features.pt") + p.add_argument( + "--image-cache", default="/tmp/yurenh2-worldalign-vg-images" + ) + p.add_argument("--model", default="facebook/dinov2-small") + p.add_argument("--device", default="cuda:1") + p.add_argument("--node-batch-size", type=int, default=8) + p.add_argument("--download-workers", type=int, default=32) + p.add_argument("--limit", type=int) + return p.parse_args() + + +def read_jsonl(path: str) -> list[dict]: + with open(path, encoding="utf-8") as handle: + return [json.loads(line) for line in handle if line.strip()] + + +def download_one(record: dict, image_cache: Path) -> Path: + image_cache.mkdir(parents=True, exist_ok=True) + output = image_cache / f"{int(record['source_image_id'])}.jpg" + if output.exists() and output.stat().st_size > 0: + return output + temporary = output.with_suffix(".jpg.part") + last_error: Exception | None = None + for attempt in range(3): + try: + response = requests.get(record["url"], timeout=30) + response.raise_for_status() + temporary.write_bytes(response.content) + with Image.open(temporary) as image: + image.verify() + temporary.replace(output) + return output + except Exception as error: + last_error = error + if temporary.exists(): + temporary.unlink() + time.sleep(1 + attempt) + raise RuntimeError(f"Failed to download {record['url']}") from last_error + + +def crop_views(record: dict, path: Path) -> tuple[Image.Image, list[Image.Image]]: + with Image.open(path) as source: + image = source.convert("RGB") + width, height = image.size + crops: list[Image.Image] = [] + for region in record["regions"]: + x0 = max(0, min(int(region["x"]), width - 1)) + y0 = max(0, min(int(region["y"]), height - 1)) + x1 = max(x0 + 1, min(x0 + int(region["width"]), width)) + y1 = max(y0 + 1, min(y0 + int(region["height"]), height)) + crops.append(image.crop((x0, y0, x1, y1))) + return image, crops + + +@torch.inference_mode() +def main() -> None: + args = parse_args() + records = read_jsonl(args.nodes) + if args.limit: + records = records[: args.limit] + if not records: + raise ValueError("No vision nodes found") + view_count = len(records[0]["regions"]) + if any(len(record["regions"]) != view_count for record in records): + raise ValueError("All nodes must have the same number of region views") + + image_cache = Path(args.image_cache) + paths: dict[str, Path] = {} + with ThreadPoolExecutor(max_workers=args.download_workers) as executor: + futures = { + executor.submit(download_one, record, image_cache): record["node_id"] + for record in records + } + for future in tqdm( + as_completed(futures), total=len(futures), desc="VG image download" + ): + paths[futures[future]] = future.result() + + processor = AutoImageProcessor.from_pretrained(args.model) + dtype = dtype_for_device(args.device) + model = AutoModel.from_pretrained(args.model, torch_dtype=dtype).to( + args.device + ) + model.eval() + + region_outputs: list[torch.Tensor] = [] + global_outputs: list[torch.Tensor] = [] + for start in tqdm( + range(0, len(records), args.node_batch_size), + desc="DINO VG bundles", + ): + batch = records[start : start + args.node_batch_size] + images: list[Image.Image] = [] + for record in batch: + global_image, crops = crop_views(record, paths[record["node_id"]]) + images.append(global_image) + images.extend(crops) + pixels = processor(images=images, return_tensors="pt")[ + "pixel_values" + ].to(args.device, dtype=dtype) + hidden = model(pixel_values=pixels, return_dict=True).last_hidden_state[ + :, 0 + ] + hidden = hidden.float().cpu().reshape( + len(batch), view_count + 1, -1 + ) + global_outputs.append(hidden[:, 0]) + region_outputs.append(hidden[:, 1:]) + + state = { + "model": args.model, + "node_ids": [record["node_id"] for record in records], + "region_features": torch.cat(region_outputs), + "global_features": torch.cat(global_outputs), + "views_per_node": view_count, + "source_metadata_removed": True, + } + Path(args.output).parent.mkdir(parents=True, exist_ok=True) + torch.save(state, args.output) + print( + f"Wrote {args.output}: regions={tuple(state['region_features'].shape)}, " + f"global={tuple(state['global_features'].shape)}" + ) + + +if __name__ == "__main__": + main() diff --git a/worldalign/vg_graph_diagnose.py b/worldalign/vg_graph_diagnose.py new file mode 100644 index 0000000..3a0ad2e --- /dev/null +++ b/worldalign/vg_graph_diagnose.py @@ -0,0 +1,222 @@ +from __future__ import annotations + +import argparse +import json +from pathlib import Path + +import numpy as np +from scipy import sparse +from scipy.sparse.linalg import eigsh +import torch + +from .common import normalized, write_json +from .vg_diagnose import ( + aligned_text_indices, + arbitrary_target_retrieval, + bundle_signature, + rank_normalize_columns, +) + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument( + "--vision", default="artifacts/vg/vision_features.pt" + ) + parser.add_argument("--text", default="artifacts/vg/text_features.pt") + parser.add_argument( + "--truth", default="artifacts/vg/ground_truth.private.jsonl" + ) + parser.add_argument("--max-nodes", type=int, default=5_000) + parser.add_argument("--knn", type=int, default=32) + parser.add_argument("--eigenpairs", type=int, default=96) + parser.add_argument("--heat-scales", type=int, default=24) + parser.add_argument("--message-passing-steps", type=int, default=3) + parser.add_argument("--output", default="artifacts/vg/graph_diagnostics.json") + return parser.parse_args() + + +def scene_features(state: dict, indices: torch.Tensor | None = None) -> torch.Tensor: + if "global_features" in state: + features = state["global_features"] + else: + features = state["region_features"].mean(1) + return features if indices is None else features[indices] + + +def exact_rank_knn_graph( + features: torch.Tensor, neighbors: int +) -> tuple[sparse.csr_matrix, torch.Tensor]: + """Build a symmetrized rank-weighted graph and retain density statistics.""" + x = normalized(features.float()) + similarity = x @ x.T + similarity.fill_diagonal_(-torch.inf) + k = min(neighbors, len(x) - 1) + values, indices = similarity.topk(k, dim=1) + + # Rank weights avoid assuming that cosine scales agree across modalities. + rank_weight = np.exp( + -np.arange(k, dtype=np.float64) / max(k / 4.0, 1.0) + ) + rows = np.repeat(np.arange(len(x), dtype=np.int64), k) + cols = indices.cpu().numpy().reshape(-1) + data = np.tile(rank_weight, len(x)) + directed = sparse.csr_matrix( + (data, (rows, cols)), shape=(len(x), len(x)) + ) + adjacency = directed.maximum(directed.T) + adjacency.setdiag(0) + adjacency.eliminate_zeros() + + density_quantiles = torch.quantile( + values, + torch.tensor([0.0, 0.25, 0.5, 0.75, 1.0]), + dim=1, + ).T + return adjacency, density_quantiles + + +def heat_kernel_signature( + adjacency: sparse.csr_matrix, eigenpairs: int, scales: int +) -> tuple[torch.Tensor, np.ndarray]: + """Return basis-invariant diagonal heat-kernel signatures.""" + degree = np.asarray(adjacency.sum(axis=1)).reshape(-1) + inv_sqrt_degree = 1.0 / np.sqrt(np.maximum(degree, 1e-12)) + normalizer = sparse.diags(inv_sqrt_degree) + normalized_adjacency = normalizer @ adjacency @ normalizer + k = min(eigenpairs, adjacency.shape[0] - 2) + eigenvalues, eigenvectors = eigsh( + normalized_adjacency, k=k, which="LA", tol=1e-4 + ) + order = np.argsort(eigenvalues)[::-1] + eigenvalues = eigenvalues[order] + eigenvectors = eigenvectors[:, order] + laplacian_eigenvalues = np.clip(1.0 - eigenvalues, 0.0, None) + + positive = laplacian_eigenvalues[laplacian_eigenvalues > 1e-8] + lower = 0.25 / max(float(positive.max(initial=1.0)), 1e-8) + upper = 8.0 / max(float(positive.min(initial=1.0)), 1e-8) + times = np.geomspace(lower, upper, scales) + heat = (eigenvectors**2) @ np.exp( + -laplacian_eigenvalues[:, None] * times[None, :] + ) + return torch.from_numpy(heat).float(), degree + + +def graph_signature( + features: torch.Tensor, + neighbors: int, + eigenpairs: int, + heat_scales: int, + message_passing_steps: int, +) -> torch.Tensor: + adjacency, density = exact_rank_knn_graph(features, neighbors) + heat, degree = heat_kernel_signature( + adjacency, eigenpairs=eigenpairs, scales=heat_scales + ) + base = rank_normalize_columns( + torch.cat( + [ + density, + torch.from_numpy(np.log1p(degree)).float()[:, None], + heat, + ], + dim=1, + ) + ) + transition = sparse.diags( + 1.0 / np.maximum(np.asarray(adjacency.sum(axis=1)).reshape(-1), 1e-12) + ) @ adjacency + current = base + for _ in range(message_passing_steps): + current_np = current.numpy() + neighbor_mean = torch.from_numpy(transition @ current_np).float() + neighbor_square_mean = torch.from_numpy( + transition @ np.square(current_np) + ).float() + neighbor_std = torch.sqrt( + torch.clamp(neighbor_square_mean - neighbor_mean.square(), min=0) + ) + current = rank_normalize_columns( + torch.cat([base, neighbor_mean, neighbor_std], dim=1) + ) + return current + + +def main() -> None: + args = parse_args() + vision = torch.load(args.vision, map_location="cpu", weights_only=False) + text = torch.load(args.text, map_location="cpu", weights_only=False) + n = min(args.max_nodes, len(vision["node_ids"]), len(text["node_ids"])) + vision_ids = vision["node_ids"][:n] + target_full = aligned_text_indices( + vision_ids, text["node_ids"], args.truth + ) + candidate_indices = torch.unique(target_full, sorted=False) + if len(candidate_indices) != n: + raise ValueError("Ground truth is not a one-to-one permutation") + candidate_lookup = { + int(old): new for new, old in enumerate(candidate_indices.tolist()) + } + targets = torch.tensor( + [candidate_lookup[int(old)] for old in target_full.tolist()] + ) + + vision_scene = scene_features(vision)[:n] + text_scene = scene_features(text, candidate_indices) + vision_graph = graph_signature( + vision_scene, + args.knn, + args.eigenpairs, + args.heat_scales, + args.message_passing_steps, + ) + text_graph = graph_signature( + text_scene, + args.knn, + args.eigenpairs, + args.heat_scales, + args.message_passing_steps, + ) + + vision_bundle = rank_normalize_columns( + bundle_signature(vision["region_features"][:n]) + ) + text_bundle = rank_normalize_columns( + bundle_signature(text["region_features"][candidate_indices]) + ) + combined_vision = torch.cat( + [normalized(vision_graph), normalized(vision_bundle)], dim=1 + ) + combined_text = torch.cat( + [normalized(text_graph), normalized(text_bundle)], dim=1 + ) + + result = { + "nodes": n, + "knn": min(args.knn, n - 1), + "eigenpairs": min(args.eigenpairs, n - 2), + "heat_scales": args.heat_scales, + "message_passing_steps": args.message_passing_steps, + "graph_signature_retrieval": arbitrary_target_retrieval( + vision_graph, text_graph, targets + ), + "bundle_signature_retrieval": arbitrary_target_retrieval( + vision_bundle, text_bundle, targets + ), + "combined_signature_retrieval": arbitrary_target_retrieval( + combined_vision, combined_text, targets + ), + "evaluation_note": ( + "Each signature is computed independently from one modality. " + "Ground truth is loaded only after graph construction to score " + "the hidden permutation." + ), + } + Path(args.output).parent.mkdir(parents=True, exist_ok=True) + write_json(args.output, result) + print(json.dumps(result, indent=2)) + + +if __name__ == "__main__": + main() diff --git a/worldalign/vg_prepare.py b/worldalign/vg_prepare.py new file mode 100644 index 0000000..7d553a8 --- /dev/null +++ b/worldalign/vg_prepare.py @@ -0,0 +1,265 @@ +from __future__ import annotations + +import argparse +import json +from pathlib import Path +import random +from typing import Any, Iterable + +from datasets import Dataset, load_dataset + +from .common import write_json + + +DATASET = "ranjaykrishna/visual_genome" + + +def parse_args() -> argparse.Namespace: + p = argparse.ArgumentParser() + p.add_argument("--output-dir", default="artifacts/vg") + p.add_argument("--cache-dir", default="/tmp/yurenh2-worldalign-vg-hf") + p.add_argument("--nodes", type=int, default=100_000) + p.add_argument("--views", type=int, default=32) + p.add_argument("--min-regions", type=int, default=32) + p.add_argument("--seed", type=int, default=20260728) + p.add_argument( + "--with-extra-tiers", + action="store_true", + help="Also load visible relation/attribute and QA annotations.", + ) + return p.parse_args() + + +def load_config(name: str, cache_dir: str) -> Dataset: + return load_dataset( + DATASET, + name, + split="train", + trust_remote_code=True, + with_image=False, + cache_dir=cache_dir, + ) + + +def by_image_id(dataset: Dataset, field: str) -> dict[int, list[dict[str, Any]]]: + return { + int(image_id): values + for image_id, values in zip(dataset["image_id"], dataset[field]) + } + + +def first_name(obj: dict[str, Any]) -> str: + names = obj.get("names") or [] + return str(names[0]).strip() if names else "" + + +def visible_relation_texts( + image_id: int, + relationships: dict[int, list[dict[str, Any]]], + attributes: dict[int, list[dict[str, Any]]], +) -> list[str]: + texts: list[str] = [] + for rel in relationships.get(image_id, []): + subject = first_name(rel["subject"]) + predicate = str(rel.get("predicate") or "").strip() + object_ = first_name(rel["object"]) + if subject and predicate and object_: + texts.append(f"{subject} {predicate} {object_}.") + for obj in attributes.get(image_id, []): + name = first_name(obj) + for attribute in obj.get("attributes") or []: + attribute = str(attribute).strip() + if name and attribute: + texts.append(f"{name} is {attribute}.") + return texts + + +def qa_texts( + image_id: int, questions: dict[int, list[dict[str, Any]]] +) -> list[str]: + texts: list[str] = [] + for item in questions.get(image_id, []): + question = str(item.get("question") or "").strip() + answer = str(item.get("answer") or "").strip() + if question and answer: + texts.append(f"Question: {question} Answer: {answer}") + return texts + + +def fixed_mixture( + groups: list[list[str]], budget: int, rng: random.Random +) -> list[str]: + """Draw a fixed-size, approximately balanced mixture without count leakage.""" + groups = [list(dict.fromkeys(group)) for group in groups if group] + for group in groups: + rng.shuffle(group) + selected: list[str] = [] + while len(selected) < budget and groups: + next_groups: list[list[str]] = [] + for group in groups: + if group and len(selected) < budget: + selected.append(group.pop()) + if group: + next_groups.append(group) + groups = next_groups + if not selected: + raise ValueError("Cannot construct a text bundle from empty groups") + seed_values = selected.copy() + while len(selected) < budget: + selected.append(rng.choice(seed_values)) + rng.shuffle(selected) + return selected + + +def write_jsonl(path: Path, records: Iterable[dict[str, Any]]) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + with path.open("w", encoding="utf-8") as handle: + for record in records: + handle.write(json.dumps(record, ensure_ascii=False) + "\n") + + +def main() -> None: + args = parse_args() + if args.views < 1: + raise ValueError("--views must be positive") + output_dir = Path(args.output_dir) + output_dir.mkdir(parents=True, exist_ok=True) + rng = random.Random(args.seed) + + region_ds = load_config("region_descriptions_v1.2.0", args.cache_dir) + relation_lookup: dict[int, list[dict[str, Any]]] = {} + attribute_lookup: dict[int, list[dict[str, Any]]] = {} + question_lookup: dict[int, list[dict[str, Any]]] = {} + if args.with_extra_tiers: + relation_lookup = by_image_id( + load_config("relationships_v1.2.0", args.cache_dir), + "relationships", + ) + attribute_lookup = by_image_id( + load_config("attributes_v1.2.0", args.cache_dir), + "attributes", + ) + question_lookup = by_image_id( + load_config("question_answers_v1.2.0", args.cache_dir), + "qas", + ) + + eligible = [ + idx + for idx, regions in enumerate(region_ds["regions"]) + if len(regions) >= args.min_regions + ] + rng.shuffle(eligible) + selected = eligible[: min(args.nodes, len(eligible))] + + # The two orders and all opaque identifiers are independent. Source IDs + # appear only in the vision preprocessing file and private truth file. + vision_order = selected.copy() + text_order = selected.copy() + rng.shuffle(vision_order) + rng.shuffle(text_order) + vision_opaque = {idx: f"v{rank:07d}" for rank, idx in enumerate(vision_order)} + text_opaque = {idx: f"t{rank:07d}" for rank, idx in enumerate(text_order)} + + vision_records: list[dict[str, Any]] = [] + text_records: list[dict[str, Any]] = [] + truth: list[dict[str, Any]] = [] + + for idx in selected: + row = region_ds[int(idx)] + regions = list(row["regions"]) + # Fix the view count and use a deterministic per-node sample. This + # removes region-count and ordering side channels. + node_rng = random.Random(args.seed ^ int(row["image_id"])) + node_rng.shuffle(regions) + regions = regions[: args.views] + + visual_regions = [ + { + "x": int(region["x"]), + "y": int(region["y"]), + "width": int(region["width"]), + "height": int(region["height"]), + } + for region in regions + ] + phrases = [str(region["phrase"]).strip() for region in regions] + node_rng.shuffle(phrases) + image_id = int(row["image_id"]) + v_id = vision_opaque[idx] + t_id = text_opaque[idx] + vision_records.append( + { + "node_id": v_id, + "source_image_id": image_id, + "url": row["url"], + "width": int(row["width"]), + "height": int(row["height"]), + "regions": visual_regions, + } + ) + text_record: dict[str, Any] = { + "node_id": t_id, + "region_closed": phrases, + } + if args.with_extra_tiers: + relation_text = visible_relation_texts( + image_id, relation_lookup, attribute_lookup + ) + qa_text = qa_texts(image_id, question_lookup) + text_record["visible_relations"] = fixed_mixture( + [phrases, relation_text], args.views, node_rng + ) + text_record["qa_expanded"] = fixed_mixture( + [phrases, relation_text, qa_text], args.views, node_rng + ) + text_records.append(text_record) + truth.append( + { + "vision_node_id": v_id, + "text_node_id": t_id, + "source_image_id": image_id, + } + ) + + rng.shuffle(vision_records) + rng.shuffle(text_records) + rng.shuffle(truth) + write_jsonl(output_dir / "vision_nodes.private.jsonl", vision_records) + write_jsonl(output_dir / "text_nodes.jsonl", text_records) + write_jsonl(output_dir / "ground_truth.private.jsonl", truth) + + nested_sizes = [ + size for size in (5_000, 20_000, 50_000, 100_000) if size <= len(selected) + ] + if len(selected) not in nested_sizes: + nested_sizes.append(len(selected)) + manifest = { + "dataset": DATASET, + "seed": args.seed, + "selected_nodes": len(selected), + "eligible_nodes": len(eligible), + "views_per_node": args.views, + "minimum_available_regions": args.min_regions, + "extra_closure_tiers": args.with_extra_tiers, + "nested_scaling_sizes": sorted(set(nested_sizes)), + "training_files": { + "vision": "vision_nodes.private.jsonl", + "text": "text_nodes.jsonl", + }, + "evaluation_only": "ground_truth.private.jsonl", + "leakage_note": ( + "The source image ID and URL are required only by the vision " + "feature preprocessor. Alignment training must consume extracted " + "features with source metadata removed." + ), + } + write_json(output_dir / "manifest.json", manifest) + print( + f"Wrote {len(selected):,} hidden-pair nodes to {output_dir}; " + f"{len(eligible):,} scenes met the region threshold." + ) + + +if __name__ == "__main__": + main() diff --git a/worldalign/vg_view_probe.py b/worldalign/vg_view_probe.py new file mode 100644 index 0000000..322644f --- /dev/null +++ b/worldalign/vg_view_probe.py @@ -0,0 +1,292 @@ +"""View-level shared-structure probe for Visual Genome nodes. + +The node-level manifold gate shows that pooled node states expose only a +low-dimensional coarse shared subspace, which cannot identify individual +assignments. This probe measures whether additional cross-modal signal +exists inside nodes, at the level of the sixteen region views, where both +modalities observe the same sixteen regions of the same scene. + +The per-node view permutation is replayed from the preparation seed and the +private image IDs, verified exactly against the released text bundles, and +used only for evaluation. No cross-modal map is trained. + +Three quantities are reported: + +1. within-node relational alignment: correlation between the visual and the + truth-aligned text view-relation fields, against a shuffled-view control; +2. coarse-frame view matching: population-context profiles make single + views cross-modally comparable through the node-level frame; assignment + accuracy is chance-normalized (16 candidates per node); +3. a within-node relational-only floor with no population context. +""" + +from __future__ import annotations + +import argparse +import json +import random +from pathlib import Path + +import torch +import torch.nn.functional as F +from scipy.optimize import linear_sum_assignment + +from .common import write_json +from .vg_prepare import load_config + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--vg-dir", default="artifacts/vg_5k") + parser.add_argument("--cache-dir", default="/tmp/yurenh2-worldalign-vg-hf") + parser.add_argument("--seed", type=int, default=20260728) + parser.add_argument("--views", type=int, default=16) + parser.add_argument( + "--context-nodes", + type=int, + default=1000, + help="Node-level frame size for population-context profiles.", + ) + parser.add_argument("--nodes", type=int, default=0, help="0 keeps all nodes.") + parser.add_argument( + "--skip-floor", + action="store_true", + help="Skip the slow within-node relational-only descent floor.", + ) + parser.add_argument( + "--output", default="artifacts/manifold_gate/vg_view_probe.json" + ) + return parser.parse_args() + + +def replay_view_permutations(args: argparse.Namespace) -> dict: + """Recover text-view -> vision-view mappings from the preparation RNG. + + ``random.shuffle`` consumes the generator identically for equal-length + lists, so shuffling index lists replays the region subsample and the + phrase shuffle exactly. Every node is verified against the released + phrase bundle before use. + """ + truth = [ + json.loads(line) + for line in Path(args.vg_dir, "ground_truth.private.jsonl").read_text( + encoding="utf-8" + ).splitlines() + if line.strip() + ] + text_nodes = { + record["node_id"]: record + for line in Path(args.vg_dir, "text_nodes.jsonl").read_text( + encoding="utf-8" + ).splitlines() + if line.strip() + for record in [json.loads(line)] + } + region_ds = load_config("region_descriptions_v1.2.0", args.cache_dir) + image_rows = { + int(image_id): row for row, image_id in enumerate(region_ds["image_id"]) + } + mappings: dict[str, dict] = {} + mismatched = 0 + for record in truth: + image_id = int(record["source_image_id"]) + row = region_ds[image_rows[image_id]] + regions = list(row["regions"]) + node_rng = random.Random(args.seed ^ image_id) + order = list(range(len(regions))) + node_rng.shuffle(order) + selected = order[: args.views] + phrase_perm = list(range(args.views)) + node_rng.shuffle(phrase_perm) + phrases_view_order = [ + str(regions[index]["phrase"]).strip() for index in selected + ] + reconstructed = [phrases_view_order[j] for j in phrase_perm] + released = text_nodes[record["text_node_id"]]["region_closed"] + if reconstructed != released: + mismatched += 1 + continue + mappings[record["vision_node_id"]] = { + "text_node_id": record["text_node_id"], + # text view j describes vision view phrase_perm[j] + "text_to_vision": phrase_perm, + } + if mismatched: + print(json.dumps({"replay_mismatched_nodes": mismatched})) + return mappings + + +def offdiag(matrix: torch.Tensor) -> torch.Tensor: + mask = ~torch.eye(len(matrix), dtype=torch.bool) + return matrix[mask] + + +def spearman(x: torch.Tensor, y: torch.Tensor) -> float: + ranks = torch.stack( + [x.argsort().argsort().double(), y.argsort().argsort().double()] + ) + return float(torch.corrcoef(ranks)[0, 1]) + + +def main() -> None: + args = parse_args() + mappings = replay_view_permutations(args) + + vision = torch.load( + Path(args.vg_dir, "vision_features.pt"), map_location="cpu", weights_only=False + ) + text = torch.load( + Path(args.vg_dir, "text_features.pt"), map_location="cpu", weights_only=False + ) + vision_index = {node: i for i, node in enumerate(vision["node_ids"])} + text_index = {node: i for i, node in enumerate(text["node_ids"])} + visual_views = F.normalize(vision["region_features"].double(), dim=-1) + text_views = F.normalize(text["region_features"].double(), dim=-1) + + node_ids = sorted(mappings) + if args.nodes: + node_ids = node_ids[: args.nodes] + + # Node-level coarse frame: view-mean states over a fixed context set, + # truth-aligned so both profile axes index the same underlying scenes. + # The frame is evaluation-side scaffolding for an upper bound; a blind + # system would substitute its recovered coarse alignment here. + context_ids = node_ids[: args.context_nodes] + visual_frame = F.normalize( + torch.stack( + [visual_views[vision_index[node]].mean(0) for node in context_ids] + ), + dim=-1, + ) + text_frame = F.normalize( + torch.stack( + [ + text_views[text_index[mappings[node]["text_node_id"]]].mean(0) + for node in context_ids + ] + ), + dim=-1, + ) + + generator = torch.Generator().manual_seed(args.seed) + relation_matched: list[float] = [] + relation_shuffled: list[float] = [] + profile_hits = 0 + profile_top3 = 0 + profile_total = 0 + floor_hits = 0 + floor_total = 0 + matched_profile_corr: list[float] = [] + mismatched_profile_corr: list[float] = [] + + for node in node_ids: + mapping = mappings[node] + v = visual_views[vision_index[node]] + t = text_views[text_index[mapping["text_node_id"]]] + text_to_vision = torch.tensor(mapping["text_to_vision"]) + vision_to_text = torch.empty_like(text_to_vision) + vision_to_text[text_to_vision] = torch.arange(len(text_to_vision)) + + relation_visual = v @ v.T + relation_text = t @ t.T + aligned = relation_text[vision_to_text][:, vision_to_text] + relation_matched.append(spearman(offdiag(relation_visual), offdiag(aligned))) + shuffle = torch.randperm(len(v), generator=generator) + shuffled = relation_text[shuffle][:, shuffle] + relation_shuffled.append(spearman(offdiag(relation_visual), offdiag(shuffled))) + + # Population-context profiles through the coarse frame. + profile_visual = v @ visual_frame.T + profile_text = t @ text_frame.T + rank_visual = profile_visual.argsort(-1).argsort(-1).double() + rank_text = profile_text.argsort(-1).argsort(-1).double() + rank_visual = rank_visual - rank_visual.mean(-1, keepdim=True) + rank_text = rank_text - rank_text.mean(-1, keepdim=True) + rank_visual = F.normalize(rank_visual, dim=-1) + rank_text = F.normalize(rank_text, dim=-1) + similarity = rank_visual @ rank_text.T # [vision view, text view] + rows, cols = linear_sum_assignment(-similarity.numpy()) + truth_cols = vision_to_text.numpy() + profile_hits += int((cols == truth_cols).sum()) + ranks = ( + similarity + >= similarity.gather(1, vision_to_text[:, None]) + ).sum(-1) + profile_top3 += int((ranks <= 3).sum()) + profile_total += len(rows) + matched_profile_corr.extend( + similarity.gather(1, vision_to_text[:, None]).squeeze(1).tolist() + ) + mismatched_profile_corr.extend( + similarity[~torch.eye(len(v), dtype=torch.bool)][:64].tolist() + ) + + # Relational-only floor: 2-swap descent on within-node fields from a + # random start, no population context. + if args.skip_floor: + continue + state = torch.randperm(len(v), generator=generator) + for _ in range(200): + improved = False + current = relation_text[state][:, state] + base = float(((relation_visual - current) ** 2).mean()) + for p in range(len(v)): + for q in range(p + 1, len(v)): + trial = state.clone() + trial[[p, q]] = trial[[q, p]] + candidate = relation_text[trial][:, trial] + if float(((relation_visual - candidate) ** 2).mean()) < base - 1e-12: + state = trial + improved = True + break + if improved: + break + if not improved: + break + floor_hits += int((state == vision_to_text).sum()) + floor_total += len(v) + + matched = torch.tensor(relation_matched) + shuffled = torch.tensor(relation_shuffled) + matched_corr = torch.tensor(matched_profile_corr) + mismatched_corr = torch.tensor(mismatched_profile_corr) + report = { + "protocol": ( + "View permutations are replayed from the preparation seed and " + "verified against released phrase bundles; they are used only " + "for evaluation. The coarse frame uses truth-aligned node " + "states, so profile matching is an upper bound on what a " + "recovered node alignment would enable." + ), + "nodes": len(node_ids), + "verified_fraction": len(mappings) / 5000.0, + "within_node_relation": { + "matched_spearman_mean": float(matched.mean()), + "matched_spearman_std": float(matched.std()), + "shuffled_spearman_mean": float(shuffled.mean()), + "shuffled_spearman_std": float(shuffled.std()), + "gap_z_across_nodes": float( + (matched.mean() - shuffled.mean()) + / (matched - shuffled).std().clamp_min(1e-12) + * len(matched) ** 0.5 + ), + }, + "coarse_frame_view_matching": { + "context_nodes": len(context_ids), + "accuracy": profile_hits / max(profile_total, 1), + "top3_accuracy": profile_top3 / max(profile_total, 1), + "chance": 1.0 / args.views, + "matched_profile_corr_mean": float(matched_corr.mean()), + "mismatched_profile_corr_mean": float(mismatched_corr.mean()), + }, + "within_node_relational_floor": { + "accuracy": floor_hits / max(floor_total, 1), + "chance": 1.0 / args.views, + }, + } + write_json(args.output, report) + print(json.dumps(report, indent=2)) + + +if __name__ == "__main__": + main() diff --git a/worldalign/view_anchor_affinity.py b/worldalign/view_anchor_affinity.py new file mode 100644 index 0000000..3c921bb --- /dev/null +++ b/worldalign/view_anchor_affinity.py @@ -0,0 +1,271 @@ +"""C5 cascade step: anchored node affinity from view-level matching. + +For every node pair the sixteen-view anchor similarity matrix is built +from per-view features and frozen world maps, solved by the Hungarian +method, and the optimal matching value becomes the node affinity. The +anchor-optimal view assignment can also score the within-node relational +agreement at that fixed assignment -- fixed-sigma evaluation avoids the +free-permutation overfit that killed the min-fit affinity. + +Hidden pairs score retrieval and seed precision only; the affinity uses +per-node features and frozen maps, nothing cross-modal that is learned. +""" + +from __future__ import annotations + +import argparse +import json + +import numpy as np +import torch +import torch.nn.functional as F +from scipy.optimize import linear_sum_assignment + +from .common import write_json + +CHANNELS = ("color", "size", "light", "horizontal", "vertical") + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--vg-dir", default="artifacts/vg_5k") + parser.add_argument( + "--features", default="artifacts/manifold_gate/view_anchor_features.pt" + ) + parser.add_argument( + "--structure", default="artifacts/manifold_gate/attn_text_full.pt" + ) + parser.add_argument( + "--vision-structure", default="artifacts/manifold_gate/attn_vision_full.pt" + ) + parser.add_argument("--samples", type=int, default=512) + parser.add_argument("--subset-seeds", default="0,1,2") + parser.add_argument("--relational-weight", type=float, default=0.5) + parser.add_argument("--chunk", type=int, default=64) + parser.add_argument( + "--output", default="artifacts/manifold_gate/view_anchor_affinity.json" + ) + return parser.parse_args() + + +def rank_z_rows(values: torch.Tensor) -> torch.Tensor: + order = values.argsort(-1).argsort(-1).double() + order = order - order.mean(-1, keepdim=True) + std = order.std(-1, keepdim=True).clamp_min(1e-9) + return order / std + + +def standardized_fields(views: torch.Tensor) -> torch.Tensor: + views = F.normalize(views.double(), dim=-1) + relation = views @ views.transpose(-2, -1) + size = relation.shape[-1] + mask = ~torch.eye(size, dtype=torch.bool) + values = relation[..., mask] + mean = values.mean(-1, keepdim=True) + std = values.std(-1, keepdim=True).clamp_min(1e-9) + standardized = (relation - mean[..., None]) / std[..., None] + return standardized.masked_fill(~mask, 0.0) + + +def evaluate_subset( + data: dict, subset: torch.Tensor, args: argparse.Namespace +) -> dict: + n = len(subset) + views = data["vision_color"].shape[1] + vision_color = data["vision_color"][subset] + text_color = data["text_color"][subset] + text_informative = data["text_informative"][subset] # [n, C, V] bool + vision_scalar = data["vision_scalar"][subset] # [n, 3, V] rank-z + text_scalar = data["text_scalar"][subset] # [n, 3, V] rank-z + text_scalar_mask = data["text_scalar_mask"][subset] # [n, 3, V] + text_fields = data["text_fields"][subset] + vision_fields = data["vision_fields"][subset] + + anchor_cost = torch.zeros(n, n) + combined_cost = torch.zeros(n, n) + for start in range(0, n, args.chunk): + stop = min(start + args.chunk, n) + block = slice(start, stop) + color = torch.einsum( + "avc,btc->abvt", vision_color[block].double(), text_color.double() + ) + color = color * text_informative[None, :, 0, None, :] + channels = [color / max(float(color[color != 0].std()), 1e-9) if (color != 0).any() else color] + for c in range(4): + vision_values = vision_scalar[block, c] # [a, V] + text_values = text_scalar[:, c] # [n, V] + difference = -( + vision_values[:, None, :, None] - text_values[None, :, None, :] + ).abs() + difference = difference * text_scalar_mask[:, c][None, :, None, :] + scale = float(difference[difference != 0].std()) if (difference != 0).any() else 1.0 + channels.append(difference / max(scale, 1e-9)) + stacked = torch.stack(channels) # [C, a, b, V, V] + present = torch.stack( + [ + text_informative[:, 0].any(-1), + text_scalar_mask[:, 0].any(-1), + text_scalar_mask[:, 1].any(-1), + text_scalar_mask[:, 2].any(-1), + text_scalar_mask[:, 3].any(-1), + ] + ).double() # [C, b] + weight = present / present.sum(0, keepdim=True).clamp_min(1.0) + anchor_matrix = torch.einsum("cabvt,cb->abvt", stacked, weight) + for a in range(stop - start): + for b in range(n): + matrix = anchor_matrix[a, b].numpy() + rows, cols = linear_sum_assignment(-matrix) + value = float(matrix[rows, cols].mean()) + anchor_cost[start + a, b] = -value + if args.relational_weight: + aligned = text_fields[b][cols][:, cols] + relational = float( + ((aligned - vision_fields[start + a]) ** 2)[ + ~torch.eye(views, dtype=torch.bool) + ].mean() + ) + combined_cost[start + a, b] = ( + -value + args.relational_weight * relational + ) + truth = torch.arange(n) + + def metrics(cost: torch.Tensor) -> dict: + centered = cost - cost.mean(0, keepdim=True) + ranks = (centered <= centered.gather(1, truth[:, None])).sum(-1) + rows, cols = linear_sum_assignment(cost.numpy()) + forward = centered.argmin(-1) + backward = centered.argmin(0) + mutual = backward[forward] == truth + sorted_cost = centered.sort(-1).values + margin = sorted_cost[:, 1] - sorted_cost[:, 0] + top = margin.argsort(descending=True)[: max(1, n // 10)] + true_costs = cost.diagonal() + permuted = [] + rng = np.random.default_rng(0) + for _ in range(200): + permutation = torch.from_numpy(rng.permutation(n)) + permuted.append(float(cost[truth, permutation].mean())) + permuted = torch.tensor(permuted) + return { + "r@1": float((ranks <= 1).double().mean()), + "r@5": float((ranks <= 5).double().mean()), + "r@10": float((ranks <= 10).double().mean()), + "median_rank": float(ranks.double().median()), + "true_z": float( + (permuted.mean() - true_costs.mean()) / permuted.std().clamp_min(1e-12) + ), + "hungarian_accuracy": float( + (torch.from_numpy(cols) == truth).double().mean() + ), + "mutual_nn_count": int(mutual.sum()), + "mutual_nn_precision": float( + (forward[mutual] == truth[mutual]).double().mean() + ) + if mutual.any() + else None, + "top_margin_decile_precision": float( + (forward[top] == truth[top]).double().mean() + ), + } + + result = {"anchors_only": metrics(anchor_cost)} + if args.relational_weight: + result["anchors_plus_relational"] = metrics(combined_cost) + return result + + +def main() -> None: + args = parse_args() + state = torch.load(args.features, map_location="cpu", weights_only=False) + truth_pairs = [ + json.loads(line) + for line in open(f"{args.vg_dir}/ground_truth.private.jsonl", encoding="utf-8") + if line.strip() + ] + text_structure = torch.load(args.structure, map_location="cpu", weights_only=False) + vision_structure = torch.load( + args.vision_structure, map_location="cpu", weights_only=False + ) + text_index = {node: i for i, node in enumerate(text_structure["node_ids"])} + vision_index = {node: i for i, node in enumerate(vision_structure["node_ids"])} + + vision_color, text_color, text_informative = [], [], [] + vision_scalar, text_scalar, text_scalar_mask = [], [], [] + text_fields, vision_fields = [], [] + for pair in truth_pairs: + vision_features = state["vision_view_features"][pair["vision_node_id"]] + text_features = state["text_view_features"][pair["text_node_id"]] + hue = torch.from_numpy(vision_features["hue"]) + color = torch.from_numpy(text_features["color"]) + vision_color.append(F.normalize(hue, dim=-1)) + text_color.append(F.normalize(color, dim=-1)) + text_informative.append((color.sum(-1) > 0)[None, :]) + vision_scalar.append( + torch.stack( + [ + rank_z_rows(torch.from_numpy(vision_features["area"])[None])[0], + rank_z_rows(torch.from_numpy(vision_features["luminance"])[None])[0], + rank_z_rows(torch.from_numpy(vision_features["x"])[None])[0], + rank_z_rows(torch.from_numpy(vision_features["y"])[None])[0], + ] + ) + ) + scalars, masks = [], [] + for key in ("size", "light", "horizontal", "vertical"): + values = torch.from_numpy(text_features[key]) + scalars.append(rank_z_rows(values[None])[0]) + masks.append(values != 0) + text_scalar.append(torch.stack(scalars)) + text_scalar_mask.append(torch.stack(masks)) + text_fields.append( + standardized_fields( + torch.as_tensor( + text_structure["context_states"][ + text_index[pair["text_node_id"]] + ] + )[None] + )[0] + ) + vision_fields.append( + standardized_fields( + torch.as_tensor( + vision_structure["context_states"][ + vision_index[pair["vision_node_id"]] + ] + )[None] + )[0] + ) + data = { + "vision_color": torch.stack(vision_color), + "text_color": torch.stack(text_color), + "text_informative": torch.stack(text_informative), + "vision_scalar": torch.stack(vision_scalar), + "text_scalar": torch.stack(text_scalar), + "text_scalar_mask": torch.stack(text_scalar_mask), + "text_fields": torch.stack(text_fields), + "vision_fields": torch.stack(vision_fields), + } + report: dict = { + "protocol": ( + "Node affinity is the Hungarian value of the per-pair view " + "anchor matrix (plus optional fixed-assignment relational " + "agreement). Hidden pairs score retrieval and seeds only." + ), + "samples": args.samples, + "subsets": [], + } + for seed in (int(s) for s in args.subset_seeds.split(",")): + generator = torch.Generator().manual_seed(seed) + subset = torch.randperm(len(truth_pairs), generator=generator)[ + : args.samples + ] + result = {"subset_seed": seed, **evaluate_subset(data, subset, args)} + report["subsets"].append(result) + print(json.dumps(result)) + write_json(args.output, report) + print(f"Wrote {args.output}") + + +if __name__ == "__main__": + main() diff --git a/worldalign/view_anchor_battery.py b/worldalign/view_anchor_battery.py new file mode 100644 index 0000000..f039c63 --- /dev/null +++ b/worldalign/view_anchor_battery.py @@ -0,0 +1,365 @@ +"""C5 second wave: fine-grained anchors at view level. + +The node-level anchor wave failed because scene-level attributes are +coarse-redundant with the shared subspace. Region-level attributes are +not: a phrase names the color, relative size, lighting, or position of +one region, and the region's own pixels and box realize them. With +sixteen candidates per node, weak per-view anchors carry usable bits. + +Per-view anchors, computed independently per modality: + +- text, per phrase: color-lexicon histogram, size score, light score, + horizontal/vertical position scores, numeral content; +- vision, per region: crop hue-band histogram, crop luminance, box area + fraction, box center coordinates. + +Scalar channels are rank-standardized within the node (sixteen values), +so they compare relative attributes inside one scene; lexicon-to-physics +maps are frozen world knowledge. Channels contribute only where the +phrase carries the attribute. Replayed view truth scores matching; +nothing here uses node pairs or the coarse frame. +""" + +from __future__ import annotations + +import argparse +import json +import re +from concurrent.futures import ThreadPoolExecutor +from pathlib import Path + +import numpy as np +import torch +from PIL import Image +from scipy.optimize import linear_sum_assignment + +from .anchor_battery import ( + COLOR_BANDS, + LIGHT_WORDS, + NUMBER_WORDS, + SIZE_WORDS, + read_jsonl, +) +from .common import write_json +from .vg_view_probe import replay_view_permutations + +HORIZONTAL_WORDS = {"left": -1.0, "right": 1.0} +VERTICAL_WORDS = {"top": -1.0, "upper": -1.0, "above": -1.0, "bottom": 1.0, "lower": 1.0, "below": 1.0} +# Twelve color classes: eight hue bands plus achromatic and brown classes +# with pixel rules on value and saturation. Frozen world knowledge. +COLOR_CLASSES = list(COLOR_BANDS) + ["white", "black", "gray", "brown"] +COLOR_SYNONYMS = {"grey": "gray", "tan": "brown", "beige": "brown", "golden": "yellow", "gold": "yellow", "silver": "gray"} + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--vg-dir", default="artifacts/vg_5k") + parser.add_argument("--cache-dir", default="/tmp/yurenh2-worldalign-vg-hf") + parser.add_argument("--seed", type=int, default=20260728) + parser.add_argument("--views", type=int, default=16) + parser.add_argument("--image-cache", default="/tmp/yurenh2-worldalign-vg-images") + parser.add_argument("--workers", type=int, default=16) + parser.add_argument("--structure", default="artifacts/manifold_gate/attn_text_full.pt") + parser.add_argument( + "--vision-structure", default="artifacts/manifold_gate/attn_vision_full.pt" + ) + parser.add_argument("--nodes", type=int, default=0) + parser.add_argument( + "--output", default="artifacts/manifold_gate/view_anchor_battery.json" + ) + parser.add_argument( + "--anchors-output", default="artifacts/manifold_gate/view_anchor_maps.pt" + ) + return parser.parse_args() + + +def phrase_features(phrase: str) -> dict: + tokens = [ + COLOR_SYNONYMS.get(token, token) + for token in re.findall(r"[a-z]+|\d+", phrase.lower()) + ] + color = np.zeros(len(COLOR_CLASSES)) + for index, name in enumerate(COLOR_CLASSES): + color[index] = sum(token == name for token in tokens) + numeral = 0.0 + for token in tokens: + if token.isdigit() and int(token) <= 50: + numeral += int(token) + elif token in NUMBER_WORDS: + numeral += NUMBER_WORDS[token] + return { + "color": color, + "size": float(sum(SIZE_WORDS.get(token, 0.0) for token in tokens)), + "light": float(sum(LIGHT_WORDS.get(token, 0.0) for token in tokens)), + "horizontal": float(sum(HORIZONTAL_WORDS.get(token, 0.0) for token in tokens)), + "vertical": float(sum(VERTICAL_WORDS.get(token, 0.0) for token in tokens)), + "numeral": numeral, + } + + +def region_features(record: dict, args: argparse.Namespace) -> dict: + path = Path(args.image_cache, f"{record['source_image_id']}.jpg") + with Image.open(path) as image: + hsv = np.asarray(image.convert("HSV"), dtype=np.float64) + height, width = hsv.shape[:2] + hues, lums, areas, xs, ys = [], [], [], [], [] + for region in record["regions"]: + x0 = max(0, min(int(region["x"]), width - 1)) + y0 = max(0, min(int(region["y"]), height - 1)) + x1 = max(x0 + 1, min(x0 + int(region["width"]), width)) + y1 = max(y0 + 1, min(y0 + int(region["height"]), height)) + crop = hsv[y0:y1, x0:x1] + hue = crop[..., 0] * (360.0 / 255.0) + saturation = crop[..., 1] / 255.0 + value = crop[..., 2] / 255.0 + white = (value > 0.75) & (saturation < 0.25) + black = value < 0.25 + gray = (~white) & (~black) & (saturation < 0.25) + brown_hue = np.minimum(np.abs(hue - 30.0), 360.0 - np.abs(hue - 30.0)) < 25.0 + brown = brown_hue & (saturation >= 0.25) & (value >= 0.2) & (value < 0.55) + chromatic = (saturation >= 0.25) & (value >= 0.2) & (~brown) + histogram = np.zeros(len(COLOR_CLASSES)) + for index, center in enumerate(COLOR_BANDS.values()): + distance = np.minimum(np.abs(hue - center), 360.0 - np.abs(hue - center)) + histogram[index] = float(((distance < 25.0) & chromatic).mean()) + histogram[len(COLOR_BANDS) + 0] = float(white.mean()) + histogram[len(COLOR_BANDS) + 1] = float(black.mean()) + histogram[len(COLOR_BANDS) + 2] = float(gray.mean()) + histogram[len(COLOR_BANDS) + 3] = float(brown.mean()) + hues.append(histogram) + lums.append(float(value.mean())) + areas.append((x1 - x0) * (y1 - y0) / (width * height)) + xs.append((x0 + x1) / 2.0 / width) + ys.append((y0 + y1) / 2.0 / height) + return { + "hue": np.stack(hues), + "luminance": np.array(lums), + "area": np.array(areas), + "x": np.array(xs), + "y": np.array(ys), + } + + +def rank_z(values: np.ndarray) -> np.ndarray: + order = values.argsort().argsort().astype(np.float64) + std = order.std() + return (order - order.mean()) / (std if std > 1e-9 else 1.0) + + +def node_anchor_matrix( + text: list[dict], vision: dict +) -> tuple[np.ndarray, np.ndarray]: + """[V, V] anchor similarity (vision rows, text columns) and coverage.""" + views = len(text) + channels = [] + coverage = np.zeros(5) + + text_color = np.stack([p["color"] for p in text]) + informative = text_color.sum(1) > 0 + if informative.any(): + tn = text_color / np.linalg.norm(text_color, axis=1, keepdims=True).clip(min=1e-9) + vn = vision["hue"] / np.linalg.norm(vision["hue"], axis=1, keepdims=True).clip(min=1e-9) + color = vn @ tn.T + color[:, ~informative] = 0.0 + flat = color[:, informative] + scale = flat.std() + channels.append(color / (scale if scale > 1e-9 else 1.0)) + coverage[0] = informative.mean() + + for slot, (text_key, vision_values, sign) in enumerate( + ( + ("size", rank_z(vision["area"]), 1.0), + ("light", rank_z(vision["luminance"]), 1.0), + ("horizontal", rank_z(vision["x"]), 1.0), + ("vertical", rank_z(vision["y"]), 1.0), + ), + start=1, + ): + scores = np.array([p[text_key] for p in text]) + informative = scores != 0 + if informative.sum() < 2: + continue + text_rank = rank_z(scores) + similarity = -np.abs(vision_values[:, None] - sign * text_rank[None, :]) + similarity[:, ~informative] = 0.0 + flat = similarity[:, informative] + scale = flat.std() + channels.append(similarity / (scale if scale > 1e-9 else 1.0)) + coverage[slot] = informative.mean() + if not channels: + return np.zeros((views, views)), coverage + return np.mean(channels, axis=0), coverage + + +def structure_signature_matrix( + text_views: torch.Tensor, visual_views: torch.Tensor +) -> np.ndarray: + """Frame-free structural channel: sorted within-node relation rows.""" + def signatures(views: torch.Tensor) -> torch.Tensor: + views = torch.nn.functional.normalize(views.double(), dim=-1) + relation = views @ views.T + size = len(relation) + mask = ~torch.eye(size, dtype=torch.bool) + rows = relation.masked_select(mask).reshape(size, size - 1) + rows = (rows - rows.mean()) / rows.std().clamp_min(1e-9) + return rows.sort(-1).values + + a = signatures(visual_views) + b = signatures(text_views) + cost = ((a[:, None, :] - b[None, :, :]) ** 2).sum(-1) + similarity = -cost + return ((similarity - similarity.mean()) / similarity.std().clamp_min(1e-9)).numpy() + + +def main() -> None: + args = parse_args() + mappings = replay_view_permutations(args) + text_nodes = { + record["node_id"]: record + for record in read_jsonl(Path(args.vg_dir, "text_nodes.jsonl")) + } + vision_nodes = { + record["node_id"]: record + for record in read_jsonl(Path(args.vg_dir, "vision_nodes.private.jsonl")) + } + text_structure = torch.load(args.structure, map_location="cpu", weights_only=False) + vision_structure = torch.load( + args.vision_structure, map_location="cpu", weights_only=False + ) + text_index = {node: i for i, node in enumerate(text_structure["node_ids"])} + vision_index = {node: i for i, node in enumerate(vision_structure["node_ids"])} + + node_ids = sorted(mappings) + if args.nodes: + node_ids = node_ids[: args.nodes] + + vision_records = [vision_nodes[node] for node in node_ids] + with ThreadPoolExecutor(max_workers=args.workers) as pool: + vision_results = list( + pool.map(lambda record: region_features(record, args), vision_records) + ) + + generator = np.random.default_rng(args.seed) + accuracies = {"anchors": [], "structure": [], "anchors_plus_structure": []} + matched_scores, shuffled_scores = [], [] + coverage_totals = np.zeros(5) + anchor_maps = {} + informative_view_hits: list[bool] = [] + seed_pool: list[tuple[float, bool]] = [] + text_view_features: dict[str, dict] = {} + vision_view_features: dict[str, dict] = {} + + for node, vision_features_one in zip(node_ids, vision_results): + mapping = mappings[node] + phrases = text_nodes[mapping["text_node_id"]]["region_closed"] + text_features_list = [phrase_features(p) for p in phrases] + anchors, coverage = node_anchor_matrix(text_features_list, vision_features_one) + coverage_totals += coverage + text_to_vision = np.array(mapping["text_to_vision"]) + vision_to_text = np.empty_like(text_to_vision) + vision_to_text[text_to_vision] = np.arange(len(text_to_vision)) + + structure = structure_signature_matrix( + torch.as_tensor( + text_structure["context_states"][text_index[mapping["text_node_id"]]] + ), + torch.as_tensor(vision_structure["context_states"][vision_index[node]]), + ) + combined = anchors + structure + + informative_columns = np.abs(anchors).sum(0) > 1e-9 + for name, matrix in ( + ("anchors", anchors), + ("structure", structure), + ("anchors_plus_structure", combined), + ): + rows, cols = linear_sum_assignment(-matrix) + accuracies[name].append(float((cols == vision_to_text).mean())) + if name == "anchors" and informative_columns.any(): + assigned_text = cols # vision row i -> text column + correct = assigned_text == vision_to_text + text_informative_hit = informative_columns[assigned_text] + informative_view_hits.extend( + correct[text_informative_hit].tolist() + ) + sorted_scores = np.sort(matrix, axis=1)[:, ::-1] + margins = sorted_scores[:, 0] - sorted_scores[:, 1] + for row in range(len(matrix)): + if informative_columns[assigned_text[row]]: + seed_pool.append( + (float(margins[row]), bool(correct[row])) + ) + matched = anchors[np.arange(len(anchors)), vision_to_text] + shuffle = generator.permutation(len(anchors)) + matched_scores.append(float(matched.mean())) + shuffled_scores.append( + float(anchors[np.arange(len(anchors)), shuffle].mean()) + ) + anchor_maps[node] = torch.from_numpy(anchors).float() + text_view_features[mapping["text_node_id"]] = { + "color": np.stack([p["color"] for p in text_features_list]), + "size": np.array([p["size"] for p in text_features_list]), + "light": np.array([p["light"] for p in text_features_list]), + "horizontal": np.array([p["horizontal"] for p in text_features_list]), + "vertical": np.array([p["vertical"] for p in text_features_list]), + "numeral": np.array([p["numeral"] for p in text_features_list]), + } + vision_view_features[node] = vision_features_one + + matched_arr = np.array(matched_scores) + shuffled_arr = np.array(shuffled_scores) + report = { + "protocol": ( + "Per-view anchors and frame-free structural signatures; " + "replayed view truth scores matching only. No node pairs, no " + "coarse frame." + ), + "nodes": len(node_ids), + "chance": 1.0 / args.views, + "true_frame_profile_baseline": 0.1448, + "channel_coverage_mean": { + "color": float(coverage_totals[0] / len(node_ids)), + "size": float(coverage_totals[1] / len(node_ids)), + "light": float(coverage_totals[2] / len(node_ids)), + "horizontal": float(coverage_totals[3] / len(node_ids)), + "vertical": float(coverage_totals[4] / len(node_ids)), + }, + "anchor_matched_vs_shuffled_z": float( + (matched_arr - shuffled_arr).mean() + / (matched_arr - shuffled_arr).std().clip(min=1e-12) + * np.sqrt(len(matched_arr)) + ), + "view_matching_accuracy": { + name: float(np.mean(values)) for name, values in accuracies.items() + }, + "informative_view_accuracy": float(np.mean(informative_view_hits)) + if informative_view_hits + else None, + "informative_view_count": len(informative_view_hits), + } + if seed_pool: + seed_pool.sort(key=lambda item: -item[0]) + calibration = {} + for fraction in (0.01, 0.05, 0.10, 0.25): + k = max(1, int(len(seed_pool) * fraction)) + calibration[f"top_{fraction:.0%}_margin"] = { + "count": k, + "precision": float(np.mean([hit for _, hit in seed_pool[:k]])), + } + report["seed_calibration"] = calibration + torch.save( + { + "anchor_maps": anchor_maps, + "text_view_features": text_view_features, + "vision_view_features": vision_view_features, + "color_classes": COLOR_CLASSES, + "order_note": "vision rows, text columns, released view order", + }, + args.anchors_output, + ) + write_json(args.output, report) + print(json.dumps(report, indent=2)) + + +if __name__ == "__main__": + main() diff --git a/worldalign/view_gate.py b/worldalign/view_gate.py new file mode 100644 index 0000000..a09fdca --- /dev/null +++ b/worldalign/view_gate.py @@ -0,0 +1,436 @@ +"""View-level assignment gate on structured Visual Genome node states. + +A node state here is the set of sixteen region views with its internal +standardized relation field, not a pooled vector. Two questions are gated, +per the current experimental gate in LAB_NOTES: + +A. Ordering at view level, given the true coarse frame (evaluation upper + bound): is the true view assignment a local optimum of (1) the + within-node relational energy, (2) the population-context profile + energy, (3) their combination? All 120 transpositions per node are + enumerated exactly. + +B. Blind structured-state node affinity, frame-free: the matching cost + between two nodes' internal relation fields (signature-initialized + Hungarian plus exact batched 2-swap descent) is used as a node-level + linear assignment energy. Retrieval, global Hungarian recovery, and + margin-calibrated precision are compared against the pooled-state and + graph-signature baselines. + +Replayed view truth and node pairs are evaluation-only throughout. +""" + +from __future__ import annotations + +import argparse +import json + +import torch +import torch.nn.functional as F +from scipy.optimize import linear_sum_assignment + +from .common import seed_everything, write_json +from .vg_view_probe import replay_view_permutations + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--vg-dir", default="artifacts/vg_5k") + parser.add_argument("--cache-dir", default="/tmp/yurenh2-worldalign-vg-hf") + parser.add_argument("--seed", type=int, default=20260728) + parser.add_argument("--views", type=int, default=16) + parser.add_argument("--part-a-nodes", type=int, default=5000) + parser.add_argument("--context-nodes", type=int, default=1000) + parser.add_argument("--joint-weights", default="0.5,2.0") + parser.add_argument("--part-b-nodes", type=int, default=512) + parser.add_argument("--part-b-seeds", default="0,1,2") + parser.add_argument("--descent-sweeps", type=int, default=40) + parser.add_argument("--chunk", type=int, default=48) + parser.add_argument("--output", default="artifacts/manifold_gate/view_gate.json") + parser.add_argument( + "--matrices-output", + default="", + help="Optional .pt path stem for saving raw/null cost matrices.", + ) + return parser.parse_args() + + +def standardized_field(views: torch.Tensor) -> torch.Tensor: + """Within-node standardized relation field with zeroed diagonal.""" + relation = views @ views.transpose(-2, -1) + size = relation.shape[-1] + mask = ~torch.eye(size, dtype=torch.bool) + values = relation[..., mask] + mean = values.mean(-1, keepdim=True) + std = values.std(-1, keepdim=True).clamp_min(1e-9) + standardized = (relation - mean[..., None]) / std[..., None] + return standardized.masked_fill(~mask, 0.0) + + +def batched_transposition_delta( + text_fields: torch.Tensor, visual_fields: torch.Tensor +) -> torch.Tensor: + """Exact MSE delta of every transposition for a batch of field pairs. + + Same cancellation as manifold_gate.all_transposition_delta_mse, with a + batch dimension: delta[b] applies to (text_fields[b], visual_fields[b]). + """ + size = text_fields.shape[-1] + count = size * (size - 1) + cross = torch.bmm(text_fields, visual_fields) + self_terms = (text_fields * visual_fields).sum(-1) + corrections = 2.0 * text_fields * visual_fields + total = ( + self_terms[:, :, None] + + self_terms[:, None, :] + - cross + - cross.transpose(-2, -1) + - corrections + ) + delta = (4.0 / count) * total + diagonal = torch.eye(size, dtype=torch.bool) + return delta.masked_fill(diagonal, 0.0) + + +def field_mse(text_fields: torch.Tensor, visual_fields: torch.Tensor) -> torch.Tensor: + size = text_fields.shape[-1] + mask = ~torch.eye(size, dtype=torch.bool) + return ((text_fields - visual_fields) ** 2)[..., mask].mean(-1) + + +def batched_descent( + text_fields: torch.Tensor, + visual_fields: torch.Tensor, + permutations: torch.Tensor, + sweeps: int, +) -> tuple[torch.Tensor, torch.Tensor]: + """Steepest 2-swap descent run in parallel over a batch of field pairs.""" + batch, size = permutations.shape + index = torch.arange(batch) + for _ in range(sweeps): + permuted = text_fields[ + index[:, None, None], + permutations[:, :, None], + permutations[:, None, :], + ] + delta = batched_transposition_delta(permuted, visual_fields) + upper = torch.triu(torch.ones(size, size, dtype=torch.bool), diagonal=1) + masked = delta.masked_fill(~upper, float("inf")) + flat = masked.reshape(batch, -1) + best = flat.argmin(-1) + best_value = flat.gather(-1, best[:, None]).squeeze(-1) + improvable = best_value < -1e-12 + if not improvable.any(): + break + p = best // size + q = best % size + rows = index[improvable] + permutations[rows] = permutations[rows].scatter( + 1, + torch.stack([p[improvable], q[improvable]], -1), + torch.stack( + [ + permutations[rows, q[improvable]], + permutations[rows, p[improvable]], + ], + -1, + ), + ) + final = text_fields[ + index[:, None, None], permutations[:, :, None], permutations[:, None, :] + ] + return field_mse(final, visual_fields), permutations + + +def load_aligned_views(args: argparse.Namespace) -> dict: + mappings = replay_view_permutations(args) + vision = torch.load( + f"{args.vg_dir}/vision_features.pt", map_location="cpu", weights_only=False + ) + text = torch.load( + f"{args.vg_dir}/text_features.pt", map_location="cpu", weights_only=False + ) + vision_index = {node: i for i, node in enumerate(vision["node_ids"])} + text_index = {node: i for i, node in enumerate(text["node_ids"])} + node_ids = sorted(mappings) + visual_views = [] + text_views_aligned = [] + text_views_raw = [] + for node in node_ids: + mapping = mappings[node] + v = F.normalize( + vision["region_features"][vision_index[node]].double(), dim=-1 + ) + t = F.normalize( + text["region_features"][text_index[mapping["text_node_id"]]].double(), + dim=-1, + ) + text_to_vision = torch.tensor(mapping["text_to_vision"]) + vision_to_text = torch.empty_like(text_to_vision) + vision_to_text[text_to_vision] = torch.arange(len(text_to_vision)) + visual_views.append(v) + text_views_aligned.append(t[vision_to_text]) + text_views_raw.append(t) + return { + "node_ids": node_ids, + "visual_views": torch.stack(visual_views), + "text_views_aligned": torch.stack(text_views_aligned), + "text_views_raw": torch.stack(text_views_raw), + } + + +def part_a(data: dict, args: argparse.Namespace) -> dict: + """Ordering gates at the true view assignment, true coarse frame.""" + nodes = min(args.part_a_nodes, len(data["node_ids"])) + visual = data["visual_views"][:nodes] + text = data["text_views_aligned"][:nodes] + visual_fields = standardized_field(visual) + text_fields = standardized_field(text) + size = visual.shape[1] + upper = torch.triu(torch.ones(size, size, dtype=torch.bool), diagonal=1) + + delta_e1 = batched_transposition_delta(text_fields, visual_fields) + e1_improving = (delta_e1[:, upper] < 0).double().mean(-1) + + context = min(args.context_nodes, nodes) + visual_frame = F.normalize(visual[:context].mean(1), dim=-1) + text_frame = F.normalize(text[:context].mean(1), dim=-1) + profile_visual = visual @ visual_frame.T + profile_text = text @ text_frame.T + + def centered_ranks(profiles: torch.Tensor) -> torch.Tensor: + ranks = profiles.argsort(-1).argsort(-1).double() + ranks = ranks - ranks.mean(-1, keepdim=True) + return F.normalize(ranks, dim=-1) + + similarity = torch.bmm( + centered_ranks(profile_visual), centered_ranks(profile_text).transpose(-2, -1) + ) + matched = similarity.diagonal(dim1=-2, dim2=-1) + delta_e2 = ( + matched[:, :, None] + matched[:, None, :] + - similarity - similarity.transpose(-2, -1) + ) + e2_improving = (delta_e2[:, upper] < 0).double().mean(-1) + + report = { + "nodes": nodes, + "context_nodes": context, + "e1_within_node_relational": { + "improving_fraction_mean": float(e1_improving.mean()), + "strict_local_min_nodes": float((e1_improving == 0).double().mean()), + }, + "e2_coarse_frame_profile": { + "improving_fraction_mean": float(e2_improving.mean()), + "strict_local_min_nodes": float((e2_improving == 0).double().mean()), + }, + "joint": {}, + } + for weight in (float(w) for w in args.joint_weights.split(",")): + delta_joint = delta_e1 + weight * delta_e2 + joint_improving = (delta_joint[:, upper] < 0).double().mean(-1) + report["joint"][f"lambda_{weight}"] = { + "improving_fraction_mean": float(joint_improving.mean()), + "strict_local_min_nodes": float((joint_improving == 0).double().mean()), + } + return report + + +def part_b_one_seed( + data: dict, args: argparse.Namespace, subset_seed: int +) -> dict: + """Frame-free structured node affinity on one node subset.""" + generator = torch.Generator().manual_seed(subset_seed) + subset = torch.randperm(len(data["node_ids"]), generator=generator)[ + : args.part_b_nodes + ] + visual_fields = standardized_field(data["visual_views"][subset]) + text_fields = standardized_field(data["text_views_raw"][subset]) + n, size = visual_fields.shape[0], visual_fields.shape[-1] + + columns_without_diag = torch.stack( + [ + torch.tensor([j for j in range(size) if j != i], dtype=torch.long) + for i in range(size) + ] + ) + row_index = torch.arange(size)[:, None] + signature_visual = visual_fields[ + :, row_index, columns_without_diag + ].sort(-1).values + signature_text = text_fields[:, row_index, columns_without_diag].sort(-1).values + visual_norms = (signature_visual**2).sum(-1) + text_norms = (signature_text**2).sum(-1) + + # Structure-scrambled null per text node: off-diagonal values are + # randomly reassigned to positions (marginals preserved, metric + # consistency destroyed). The excess of structured over null matching + # cost removes the free-permutation overfitting capacity that dominates + # raw min-cost at sixteen views. + null_generator = torch.Generator().manual_seed(subset_seed + 1) + upper = torch.triu(torch.ones(size, size, dtype=torch.bool), diagonal=1) + null_text_fields = torch.empty_like(text_fields) + for node in range(n): + values = text_fields[node][upper] + scrambled = values[ + torch.randperm(len(values), generator=null_generator) + ] + field = torch.zeros(size, size, dtype=text_fields.dtype) + field[upper] = scrambled + null_text_fields[node] = field + field.T + + def matched_costs( + target_fields: torch.Tensor, + ) -> tuple[torch.Tensor, torch.Tensor]: + """Descended min-fit cost and zero-capacity signature-Hungarian cost.""" + target_signature = target_fields[:, row_index, columns_without_diag].sort( + -1 + ).values + target_norms = (target_signature**2).sum(-1) + result = torch.zeros(n, n) + signature_result = torch.zeros(n, n) + for start in range(0, n, args.chunk): + stop = min(start + args.chunk, n) + block = signature_visual[start:stop] + init_cost = ( + visual_norms[start:stop, None, :, None] + + target_norms[None, :, None, :] + - 2.0 * torch.einsum("cvs,nws->cnvw", block, target_signature) + ) + chunk_pairs = init_cost.shape[0] * init_cost.shape[1] + permutations = torch.empty(chunk_pairs, size, dtype=torch.long) + flat_cost = init_cost.reshape(chunk_pairs, size, size) + for pair in range(chunk_pairs): + pair_rows, columns = linear_sum_assignment(flat_cost[pair].numpy()) + permutations[pair] = torch.from_numpy(columns) + signature_result.view(-1)[ + (start * n) + pair + ] = float(flat_cost[pair][pair_rows, columns].sum()) + pair_text = ( + target_fields[None, :, :, :] + .expand(stop - start, n, size, size) + .reshape(chunk_pairs, size, size) + ) + pair_visual = ( + visual_fields[start:stop, None, :, :] + .expand(stop - start, n, size, size) + .reshape(chunk_pairs, size, size) + ) + final_cost, _ = batched_descent( + pair_text, pair_visual, permutations, args.descent_sweeps + ) + result[start:stop] = final_cost.reshape(stop - start, n) + return result, signature_result + + raw_cost, signature_cost = matched_costs(text_fields) + null_cost, _ = matched_costs(null_text_fields) + cost = raw_cost - null_cost + + truth = torch.arange(n) + centered = cost - cost.mean(0, keepdim=True) + ranks = (centered <= centered.gather(1, truth[:, None])).sum(-1) + raw_ranks = (raw_cost <= raw_cost.gather(1, truth[:, None])).sum(-1) + rows, cols = linear_sum_assignment(cost.numpy()) + true_cost = cost.diagonal() + random_costs = [] + for _ in range(300): + permutation = torch.argsort(torch.rand(n, generator=generator)) + random_costs.append(float(cost[truth, permutation].mean())) + random_costs = torch.tensor(random_costs) + sorted_cost = centered.sort(-1).values + margin = sorted_cost[:, 1] - sorted_cost[:, 0] + order = margin.argsort(descending=True) + top = order[: max(1, n // 10)] + mutual = centered.argmin(-1)[centered.argmin(0)] == torch.arange(n) + forward = centered.argmin(-1) + if args.matrices_output: + torch.save( + { + "subset_seed": subset_seed, + "raw_cost": raw_cost, + "null_cost": null_cost, + }, + args.matrices_output.replace(".pt", f"_seed{subset_seed}.pt"), + ) + return { + "subset_seed": subset_seed, + "nodes": n, + "retrieval_null_corrected_centered": { + "r@1": float((ranks <= 1).double().mean()), + "r@5": float((ranks <= 5).double().mean()), + "r@10": float((ranks <= 10).double().mean()), + "median_rank": float(ranks.double().median()), + "chance_r@1": 1.0 / n, + }, + "retrieval_raw": { + "r@10": float((raw_ranks <= 10).double().mean()), + "median_rank": float(raw_ranks.double().median()), + }, + "retrieval_signature_only": { + "r@10": float( + ( + ( + signature_cost + <= signature_cost.gather(1, truth[:, None]) + ).sum(-1) + <= 10 + ) + .double() + .mean() + ), + "true_z": float( + ( + signature_cost.mean() + - signature_cost.diagonal().mean() + ) + / signature_cost.std().clamp_min(1e-12) + ), + }, + "linear_assignment_energy_null_corrected": { + "true_mean_cost": float(true_cost.mean()), + "random_mean_cost": float(random_costs.mean()), + "random_std": float(random_costs.std()), + "true_z": float( + (random_costs.mean() - true_cost.mean()) + / random_costs.std().clamp_min(1e-12) + ), + }, + "hungarian_recovery_accuracy": float( + (torch.from_numpy(cols) == truth).double().mean() + ), + "top_margin_decile_precision": float((forward[top] == top).double().mean()), + "mutual_nn_count": int(mutual.sum()), + "mutual_nn_precision": float( + (forward[mutual] == torch.arange(n)[mutual]).double().mean() + ) + if mutual.any() + else None, + } + + +def main() -> None: + args = parse_args() + seed_everything(args.seed) + data = load_aligned_views(args) + report: dict = { + "protocol": ( + "Node states are sixteen-view sets with internal standardized " + "relation fields. Part A gates orderings at the true view " + "assignment given the true coarse frame (upper bound); part B " + "matches internal fields blind, with no frame and no pooled " + "vectors, and uses hidden pairs only for scoring." + ), + "part_a_view_ordering": part_a(data, args), + } + print(json.dumps({"part_a": report["part_a_view_ordering"]})) + report["part_b_structured_affinity"] = [] + for subset_seed in (int(s) for s in args.part_b_seeds.split(",")): + result = part_b_one_seed(data, args, subset_seed) + report["part_b_structured_affinity"].append(result) + print(json.dumps({"part_b": result})) + write_json(args.output, report) + print(f"Wrote {args.output}") + + +if __name__ == "__main__": + main() diff --git a/worldalign/world_match.py b/worldalign/world_match.py new file mode 100644 index 0000000..e719a9f --- /dev/null +++ b/worldalign/world_match.py @@ -0,0 +1,138 @@ +"""World matching: spectral initialisation plus energy refinement. + +The two solvers are complementary. A spectral solver reads the coarse +correspondence out of the eigenstructure of two relation fields in +polynomial time and without any search, but its rounding is noisy. Exact +steepest descent on the closed-form energy repairs the rounding but +cannot find the basin from a random start. Composed, they recover the +hidden assignment. + +Neither stage sees a pair. The relation fields are built independently +per modality; the hidden permutation is applied to the text field and +read back only to score. +""" + +from __future__ import annotations + +import argparse +import json +from pathlib import Path + +import numpy as np +import torch + +from .common import seed_everything, write_json +from .spectral_match import grampa, spectral_profile, umeyama +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", required=True) + parser.add_argument("--eta", type=float, default=1.0) + parser.add_argument("--pair-weight", type=float, default=1.0) + parser.add_argument("--triangle-weight", type=float, default=1.0) + parser.add_argument("--refine-steps", type=int, default=4000) + parser.add_argument("--chunk", type=int, default=256) + parser.add_argument("--restarts", type=int, default=1) + parser.add_argument("--device", default="cuda:3") + parser.add_argument("--seed", type=int, default=0) + parser.add_argument("--output", default="") + return parser.parse_args() + + +def normalise(matrix: np.ndarray) -> np.ndarray: + size = len(matrix) + mask = ~np.eye(size, dtype=bool) + values = matrix[mask] + out = (matrix - values.mean()) / values.std() + np.fill_diagonal(out, 0.0) + return out + + +def main() -> None: + args = parse_args() + seed_everything(args.seed) + state = torch.load(args.fields, map_location="cpu", weights_only=False) + visual = normalise(state["visual_field"].double().numpy().copy()) + text = normalise(state["text_field"].double().numpy().copy()) + size = len(visual) + + generator = np.random.default_rng(args.seed) + hidden = generator.permutation(size) + shuffled = text[np.ix_(hidden, hidden)] + truth = np.arange(size) + + def accuracy(assignment: np.ndarray) -> float: + return float((hidden[assignment] == truth).mean()) + + device = torch.device(args.device) + text_tensor = standardized(torch.from_numpy(shuffled).to(device)).float() + visual_tensor = standardized(torch.from_numpy(visual).to(device)).float() + energy = ClosedFormEnergy( + text_tensor, + visual_tensor, + args.pair_weight, + args.triangle_weight, + args.chunk, + ) + swaps = all_swaps(size, device) + truth_permutation = torch.from_numpy(np.argsort(hidden)).to(device) + true_energy = float(energy.energy(truth_permutation[None])[0]) + + report = { + "protocol": ( + "Fields are built per modality without pairs; the hidden " + "permutation is applied to the text field and used only to " + "score. Spectral initialisation is followed by exact " + "steepest descent on the closed-form energy." + ), + "samples": size, + "field_correlation_at_truth": float( + np.corrcoef( + visual[~np.eye(size, dtype=bool)], + text[~np.eye(size, dtype=bool)], + )[0, 1] + ), + "true_energy": true_energy, + "chance": 1.0 / size, + "spectra": { + "visual": spectral_profile(visual), + "text": spectral_profile(shuffled), + }, + "stages": {}, + } + + spectral = grampa(visual, shuffled, args.eta) + report["stages"]["spectral"] = {"accuracy": accuracy(spectral)} + initial = torch.from_numpy(spectral.copy()).to(device) + refined, value = steepest_descent(energy, initial, swaps, args.refine_steps) + report["stages"]["spectral_plus_refinement"] = { + "accuracy": accuracy(refined.cpu().numpy()), + "energy": value, + "energy_over_true": value / abs(true_energy) - true_energy / abs(true_energy), + } + + quenches = [] + for restart in range(args.restarts): + start = torch.from_numpy(generator.permutation(size)).to(device) + final, final_value = steepest_descent( + energy, start, swaps, args.refine_steps + ) + quenches.append( + { + "accuracy": accuracy(final.cpu().numpy()), + "energy": final_value, + } + ) + report["stages"]["random_start_refinement"] = quenches + + print(json.dumps({key: value for key, value in report["stages"].items()})) + if args.output: + write_json(args.output, report) + print(f"Wrote {args.output}") + + +if __name__ == "__main__": + main() |
