latest (dev)
Copy
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: OptimStateDictConfigHelp improve this page
Found an error, an unclear step, or a missing example?
Was this page helpful?

