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
« 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
18from __future__ import annotations
20import contextlib
21from dataclasses import dataclass, field
22from typing import Any, Dict, Iterable, List, Optional, Sequence
24import torch
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)
35_PACKED_ALIGNMENT_BYTES = 512
36platform = get_platform()
39@dataclass
40class _PackedBatchRegion:
41 """One dtype-contiguous host range transferred for a pipeline batch."""
43 dtype: Any
44 host_offset: int
45 numel: int
46 slots: List[SwapSlot]
49@dataclass
50class _PackedBatchPlan:
51 """Packed transfer regions for one optimizer pipeline batch."""
53 regions: Dict[Any, _PackedBatchRegion] = field(default_factory=dict)
56@dataclass
57class _StagingArena:
58 """One raw device allocation and its dtype-specific views."""
60 raw_buffer: Any
61 dtype_views: Dict[Any, Any] = field(default_factory=dict)
62 layout_signature: Any = None
65class TorchSwapRuntime(PipelineSwapRuntime):
66 """Torch tensor storage/copy runtime."""
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] = {}
80 @property
81 def packed_enabled(self) -> bool:
82 """Return whether this runtime may build packed state candidates."""
83 return self._packed_enabled
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())
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)
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))
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
114 slots_by_dtype: Dict[Any, List[SwapSlot]] = {}
115 for slot in packed_slots:
116 slots_by_dtype.setdefault(slot.dtype, []).append(slot)
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
142 self._host_buffers = new_buffers
143 self._host_layout_signature = signature
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
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())
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}.")
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
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()
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)
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"
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"
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"
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"
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)
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)
262 def current_stream(self) -> Any:
263 """Return the current compute stream."""
264 return platform.get_current_stream()
266 def new_stream(self) -> Any:
267 """Create the copy stream."""
268 return platform.new_stream()
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)
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
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)
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
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
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.")
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)
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 = {}
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
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())
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
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
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())
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
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
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))
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)
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
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
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 }
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)
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)
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
570 @staticmethod
571 def _align_bytes(num_bytes: int) -> int:
572 return ((num_bytes + _PACKED_ALIGNMENT_BYTES - 1) // _PACKED_ALIGNMENT_BYTES) * _PACKED_ALIGNMENT_BYTES
575class TorchSwapOptimizer(torch.optim.Optimizer):
576 """Torch optimizer wrapper for Adam/AdamW state swap."""
578 _is_swap_optimizer = True
579 _adapters = (TorchHyperAdamWAdapter, TorchNativeAdamAdapter, TorchNativeAdamWAdapter)
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()
598 def __getattr__(self, name: str) -> Any:
599 """Delegate unknown attributes to the base optimizer."""
600 return getattr(self.optimizer, name)
602 @property
603 def param_groups(self):
604 """Proxy parameter groups."""
605 return self.optimizer.param_groups
607 @param_groups.setter
608 def param_groups(self, value) -> None:
609 self.optimizer.param_groups = value
611 @property
612 def state(self):
613 """Proxy optimizer state."""
614 return self.optimizer.state
616 @property
617 def defaults(self):
618 """Proxy optimizer defaults."""
619 return self.optimizer.defaults
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)
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)
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)
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()
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)
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 )
658 @contextlib.contextmanager
659 def _no_grad_context(self) -> Iterable[None]:
660 with torch.no_grad():
661 yield
664def get_swap_optimizer():
665 """Return the Torch optimizer-state swap wrapper class."""
666 return TorchSwapOptimizer