diff --git a/caikit_nlp/modules/text_generation/peft_prompt_tuning.py b/caikit_nlp/modules/text_generation/peft_prompt_tuning.py index 435e439dc..a0135cc5b 100644 --- a/caikit_nlp/modules/text_generation/peft_prompt_tuning.py +++ b/caikit_nlp/modules/text_generation/peft_prompt_tuning.py @@ -490,7 +490,9 @@ def train( num_epochs, cls.RANDOM_SEED, learning_rate, - max_steps=infer_max_steps(num_epochs, batch_size, training_dataset), + max_steps=infer_max_steps( + num_epochs, batch_size, accumulate_steps, training_dataset + ), silence_progress_bars=silence_progress_bars, accumulate_steps=accumulate_steps, # NOTE: following can override above arguments in order diff --git a/caikit_nlp/toolkit/text_generation/training_utils.py b/caikit_nlp/toolkit/text_generation/training_utils.py index ac905fadd..411d75cf7 100644 --- a/caikit_nlp/toolkit/text_generation/training_utils.py +++ b/caikit_nlp/toolkit/text_generation/training_utils.py @@ -223,8 +223,10 @@ def launch_training( def infer_max_steps( num_epochs: int, batch_size: int, + ga_steps: int, training_dataset: Union[Dataset, TransformersIterableDataset], ): + batch_size = batch_size * ga_steps # Calculate the number of samples that we have if isinstance(training_dataset, Dataset): data_len = len(training_dataset)