mirror of
https://github.com/wassname/alignment-handbook.git
synced 2026-08-09 11:50:32 +08:00
Clean deprecated max_samples arguments (#89)
This commit is contained in:
+2
-4
@@ -151,8 +151,7 @@ def main():
|
||||
logger.info("*** Train ***")
|
||||
train_result = trainer.train()
|
||||
metrics = train_result.metrics
|
||||
max_train_samples = data_args.max_train_samples if data_args.max_train_samples is not None else len(train_dataset)
|
||||
metrics["train_samples"] = min(max_train_samples, len(train_dataset))
|
||||
metrics["train_samples"] = len(train_dataset)
|
||||
trainer.log_metrics("train", metrics)
|
||||
trainer.save_metrics("train", metrics)
|
||||
trainer.save_state()
|
||||
@@ -163,8 +162,7 @@ def main():
|
||||
if training_args.do_eval:
|
||||
logger.info("*** Evaluate ***")
|
||||
metrics = trainer.evaluate()
|
||||
max_eval_samples = data_args.max_eval_samples if data_args.max_eval_samples is not None else len(eval_dataset)
|
||||
metrics["eval_samples"] = min(max_eval_samples, len(eval_dataset))
|
||||
metrics["eval_samples"] = len(eval_dataset)
|
||||
trainer.log_metrics("eval", metrics)
|
||||
trainer.save_metrics("eval", metrics)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user