Source code for leap_c.autograd.torch
"""This module creates PyTorch autograd functions."""
from typing import TYPE_CHECKING
from leap_c.autograd.function import DiffFunction
from leap_c.utils.dependencies import require_torch
if TYPE_CHECKING:
import torch
else:
torch = require_torch()
[docs]
def create_autograd_function(fun: DiffFunction) -> type[torch.autograd.Function]:
"""Creates a PyTorch autograd function from an object implementing forward and backward methods.
The `fun` object must implement the `forward` and `backward` methods as described in the
`Function` class. During the backward pass, the custom context object gets the information about
which inputs need gradients via `torch_ctx.needs_input_grad`.
Args:
fun: An object implementing `forward(custom_ctx, *args)` and
`backward(custom_ctx, *grad_outputs)` methods, where `custom_ctx` is a context object
that can be used to store intermediate values for the backward pass.
Returns:
A PyTorch autograd function, wrapping the object.
Examples:
>>> fn = create_autograd_function(obj)
>>> ctx, y = fn(*inputs)
"""
class AutogradFunction(torch.autograd.Function):
@staticmethod
def forward(torch_ctx, *args):
custom_ctx, *non_ctx_args = args
device = non_ctx_args[0].device
np_args = _to_np(non_ctx_args)
custom_ctx, *outputs = fun.forward(custom_ctx, *np_args) # type: ignore
torch_ctx.custom_ctx = custom_ctx
if len(outputs) == 1:
return custom_ctx, _to_tensor(outputs[0], device)
return custom_ctx, *_to_tensor(outputs, device)
@staticmethod
def backward(torch_ctx, grad_ctx, *grad_outputs): # type: ignore
device = grad_outputs[0].device
custom_ctx = torch_ctx.custom_ctx
custom_ctx.needs_input_grad = torch_ctx.needs_input_grad
np_grad_outputs = _to_np(grad_outputs)
grad_inputs = fun.backward(custom_ctx, *np_grad_outputs) # type: ignore
torch_grad_inputs = _to_tensor(grad_inputs, device)
if isinstance(torch_grad_inputs, torch.Tensor):
return None, torch_grad_inputs
return None, *torch_grad_inputs
return AutogradFunction
def _to_np(data):
if data is None:
return None
if isinstance(data, (tuple, list)):
return tuple(_to_np(item) for item in data)
try:
return data.detach().cpu().numpy()
except AttributeError:
return data
def _to_tensor(data, device):
if data is None:
return None
if isinstance(data, (tuple, list)):
return tuple(_to_tensor(item, device) for item in data)
return torch.as_tensor(data, device=device)