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

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