latest (dev)
Copy
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
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?

