latest (dev)
Copy
Latest development documentation · Updated 2026-10-08
tensorplay.func API
Functions 17
batch_norm_without_running_stats
functionFull reference ↗- tensorplay.func.batch_norm_without_running_stats(module: Module) None[source]
Drops the running statistics of one batch-normalization module.
Running statistics are updated in place on every forward pass. Under a transform that evaluates the module more than once – mapping over an ensemble, or differentiating through it – those updates would be applied repeatedly and in an order the caller never asked for, so the module is switched to normalizing with the batch’s own statistics instead.
chunk_vmap
functionFull reference ↗- tensorplay.func.chunk_vmap(func: Callable, in_dims: int | tuple = 0, out_dims: int | tuple[int, ...] | None = 0, randomness: str = 'error', chunks: int = 2) Callable
vmap()splitting the batch intochunkspieces.Prefer
vmap(..., chunk_size=...), which states the bound in samples rather than in pieces and so does not change meaning with the batch size.
debug_unwrap
functionFull reference ↗- tensorplay.func.debug_unwrap(tensor: Any, *, recurse: bool = True) Any
Returns the plain tensor underlying a transform’s temporary wrapper.
Transforms in this build hand ordinary tensors to the function they wrap, so there is nothing to strip and the argument is returned unchanged. The entry point exists so that debugging code written against it keeps working if a wrapper representation is introduced later.
functional_call
functionFull reference ↗- tensorplay.func.functional_call(module: Module, parameter_and_buffer_dicts: dict[str, Any] | Sequence[dict[str, Any]], args: Any | tuple | None = None, kwargs: dict[str, Any] | None = None, *, tie_weights: bool = True, strict: bool = False)
Runs
modulewith the parameters and buffers given, not its own.The module’s own state is put back afterwards, including if
moduleraises, so this is safe to call on a live model.- Parameters:
module (tensorplay.nn.Module) – the module to call.
parameter_and_buffer_dicts (dict or sequence of dicts) – the state to substitute, keyed by the names
named_parametersandnamed_buffersreport. Several dicts are merged; overlapping keys are an error, since which one wins would be arbitrary.args (Any or tuple) – positional arguments for the module. A non-tuple value is passed as the single argument.
kwargs (dict) – keyword arguments for the module.
tie_weights (bool) – when the module ties two names to one tensor, keep them tied by requiring the replacement to be shared as well. Default:
True.strict (bool) – reject names that the module does not have. Default:
False.
Example
>>> params = dict(model.named_parameters()) >>> functional_call(model, params, (x,))
functionalize
functionFull reference ↗- tensorplay.func.functionalize(func: Callable, *, remove: str = 'mutations') Callable
Returns a version of
functhat leaves its arguments untouched.funcmay write into its inputs in place; the wrapper hands it private copies instead, so the caller’s tensors – and any views onto them – come back exactly as they went in, while the returned value is unchanged.- Parameters:
func (Callable) – the function to make side-effect free.
remove (str) –
"mutations"copies every tensor argument."mutations_and_views"additionally detaches the copies from any aliasing they arrived with, so writes through a view of an argument cannot reach another argument either.
Example
>>> def f(x): ... x.add_(1) ... return x >>> x = tensorplay.zeros(3) >>> out = functionalize(f)(x) >>> x # unchanged tensor([0., 0., 0.])
grad_and_value
functionFull reference ↗- tensorplay.func.grad_and_value(func: Callable, argnums: int | tuple[int, ...] = 0, has_aux: bool = False) Callable
Returns a function computing both the gradient of
funcand its value.The value comes from the same forward pass the gradient needs, so this costs no more than
grad()alone and saves evaluatingfunctwice.Returns a function producing
(gradient, value), or(gradient, (value, aux))whenhas_auxis set.
grad
functionFull reference ↗- tensorplay.func.grad(func: Callable, argnums: int | tuple[int, ...] = 0, has_aux: bool = False) Callable
Returns a function computing the gradient of
func.funcmust return a scalar tensor; the returned function has the same signature and returns the gradient with respect toargnums. Because it is again an ordinary function of the same inputs,grad(grad(f))is the second derivative.- Parameters:
Example
>>> x = tensorplay.randn([]) >>> grad(tensorplay.sin)(x)
hessian
functionFull reference ↗- tensorplay.func.hessian(func: Callable, argnums: int | tuple[int, ...] = 0)
Returns a function computing the Hessian of
func.Built as the reverse-mode jacobian of the reverse-mode gradient. The forward-over-reverse ordering would need tangents to flow through the backward pass itself, which this engine’s reverse mode does not carry; reverse-over-reverse composes two plain jacobians instead.
Example
>>> def f(x): ... return x.sin().sum() >>> hess = hessian(f)(tensorplay.randn(5)) >>> hess.shape tensorplay.Size(5, 5)
jacfwd
functionFull reference ↗- tensorplay.func.jacfwd(func: Callable, argnums: int | tuple[int, ...] = 0, has_aux: bool = False, *, randomness: str = 'error')
Returns a function computing the Jacobian of
funcby forward mode.Forward mode costs one pass per input element, so
jacfwdis the cheaper direction when the input is smaller than the output – the mirror image ofjacrev().- Parameters:
func (Callable) – a Python function returning one or more tensors.
argnums (int or Tuple[int]) – which positional arguments to differentiate with respect to. Default: 0.
has_aux (bool) – whether
funcreturns(output, aux).randomness (str) – how the underlying map treats random operations; one of
"error","different"or"same".
Example
>>> jacobian = jacfwd(tensorplay.sin)(tensorplay.randn(5))
jacrev
functionFull reference ↗- tensorplay.func.jacrev(func: Callable, argnums: int | tuple[int, ...] = 0, *, has_aux: bool = False, chunk_size: int | None = None, _preallocate_and_copy: bool = False)
Returns a function computing the Jacobian of
funcby reverse mode.Reverse mode costs one backward pass per output element, so
jacrevis the cheaper direction when the output is smaller than the input.- Parameters:
func (Callable) – a Python function returning one or more tensors.
argnums (int or Tuple[int]) – which positional arguments to differentiate with respect to. Default: 0.
has_aux (bool) – whether
funcreturns(output, aux).chunk_size (int, optional) – compute the Jacobian
chunk_sizerows at a time to bound peak memory.Nonecomputes it in one shot.
Example
>>> jacobian = jacrev(tensorplay.sin)(tensorplay.randn(5)) >>> jacobian.shape tensorplay.Size(5, 5)
jvp
functionFull reference ↗- tensorplay.func.jvp(func: Callable, primals: Any, tangents: Any, *, strict: bool = False, has_aux: bool = False)
Evaluates
funcatprimalstogether with its directional derivative alongtangents.- Parameters:
func (Callable) – a Python function taking one or more tensor arguments.
primals (Tensors) – a tuple of positional arguments to evaluate at.
tangents (Tensors) – the direction to differentiate along. Must have the same python structure, shapes and dtypes as
primals.strict (bool) – raise instead of returning zeros when the output turns out to be independent of the inputs.
has_aux (bool) – whether
funcreturns(output, aux).
- Returns:
(output, jvp_out), or(output, jvp_out, aux)whenhas_aux.
Example
>>> x = tensorplay.randn(5) >>> out, tangent_out = jvp(tensorplay.sin, (x,), (tensorplay.ones(5),))
linearize
functionFull reference ↗- tensorplay.func.linearize(func: Callable, *primals) tuple[Any, Callable]
Evaluates
funcatprimalsand returns its linear approximation there.The primal evaluation happens once, here; the returned
jvp_fnreuses that linearization, so applying it to many tangents costs only the tangent passes. Use it in place of repeatedjvp()calls at a fixed point.- Returns:
(output, jvp_fn).jvp_fntakes tangents matching the structure ofprimalsand returns the directional derivative of the output.
Example
>>> x = tensorplay.randn(5) >>> out, jvp_fn = linearize(tensorplay.sin, x) >>> tangent_out = jvp_fn(tensorplay.ones(5))
rearrange
functionFull reference ↗- tensorplay.func.rearrange(tensor: Any, pattern: str, **axes_lengths: int) Any[source]
Reshapes and permutes
tensoraccording topattern.- Parameters:
tensor – the tensor to rearrange, or a sequence of tensors, which is stacked along a new leading axis first.
pattern (str) –
"<input axes> -> <output axes>". Names bind by position on the left and select by name on the right. Parentheses group axes;...matches any number of axes.**axes_lengths – sizes for axes the pattern splits an input axis into, which cannot be inferred from the shape alone.
Example
>>> x = tensorplay.randn(2, 3, 4) >>> rearrange(x, "b h w -> b w h").shape tensorplay.Size(2, 4, 3) >>> rearrange(x, "b h w -> b (h w)").shape tensorplay.Size(2, 12) >>> rearrange(x, "(b1 b2) h w -> b1 b2 h w", b1=1).shape tensorplay.Size(1, 2, 3, 4)
replace_all_batch_norm_modules_
functionFull reference ↗- tensorplay.func.replace_all_batch_norm_modules_(root: Module) Module
Drops the running statistics of every batch-normalization module in
root, in place, and returnsroot.
stack_module_state
functionFull reference ↗- tensorplay.func.stack_module_state(models: list[Module]) tuple[dict[str, Any], dict[str, Any]]
Stacks the state of several identical modules into batched tensors.
Pair the result with
functional_call()undervmap()to evaluate a whole ensemble in one call instead of looping over the models.All models must be the same class and in the same training mode – a mix would make the stacked call mean two different things at once.
- Returns:
(stacked_params, stacked_buffers), each keyed as the modules’named_parameters/named_buffersare, with a new leading dimension of lengthlen(models).
Example
>>> params, buffers = stack_module_state(models) >>> def call(p, b, x): ... return functional_call(base_model, (p, b), (x,)) >>> vmap(call)(params, buffers, batched_x)
vjp
functionFull reference ↗- tensorplay.func.vjp(func: Callable, *primals, has_aux: bool = False)
Evaluates
funcatprimalsand returns a function computing the vector-Jacobian product.- Parameters:
func (Callable) – a Python function taking one or more tensor arguments.
primals (Tensors) – positional arguments to evaluate
funcat. The returned function differentiates with respect to all of them.has_aux (bool) – whether
funcreturns(output, aux), whereauxis carried through undifferentiated.
- Returns:
(output, vjp_fn), or(output, vjp_fn, aux)whenhas_aux.vjp_fntakes a cotangent with the same structure asoutputand returns the gradients with respect toprimals.
Example
>>> x = tensorplay.randn(5) >>> out, vjp_fn = vjp(tensorplay.sin, x) >>> (grad,) = vjp_fn(tensorplay.ones_like(out))
vmap
functionFull reference ↗- tensorplay.func.vmap(func: Callable, in_dims: int | tuple = 0, out_dims: int | tuple[int, ...] | None = 0, randomness: str = 'error', *, chunk_size: int | None = None) Callable
Returns a function that maps
funcover an added batch dimension.Write the function for a single sample;
vmaphandles the batch. That keeps the single-sample logic readable and removes the reshaping and unsqueezing that hand-batching otherwise scatters through it.- Parameters:
func (Callable) – a function taking one or more arguments, returning one or more tensors.
in_dims (int or nested structure) – which dimension of each input to map over.
Nonemarks an argument that is not batched and is passed through whole. The structure must be a prefix of the argument structure. Default: 0.out_dims (int or python collection) – where the mapped dimension should appear in each output. Default: 0.
randomness (str) – how random operations inside
funcbehave. With"error"(the default) they raise, because the intent is ambiguous;"different"draws fresh values per sample, and"same"replays the same values for every sample.chunk_size (int, optional) – process the batch
chunk_sizesamples at a time to bound peak memory.Noneprocesses it in one go.
Example
>>> def dot(x, y): ... return (x * y).sum() >>> x, y = tensorplay.randn(4, 3), tensorplay.randn(4, 3) >>> vmap(dot)(x, y).shape tensorplay.Size(4)
Help improve this page
Found an error, an unclear step, or a missing example?

