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