latest (dev)
Copy
Latest development documentation · Updated 2026-10-08
tensorplay.func API reference
Composable function transforms.
A transform takes a function and returns a function, so they compose:
grad(grad(f)) is a second derivative, vmap(grad(f)) is a batch of
gradients, and jacfwd(jacrev(f)) is a Hessian. Everything here works on
plain Python callables over tensors – no module, no mutable state, no
grad field to clear between calls.
Modules do hold state, so functional_call() turns one into a function of
its parameters, and stack_module_state() batches an ensemble of them into
a single set of stacked tensors for vmap().
Function transforms
Returns a function that maps |
|
|
|
Returns a function computing the gradient of |
|
Returns a function computing both the gradient of |
|
Evaluates |
|
Evaluates |
|
Evaluates |
|
Returns a function computing the Jacobian of |
|
Returns a function computing the Jacobian of |
|
Returns a function computing the Hessian of |
|
Returns a version of |
|
Reshapes and permutes |
vmap() maps a function over a new batch dimension and grad()
differentiates it; the rest are built from or alongside those two.
grad_and_value() returns the value together with the gradient from a
single evaluation. vjp() and jvp() expose the vector-Jacobian
and Jacobian-vector products directly; jacrev() and jacfwd()
assemble full Jacobians with reverse- and forward-mode AD; hessian()
composes the two into a Hessian; and linearize() evaluates a function
once and returns its value together with a callable forward-mode
linearization at that point. functionalize() removes mutations from a
function, and rearrange() reshapes tensors by naming their axes.
chunk_vmap() is vmap() with the batch split into a fixed number
of chunks — prefer vmap(..., chunk_size=...), which bounds peak memory in
samples rather than in pieces and so does not change meaning with the batch
size.
Utilities for working with tensorplay.nn modules
A transform needs a function; a module holds state. In general you can transform over a function that calls a module directly:
import tensorplay as tp
from tensorplay.func import jacrev
model = tp.nn.Linear(3, 3)
def f(x):
return model(x)
x = tp.randn(3)
jacobian = jacrev(f)(x)
assert jacobian.shape == (3, 3)
To differentiate with respect to the module’s parameters, build a function
whose inputs are the parameters. That is what functional_call() is
for: it accepts a module, replacement state, and the inputs to the module’s
forward pass, and runs the module with the replacement state instead of its
own:
import tensorplay as tp
from tensorplay.func import functional_call, jacrev
model = tp.nn.Linear(3, 3)
def f(params, x):
return functional_call(model, params, (x,))
x = tp.randn(3)
jacobian = jacrev(f)(dict(model.named_parameters()), x)
assert jacobian["weight"].shape == (3, 3, 3)
assert jacobian["bias"].shape == (3, 3)
stack_module_state() stacks the state of several identical modules
into batched tensors, so vmap() can evaluate a whole ensemble in one
call instead of looping over the models:
import tensorplay as tp
import tensorplay.nn as nn
from tensorplay.func import functional_call, stack_module_state, vmap
models = [nn.Linear(3, 3) for _ in range(4)]
stacked_params, stacked_buffers = stack_module_state(models)
x = tp.randn(5, 3)
def call_one(params, buffers, x):
return functional_call(models[0], (params, buffers), (x,))
ensemble_out = vmap(call_one, in_dims=(0, 0, None))(
stacked_params, stacked_buffers, x
)
assert ensemble_out.shape == (4, 5, 3)
For batch norm modules under transforms, see patching batch norm.
Runs |
|
Stacks the state of several identical modules into batched tensors. |
|
Drops the running statistics of every batch-normalization module in |
|
Drops the running statistics of one batch-normalization module. |
Debug utilities
Returns the plain tensor underlying a transform's temporary wrapper. |
Inside a transformed function, tensors carry an invisible transform level.
debug_unwrap() removes it so the underlying tensor can be inspected —
printing its shape, checking a value in a debugger — without disturbing the
transform. Continue computing with the original argument, not with the
unwrapped tensor:
import tensorplay as tp
from tensorplay.func import debug_unwrap, vmap
def f(x):
print(debug_unwrap(x).shape) # the per-sample tensor, e.g. (3,)
return x * 2
out = vmap(f)(tp.randn(2, 3))
assert out.shape == (2, 3)
Help improve this page
Found an error, an unclear step, or a missing example?

