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
« 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# ============================================================================
16"""HyperParallel optimizer module."""
18import inspect
19import logging
20from typing import Any, Dict, List, Optional
22from torch import nn
24import hyper_parallel.core.optimizer.utils # noqa: F401 - install rank0 logging helpers on logging.Logger
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
32logger = logging.getLogger(__name__)
33logger.setLevel(logging.INFO)
35__all__ = ['get_hyper_optimizer', 'get_hyper_lr_scheduler']
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.
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.
54 Example:
55 from hyper_parallel.core.optimizer import get_hyper_optimizer
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 }
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 )
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))
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))
107 # 2. Optimizer Creation
108 optimizers = {}
109 detect_dtensor_backend(adamw_params, muon_params)
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)
116 if muon_params:
117 optimizers["muon"] = Muon(muon_params, **filtered_muon_config)
118 logger.info_rank0("Using muon config: %s", filtered_muon_config)
120 flatten = bool(adamw_params and muon_params)
122 return ChainedOptimizer(model, optimizers=optimizers, flatten=flatten)