Skip to content

Commit 8784646

Browse files
committed
fix: allow to set eval datasets, fix shutdown of processes
Signed-off-by: Ettore Di Giacinto <mudler@localai.io>
1 parent 3451dbd commit 8784646

10 files changed

Lines changed: 433 additions & 215 deletions

File tree

backend/python/trl/backend.py

Lines changed: 47 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -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

Comments
 (0)