# tensorplay.func API reference Source: https://www.tensorplay.cn/docs/func.api.html 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()](/docs/generated/tensorplay.func.functional_call.html#tensorplay.func.functional_call) turns one into a function of its parameters, and [stack_module_state()](/docs/generated/tensorplay.func.stack_module_state.html#tensorplay.func.stack_module_state) batches an ensemble of them into a single set of stacked tensors for [vmap()](/docs/generated/tensorplay.func.vmap.html#tensorplay.func.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()](/docs/generated/tensorplay.func.vmap.html#tensorplay.func.vmap) maps a function over a new batch dimension and [grad()](/docs/generated/tensorplay.func.grad.html#tensorplay.func.grad) differentiates it; the rest are built from or alongside those two. [grad_and_value()](/docs/generated/tensorplay.func.grad_and_value.html#tensorplay.func.grad_and_value) returns the value together with the gradient from a single evaluation. [vjp()](/docs/generated/tensorplay.func.vjp.html#tensorplay.func.vjp) and [jvp()](/docs/generated/tensorplay.func.jvp.html#tensorplay.func.jvp) expose the vector-Jacobian and Jacobian-vector products directly; [jacrev()](/docs/generated/tensorplay.func.jacrev.html#tensorplay.func.jacrev) and [jacfwd()](/docs/generated/tensorplay.func.jacfwd.html#tensorplay.func.jacfwd) assemble full Jacobians with reverse- and forward-mode AD; [hessian()](/docs/generated/tensorplay.func.hessian.html#tensorplay.func.hessian) composes the two into a Hessian; and [linearize()](/docs/generated/tensorplay.func.linearize.html#tensorplay.func.linearize) evaluates a function once and returns its value together with a callable forward-mode linearization at that point. [functionalize()](/docs/generated/tensorplay.func.functionalize.html#tensorplay.func.functionalize) removes mutations from a function, and [rearrange()](/docs/generated/tensorplay.func.rearrange.html#tensorplay.func.rearrange) reshapes tensors by naming their axes. [chunk_vmap()](/docs/generated/tensorplay.func.chunk_vmap.html#tensorplay.func.chunk_vmap) is [vmap()](/docs/generated/tensorplay.func.vmap.html#tensorplay.func.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()](/docs/generated/tensorplay.func.functional_call.html#tensorplay.func.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()](/docs/generated/tensorplay.func.stack_module_state.html#tensorplay.func.stack_module_state) stacks the state of several identical modules into batched tensors, so [vmap()](/docs/generated/tensorplay.func.vmap.html#tensorplay.func.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](/docs/func.batch_norm.html). | 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()](/docs/generated/tensorplay.func.debug_unwrap.html#tensorplay.func.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) ```