diff options
Diffstat (limited to 'worldalign/extract_text.py')
| -rw-r--r-- | worldalign/extract_text.py | 95 |
1 files changed, 95 insertions, 0 deletions
diff --git a/worldalign/extract_text.py b/worldalign/extract_text.py new file mode 100644 index 0000000..8586e21 --- /dev/null +++ b/worldalign/extract_text.py @@ -0,0 +1,95 @@ +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() |
