Coverage for  / home / jenkins / .local / lib / python3.10 / site-packages / hyper_parallel / platform / mindspore / swap_optimizer / adapters.py: 50%

421 statements  

« 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"""MindSpore Adam/AdamW swap optimizer adapters.""" 

16# pylint: disable=protected-access 

17 

18from __future__ import annotations 

19 

20import importlib 

21from typing import Any, Dict, Iterable, List, Tuple 

22 

23import mindspore as ms 

24from mindspore import nn 

25from mindspore.common import dtype as mstype 

26from mindspore.ops import functional as F 

27 

28from hyper_parallel.core.dtensor.dtensor import SkipDTensorDispatch 

29from hyper_parallel.core.optimizer.swap_optimizer_base import ( 

30 OptimizerSwapAdapter, 

31 SUPPORTED_STATE_KEYS, 

32 SwapSlot, 

33 UpdateUnit, 

34) 

35 

36 

37def _to_tuple(value: Any) -> Tuple[Any, ...]: 

38 if isinstance(value, tuple): 

39 return value 

40 if isinstance(value, list): 

41 return tuple(value) 

42 return tuple(value) 

43 

44 

45class MindSporeAdamBaseAdapter(OptimizerSwapAdapter): 

46 """Common MindSpore optimizer adapter logic.""" 

47 

48 def __init__(self, optimizer: Any, config: Any, runtime: Any) -> None: 

49 super().__init__(optimizer, config, runtime) 

50 self._slots: Dict[Tuple[int, str], SwapSlot] = {} 

51 

52 def validate(self) -> None: 

53 """Base validation.""" 

54 if getattr(self.optimizer, "use_parallel", False): 

55 raise ValueError("MindSpore swap optimizer does not support parallel optimizer yet.") 

56 

57 def iter_update_units(self, step_context: Dict[str, Any]) -> List[UpdateUnit]: 

58 """Return units collected in prepare_step.""" 

59 return step_context["units"] 

60 

61 def all_slots(self) -> Iterable[SwapSlot]: 

62 """Iterate known slots.""" 

63 return tuple(self._slots.values()) 

64 

65 def initial_slots(self) -> Iterable[SwapSlot]: 

66 """Build optimizer state slots that can be offloaded before the first update.""" 

67 return self._checkpoint_slots() 

68 

69 def packed_layout_units(self) -> List[UpdateUnit]: 

70 """Return stable optimizer units used to build the packed host layout.""" 

71 return [] 

72 

73 def checkpoint_state_dict(self, *args: Any, **kwargs: Any) -> Dict[str, Any]: 

74 """Return checkpoint-safe optimizer state dict.""" 

75 del args, kwargs 

76 state = self._state_dict() 

77 slot_by_name = self._checkpoint_slot_map() 

78 for name, slot in slot_by_name.items(): 

79 if name not in state or not slot.swappable: 

80 continue 

81 if slot.cpu_tensor is None: 

82 if slot.state == "host": 

83 raise RuntimeError(f"Swap slot {slot.name!r} is host-resident but has no CPU mirror.") 

84 slot.cpu_tensor = self.runtime.make_cpu_tensor(slot.tensor) 

85 state[name] = ms.Parameter(self.runtime.make_cpu_tensor(slot.cpu_tensor), name=name) 

86 return state 

87 

88 def load_checkpoint_state_dict( 

89 self, 

90 state_dict: Dict[str, Any], 

91 *args: Any, 

92 **kwargs: Any, 

93 ) -> None: 

94 """Load checkpoint-safe parameter dict.""" 

95 del args, kwargs 

96 slot_by_name = self._checkpoint_slot_map(promote_checkpoint_swappable=True) 

97 remaining = dict(state_dict) 

98 for name, slot in slot_by_name.items(): 

99 if name not in remaining or not slot.swappable: 

100 continue 

101 value = remaining.pop(name) 

102 tensor = getattr(value, "data", value) 

103 cpu_tensor = self.runtime.make_cpu_tensor(tensor) 

104 if slot.packed and slot.cpu_tensor is not None: 

105 self.runtime.copy_cpu_tensor(slot.cpu_tensor, cpu_tensor) 

106 else: 

107 slot.cpu_tensor = cpu_tensor 

108 self.runtime.release_device_storage(slot) 

109 if slot.packed: 

110 slot.tensor = slot.cpu_tensor 

111 slot.state = "host" 

112 if remaining: 

113 self._load_state_dict(remaining) 

114 self.publish_packed_state() 

115 

116 def _state_dict(self) -> Dict[str, Any]: 

117 if not hasattr(self.optimizer, "state_dict"): 

118 raise RuntimeError( 

119 "The installed MindSpore version does not support optimizer.state_dict()." 

120 ) 

121 return self.optimizer.state_dict() 

122 

123 def _load_state_dict(self, state_dict: Dict[str, Any]) -> None: 

124 if not hasattr(self.optimizer, "load_state_dict"): 

125 raise RuntimeError( 

126 "The installed MindSpore version does not support optimizer.load_state_dict()." 

127 ) 

128 self.optimizer.load_state_dict(state_dict, strict=False) 

129 

130 def _checkpoint_slots(self) -> Iterable[SwapSlot]: 

131 """Return slots for the optimizer's current checkpoint-visible state.""" 

132 return tuple(self._slots.values()) 

133 

134 def _checkpoint_slot_map(self, *, promote_checkpoint_swappable: bool = False) -> Dict[str, SwapSlot]: 

135 """Build a name-to-slot map from current optimizer state Parameters.""" 

136 checkpoint_slot_ids = {id(slot) for slot in self._checkpoint_slots()} 

137 slot_by_name: Dict[str, SwapSlot] = {} 

138 for (index, key), slot in self._slots.items(): 

139 if id(slot) not in checkpoint_slot_ids: 

140 continue 

141 if ( 

142 promote_checkpoint_swappable 

143 and not slot.swappable 

144 and self._is_checkpoint_swappable_slot(slot) 

145 ): 

146 slot.swappable = True 

147 slot.storage_nbytes = self.runtime.storage_nbytes(slot.tensor) 

148 name = getattr(self._state_parameter(index, key), "name", None) 

149 if not name: 

150 continue 

151 previous = slot_by_name.get(name) 

152 if previous is not None and previous is not slot: 

153 raise ValueError(f"Duplicate optimizer state parameter name in swap slots: {name!r}.") 

154 slot_by_name[name] = slot 

155 return slot_by_name 

156 

157 def _is_checkpoint_swappable_slot(self, slot: SwapSlot) -> bool: 

158 """Return whether a checkpoint slot meets the per-tensor swap requirements.""" 

159 if slot.name not in SUPPORTED_STATE_KEYS: 

160 return False 

161 

162 tensor = slot.tensor 

163 if isinstance(tensor, ms.Parameter): 

164 tensor = tensor.data 

165 if hasattr(tensor, "to_local"): 

166 tensor = tensor.to_local() 

167 

168 dtype_text = str(getattr(tensor, "dtype", "")).lower() 

169 if "float" not in dtype_text and "bfloat" not in dtype_text: 

170 return False 

171 if int(tensor.numel()) < int(self.config.min_numel): 

172 return False 

173 if not tensor.is_contiguous(): 

174 return False 

175 try: 

176 storage = tensor.untyped_storage() 

177 if storage.size() != int(tensor.numel()) * int(tensor.itemsize): 

178 return False 

179 except (AttributeError, RuntimeError): 

180 return False 

181 return True 

182 

183 def _make_slot(self, index: int, key: str, tensor: Any) -> SwapSlot: 

184 """Return the stable swap slot for an optimizer state tensor.""" 

185 slot = self._slots.get((index, key)) 

186 if slot is not None: 

187 return slot 

188 swappable = self.runtime.is_swappable_tensor(tensor, self.config.min_numel) 

189 is_packable = getattr(self.runtime, "is_packable_tensor", None) 

190 packed = bool(getattr(self.runtime, "packed_enabled", False) 

191 and is_packable is not None 

192 and is_packable(tensor, self.config.min_numel)) 

193 swappable = swappable or packed 

194 slot = SwapSlot( 

195 name=key, 

196 tensor=tensor, 

197 cpu_tensor=None, 

198 storage_nbytes=self.runtime.storage_nbytes(tensor), 

199 swappable=swappable, 

200 state="device", 

201 packed=packed, 

202 ) 

203 populate_metadata = getattr(self.runtime, "populate_slot_metadata", None) 

204 if populate_metadata is not None: 

205 populate_metadata(slot, tensor) 

206 slot.device = self._parameter_device(index) 

207 self._slots[(index, key)] = slot 

208 return slot 

209 

210 def publish_packed_state(self) -> None: 

211 """Publish persistent packed CPU mirrors to optimizer state Parameters.""" 

212 if not getattr(self.runtime, "packed_enabled", False): 

213 return 

214 for (index, key), slot in self._slots.items(): 

215 if not slot.packed or slot.cpu_tensor is None: 

216 continue 

217 parameter = self._state_parameter(index, key) 

218 set_data = getattr(parameter, "set_data", None) 

219 if callable(set_data): 

220 set_data(slot.cpu_tensor) 

221 continue 

222 if hasattr(parameter, "data"): 

223 parameter.data = slot.cpu_tensor 

224 continue 

225 raise RuntimeError( 

226 f"MindSpore optimizer state Parameter for slot {key!r} cannot publish a packed CPU mirror." 

227 ) 

228 

229 def _state_parameter(self, index: int, key: str) -> Any: 

230 """Return the optimizer-owned Parameter for one logical state key.""" 

231 raise NotImplementedError 

232 

233 def _parameter_device(self, index: int) -> Any: 

234 """Return the target update device for an optimizer state slot.""" 

235 params = getattr(self.optimizer, "_parameters", None) 

236 if params is None: 

237 params = getattr(self.optimizer, "fp32_params") 

238 param = _to_tuple(params)[index] 

239 if hasattr(param, "to_local"): 

240 param = param.to_local() 

241 return param.device 

242 

243 @staticmethod 

244 def _slot_tensor(unit: UpdateUnit, key: str, fallback: Any) -> Any: 

245 """Return the active staging view for a logical state key.""" 

246 for slot in unit.slots: 

247 if slot.name == key: 

248 return slot.tensor 

249 return fallback 

250 

251 def _selected_keys(self, available: Tuple[str, ...]) -> Tuple[str, ...]: 

252 """Return configured state keys that are available for this optimizer.""" 

253 keys = self.config.state_keys or available 

254 result = [] 

255 for key in keys: 

256 if key == "master_param": 

257 if self.config.state_keys is not None: 

258 raise ValueError(f"Requested state key '{key}' is not available for {type(self.optimizer)!r}.") 

259 continue 

260 if key in available: 

261 result.append(key) 

262 elif self.config.state_keys is not None: 

263 raise ValueError(f"Requested state key '{key}' is not available for {type(self.optimizer)!r}.") 

264 return tuple(result) 

265 

266 def _validate_gradient_count( 

267 self, 

268 gradients: Any, 

269 params: Any, 

270 ) -> Tuple[Tuple[Any, ...], Tuple[Any, ...]]: 

271 """Normalize parameters and gradients, and require one gradient per parameter.""" 

272 grad_tuple = _to_tuple(gradients) 

273 param_tuple = _to_tuple(params) 

274 if len(grad_tuple) != len(param_tuple): 

275 raise ValueError( 

276 f"MindSpore swap optimizer expected {len(param_tuple)} gradients, but got {len(grad_tuple)}." 

277 ) 

278 return grad_tuple, param_tuple 

279 

280 

281class MindSporeNativeAdamAdapter(MindSporeAdamBaseAdapter): 

282 """Adapter for ``mindspore.nn.Adam``.""" 

283 

284 @classmethod 

285 def matches(cls, optimizer: Any) -> bool: 

286 return isinstance(optimizer, nn.Adam) 

287 

288 def validate(self) -> None: 

289 super().validate() 

290 if getattr(self.config, "packed_swap", True): 

291 raise ValueError( 

292 "MindSpore nn.Adam does not support packed_swap=True. " 

293 "Set packed_swap=False to use per-tensor swap, or use " 

294 "mindformers AdamW for packed swap." 

295 ) 

296 if getattr(self.optimizer, "use_lazy", False): 

297 raise ValueError("MindSpore Adam swap optimizer does not support use_lazy=True.") 

298 if getattr(self.optimizer, "use_offload", False): 

299 raise ValueError("MindSpore Adam swap optimizer does not support use_offload=True.") 

300 

301 def prepare_step(self, *args: Any, **kwargs: Any) -> Dict[str, Any]: 

302 """Prepare native Adam step.""" 

303 if len(args) != 1 or kwargs: 

304 raise ValueError("MindSpore swap optimizer only accepts gradients.") 

305 gradients = args[0] 

306 opt = self.optimizer 

307 grad_tuple, params = self._validate_gradient_count(gradients, opt._parameters) 

308 gradients = opt.decay_weight(grad_tuple) 

309 gradients = opt.gradients_centralization(gradients) 

310 gradients = opt.scale_grad(gradients) 

311 gradients = opt._grad_sparse_indices_deduplicate(gradients) 

312 lr = opt.get_lr() 

313 opt.assignadd(opt.global_step, opt.global_step_increase_tensor) 

314 beta1_power = opt.beta1_power * opt.beta1 

315 opt.beta1_power = beta1_power 

316 beta2_power = opt.beta2_power * opt.beta2 

317 opt.beta2_power = beta2_power 

318 

319 grad_tuple = _to_tuple(gradients) 

320 units = [] 

321 for index, (param, grad) in enumerate(zip(params, grad_tuple)): 

322 if grad is None: 

323 continue 

324 slots = self._build_slots(index) 

325 units.append(UpdateUnit( 

326 adapter_index=index, 

327 param=param, 

328 grad=grad, 

329 slots=slots, 

330 )) 

331 return { 

332 "units": units, 

333 "gradients": grad_tuple, 

334 "lr": lr, 

335 "beta1_power": beta1_power, 

336 "beta2_power": beta2_power, 

337 } 

338 

339 def step_batch(self, batch: List[UpdateUnit], step_context: Dict[str, Any]) -> Tuple[Any, ...]: 

340 """Run native Adam for one batch.""" 

341 opt = self.optimizer 

342 results = [] 

343 for unit in batch: 

344 lr = self._index_lr(step_context["lr"], unit.adapter_index) 

345 if opt.use_amsgrad: 

346 result = opt.opt( 

347 unit.param, 

348 opt.moment1[unit.adapter_index], 

349 opt.moment2[unit.adapter_index], 

350 opt.vhat[unit.adapter_index], 

351 step_context["beta1_power"], 

352 step_context["beta2_power"], 

353 lr, 

354 opt.beta1, 

355 opt.beta2, 

356 opt.eps, 

357 unit.grad, 

358 ) 

359 else: 

360 result = opt._apply_adam( 

361 (unit.param,), 

362 step_context["beta1_power"], 

363 step_context["beta2_power"], 

364 (opt.moment1[unit.adapter_index],), 

365 (opt.moment2[unit.adapter_index],), 

366 (lr,) if opt.is_group_lr else lr, 

367 (unit.grad,), 

368 ) 

369 results.append(result) 

370 return tuple(results) 

371 

372 def _build_slots(self, index: int) -> List[SwapSlot]: 

373 available = ["exp_avg", "exp_avg_sq"] 

374 if getattr(self.optimizer, "use_amsgrad", False) and hasattr(self.optimizer, "vhat"): 

375 available.append("max_exp_avg_sq") 

376 slots = [] 

377 for key in self._selected_keys(tuple(available)): 

378 slots.append(self._make_slot(index, key, self._state_parameter(index, key))) 

379 return slots 

380 

381 def _state_parameter(self, index: int, key: str) -> Any: 

382 if key == "exp_avg": 

383 return self.optimizer.moment1[index] 

384 if key == "exp_avg_sq": 

385 return self.optimizer.moment2[index] 

386 if key == "max_exp_avg_sq": 

387 return self.optimizer.vhat[index] 

388 raise ValueError(f"Unknown native Adam state key: {key!r}.") 

389 

390 def _checkpoint_slots(self) -> Iterable[SwapSlot]: 

391 """Rebuild slots from native Adam state containers for checkpoint load.""" 

392 slots = [] 

393 for index in range(len(_to_tuple(self.optimizer._parameters))): 

394 slots.extend(self._build_slots(index)) 

395 return tuple(slots) 

396 

397 @staticmethod 

398 def _index_lr(lr: Any, index: int) -> Any: 

399 try: 

400 return lr[index] 

401 except (TypeError, IndexError): 

402 return lr 

403 

404 

405class MindSporeNativeAdamWAdapter(MindSporeAdamBaseAdapter): 

406 """Adapter for ``mindspore.nn.AdamWeightDecay``.""" 

407 

408 @classmethod 

409 def matches(cls, optimizer: Any) -> bool: 

410 adamw_cls = getattr(nn, "AdamW", None) 

411 return isinstance(optimizer, nn.AdamWeightDecay) or ( 

412 adamw_cls is not None and isinstance(optimizer, adamw_cls) 

413 ) 

414 

415 def validate(self) -> None: 

416 super().validate() 

417 if getattr(self.config, "packed_swap", True): 

418 raise ValueError( 

419 "MindSpore nn.AdamWeightDecay does not support packed_swap=True. " 

420 "Set packed_swap=False to use per-tensor swap, or use " 

421 "mindformers AdamW for packed swap." 

422 ) 

423 if not getattr(self.optimizer, "use_fused_opt", False): 

424 raise ValueError("MindSpore AdamWeightDecay swap optimizer only supports use_fused_opt=True.") 

425 

426 def prepare_step(self, *args: Any, **kwargs: Any) -> Dict[str, Any]: 

427 """Prepare native AdamWeightDecay step.""" 

428 if len(args) != 1 or kwargs: 

429 raise ValueError("MindSpore swap optimizer only accepts gradients.") 

430 gradients = args[0] 

431 opt = self.optimizer 

432 grad_tuple, params = self._validate_gradient_count(gradients, opt._parameters) 

433 weight_decay = opt.get_weight_decay() 

434 lr = opt.get_lr() 

435 opt.assignadd(opt.global_step, opt.global_step_increase_tensor) 

436 units = [] 

437 for index, (param, grad) in enumerate(zip(params, grad_tuple)): 

438 if grad is None: 

439 continue 

440 slots = self._build_slots(index) 

441 units.append(UpdateUnit( 

442 adapter_index=index, 

443 param=param, 

444 grad=grad, 

445 slots=slots, 

446 )) 

447 return {"units": units, "gradients": grad_tuple, "lr": lr, "weight_decay": weight_decay} 

448 

449 def step_batch(self, batch: List[UpdateUnit], step_context: Dict[str, Any]) -> Tuple[Any, ...]: 

450 """Run AdamWeightDecay fused primitive for one batch.""" 

451 opt = self.optimizer 

452 results = [] 

453 for unit in batch: 

454 if not opt.optim_filter[unit.adapter_index]: 

455 results.append(True) 

456 continue 

457 lr = self._indexed(step_context["lr"], unit.adapter_index, opt.is_group_lr) 

458 weight_decay = self._indexed(step_context["weight_decay"], unit.adapter_index, opt.is_group) 

459 decay = weight_decay if opt.decay_flags[unit.adapter_index] else 0.0 

460 grad = F.cast(unit.grad, F.dtype(unit.param)) 

461 results.append(opt.fused_opt( 

462 unit.param, 

463 opt.moments1[unit.adapter_index], 

464 opt.moments2[unit.adapter_index], 

465 lr, 

466 opt.beta1, 

467 opt.beta2, 

468 opt.eps, 

469 decay, 

470 grad, 

471 )) 

472 return tuple(results) 

473 

474 def _build_slots(self, index: int) -> List[SwapSlot]: 

475 slots = [] 

476 for key in self._selected_keys(("exp_avg", "exp_avg_sq")): 

477 slots.append(self._make_slot(index, key, self._state_parameter(index, key))) 

478 return slots 

479 

480 def _state_parameter(self, index: int, key: str) -> Any: 

481 if key == "exp_avg": 

482 return self.optimizer.moments1[index] 

483 if key == "exp_avg_sq": 

484 return self.optimizer.moments2[index] 

485 raise ValueError(f"Unknown native AdamWeightDecay state key: {key!r}.") 

486 

487 def _checkpoint_slots(self) -> Iterable[SwapSlot]: 

488 """Rebuild slots from native AdamWeightDecay state containers for checkpoint load.""" 

489 slots = [] 

490 for index in range(len(_to_tuple(self.optimizer._parameters))): 

491 slots.extend(self._build_slots(index)) 

492 return tuple(slots) 

493 

494 @staticmethod 

495 def _indexed(value: Any, index: int, is_indexed: bool) -> Any: 

496 return value[index] if is_indexed else value 

497 

498 

499class MindFormersAdamWAdapter(MindSporeAdamBaseAdapter): 

500 """Adapter for ``mindformers.pynative.optimizer.adamw.AdamW``.""" 

501 

502 @classmethod 

503 def matches(cls, optimizer: Any) -> bool: 

504 optimizer_type = type(optimizer) 

505 return ( 

506 optimizer_type.__name__ == "AdamW" 

507 and optimizer_type.__module__ == "mindformers.pynative.optimizer.adamw" 

508 ) 

509 

510 def validate(self) -> None: 

511 super().validate() 

512 if getattr(self.optimizer, "enable_cpu_offload", False): 

513 raise ValueError("mindformers AdamW enable_cpu_offload is not supported with swap optimizer.") 

514 

515 def packed_layout_units(self) -> List[UpdateUnit]: 

516 """Return all MindFormers AdamW units in stable optimizer order.""" 

517 return [ 

518 UpdateUnit( 

519 adapter_index=index, 

520 param=param, 

521 grad=None, 

522 slots=self._build_slots(index), 

523 ) 

524 for index, param in enumerate(_to_tuple(self.optimizer.fp32_params)) 

525 ] 

526 

527 def prepare_step(self, *args: Any, **kwargs: Any) -> Dict[str, Any]: 

528 """Prepare mindformers PyNative AdamW step.""" 

529 if len(args) != 1 or kwargs: 

530 raise ValueError("MindSpore swap optimizer only accepts gradients.") 

531 gradients = args[0] 

532 opt = self.optimizer 

533 grad_tuple, params = self._validate_gradient_count(gradients, opt.fp32_params) 

534 weight_decay = opt.get_weight_decay() 

535 lr = opt.get_lr() 

536 opt._increase_global_step() 

537 

538 lr = [float(x) for x in lr] if (opt.is_group and opt.is_group_lr) else float(lr) 

539 weight_decay = [float(x) for x in weight_decay] if opt.is_group else float(weight_decay) 

540 units = [] 

541 for index, (param, grad) in enumerate(zip(params, grad_tuple)): 

542 if grad is None and not self.runtime.packed_enabled: 

543 continue 

544 slots = self._build_slots(index) 

545 units.append(UpdateUnit( 

546 adapter_index=index, 

547 param=param, 

548 grad=grad, 

549 slots=slots, 

550 )) 

551 return {"units": units, "gradients": grad_tuple, "lr": lr, "weight_decay": weight_decay} 

552 

553 def step_batch(self, batch: List[UpdateUnit], step_context: Dict[str, Any]) -> Tuple[Any, ...]: 

554 """Run mindformers AdamW helpers for one batch.""" 

555 with SkipDTensorDispatch(): 

556 opt = self.optimizer 

557 module = importlib.import_module(type(opt).__module__) 

558 results = [] 

559 is_lr_list = isinstance(step_context["lr"], list) 

560 is_wd_list = isinstance(step_context["weight_decay"], list) 

561 if getattr(opt, "enable_fused_opt", False): 

562 step = module.op_cast(opt.global_step, mstype.int64) 

563 for unit in batch: 

564 if unit.grad is None: 

565 continue 

566 if not opt.optim_filter[unit.adapter_index]: 

567 results.append(True) 

568 continue 

569 update_param = self._slot_tensor(unit, "master_param", unit.param) 

570 learning_rate = ( 

571 step_context["lr"][unit.adapter_index] 

572 if is_lr_list else step_context["lr"] 

573 ) 

574 weight_decay = ( 

575 step_context["weight_decay"][unit.adapter_index] 

576 if is_wd_list else step_context["weight_decay"] 

577 ) 

578 results.append(module._run_fused_adamw_opt( 

579 opt.fused_adamw_opt, 

580 opt.amsgrad, 

581 opt.maximize, 

582 opt.beta1_value, 

583 opt.beta2_value, 

584 opt.eps_value, 

585 step, 

586 learning_rate, 

587 weight_decay, 

588 update_param, 

589 unit.grad, 

590 self._slot_tensor(unit, "exp_avg", opt.exp_avg[unit.adapter_index]), 

591 self._slot_tensor(unit, "exp_avg_sq", opt.exp_avg_sq[unit.adapter_index]), 

592 self._slot_tensor(unit, "max_exp_avg_sq", opt.max_exp_avg_sq[unit.adapter_index]), 

593 )) 

594 self._sync_batch_master_params(batch) 

595 return tuple(results) 

596 

597 bias_correction1 = 1.0 - opt.beta1 ** opt.global_step 

598 bias_correction2 = 1.0 - opt.beta2 ** opt.global_step 

599 for unit in batch: 

600 if unit.grad is None: 

601 continue 

602 update_param = self._slot_tensor(unit, "master_param", unit.param) 

603 results.append(module._run_adamw_opt( 

604 opt.beta1, 

605 opt.beta2, 

606 opt.eps, 

607 step_context["lr"][unit.adapter_index] if is_lr_list else step_context["lr"], 

608 step_context["weight_decay"][unit.adapter_index] if is_wd_list else step_context["weight_decay"], 

609 update_param, 

610 unit.grad, 

611 self._slot_tensor(unit, "exp_avg", opt.exp_avg[unit.adapter_index]), 

612 self._slot_tensor(unit, "exp_avg_sq", opt.exp_avg_sq[unit.adapter_index]), 

613 opt.optim_filter[unit.adapter_index], 

614 bias_correction1, 

615 bias_correction2, 

616 opt.one_minus_beta2, 

617 )) 

618 self._sync_batch_master_params(batch) 

619 return tuple(results) 

620 

621 def finish_step(self, step_context: Dict[str, Any]) -> None: 

622 del step_context 

623 if not self.config.include_master_params: 

624 with SkipDTensorDispatch(): 

625 self.optimizer._copy_main_params_to_model_params() 

626 

627 def _build_slots(self, index: int) -> List[SwapSlot]: 

628 """Build swap slots for one MindFormers optimizer parameter.""" 

629 opt = self.optimizer 

630 available = ["exp_avg", "exp_avg_sq"] 

631 max_slot = getattr(opt, "max_exp_avg_sq", None) 

632 if max_slot is not None and max_slot is not opt.exp_avg_sq: 

633 available.append("max_exp_avg_sq") 

634 slots = [] 

635 selected_keys = [] 

636 for key in (self.config.state_keys or tuple(available)): 

637 if key == "master_param": 

638 continue 

639 if key not in available: 

640 raise ValueError(f"Requested state key '{key}' is not available for {type(self.optimizer)!r}.") 

641 selected_keys.append(key) 

642 for key in tuple(selected_keys): 

643 slots.append(self._make_slot(index, key, self._state_parameter(index, key))) 

644 if self.config.include_master_params and hasattr(opt, "fp32_params"): 

645 fp32_param = opt.fp32_params[index] 

646 model_param = opt._parameters[index] 

647 if fp32_param is not model_param: 

648 slots.append(self._make_slot(index, "master_param", fp32_param)) 

649 return slots 

650 

651 def _state_parameter(self, index: int, key: str) -> Any: 

652 opt = self.optimizer 

653 if key == "exp_avg": 

654 return opt.exp_avg[index] 

655 if key == "exp_avg_sq": 

656 return opt.exp_avg_sq[index] 

657 if key == "max_exp_avg_sq": 

658 return opt.max_exp_avg_sq[index] 

659 if key == "master_param": 

660 return opt.fp32_params[index] 

661 raise ValueError(f"Unknown MindFormers AdamW state key: {key!r}.") 

662 

663 def _checkpoint_slots(self) -> Iterable[SwapSlot]: 

664 """Rebuild slots from MindFormers AdamW state containers for checkpoint load.""" 

665 slots = [] 

666 for index in range(len(_to_tuple(self.optimizer.fp32_params))): 

667 slots.extend(self._build_slots(index)) 

668 return tuple(slots) 

669 

670 def _sync_batch_master_params(self, batch: List[UpdateUnit]) -> None: 

671 opt = self.optimizer 

672 if not self.config.include_master_params: 

673 return 

674 module = importlib.import_module(type(opt).__module__) 

675 for unit in batch: 

676 if unit.grad is None: 

677 continue 

678 if opt._is_low_precision_param[unit.adapter_index]: 

679 module.inplace_copy( 

680 opt._parameters[unit.adapter_index], 

681 module.op_cast( 

682 self._slot_tensor(unit, "master_param", opt.fp32_params[unit.adapter_index]), 

683 opt._parameters[unit.adapter_index].dtype, 

684 ), 

685 )