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