PyTorch Interoperability#

Introduction#

Warp provides helper functions to convert arrays to/from PyTorch:

w = wp.array([1.0, 2.0, 3.0], dtype=float, device="cpu")

# convert to Torch tensor
t = wp.to_torch(w)

# convert from Torch tensor
w = wp.from_torch(t)

These helper functions allow the conversion of Warp arrays to/from PyTorch tensors without copying the underlying data. At the same time, if available, gradient arrays and tensors are converted to/from PyTorch autograd tensors, allowing the use of Warp arrays in PyTorch autograd computations.

Warp also provides helper functions for mapping devices and dtypes between Warp and PyTorch:

import torch

torch_device = wp.device_to_torch("cpu")
torch_dtype = wp.dtype_to_torch(wp.float32)
t = torch.ones(3, device=torch_device, dtype=torch_dtype)

warp_dtype = wp.dtype_from_torch(t.dtype)
warp_device = wp.device_from_torch(t.device)

By default, Warp arrays allocated with functions such as wp.zeros() use Warp’s CUDA allocator. If an application needs those allocations to come from PyTorch’s CUDA caching allocator, see PyTorch CUDA Caching Allocator for a minimal custom allocator example.

Stream Conversion#

To convert a PyTorch CUDA stream to a Warp CUDA stream and vice versa, Warp provides the following functions:

PyTorch’s CUDA default stream is blocking, while non-default streams created with torch.cuda.Stream() are non-blocking. Streams created by Warp are blocking. Converting a stream between frameworks preserves its blocking behavior. For a detailed explanation and its implications, see Blocking and Non-Blocking Streams.

CUDA Graph Capture with PyTorch and Warp#

It is possible to capture CUDA graphs that include both PyTorch and Warp operations, as long as they run on the same CUDA stream. By default, PyTorch uses the synchronous CUDA default stream, which is not suitable for graph capture. A new stream must be created prior to capture, as shown in the examples below.

Capturing a graph using a PyTorch stream:

import torch
import warp as wp

@wp.kernel
def scale(a: wp.array[float], s: float):
    tid = wp.tid()
    a[tid] = a[tid] * s

n = 1024 * 1024
torch_device = wp.device_to_torch("cuda:0")

# Create a non-default PyTorch stream and convert it to Warp
torch_stream = torch.cuda.Stream(device=torch_device)
warp_stream = wp.stream_from_torch(torch_stream)

a = wp.ones(n, dtype=float, device="cuda:0")

# Capture a graph using the shared stream
with wp.ScopedStream(warp_stream):
    with wp.ScopedCapture() as capture:
        wp.launch(scale, dim=n, inputs=[a, 2.0])

# Replay the graph
wp.capture_launch(capture.graph, stream=warp_stream)

Capturing a graph using a Warp stream:

import torch
import warp as wp

@wp.kernel
def scale(a: wp.array[float], s: float):
    tid = wp.tid()
    a[tid] = a[tid] * s

n = 1024 * 1024

a = wp.ones(n, dtype=float, device="cuda:0")

# Make PyTorch use the Warp stream
torch_stream = wp.stream_to_torch("cuda:0")

# Capture a graph using the Warp stream
with wp.ScopedDevice("cuda:0"), torch.cuda.stream(torch_stream):
    with wp.ScopedCapture() as capture:
        wp.launch(scale, dim=n, inputs=[a, 2.0])

# Replay the graph
wp.capture_launch(capture.graph)

It can be tricky to capture arbitrary PyTorch code in CUDA graphs, because many PyTorch operations involve code that is not capturable. Some warmup steps may be required. For more information about PyTorch and CUDA graphs, see the PyTorch blog post on CUDA graphs.

Optimization Examples#

Using warp.from_torch#

An example usage of minimizing a loss function over an array of 2D points written in Warp via PyTorch’s Adam optimizer using warp.from_torch() is as follows:

import warp as wp
import torch


@wp.kernel()
def loss(xs: wp.array2d[float], l: wp.array[float]):
    tid = wp.tid()
    wp.atomic_add(l, 0, xs[tid, 0] ** 2.0 + xs[tid, 1] ** 2.0)

# indicate requires_grad so that Warp can accumulate gradients in the grad buffers
xs = torch.randn(100, 2, requires_grad=True)
l = torch.zeros(1, requires_grad=True)
opt = torch.optim.Adam([xs], lr=0.1)

wp_xs = wp.from_torch(xs)
wp_l = wp.from_torch(l)

tape = wp.Tape()
with tape:
    # record the loss function kernel launch on the tape
    wp.launch(loss, dim=len(xs), inputs=[wp_xs], outputs=[wp_l], device=wp_xs.device)

for i in range(500):
    tape.zero()
    tape.backward(loss=wp_l)  # compute gradients
    # now xs.grad will be populated with the gradients computed by Warp
    opt.step()  # update xs (and thereby wp_xs)

    # these lines are only needed for evaluating the loss
    # (the optimization just needs the gradient, not the loss value)
    wp_l.zero_()
    wp.launch(loss, dim=len(xs), inputs=[wp_xs], outputs=[wp_l], device=wp_xs.device)
    print(f"{i}\tloss: {l.item()}")

Using warp.to_torch#

Less code is needed when we declare the optimization variables directly in Warp and use warp.to_torch() to convert them to PyTorch tensors. Here, we revisit the same example from above where now only a single conversion to a PyTorch tensor is needed to supply Adam with the optimization variables:

import warp as wp
import numpy as np
import torch


@wp.kernel()
def loss(xs: wp.array2d[float], l: wp.array[float]):
    tid = wp.tid()
    wp.atomic_add(l, 0, xs[tid, 0] ** 2.0 + xs[tid, 1] ** 2.0)

# initialize the optimization variables in Warp
xs = wp.array(np.random.randn(100, 2), dtype=wp.float32, requires_grad=True)
l = wp.zeros(1, dtype=wp.float32, requires_grad=True)
# just a single wp.to_torch call is needed, Adam optimizes using the Warp array gradients
opt = torch.optim.Adam([wp.to_torch(xs)], lr=0.1)

tape = wp.Tape()
with tape:
    wp.launch(loss, dim=len(xs), inputs=[xs], outputs=[l], device=xs.device)

for i in range(500):
    tape.zero()
    tape.backward(loss=l)
    opt.step()

    l.zero_()
    wp.launch(loss, dim=len(xs), inputs=[xs], outputs=[l], device=xs.device)
    print(f"{i}\tloss: {l.numpy()[0]}")

Autograd Integration#

Using torch.autograd.Function (PyTorch <= 2.3.1)#

One can insert Warp kernel launches in a PyTorch graph by defining a torch.autograd.Function class, which requires forward and backward functions to be defined. After mapping incoming PyTorch tensors to Warp arrays, a Warp kernel may be launched in the usual way. In the backward pass, the same kernel’s adjoint may be launched by setting adjoint = True in wp.launch(). Alternatively, the user may choose to rely on Warp’s tape. In the following example, we demonstrate how Warp may be used to evaluate the Rosenbrock function in an optimization context.

Caution

When writing the backward function, remember that wp.from_torch() and wp.to_torch() are zero-copy conversions. Treat incoming PyTorch grad_output tensors as externally owned buffers, because PyTorch may reuse the same tensor across backward calls. Do not assign a wp.from_torch(grad_output) array directly to an output array’s .grad attribute.

The table below summarizes the ownership rules for common gradient-buffer patterns:

Gradient Buffer Ownership#

Pattern

Buffer used by Warp

Safe to reuse or retain?

Guidance

output.grad = wp.from_torch(grad_output)

The external PyTorch buffer

No

Avoid. Warp may consume or zero storage that PyTorch expects to reuse.

tape.backward(grads={output: external_grad}) when output.grad is None

external_grad itself; Tape adopts it as output.grad

No

Allocate an independent gradient buffer for output first.

tape.backward(grads={output: external_grad}) when output already owns .grad

The owned Warp buffer; external_grad is copied into it

Yes

Recommended for external PyTorch gradients.

wp.to_torch(input.grad)

A zero-copy view of the Warp gradient buffer

Only until that buffer is modified

Call .clone() before tape.zero() if PyTorch must retain the gradient.

For a manual adjoint launch, pass the incoming gradient as an explicit adj_outputs buffer, as shown below. When using warp.Tape, pass it through Tape.backward() using grads={...} and ensure that the output Warp array already owns an independent gradient buffer. If the backward pass depends on a PyTorch input’s forward value, save the original tensor with ctx.save_for_backward() even if Warp wraps a detached view. Accessing ctx.saved_tensors in backward() lets PyTorch detect in-place mutations before Warp reads the shared storage. On CUDA, run Warp work on the active PyTorch stream using wp.stream_from_torch() and wp.ScopedStream so operations on shared zero-copy buffers remain ordered.

Because these conversions are zero-copy, pass PyTorch tensors directly when their layout is compatible with the Warp kernel. Scalar Warp arrays preserve PyTorch strides, so a non-contiguous tensor can often be wrapped without copying. Use tensor.contiguous() deliberately only when a Warp dtype or kernel requires that layout, such as when wrapping vector or matrix dtypes whose trailing component dimensions must be contiguous, or when choosing a contiguous copy for performance.

import warp as wp
import numpy as np
import torch

def active_torch_stream(tensor):
    if tensor.is_cuda:
        return wp.ScopedStream(wp.stream_from_torch(tensor.device))

    return wp.ScopedStream(None)

# Define the Rosenbrock function
@wp.func
def rosenbrock(x: float, y: float):
    return (1.0 - x) ** 2.0 + 100.0 * (y - x**2.0) ** 2.0

@wp.kernel
def eval_rosenbrock(
    xs: wp.array[wp.vec2],
    # outputs
    z: wp.array[float],
):
    i = wp.tid()
    x = xs[i]
    z[i] = rosenbrock(x[0], x[1])


class Rosenbrock(torch.autograd.Function):
    @staticmethod
    def forward(ctx, xy, num_points):
        # Save the PyTorch input so autograd can detect in-place mutations before
        # backward. Warp wraps a detached view because gradients are returned manually.
        ctx.save_for_backward(xy)
        ctx.num_points = num_points

        with active_torch_stream(xy):
            ctx.xy = wp.from_torch(xy.detach(), dtype=wp.vec2, requires_grad=False)

            # allocate output
            ctx.z = wp.zeros(num_points, dtype=wp.float32, device=ctx.xy.device, requires_grad=False)

            wp.launch(
                kernel=eval_rosenbrock,
                dim=ctx.num_points,
                inputs=[ctx.xy],
                outputs=[ctx.z],
                device=ctx.xy.device
            )

        return wp.to_torch(ctx.z)

    @staticmethod
    def backward(ctx, adj_z):
        # Accessing saved_tensors performs PyTorch's version check for the input.
        (xy,) = ctx.saved_tensors

        # Allocate the input adjoint that this backward call will return
        # to PyTorch. Keep adj_z as an external adjoint buffer.
        adj_xy = torch.zeros_like(xy)

        with active_torch_stream(adj_z):
            wp_adj_xy = wp.from_torch(adj_xy, dtype=wp.vec2, requires_grad=False)
            wp_adj_z = wp.from_torch(adj_z, requires_grad=False)

            wp.launch(
                kernel=eval_rosenbrock,
                dim=ctx.num_points,
                inputs=[ctx.xy],
                outputs=[ctx.z],
                adj_inputs=[wp_adj_xy],
                adj_outputs=[wp_adj_z],
                device=ctx.xy.device,
                adjoint=True
            )

        # return adjoint w.r.t. inputs
        return (adj_xy, None)


num_points = 1500
learning_rate = 5e-2

torch_device = wp.device_to_torch(wp.get_device())

rng = np.random.default_rng(42)
xy = torch.tensor(rng.normal(size=(num_points, 2)), dtype=torch.float32, requires_grad=True, device=torch_device)
opt = torch.optim.Adam([xy], lr=learning_rate)

for _ in range(10000):
    # step
    opt.zero_grad()
    z = Rosenbrock.apply(xy, num_points)
    z.backward(torch.ones_like(z))

    opt.step()

# minimum at (1, 1)
xy_np = xy.numpy(force=True)
print(np.mean(xy_np, axis=0))

If replacing the manual adjoint launch with warp.Tape, keep the same ownership model: treat the PyTorch grad_output tensor as external storage and pre-allocate independent Warp adjoint buffers for arrays whose gradients must survive after Tape.backward().

Note that if Warp code is wrapped in a torch.autograd.Function that gets called in torch.compile(), it will automatically exclude that function from compiler optimizations. If your script uses torch.compile(), we recommend using PyTorch version 2.3.0+, which has improvements that address this scenario.

Using PyTorch Custom Operators (PyTorch >= 2.4.0)#

PyTorch 2.4+ introduced custom operators to replace PyTorch autograd functions. These treat arbitrary Python functions (including Warp calls) as opaque callables, which prevents torch.compile() from tracing into them. This means that forward PyTorch graph evaluations that include Warp kernel launches can be safely accelerated with torch.compile(). We can re-write the previous example using custom operators as follows:

import warp as wp
import numpy as np
import torch

# Define the Rosenbrock function
@wp.func
def rosenbrock(x: float, y: float):
    return (1.0 - x) ** 2.0 + 100.0 * (y - x**2.0) ** 2.0


@wp.kernel
def eval_rosenbrock(
    xy: wp.array[wp.vec2],
    # outputs
    z: wp.array[float],
):
    i = wp.tid()
    v = xy[i]
    z[i] = rosenbrock(v[0], v[1])


@torch.library.custom_op("wp::warp_rosenbrock", mutates_args=())
def warp_rosenbrock(xy: torch.Tensor, num_points: int) -> torch.Tensor:
    wp_xy = wp.from_torch(xy, dtype=wp.vec2, requires_grad=False)
    wp_z = wp.zeros(num_points, dtype=wp.float32, device=wp_xy.device, requires_grad=False)

    wp.launch(kernel=eval_rosenbrock, dim=num_points, inputs=[wp_xy], outputs=[wp_z], device=wp_xy.device)

    return wp.to_torch(wp_z)


@warp_rosenbrock.register_fake
def _(xy, num_points):
    return torch.empty(num_points, dtype=torch.float32)


@torch.library.custom_op("wp::warp_rosenbrock_backward", mutates_args=())
def warp_rosenbrock_backward(
    xy: torch.Tensor, num_points: int, z: torch.Tensor, adj_z: torch.Tensor
) -> torch.Tensor:
    wp_xy = wp.from_torch(xy, dtype=wp.vec2, requires_grad=False)
    wp_z = wp.from_torch(z, requires_grad=False)
    adj_xy = torch.zeros_like(xy)

    wp_adj_xy = wp.from_torch(adj_xy, dtype=wp.vec2, requires_grad=False)
    wp_adj_z = wp.from_torch(adj_z, requires_grad=False)

    wp.launch(
        kernel=eval_rosenbrock,
        dim=num_points,
        inputs=[wp_xy],
        outputs=[wp_z],
        adj_inputs=[wp_adj_xy],
        adj_outputs=[wp_adj_z],
        device=wp_xy.device,
        adjoint=True,
    )

    return adj_xy


@warp_rosenbrock_backward.register_fake
def _(xy, num_points, z, adj_z):
    return torch.empty_like(xy)


def backward(ctx, adj_z):
    return warp_rosenbrock_backward(ctx.xy, ctx.num_points, ctx.z, adj_z), None


def setup_context(ctx, inputs, output):
    ctx.xy, ctx.num_points = inputs
    ctx.z = output


warp_rosenbrock.register_autograd(backward, setup_context=setup_context)

num_points = 1500
learning_rate = 5e-2

torch_device = wp.device_to_torch(wp.get_device())

rng = np.random.default_rng(42)
xy = torch.tensor(rng.normal(size=(num_points, 2)), dtype=torch.float32, requires_grad=True, device=torch_device)
opt = torch.optim.Adam([xy], lr=learning_rate)

@torch.compile(fullgraph=True)
def forward():
    global xy, num_points

    z = warp_rosenbrock(xy, num_points)
    return z

for _ in range(10000):
    # step
    opt.zero_grad()
    z = forward()
    z.backward(torch.ones_like(z))
    opt.step()

# minimum at (1, 1)
xy_np = xy.numpy(force=True)
print(np.mean(xy_np, axis=0))

Performance Tuning#

The wp.from_torch() function creates a Warp array object that shares data with a PyTorch tensor. Although this function does not copy the data, there is always some CPU overhead during the conversion. If these conversions happen frequently, the overall program performance may suffer. As a general rule, repeated conversions of the same tensor should be avoided. Instead of:

x_t = torch.arange(n, dtype=torch.float32, device=device)
y_t = torch.ones(n, dtype=torch.float32, device=device)

for i in range(10):
    x_w = wp.from_torch(x_t)
    y_w = wp.from_torch(y_t)
    wp.launch(saxpy, dim=n, inputs=[x_w, y_w, 1.0], device=device)

Try converting the arrays only once and reuse them:

x_t = torch.arange(n, dtype=torch.float32, device=device)
y_t = torch.ones(n, dtype=torch.float32, device=device)

x_w = wp.from_torch(x_t)
y_w = wp.from_torch(y_t)

for i in range(10):
    wp.launch(saxpy, dim=n, inputs=[x_w, y_w, 1.0], device=device)

If reusing arrays is not possible (e.g., a new PyTorch tensor is constructed on every iteration), passing return_ctype=True to wp.from_torch() should yield better performance. Setting this argument to True avoids constructing a wp.array object and instead returns a low-level array descriptor. This descriptor is a simple C structure that can be passed to Warp kernels instead of a wp.array, but cannot be used in other places that require a wp.array.

for n in range(1, 10):
    # get Torch tensors for this iteration
    x_t = torch.arange(n, dtype=torch.float32, device=device)
    y_t = torch.ones(n, dtype=torch.float32, device=device)

    # get Warp array descriptors
    x_ctype = wp.from_torch(x_t, return_ctype=True)
    y_ctype = wp.from_torch(y_t, return_ctype=True)

    wp.launch(saxpy, dim=n, inputs=[x_ctype, y_ctype, 1.0], device=device)

An alternative approach is to pass the PyTorch tensors to Warp kernels directly. This avoids constructing temporary Warp arrays by leveraging standard array interfaces (like __cuda_array_interface__) supported by both PyTorch and Warp. The main advantage of this approach is convenience, since there is no need to call any conversion functions. The main limitation is that it does not handle gradients, because gradient information is not included in the standard array interfaces. This technique is therefore most suitable for algorithms that do not involve differentiation.

x = torch.arange(n, dtype=torch.float32, device=device)
y = torch.ones(n, dtype=torch.float32, device=device)

for i in range(10):
    wp.launch(saxpy, dim=n, inputs=[x, y, 1.0], device=device)
python -m warp.examples.benchmarks.benchmark_interop_torch

Sample output:

5095 ms  from_torch(...)
2113 ms  from_torch(..., return_ctype=True)
2950 ms  direct from torch

The default wp.from_torch() conversion is the slowest. Passing return_ctype=True is the fastest, because it skips creating temporary Warp array objects. Passing PyTorch tensors to Warp kernels directly falls somewhere in between. It skips creating temporary Warp arrays, but accessing the __cuda_array_interface__ attributes of PyTorch tensors adds overhead because they are initialized on-demand.

If you build a cache on top of these patterns (for example, keying on a tensor’s .data_ptr() or on a wp.array descriptor), invalidate the cache when the underlying Warp array is freed. A new allocation can reuse the same memory address with a different size, shape, or dtype, so pointer equality alone is not a safe cache key.

Case Study: PyTorch Deferred Gradient Allocation#

When writing custom PyTorch autograd functions that use Warp kernels, whether using analytic gradient kernels or the Warp tape, PyTorch’s deferred gradient allocation can cause unexpected synchronization delays. This case study demonstrates the problem and provides practical solutions.

The Problem: Deferred Gradient Allocation#

PyTorch employs a deferred allocation strategy for gradient tensors. When you create a tensor with requires_grad=True, PyTorch does not immediately allocate memory for the gradient. Instead, gradients are allocated on-demand during the backward pass.

However, when wp.from_torch() encounters a tensor with requires_grad=True but no allocated gradient, it forces the gradient to be allocated immediately. This creates overhead that can significantly impact performance.

When PyTorch later discovers that an external framework has allocated its gradient tensors, it must perform an expensive device-wide synchronization to ensure correctness. This synchronization overhead can significantly impact performance.

Here’s an example that demonstrates the problem:

import warp as wp
import torch

device = 'cuda'
N = 300_000_000

@wp.kernel(enable_backward=False)
def forward_kernel(
    a: wp.array[float],
    b: wp.array[float],
    output: wp.array[float]
):
    i = wp.tid()
    x = a[i]
    y = b[i]
    output[i] = x*x + y*y


@wp.kernel(enable_backward=False)
def backward_kernel(
    grad_output: wp.array[float],
    a: wp.array[float],
    b: wp.array[float],
    grad_a: wp.array[float],
    grad_b: wp.array[float]
):
    i = wp.tid()
    x = a[i]
    y = b[i]
    adj_z = grad_output[i]

    grad_a[i] = 2.0 * x * adj_z
    grad_b[i] = 2.0 * y * adj_z


class WarpFunction(torch.autograd.Function):

    @staticmethod
    def forward(ctx, a, b):
        ctx.save_for_backward(a, b)

        device = wp.device_from_torch(a.device)

        output = torch.empty(N, device=a.device)
        wp.launch(
            kernel=forward_kernel,
            dim=(N),
            device=device,
            inputs=[
                wp.from_torch(a),      # ⚠️ Triggers gradient allocation
                wp.from_torch(b),      # ⚠️ Triggers gradient allocation
                wp.from_torch(output),
            ],
        )

        return output

    @staticmethod
    def backward(ctx, grad_output):
        a, b = ctx.saved_tensors

        device = wp.device_from_torch(a.device)

        grad_a = torch.empty_like(a)
        grad_b = torch.empty_like(b)

        wp.launch(
            kernel=backward_kernel,
            dim=(N),
            device=device,
            inputs=[
                wp.from_torch(grad_output, requires_grad=False),
                wp.from_torch(a),
                wp.from_torch(b),
                wp.from_torch(grad_a),
                wp.from_torch(grad_b),
            ],
        )

        return grad_a, grad_b


a = torch.randn(N, device=device, dtype=torch.float32, requires_grad=True)
b = torch.randn(N, device=device, dtype=torch.float32, requires_grad=True)

torch.cuda.synchronize()

for i in range(TRIALS):
    with wp.ScopedTimer(f"TRIAL {i}", use_nvtx=True, synchronize=True):
        with wp.ScopedTimer("Create Tensors", use_nvtx=True, synchronize=True):
            a_torch = a.clone().detach().requires_grad_(True)
            b_torch = b.clone().detach().requires_grad_(True)
        with wp.ScopedTimer("Forward", use_nvtx=True, synchronize=True):
            output_warp = WarpFunction.apply(a_torch, b_torch)
        with wp.ScopedTimer("Loss", use_nvtx=True, synchronize=True):
            loss_warp = output_warp.sum()
        with wp.ScopedTimer("Backward", use_nvtx=True, synchronize=True):
            loss_warp.backward()

When profiling this code with NVIDIA Nsight Systems, significant gaps appear in the GPU timeline, indicating device-wide synchronization events:

Nsight Systems capture showing synchronization gaps between kernel launches

NVIDIA Nsight Systems timeline showing synchronization gaps between Warp kernel launches and PyTorch operations.#

The gaps represent device-wide synchronizations triggered when PyTorch discovers externally allocated gradients.

This problem is particularly severe in this benchmark because new tensors (a_torch and b_torch) are created on each iteration via .clone().detach().requires_grad_(True). Since these fresh tensors have requires_grad=True but no pre-allocated gradients, the synchronization penalty is incurred on every single iteration.

Solutions#

There are three approaches to avoid this synchronization overhead, depending on your use case:

Solution A: Disable Gradient Tracking in wp.from_torch()

The simplest solution is to pass requires_grad=False to wp.from_torch(), preventing Warp from auto-allocating gradients:

@staticmethod
def forward(ctx, a, b):
    ctx.save_for_backward(a, b)
    device = wp.device_from_torch(a.device)
    output = torch.empty(N, device=a.device)

    wp.launch(
        kernel=forward_kernel,
        dim=(N),
        device=device,
        inputs=[
            wp.from_torch(a, requires_grad=False),      # ✓ No gradient allocation
            wp.from_torch(b, requires_grad=False),      # ✓ No gradient allocation
            wp.from_torch(output, requires_grad=False),
        ],
    )
    return output

@staticmethod
def backward(ctx, grad_output):
    a, b = ctx.saved_tensors
    device = wp.device_from_torch(a.device)
    grad_a = torch.empty_like(a)
    grad_b = torch.empty_like(b)

    wp.launch(
        kernel=backward_kernel,
        dim=(N),
        device=device,
        inputs=[
            wp.from_torch(grad_output, requires_grad=False),
            wp.from_torch(a, requires_grad=False),
            wp.from_torch(b, requires_grad=False),
            wp.from_torch(grad_a, requires_grad=False),
            wp.from_torch(grad_b, requires_grad=False),
        ],
    )
    return grad_a, grad_b

This approach works well when you’re managing forward and gradient tensors separately and don’t need Warp to track gradients automatically.

Solution B: Detach Tensors from the PyTorch Graph

When managing gradients outside PyTorch’s autograd graph, detaching tensors before wrapping them is a clean approach:

@staticmethod
def forward(ctx, a, b):
    # Store detached tensors - we'll manage gradients manually
    ctx.a = a.detach()
    ctx.b = b.detach()

    device = wp.device_from_torch(a.device)
    output = torch.empty(N, device=a.device)

    wp.launch(
        kernel=forward_kernel,
        dim=(N),
        device=device,
        inputs=[
            wp.from_torch(ctx.a),      # ✓ Detached, no requires_grad
            wp.from_torch(ctx.b),      # ✓ Detached, no requires_grad
            wp.from_torch(output, requires_grad=False),
        ],
    )
    return output

@staticmethod
def backward(ctx, grad_output):
    device = wp.device_from_torch(ctx.a.device)
    grad_a = torch.empty_like(ctx.a)
    grad_b = torch.empty_like(ctx.b)

    wp.launch(
        kernel=backward_kernel,
        dim=(N),
        device=device,
        inputs=[
            wp.from_torch(grad_output.detach(), requires_grad=False),
            wp.from_torch(ctx.a),      # ✓ Already detached
            wp.from_torch(ctx.b),      # ✓ Already detached
            wp.from_torch(grad_a, requires_grad=False),
            wp.from_torch(grad_b, requires_grad=False),
        ],
    )
    return grad_a, grad_b

Detaching removes tensors from PyTorch’s computational graph (and clears requires_grad), making it clear that gradient management happens outside PyTorch’s autograd system.

Solution C: Pre-allocate Gradients with PyTorch

Alternatively, you can pre-allocate gradients using PyTorch’s allocator before passing tensors to Warp. This approach works for both analytic gradient kernels and when using the Warp tape.

Variant 1: With Analytic Gradient Kernels

@staticmethod
def forward(ctx, a, b):
    # Pre-allocate gradients using PyTorch's allocator
    if a.grad is None:
        a.grad = torch.empty_like(a)
    if b.grad is None:
        b.grad = torch.empty_like(b)

    ctx.save_for_backward(a, b)

    device = wp.device_from_torch(a.device)
    output = torch.empty(N, device=a.device)

    wp.launch(
        kernel=forward_kernel,
        dim=(N),
        device=device,
        inputs=[
            wp.from_torch(a),
            wp.from_torch(b),
            wp.from_torch(output),
        ],
    )
    return output

@staticmethod
def backward(ctx, grad_output):
    a, b = ctx.saved_tensors
    device = wp.device_from_torch(a.device)

    # Now we can use a.grad and b.grad directly
    wp.launch(
        kernel=backward_kernel,
        dim=(N),
        device=device,
        inputs=[
            wp.from_torch(grad_output, requires_grad=False),
            wp.from_torch(a),
            wp.from_torch(b),
            wp.from_torch(a.grad),  # ✓ Allocated by PyTorch
            wp.from_torch(b.grad),  # ✓ Allocated by PyTorch
        ],
    )
    return a.grad, b.grad

Variant 2: With Warp’s Tape (Automatic Differentiation)

Instead of implementing an analytic backward kernel, Warp can generate an adjoint for the forward kernel. Backward generation is enabled by default and can be invoked directly with wp.launch(..., adjoint=True) or managed across recorded launches with wp.Tape. The following example uses a tape to record the forward launch and run its generated adjoint:

@wp.kernel
def sum_squares(
    a: wp.array[float],
    b: wp.array[float],
    output: wp.array[float]
):
    i = wp.tid()
    x = a[i]
    y = b[i]
    output[i] = x*x + y*y


@staticmethod
def forward(ctx, a, b):
    ctx.save_for_backward(a, b)

    device = wp.device_from_torch(a.device)
    output = torch.zeros(N, device=a.device)

    # Pre-allocate zero-filled gradient buffers using PyTorch's allocator
    ctx.grad_a = torch.zeros_like(a)
    ctx.grad_b = torch.zeros_like(b)
    ctx.grad_output = torch.zeros_like(output)

    warp_stream = wp.stream_from_torch(a.device) if a.is_cuda else None
    with wp.ScopedStream(warp_stream):
        wp_a = wp.from_torch(a, grad=ctx.grad_a)
        wp_b = wp.from_torch(b, grad=ctx.grad_b)
        wp_output = wp.from_torch(output, requires_grad=True, grad=ctx.grad_output)

        with wp.Tape() as tape:
            wp.launch(
                kernel=sum_squares,
                dim=(N),
                device=device,
                inputs=[
                    wp_a,
                    wp_b,
                    wp_output,
                ],
            )

    ctx.tape = tape
    ctx.wp_output = wp_output

    return output

@staticmethod
def backward(ctx, grad_output):
    # Access saved_tensors so PyTorch checks that the inputs were not
    # modified in-place between forward and backward.
    _ = ctx.saved_tensors

    ctx.grad_a.zero_()
    ctx.grad_b.zero_()
    ctx.grad_output.zero_()

    warp_stream = wp.stream_from_torch(grad_output.device) if grad_output.is_cuda else None
    with wp.ScopedStream(warp_stream):
        ctx.tape.backward(
            grads={
                ctx.wp_output: wp.from_torch(grad_output, requires_grad=False),
            }
        )

        # Clone gradients before zeroing the tape, since these tensors are
        # zero-copy views of Warp/PyTorch gradient buffers.
        grad_a = ctx.grad_a.clone()
        grad_b = ctx.grad_b.clone()
        ctx.tape.zero()

    return grad_a, grad_b

This approach ensures adjoint buffers are allocated using PyTorch’s caching allocator, which properly tracks memory and stream dependencies. The input adjoints use dedicated zero-filled PyTorch tensors, separate from the input tensors’ .grad fields. The output Warp array also receives its own gradient buffer through grad=ctx.grad_output so the output adjoint passed through grads={...} is copied into owned storage instead of becoming the output array’s .grad storage. A larger wrapper can cache and reuse these dedicated adjoint buffers, but it should still zero them before each backward pass and pass them explicitly to wp.from_torch(..., grad=...):

# Allocate once, then reuse from the wrapper that calls wp.from_torch()
grad_a_buffer = torch.zeros_like(a_torch)
grad_b_buffer = torch.zeros_like(b_torch)
grad_output_buffer = torch.zeros_like(output_torch)

# Before each backward pass
grad_a_buffer.zero_()
grad_b_buffer.zero_()
grad_output_buffer.zero_()
wp_a = wp.from_torch(a_torch, grad=grad_a_buffer)
wp_b = wp.from_torch(b_torch, grad=grad_b_buffer)
wp_output = wp.from_torch(output_torch, requires_grad=True, grad=grad_output_buffer)

The following minimal, runnable example runs repeated backward with a reused external grad_outputs tensor and PyTorch gradcheck without letting the tape adopt external gradient storage:

import torch
import warp as wp


def active_torch_stream(tensor):
    if tensor.is_cuda:
        return wp.ScopedStream(wp.stream_from_torch(tensor.device))

    return wp.ScopedStream(None)


@wp.kernel
def square_kernel(x: wp.array[wp.float64], y: wp.array[wp.float64]):
    tid = wp.tid()
    y[tid] = x[tid] * x[tid]


class WarpSquare(torch.autograd.Function):

    @staticmethod
    def forward(ctx, x):
        ctx.save_for_backward(x)

        y = torch.empty_like(x)
        ctx.grad_x = torch.zeros_like(x)
        ctx.grad_y = torch.zeros_like(y)

        with active_torch_stream(x):
            # Warp records the detached input, then returns gradients manually to PyTorch.
            x_wp = wp.from_torch(
                x.detach(),
                dtype=wp.float64,
                requires_grad=True,
                grad=ctx.grad_x,
            )
            y_wp = wp.from_torch(
                y,
                dtype=wp.float64,
                requires_grad=True,
                grad=ctx.grad_y,
            )

            tape = wp.Tape()
            with tape:
                wp.launch(square_kernel, dim=x.numel(), inputs=[x_wp], outputs=[y_wp], device=x_wp.device)

        ctx.tape = tape
        # Keep Warp array wrappers alive until backward.
        ctx.x_wp = x_wp
        ctx.y_wp = y_wp

        return y

    @staticmethod
    def backward(ctx, grad_y):
        _ = ctx.saved_tensors

        ctx.grad_x.zero_()
        ctx.grad_y.zero_()

        with active_torch_stream(grad_y):
            # Since y_wp already has ctx.grad_y attached, Tape.backward()
            # copies grad_y into that buffer instead of retaining grad_y.
            ctx.tape.backward(
                grads={
                    ctx.y_wp: wp.from_torch(grad_y, dtype=wp.float64, requires_grad=False),
                }
            )

            grad_x = ctx.grad_x.clone()
            ctx.tape.zero()

        return grad_x


device = wp.get_device()
torch_device = wp.device_to_torch(device)

x = torch.tensor([1.0, -2.0, 3.0], dtype=torch.float64, device=torch_device, requires_grad=True)
y = WarpSquare.apply(x)

# Use a strided grad_outputs tensor to model external storage that PyTorch may reuse.
grad_y = torch.ones(
    6,
    dtype=torch.float64,
    device=torch_device,
)[::2]
expected_grad_y = grad_y.clone()

for i in range(3):
    (grad_x,) = torch.autograd.grad(y, (x,), grad_outputs=grad_y, retain_graph=i < 2)
    torch.testing.assert_close(grad_x, 2.0 * x.detach())
    torch.testing.assert_close(grad_y, expected_grad_y)

x_check = x.detach().clone().requires_grad_()
assert torch.autograd.gradcheck(WarpSquare.apply, (x_check,), eps=1.0e-6, atol=1.0e-5)

Performance Comparison#

Benchmarking these approaches on a workload with N=300,000,000 elements shows:

Baseline (with synchronization overhead):  98.02 ms
Solution A (requires_grad=False):          22.59 ms  (4.3x faster)
Solution B (detach):                       22.11 ms  (4.4x faster)
Solution C (pre-allocate):                 28.62 ms  (3.4x faster)

All three solutions eliminate the synchronization overhead:

  • Solutions A and B are fastest because they allocate gradients in the backward pass as simple standalone tensors

  • Solution C is slightly slower because it pre-allocates PyTorch-owned gradient buffers in the forward pass. The analytic-kernel path attaches them as .grad attributes, while the tape path passes them through grad=.

Choose based on your workflow:

  • Solution A: Most explicit about disabling gradient tracking

  • Solution B: Cleanest for manual gradient management

  • Solution C: Required when using Warp’s tape or when you need .grad access

These solutions are primarily needed when working with newly created tensors in scenarios like:

  • Training loops that create fresh tensors each iteration

  • Repeated inference with dynamically allocated tensors

  • Any workflow using .clone().detach().requires_grad_(True) patterns

If you recycle the same tensors across iterations (whose gradients have already been allocated), there will be no need for gradient allocation (deferred or otherwise) and therefore no synchronization overhead.