adding bigvae

This commit is contained in:
wassname
2023-11-12 12:56:55 +08:00
parent ac6be401fe
commit 075de416c8
10 changed files with 4535 additions and 108 deletions
+486
View File
@@ -0,0 +1,486 @@
{
"cells": [
{
"cell_type": "code",
"execution_count": 1,
"metadata": {},
"outputs": [
{
"name": "stderr",
"output_type": "stream",
"text": [
"/media/wassname/SGIronWolf/projects5/worldmodels/vae_llm_worldmodel/.venv/lib/python3.11/site-packages/tqdm/auto.py:21: TqdmWarning: IProgress not found. Please update jupyter and ipywidgets. See https://ipywidgets.readthedocs.io/en/stable/user_install.html\n",
" from .autonotebook import tqdm as notebook_tqdm\n"
]
}
],
"source": [
"import argparse\n",
"from contextlib import contextmanager\n",
"from itertools import chain, islice\n",
"import json\n",
"import math\n",
"from pathlib import Path\n",
"import random\n",
"import os\n",
"import sys\n",
"import zipfile\n",
"\n",
"import accelerate\n",
"from datasets import load_dataset\n",
"import peft\n",
"import safetensors.torch as safetorch\n",
"import torch\n",
"from torch import nn, optim\n",
"from torch.nn import functional as F\n",
"from torch.utils import data\n",
"from tqdm import trange, tqdm\n",
"from transformers import AutoTokenizer, AutoModelForCausalLM, BitsAndBytesConfig\n",
"from loguru import logger\n"
]
},
{
"cell_type": "code",
"execution_count": 2,
"metadata": {},
"outputs": [],
"source": [
"# autoreload import your package\n",
"%load_ext autoreload\n",
"%autoreload 2\n",
"\n",
"from vae_llm_worldmodels.models.bigvae.bigvae import set_adapter, DecoderOnlyTransformerVAE, VAERouter\n"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"## Params"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": []
},
{
"cell_type": "code",
"execution_count": 6,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"['--rank', '16', '--context=96', '--vae_context=32', '--batch_size=1', '--output=./output/adapter']\n"
]
},
{
"data": {
"text/plain": [
"Namespace(batch_size=1, dropout=0.0, epochs=1, gradient_accumulation_steps=1, gradient_checkpointing=False, lr=0.0001, model='stabilityai/stablelm-3b-4e1t', context=96, vae_context=32, output=PosixPath('output/adapter'), rank=16, save_every=1000, start_from=None, z_dim=768)"
]
},
"execution_count": 6,
"metadata": {},
"output_type": "execute_result"
}
],
"source": [
"parser = argparse.ArgumentParser(description=__doc__)\n",
"parser.add_argument(\"--batch_size\", type=int, default=2, help=\"microbatch size\")\n",
"parser.add_argument(\"--dropout\", type=float, default=0.0, help=\"dropout rate\")\n",
"parser.add_argument(\"--epochs\", type=int, default=1, help=\"number of epochs\")\n",
"parser.add_argument(\n",
" \"--gradient_accumulation_steps\", type=int, default=1, help=\"gradient accumulation steps\"\n",
")\n",
"parser.add_argument(\n",
" \"--gradient_checkpointing\",\n",
" action=\"store_true\",\n",
" default=False,\n",
" help=\"use gradient checkpointing\",\n",
")\n",
"parser.add_argument(\"--lr\", type=float, default=1e-4, help=\"learning rate\")\n",
"parser.add_argument(\n",
" \"--model\",\n",
" type=str,\n",
" # default=\"mistralai/Mistral-7B-v0.1\",\n",
" # default=\"yichunkuo/stablelm-3b-4e1t-gptq\",\n",
" default=\"stabilityai/stablelm-3b-4e1t\", \n",
" # default=\"gpt2\", \n",
" # default=\"mlabonne/gpt2-GPTQ-4bit\",\n",
" help=\"model name\",\n",
")\n",
"parser.add_argument(\"--context\", type=int, default=2048, help=\"context window length\")\n",
"parser.add_argument(\"--vae_context\", type=int, default=64, help=\"vae embed context\")\n",
"parser.add_argument(\"--output\", type=Path, required=True, help=\"path to save adapter\")\n",
"parser.add_argument(\"--rank\", type=int, default=32, help=\"the lora rank\")\n",
"parser.add_argument(\"--save_every\", type=int, default=1000, help=\"save every n steps\")\n",
"parser.add_argument(\"--start_from\", type=str, help=\"start from existing lora\")\n",
"parser.add_argument(\"--z_dim\", type=int, default=768, help=\"the latent dimension\")\n",
"\n",
"\n",
"argvs = \"\"\"\n",
"--rank 16 \n",
"--context=96 \n",
"--vae_context=32 \n",
"--batch_size=1 \n",
"--output=./output/adapter \n",
"\"\"\"\n",
"argvs = argvs.replace('\\n', ' ').strip()\n",
"argv = [s.strip() for s in argvs.split(\" \") if s and not s.startswith(\"#\")]\n",
"print(argv)\n",
"args = parser.parse_args(argv)\n",
"args\n"
]
},
{
"cell_type": "code",
"execution_count": 7,
"metadata": {},
"outputs": [],
"source": [
"max_length = 32\n",
"tokenizer_args = dict(\n",
" padding='max_length', max_length=max_length,\n",
" truncation=True,\n",
")\n"
]
},
{
"cell_type": "code",
"execution_count": 8,
"metadata": {},
"outputs": [],
"source": [
"from transformers.utils.logging import _get_library_root_logger\n",
"\n",
"os.environ[\"TOKENIZERS_PARALLELISM\"] = \"false\"\n",
"os.environ[\"TRANSFORMERS_VERBOSITY\"] = \"detail\"\n",
"\n",
"library_root_logger = _get_library_root_logger()\n",
"library_root_logger.propagate = True\n",
"\n"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"## Load\n"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": []
},
{
"cell_type": "code",
"execution_count": 9,
"metadata": {},
"outputs": [
{
"name": "stderr",
"output_type": "stream",
"text": [
"\u001b[32m2023-11-12 09:44:45.920\u001b[0m | \u001b[1mINFO \u001b[0m | \u001b[36mcontextlib\u001b[0m:\u001b[36minner\u001b[0m:\u001b[36m81\u001b[0m - \u001b[1mLoading model: stabilityai/stablelm-3b-4e1t\u001b[0m\n",
"Using pad_token, but it is not set yet.\n"
]
},
{
"ename": "TypeError",
"evalue": "VAERouter.__init__() takes from 3 to 4 positional arguments but 5 were given",
"output_type": "error",
"traceback": [
"\u001b[0;31m---------------------------------------------------------------------------\u001b[0m",
"\u001b[0;31mTypeError\u001b[0m Traceback (most recent call last)",
"\u001b[1;32m/media/wassname/SGIronWolf/projects5/worldmodels/vae_llm_worldmodel/notebooks/mjc-002-vae_train.ipynb Cell 10\u001b[0m line \u001b[0;36m5\n\u001b[1;32m <a href='vscode-notebook-cell:/media/wassname/SGIronWolf/projects5/worldmodels/vae_llm_worldmodel/notebooks/mjc-002-vae_train.ipynb#X11sZmlsZQ%3D%3D?line=56'>57</a>\u001b[0m vae_model\u001b[39m.\u001b[39mvae\u001b[39m.\u001b[39mrequires_grad_(\u001b[39mFalse\u001b[39;00m)\n\u001b[1;32m <a href='vscode-notebook-cell:/media/wassname/SGIronWolf/projects5/worldmodels/vae_llm_worldmodel/notebooks/mjc-002-vae_train.ipynb#X11sZmlsZQ%3D%3D?line=57'>58</a>\u001b[0m vae_model\u001b[39m.\u001b[39mvae\u001b[39m.\u001b[39mw_d\u001b[39m.\u001b[39mrequires_grad_()\n\u001b[0;32m---> <a href='vscode-notebook-cell:/media/wassname/SGIronWolf/projects5/worldmodels/vae_llm_worldmodel/notebooks/mjc-002-vae_train.ipynb#X11sZmlsZQ%3D%3D?line=58'>59</a>\u001b[0m router \u001b[39m=\u001b[39m VAERouter(base_model_peft, vae_model, device, peft_config)\n\u001b[1;32m <a href='vscode-notebook-cell:/media/wassname/SGIronWolf/projects5/worldmodels/vae_llm_worldmodel/notebooks/mjc-002-vae_train.ipynb#X11sZmlsZQ%3D%3D?line=59'>60</a>\u001b[0m \u001b[39mif\u001b[39;00m args\u001b[39m.\u001b[39mstart_from:\n\u001b[1;32m <a href='vscode-notebook-cell:/media/wassname/SGIronWolf/projects5/worldmodels/vae_llm_worldmodel/notebooks/mjc-002-vae_train.ipynb#X11sZmlsZQ%3D%3D?line=60'>61</a>\u001b[0m router\u001b[39m.\u001b[39mload_pretrained(args\u001b[39m.\u001b[39mstart_from, is_trainable\u001b[39m=\u001b[39m\u001b[39mTrue\u001b[39;00m)\n",
"\u001b[0;31mTypeError\u001b[0m: VAERouter.__init__() takes from 3 to 4 positional arguments but 5 were given"
]
}
],
"source": [
"accelerator = accelerate.Accelerator(\n",
" mixed_precision=\"bf16\", gradient_accumulation_steps=args.gradient_accumulation_steps\n",
")\n",
"device = accelerator.device if accelerator.num_processes > 1 else \"cuda:0\"\n",
"is_main = accelerator.is_main_process\n",
"\n",
"print = tqdm.external_write_mode()(logger.info)\n",
"print0 = accelerator.on_main_process(print)\n",
"\n",
"if Path(args.model).exists():\n",
" model_name = Path(args.model).resolve()\n",
"else:\n",
" model_name = args.model\n",
"\n",
"print0(f\"Loading model: {model_name}\")\n",
"with accelerator.main_process_first():\n",
" tokenizer = AutoTokenizer.from_pretrained(model_name, trust_remote_code=True)\n",
" tokenizer.padding_side = \"left\"\n",
" if tokenizer.pad_token is None:\n",
" tokenizer.pad_token = tokenizer.eos_token\n",
" bnb_config = BitsAndBytesConfig(\n",
" load_in_4bit=True,\n",
" bnb_4bit_compute_dtype=torch.bfloat16,\n",
" bnb_4bit_quant_type=\"nf4\",\n",
" bnb_4bit_use_double_quant=True,\n",
" )\n",
" base_model = AutoModelForCausalLM.from_pretrained(\n",
" model_name,\n",
" device_map={\"\": device},\n",
" quantization_config=bnb_config,\n",
" torch_dtype=torch.bfloat16, \n",
" trust_remote_code=True\n",
" )\n",
" peft_config = peft.LoraConfig(\n",
" peft.TaskType.CAUSAL_LM,\n",
" inference_mode=False,\n",
" r=args.rank,\n",
" lora_alpha=8,\n",
" lora_dropout=args.dropout,\n",
" target_modules=[\n",
" \"self_attn.q_proj\",\n",
" \"self_attn.k_proj\",\n",
" \"self_attn.v_proj\",\n",
" \"self_attn.o_proj\",\n",
" \"mlp.gate_proj\",\n",
" \"mlp.up_proj\",\n",
" \"mlp.down_proj\",\n",
" ],\n",
" )\n",
" base_model_peft = peft.get_peft_model(base_model, peft_config)\n",
" vae_model = DecoderOnlyTransformerVAE(\n",
" base_model_peft, peft_config, device=device, z_dim=args.z_dim,\n",
" )\n",
" if args.start_from:\n",
" vae_model.load_pretrained(args.start_from)\n",
" base_model_peft.requires_grad_(False)\n",
" vae_model.vae.requires_grad_(False)\n",
" vae_model.vae.w_d.requires_grad_()\n",
" router = VAERouter(base_model_peft, vae_model, device)\n",
" if args.start_from:\n",
" router.load_pretrained(args.start_from, is_trainable=True)\n",
"accelerator.wait_for_everyone()\n",
"\n",
"router.train()\n",
"if args.gradient_checkpointing:\n",
" router.model.gradient_checkpointing_enable()\n",
" router.model.enable_input_require_grads()\n",
"\n",
"if is_main:\n",
" router.model.print_trainable_parameters()\n",
"\n",
"router.model.set_adapter(\"router\")\n",
"opt = optim.Adam(router.model.parameters(),\n",
" lr=args.lr,\n",
" betas=(0.9, 0.99))\n",
"accelerator.wait_for_everyone()\n"
]
},
{
"cell_type": "code",
"execution_count": 17,
"metadata": {},
"outputs": [],
"source": []
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"https://github.com/JD-P/minihf/blob/adavae-moe/train_vae_router.py#L277\n"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
"# prepare dataset\n",
"input_ids_all, attention_mask_all = [], []\n",
"for shard_name in os.listdir(args.preprocessed):\n",
" data_path = os.path.join(args.preprocessed, shard_name)\n",
" data_file = safetorch.load_file(data_path)\n",
" input_ids = torch.split(data_file[\"input_ids\"], args.context, dim=1)\n",
" attention_mask = torch.split(data_file[\"attention_mask\"], args.context, dim=1)\n",
" if input_ids[-1].shape[1] != args.context:\n",
" input_ids = input_ids[:-1]\n",
" attention_mask = attention_mask[:-1]\n",
" input_ids_all.extend(input_ids)\n",
" attention_mask_all.extend(attention_mask)\n",
"del data_file, input_ids, attention_mask\n",
"input_ids_all = torch.cat(input_ids_all)\n",
"attention_mask_all = torch.cat(attention_mask_all)\n",
"valid_indices = attention_mask_all.sum(dim=1) == args.context\n",
"input_ids_all = input_ids_all[valid_indices]\n",
"attention_mask_all = attention_mask_all[valid_indices]\n",
"del valid_indices\n",
"\n",
"preprocessed = data.TensorDataset(input_ids_all, attention_mask_all)\n",
"\n",
"dataloader = data.DataLoader(\n",
" preprocessed,\n",
" batch_size=args.batch_size,\n",
" shuffle=True,\n",
" drop_last=True,\n",
")\n",
"\n",
"router, opt, dataloader = accelerator.prepare(router, opt, dataloader)\n"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
"from vae_llm_worldmodels.utils import cosine_warmup\n",
"\n",
"\n",
"@torch.no_grad()\n",
"@torch.cuda.amp.autocast(dtype=torch.bfloat16)\n",
"def demo(model, input_ids, attention_mask, n_tokens):\n",
" \"\"\"inference.\"\"\"\n",
" bs = min(input_ids.shape[0], 2)\n",
" n_outputs = 2\n",
" tau = 0.8\n",
"\n",
" index = random.randrange(args.context - (args.vae_context * 2))\n",
" context_ids = input_ids[:,:index]\n",
" context_mask = attention_mask[:,:index]\n",
" embed_ids = input_ids[:,index:index + args.vae_context]\n",
" embed_mask = attention_mask[:,index:index + args.vae_context]\n",
" target_ids = input_ids[:,index:index + args.vae_context * 2]\n",
" target_mask = input_ids[:,index:index + args.vae_context * 2]\n",
"\n",
" in_texts = [tokenizer.decode(toks, skip_special_tokens=True)\n",
" for toks in torch.cat([context_ids, embed_ids], dim=1)]\n",
" mean = model.encode(embed_ids[:bs], embed_mask[:bs])\n",
" z = model.vae.vae.sample(mean.repeat_interleave(n_outputs, 0), tau=tau)\n",
" context_ids = context_ids[:bs].repeat_interleave(n_outputs, 0)\n",
" context_mask = context_mask[:bs].repeat_interleave(n_outputs, 0)\n",
" # empty = z.new_zeros([z.shape[0], 0], dtype=torch.long)\n",
" output_ids = model.generate(z, context_ids, context_mask, n_tokens, tau=tau)\n",
" out_texts = [tokenizer.decode(toks, skip_special_tokens=True) for toks in output_ids]\n",
" out_texts = list(batched(out_texts, n_outputs))\n",
" print(\"======\")\n",
" for in_text, out_batch in zip(in_texts, out_texts):\n",
" print(\"=== Input ===\")\n",
" print(in_text)\n",
" print(\"=== Outputs ===\")\n",
" for i, out_text in enumerate(out_batch):\n",
" print(out_text)\n",
" if i < len(out_batch) - 1:\n",
" print(\"===\")\n",
" print(\"======\")\n",
"\n",
"def save():\n",
" print0(f\"### Saving model to {args.output}\", file=sys.stderr)\n",
" accelerator.wait_for_everyone()\n",
" if accelerator.is_main_process:\n",
" unwrapped_model = accelerator.unwrap_model(router)\n",
" unwrapped_model.save_pretrained(args.output)\n",
" state_obj = {\"step\": i, \"last_kl_weight\": kl_sched(i)}\n",
" with open(args.output / \"state.json\", \"w\") as f:\n",
" json.dump(state_obj, f)\n"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
"# train\n",
"i = 0\n",
"kl_sched = cosine_warmup(5000, 0.01)\n",
"\n",
"\n",
"accelerator.wait_for_everyone()\n",
"for epoch in trange(args.epochs, disable=not is_main):\n",
" for input_ids, attention_mask in tqdm(dataloader, disable=not is_main):\n",
" input_ids = input_ids.long()\n",
" if is_main and i % 100 == 0:\n",
" demo(accelerator.unwrap_model(router), input_ids, attention_mask, args.vae_context)\n",
" pass\n",
" with accelerator.accumulate(router):\n",
" index = random.randrange(args.context - (args.vae_context * 2))\n",
" context_ids = input_ids[:,:index]\n",
" context_mask = attention_mask[:,:index]\n",
" embed_ids = input_ids[:,index:index + args.vae_context]\n",
" embed_mask = attention_mask[:,index:index + args.vae_context]\n",
" target_ids = input_ids[:,index:index + args.vae_context * 2]\n",
" target_mask = attention_mask[:,index:index + args.vae_context * 2]\n",
"\n",
" drop_mask = torch.rand([context_ids.shape[0], 1], device=device) < 0.5\n",
" context_ids = torch.where(drop_mask, torch.zeros_like(context_ids), context_ids)\n",
" context_mask = torch.where(drop_mask, torch.zeros_like(context_mask), context_mask)\n",
" outputs = router(embed_ids, embed_mask,\n",
" target_ids[:,:-1], target_mask[:,:-1],\n",
" context_ids, context_mask)\n",
" rec_losses = F.cross_entropy(\n",
" outputs.logits[:, -args.vae_context * 2:].transpose(-1, -2),\n",
" target_ids,\n",
" reduction=\"none\",\n",
" )\n",
" n_toks = target_mask.sum()\n",
" rec_loss = torch.sum(rec_losses * target_mask, dtype=torch.float32) / n_toks\n",
" # kl_loss = torch.sum(mean**2 / 2, dtype=torch.float32) * kl_sched(i) / n_toks\n",
" loss = rec_loss # + kl_loss\n",
"\n",
" # accelerator.backward(loss, inputs=list(p for p in accelerator.unwrap_model(router).model.parameters() if p.requires_grad))\n",
" accelerator.backward(loss)\n",
" # for n, p in router.named_parameters():\n",
" # if p.grad is not None:\n",
" # grad_norm = torch.norm(p.grad, dtype=torch.float32)\n",
" # if grad_norm != 0:\n",
" # print(f\"{n}: {grad_norm:g}\", file=sys.stderr)\n",
" opt.step()\n",
" opt.zero_grad()\n",
"\n",
" loss_global, rec_global = accelerator.reduce(\n",
" (loss, rec_loss), \"mean\"\n",
" )\n",
" print0(\n",
" f\"epoch: {epoch}, step: {i}, loss: {loss_global.item():g}, rec: {rec_global.item():g}\",\n",
" file=sys.stderr,\n",
" )\n",
" i += 1\n",
"\n",
" if i % args.save_every == 0:\n",
" save()\n",
"\n",
" save()\n"
]
}
],
"metadata": {
"kernelspec": {
"display_name": ".venv",
"language": "python",
"name": "python3"
},
"language_info": {
"codemirror_mode": {
"name": "ipython",
"version": 3
},
"file_extension": ".py",
"mimetype": "text/x-python",
"name": "python",
"nbconvert_exporter": "python",
"pygments_lexer": "ipython3",
"version": "3.11.0rc1"
}
},
"nbformat": 4,
"nbformat_minor": 2
}