From 2960bf8d948837a6cb16c394060717880a613306 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Quentin=20Gallou=C3=A9dec?= <45557362+qgallouedec@users.noreply.github.com> Date: Thu, 17 Jul 2025 15:08:38 -0700 Subject: [PATCH] Update unsloth-cli.py (#2985) --- unsloth-cli.py | 13 ++++++------- 1 file changed, 6 insertions(+), 7 deletions(-) diff --git a/unsloth-cli.py b/unsloth-cli.py index 86f02075f1..a782cfd743 100644 --- a/unsloth-cli.py +++ b/unsloth-cli.py @@ -38,7 +38,7 @@ def run(args): from unsloth import FastLanguageModel from datasets import load_dataset from transformers.utils import strtobool - from trl import SFTTrainer + from trl import SFTTrainer, SFTConfig from transformers import TrainingArguments from unsloth import is_bfloat16_supported import logging @@ -100,7 +100,7 @@ def run(args): print("Data is formatted and ready!") # Configure training arguments - training_args = TrainingArguments( + training_args = SFTConfig( per_device_train_batch_size=args.per_device_train_batch_size, gradient_accumulation_steps=args.gradient_accumulation_steps, warmup_steps=args.warmup_steps, @@ -115,17 +115,16 @@ def run(args): seed=args.seed, output_dir=args.output_dir, report_to=args.report_to, + max_length=args.max_seq_length, + dataset_num_proc=2, + packing=False, ) # Initialize trainer trainer = SFTTrainer( model=model, - tokenizer=tokenizer, + processing_class=tokenizer, train_dataset=dataset, - dataset_text_field="text", - max_seq_length=args.max_seq_length, - dataset_num_proc=2, - packing=False, args=training_args, )