diff --git a/init2winit/checkpoint.py b/init2winit/checkpoint.py index 9c48fc19..de541781 100644 --- a/init2winit/checkpoint.py +++ b/init2winit/checkpoint.py @@ -19,6 +19,7 @@ are nested numpy arrays. """ +import gc from absl import flags from absl import logging from init2winit.dataset_lib import data_utils @@ -146,14 +147,32 @@ def maybe_restore_checkpoint( False, ) # is_restored + restored_optimizer_state = ckpt_to_return['optimizer_state'] + restored_params = ckpt_to_return['params'] + restored_batch_stats = ckpt_to_return['batch_stats'] + restored_metrics_state = ckpt_to_return['training_metrics_grabber'] + restored_global_step = ckpt_to_return['global_step'] + restored_sum_train_cost = ckpt_to_return['sum_train_cost'] + restored_preemption_count = ckpt_to_return['preemption_count'] + + del latest_ckpt + del ckpt_to_return + del unreplicated_checkpoint_state + del unwrapped_optimizer_state + del unreplicated_optimizer_state + del unreplicated_params + del unreplicated_batch_stats + del unreplicated_training_metrics_state + gc.collect() + return ( - ckpt_to_return['optimizer_state'], - ckpt_to_return['params'], - ckpt_to_return['batch_stats'], - ckpt_to_return['training_metrics_grabber'], - ckpt_to_return['global_step'], # global_step - ckpt_to_return['sum_train_cost'], - ckpt_to_return['preemption_count'], # preemption_count + restored_optimizer_state, + restored_params, + restored_batch_stats, + restored_metrics_state, + restored_global_step, + restored_sum_train_cost, + restored_preemption_count, is_restored, ) # is_restored diff --git a/init2winit/trainer_lib/base_trainer.py b/init2winit/trainer_lib/base_trainer.py index 22252fbf..4fc5f42e 100644 --- a/init2winit/trainer_lib/base_trainer.py +++ b/init2winit/trainer_lib/base_trainer.py @@ -16,6 +16,7 @@ """Abstract parent class for all trainers.""" import abc +import gc import itertools import multiprocessing import os.path @@ -690,6 +691,12 @@ def setup_and_maybe_restore(self, init_rng, data_rng, callback_rng): 'Training state sharded in %f seconds', time.time() - start_time ) + del unreplicated_params + del unreplicated_optimizer_state + del unreplicated_batch_stats + del unreplicated_metrics_state + gc.collect() + self._dataset = self.setup_data_loader(data_rng, self._global_step) self._eval_callbacks = self._setup_eval_callbacks(callback_rng) logging.info('Training state setup complete') @@ -719,6 +726,7 @@ def train(self): logging.info('Hyperparameters: %s', self._hps) self.setup_and_maybe_restore(init_rng, data_rng, callback_rng) + gc.collect() trainer_utils.log_message( 'Setup and maybe restore completed!', diff --git a/init2winit/trainer_lib/training_algorithm.py b/init2winit/trainer_lib/training_algorithm.py index 8857b319..2d3f8435 100644 --- a/init2winit/trainer_lib/training_algorithm.py +++ b/init2winit/trainer_lib/training_algorithm.py @@ -209,6 +209,8 @@ def restore_optimizer_state(self, optimizer_state): Returns: The post-processed optimizer state, ready for sharding. """ + if hasattr(self, '_optimizer_state'): + self._optimizer_state = None return optimizer_state @@ -945,7 +947,6 @@ def update_params( grad_norm=grad_norm.item(), update_norm=update_norm.item(), ) - self._optimizer_state = new_optimizer_state return new_optimizer_state, new_params, new_batch_stats, cost_value, grad @@ -988,7 +989,7 @@ def init_optimizer_state( # Wrapping init in jax.jit fuses per-parameter state creation ops into # a single compilation instead of compiling each one individually. optax_optimizer_state = jax.jit(optimizer_init_fn)(params) - self._optimizer_state = optax_optimizer_state + self._optimizer_state = None self._update_fn = optax_optimizer_update_fn return optax_optimizer_state