Spaces:
Runtime error
Runtime error
Fix Modal training: use SFTConfig to avoid pickling error
Browse files- 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
|
| 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 =
|
| 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 |
-
|
| 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...")
|