From 8cc1f10bbd2c53b2bc067bfcf696ff7cb5ee67c3 Mon Sep 17 00:00:00 2001 From: AdrianAbeyta Date: Wed, 7 Sep 2022 22:46:12 +0000 Subject: [PATCH] removed hardcoded warmup steps, use global steps in place of max_steps --- src/transformers/trainer.py | 4 ++-- src/transformers/training_args.py | 2 +- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/src/transformers/trainer.py b/src/transformers/trainer.py index 66879d659654..d4e4a053ceb8 100755 --- a/src/transformers/trainer.py +++ b/src/transformers/trainer.py @@ -1876,8 +1876,8 @@ def _inner_training_loop( metrics = speed_metrics("train", start_time, num_samples=num_train_samples, num_steps=self.state.max_steps) - total_samples = args.max_steps*total_train_batch_size if args.max_steps > 0 else num_examples*num_train_epochs - perf_samples = total_samples - 10*total_train_batch_size + total_samples = self.state.global_step*total_train_batch_size if args.max_steps > 0 else num_examples*num_train_epochs + perf_samples = total_samples - self.args.warmup_steps*total_train_batch_size stable_train_metrics = speed_metrics("stable_train", start_train_stable_time, perf_samples) self.store_flos() diff --git a/src/transformers/training_args.py b/src/transformers/training_args.py index e662d6fca4fd..158cc893e315 100644 --- a/src/transformers/training_args.py +++ b/src/transformers/training_args.py @@ -568,7 +568,7 @@ class TrainingArguments: warmup_ratio: float = field( default=0.0, metadata={"help": "Linear warmup over warmup_ratio fraction of total steps."} ) - warmup_steps: int = field(default=0, metadata={"help": "Linear warmup over warmup_steps."}) + warmup_steps: int = field(default=10, metadata={"help": "Linear warmup over warmup_steps."}) log_level: Optional[str] = field( default="passive",