summaryrefslogtreecommitdiff
path: root/worldalign/extract_text.py
blob: 8586e21754c73533bc3af8b6dd19d893945c79b5 (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
from __future__ import annotations

import argparse
from pathlib import Path

import torch
from datasets import load_dataset
from tqdm import tqdm
from transformers import AutoModel, AutoTokenizer

from .common import batch_indices, dtype_for_device, read_json


def parse_args() -> argparse.Namespace:
    p = argparse.ArgumentParser()
    p.add_argument("--manifest", default="artifacts/manifest.json")
    p.add_argument("--output", default="artifacts/text.pt")
    p.add_argument("--model", default="Qwen/Qwen2.5-0.5B")
    p.add_argument("--device", default="cuda:3")
    p.add_argument("--batch-size", type=int, default=96)
    p.add_argument("--max-length", type=int, default=64)
    p.add_argument(
        "--layer",
        type=int,
        default=-1,
        help="Hidden-state index; -1 is the final transformer output.",
    )
    p.add_argument("--limit", type=int)
    return p.parse_args()


def mean_pool(hidden: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
    weights = mask.to(hidden.dtype).unsqueeze(-1)
    return (hidden * weights).sum(1) / weights.sum(1).clamp_min(1)


@torch.inference_mode()
def main() -> None:
    args = parse_args()
    manifest = read_json(args.manifest)
    rows = manifest["all_rows"]
    if args.limit:
        rows = rows[: args.limit]
    dataset = load_dataset(manifest["dataset"], split=manifest["dataset_split"])
    # Avoid decoding the 4.3GB image column when only captions are requested.
    dataset = dataset.remove_columns("image")
    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)
    # The backbone is sufficient for hidden states. Using a causal-LM wrapper
    # would also materialize vocabulary logits that are discarded here.
    model = AutoModel.from_pretrained(args.model, torch_dtype=dtype).to(args.device)
    model.eval()

    outputs: list[torch.Tensor] = []
    captions: list[str] = []
    all_captions: list[list[str]] = []
    for ids in tqdm(
        batch_indices(len(rows), args.batch_size), desc="Qwen text features"
    ):
        batch_captions = [dataset[int(rows[i])]["caption"] for i in ids]
        texts = [caps[0] for caps in batch_captions]
        tokens = tokenizer(
            texts,
            padding=True,
            truncation=True,
            max_length=args.max_length,
            return_tensors="pt",
        )
        tokens = {k: v.to(args.device) for k, v in tokens.items()}
        result = model(**tokens, output_hidden_states=True, return_dict=True)
        feature = mean_pool(
            result.hidden_states[args.layer], tokens["attention_mask"]
        )
        outputs.append(feature.float().cpu())
        captions.extend(texts)
        all_captions.extend(batch_captions)

    value = {
        "model": args.model,
        "layer": args.layer,
        "rows": rows,
        "features": torch.cat(outputs),
        "captions": captions,
        "all_captions": all_captions,
    }
    Path(args.output).parent.mkdir(parents=True, exist_ok=True)
    torch.save(value, args.output)
    print(f"Wrote {args.output}: {tuple(value['features'].shape)}")


if __name__ == "__main__":
    main()