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
« 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
18from __future__ import annotations
20import importlib
21from typing import Any, Dict, Iterable, List, Tuple
23import mindspore as ms
24from mindspore import nn
25from mindspore.common import dtype as mstype
26from mindspore.ops import functional as F
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)
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)
45class MindSporeAdamBaseAdapter(OptimizerSwapAdapter):
46 """Common MindSpore optimizer adapter logic."""
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] = {}
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.")
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"]
61 def all_slots(self) -> Iterable[SwapSlot]:
62 """Iterate known slots."""
63 return tuple(self._slots.values())
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()
69 def packed_layout_units(self) -> List[UpdateUnit]:
70 """Return stable optimizer units used to build the packed host layout."""
71 return []
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
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()
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()
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)
130 def _checkpoint_slots(self) -> Iterable[SwapSlot]:
131 """Return slots for the optimizer's current checkpoint-visible state."""
132 return tuple(self._slots.values())
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
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
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()
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
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
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 )
229 def _state_parameter(self, index: int, key: str) -> Any:
230 """Return the optimizer-owned Parameter for one logical state key."""
231 raise NotImplementedError
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
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
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)
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
281class MindSporeNativeAdamAdapter(MindSporeAdamBaseAdapter):
282 """Adapter for ``mindspore.nn.Adam``."""
284 @classmethod
285 def matches(cls, optimizer: Any) -> bool:
286 return isinstance(optimizer, nn.Adam)
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.")
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
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 }
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)
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
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}.")
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)
397 @staticmethod
398 def _index_lr(lr: Any, index: int) -> Any:
399 try:
400 return lr[index]
401 except (TypeError, IndexError):
402 return lr
405class MindSporeNativeAdamWAdapter(MindSporeAdamBaseAdapter):
406 """Adapter for ``mindspore.nn.AdamWeightDecay``."""
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 )
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.")
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}
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)
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
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}.")
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)
494 @staticmethod
495 def _indexed(value: Any, index: int, is_indexed: bool) -> Any:
496 return value[index] if is_indexed else value
499class MindFormersAdamWAdapter(MindSporeAdamBaseAdapter):
500 """Adapter for ``mindformers.pynative.optimizer.adamw.AdamW``."""
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 )
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.")
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 ]
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()
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}
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)
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)
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()
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
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}.")
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)
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 )