diff options
Diffstat (limited to 'worldalign/vg_attention_extract.py')
| -rw-r--r-- | worldalign/vg_attention_extract.py | 269 |
1 files changed, 269 insertions, 0 deletions
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() |
