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