diff options
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() |
