From 67aeed2cd743506303b746697c1788caf5ef7e79 Mon Sep 17 00:00:00 2001 From: mrcabbage972 Date: Mon, 9 Jan 2023 23:03:29 -0500 Subject: [PATCH] Adding override of 32-bit optimization for embedding layer --- model/supervised_finetuning/trainer.py | 9 ++++++++- 1 file changed, 8 insertions(+), 1 deletion(-) diff --git a/model/supervised_finetuning/trainer.py b/model/supervised_finetuning/trainer.py index ae7fb3c3..450854f1 100644 --- a/model/supervised_finetuning/trainer.py +++ b/model/supervised_finetuning/trainer.py @@ -3,11 +3,11 @@ import os from distutils.util import strtobool from typing import Any, Dict, List, Optional, Tuple, Union +import bitsandbytes import torch from torch import nn from transformers import PreTrainedModel, Trainer, TrainingArguments from transformers.training_args import OptimizerNames - from utils import get_dataset, get_loss, get_model, get_tokenizer, read_yamls os.environ["WANDB_PROJECT"] = "supervised-finetuning" @@ -134,6 +134,13 @@ if __name__ == "__main__": optimizer = OptimizerNames.ADAMW_BNB if training_conf.quantization else None + if training_conf.quantization: + for module in model.modules(): + if isinstance(module, torch.nn.Embedding): + bitsandbytes.optim.GlobalOptimManager.get_instance().register_module_override( + module, "weight", {"optim_bits": 32} + ) + args = TrainingArguments( output_dir=f"{training_conf.model_name}-{training_conf.log_dir}-finetuned", num_train_epochs=training_conf.num_train_epochs,