Coverage for / home / jenkins / .local / lib / python3.10 / site-packages / hyper_parallel / core / optimizer / adamw.py: 0%
55 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# ============================================================================
15"""Adamw optimizer."""
17from typing import List
19import torch
21from hyper_parallel.core.optimizer.utils import get_current_device
24def adamw(
25 params: List[torch.Tensor],
26 grads: List[torch.Tensor],
27 exp_avgs: List[torch.Tensor],
28 exp_avg_sqs: List[torch.Tensor],
29 max_exp_avg_sqs: List[torch.Tensor],
30 step: int,
31 *,
32 amsgrad: bool,
33 beta1: float,
34 beta2: float,
35 lr: float,
36 weight_decay: float,
37 eps: float,
38 maximize: bool
39) -> None:
40 r"""Functional API that performs AdamW algorithm computation.
41 See :class:`~torch.optim.AdamW` for details.
42 """
43 device = get_current_device()
44 step_tensor = torch.tensor(step, dtype=torch.int64, device=device)
45 state_steps = [step_tensor] * len(params)
47 torch._fused_adamw_( # pylint: disable=protected-access
48 params,
49 grads,
50 exp_avgs,
51 exp_avg_sqs,
52 max_exp_avg_sqs if amsgrad else [],
53 state_steps,
54 amsgrad=amsgrad,
55 lr=lr,
56 beta1=beta1,
57 beta2=beta2,
58 weight_decay=weight_decay,
59 eps=eps,
60 maximize=maximize
61 )
64class AdamW(torch.optim.Optimizer):
65 """AdamW optimizer implementation."""
67 def __init__(
68 self,
69 params,
70 lr=1e-3,
71 betas=(0.9, 0.999),
72 eps=1e-8,
73 weight_decay=0.01,
74 amsgrad=False,
75 maximize=False
76 ):
77 defaults = {
78 "lr": lr,
79 "betas": betas,
80 "eps": eps,
81 "weight_decay": weight_decay,
82 "amsgrad": amsgrad,
83 "maximize": maximize
84 }
85 super().__init__(params, defaults)
87 def __setstate__(self, state):
88 """Set optimizer state."""
89 super().__setstate__(state)
90 for group in self.param_groups:
91 group.setdefault('amsgrad', False)
92 group.setdefault('maximize', False)
94 def __str__(self):
95 return super().__repr__()
97 __repr__ = __str__
99 def step(self, closure=None):
100 """Performs a single optimization step."""
101 loss = None
102 if closure is not None:
103 with torch.enable_grad():
104 loss = closure()
106 for group in self.param_groups:
107 params_with_grad = []
108 grads = []
109 exp_avgs = []
110 exp_avg_sqs = []
111 max_exp_avg_sqs = []
113 amsgrad = group['amsgrad']
114 beta1, beta2 = group['betas']
115 group['step'] = (group.get('step') or 0) + 1
117 current_rank_params = group['params']
118 for p in current_rank_params:
119 if p.grad is None:
120 continue
122 if p.grad.data.is_sparse:
123 raise RuntimeError('AdamW does not support sparse gradients')
125 state = self.state[p]
127 if len(state) == 0:
128 state['exp_avg'] = torch.zeros_like(p.grad, memory_format=torch.preserve_format)
129 state['exp_avg_sq'] = torch.zeros_like(p.grad, memory_format=torch.preserve_format)
130 if amsgrad:
131 state['max_exp_avg_sq'] = torch.zeros_like(p.grad, memory_format=torch.preserve_format)
133 params_with_grad.append(p)
134 grads.append(p.grad)
135 exp_avgs.append(state['exp_avg'])
136 exp_avg_sqs.append(state['exp_avg_sq'])
138 if amsgrad:
139 max_exp_avg_sqs.append(state['max_exp_avg_sq'])
141 if params_with_grad:
142 adamw(
143 params_with_grad,
144 grads,
145 exp_avgs,
146 exp_avg_sqs,
147 max_exp_avg_sqs,
148 group['step'],
149 amsgrad=amsgrad,
150 beta1=beta1,
151 beta2=beta2,
152 lr=group['lr'],
153 weight_decay=group['weight_decay'],
154 eps=group['eps'],
155 maximize=group['maximize']
156 )
158 return loss