core.models.mimo.optimizer#
Optimizer for MIMO models with heterogeneous parallelism.
Module Contents#
Classes#
Optimizer info for a single module. |
|
Optimizer for MimoModel with heterogeneous parallelism. |
Functions#
Yield (sub_state_dict, inner_optimizer) pairs. |
|
Save: extract param_groups from optimizer sub-dict into a ShardedObject. |
|
Save: extract grad_scaler into a ShardedObject. |
|
Save: extract param_state_sharding_type into a ShardedObject. |
|
Load: restore param_groups with current param IDs from the inner optimizer. |
|
Load: restore param_state_sharding_type from ShardedObject key. |
|
Load: restore grad_scaler from ShardedObject key. |
|
Build replica_id tuple for ShardedObject deduplication. |
|
Derive the optimizer’s ProcessGroupCollection from a populated HyperCommGrid. |
|
Create optimizer for MimoModel with heterogeneous parallelism. |
Data#
API#
- class core.models.mimo.optimizer.ModuleOptimizerInfo#
Optimizer info for a single module.
- optimizer: Optional[megatron.core.optimizer.optimizer.MegatronOptimizer]#
None
- grid: Optional[megatron.core.hyper_comm_grid.HyperCommGrid]#
None
- pg_collection: Optional[megatron.core.process_groups_config.ProcessGroupCollection]#
None
- is_active: bool#
None
- class core.models.mimo.optimizer.MimoOptimizer(
- module_infos: Dict[str, core.models.mimo.optimizer.ModuleOptimizerInfo],
- config: megatron.core.optimizer.optimizer_config.OptimizerConfig,
Bases:
megatron.core.optimizer.optimizer.MegatronOptimizerOptimizer for MimoModel with heterogeneous parallelism.
Each module gets its own optimizer. Global gradient norm is computed across all modules via all_reduce MAX.
Initialization
Input optimizer is the base optimizer (e.g., Adam).
- prepare_grads() bool#
Prepare gradients for all active module optimizers.
- get_grad_norm() float#
Compute global gradient norm across all modules via all_reduce MAX.
- step() Tuple[bool, Optional[float], Optional[int]]#
Run one optimizer step across all active module optimizers.
- step_with_ready_grads() bool#
Step active optimizers after gradients have been prepared.
- zero_grad(set_to_none: bool = True)#
Clear gradients on all active module optimizers.
- prepare_model_params_for_param_sync() None#
Stage parameters for explicit synchronization in all active module optimizers.
- get_loss_scale() torch.Tensor#
Return the loss scale tensor from the first active optimizer.
- count_zeros() int#
Count zero gradients per module (world-MAX so disjoint grids agree), then sum.
- property param_groups: List[dict]#
Combined param groups from all active module optimizers.
- state_dict()#
Return per-module optimizer state dicts.
- load_state_dict(state_dict: Dict)#
Load per-module optimizer state dicts.
Reassembles param_groups and grad_scaler that were extracted and saved as ShardedObjects by sharded_state_dict(), then delegates to each per-module optimizer’s load_state_dict.
- sharded_state_dict(
- model_sharded_state_dict,
- is_loading: bool = False,
- **kwargs,
Build sharded state dict, routing param_groups and grad_scaler through distributed save as ShardedObjects (common.pt is rank-0 only, which misses LLM optimizer state in non-colocated mode).
- reload_model_params(state_dict=None)#
Reload model parameters in all active module optimizers.
- core.models.mimo.optimizer._iter_optimizer_sub_dicts(module_sd, optimizer)#
Yield (sub_state_dict, inner_optimizer) pairs.
For a single optimizer, yields (module_sd, optimizer) once. For ChainedOptimizer with N>1 inner optimizers, yields (module_sd[i], chained_optimizers[i]) for each.
- core.models.mimo.optimizer._extract_param_groups(sub_sd, module_name, suffix, replica_id)#
Save: extract param_groups from optimizer sub-dict into a ShardedObject.
- core.models.mimo.optimizer._extract_grad_scaler(sub_sd, module_name, suffix, replica_id)#
Save: extract grad_scaler into a ShardedObject.
- core.models.mimo.optimizer._extract_param_state_sharding_type(
- sub_sd,
- module_name,
- suffix,
- replica_id,
Save: extract param_state_sharding_type into a ShardedObject.
- core.models.mimo.optimizer._restore_param_groups(sub_sd, inner_optimizer, module_name)#
Load: restore param_groups with current param IDs from the inner optimizer.
- core.models.mimo.optimizer._restore_param_state_sharding_type(sub_sd)#
Load: restore param_state_sharding_type from ShardedObject key.
- core.models.mimo.optimizer._restore_grad_scaler(sub_sd)#
Load: restore grad_scaler from ShardedObject key.
- core.models.mimo.optimizer._get_replica_id(
- pg_collection: Optional[megatron.core.process_groups_config.ProcessGroupCollection],
Build replica_id tuple for ShardedObject deduplication.
Returns (tp_rank, pp_rank, dp_rank) so only (0, 0, 0) within each module’s parallelism group is the main replica; all other ranks in the same module are non-main replicas of the same object.
- core.models.mimo.optimizer._EXPERT_VIEW#
‘expert’
- core.models.mimo.optimizer._get_pg_collection_for_optimizer(
- grid,
Derive the optimizer’s ProcessGroupCollection from a populated HyperCommGrid.
Dense groups come from the base view; expert-parallel groups (tp_ep_pp, expt_dp) come from the grid’s dedicated expert view – expert parallelism is always factored into a separate view (expt_tp/ep/expt_dp), never the base view. All groups must be pre-created on the grid.
- core.models.mimo.optimizer.get_mimo_optimizer(
- mimo_model: MimoModel,
- config: megatron.core.optimizer.optimizer_config.OptimizerConfig,
Create optimizer for MimoModel with heterogeneous parallelism.