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