diff options
Diffstat (limited to 'worldalign/vg_extract_text.py')
| -rw-r--r-- | worldalign/vg_extract_text.py | 111 |
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() |
