AxoSim

API reference

CUDA Graph adaptation

Signatures, parameters, return contracts, and source for cuda graph adaptation.

Module contract

CudaGraphAdaptationStep captures one fixed-shape forward, loss, backward, optional gradient transform, and optimizer update. It restores module and optimizer state after warmup and capture. Replay copies source tensors into retained static buffers. CUDA inputs and optimizer parameters are required; shape, dtype, parameter groups, and allocations must remain compatible.

Source revision: 306a51ed950b. Public export index.

CudaGraphAdaptationStep

Source

A captured full-forward, backward, and optimizer update.

CudaGraphAdaptationStep(module: nn.Module, optimizer: torch.optim.Optimizer, static_inputs: torch.Tensor, static_targets: torch.Tensor, graph: torch.cuda.CUDAGraph, static_loss: torch.Tensor) -> None

Fields

Parameter Type Default
module nn.Module required
optimizer torch.optim.Optimizer required
static_inputs torch.Tensor required
static_targets torch.Tensor required
graph torch.cuda.CUDAGraph required
static_loss torch.Tensor required

optimizer: Optimizer attached to the selected parameters.

CudaGraphAdaptationStep.capture

Source

@classmethod
capture(cls, module: nn.Module, optimizer: torch.optim.Optimizer, example_inputs: torch.Tensor, example_targets: torch.Tensor, loss_function: TensorLoss, *, warmup_steps: int=3, gradient_transform: GradientTransform | None=None) -> 'CudaGraphAdaptationStep'

Capture a fixed-shape adaptation update without consuming it.

module and optimizer parameters must be CUDA-resident; example_inputs and example_targets establish fixed shape/dtype. loss_function(prediction,targets) must return a scalar differentiable tensor. gradient_transform is a zero-argument callable executed after backward. Returns a captured step while restoring pre-capture module/optimizer values.

Parameter Type Default
module nn.Module required
optimizer torch.optim.Optimizer required
example_inputs torch.Tensor required
example_targets torch.Tensor required
loss_function TensorLoss required
warmup_steps int 3
gradient_transform GradientTransform | None None

optimizer: Optimizer attached to the selected parameters. warmup_steps: Warmup updates executed and restored during capture. gradient_transform: Optional zero-argument gradient transform after backward.

Returns 'CudaGraphAdaptationStep'.

CudaGraphAdaptationStep.call

Source

__call__(self, inputs: torch.Tensor, targets: torch.Tensor, *, synchronize: bool=False) -> torch.Tensor

inputs and targets must match captured shapes and dtypes. Copies values to retained static CUDA buffers, replays the complete update, and returns a detached cloned scalar loss. synchronize=True waits for CUDA completion.

Parameter Type Default
inputs torch.Tensor required
targets torch.Tensor required
synchronize bool False

synchronize: Wait for CUDA completion before returning.

Returns torch.Tensor.

Search guides, examples, and API signatures.