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

441 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"""Torch optimizer state swap wrapper.""" 

16# pylint: disable=protected-access 

17 

18from __future__ import annotations 

19 

20import contextlib 

21from dataclasses import dataclass, field 

22from typing import Any, Dict, Iterable, List, Optional, Sequence 

23 

24import torch 

25 

26from hyper_parallel.core.optimizer.swap_optimizer_base import PipelineSwapRuntime, SwapSlot 

27from hyper_parallel.platform import get_platform 

28from hyper_parallel.platform.torch.swap_optimizer.adapters import ( 

29 TorchHyperAdamWAdapter, 

30 TorchNativeAdamAdapter, 

31 TorchNativeAdamWAdapter, 

32) 

33 

34 

35_PACKED_ALIGNMENT_BYTES = 512 

36platform = get_platform() 

37 

38 

39@dataclass 

40class _PackedBatchRegion: 

41 """One dtype-contiguous host range transferred for a pipeline batch.""" 

42 

43 dtype: Any 

44 host_offset: int 

45 numel: int 

46 slots: List[SwapSlot] 

47 

48 

49@dataclass 

50class _PackedBatchPlan: 

51 """Packed transfer regions for one optimizer pipeline batch.""" 

52 

53 regions: Dict[Any, _PackedBatchRegion] = field(default_factory=dict) 

54 

55 

56@dataclass 

57class _StagingArena: 

58 """One raw device allocation and its dtype-specific views.""" 

59 

60 raw_buffer: Any 

61 dtype_views: Dict[Any, Any] = field(default_factory=dict) 

62 layout_signature: Any = None 

63 

64 

65class TorchSwapRuntime(PipelineSwapRuntime): 

66 """Torch tensor storage/copy runtime.""" 

67 

68 def __init__(self, config: Any) -> None: 

69 super().__init__(config) 

70 self._packed_enabled = bool(getattr(config, "packed_swap", True)) 

71 self._host_buffers: Dict[Any, Any] = {} 

72 self._host_layout_signature = () 

73 self._packed_batch_plans: List[_PackedBatchPlan] = [] 

74 self._staging_arenas: List[Optional[_StagingArena]] = [None, None] 

75 self._packed_ready_events: Dict[int, Any] = {} 

76 self._packed_offload_events: Dict[int, Any] = {} 

77 self._packed_tail_event: Optional[Any] = None 

78 self._packed_device_views: Dict[Any, Any] = {} 

79 

80 @property 

81 def packed_enabled(self) -> bool: 

82 """Return whether this runtime may build packed state candidates.""" 

83 return self._packed_enabled 

84 

85 def populate_slot_metadata(self, slot: SwapSlot, template: Any) -> None: 

86 """Populate stable logical tensor metadata without allocating device state.""" 

87 storage_tensor = self._storage_tensor(template) 

88 slot.shape = tuple(storage_tensor.shape) 

89 slot.dtype = storage_tensor.dtype 

90 slot.device = storage_tensor.device 

91 slot.numel = int(storage_tensor.numel()) 

92 slot.storage_nbytes = slot.numel * int(storage_tensor.element_size()) 

93 

94 def is_packable_template(self, tensor: Any, min_numel: int) -> bool: 

95 """Return whether a state shaped like ``tensor`` can use packed staging.""" 

96 if not self._packed_enabled: 

97 return False 

98 return self.is_swappable_tensor(tensor, min_numel) 

99 

100 @staticmethod 

101 def is_distributed_tensor(tensor: Any) -> bool: 

102 """Return whether ``tensor`` exposes a DTensor local shard.""" 

103 return tensor is not None and callable(getattr(tensor, "to_local", None)) 

104 

105 def prepare_packed_host(self, slots: Sequence[SwapSlot]) -> None: 

106 """Pack logical optimizer states into persistent pinned buffers by dtype.""" 

107 if not self._packed_enabled: 

108 return 

109 packed_slots = [slot for slot in slots if slot.swappable and slot.packed] 

110 signature = tuple((id(slot), slot.dtype, slot.numel) for slot in packed_slots) 

111 if signature == self._host_layout_signature: 

112 return 

113 

114 slots_by_dtype: Dict[Any, List[SwapSlot]] = {} 

115 for slot in packed_slots: 

116 slots_by_dtype.setdefault(slot.dtype, []).append(slot) 

117 

118 new_buffers: Dict[Any, Any] = {} 

119 for dtype, dtype_slots in slots_by_dtype.items(): 

120 total_numel = sum(slot.numel for slot in dtype_slots) 

121 host_buffer = torch.empty(total_numel, dtype=dtype, device="cpu", pin_memory=True) 

122 new_buffers[dtype] = host_buffer 

123 host_offset = 0 

124 for slot in dtype_slots: 

125 flat_view = host_buffer.narrow(0, host_offset, slot.numel) 

126 host_view = flat_view.view(slot.shape) 

127 source = slot.cpu_tensor if slot.cpu_tensor is not None else slot.tensor 

128 if source is None: 

129 host_view.zero_() 

130 else: 

131 source_tensor = self._storage_tensor(source) 

132 host_view.copy_(source_tensor.detach().reshape(-1).view(slot.shape), non_blocking=False) 

133 if source_tensor.device.type != "cpu" and source is slot.tensor: 

134 self.release_device_storage(slot) 

135 slot.host_offset = host_offset 

136 slot.cpu_tensor = host_view 

137 slot.bind_tensor(host_view) 

138 slot.state = "host" 

139 slot.event = None 

140 host_offset += slot.numel 

141 

142 self._host_buffers = new_buffers 

143 self._host_layout_signature = signature 

144 

145 def is_swappable_tensor(self, tensor: Any, min_numel: int) -> bool: 

146 """Return whether ``tensor`` can participate in swap.""" 

147 storage_tensor = self._storage_tensor(tensor) 

148 if not isinstance(storage_tensor, torch.Tensor): 

149 return False 

150 if not storage_tensor.is_floating_point(): 

151 return False 

152 if int(storage_tensor.numel()) < int(min_numel): 

153 return False 

154 if storage_tensor.is_sparse: 

155 return False 

156 if not storage_tensor.is_contiguous(): 

157 return False 

158 if storage_tensor.device.type == "cpu": 

159 return False 

160 try: 

161 storage_size = int(storage_tensor.untyped_storage().size()) 

162 expected_size = int(storage_tensor.numel()) * int(storage_tensor.element_size()) 

163 if storage_size != expected_size: 

164 return False 

165 except RuntimeError: 

166 return False 

167 return True 

168 

169 def storage_nbytes(self, tensor: Any) -> int: 

170 """Return storage bytes for a Torch tensor.""" 

171 storage_tensor = self._storage_tensor(tensor) 

172 if not isinstance(storage_tensor, torch.Tensor): 

173 return 0 

174 try: 

175 return int(storage_tensor.untyped_storage().size()) 

176 except RuntimeError: 

177 return int(storage_tensor.numel()) * int(storage_tensor.element_size()) 

178 

179 def make_cpu_tensor(self, tensor: Any) -> Any: 

180 """Create a CPU mirror tensor.""" 

181 storage_tensor = self._storage_tensor(tensor) 

182 if isinstance(storage_tensor, torch.Tensor): 

183 source = storage_tensor.detach() 

184 try: 

185 cpu_tensor = torch.empty_like(source, device="cpu", pin_memory=True) 

186 except RuntimeError: 

187 cpu_tensor = torch.empty_like(source, device="cpu") 

188 cpu_tensor.copy_(source, non_blocking=True) 

189 return cpu_tensor 

190 raise ValueError(f"Expected torch.Tensor for CPU mirror, got {type(tensor)!r}.") 

191 

192 def make_zero_cpu_tensor_like(self, tensor: Any) -> Any: 

193 """Create a zero-valued CPU mirror without materializing device state.""" 

194 storage_tensor = self._storage_tensor(tensor) 

195 if not isinstance(storage_tensor, torch.Tensor): 

196 raise ValueError(f"Expected torch.Tensor for CPU mirror, got {type(tensor)!r}.") 

197 try: 

198 cpu_tensor = torch.empty_like(storage_tensor, device="cpu", pin_memory=True) 

199 except RuntimeError: 

200 cpu_tensor = torch.empty_like(storage_tensor, device="cpu") 

201 cpu_tensor.zero_() 

202 return cpu_tensor 

203 

204 def make_device_tensor_like(self, param: Any, saved_tensor: Any) -> Any: 

205 """Create a live state tensor on the parameter device.""" 

206 if not isinstance(saved_tensor, torch.Tensor): 

207 raise ValueError(f"Expected torch.Tensor in optimizer state, got {type(saved_tensor)!r}.") 

208 return saved_tensor.detach().to(device=param.device, dtype=saved_tensor.dtype).clone() 

209 

210 def make_empty_device_tensor_like(self, param: Any, saved_tensor: Any) -> Any: 

211 """Create an uninitialized live state tensor shell on the parameter device.""" 

212 if not isinstance(saved_tensor, torch.Tensor): 

213 raise ValueError(f"Expected torch.Tensor in optimizer state, got {type(saved_tensor)!r}.") 

214 return torch.empty_like(saved_tensor, device=param.device, dtype=saved_tensor.dtype) 

215 

216 def copy_to_device(self, slot: SwapSlot) -> None: 

217 """Copy one CPU mirror to device tensor.""" 

218 if slot.state != "host": 

219 return 

220 if slot.cpu_tensor is None: 

221 return 

222 self._storage_tensor(slot.tensor).copy_(slot.cpu_tensor, non_blocking=True) 

223 slot.state = "h2d" 

224 

225 def wait_prefetch_slot(self, slot: SwapSlot) -> None: 

226 """Torch fallback copies are synchronous on CPU and stream-ordered on device.""" 

227 slot.state = "device" 

228 

229 def copy_to_cpu(self, slot: SwapSlot) -> None: 

230 """Copy one device tensor to CPU mirror.""" 

231 source = self._storage_tensor(slot.tensor) 

232 if slot.cpu_tensor is None: 

233 slot.cpu_tensor = self.make_cpu_tensor(source) 

234 else: 

235 slot.cpu_tensor.copy_(source.detach(), non_blocking=True) 

236 slot.state = "d2h" 

237 

238 def wait_offload_slot(self, slot: SwapSlot) -> None: 

239 """Release device storage after D2H copy completes.""" 

240 if self._storage_tensor(slot.tensor).device.type != "cpu": 

241 self.release_device_storage(slot) 

242 slot.state = "host" 

243 

244 def restore_device_storage(self, slot: SwapSlot) -> None: 

245 """Restore device tensor storage before H2D.""" 

246 storage_tensor = self._storage_tensor(slot.tensor) 

247 if storage_tensor.device.type == "cpu": 

248 return 

249 storage = storage_tensor.untyped_storage() 

250 if storage.size() != slot.storage_nbytes: 

251 storage.resize_(slot.storage_nbytes) 

252 

253 def release_device_storage(self, slot: SwapSlot) -> None: 

254 """Release device storage for a swappable tensor.""" 

255 storage_tensor = self._storage_tensor(slot.tensor) 

256 if storage_tensor.device.type == "cpu": 

257 return 

258 storage = storage_tensor.untyped_storage() 

259 if storage.size() != 0: 

260 storage.resize_(0) 

261 

262 def current_stream(self) -> Any: 

263 """Return the current compute stream.""" 

264 return platform.get_current_stream() 

265 

266 def new_stream(self) -> Any: 

267 """Create the copy stream.""" 

268 return platform.new_stream() 

269 

270 def stream_context(self, stream: Any): 

271 """Return a stream context for ``stream``.""" 

272 if stream is None: 

273 return contextlib.nullcontext() 

274 return platform.get_stream_context()(stream) 

275 

276 def record_event(self, stream: Any = None) -> Any: 

277 """Record an event on ``stream``.""" 

278 event = platform.new_event() 

279 if stream is None: 

280 event.record() 

281 else: 

282 event.record(stream) 

283 return event 

284 

285 def wait_event(self, event: Any, stream: Any = None) -> None: 

286 """Make ``stream`` wait for ``event``.""" 

287 if event is None: 

288 return 

289 if stream is None: 

290 event.synchronize() 

291 return 

292 event.wait(stream) 

293 

294 def supports_packed_pipeline(self, batches: Sequence[Sequence[Any]]) -> bool: 

295 """Make the final, step-specific decision to use the packed pipeline.""" 

296 if not self._packed_enabled or not batches: 

297 return False 

298 swappable_slots = [ 

299 slot 

300 for batch in batches 

301 for unit in batch 

302 for slot in unit.slots 

303 if slot.swappable 

304 ] 

305 if not swappable_slots: 

306 return False 

307 devices = {slot.device for slot in swappable_slots} 

308 eligible = len(devices) == 1 and None not in devices and all( 

309 slot.packed 

310 and slot.cpu_tensor is not None 

311 and slot.dtype in self._host_buffers 

312 for slot in swappable_slots 

313 ) 

314 if not eligible and any( 

315 slot.packed and slot.logical_tensor is not None 

316 for slot in swappable_slots 

317 ): 

318 raise RuntimeError( 

319 "Packed DTensor optimizer states cannot fall back to per-tensor swap after " 

320 f"host packing; packed pipeline eligibility failed for local devices {sorted(map(str, devices))}." 

321 ) 

322 return eligible 

323 

324 @staticmethod 

325 def _storage_tensor(tensor: Any) -> Any: 

326 """Return the local tensor whose storage is managed by this runtime.""" 

327 to_local = getattr(tensor, "to_local", None) 

328 return to_local() if callable(to_local) else tensor 

329 

330 def begin_packed_step(self, batches: Sequence[Sequence[Any]]) -> None: 

331 """Build batch transfer plans and materialize two raw staging buffers.""" 

332 first_slot = next( 

333 slot 

334 for batch in batches 

335 for unit in batch 

336 for slot in unit.slots 

337 if slot.swappable 

338 ) 

339 # FSDP may leave gradient reductions/reshards queued on auxiliary 

340 # streams. A current-stream event cannot cover those streams, so the 

341 # device-wide synchronization is required before optimizer reads state. 

342 getattr(torch, first_slot.device.type).synchronize(first_slot.device) 

343 self._packed_tail_event = None 

344 self._packed_batch_plans = [self._build_packed_batch_plan(batch) for batch in batches] 

345 max_numel_by_dtype: Dict[Any, int] = {} 

346 device = None 

347 for batch_plan in self._packed_batch_plans: 

348 for dtype, region in batch_plan.regions.items(): 

349 max_numel_by_dtype[dtype] = max(max_numel_by_dtype.get(dtype, 0), region.numel) 

350 if device is None and region.slots: 

351 device = region.slots[0].device 

352 if device is None: 

353 raise RuntimeError("Packed optimizer pipeline has no device-resident state metadata.") 

354 

355 dtype_layouts = {} 

356 byte_offset = 0 

357 for dtype in sorted(max_numel_by_dtype, key=str): 

358 byte_offset = self._align_bytes(byte_offset) 

359 element_size = int(self._host_buffers[dtype].element_size()) 

360 num_bytes = max_numel_by_dtype[dtype] * element_size 

361 dtype_layouts[dtype] = (byte_offset, num_bytes) 

362 byte_offset += num_bytes 

363 total_bytes = self._align_bytes(byte_offset) 

364 

365 for staging_index in range(2): 

366 arena = self._materialize_staging_arena(staging_index, total_bytes, device) 

367 layout_signature = tuple( 

368 (dtype, offset, num_bytes) for dtype, (offset, num_bytes) in dtype_layouts.items() 

369 ) 

370 if arena.layout_signature != layout_signature: 

371 arena.dtype_views = { 

372 dtype: arena.raw_buffer.narrow(0, offset, num_bytes).view(dtype) 

373 for dtype, (offset, num_bytes) in dtype_layouts.items() 

374 } 

375 arena.layout_signature = layout_signature 

376 self._drop_packed_views(staging_index) 

377 self._packed_ready_events = {} 

378 self._packed_offload_events = {} 

379 

380 def enqueue_packed_prefetch(self, batch_index: int, staging_index: int) -> None: 

381 """Enqueue one packed H2D and record its ready event.""" 

382 copy_stream = self._get_copy_stream() 

383 with self.stream_context(copy_stream): 

384 self._copy_packed_to_device(batch_index, staging_index) 

385 ready_event = self.record_event(copy_stream) 

386 self._packed_ready_events[batch_index] = ready_event 

387 self._packed_tail_event = ready_event 

388 

389 def wait_packed_prefetch(self, batch_index: int, staging_index: int) -> None: 

390 """Order the compute stream after the batch's packed transfer chain.""" 

391 del staging_index 

392 ready_event = self._packed_ready_events.get(batch_index) 

393 if ready_event is None: 

394 raise RuntimeError(f"Packed optimizer batch {batch_index} has no ready event.") 

395 self.wait_event(ready_event, self.current_stream()) 

396 

397 def activate_packed_batch(self, batch_index: int, staging_index: int) -> None: 

398 """Bind each swap slot to its slice of one staging arena.""" 

399 batch_plan = self._packed_batch_plans[batch_index] 

400 arena = self._require_staging_arena(staging_index) 

401 for dtype, region in batch_plan.regions.items(): 

402 dtype_view = arena.dtype_views[dtype] 

403 for slot in region.slots: 

404 relative_offset = slot.host_offset - region.host_offset 

405 cache_key = ( 

406 staging_index, 

407 id(arena.raw_buffer), 

408 id(dtype_view), 

409 id(slot), 

410 dtype, 

411 relative_offset, 

412 slot.numel, 

413 slot.shape, 

414 ) 

415 device_view = self._packed_device_views.get(cache_key) 

416 if device_view is None: 

417 device_view = dtype_view.narrow(0, relative_offset, slot.numel).view(slot.shape) 

418 self._packed_device_views[cache_key] = device_view 

419 slot.bind_tensor(device_view) 

420 slot.state = "device" 

421 slot.event = None 

422 

423 def enqueue_packed_offload_prefetch( 

424 self, 

425 batch_index: int, 

426 next_index: Optional[int], 

427 staging_index: int, 

428 ) -> None: 

429 """Serialize current D2H and next same-parity H2D on the copy stream.""" 

430 copy_stream = self._get_copy_stream() 

431 compute_event = self._record_current_stream_event() 

432 with self.stream_context(copy_stream): 

433 self.wait_event(compute_event, copy_stream) 

434 self._copy_packed_to_host(batch_index, staging_index) 

435 if next_index is not None: 

436 self._copy_packed_to_device(next_index, staging_index) 

437 chain_event = self.record_event(copy_stream) 

438 self._packed_offload_events[batch_index] = chain_event 

439 self._packed_tail_event = chain_event 

440 if next_index is not None: 

441 self._packed_ready_events[next_index] = chain_event 

442 

443 def wait_packed_offload(self, batch_index: int) -> None: 

444 """Order the compute stream after a packed transfer chain.""" 

445 offload_event = self._packed_offload_events.get(batch_index) 

446 if offload_event is None: 

447 raise RuntimeError(f"Packed optimizer batch {batch_index} has no offload event.") 

448 self.wait_event(offload_event, self.current_stream()) 

449 

450 def finish_packed_offload(self, batch_index: int) -> None: 

451 """Make persistent pinned views authoritative after D2H completion.""" 

452 batch_plan = self._packed_batch_plans[batch_index] 

453 for region in batch_plan.regions.values(): 

454 for slot in region.slots: 

455 slot.bind_tensor(slot.cpu_tensor) 

456 slot.state = "host" 

457 slot.event = None 

458 

459 def end_packed_step(self) -> None: 

460 """Drain the packed copy chain and release step-local staging storage.""" 

461 if self._packed_tail_event is not None: 

462 # All D2H copies are serialized on the copy stream. Waiting only 

463 # for its tail is sufficient before detaching views and releasing 

464 # the step-local device allocation. 

465 self.wait_event(self._packed_tail_event, None) 

466 active_slots = { 

467 id(slot): slot 

468 for batch_plan in self._packed_batch_plans 

469 for region in batch_plan.regions.values() 

470 for slot in region.slots 

471 if slot.state == "device" 

472 } 

473 if active_slots: 

474 compute_stream = self.current_stream() 

475 synchronize = getattr(compute_stream, "synchronize", None) 

476 if synchronize is not None: 

477 synchronize() 

478 for slot in active_slots.values(): 

479 slot.cpu_tensor.copy_(self._storage_tensor(slot.tensor).detach(), non_blocking=False) 

480 slot.bind_tensor(slot.cpu_tensor) 

481 slot.state = "host" 

482 slot.event = None 

483 for arena in self._staging_arenas: 

484 if arena is None: 

485 continue 

486 # Drop views before shrinking the raw storage; otherwise a view 

487 # can keep the device allocation alive after the step ends. 

488 arena.dtype_views = {} 

489 arena.layout_signature = None 

490 storage = arena.raw_buffer.untyped_storage() 

491 if storage.size() != 0: 

492 storage.resize_(0) 

493 self._packed_device_views = {} 

494 self._packed_batch_plans = [] 

495 self._packed_ready_events = {} 

496 self._packed_offload_events = {} 

497 self._packed_tail_event = None 

498 

499 def _build_packed_batch_plan(self, batch: Sequence[Any]) -> _PackedBatchPlan: 

500 """Group a batch's packed slots into contiguous host regions by dtype.""" 

501 slots_by_dtype: Dict[Any, List[SwapSlot]] = {} 

502 seen_slots = set() 

503 for unit in batch: 

504 for slot in unit.slots: 

505 if not slot.swappable or not slot.packed or id(slot) in seen_slots: 

506 continue 

507 slots_by_dtype.setdefault(slot.dtype, []).append(slot) 

508 seen_slots.add(id(slot)) 

509 

510 regions = {} 

511 for dtype, slots in slots_by_dtype.items(): 

512 slots.sort(key=lambda slot: slot.host_offset) 

513 host_offset = slots[0].host_offset 

514 expected_offset = host_offset 

515 for slot in slots: 

516 if slot.host_offset != expected_offset: 

517 raise RuntimeError( 

518 f"Packed optimizer batch has a non-contiguous {dtype} host range at slot {slot.name!r}." 

519 ) 

520 expected_offset += slot.numel 

521 regions[dtype] = _PackedBatchRegion(dtype, host_offset, expected_offset - host_offset, slots) 

522 return _PackedBatchPlan(regions) 

523 

524 def _materialize_staging_arena(self, staging_index: int, total_bytes: int, device: Any) -> _StagingArena: 

525 """Allocate or resize one packed device staging arena for the requested layout.""" 

526 arena = self._staging_arenas[staging_index] 

527 if arena is None or arena.raw_buffer.device != device: 

528 raw_buffer = torch.empty(total_bytes, dtype=torch.uint8, device=device) 

529 arena = _StagingArena(raw_buffer) 

530 self._staging_arenas[staging_index] = arena 

531 self._drop_packed_views(staging_index) 

532 return arena 

533 

534 raw_buffer = arena.raw_buffer 

535 storage = raw_buffer.untyped_storage() 

536 if storage.size() < total_bytes: 

537 storage.resize_(total_bytes) 

538 arena.dtype_views = {} 

539 arena.layout_signature = None 

540 self._drop_packed_views(staging_index) 

541 raw_buffer.set_(storage, 0, (total_bytes,), (1,)) 

542 return arena 

543 

544 def _drop_packed_views(self, staging_index: int) -> None: 

545 """Drop cached views for one arena after its storage/layout changes.""" 

546 self._packed_device_views = { 

547 key: value for key, value in self._packed_device_views.items() if key[0] != staging_index 

548 } 

549 

550 def _copy_packed_to_device(self, batch_index: int, staging_index: int) -> None: 

551 batch_plan = self._packed_batch_plans[batch_index] 

552 arena = self._require_staging_arena(staging_index) 

553 for dtype, region in batch_plan.regions.items(): 

554 host_view = self._host_buffers[dtype].narrow(0, region.host_offset, region.numel) 

555 arena.dtype_views[dtype].narrow(0, 0, region.numel).copy_(host_view, non_blocking=True) 

556 

557 def _copy_packed_to_host(self, batch_index: int, staging_index: int) -> None: 

558 batch_plan = self._packed_batch_plans[batch_index] 

559 arena = self._require_staging_arena(staging_index) 

560 for dtype, region in batch_plan.regions.items(): 

561 host_view = self._host_buffers[dtype].narrow(0, region.host_offset, region.numel) 

562 host_view.copy_(arena.dtype_views[dtype].narrow(0, 0, region.numel), non_blocking=True) 

563 

564 def _require_staging_arena(self, staging_index: int) -> _StagingArena: 

565 arena = self._staging_arenas[staging_index] 

566 if arena is None: 

567 raise RuntimeError(f"Packed optimizer staging arena {staging_index} is not materialized.") 

568 return arena 

569 

570 @staticmethod 

571 def _align_bytes(num_bytes: int) -> int: 

572 return ((num_bytes + _PACKED_ALIGNMENT_BYTES - 1) // _PACKED_ALIGNMENT_BYTES) * _PACKED_ALIGNMENT_BYTES 

573 

574 

575class TorchSwapOptimizer(torch.optim.Optimizer): 

576 """Torch optimizer wrapper for Adam/AdamW state swap.""" 

577 

578 _is_swap_optimizer = True 

579 _adapters = (TorchHyperAdamWAdapter, TorchNativeAdamAdapter, TorchNativeAdamWAdapter) 

580 

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

582 # Do not call ``torch.optim.Optimizer.__init__``: the wrapped base 

583 # optimizer already owns param_groups/state/defaults. Inheriting keeps 

584 # PyTorch LR schedulers and isinstance checks happy while this wrapper 

585 # delegates all optimizer state to ``self.optimizer``. 

586 self.optimizer = optimizer 

587 self.config = config 

588 self.runtime = TorchSwapRuntime(config) 

589 self.adapter = self._build_adapter() 

590 self.adapter.validate() 

591 # Torch Adam states are normally lazy, but callers may materialize them 

592 # before wrapping to avoid first-step initialization in the measured loop. 

593 initial_slots = tuple(self.adapter.initial_slots()) 

594 self.runtime.offload_initial_slots(initial_slots) 

595 self.runtime.prepare_packed_host(initial_slots) 

596 self.adapter.publish_packed_state() 

597 

598 def __getattr__(self, name: str) -> Any: 

599 """Delegate unknown attributes to the base optimizer.""" 

600 return getattr(self.optimizer, name) 

601 

602 @property 

603 def param_groups(self): 

604 """Proxy parameter groups.""" 

605 return self.optimizer.param_groups 

606 

607 @param_groups.setter 

608 def param_groups(self, value) -> None: 

609 self.optimizer.param_groups = value 

610 

611 @property 

612 def state(self): 

613 """Proxy optimizer state.""" 

614 return self.optimizer.state 

615 

616 @property 

617 def defaults(self): 

618 """Proxy optimizer defaults.""" 

619 return self.optimizer.defaults 

620 

621 def add_param_group(self, param_group: Dict[str, Any]) -> None: 

622 """Proxy param group addition.""" 

623 self.optimizer.add_param_group(param_group) 

624 

625 def zero_grad(self, set_to_none: bool = True) -> None: 

626 """Proxy gradient clearing.""" 

627 self.optimizer.zero_grad(set_to_none=set_to_none) 

628 

629 def step(self, closure: Optional[Any] = None) -> Any: 

630 """Run one optimizer step with pipeline state swap.""" 

631 if closure is not None: 

632 raise ValueError("Swap optimizer does not support closure.") 

633 with self._no_grad_context(): 

634 step_context = self.adapter.prepare_step() 

635 units = self.adapter.iter_update_units(step_context) 

636 batches = self.runtime.partition(units) 

637 self.runtime.run_pipeline(batches, step_context, self.adapter.step_batch) 

638 return self.adapter.finish_step(step_context) 

639 

640 def state_dict(self) -> Dict[str, Any]: 

641 """Return optimizer state dict using CPU mirrors for swappable tensors.""" 

642 return self.adapter.checkpoint_state_dict() 

643 

644 def load_state_dict(self, state_dict: Dict[str, Any]) -> None: 

645 """Load optimizer state dict while keeping swappable tensors on CPU mirrors.""" 

646 self.adapter.load_checkpoint_state_dict(state_dict) 

647 

648 def _build_adapter(self): 

649 for adapter_cls in self._adapters: 

650 if adapter_cls.matches(self.optimizer): 

651 return adapter_cls(self.optimizer, self.config, self.runtime) 

652 raise ValueError( 

653 "Swap optimizer only supports torch.optim.Adam, torch.optim.AdamW, " 

654 "and hyper_parallel.core.optimizer.adamw.AdamW on the Torch backend. " 

655 f"Got {type(self.optimizer)!r}." 

656 ) 

657 

658 @contextlib.contextmanager 

659 def _no_grad_context(self) -> Iterable[None]: 

660 with torch.no_grad(): 

661 yield 

662 

663 

664def get_swap_optimizer(): 

665 """Return the Torch optimizer-state swap wrapper class.""" 

666 return TorchSwapOptimizer