sankalphs commited on
Commit
bf88014
·
1 Parent(s): 94d68b4

Fix Modal training: use SFTConfig to avoid pickling error

Browse files
Files changed (1) hide show
  1. training/modal_train.py +6 -6
training/modal_train.py CHANGED
@@ -62,10 +62,10 @@ def train(
62
  lora_alpha: int = 32,
63
  max_seq_length: int = 1024,
64
  ):
 
65
  import torch
66
  from datasets import load_dataset
67
- from transformers import TrainingArguments
68
- from trl import SFTTrainer
69
  from unsloth import FastLanguageModel
70
 
71
  hf_token = os.environ.get("HF_TOKEN")
@@ -118,7 +118,7 @@ def train(
118
  dataset = dataset.map(format_messages, batched=True)
119
 
120
  output_dir = "/tmp/retro-alpha-lora"
121
- training_args = TrainingArguments(
122
  output_dir=output_dir,
123
  num_train_epochs=num_epochs,
124
  per_device_train_batch_size=per_device_batch_size,
@@ -133,15 +133,15 @@ def train(
133
  group_by_length=True,
134
  report_to="none",
135
  remove_unused_columns=False,
 
 
136
  )
137
 
138
  trainer = SFTTrainer(
139
  model=model,
140
  train_dataset=dataset,
141
- tokenizer=tokenizer,
142
  args=training_args,
143
- max_seq_length=max_seq_length,
144
- dataset_text_field="text",
145
  )
146
 
147
  print("Starting training...")
 
62
  lora_alpha: int = 32,
63
  max_seq_length: int = 1024,
64
  ):
65
+ import unsloth # noqa: F401, must import first per Unsloth warning
66
  import torch
67
  from datasets import load_dataset
68
+ from trl import SFTConfig, SFTTrainer
 
69
  from unsloth import FastLanguageModel
70
 
71
  hf_token = os.environ.get("HF_TOKEN")
 
118
  dataset = dataset.map(format_messages, batched=True)
119
 
120
  output_dir = "/tmp/retro-alpha-lora"
121
+ training_args = SFTConfig(
122
  output_dir=output_dir,
123
  num_train_epochs=num_epochs,
124
  per_device_train_batch_size=per_device_batch_size,
 
133
  group_by_length=True,
134
  report_to="none",
135
  remove_unused_columns=False,
136
+ dataset_text_field="text",
137
+ max_seq_length=max_seq_length,
138
  )
139
 
140
  trainer = SFTTrainer(
141
  model=model,
142
  train_dataset=dataset,
143
+ processing_class=tokenizer,
144
  args=training_args,
 
 
145
  )
146
 
147
  print("Starting training...")