Rayugacodes commited on
Commit
1572306
·
verified ·
1 Parent(s): 0d9e780

Fix: batch_size=16, 10K samples, unbuffered output, 2 epochs

Browse files
Files changed (1) hide show
  1. 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 = 50000):
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=3,
108
- per_device_train_batch_size=8,
109
  gradient_accumulation_steps=2,
110
  learning_rate=2e-4,
111
  lr_scheduler_type="cosine",
112
  warmup_ratio=0.1,
113
- logging_steps=10,
114
  eval_strategy="steps",
115
- eval_steps=200,
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(