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