summaryrefslogtreecommitdiff
path: root/worldalign/extract_text.py
diff options
context:
space:
mode:
Diffstat (limited to 'worldalign/extract_text.py')
-rw-r--r--worldalign/extract_text.py95
1 files changed, 95 insertions, 0 deletions
diff --git a/worldalign/extract_text.py b/worldalign/extract_text.py
new file mode 100644
index 0000000..8586e21
--- /dev/null
+++ b/worldalign/extract_text.py
@@ -0,0 +1,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()