Diff Coverage

Diff: origin/r1.0.0...HEAD, staged and unstaged changes

Source File Diff Coverage (%) Missing Lines
hyper_parallel/platform/mindspore/activation_checkpoint/checkpoint_exclude_wrapper.py 89.3% 151,158,165,179-180,186,190,193,249,251,270,370,390,398,407,412,420,436,472,476-485
hyper_parallel/platform/mindspore/activation_checkpoint/checkpoint_exclude_wrapper.py
147
148
149
150
151
152
153
154
155

def _canonical_input_tensor(tensor: Any) -> Any:
    """Return the local tensor carried by a DTensor wrapper."""
    if isinstance(tensor, DTensor):
        return tensor.to_local()
    return tensor


def _tensor_metadata(tensor: Any) -> Any:
154
155
156
157
158
159
160
161
162

def _tensor_metadata(tensor: Any) -> Any:
    """Return storage/layout metadata, or ``None`` when aliasing is unavailable."""
    if not isinstance(tensor, ms.Tensor) or tensor.numel() == 0:
        return None
    try:
        storage = tensor.untyped_storage()
        storage_ptr = storage.data_ptr()
        storage_nbytes = storage.size()
161
162
163
164
165
166
167
168
169
        storage_ptr = storage.data_ptr()
        storage_nbytes = storage.size()
        itemsize = tensor.itemsize
        if storage_ptr == 0 or storage_nbytes == 0 or itemsize <= 0:
            return None
        return _TensorMetadata(
            tensor_id=id(tensor),
            storage_ptr=storage_ptr,
            storage_nbytes=storage_nbytes,
175
176
177
178
179
180
181
182
183
184
            itemsize=itemsize,
            numel=tensor.numel(),
            is_contiguous=tensor.is_contiguous(),
        )
    except (AttributeError, RuntimeError, TypeError, ValueError):
        return None


def _storage_span(metadata: _TensorMetadata) -> Any:
    """Return inclusive element offsets touched by a non-negative-stride tensor."""
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197

def _storage_span(metadata: _TensorMetadata) -> Any:
    """Return inclusive element offsets touched by a non-negative-stride tensor."""
    if metadata.numel == 0 or any(step < 0 for step in metadata.stride):
        return None
    max_offset = metadata.storage_offset
    for size, step in zip(metadata.shape, metadata.stride):
        if size == 0:
            return None
        max_offset += (size - 1) * step
    if metadata.storage_offset < 0 or (max_offset + 1) * metadata.itemsize > metadata.storage_nbytes:
        return None
    return metadata.storage_offset, max_offset


def _exact_layout(saved: _TensorMetadata, base: _TensorMetadata) -> bool:
245
246
247
248
249
250
251
252
253
254
255
            return None
        saved_span = _storage_span(saved)
        base_span = _storage_span(base)
        if saved_span is None or base_span is None:
            return None
        if saved_span[0] < base_span[0] or saved_span[1] > base_span[1]:
            return None

    return _ViewRecipe(
        shape=saved.shape,
        stride=saved.stride,
266
267
268
269
270
271
272
273
274
def _match_recompute_input(tensor: Any, inputs: Tuple[_InputInfo, ...]) -> Any:
    """Match a saved tensor to an excluded-call input by storage and layout."""
    saved = _tensor_metadata(tensor)
    if saved is None:
        return None

    exact_identity = []
    exact_layout = []
    views = []
366
367
368
369
370
371
372
373
        if id(canonical) in seen_tensor_ids:
            continue
        metadata = _tensor_metadata(canonical)
        if metadata is None:
            continue
        inputs.append(_InputInfo(path, metadata))
        seen_tensor_ids.add(id(canonical))
    return tuple(inputs)
386
387
388
389
390
391
392
393
394
    """Rebuild a detached saved input/view from one replay-produced base."""
    canonical = _canonical_input_tensor(tensor)
    metadata = _tensor_metadata(canonical)
    if metadata is None:
        raise RuntimeError("Checkpoint replay input does not expose usable storage metadata")
    if (
        metadata.dtype != recipe.dtype
        or metadata.shape != recipe.base_shape
        or metadata.stride != recipe.base_stride
394
395
396
397
398
399
400
401
402
        or metadata.stride != recipe.base_stride
        or metadata.version != recipe.base_version
        or metadata.is_contiguous != recipe.base_is_contiguous
    ):
        raise RuntimeError(
            "Checkpoint replay input layout/version changed for a checkpoint-excluded saved tensor"
        )

    if recipe.exact_input:
403
404
405
406
407
408
409
410
411
412
413
414
415
416
        return canonical.detach()

    view_offset = metadata.storage_offset + recipe.relative_offset
    if any(step < 0 for step in recipe.stride):
        raise RuntimeError("Checkpoint-excluded saved tensor has unsupported negative stride")
    max_offset = view_offset
    numel = 1
    for size, step in zip(recipe.shape, recipe.stride):
        if size == 0:
            raise RuntimeError("Checkpoint-excluded empty saved views should have been saved normally")
        numel *= size
        max_offset += (size - 1) * step
    if (
        numel == 0
416
417
418
419
420
421
422
423
424
        numel == 0
        or view_offset < 0
        or (max_offset + 1) * metadata.itemsize > metadata.storage_nbytes
    ):
        raise RuntimeError("Checkpoint-excluded saved view is outside the replay input storage")

    restored = canonical.new_empty((0,))
    restored.set_(canonical.untyped_storage(), view_offset, recipe.shape, recipe.stride)
    return restored
432
433
434
435
436
437
438
439
    """Bind saved-alias handles to tensors rebuilt from checkpoint replay inputs."""
    for binding in entry.input_bindings:
        tensor = _resolve_input(args, kwargs, binding.path)
        if not isinstance(tensor, ms.Tensor):
            raise RuntimeError(
                "Checkpoint replay did not reproduce a tensor input required by a checkpoint-excluded region"
            )
        binding.handle.materialize(_rebuild_saved_alias(tensor, binding.handle.recipe))
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489

def _apply_recompute_boundary(output: Any, input_bindings: List[_InputBinding]) -> Any:
    """Wrap output tensor leaves when deferred inputs require checkpoint replay."""
    if not _has_used_input(input_bindings):
        return output

    if isinstance(output, ms.Tensor):
        return _RecomputeBoundary.apply(output)
    if isinstance(output, list):
        return [_apply_recompute_boundary(item, input_bindings) for item in output]
    if isinstance(output, tuple):
        items = [_apply_recompute_boundary(item, input_bindings) for item in output]
        if hasattr(output, "_fields"):
            return type(output)(*items)
        return tuple(items)
    if isinstance(output, dict):
        return type(output)((key, _apply_recompute_boundary(value, input_bindings)) for key, value in output.items())
    return output


class CheckpointExcludeWrapper(ActivationWrapper):
    """Exclude a callable region from checkpoint recomputation."""