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

36 statements  

« prev     ^ index     » next       coverage.py v7.13.1, created at 2026-08-04 05:18 +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 

16"""HyperParallel optimizer module.""" 

17 

18import inspect 

19import logging 

20from typing import Any, Dict, List, Optional 

21 

22from torch import nn 

23 

24import hyper_parallel.core.optimizer.utils # noqa: F401 - install rank0 logging helpers on logging.Logger 

25 

26from hyper_parallel.core.optimizer.adamw import AdamW 

27from hyper_parallel.core.optimizer.lr_scheduler import get_hyper_lr_scheduler 

28from hyper_parallel.core.optimizer.muon import Muon 

29from hyper_parallel.core.optimizer.optimizer import ChainedOptimizer 

30from hyper_parallel.core.optimizer.dtensor_compat import detect_dtensor_backend 

31 

32logger = logging.getLogger(__name__) 

33logger.setLevel(logging.INFO) 

34 

35__all__ = ['get_hyper_optimizer', 'get_hyper_lr_scheduler'] 

36 

37 

38def get_hyper_optimizer( 

39 model: nn.Module, 

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

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

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

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

44) -> ChainedOptimizer: 

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

46 

47 Args: 

48 model: The neural network model. 

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

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

51 muon_kwargs: Dedicated configurations dict for Muon. 

52 adamw_kwargs: Dedicated configurations dict for AdamW. 

53 

54 Example: 

55 from hyper_parallel.core.optimizer import get_hyper_optimizer 

56 

57 _adamw_legacy = { 

58 'adamw_lr': 1e-3,  

59 'adamw_weight_decay': 1e-2,  

60 'adamw_betas': (0.9, 0.95),  

61 'adamw_eps': 1e-8, 

62 'fused': True 

63 } 

64 _muon_legacy = { 

65 'muon_lr': 2e-2, 

66 'muon_weight_decay': 0.1, 

67 'muon_momentum': 0.95, 

68 'muon_ns_steps': 5, 

69 'muon_ns_variant': 'asym5', 

70 'muon_nesterov': True, 

71 'muon_hsdp_replica_count': 2 

72 } 

73 

74 optimizer = get_hyper_optimizer( 

75 model=model, 

76 muon_params=muon_groups, 

77 adamw_params=adamw_groups, 

78 adamw_kwargs=_adamw_legacy, 

79 muon_kwargs=_muon_legacy, 

80 ) 

81 

82 optimizer.step() 

83 """ 

84 # 1. Arguments Preparation 

85 # 1.1 adamw 

86 adamw_raw = adamw_kwargs or {} 

87 adamw_config = { 

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

89 for k, v in adamw_raw.items() 

90 } 

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

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

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

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

95 

96 # 1.2 muon 

97 muon_raw = muon_kwargs or {} 

98 muon_config = { 

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

100 for k, v in muon_raw.items() 

101 } 

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

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

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

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

106 

107 # 2. Optimizer Creation 

108 optimizers = {} 

109 detect_dtensor_backend(adamw_params, muon_params) 

110 

111 # build optimizer 

112 if adamw_params: 

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

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

115 

116 if muon_params: 

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

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

119 

120 flatten = bool(adamw_params and muon_params) 

121 

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