TensorPlay
latest (dev)
Copy
View Markdown

Latest development documentation · Updated 2026-10-08

tensorplay.func.stack_module_state

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)

On this page

Ask DeepWiki