core.optimizer.fully_sharded_optimizer#
MCore optimizer wrapper for experimental Megatron-FSDP v2.
Module Contents#
Classes#
MCore optimizer wrapper for MFSDP-owned sharded parameters and gradients. |
API#
- class core.optimizer.fully_sharded_optimizer.FullyShardedOptimizer(
- optimizer: torch.optim.Optimizer,
- config: core.optimizer.optimizer_config.OptimizerConfig,
- grad_scaler: Optional[core.optimizer.grad_scaler.MegatronGradScaler],
- init_state_fn: Callable,
- model_chunks: List[core.transformer.module.MegatronModule],
Bases:
core.optimizer.optimizer.MixedPrecisionOptimizerMCore optimizer wrapper for MFSDP-owned sharded parameters and gradients.
MFSDP v2 owns the optimizer-facing parameter and gradient shards directly. Unlike :class:
DistributedOptimizer, this wrapper does not build DDP param-and-grad-buffer range maps or allocate separate main-parameter shards. It preserves MCore’s mixed-precision optimizer step contract while making MFSDP-specific storage operations explicit.Initialization
Initialize the MFSDP optimizer wrapper.
- Parameters:
optimizer – Base optimizer such as Adam or SGD.
config – Optimizer configuration.
grad_scaler – Optional loss scaler. Currently unsupported for MFSDP v2, but accepted to match the MCore optimizer construction contract.
init_state_fn – Function used to initialize optimizer state.
model_chunks – MFSDP v2 model chunks optimized by this wrapper.
- static _validate_config(
- config: core.optimizer.optimizer_config.OptimizerConfig,
- model_chunks: List[core.transformer.module.MegatronModule],
Validate the MFSDP v2 optimizer support contract.
- abstractmethod state_dict()#
Return optimizer state.
MFSDP v2 optimizer checkpointing needs an FSDP-native DTensor state contract. Keep this intentionally unsupported for the prototype instead of falling back to DDP-buffer assumptions.
- abstractmethod load_state_dict(state_dict)#
Load optimizer state.
- abstractmethod sharded_state_dict(
- model_sharded_state_dict: core.dist_checkpointing.mapping.ShardedStateDict,
- is_loading: bool = False,
- metadata: Optional[dict] = None,
Build a sharded optimizer state dict.
- zero_grad(set_to_none: bool = True) None#
Clear optimizer-visible sharded grads and any grads filtered from local groups.
- _copy_model_grads_to_main_grads() None#
No-op: MFSDP v2 reduces directly into optimizer-visible sharded grads.
- _copy_main_params_to_model_params() None#
No-op: MFSDP v2 currently syncs compute weights in its forward pre-hook.
- _copy_model_params_to_main_params(state_dict=None) None#
No-op: model loads already write into MFSDP v2’s main weights.