TensorPlay
latest (dev)
Copy
View Markdown

Latest development documentation · Updated 2026-10-08

Source code for tensorplay.distributed.fsdp.api

"""Configuration objects for fully sharded module execution."""

from collections.abc import Sequence
from dataclasses import dataclass
from enum import Enum, auto
from typing import Any

from tensorplay.nn.modules.batchnorm import _BatchNorm

__all__ = [
    "ShardingStrategy",
    "BackwardPrefetch",
    "MixedPrecision",
    "CPUOffload",
    "StateDictType",
    "StateDictConfig",
    "FullStateDictConfig",
    "LocalStateDictConfig",
    "ShardedStateDictConfig",
    "OptimStateDictConfig",
    "FullOptimStateDictConfig",
    "LocalOptimStateDictConfig",
    "ShardedOptimStateDictConfig",
    "StateDictSettings",
]



[docs]
class ShardingStrategy(Enum):
    FULL_SHARD = auto()
    SHARD_GRAD_OP = auto()
    NO_SHARD = auto()
    HYBRID_SHARD = auto()
    _HYBRID_SHARD_ZERO2 = auto()




[docs]
class BackwardPrefetch(Enum):
    BACKWARD_PRE = auto()
    BACKWARD_POST = auto()




[docs]
@dataclass
class MixedPrecision:
    param_dtype: Any = None
    reduce_dtype: Any = None
    buffer_dtype: Any = None
    keep_low_precision_grads: bool = False
    cast_forward_inputs: bool = False
    cast_root_forward_inputs: bool = True
    _module_classes_to_ignore: Sequence[type] = (_BatchNorm,)




[docs]
@dataclass
class CPUOffload:
    offload_params: bool = False




[docs]
class StateDictType(Enum):
    FULL_STATE_DICT = auto()
    LOCAL_STATE_DICT = auto()
    SHARDED_STATE_DICT = auto()




[docs]
@dataclass
class StateDictConfig:
    offload_to_cpu: bool = False




[docs]
@dataclass
class FullStateDictConfig(StateDictConfig):
    rank0_only: bool = False




[docs]
@dataclass
class LocalStateDictConfig(StateDictConfig):
    pass




[docs]
@dataclass
class ShardedStateDictConfig(StateDictConfig):
    _use_dtensor: bool = False




[docs]
@dataclass
class OptimStateDictConfig:
    offload_to_cpu: bool = True




[docs]
@dataclass
class FullOptimStateDictConfig(OptimStateDictConfig):
    rank0_only: bool = False




[docs]
@dataclass
class LocalOptimStateDictConfig(OptimStateDictConfig):
    offload_to_cpu: bool = False




[docs]
@dataclass
class ShardedOptimStateDictConfig(OptimStateDictConfig):
    _use_dtensor: bool = False




[docs]
@dataclass
class StateDictSettings:
    state_dict_type: StateDictType
    state_dict_config: StateDictConfig
    optim_state_dict_config: OptimStateDictConfig
Ask DeepWiki