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