diff options
| author | Yuren Hao <blackhao0426@gmail.com> | 2026-08-01 14:10:03 -0500 |
|---|---|---|
| committer | Yuren Hao <blackhao0426@gmail.com> | 2026-08-01 14:10:03 -0500 |
| commit | a62cf4d2a99b4a7985c61b2a7feb92a82a8218b7 (patch) | |
| tree | ee2248078db7edf3812a07f195afa3d9bd6f10c6 /worldalign/view_anchor_battery.py | |
World Alignment: unpaired cross-modal correspondence by relational identifiability
Method: scene states are sets of part states; relation fields are built
within each modality and are invariant to how each side labels its own
features; the cross-modal bridge is a coupling searched under an energy
that is a closed-form functional of one matrix; solving is spectral
initialisation followed by exact local refinement.
Evidence: in a procedurally generated closed world, blind recovery of a
hidden image-caption correspondence reaches 95.3% at 256 scenes against
0.39% chance, and the recovered pairs transfer to 200 held-out scenes at
93.0% exact retrieval with random-pair and shuffled-image controls at or
near chance. Cross-modal value correspondence is derived from disjoint
corpora rather than declared. On Visual Genome the field correlation
reaches 0.656 against the 0.9 that polynomial recovery needs, with the
deficit attributed away from segmentation and discretisation.
Protocol: no image-text pair enters any objective, optimiser,
initialisation, or model selection; hidden pairs score orderings only.
Co-Authored-By: Claude <noreply@anthropic.com>
Diffstat (limited to 'worldalign/view_anchor_battery.py')
| -rw-r--r-- | worldalign/view_anchor_battery.py | 365 |
1 files changed, 365 insertions, 0 deletions
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() |
