latest (dev)
Copy
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()undervmap()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_buffersare, with a new leading dimension of lengthlen(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)
Help improve this page
Found an error, an unclear step, or a missing example?
Was this page helpful?

