summaryrefslogtreecommitdiff
path: root/worldalign/train_prefix.py
diff options
context:
space:
mode:
Diffstat (limited to 'worldalign/train_prefix.py')
-rw-r--r--worldalign/train_prefix.py141
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()
+