Engines

.. py:class:: ModuleCaptureEngine(model: ~torch.nn.modules.module.Module, modules_to_capture: list[~torch.nn.modules.module.Module], deepcopy_captured_args_and_kwargs: bool = False, output_callback: ~collections.abc.Callable[[int, ~typing.Any, ~typing.Any], None] | None = None, detach_outputs: bool = True, capture_inputs: bool = True, auto_truncate: bool = False)

module:

silverspoon_kd.engines

canonical:

silverspoon_kd.engines.module_capture_engine.ModuleCaptureEngine

Bases: :py:class:object

Engine for capturing inputs and outputs of neural network modules during forward passes.

This class provides a general-purpose mechanism to:

  1. Capture inputs to specified modules

  2. Capture outputs from specified modules

  3. Optionally truncate the forward pass after a terminal module

The captured data can be used for various purposes such as knowledge distillation, analysis, or debugging.

.. py:method:: ModuleCaptureEngine.init(model: ~torch.nn.modules.module.Module, modules_to_capture: list[~torch.nn.modules.module.Module], deepcopy_captured_args_and_kwargs: bool = False, output_callback: ~collections.abc.Callable[[int, ~typing.Any, ~typing.Any], None] | None = None, detach_outputs: bool = True, capture_inputs: bool = True, auto_truncate: bool = False)

module:

silverspoon_kd.engines

Initialize the ModuleCaptureEngine.

param model:

The model containing the modules to capture

param modules_to_capture:

List of modules whose inputs and outputs should be captured

param deepcopy_captured_args_and_kwargs:

Whether to deepcopy captured args/kwargs (slower but safer)

param output_callback:

Optional callback function called after each module’s forward pass. Signature: callback(module_id, input, output) -> None

param detach_outputs:

Whether to detach captured outputs from computation graph (default True). Set to False when outputs need to retain gradients for backpropagation.

param capture_inputs:

Whether to capture module inputs (args/kwargs). Set to False when only outputs are needed (e.g. HKD) to save memory.

param auto_truncate:

If True, automatically stop the forward pass as soon as all modules_to_capture have been captured. This is a performance optimization — layers after the last captured module are skipped entirely, saving both compute and memory. Defaults to False because exception-based truncation is incompatible with FSDP and torch.compile.

.. py:method:: ModuleCaptureEngine.register() -> None

module:

silverspoon_kd.engines

Register forward hooks and input capture wrappers for all modules.

This sets up the infrastructure to capture module inputs/outputs during forward passes.

.. py:method:: ModuleCaptureEngine.deregister() -> None

module:

silverspoon_kd.engines

Remove all forward hooks and restore original forward methods.

.. py:method:: ModuleCaptureEngine.clear_captured() -> None

module:

silverspoon_kd.engines

Clear all captured data.

.. py:method:: ModuleCaptureEngine.get_captured(module_id: int) -> tuple[~typing.Any, ~typing.Any, ~typing.Any]

module:

silverspoon_kd.engines

Get the captured data for a specific module.

param module_id:

The module identifier

returns:

Tuple of (args, kwargs, output)

raises KeyError:

If no data was captured for the given module_id

.. py:method:: ModuleCaptureEngine.pop_captured_inputs(module_id: int) -> tuple[~typing.Any, ~typing.Any]

module:

silverspoon_kd.engines

Pop (retrieve and remove) the captured inputs for a specific module.

param module_id:

The module identifier

returns:

Tuple of (args, kwargs)

raises KeyError:

If no inputs were captured for the given module_id

.. py:method:: ModuleCaptureEngine.pop_captured_output(module_id: int) -> ~typing.Any

module:

silverspoon_kd.engines

Pop (retrieve and remove) the captured output for a specific module.

param module_id:

The module identifier

returns:

The captured output

raises KeyError:

If no output was captured for the given module_id

.. py:function:: create_profiler(output_dir: str, wait: int = 20, warmup: int = 3, active: int = 3, repeat: int = 1, with_stack: bool = False) -> ~torch.profiler.profiler.profile

module:

silverspoon_kd.engines

Create a torch.profiler.profile instance with a configurable schedule.

Traces are exported as Chrome Trace JSON files via the on_trace_ready callback.

param output_dir:

Directory to save trace files.

param wait:

Number of steps to skip before profiling begins.

param warmup:

Number of warmup steps (not traced).

param active:

Number of steps to actively trace.

param repeat:

Number of times to repeat the wait/warmup/active cycle.

param with_stack:

Whether to record Python call stacks in traces.

returns:

A torch.profiler.profile instance (use as context manager).

Memory diagnostics

.. py:class:: MemoryLeakDetectionCallback(diff_between_batches: tuple[int, int] | None = None, snapshot_at_batch: int | None = None, snapshot_path: str | None = None, log_every_n_batches: int = 0, log_first_n_batches: int = 10, max_diff_buckets: int = 30, logger_: ~logging.Logger | None = None)

module:

silverspoon_kd.engines

canonical:

silverspoon_kd.engines.memory_diagnostics.MemoryLeakDetectionCallback

Bases: :py:class:~transformers.trainer_callback.TrainerCallback

Trainer callback that runs tensor-inventory diffs during eval.

Logs CUDA memory stats at every eval batch. Optionally:

  • Takes a tensor-inventory snapshot at diff_between_batches[0] and diffs it against a second snapshot at diff_between_batches[1], logging every shape+dtype bucket that appears new.

  • Records a CUDA memory history pickle covering a window around snapshot_at_batch, viewable at https://pytorch.org/memory_viz.

Both features are opt-in. With their defaults of None, the callback only logs per-batch memory stats — a lightweight sanity check that is cheap enough to enable by default in tests and CI.

param diff_between_batches:

Pair (first, second) of eval batch indices between which to diff live tensors. The inventories are taken at the START of each named batch (before that batch’s forward runs) so the diff reflects state that persisted between batches. Set to None to disable diffing.

param snapshot_at_batch:

Eval batch index at which to dump the CUDA memory history pickle. The recording starts 5 batches earlier so the batches leading up to the snapshot are covered. Set to None to disable snapshot dumping.

param snapshot_path:

Where to write the memory history pickle. Required when snapshot_at_batch is set.

param log_every_n_batches:

Log memory stats every N eval batches. Set to 0 to log every batch (default).

param log_first_n_batches:

Always log memory stats for the first N batches regardless of log_every_n_batches (default 10). Useful for catching warmup-phase allocations.

param max_diff_buckets:

Maximum diff buckets to log per diff. Use 0 for unlimited.

param logger_:

Logger to write to. Defaults to this module’s logger.

.. py:method:: MemoryLeakDetectionCallback.init(diff_between_batches: tuple[int, int] | None = None, snapshot_at_batch: int | None = None, snapshot_path: str | None = None, log_every_n_batches: int = 0, log_first_n_batches: int = 10, max_diff_buckets: int = 30, logger_: ~logging.Logger | None = None)

module:

silverspoon_kd.engines

.. py:method:: MemoryLeakDetectionCallback.on_prediction_step(args, state, control, **kwargs) -> None

module:

silverspoon_kd.engines

Per-eval-batch hook: log stats, snapshot, and diff tensors.

.. py:method:: MemoryLeakDetectionCallback.on_evaluate(args, state, control, **kwargs) -> None

module:

silverspoon_kd.engines

End-of-evaluate hook: reset per-eval-loop counters and inventories.

.. py:function:: format_memory_stats(device: ~torch.device | None = None) -> str

module:

silverspoon_kd.engines

Return a one-line summary of current CUDA memory state.

Format: alloc=X.XXXGiB reserved=X.XXXGiB peak=X.XXXGiB

param device:

CUDA device to query. Defaults to the current device.

.. py:function:: record_cuda_memory_history(snapshot_path: str, max_entries: int = 100000) -> ~collections.abc.Iterator[None]

module:

silverspoon_kd.engines

Context manager that records CUDA allocation history and dumps a pickle.

Uses the (private but long-stable) PyTorch memory history API: torch.cuda.memory._record_memory_history and _dump_snapshot. The resulting pickle can be opened at https://pytorch.org/memory_viz to view a timeline of every allocation with full backtraces.

param snapshot_path:

Where to write the pickle. The parent directory is created if missing.

param max_entries:

Maximum number of allocation events to record. The default is generous; reduce on memory-constrained systems.

Yields:

None. Allocations inside the with block are recorded.

Example:

.. code-block:: python

with record_cuda_memory_history("/tmp/leak.pickle"):
    for batch in eval_loader:
        distiller.prediction_step(...)

.. py:class:: TensorInfo(shape: tuple[int, …], dtype: str, device: str, nbytes: int)

module:

silverspoon_kd.engines

canonical:

silverspoon_kd.engines.memory_diagnostics.TensorInfo

Bases: :py:class:object

Summary of a single live CUDA tensor.

Only contains lightweight metadata — not a reference to the tensor itself — so that collecting an inventory cannot itself prevent garbage collection.

.. py:attribute:: TensorInfo.shape

module:

silverspoon_kd.engines

type:

tuple[int, …]

.. py:attribute:: TensorInfo.dtype

module:

silverspoon_kd.engines

type:

str

.. py:attribute:: TensorInfo.device

module:

silverspoon_kd.engines

type:

str

.. py:attribute:: TensorInfo.nbytes

module:

silverspoon_kd.engines

type:

int

.. py:method:: TensorInfo.init(shape: tuple[int, …], dtype: str, device: str, nbytes: int) -> None

module:

silverspoon_kd.engines

.. py:class:: TensorDiffBucket(shape: tuple[int, …], dtype: str, device: str, count: int, total_bytes: int)

module:

silverspoon_kd.engines

canonical:

silverspoon_kd.engines.memory_diagnostics.TensorDiffBucket

Bases: :py:class:object

A group of new tensors with the same shape/dtype/device.

.. py:attribute:: TensorDiffBucket.shape

module:

silverspoon_kd.engines

type:

tuple[int, …]

.. py:attribute:: TensorDiffBucket.dtype

module:

silverspoon_kd.engines

type:

str

.. py:attribute:: TensorDiffBucket.device

module:

silverspoon_kd.engines

type:

str

.. py:attribute:: TensorDiffBucket.count

module:

silverspoon_kd.engines

type:

int

.. py:attribute:: TensorDiffBucket.total_bytes

module:

silverspoon_kd.engines

type:

int

.. py:method:: TensorDiffBucket.init(shape: tuple[int, …], dtype: str, device: str, count: int, total_bytes: int) -> None

module:

silverspoon_kd.engines

.. py:function:: collect_tensor_inventory(include_cpu: bool = False, force_gc: bool = True) -> dict[int, ~silverspoon_kd.engines.memory_diagnostics.TensorInfo]

module:

silverspoon_kd.engines

Return a dict of id(tensor) -> TensorInfo for all live tensors.

Uses gc.get_objects() plus isinstance(obj, torch.Tensor) to walk every Python-reachable tensor. Tensors that are only kept alive via C++ references (e.g. inside custom autograd function state) will not appear.

param include_cpu:

If True, CPU tensors are included too. Defaults to False since CPU tensors are generally cheap and the goal is to find GPU memory leaks.

param force_gc:

If True, runs gc.collect() first so that unreachable cycles are cleared before inventory collection. This avoids false positives where an object is reachable only through a cycle that hasn’t been collected yet.

returns:

A dict keyed by id(tensor). Use id() as the key so inventories can be diffed by identity without depending on Python’s __hash__ / __eq__ semantics (which PyTorch tensors do not implement in the usual way).

.. py:function:: diff_tensor_inventories(before: dict[int, ~silverspoon_kd.engines.memory_diagnostics.TensorInfo], after: dict[int, ~silverspoon_kd.engines.memory_diagnostics.TensorInfo]) -> tuple[list[~silverspoon_kd.engines.memory_diagnostics.TensorDiffBucket], int]

module:

silverspoon_kd.engines

Return (buckets, total_bytes) describing NEW tensors in after.

Tensors are grouped by (shape, dtype, device) so that e.g. “28 new (1, 1024, 2048) bf16 tensors on cuda:0” appears as a single row rather than 28 individual entries. Buckets are sorted by total bytes descending so the largest leaks come first.

param before:

Inventory snapshot taken earlier (e.g. at eval batch 5).

param after:

Inventory snapshot taken later (e.g. at eval batch 15).

returns:

(buckets, total_bytes) where buckets is a list of

class:

TensorDiffBucket sorted by total_bytes desc, and total_bytes is the total bytes of all new tensors.

.. py:function:: log_tensor_diff(before: dict[int, ~silverspoon_kd.engines.memory_diagnostics.TensorInfo], after: dict[int, ~silverspoon_kd.engines.memory_diagnostics.TensorInfo], label: str = ‘diff’, max_buckets: int = 30, log: ~logging.Logger | None = None) -> tuple[list[~silverspoon_kd.engines.memory_diagnostics.TensorDiffBucket], int]

module:

silverspoon_kd.engines

Pretty-print the tensor diff and return the buckets/total.

param before:

Inventory snapshot taken earlier.

param after:

Inventory snapshot taken later.

param label:

String label used in log lines to identify this diff.

param max_buckets:

Maximum number of bucket rows to log (the largest are logged first). Use 0 for no limit.

param log:

Logger to write to. Defaults to this module’s logger.

returns:

Same as :func:diff_tensor_inventories.

.. py:class:: StorageInfo(nbytes: int, device: str, representative_shape: tuple[int, …], representative_dtype: str, num_views: int, view_shapes: tuple[tuple[int, …], …] = ())

module:

silverspoon_kd.engines

canonical:

silverspoon_kd.engines.memory_diagnostics.StorageInfo

Bases: :py:class:object

Summary of a unique CUDA tensor storage.

Different from :class:TensorInfo because it deduplicates by underlying storage data pointer, capturing the actual memory footprint regardless of how many tensor views refer to it. This catches the common case where a small “view” tensor (e.g. a slice of a layer output) keeps a much larger underlying buffer alive.

The view_shapes field records ALL distinct tensor shapes that share this storage. Inspecting this is invaluable when debugging leaks: if a 32 MiB storage is referenced only by tensors with shape () (scalars), something is keeping a tiny scalar view of a large buffer alive.

.. py:attribute:: StorageInfo.nbytes

module:

silverspoon_kd.engines

type:

int

.. py:attribute:: StorageInfo.device

module:

silverspoon_kd.engines

type:

str

.. py:attribute:: StorageInfo.representative_shape

module:

silverspoon_kd.engines

type:

tuple[int, …]

.. py:attribute:: StorageInfo.representative_dtype

module:

silverspoon_kd.engines

type:

str

.. py:attribute:: StorageInfo.num_views

module:

silverspoon_kd.engines

type:

int

.. py:attribute:: StorageInfo.view_shapes

module:

silverspoon_kd.engines

type:

tuple[tuple[int, …], …]

value:

()

.. py:method:: StorageInfo.init(nbytes: int, device: str, representative_shape: tuple[int, …], representative_dtype: str, num_views: int, view_shapes: tuple[tuple[int, …], …] = ()) -> None

module:

silverspoon_kd.engines

.. py:class:: StorageDiffBucket(nbytes: int, device: str, representative_shape: tuple[int, …], representative_dtype: str, count: int, total_bytes: int, example_view_shapes: tuple[tuple[int, …], …] = ())

module:

silverspoon_kd.engines

canonical:

silverspoon_kd.engines.memory_diagnostics.StorageDiffBucket

Bases: :py:class:object

A group of new storages with the same size/device/dtype.

.. py:attribute:: StorageDiffBucket.nbytes

module:

silverspoon_kd.engines

type:

int

.. py:attribute:: StorageDiffBucket.device

module:

silverspoon_kd.engines

type:

str

.. py:attribute:: StorageDiffBucket.representative_shape

module:

silverspoon_kd.engines

type:

tuple[int, …]

.. py:attribute:: StorageDiffBucket.representative_dtype

module:

silverspoon_kd.engines

type:

str

.. py:attribute:: StorageDiffBucket.count

module:

silverspoon_kd.engines

type:

int

.. py:attribute:: StorageDiffBucket.total_bytes

module:

silverspoon_kd.engines

type:

int

.. py:attribute:: StorageDiffBucket.example_view_shapes

module:

silverspoon_kd.engines

type:

tuple[tuple[int, …], …]

value:

()

.. py:method:: StorageDiffBucket.init(nbytes: int, device: str, representative_shape: tuple[int, …], representative_dtype: str, count: int, total_bytes: int, example_view_shapes: tuple[tuple[int, …], …] = ()) -> None

module:

silverspoon_kd.engines

.. py:function:: collect_storage_inventory(include_cpu: bool = False, force_gc: bool = True) -> dict[int, ~silverspoon_kd.engines.memory_diagnostics.StorageInfo]

module:

silverspoon_kd.engines

Return a dict of storage_data_ptr -> StorageInfo for live storages.

This is the memory-faithful counterpart to

func:

collect_tensor_inventory: it deduplicates tensors that share an underlying storage and reports the storage’s true nbytes (which can be much larger than any single view’s element count × element size).

This function is the right tool for finding leaks where the visible Python tensors are small but the underlying CUDA memory is large — e.g. a 1 MB view of a 1 GB buffer that nothing else has dropped.

param include_cpu:

If True, CPU storages are included.

param force_gc:

If True, runs gc.collect() first.

returns:

Dict keyed by storage’s data_ptr() (an int). Each value is a

class:

StorageInfo describing the storage’s true memory footprint and one representative tensor view.

.. py:function:: diff_storage_inventories(before: dict[int, ~silverspoon_kd.engines.memory_diagnostics.StorageInfo], after: dict[int, ~silverspoon_kd.engines.memory_diagnostics.StorageInfo]) -> tuple[list[~silverspoon_kd.engines.memory_diagnostics.StorageDiffBucket], int]

module:

silverspoon_kd.engines

Return (buckets, total_bytes) describing NEW storages in after.

Storages are grouped by (nbytes, device, representative_dtype) so that e.g. “28 new 32 MiB bf16 storages on cuda:0” appears as a single row. Buckets are sorted by total_bytes descending.

Note: storages that disappear are intentionally not reported here. The goal is to find tensors that are leaking, not ones that were freed.

.. py:function:: log_storage_diff(before: dict[int, ~silverspoon_kd.engines.memory_diagnostics.StorageInfo], after: dict[int, ~silverspoon_kd.engines.memory_diagnostics.StorageInfo], label: str = ‘storage_diff’, max_buckets: int = 30, log: ~logging.Logger | None = None) -> tuple[list[~silverspoon_kd.engines.memory_diagnostics.StorageDiffBucket], int]

module:

silverspoon_kd.engines

Pretty-print the storage diff and return the buckets/total.

.. py:function:: find_tensors_for_data_ptrs(target_ptrs: set, include_cpu: bool = False) -> list[~torch.Tensor]

module:

silverspoon_kd.engines

Return all live torch.Tensor objects whose storage data_ptr is in target_ptrs.

Useful for inspecting WHICH Python tensor objects are keeping a leaked storage alive — combine with gc.get_referrers() on the result to walk back to the Python container holding the reference.

param target_ptrs:

Set of integer storage.data_ptr() values to find.

param include_cpu:

If True, CPU tensors are also returned.

returns:

A list of live tensors. The list itself is a fresh allocation; the tensors retain whatever existing references they have.

.. py:function:: describe_referrers(obj: object, max_depth: int = 2, max_referrers: int = 10) -> list[str]

module:

silverspoon_kd.engines

Return human-readable descriptions of objects referring to obj.

Walks gc.get_referrers(obj) to find Python containers holding a reference to obj. The result is a list of short descriptive strings (truncated to max_referrers per level) that include the type, repr (truncated), and any helpful identifying info.

param obj:

The Python object to find referrers of.

param max_depth:

How many levels of indirection to walk. 1 lists direct referrers; 2 also lists referrers-of-referrers; etc.

param max_referrers:

Maximum referrers to describe per level.

returns:

A flat list of description strings, one per referrer at any level within max_depth.