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