mirror of
https://github.com/wassname/iris_bigvae.git
synced 2026-09-09 11:24:31 +08:00
adding bigvae
This commit is contained in:
@@ -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
|
||||
}
|
||||
Reference in New Issue
Block a user