from __future__ import annotations import argparse from pathlib import Path import torch from datasets import load_dataset from tqdm import tqdm from transformers import AutoModel, AutoTokenizer from .common import batch_indices, dtype_for_device, read_json def parse_args() -> argparse.Namespace: p = argparse.ArgumentParser() p.add_argument("--manifest", default="artifacts/manifest.json") p.add_argument("--output", default="artifacts/text.pt") p.add_argument("--model", default="Qwen/Qwen2.5-0.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=64) p.add_argument( "--layer", type=int, default=-1, help="Hidden-state index; -1 is the final transformer output.", ) p.add_argument("--limit", type=int) return p.parse_args() def mean_pool(hidden: torch.Tensor, mask: torch.Tensor) -> torch.Tensor: weights = mask.to(hidden.dtype).unsqueeze(-1) return (hidden * weights).sum(1) / weights.sum(1).clamp_min(1) @torch.inference_mode() def main() -> None: args = parse_args() manifest = read_json(args.manifest) rows = manifest["all_rows"] if args.limit: rows = rows[: args.limit] dataset = load_dataset(manifest["dataset"], split=manifest["dataset_split"]) # Avoid decoding the 4.3GB image column when only captions are requested. dataset = dataset.remove_columns("image") 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) # The backbone is sufficient for hidden states. Using a causal-LM wrapper # would also materialize vocabulary logits that are discarded here. model = AutoModel.from_pretrained(args.model, torch_dtype=dtype).to(args.device) model.eval() outputs: list[torch.Tensor] = [] captions: list[str] = [] all_captions: list[list[str]] = [] for ids in tqdm( batch_indices(len(rows), args.batch_size), desc="Qwen text features" ): batch_captions = [dataset[int(rows[i])]["caption"] for i in ids] texts = [caps[0] for caps in batch_captions] tokens = tokenizer( texts, padding=True, truncation=True, max_length=args.max_length, return_tensors="pt", ) tokens = {k: v.to(args.device) for k, v in tokens.items()} result = model(**tokens, output_hidden_states=True, return_dict=True) feature = mean_pool( result.hidden_states[args.layer], tokens["attention_mask"] ) outputs.append(feature.float().cpu()) captions.extend(texts) all_captions.extend(batch_captions) value = { "model": args.model, "layer": args.layer, "rows": rows, "features": torch.cat(outputs), "captions": captions, "all_captions": all_captions, } Path(args.output).parent.mkdir(parents=True, exist_ok=True) torch.save(value, args.output) print(f"Wrote {args.output}: {tuple(value['features'].shape)}") if __name__ == "__main__": main()