Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
33 changes: 26 additions & 7 deletions init2winit/checkpoint.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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

Expand Down
8 changes: 8 additions & 0 deletions init2winit/trainer_lib/base_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@
"""Abstract parent class for all trainers."""

import abc
import gc
import itertools
import multiprocessing
import os.path
Expand Down Expand Up @@ -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')
Expand Down Expand Up @@ -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!',
Expand Down
5 changes: 3 additions & 2 deletions init2winit/trainer_lib/training_algorithm.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down Expand Up @@ -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

Expand Down Expand Up @@ -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

Expand Down