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