TensorPlay
Reference guides
latest (dev)
Copy
View Markdown

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

vmap

Returns a function that maps func over an added batch dimension.

chunk_vmap

vmap() splitting the batch into chunks pieces.

grad

Returns a function computing the gradient of func.

grad_and_value

Returns a function computing both the gradient of func and its value.

vjp

Evaluates func at primals and returns a function computing the vector-Jacobian product.

jvp

Evaluates func at primals together with its directional derivative along tangents.

linearize

Evaluates func at primals and returns its linear approximation there.

jacrev

Returns a function computing the Jacobian of func by reverse mode.

jacfwd

Returns a function computing the Jacobian of func by forward mode.

hessian

Returns a function computing the Hessian of func.

functionalize

Returns a version of func that leaves its arguments untouched.

rearrange

Reshapes and permutes tensor according to pattern.

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.

functional_call

Runs module with the parameters and buffers given, not its own.

stack_module_state

Stacks the state of several identical modules into batched tensors.

replace_all_batch_norm_modules_

Drops the running statistics of every batch-normalization module in root, in place, and returns root.

batch_norm_without_running_stats

Drops the running statistics of one batch-normalization module.

Debug utilities

debug_unwrap

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)

On this page

Ask DeepWiki