mirror of
https://github.com/SakanaAI/doc-to-lora.git
synced 2026-07-23 17:01:04 +02:00
custom liger patch
This commit is contained in:
parent
59370b175f
commit
284d35fb12
1 changed files with 10 additions and 3 deletions
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue