diff options
Diffstat (limited to 'worldalign/train_prefix.py')
| -rw-r--r-- | worldalign/train_prefix.py | 141 |
1 files changed, 141 insertions, 0 deletions
diff --git a/worldalign/train_prefix.py b/worldalign/train_prefix.py new file mode 100644 index 0000000..76607e6 --- /dev/null +++ b/worldalign/train_prefix.py @@ -0,0 +1,141 @@ +from __future__ import annotations + +import argparse +from pathlib import Path + +import torch +from torch.optim import AdamW +from tqdm import tqdm +from transformers import AutoModelForCausalLM, AutoTokenizer + +from .common import ( + cosine_schedule, + dtype_for_device, + parameter_count, + read_json, + seed_everything, +) +from .io import load_feature_pair, select_rows +from .models import PrefixAdapter + + +def parse_args() -> argparse.Namespace: + p = argparse.ArgumentParser() + p.add_argument("--manifest", default="artifacts/manifest.json") + p.add_argument("--vision", default="artifacts/vision.pt") + p.add_argument("--text", default="artifacts/text.pt") + p.add_argument("--output", default="artifacts/prefix.pt") + p.add_argument("--device", default="cuda:1") + p.add_argument("--steps", type=int, default=3_000) + p.add_argument("--batch-size", type=int, default=32) + p.add_argument("--prefix-length", type=int, default=8) + p.add_argument("--hidden-dim", type=int, default=2048) + p.add_argument("--max-length", type=int, default=48) + p.add_argument("--lr", type=float, default=3e-4) + p.add_argument("--warmup", type=int, default=200) + p.add_argument("--seed", type=int, default=20260728) + return p.parse_args() + + +def main() -> None: + args = parse_args() + seed_everything(args.seed) + manifest = read_json(args.manifest) + _, text, _, tlookup = load_feature_pair(args.vision, args.text) + train_rows = manifest["text_only_train"] + semantic = select_rows(text["features"], tlookup, train_rows) + captions_by_row = { + int(row): caption for row, caption in zip(text["rows"], text["captions"]) + } + captions = [captions_by_row[int(row)] for row in train_rows] + + tokenizer = AutoTokenizer.from_pretrained(text["model"]) + if tokenizer.pad_token_id is None: + tokenizer.pad_token = tokenizer.eos_token + tokenizer.padding_side = "right" + dtype = dtype_for_device(args.device) + lm = AutoModelForCausalLM.from_pretrained( + text["model"], torch_dtype=dtype + ).to(args.device) + lm.eval() + for parameter in lm.parameters(): + parameter.requires_grad_(False) + lm_dim = lm.get_input_embeddings().embedding_dim + + adapter = PrefixAdapter( + semantic_dim=semantic.shape[-1], + lm_dim=lm_dim, + prefix_length=args.prefix_length, + hidden_dim=args.hidden_dim, + ).to(args.device) + print(f"Prefix adapter parameters: {parameter_count(adapter):,}") + optimizer = AdamW(adapter.parameters(), lr=args.lr, weight_decay=1e-4) + generator = torch.Generator().manual_seed(args.seed) + + history = [] + progress = tqdm(range(args.steps), desc="text-only prefix") + for step in progress: + ids = torch.randint( + len(semantic), (args.batch_size,), generator=generator + ) + batch_captions = [captions[int(i)] for i in ids] + tokens = tokenizer( + batch_captions, + padding=True, + truncation=True, + max_length=args.max_length, + return_tensors="pt", + ) + input_ids = tokens["input_ids"].to(args.device) + attention = tokens["attention_mask"].to(args.device) + prefix = adapter(semantic[ids].to(args.device)).to(dtype) + token_embeddings = lm.get_input_embeddings()(input_ids) + inputs_embeds = torch.cat([prefix, token_embeddings], dim=1) + prefix_attention = torch.ones( + prefix.shape[:2], dtype=attention.dtype, device=args.device + ) + full_attention = torch.cat([prefix_attention, attention], dim=1) + labels = input_ids.clone() + labels[attention == 0] = -100 + prefix_labels = torch.full( + prefix.shape[:2], -100, dtype=labels.dtype, device=args.device + ) + full_labels = torch.cat([prefix_labels, labels], dim=1) + result = lm( + inputs_embeds=inputs_embeds, + attention_mask=full_attention, + labels=full_labels, + use_cache=False, + return_dict=True, + ) + loss = result.loss + optimizer.zero_grad(set_to_none=True) + loss.backward() + torch.nn.utils.clip_grad_norm_(adapter.parameters(), 1.0) + optimizer.step() + scale = cosine_schedule(step, args.steps, args.warmup) + for group in optimizer.param_groups: + group["lr"] = args.lr * scale + + if step % 20 == 0: + progress.set_postfix(loss=f"{loss.item():.3f}") + if step % 100 == 0 or step == args.steps - 1: + history.append({"step": step, "loss": float(loss.item())}) + + state = { + "config": adapter.config(), + "state_dict": adapter.state_dict(), + "text_model": text["model"], + "text_layer": text["layer"], + "args": vars(args), + "history": history, + "training": "text only; no image features or image-text pairs", + } + Path(args.output).parent.mkdir(parents=True, exist_ok=True) + torch.save(state, args.output) + print(f"Wrote {args.output}") + + +if __name__ == "__main__": + main() + |
