Diff Coverage

Diff: origin/master...HEAD, staged and unstaged changes

Source File Diff Coverage (%) Missing Lines
hyper_parallel/core/activation_checkpoint/__init__.py 100%  
hyper_parallel/core/activation_checkpoint/activation_checkpoint.py 70.6% 155-156,178,191,238
hyper_parallel/core/activation_checkpoint/pinned_memory_pool.py 88.0% 117,144,188-191,223-226,232-237,242-246
hyper_parallel/core/activation_checkpoint/swap.py 44.8% 204-207,258-264,269-276,409-412,414-416,624-629,632-638,640-641,715-718,720-721,859
hyper_parallel/platform/platform.py 100%  
hyper_parallel/platform/torch/activation_checkpoint/activation_swap.py 100%  
hyper_parallel/platform/torch/activation_checkpoint/sac.py 100%  
hyper_parallel/platform/torch/platform.py 50.0% 1629,1635,1659,1676
hyper_parallel/core/activation_checkpoint/activation_checkpoint.py
151
152
153
154
155
156
157
158
159
160
        if swap_inputs:
            unsupported.append("swap_inputs")
        if group_swap:
            unsupported.append("group_swap")
        if cpu_pool is not None:
            unsupported.append("cpu_pool")
        if context_fn is not None:
            unsupported.append("custom context_fn")
        if kwargs.get("use_reentrant", False):
            unsupported.append("use_reentrant=True")
174
175
176
177
178
179
180
181
        factories: list = [create_recompute_contexts]
        if policy_fn is not None:
            selective_kwargs = {"group_swap": group_swap}
            if cpu_pool is not None:
                selective_kwargs["cpu_pool"] = cpu_pool
            factories.append(partial(plat.create_selective_checkpoint_contexts, policy_fn, **selective_kwargs))
        if context_fn is not None:
            factories.append(context_fn)
187
188
189
190
191
192
193
194
195

    if swap_inputs:
        async_kwargs = {"group_swap": group_swap}
        if cpu_pool is not None:
            async_kwargs["cpu_pool"] = cpu_pool
        context = partial(plat.async_save_on_cpu, **async_kwargs)
    else:
        context = contextlib.nullcontext
    with context():
234
235
236
237
238
239
240
241
            "Use Torch-native non-reentrant checkpointing with SAVE/RECOMPUTE policies."
        )
    async_kwargs = {"policy_fn": policy_fn, "group_swap": group_swap}
    if cpu_pool is not None:
        async_kwargs["cpu_pool"] = cpu_pool
    with plat.async_save_on_cpu(**async_kwargs):
        return function(*args, **kwargs)

hyper_parallel/core/activation_checkpoint/pinned_memory_pool.py
113
114
115
116
117
118
119
120
121
        # event objects so each runtime event is queried only once per scan.
        event_results: Dict[int, Tuple[Any, bool]] = {}
        for capacity in list(self._pending):
            if capacity < minimum_capacity:
                continue
            still_pending = []
            for block in self._pending[capacity]:
                event = block.event
                event_key = id(event)
140
141
142
143
144
145
146
147
148
        """Associate a checked-out view with its owning allocation."""
        key = _storage_key(view)
        owner = self._blocks.get(key)
        if owner is not None and owner is not block:
            raise RuntimeError("Two host buffers exposed the same storage identity.")
        old_key = block.storage_key
        if old_key is not None and old_key != key and self._blocks.get(old_key) is block:
            del self._blocks[old_key]
        self._blocks[key] = block
184
185
186
187
188
189
190
191
192
193
194
195
            block = self._find_available_locked(aligned_size)
            if block is not None:
                try:
                    return self._checkout_locked(block, size)
                except Exception:
                    block.state = _AVAILABLE
                    self._available[block.capacity].append(block)
                    raise

            if self._total_allocated + aligned_size <= self._max_host_bytes:
                self._total_allocated += aligned_size
                reserved = True
219
220
221
222
223
224
225
226
227
228
229
230
            block = _Block(buffer, aligned_size, _IN_USE)
            try:
                with self._lock:
                    view = self._checkout_locked(block, size)
            except Exception:
                with self._lock:
                    self._total_allocated -= aligned_size
                raise
            return view

        event = block.event
        try:
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250

        event = block.event
        try:
            event.synchronize()
        except Exception:
            with self._lock:
                block.state = _PENDING
                block.event = event
                self._pending[block.capacity].insert(0, block)
            raise
        block.event = None
        try:
            with self._lock:
                return self._checkout_locked(block, size)
        except Exception:
            with self._lock:
                block.state = _AVAILABLE
                self._available[block.capacity].append(block)
            raise

    def release(self, tensor: Any, event: Optional[Any] = None) -> None:
        """Return an acquired view to the pool, optionally after an async event."""
        key = _storage_key(tensor)
hyper_parallel/core/activation_checkpoint/swap.py
200
201
202
203
204
205
206
207
208
209
210
211
    def release_cpu_buffer(self, event=None):
        """Release an explicitly pooled host buffer exactly once."""
        if self.cpu_pool is None or self._cpu_pool_buffer is None:
            return
        release_tensor = self.val_cpu if self.val_cpu is not None else self._cpu_pool_buffer
        self.cpu_pool.release(release_tensor, event=event)
        self._cpu_pool_buffer = None
        self.val_cpu = None

    def wait_load(self, release_event=None):
        """change state to device after async load is done"""
        if self._state == self.STATE_NON_TENSOR or self._keep_on_device or self._duplicate_swap:
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
                self.val_cpu = platform.empty_like(
                    self.val, device="cpu", pin_memory=True
                )
            else:
                logical_bytes = self.val.numel() * platform.get_element_size(self.val)
                self._cpu_pool_buffer = self.cpu_pool.acquire(logical_bytes)
                try:
                    self.val_cpu = self._cpu_pool_buffer.view(self.val.dtype).reshape(self.val.shape)
                except Exception:
                    self.release_cpu_buffer()
                    raise
        try:
            if self.cpu_pool is not None or self.is_slice_tensor:
                self.val_cpu.copy_(self.val, non_blocking=True)
            else:
                self.val_cpu.untyped_storage().copy_(self.val.untyped_storage(), non_blocking=True)
        except Exception:
            if self.cpu_pool is not None and self._cpu_pool_buffer is not None:
                release_event = platform.new_event()
                release_event.record(platform.get_current_stream())
                self.release_cpu_buffer(release_event)
            self.val_cpu = None
            raise
        self._state = self.STATE_D2H

    def wait_offload(self):
        """wait offload to host and free device memory"""
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
        self.clear()

    def release_cpu_buffers(self, event=None):
        """Release all explicitly pooled host buffers held by this storage."""
        def _release(x):
            if isinstance(x, SwapTensor):
                x.release_cpu_buffer(event=event)
            return x

        for storage_list in self.values():
            for item in storage_list:
                platform.tree_map(_release, item)

    def wait_offload(self):
        """wait offload for all tensors in swap storage"""
        def _wait_offload(x):
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
                        cpu_pool = bucket["cpu_pool"]
                        if cpu_pool is None:
                            cpu_buf = _get_cpu_pinned_buf(dtype_key, numel, bucket["dtype"])
                        else:
                            raw_buf = cpu_pool.acquire(bucket["total_bytes"])
                            try:
                                cpu_buf = raw_buf.view(bucket["dtype"])
                            except Exception:
                                cpu_pool.release(raw_buf)
                                raise
                        group_cpu_bufs[bucket_key] = cpu_buf
                        cpu_buf[:numel].copy_(group_device_bufs[bucket_key], non_blocking=True)
                except Exception:
                    release_event = platform.new_event()
                    release_event.record(copy_stream)
                    for bucket_key, cpu_buf in group_cpu_bufs.items():
                        bucket = self._packed_buckets[bucket_key]
                        if bucket["cpu_pool"] is not None:
                            bucket["cpu_pool"].release(cpu_buf, event=release_event)
                        else:
                            _return_cpu_pinned_buf(cpu_buf)
                    raise
                self._group_device_buf = group_device_bufs
                self._group_cpu_buf = group_cpu_bufs

            # Slice tensors use the existing per-tensor path.
711
712
713
714
715
716
717
718
719
720
721
722
723
724
    def release_cpu_buffers(self, event=None):
        """Release staging buffers immediately or defer until ``event`` completes."""
        if self._group_cpu_buf is None:
            return
        for bucket_key, buf in self._group_cpu_buf.items():
            bucket = self._packed_buckets.get(bucket_key)
            if bucket is not None and bucket["cpu_pool"] is not None:
                bucket["cpu_pool"].release(buf, event=event)
            else:
                _return_cpu_pinned_buf(buf)
        self._group_cpu_buf = None

    def wait_load(self):
        """Wait for load to complete for all storages in the group.
855
856
857
858
859
860
861
862
863
        for event in (group._offload_event, group._load_event):
            if event is not None:
                event.synchronize()
        for storage in group._storages:
            storage.release_cpu_buffers()
        group.release_cpu_buffers()
        group._storages.clear()

    def get_current_group_name(self) -> str:
hyper_parallel/platform/torch/platform.py
1625
1626
1627
1628
1629
1630
1631
1632
1633
    @staticmethod
    def swap_wrapper(module, policy_fn=None, group_swap=False, cpu_pool=None):
        # pylint: disable=C0415
        from hyper_parallel.platform.torch.activation_checkpoint.activation_swap import swap_wrapper
        return swap_wrapper(module, policy_fn=policy_fn, group_swap=group_swap, cpu_pool=cpu_pool)

    @staticmethod
    def swap_tensor_wrapper(target, tag=None, group_swap=False, cpu_pool=None):
        # pylint: disable=C0415
1631
1632
1633
1634
1635
1636
1637
1638
1639
    @staticmethod
    def swap_tensor_wrapper(target, tag=None, group_swap=False, cpu_pool=None):
        # pylint: disable=C0415
        from hyper_parallel.platform.torch.activation_checkpoint.activation_swap import swap_tensor_wrapper
        return swap_tensor_wrapper(target, tag=tag, group_swap=group_swap, cpu_pool=cpu_pool)

    @staticmethod
    def get_class_activation_wrapper():
        # pylint: disable=C0415
1655
1656
1657
1658
1659
1660
1661
1662
1663
        policy_fn_or_list, allow_cache_entry_mutation=False, group_swap=False, cpu_pool=None
    ):
        # pylint: disable=C0415
        from hyper_parallel.platform.torch.activation_checkpoint.sac import create_selective_checkpoint_contexts
        return create_selective_checkpoint_contexts(
            policy_fn_or_list, allow_cache_entry_mutation, group_swap, cpu_pool
        )

    @staticmethod
1672
1673
1674
1675
1676
1677
1678
1679
1680
    @staticmethod
    def async_save_on_cpu(policy_fn=None, group_swap: bool = False, cpu_pool=None):
        # pylint: disable=C0415
        from hyper_parallel.platform.torch.activation_checkpoint.activation_swap import AsyncSaveOnCpu
        return AsyncSaveOnCpu(policy_fn, group_swap=group_swap, cpu_pool=cpu_pool)

    @staticmethod
    def get_element_size(tensor):
        """Get Tensor Element Size"""