core.optimizer.fully_sharded_optimizer#

MCore optimizer wrapper for experimental Megatron-FSDP v2.

Module Contents#

Classes#

FullyShardedOptimizer

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.MixedPrecisionOptimizer

MCore 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],
) None#

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,
) core.dist_checkpointing.mapping.ShardedStateDict#

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.