TensorPlay
API reference
latest (dev)
Copy
View Markdown

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 into chunks pieces.

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 module with the parameters and buffers given, not its own.

The module’s own state is put back afterwards, including if module raises, 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_parameters and named_buffers report. 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 func that leaves its arguments untouched.

func may 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 func and its value.

The value comes from the same forward pass the gradient needs, so this costs no more than grad() alone and saves evaluating func twice.

Returns a function producing (gradient, value), or (gradient, (value, aux)) when has_aux is 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.

func must return a scalar tensor; the returned function has the same signature and returns the gradient with respect to argnums. Because it is again an ordinary function of the same inputs, grad(grad(f)) is the second derivative.

Parameters:
  • func (Callable) – a function returning a single-element tensor.

  • argnums (int or Tuple[int]) – which positional arguments to differentiate with respect to. Default: 0.

  • has_aux (bool) – whether func returns (output, aux), where aux is carried through undifferentiated.

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 func by forward mode.

Forward mode costs one pass per input element, so jacfwd is the cheaper direction when the input is smaller than the output – the mirror image of jacrev().

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 func returns (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 func by reverse mode.

Reverse mode costs one backward pass per output element, so jacrev is 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 func returns (output, aux).

  • chunk_size (int, optional) – compute the Jacobian chunk_size rows at a time to bound peak memory. None computes 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 func at primals together with its directional derivative along tangents.

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 func returns (output, aux).

Returns:

(output, jvp_out), or (output, jvp_out, aux) when has_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 func at primals and returns its linear approximation there.

The primal evaluation happens once, here; the returned jvp_fn reuses that linearization, so applying it to many tangents costs only the tangent passes. Use it in place of repeated jvp() calls at a fixed point.

Returns:

(output, jvp_fn). jvp_fn takes tangents matching the structure of primals and 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 tensor according to pattern.

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 returns root.

#

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() under vmap() 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_buffers are, with a new leading dimension of length len(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 func at primals and 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 func at. The returned function differentiates with respect to all of them.

  • has_aux (bool) – whether func returns (output, aux), where aux is carried through undifferentiated.

Returns:

(output, vjp_fn), or (output, vjp_fn, aux) when has_aux. vjp_fn takes a cotangent with the same structure as output and returns the gradients with respect to primals.

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 func over an added batch dimension.

Write the function for a single sample; vmap handles 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. None marks 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 func behave. 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_size samples at a time to bound peak memory. None processes 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)

On this page

Ask DeepWiki