Spaces:
Sleeping
Sleeping
Fix: batch_size=16, 10K samples, unbuffered output, 2 epochs
Browse files- train_on_hf.py +12 -5
train_on_hf.py
CHANGED
|
@@ -13,8 +13,13 @@ Usage (on HF with GPU):
|
|
| 13 |
|
| 14 |
import argparse
|
| 15 |
import json
|
|
|
|
| 16 |
from pathlib import Path
|
| 17 |
|
|
|
|
|
|
|
|
|
|
|
|
|
| 18 |
|
| 19 |
def setup(hf_token: str):
|
| 20 |
"""Login and download data from HF."""
|
|
@@ -47,7 +52,7 @@ def setup(hf_token: str):
|
|
| 47 |
return data_dir
|
| 48 |
|
| 49 |
|
| 50 |
-
def train_world_model(data_dir: Path, max_samples: int =
|
| 51 |
"""Stage 2: Train World Model via SFT."""
|
| 52 |
from datasets import Dataset
|
| 53 |
from transformers import AutoModelForCausalLM, AutoTokenizer
|
|
@@ -104,20 +109,22 @@ def train_world_model(data_dir: Path, max_samples: int = 50000):
|
|
| 104 |
|
| 105 |
training_args = SFTConfig(
|
| 106 |
output_dir="./world_model_checkpoints",
|
| 107 |
-
num_train_epochs=
|
| 108 |
-
per_device_train_batch_size=
|
| 109 |
gradient_accumulation_steps=2,
|
| 110 |
learning_rate=2e-4,
|
| 111 |
lr_scheduler_type="cosine",
|
| 112 |
warmup_ratio=0.1,
|
| 113 |
-
logging_steps=
|
| 114 |
eval_strategy="steps",
|
| 115 |
-
eval_steps=
|
| 116 |
save_steps=500,
|
| 117 |
save_total_limit=2,
|
| 118 |
fp16=True,
|
| 119 |
max_length=512,
|
| 120 |
report_to="none",
|
|
|
|
|
|
|
| 121 |
)
|
| 122 |
|
| 123 |
trainer = SFTTrainer(
|
|
|
|
| 13 |
|
| 14 |
import argparse
|
| 15 |
import json
|
| 16 |
+
import sys
|
| 17 |
from pathlib import Path
|
| 18 |
|
| 19 |
+
# Force unbuffered output so HF Spaces logs show immediately
|
| 20 |
+
sys.stdout.reconfigure(line_buffering=True)
|
| 21 |
+
sys.stderr.reconfigure(line_buffering=True)
|
| 22 |
+
|
| 23 |
|
| 24 |
def setup(hf_token: str):
|
| 25 |
"""Login and download data from HF."""
|
|
|
|
| 52 |
return data_dir
|
| 53 |
|
| 54 |
|
| 55 |
+
def train_world_model(data_dir: Path, max_samples: int = 10000):
|
| 56 |
"""Stage 2: Train World Model via SFT."""
|
| 57 |
from datasets import Dataset
|
| 58 |
from transformers import AutoModelForCausalLM, AutoTokenizer
|
|
|
|
| 109 |
|
| 110 |
training_args = SFTConfig(
|
| 111 |
output_dir="./world_model_checkpoints",
|
| 112 |
+
num_train_epochs=2,
|
| 113 |
+
per_device_train_batch_size=16,
|
| 114 |
gradient_accumulation_steps=2,
|
| 115 |
learning_rate=2e-4,
|
| 116 |
lr_scheduler_type="cosine",
|
| 117 |
warmup_ratio=0.1,
|
| 118 |
+
logging_steps=5,
|
| 119 |
eval_strategy="steps",
|
| 120 |
+
eval_steps=100,
|
| 121 |
save_steps=500,
|
| 122 |
save_total_limit=2,
|
| 123 |
fp16=True,
|
| 124 |
max_length=512,
|
| 125 |
report_to="none",
|
| 126 |
+
disable_tqdm=False,
|
| 127 |
+
dataloader_num_workers=0,
|
| 128 |
)
|
| 129 |
|
| 130 |
trainer = SFTTrainer(
|