custom liger patch

This commit is contained in:
51616 2025-01-13 15:20:58 +00:00
parent 59370b175f
commit 284d35fb12

View file

@ -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,