@@ -256,6 +256,39 @@ def _do_training(self, request, job):
256256 else :
257257 dataset = load_dataset (request .dataset_source , split = dataset_split )
258258
259+ # Evaluation dataset setup
260+ save_steps = request .save_steps if request .save_steps > 0 else 500
261+ eval_dataset = None
262+ eval_strategy = extra .get ("eval_strategy" , "steps" )
263+ eval_steps = int (extra .get ("eval_steps" , str (save_steps )))
264+
265+ if eval_strategy != "no" :
266+ eval_split = extra .get ("eval_split" )
267+ eval_dataset_source = extra .get ("eval_dataset_source" )
268+ if eval_split :
269+ # Load a specific split as eval dataset
270+ if os .path .exists (request .dataset_source ):
271+ if request .dataset_source .endswith ('.json' ) or request .dataset_source .endswith ('.jsonl' ):
272+ eval_dataset = load_dataset ("json" , data_files = request .dataset_source , split = eval_split )
273+ elif request .dataset_source .endswith ('.csv' ):
274+ eval_dataset = load_dataset ("csv" , data_files = request .dataset_source , split = eval_split )
275+ else :
276+ eval_dataset = load_dataset (request .dataset_source , split = eval_split )
277+ else :
278+ eval_dataset = load_dataset (request .dataset_source , split = eval_split )
279+ elif eval_dataset_source :
280+ # Load eval dataset from a separate source
281+ eval_dataset = load_dataset (eval_dataset_source , split = "train" )
282+ else :
283+ # Auto-split the training set
284+ eval_split_ratio = float (extra .get ("eval_split_ratio" , "0.1" ))
285+ split = dataset .train_test_split (test_size = eval_split_ratio )
286+ dataset = split ["train" ]
287+ eval_dataset = split ["test" ]
288+
289+ if eval_strategy == "no" :
290+ eval_dataset = None
291+
259292 # Training config
260293 output_dir = request .output_dir or f"./output-{ job .job_id } "
261294 num_epochs = request .num_epochs if request .num_epochs > 0 else 3
@@ -265,7 +298,6 @@ def _do_training(self, request, job):
265298 warmup_steps = request .warmup_steps if request .warmup_steps > 0 else 5
266299 weight_decay = request .weight_decay if request .weight_decay > 0 else 0.01
267300 max_steps = request .max_steps if request .max_steps > 0 else - 1
268- save_steps = request .save_steps if request .save_steps > 0 else 500
269301 seed = request .seed if request .seed > 0 else 3407
270302 optimizer = request .optimizer or "adamw_torch"
271303
@@ -308,6 +340,12 @@ def _do_training(self, request, job):
308340 if save_total_limit :
309341 _save_kwargs ["save_total_limit" ] = save_total_limit
310342
343+ # Eval arguments
344+ _eval_kwargs = {}
345+ if eval_dataset is not None :
346+ _eval_kwargs ["eval_strategy" ] = eval_strategy
347+ _eval_kwargs ["eval_steps" ] = eval_steps
348+
311349 # Common training arguments shared by all methods
312350 _common_args = dict (
313351 output_dir = output_dir ,
@@ -323,6 +361,7 @@ def _do_training(self, request, job):
323361 logging_steps = 1 ,
324362 report_to = "none" ,
325363 ** _save_kwargs ,
364+ ** _eval_kwargs ,
326365 ** common_train_kwargs ,
327366 )
328367
@@ -343,6 +382,7 @@ def _do_training(self, request, job):
343382 model = model ,
344383 args = training_args ,
345384 train_dataset = dataset ,
385+ eval_dataset = eval_dataset ,
346386 processing_class = tokenizer ,
347387 callbacks = [progress_cb .get_callback ()],
348388 )
@@ -365,6 +405,7 @@ def _do_training(self, request, job):
365405 model = model ,
366406 args = training_args ,
367407 train_dataset = dataset ,
408+ eval_dataset = eval_dataset ,
368409 processing_class = tokenizer ,
369410 callbacks = [progress_cb .get_callback ()],
370411 )
@@ -399,6 +440,7 @@ def _do_training(self, request, job):
399440 model = model ,
400441 args = training_args ,
401442 train_dataset = dataset ,
443+ eval_dataset = eval_dataset ,
402444 processing_class = tokenizer ,
403445 reward_funcs = reward_funcs ,
404446 callbacks = [progress_cb .get_callback ()],
@@ -420,6 +462,7 @@ def _do_training(self, request, job):
420462 model = model ,
421463 args = training_args ,
422464 train_dataset = dataset ,
465+ eval_dataset = eval_dataset ,
423466 processing_class = tokenizer ,
424467 callbacks = [progress_cb .get_callback ()],
425468 )
@@ -440,6 +483,7 @@ def _do_training(self, request, job):
440483 model = model ,
441484 args = training_args ,
442485 train_dataset = dataset ,
486+ eval_dataset = eval_dataset ,
443487 processing_class = tokenizer ,
444488 callbacks = [progress_cb .get_callback ()],
445489 )
@@ -460,6 +504,7 @@ def _do_training(self, request, job):
460504 model = model ,
461505 args = training_args ,
462506 train_dataset = dataset ,
507+ eval_dataset = eval_dataset ,
463508 processing_class = tokenizer ,
464509 callbacks = [progress_cb .get_callback ()],
465510 )
@@ -478,6 +523,7 @@ def _do_training(self, request, job):
478523 model = model ,
479524 args = training_args ,
480525 train_dataset = dataset ,
526+ eval_dataset = eval_dataset ,
481527 processing_class = tokenizer ,
482528 callbacks = [progress_cb .get_callback ()],
483529 )
0 commit comments