From 284d35fb126817d2d620ff3c011d82f6334d8897 Mon Sep 17 00:00:00 2001 From: 51616 Date: Mon, 13 Jan 2025 15:20:58 +0000 Subject: [PATCH] custom liger patch --- hyperlora/intx_sft.py | 13 ++++++++++--- 1 file changed, 10 insertions(+), 3 deletions(-) diff --git a/hyperlora/intx_sft.py b/hyperlora/intx_sft.py index cddd269..a7ff2fb 100755 --- a/hyperlora/intx_sft.py +++ b/hyperlora/intx_sft.py @@ -46,6 +46,8 @@ from transformers import ( HfArgumentParser, set_seed, ) +from peft import PeftModel +from transformers.utils import is_liger_kernel_available from torch.utils.data import DataLoader from utils import ( extract_cli_args, @@ -461,12 +463,17 @@ def main(): # TODO: use SFTTrainer instead? https://huggingface.co/docs/trl/en/sft_trainer # TODO: use packing with SFTTrainer - # HACK: see transformers/trainer.py for liger-kernel patch - # slows down training speed w/ short inputs - # might improve/decrease training speed w/ longer inputs # HACK: see transformers/trainer_seq2seq.py for supressing # "Trainer.tokenizer is now deprecated. You should use Trainer.processing_class instead." + if training_args.use_liger_kernel and is_liger_kernel_available(): + from liger_kernel.transformers import _apply_liger_kernel_to_instance + if isinstance(model, ModulatedPretrainedModel): + logger.info("Applying liger-kernel to ModulatedPretrainedModel") + _apply_liger_kernel_to_instance(model=model.base_model.base_model.model) + elif isinstance(model, PeftModel): + logger.info("Applying liger-kernel to PeftModel") + _apply_liger_kernel_to_instance(model=model.base_model.model) wandb.init( project=os.getenv("WANDB_PROJECT"), name=run_name,