TensorPlay
latest (dev)
Copy
View Markdown

Latest development documentation · Updated 2026-10-08

tensorplay.func.vmap

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