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
« 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
18import inspect
19import logging
20from typing import Any, Dict, List, Optional
22from hyper_parallel.core.optimizer.swap_optimizer import (
23 SwapOptimizer,
24 SwapOptimizerConfig,
25 swap_optimizer,
26)
28logger = logging.getLogger(__name__)
29logger.setLevel(logging.INFO)
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}
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
51def __dir__(): # pylint: disable=invalid-name
52 """Include lazy torch-only optimizer exports in ``dir()``."""
53 return sorted(set(globals()) | set(_LAZY_EXPORTS))
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
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
66 return AdamW, Muon, ChainedOptimizer, detect_dtensor_backend
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)
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.
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.
94 Example:
95 from hyper_parallel.core.optimizer import get_hyper_optimizer
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 }
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 )
122 optimizer.step()
123 """
124 AdamW, Muon, ChainedOptimizer, detect_dtensor_backend = _load_torch_optimizer_runtime()
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))
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))
149 # 2. Optimizer Creation
150 optimizers = {}
151 detect_dtensor_backend(adamw_params, muon_params)
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)
158 if muon_params:
159 optimizers["muon"] = Muon(muon_params, **filtered_muon_config)
160 logger.info_rank0("Using muon config: %s", filtered_muon_config)
162 flatten = bool(adamw_params and muon_params)
164 return ChainedOptimizer(model, optimizers=optimizers, flatten=flatten)
167__all__ = [
168 'SwapOptimizer',
169 'SwapOptimizerConfig',
170 'get_hyper_optimizer',
171 'get_hyper_lr_scheduler',
172 'swap_optimizer',
173]