Coverage for  / home / jenkins / .local / lib / python3.10 / site-packages / hyper_parallel / core / optimizer / __init__.py: 38%

52 statements  

« prev     ^ index     » next       coverage.py v7.13.1, created at 2026-08-21 04:29 +0800

1# Copyright 2026 Huawei Technologies Co., Ltd 

2# 

3# Licensed under the Apache License, Version 2.0 (the "License"); 

4# you may not use this file except in compliance with the License. 

5# You may obtain a copy of the License at 

6# 

7# http://www.apache.org/licenses/LICENSE-2.0 

8# 

9# Unless required by applicable law or agreed to in writing, software 

10# distributed under the License is distributed on an "AS IS" BASIS, 

11# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. 

12# See the License for the specific language governing permissions and 

13# limitations under the License. 

14# ============================================================================ 

15"""HyperParallel optimizer module.""" 

16from importlib import import_module as _import_module # pylint: disable=invalid-name 

17 

18import inspect 

19import logging 

20from typing import Any, Dict, List, Optional 

21 

22from hyper_parallel.core.optimizer.swap_optimizer import ( 

23 SwapOptimizer, 

24 SwapOptimizerConfig, 

25 swap_optimizer, 

26) 

27 

28logger = logging.getLogger(__name__) 

29logger.setLevel(logging.INFO) 

30 

31# Torch-only optimizer implementations import torch at module load. Keep them 

32# off the eager path so MindSpore-only environments can import SwapOptimizer. 

33_LAZY_EXPORTS = { 

34 "AdamW": ".adamw", 

35 "Muon": ".muon", 

36 "ChainedOptimizer": ".optimizer", 

37 "detect_dtensor_backend": ".dtensor_compat", 

38} 

39 

40 

41def __getattr__(name): # pylint: disable=invalid-name 

42 """Lazily import torch-only optimizer symbols.""" 

43 if name not in _LAZY_EXPORTS: 

44 raise AttributeError(f"module {__name__!r} has no attribute {name!r}") 

45 module = _import_module(_LAZY_EXPORTS[name], __name__) 

46 value = getattr(module, name) 

47 globals()[name] = value 

48 return value 

49 

50 

51def __dir__(): # pylint: disable=invalid-name 

52 """Include lazy torch-only optimizer exports in ``dir()``.""" 

53 return sorted(set(globals()) | set(_LAZY_EXPORTS)) 

54 

55 

56def _load_torch_optimizer_runtime(): 

57 """Import torch-only optimizer helpers used by the factory APIs.""" 

58 # pylint: disable=import-outside-toplevel,unused-import 

59 import hyper_parallel.core.optimizer.utils # noqa: F401 - install rank0 logging helpers 

60 

61 from hyper_parallel.core.optimizer.adamw import AdamW 

62 from hyper_parallel.core.optimizer.dtensor_compat import detect_dtensor_backend 

63 from hyper_parallel.core.optimizer.muon import Muon 

64 from hyper_parallel.core.optimizer.optimizer import ChainedOptimizer 

65 

66 return AdamW, Muon, ChainedOptimizer, detect_dtensor_backend 

67 

68 

69def get_hyper_lr_scheduler(*args: Any, **kwargs: Any) -> Any: 

70 """Create the HyperParallel LR scheduler.""" 

71 # pylint: disable=import-outside-toplevel 

72 from hyper_parallel.core.optimizer.lr_scheduler import ( 

73 get_hyper_lr_scheduler as _get_hyper_lr_scheduler, 

74 ) 

75 return _get_hyper_lr_scheduler(*args, **kwargs) 

76 

77 

78def get_hyper_optimizer( 

79 model: Any, 

80 muon_params: List[Dict[str, Any]], 

81 adamw_params: List[Dict[str, Any]], 

82 muon_kwargs: Optional[Dict[str, Any]] = None, 

83 adamw_kwargs: Optional[Dict[str, Any]] = None, 

84) -> Any: 

85 """Create a chained Muon + AdamW optimizer. 

86 

87 Args: 

88 model: The neural network model. 

89 muon_params: Param groups for Muon. Empty list disables Muon. 

90 adamw_params: Param groups for AdamW. Empty list disables AdamW. 

91 muon_kwargs: Dedicated configurations dict for Muon. 

92 adamw_kwargs: Dedicated configurations dict for AdamW. 

93 

94 Example: 

95 from hyper_parallel.core.optimizer import get_hyper_optimizer 

96 

97 _adamw_legacy = { 

98 'adamw_lr': 1e-3,  

99 'adamw_weight_decay': 1e-2,  

100 'adamw_betas': (0.9, 0.95),  

101 'adamw_eps': 1e-8, 

102 'fused': True 

103 } 

104 _muon_legacy = { 

105 'muon_lr': 2e-2, 

106 'muon_weight_decay': 0.1, 

107 'muon_momentum': 0.95, 

108 'muon_ns_steps': 5, 

109 'muon_ns_variant': 'asym5', 

110 'muon_nesterov': True, 

111 'muon_hsdp_replica_count': 2 

112 } 

113 

114 optimizer = get_hyper_optimizer( 

115 model=model, 

116 muon_params=muon_groups, 

117 adamw_params=adamw_groups, 

118 adamw_kwargs=_adamw_legacy, 

119 muon_kwargs=_muon_legacy, 

120 ) 

121 

122 optimizer.step() 

123 """ 

124 AdamW, Muon, ChainedOptimizer, detect_dtensor_backend = _load_torch_optimizer_runtime() 

125 

126 # 1. Arguments Preparation 

127 # 1.1 adamw 

128 adamw_raw = adamw_kwargs or {} 

129 adamw_config = { 

130 k[6:] if k.startswith("adamw_") else k: v 

131 for k, v in adamw_raw.items() 

132 } 

133 allowed_keys_adamw = inspect.signature(AdamW.__init__).parameters.keys() - {'self', 'params'} 

134 filtered_adamw_config = {k: v for k, v in adamw_config.items() if k in allowed_keys_adamw} 

135 if excluded_adamw_keys := adamw_config.keys() - allowed_keys_adamw: 

136 logger.info_rank0("Excluded adamw config: %s", list(excluded_adamw_keys)) 

137 

138 # 1.2 muon 

139 muon_raw = muon_kwargs or {} 

140 muon_config = { 

141 k[5:] if k.startswith("muon_") else k: v 

142 for k, v in muon_raw.items() 

143 } 

144 allowed_keys_muon = inspect.signature(Muon.__init__).parameters.keys() - {'self', 'params'} 

145 filtered_muon_config = {k: v for k, v in muon_config.items() if k in allowed_keys_muon} 

146 if excluded_muon_keys := muon_config.keys() - allowed_keys_muon: 

147 logger.info_rank0("Excluded muon config: %s", list(excluded_muon_keys)) 

148 

149 # 2. Optimizer Creation 

150 optimizers = {} 

151 detect_dtensor_backend(adamw_params, muon_params) 

152 

153 # build optimizer 

154 if adamw_params: 

155 optimizers["adamw"] = AdamW(adamw_params, **filtered_adamw_config) 

156 logger.info_rank0("Using adamw config: %s", filtered_adamw_config) 

157 

158 if muon_params: 

159 optimizers["muon"] = Muon(muon_params, **filtered_muon_config) 

160 logger.info_rank0("Using muon config: %s", filtered_muon_config) 

161 

162 flatten = bool(adamw_params and muon_params) 

163 

164 return ChainedOptimizer(model, optimizers=optimizers, flatten=flatten) 

165 

166 

167__all__ = [ 

168 'SwapOptimizer', 

169 'SwapOptimizerConfig', 

170 'get_hyper_optimizer', 

171 'get_hyper_lr_scheduler', 

172 'swap_optimizer', 

173]