summaryrefslogtreecommitdiff
path: root/worldalign/extract_text_orbits.py
blob: 43b9020cd1bbd32036fff2c607f8fef966b65e6a (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
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()