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.full_optim_state_dict(model, optimizer)
