summaryrefslogtreecommitdiff
path: root/worldalign/synth_extract.py
blob: 3186c3ee97c1102a6ff96df156eb70503edef598 (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
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()