Skip to content
Merged
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
10 changes: 9 additions & 1 deletion deepspeed/runtime/engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -727,6 +727,13 @@ def zero_optimization_partition_gradients(self):
def zero_optimization_partition_weights(self):
return self.zero_optimization_stage() >= ZeroStageEnum.weights

def is_first_weights_partition_group(self):
ret = True if self.mics_shard_size() < 0 \
and self.zero_optimization_partition_weights() else False
if self.mics_shard_size() > 0 and self.global_rank < self.mics_shard_size():
ret = True
return ret

def zero_contiguous_gradients(self):
return self._config.zero_config.contiguous_gradients

Expand Down Expand Up @@ -906,7 +913,8 @@ def _configure_checkpointing(self, dist_init_required):
# only the first data parallel process needs to store the model checkpoint
# if you want to use node local storage this must be done by rank 0 on each
# node
self.save_non_zero_checkpoint = (rank == 0) or self.zero_optimization_partition_weights()
self.save_non_zero_checkpoint = (rank == 0) or (self.zero_optimization_partition_weights()
and self.is_first_weights_partition_group())

if self.zero_optimization() or self.bfloat16_enabled():
param_rank = dist.get_rank(group=self.optimizer.dp_process_group)
Expand Down
14 changes: 14 additions & 0 deletions tests/unit/checkpoint/test_mics_optimizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -64,3 +64,17 @@ def test_not_load_optimizer_state(self, tmpdir, shard_size):
def test_load_module_only(self, tmpdir, shard_size):
config_dict, hidden_dim, models = self._toy_model_config(shard_size)
checkpoint_correctness_verification(config_dict, models, hidden_dim, tmpdir, load_module_only=True)

@pytest.mark.parametrize('shard_size', [1, 2, 4])
def test_save_checkpoint_on_first_partition_group(self, tmpdir, shard_size):
config_dict, _, models = self._toy_model_config(shard_size)
ds_engine, _, _, _ = deepspeed.initialize(config=config_dict,
model=models[0],
model_parameters=models[0].parameters(),
optimizer=None)

ds_engine.save_checkpoint(tmpdir)
if ds_engine.global_rank < shard_size:
assert ds_engine.save_non_zero_checkpoint == True
else:
assert ds_engine.save_non_zero_checkpoint == False