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

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.""" 

16 

17from typing import List 

18 

19import torch 

20 

21from hyper_parallel.core.optimizer.utils import get_current_device 

22 

23 

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) 

46 

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 ) 

62 

63 

64class AdamW(torch.optim.Optimizer): 

65 """AdamW optimizer implementation.""" 

66 

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) 

86 

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) 

93 

94 def __str__(self): 

95 return super().__repr__() 

96 

97 __repr__ = __str__ 

98 

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() 

105 

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 = [] 

112 

113 amsgrad = group['amsgrad'] 

114 beta1, beta2 = group['betas'] 

115 group['step'] = (group.get('step') or 0) + 1 

116 

117 current_rank_params = group['params'] 

118 for p in current_rank_params: 

119 if p.grad is None: 

120 continue 

121 

122 if p.grad.data.is_sparse: 

123 raise RuntimeError('AdamW does not support sparse gradients') 

124 

125 state = self.state[p] 

126 

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) 

132 

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']) 

137 

138 if amsgrad: 

139 max_exp_avg_sqs.append(state['max_exp_avg_sq']) 

140 

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 ) 

157 

158 return loss