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

import argparse
import json
from pathlib import Path

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

from .common import dtype_for_device
from .extract_text import mean_pool


def parse_args() -> argparse.Namespace:
    p = argparse.ArgumentParser()
    p.add_argument("--nodes", default="artifacts/vg/text_nodes.jsonl")
    p.add_argument("--output", default="artifacts/vg/text_features.pt")
    p.add_argument(
        "--tier",
        choices=["region_closed", "visible_relations", "qa_expanded"],
        default="region_closed",
    )
    p.add_argument("--model", default="Qwen/Qwen2.5-1.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=48)
    p.add_argument("--layer", type=int, default=-1)
    p.add_argument("--views", type=int, default=32)
    p.add_argument("--limit", type=int)
    return p.parse_args()


def read_jsonl(path: str) -> list[dict]:
    with open(path, encoding="utf-8") as handle:
        return [json.loads(line) for line in handle if line.strip()]


@torch.inference_mode()
def main() -> None:
    args = parse_args()
    records = read_jsonl(args.nodes)
    if args.limit:
        records = records[: args.limit]
    if not records:
        raise ValueError("No text nodes found")
    for record in records:
        if args.tier not in record:
            raise ValueError(
                f"Tier {args.tier!r} is absent; regenerate bundles with "
                "--with-extra-tiers if needed"
            )
        if len(record[args.tier]) < args.views:
            raise ValueError(
                f"Node {record['node_id']} has fewer than {args.views} views"
            )

    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()

    flat_texts = [
        text
        for record in records
        for text in record[args.tier][: args.views]
    ]
    outputs: list[torch.Tensor] = []
    for start in tqdm(
        range(0, len(flat_texts), args.batch_size),
        desc=f"Qwen VG text:{args.tier}",
    ):
        tokens = tokenizer(
            flat_texts[start : start + args.batch_size],
            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
        )
        outputs.append(
            mean_pool(
                result.hidden_states[args.layer], tokens["attention_mask"]
            )
            .float()
            .cpu()
        )
    features = torch.cat(outputs).reshape(len(records), args.views, -1)
    state = {
        "model": args.model,
        "layer": args.layer,
        "tier": args.tier,
        "node_ids": [record["node_id"] for record in records],
        "region_features": features,
        "views_per_node": args.views,
    }
    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()