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 optimizer state swap wrapper."""
16# pylint: disable=protected-access
17
18from __future__ import annotations
19
20import contextlib
21import ctypes
22from dataclasses import dataclass, field
23from typing import Any, Dict, List, Optional, Sequence
24
25import mindspore as ms
26from mindspore.graph.api import _no_grad
27
28from hyper_parallel.core.optimizer.swap_optimizer_base import (
29 STATE_KEYS,
30 PipelineSwapRuntime,
31 SwapSlot,
32 _iter_unique_slots,
33)
34from hyper_parallel.platform import get_platform
35from hyper_parallel.platform.mindspore.swap_optimizer.adapters import (
36 MindFormersAdamWAdapter,
37 MindSporeNativeAdamAdapter,
38 MindSporeNativeAdamWAdapter,
39)
40
41platform = get_platform()
42_PACKED_ALIGNMENT_BYTES = 512
43
44
45@dataclass
46class _PackedBatchRegion:
47 """One persistent dtype buffer transferred for a pipeline batch."""
48
49 dtype: Any
50 numel: int
51 slots: List[SwapSlot]
52
53
54@dataclass
55class _PackedBatchPlan:
56 """Persistent packed transfer layout for one optimizer pipeline batch."""
57
58 regions: Dict[Any, _PackedBatchRegion] = field(default_factory=dict)
59
60
61@dataclass
62class _StagingArena:
63 """One step-local raw NPU allocation and its dtype-specific views."""
64
65 raw_buffer: Any
66 dtype_views: Dict[Any, Any] = field(default_factory=dict)
67
68
69class MindSporeSwapRuntime(PipelineSwapRuntime):
70 """MindSpore tensor storage/copy runtime."""
71
72 def __init__(self, config: Any) -> None:
73 super().__init__(config)
74 self._packed_enabled = bool(getattr(config, "packed_swap", True))
75 self._host_buffers: List[Dict[Any, Any]] = []
76 self._host_layout_signature = ()
77 self._packed_batch_plans: List[_PackedBatchPlan] = []
78 self._staging_arenas: List[Optional[_StagingArena]] = [None, None]
79 self._packed_ready_events: Dict[int, Any] = {}
80 self._packed_offload_events: Dict[int, Any] = {}
81 self._packed_tail_event: Optional[Any] = None
82
83 @property
84 def packed_enabled(self) -> bool:
85 """Return whether packed MindSpore optimizer swap is enabled."""
86 return self._packed_enabled
87
88 def populate_slot_metadata(self, slot: SwapSlot, template: Any) -> None:
89 """Populate stable logical metadata before the source storage is released."""
90 storage_tensor = self._storage_tensor(template)
91 slot.shape = tuple(storage_tensor.shape)
92 slot.dtype = storage_tensor.dtype
93 slot.device = storage_tensor.device
94 slot.numel = int(storage_tensor.numel())
95 slot.storage_nbytes = slot.numel * int(storage_tensor.itemsize)
96
97 def prepare_packed_host(self, batches: Sequence[Sequence[Any]]) -> bool:
98 """Pack optimizer states into persistent pinned CPU buffers by batch and dtype.
99
100 Returns:
101 Whether the persistent host layout was rebuilt.
102 """
103 if not self._packed_enabled:
104 return False
105 batch_lists = [list(batch) for batch in batches]
106 signature = self._packed_layout_signature(batch_lists)
107 if signature == self._host_layout_signature:
108 return False
109
110 new_buffers: List[Dict[Any, Any]] = []
111 new_plans: List[_PackedBatchPlan] = []
112 slot_bindings: List[tuple[SwapSlot, int, Any]] = []
113 layout_slot_ids = set()
114 for batch in batch_lists:
115 slots_by_dtype: Dict[Any, List[SwapSlot]] = {}
116 for unit in batch:
117 for slot in unit.slots:
118 if not slot.swappable or not slot.packed:
119 continue
120 slot_id = id(slot)
121 if slot_id in layout_slot_ids:
122 raise ValueError(
123 f"Packed swap slot {slot.name!r} appears in more than one host batch."
124 )
125 if slot.state != "host":
126 raise RuntimeError(
127 f"Packed swap slot {slot.name!r} must be idle on host before layout preparation."
128 )
129 # The D2H copy may still be pending after the compute stream
130 # ordered storage release. This is the host-consumption
131 # boundary for the mirror, so wait here before reading it.
132 if slot.event is not None:
133 self.wait_event(slot.event, None)
134 slot.event = None
135 if slot.cpu_tensor is None:
136 raise RuntimeError(
137 f"Packed swap slot {slot.name!r} has no CPU mirror before host layout preparation."
138 )
139 layout_slot_ids.add(slot_id)
140 slots_by_dtype.setdefault(slot.dtype, []).append(slot)
141
142 batch_buffers: Dict[Any, Any] = {}
143 batch_regions: Dict[Any, _PackedBatchRegion] = {}
144 for dtype, dtype_slots in slots_by_dtype.items():
145 total_numel = sum(slot.numel for slot in dtype_slots)
146 host_buffer = ms.mint.empty(
147 (total_numel,), dtype=dtype, device="cpu", pin_memory=True
148 )
149 batch_buffers[dtype] = host_buffer
150 host_offset = 0
151 for slot in dtype_slots:
152 host_view = self._make_cpu_storage_view(host_buffer, host_offset, slot.shape)
153 self._copy_cpu_tensor(host_view, slot.cpu_tensor)
154 slot_bindings.append((slot, host_offset, host_view))
155 host_offset += slot.numel
156 batch_regions[dtype] = _PackedBatchRegion(dtype, total_numel, dtype_slots)
157 new_buffers.append(batch_buffers)
158 new_plans.append(_PackedBatchPlan(batch_regions))
159
160 for slot, host_offset, host_view in slot_bindings:
161 slot.host_offset = host_offset
162 slot.cpu_tensor = host_view
163 slot.tensor = host_view
164 slot.state = "host"
165 slot.event = None
166 self._host_buffers = new_buffers
167 self._packed_batch_plans = new_plans
168 self._host_layout_signature = signature
169 return True
170
171 def is_swappable_tensor(self, tensor: Any, min_numel: int) -> bool:
172 """Return whether ``tensor`` can participate in swap."""
173 storage_tensor = self._storage_tensor(tensor)
174 if not isinstance(storage_tensor, ms.Tensor):
175 return False
176 if not _is_float_dtype(storage_tensor.dtype):
177 return False
178 if int(storage_tensor.numel()) < int(min_numel):
179 return False
180 if not storage_tensor.is_contiguous():
181 return False
182 if _is_cpu_tensor(storage_tensor):
183 return False
184 try:
185 storage = storage_tensor.untyped_storage()
186 expected = int(storage_tensor.numel()) * int(storage_tensor.itemsize)
187 if storage.size() != expected:
188 return False
189 except (AttributeError, RuntimeError):
190 return False
191 return True
192
193 def is_packable_tensor(self, tensor: Any, min_numel: int) -> bool:
194 """Return whether a CPU- or NPU-resident state can enter packed storage."""
195 storage_tensor = self._storage_tensor(tensor)
196 if not isinstance(storage_tensor, ms.Tensor):
197 return False
198 if not _is_float_dtype(storage_tensor.dtype):
199 return False
200 if int(storage_tensor.numel()) < int(min_numel) or not storage_tensor.is_contiguous():
201 return False
202 try:
203 expected = int(storage_tensor.numel()) * int(storage_tensor.itemsize)
204 return storage_tensor.untyped_storage().size() == expected
205 except (AttributeError, RuntimeError):
206 return False
207
208 def storage_nbytes(self, tensor: Any) -> int:
209 """Return storage bytes for a MindSpore tensor."""
210 if not isinstance(tensor, ms.Tensor):
211 return 0
212 try:
213 return int(tensor.untyped_storage().size())
214 except (AttributeError, RuntimeError):
215 return int(tensor.numel()) * int(tensor.itemsize)
216
217 def refresh_swappable_slots(self, batch: Any) -> None:
218 """Mark state slots that become device-resident after an update."""
219 for slot in _iter_unique_slots(batch):
220 if slot.swappable or slot.name not in STATE_KEYS:
221 continue
222 if not self.is_swappable_tensor(slot.tensor, self.config.min_numel):
223 continue
224 slot.swappable = True
225 slot.storage_nbytes = self.storage_nbytes(slot.tensor)
226 slot.state = "device"
227
228 def make_cpu_tensor(self, tensor: Any) -> Any:
229 """Create a CPU mirror tensor and copy live storage into it."""
230 source = self._storage_tensor(tensor)
231 cpu_tensor = ms.mint.empty(
232 tuple(source.shape), dtype=source.dtype, device="cpu", pin_memory=True
233 )
234 if _is_cpu_tensor(source):
235 # In an Ascend context MindSpore's Tensor.copy()/clone() may place
236 # the result on Ascend even when the source tensor is on CPU.
237 self._copy_cpu_tensor(cpu_tensor, source)
238 else:
239 self._copy_storage(cpu_tensor, source)
240 return cpu_tensor
241
242 def copy_cpu_tensor(self, target: Any, source: Any) -> None:
243 """Copy matching CPU tensors without dispatching a device operator."""
244 self._copy_cpu_tensor(target, source)
245
246 def copy_to_device(self, slot: SwapSlot) -> None:
247 """Copy CPU mirror into live tensor."""
248 if slot.state != "host":
249 return
250 if slot.cpu_tensor is None:
251 return
252 self.load_into_tensor(slot.tensor, slot.cpu_tensor)
253 slot.state = "device"
254
255 def wait_prefetch_slot(self, slot: SwapSlot) -> None:
256 """MindSpore fallback copies are stream ordered."""
257 slot.state = "device"
258
259 def copy_to_cpu(self, slot: SwapSlot) -> None:
260 """Copy live tensor into CPU mirror."""
261 if slot.cpu_tensor is None:
262 slot.cpu_tensor = self.make_cpu_tensor(slot.tensor)
263 else:
264 self._copy_storage(slot.cpu_tensor, slot.tensor)
265 slot.state = "d2h"
266
267 def wait_offload_slot(self, slot: SwapSlot) -> None:
268 """Release storage after offload."""
269 self.release_device_storage(slot)
270 slot.state = "host"
271
272 def load_into_tensor(self, tensor: Any, value: Any) -> None:
273 """Copy CPU mirror storage into the live tensor storage."""
274 self._copy_storage(tensor, value)
275
276 def restore_device_storage(self, slot: SwapSlot) -> None:
277 """Restore live tensor storage."""
278 storage_tensor = self._storage_tensor(slot.tensor)
279 if _is_cpu_tensor(storage_tensor):
280 return
281 try:
282 storage = storage_tensor.untyped_storage()
283 if storage.size() != slot.storage_nbytes:
284 storage.resize_(slot.storage_nbytes)
285 except (AttributeError, RuntimeError) as exc:
286 raise RuntimeError(
287 f"Failed to restore device storage for swap slot {slot.name!r}: {exc}"
288 ) from exc
289
290 def release_device_storage(self, slot: SwapSlot) -> None:
291 """Release live tensor storage."""
292 storage_tensor = self._storage_tensor(slot.tensor)
293 if _is_cpu_tensor(storage_tensor):
294 return
295 try:
296 storage = storage_tensor.untyped_storage()
297 if storage.size() != 0:
298 storage.resize_(0)
299 except (AttributeError, RuntimeError) as exc:
300 raise RuntimeError(
301 f"Failed to release device storage for swap slot {slot.name!r}: {exc}"
302 ) from exc
303
304 def supports_packed_pipeline(self, batches: Sequence[Sequence[Any]]) -> bool:
305 """Return whether every swappable batch slot uses packed host storage."""
306 if not self._packed_enabled:
307 return False
308 if not batches:
309 return False
310
311 slots = [slot for batch in batches for unit in batch for slot in unit.slots if slot.swappable]
312 if not slots:
313 return False
314
315 unpacked_slots = [slot for slot in slots if not slot.packed]
316 missing_cpu_slots = [slot for slot in slots if slot.cpu_tensor is None]
317 if unpacked_slots or missing_cpu_slots:
318 return False
319
320 if self._packed_layout_signature(batches) != self._host_layout_signature:
321 raise RuntimeError(
322 "Packed MindSpore optimizer batches do not match the persistent host layout."
323 )
324
325 return True
326
327 def begin_packed_step(self, batches: Sequence[Sequence[Any]]) -> None:
328 """Validate the persistent layout and materialize two step-local staging buffers."""
329 ms.runtime.synchronize()
330 self._packed_tail_event = None
331 if self._packed_layout_signature(batches) != self._host_layout_signature:
332 raise RuntimeError(
333 "Packed MindSpore optimizer batches changed after host layout preparation."
334 )
335 self._staging_arenas = [None, None]
336 max_numel_by_dtype: Dict[Any, int] = {}
337 element_size_by_dtype: Dict[Any, int] = {}
338 device = None
339 for batch_index, batch_plan in enumerate(self._packed_batch_plans):
340 for dtype, region in batch_plan.regions.items():
341 max_numel_by_dtype[dtype] = max(max_numel_by_dtype.get(dtype, 0), region.numel)
342 element_size_by_dtype[dtype] = int(self._host_buffers[batch_index][dtype].itemsize)
343 if device is None and region.slots:
344 device = region.slots[0].device
345 if device is None:
346 raise RuntimeError("Packed MindSpore optimizer pipeline has no state metadata.")
347
348 dtype_layouts = {}
349 byte_offset = 0
350 for dtype in sorted(max_numel_by_dtype, key=str):
351 byte_offset = self._align_bytes(byte_offset)
352 element_size = element_size_by_dtype[dtype]
353 num_bytes = max_numel_by_dtype[dtype] * element_size
354 dtype_layouts[dtype] = (byte_offset, num_bytes)
355 byte_offset += num_bytes
356 total_bytes = self._align_bytes(byte_offset)
357
358 for staging_index in range(2):
359 arena = self._materialize_staging_arena(staging_index, total_bytes, device)
360 arena.dtype_views = {
361 dtype: arena.raw_buffer.narrow(0, offset, num_bytes).view(dtype)
362 for dtype, (offset, num_bytes) in dtype_layouts.items()
363 }
364 self._packed_ready_events = {}
365 self._packed_offload_events = {}
366
367 def enqueue_packed_prefetch(self, batch_index: int, staging_index: int) -> None:
368 """Enqueue one dtype-packed H2D chain."""
369 copy_stream = self._get_copy_stream()
370 with self.stream_context(copy_stream):
371 self._copy_packed_to_device(batch_index, staging_index)
372 ready_event = self.record_event(copy_stream)
373 self._packed_ready_events[batch_index] = ready_event
374 self._packed_tail_event = ready_event
375
376 def wait_packed_prefetch(self, batch_index: int, staging_index: int) -> None:
377 """Order the compute stream after a packed prefetch."""
378 del staging_index
379 event = self._packed_ready_events.get(batch_index)
380 if event is None:
381 raise RuntimeError(f"Packed MindSpore batch {batch_index} has no ready event.")
382 self.wait_event(event, self.current_stream())
383
384 def activate_packed_batch(self, batch_index: int, staging_index: int) -> None:
385 """Bind logical optimizer slots to views of one staging arena."""
386 batch_plan = self._packed_batch_plans[batch_index]
387 arena = self._require_staging_arena(staging_index)
388 for region in batch_plan.regions.values():
389 dtype_view = arena.dtype_views[region.dtype]
390 for slot in region.slots:
391 device_view = dtype_view.narrow(0, slot.host_offset, slot.numel).view(slot.shape)
392 slot.tensor = device_view
393 slot.state = "device"
394 slot.event = None
395
396 def enqueue_packed_offload_prefetch(
397 self,
398 batch_index: int,
399 next_index: Optional[int],
400 staging_index: int,
401 ) -> None:
402 """Serialize D2H and the next same-parity H2D on the copy stream."""
403 copy_stream = self._get_copy_stream()
404 compute_event = self._record_current_stream_event()
405 with self.stream_context(copy_stream):
406 self.wait_event(compute_event, copy_stream)
407 self._copy_packed_to_host(batch_index, staging_index)
408 if next_index is not None:
409 self._copy_packed_to_device(next_index, staging_index)
410 chain_event = self.record_event(copy_stream)
411 self._packed_offload_events[batch_index] = chain_event
412 self._packed_tail_event = chain_event
413 if next_index is not None:
414 self._packed_ready_events[next_index] = chain_event
415
416 def wait_packed_offload(self, batch_index: int) -> None:
417 """Order the compute stream after a trailing packed transfer chain."""
418 event = self._packed_offload_events.get(batch_index)
419 if event is None:
420 raise RuntimeError(f"Packed MindSpore batch {batch_index} has no offload event.")
421 self.wait_event(event, self.current_stream())
422
423 def finish_packed_offload(self, batch_index: int) -> None:
424 """Make persistent pinned views authoritative after D2H."""
425 for region in self._packed_batch_plans[batch_index].regions.values():
426 for slot in region.slots:
427 slot.tensor = slot.cpu_tensor
428 slot.state = "host"
429 slot.event = None
430
431 def release_packed_step_results(self, results: List[Any]) -> None:
432 """Drop PyNative update stubs before releasing their staging inputs."""
433 results.clear()
434
435 def end_packed_step(self) -> None:
436 """Drain transfers, detach all slot views, and destroy both staging arenas."""
437 tail_event = getattr(self, "_packed_tail_event", None)
438 if tail_event is not None:
439 self.wait_event(tail_event, None)
440 active_slots = {
441 id(slot): slot
442 for plan in self._packed_batch_plans
443 for region in plan.regions.values()
444 for slot in region.slots
445 if slot.state == "device"
446 }
447 # Normal batches offload through _copy_packed_to_host(); recover only interrupted updates here.
448 if active_slots:
449 compute_stream = self.current_stream()
450 synchronize = getattr(compute_stream, "synchronize", None)
451 if synchronize is not None:
452 synchronize()
453 for slot in active_slots.values():
454 bounce_buffer = ms.mint.empty(
455 slot.shape, dtype=slot.dtype, device="cpu", pin_memory=True
456 )
457 self._copy_tensor(bounce_buffer, slot.tensor, non_blocking=False)
458 self._copy_cpu_tensor(slot.cpu_tensor, bounce_buffer)
459 slot.tensor = slot.cpu_tensor
460 slot.state = "host"
461 slot.event = None
462 all_slots = {
463 id(slot): slot
464 for plan in self._packed_batch_plans
465 for region in plan.regions.values()
466 for slot in region.slots
467 }
468 # A slot view keeps the arena storage alive even after raw_buffer.resize_(0).
469 # Detach every slot first, including batches already marked as host-resident.
470 for slot in all_slots.values():
471 if slot.cpu_tensor is not None:
472 slot.tensor = slot.cpu_tensor
473 slot.state = "host"
474 slot.event = None
475 for arena in self._staging_arenas:
476 if arena is None:
477 continue
478 arena.dtype_views = {}
479 storage = arena.raw_buffer.untyped_storage()
480 if storage.size() != 0:
481 storage.resize_(0)
482 arena.raw_buffer = None
483 # Do not retain zero-sized raw tensors between optimizer steps. MindSpore
484 # views may otherwise keep their previous device allocation in the pool.
485 self._staging_arenas = [None, None]
486 self._packed_ready_events = {}
487 self._packed_offload_events = {}
488 self._packed_tail_event = None
489
490 def current_stream(self) -> Any:
491 """Return the current compute stream."""
492 return platform.get_current_stream()
493
494 def new_stream(self) -> Any:
495 """Create the copy stream."""
496 return platform.new_stream()
497
498 def stream_context(self, stream: Any) -> Any:
499 """Return a stream context for ``stream``."""
500 if stream is None:
501 return contextlib.nullcontext()
502 return platform.get_stream_context()(stream)
503
504 def record_event(self, stream: Any = None) -> Any:
505 """Record an event on ``stream``."""
506 event = platform.new_event()
507 if stream is None:
508 event.record()
509 else:
510 event.record(stream)
511 return event
512
513 def wait_event(self, event: Any, stream: Any = None) -> None:
514 """Make ``stream`` wait for ``event``."""
515 if event is None:
516 return
517 if stream is None:
518 event.synchronize()
519 return
520 event.wait(stream)
521
522 def _copy_storage(self, target: Any, source: Any) -> None:
523 """Copy tensor storage without allocating a new tensor object."""
524 if isinstance(target, ms.Parameter):
525 target = target.data
526 if isinstance(source, ms.Parameter):
527 source = source.data
528 target.untyped_storage().copy_(source.untyped_storage(), non_blocking=True)
529
530 def _copy_tensor(self, target: Any, source: Any, *, non_blocking: bool = True) -> None:
531 """Copy logical tensor views while respecting their offsets and lengths."""
532 if isinstance(target, ms.Parameter):
533 target = target.data
534 if isinstance(source, ms.Parameter):
535 source = source.data
536 target.copy_(source, non_blocking=non_blocking)
537
538 @staticmethod
539 def _packed_layout_signature(batches: Sequence[Sequence[Any]]) -> tuple:
540 """Return the batch-sensitive identity of a packed host layout."""
541 return tuple(
542 tuple(
543 (
544 getattr(unit, "adapter_index", None),
545 tuple(
546 (id(slot), slot.name, slot.dtype, slot.shape, slot.numel)
547 for slot in unit.slots
548 if slot.swappable and slot.packed
549 ),
550 )
551 for unit in batch
552 )
553 for batch in batches
554 )
555
556 def _materialize_staging_arena(self, staging_index: int, total_bytes: int, device: Any) -> _StagingArena:
557 raw_buffer = platform.empty((total_bytes,), dtype=ms.uint8, device=device)
558 arena = _StagingArena(raw_buffer)
559 self._staging_arenas[staging_index] = arena
560 return arena
561
562 def _copy_packed_to_device(self, batch_index: int, staging_index: int) -> None:
563 plan = self._packed_batch_plans[batch_index]
564 arena = self._require_staging_arena(staging_index)
565 for dtype, region in plan.regions.items():
566 host_buffer = self._host_buffers[batch_index][dtype]
567 staging_view = arena.dtype_views[dtype].narrow(0, 0, region.numel)
568 self._copy_tensor(staging_view, host_buffer)
569
570 def _copy_packed_to_host(self, batch_index: int, staging_index: int) -> None:
571 plan = self._packed_batch_plans[batch_index]
572 arena = self._require_staging_arena(staging_index)
573 for dtype, region in plan.regions.items():
574 host_buffer = self._host_buffers[batch_index][dtype]
575 staging_view = arena.dtype_views[dtype].narrow(0, 0, region.numel)
576 self._copy_tensor(host_buffer, staging_view)
577
578 def _require_staging_arena(self, staging_index: int) -> _StagingArena:
579 arena = self._staging_arenas[staging_index]
580 if arena is None:
581 raise RuntimeError(f"Packed MindSpore staging arena {staging_index} is not materialized.")
582 return arena
583
584 def _make_cpu_storage_view(self, host_buffer: Any, element_offset: int, shape: Sequence[int]) -> Any:
585 """Create a pinned CPU view directly from the host buffer storage."""
586 shape = tuple(shape)
587 stride = []
588 running_stride = 1
589 for dim in reversed(shape):
590 stride.append(running_stride)
591 running_stride *= int(dim)
592 host_view = ms.mint.empty(
593 (0,), dtype=host_buffer.dtype, device=host_buffer.device
594 ).set_(host_buffer.untyped_storage(), element_offset, shape, tuple(reversed(stride)))
595 if not _is_cpu_tensor(host_view):
596 raise RuntimeError(
597 f"Packed optimizer host view must remain on CPU, but got device {host_view.device!r}."
598 )
599 return host_view
600
601 def _copy_cpu_tensor(self, target: Any, source: Any) -> None:
602 """Copy CPU tensor data directly, including non-zero storage offsets."""
603 if isinstance(target, ms.Parameter):
604 target = target.data
605 if isinstance(source, ms.Parameter):
606 source = source.data
607 if not _is_cpu_tensor(target) or not _is_cpu_tensor(source):
608 raise RuntimeError(
609 "Packed host memcpy requires CPU source and target tensors, but got "
610 f"target device {getattr(target, 'device', None)!r} and "
611 f"source device {getattr(source, 'device', None)!r}."
612 )
613 if int(target.numel()) != int(source.numel()) or target.dtype != source.dtype:
614 raise RuntimeError(
615 "Packed host memcpy requires matching tensor sizes and dtypes, but got "
616 f"target=({target.numel()}, {target.dtype}) and source=({source.numel()}, {source.dtype})."
617 )
618 itemsize = int(target.itemsize)
619 target_ptr = int(target.untyped_storage().data_ptr()) + int(target.storage_offset()) * itemsize
620 source_ptr = int(source.untyped_storage().data_ptr()) + int(source.storage_offset()) * itemsize
621 ctypes.memmove(target_ptr, source_ptr, int(target.numel()) * itemsize)
622
623 def _storage_tensor(self, tensor: Any) -> Any:
624 """Return the local storage-owning tensor for a parameter or DTensor."""
625 # ParameterDTensor owns its real NPU allocation through _local_tensor.
626 # Inspect it before unwrapping Parameter.data, which may expose a
627 # different wrapper storage and leave the backing allocation resident.
628 local_tensor = getattr(tensor, "_local_tensor", None)
629 if local_tensor is not None:
630 return local_tensor
631 if isinstance(tensor, ms.Parameter):
632 tensor = tensor.data
633 local_tensor = getattr(tensor, "_local_tensor", None)
634 if local_tensor is not None:
635 return local_tensor
636 if hasattr(tensor, "to_local"):
637 tensor = tensor.to_local()
638 return tensor
639
640 @staticmethod
641 def _debug_storage_size(tensor: Any) -> Any:
642 """Return storage size for diagnostics without masking runtime behavior."""
643 if tensor is None:
644 return None
645 try:
646 return tensor.untyped_storage().size()
647 except (AttributeError, RuntimeError):
648 return "error"
649
650 @staticmethod
651 def _align_bytes(num_bytes: int) -> int:
652 return ((num_bytes + _PACKED_ALIGNMENT_BYTES - 1) // _PACKED_ALIGNMENT_BYTES) * _PACKED_ALIGNMENT_BYTES
653
654
655class MindSporeSwapOptimizer:
656 """MindSpore callable optimizer wrapper for state swap."""
657
658 _is_swap_optimizer = True
659 _adapters = (MindSporeNativeAdamAdapter, MindSporeNativeAdamWAdapter, MindFormersAdamWAdapter)
660
661 def __init__(self, optimizer: Any, config: Any) -> None:
662 self.optimizer = optimizer
663 self.config = config
664 self.runtime = MindSporeSwapRuntime(config)
665 self.adapter = self._build_adapter()
666 self.adapter.validate()
667 # MindSpore Adam states already exist at optimizer construction, so move them off device
668 # before the first forward/backward peak and let the first optimizer update prefetch them.
669 initial_slots = tuple(self.adapter.initial_slots())
670 self.runtime.offload_initial_slots(initial_slots)
671 if self.runtime.packed_enabled:
672 layout_units = self.adapter.packed_layout_units()
673 layout_batches = self.runtime.partition(layout_units)
674 self.runtime.prepare_packed_host(layout_batches)
675 self.adapter.publish_packed_state()
676
677 def __getattr__(self, name: str) -> Any:
678 """Delegate unknown attributes to the base optimizer."""
679 return getattr(self.optimizer, name)
680
681 def __call__(self, *args: Any, **kwargs: Any) -> Any:
682 """Call construct for MindSpore optimizer compatibility."""
683 return self.construct(*args, **kwargs)
684
685 def construct(self, *args: Any, **kwargs: Any) -> Any:
686 """Run one optimizer update with pipeline state swap."""
687 with _no_grad():
688 step_context = self.adapter.prepare_step(*args, **kwargs)
689 units = self.adapter.iter_update_units(step_context)
690 batches = self.runtime.partition(units)
691 if self.runtime.packed_enabled and self.runtime.prepare_packed_host(batches):
692 self.adapter.publish_packed_state()
693 result = self.runtime.run_pipeline(batches, step_context, self.adapter.step_batch)
694 self.adapter.finish_step(step_context)
695 return tuple(result)
696
697 def state_dict(self) -> Dict[str, Any]:
698 """Return checkpoint-safe state dict using CPU mirrors for swappable tensors."""
699 self.runtime.synchronize_cpu_mirrors(self.adapter.all_slots())
700 return self.adapter.checkpoint_state_dict()
701
702 def load_state_dict(self, state_dict: Dict[str, Any]) -> None:
703 """Load checkpoint-safe state dict while keeping swappable tensors on CPU mirrors."""
704 self.adapter.load_checkpoint_state_dict(state_dict)
705
706 def _build_adapter(self):
707 for adapter_cls in self._adapters:
708 if adapter_cls.matches(self.optimizer):
709 return adapter_cls(self.optimizer, self.config, self.runtime)
710 raise ValueError(
711 "Swap optimizer only supports mindspore.nn.Adam, mindspore.nn.AdamWeightDecay "
712 "(and compatible nn.AdamW aliases), and mindformers.pynative.optimizer.adamw.AdamW. "
713 f"Got {type(self.optimizer)!r}."
714 )
715
716
717def get_swap_optimizer():
718 """Return the MindSpore optimizer-state swap wrapper class."""
719 return MindSporeSwapOptimizer
720
721
722def _is_float_dtype(dtype: Any) -> bool:
723 text = str(dtype).lower()
724 return "float" in text or "bfloat" in text
725
726
727def _is_cpu_tensor(tensor: Any) -> bool:
728 device = getattr(tensor, "device", None)
729 if device is None:
730 return False
731 # MindSpore may render a host device as ``CPU`` or ``CPU:0`` depending
732 # on the backend/version. Both identify host memory for raw memcpy.
733 return str(device).strip().lower().split(":", 1)[0] == "cpu"