from torch.distributed.fsdp import ( FullyShardedDataParallel as FSDP, FullStateDictConfig, StateDictType, ) from torch.nn.parallel import DistributedDataParallel as DDP fullstate_save_policy = FullStateDictConfig(offload_to_cpu=True, rank0_only=True) def _is_fsdp(model): return isinstance(model, FSDP) def _is_ddp(model): return isinstance(model, DDP) def get_model_state_dict_on_rank_0( model, rank, ): """Returns the model state dict only on rank 0 (called from all workers)""" if _is_fsdp(model): model_state = _get_model_state_dict_for_fsdp_rank_0(model, rank) else: model_state = None if rank == 0: if _is_ddp(model): model_state = model.module.state_dict() else: model_state = model.state_dict() return model_state def get_optimizer_state_dict_on_rank_0( model, optimizer, rank, ): """Returns the optimizer state dict only on rank 0 (called from all workers)""" if _is_fsdp(model): optim_state = _get_optimizer_state_dict_for_fsdp_rank_0(model, optimizer) else: optim_state = optimizer.state_dict() if rank == 0 else None return optim_state def _get_model_state_dict_for_fsdp_rank_0(model, rank): with FSDP.state_dict_type(model, StateDictType.FULL_STATE_DICT, fullstate_save_policy): model_state = model.state_dict() return model_state def _get_optimizer_state_dict_for_fsdp_rank_0(model, optimizer): return FSDP.optim_state_dict(model, optimizer)