latest (dev)
Copy
Latest development documentation · Updated 2026-10-08
UX limitations
The transforms work best on pure functions: functions whose output is completely determined by their inputs and that do not perform side effects. Certain in-place operations are also supported. Writing code compatible with the transforms may involve changing habits, but in exchange the transforms let you express quantities — per-sample gradients, ensembles, higher-order derivatives — that are awkward to compute any other way.
General limitations
All transforms share one rule: everything you want out of a function must be returned from it. Assigning to a global variable (or to a mutable object captured from outside) does compute — the assignments happen — but the transform does not track them, so they cannot be differentiated or mapped over.
So, instead of the following:
import tensorplay as tp
from tensorplay.func import grad, vmap
intermediate = None
def f(x):
global intermediate
intermediate = x.sin() # computed, but invisible to the transform
return intermediate.sin()
x = tp.randn([])
grad_x = grad(f)(x) # fine, but `intermediate` is not a result
return the intermediate and mark it as auxiliary data:
def f(x):
intermediate = x.sin()
return intermediate.sin(), intermediate
grad_x, intermediate = grad(f, has_aux=True)(x)
tensorplay.autograd APIs
Using tensorplay.autograd.grad or Tensor.backward inside a function
being transformed is not supported. Under vmap() the call raises
IndexError (the batched tensor does not survive the plain-autograd
boundary), and under grad() the graph is silently severed — the inner
call runs on detached tensors and the outer gradient comes out zero. Use
the transform equivalents instead:
vmap limitations
vmap() is the most restrictive transform. The gradient-related
transforms (grad(), vjp(), jvp()) do not have the
restrictions listed here; jacfwd() and hessian() build on
forward-mode machinery and are subject to their own coverage gaps.
vmap(func) returns a function that maps func over some new dimension of
each tensor input. The mental model is a for-loop: for pure functions,
vmap(f)(x) is equivalent to tp.stack([f(x_i) for x_i in x.unbind(0)]).
Operator coverage
vmap() executes func once over batched inputs, so every operator
func calls needs a batching rule. An operator without one raises
NotImplementedError naming the operator, for example:
vmap(lambda a, b: a.dot(b))(tp.randn(3, 4), tp.randn(3, 4))
# NotImplementedError: Kernel not found for op: dot on backend: VmapCPU
Verified to work under vmap: the arithmetic operators (+, -, *, /,
**, unary -), the elementwise math functions (abs, exp, log,
sqrt, sin, cos, tanh, sigmoid, relu), matmul, mm, bmm,
sum, cumsum, clamp, where, stack, cat, index_select,
narrow, select, tril, maximum, minimum, logsumexp, the shape
operations (reshape, view, transpose, permute, movedim,
squeeze, unsqueeze, expand, contiguous), basic and advanced
indexing, the new_* factories plus tp.zeros/tp.ones, and randn.
Not yet covered — they raise the NotImplementedError above: dot,
mean, max, min, argmax, softmax, log_softmax, norm,
masked_fill, clone, sort, argsort, topk, einsum, dropout,
randn_like, the norm layers (batch_norm, group_norm, layer_norm),
nonzero, and item.
In-place operations
In-place arithmetic has no batching rule and raises cleanly:
vmap(lambda x, y: x.add_(y), in_dims=(0, 0))(tp.randn(3, 1), tp.randn(3, 1))
# NotImplementedError: Kernel not found for op: add_.Tensor on backend: VmapCPU
This holds whether the mutated tensor is batched or not — prefer the
out-of-place form under vmap, or x + y directly.
Data-dependent operations
.item() raises RuntimeError under vmap — converting a batched value
to a scalar is rejected because the result would depend on which sample is
looked at — and nonzero raises NotImplementedError; the output of a
data-dependent op varies per sample, so no single stacked result exists.
Rewrite the code to avoid materializing per-sample shapes.
Data-dependent Python control flow
Warning
The condition of an if (or a while/for test) must not be a tensor
being mapped over. The boolean conversion fails with a RuntimeError that
explains the problem:
def relu(x):
if x > 0: # x is a mapped tensor here: do not do this
return x
return 0 * x
vmap(relu)(tp.randn(3))
# RuntimeError: item() is not supported on a vmap-batched tensor;
# data-dependent control flow cannot see per-batch values. ...
Re-express value-dependent branches with tensorplay.where(), whose
comparison and selection rules are supported.
The out= keyword
The public operator signatures do not accept out= (passing one raises
TypeError about the argument combination), so there is no out= form
for vmap to transform.
Randomness
The intent of a random operation under vmap is ambiguous — should every
sample see fresh values, or the same ones? — so vmap takes a
randomness flag instead of guessing:
"error"(the default): any random operation raisesRuntimeError: random operations are not allowed under randomness=error."different": elements of the batch draw different values."same": every element of the batch sees the same value.
def add_noise(x):
return x + tp.randn(())
x = tp.ones(3)
r = vmap(add_noise, randomness="same")(x)
assert (r == r[0]).all() # one value, repeated
r = vmap(add_noise, randomness="different")(x)
assert (r != r[0]).any() # fresh value per sample
An invalid value raises before the call: only "error", "different",
and "same" are accepted. The flag governs the tensor factories such as
tp.randn; dropout and randn_like do not have batching rules yet and
raise regardless of the flag.
Composability
Verified compositions of the transforms with each other:
Composition |
Status |
|---|---|
|
works — per-sample gradients |
|
raises |
|
works — two independent batch dimensions |
|
works |
|
raises — the internal cotangent pass loses the batch shape |
|
raises — the forward-mode rule for |
Norm layers
batch_norm, group_norm, and layer_norm all lack batching rules, so
vmap() over any of them raises. The reverse-mode transforms
(grad() and jacrev()) do work through batch_norm — the
running-stat update is in-place, so it bypasses the transform machinery,
but the values themselves are correct. See
patching batch norm before transforming a model that
contains normalization.
Help improve this page
Found an error, an unclear step, or a missing example?

