Coverage for src/bioimageio/core/_restore_batch_multi_index.py: 93%
14 statements
« prev ^ index » next coverage.py v7.15.0, created at 2026-07-08 15:59 +0000
« prev ^ index » next coverage.py v7.15.0, created at 2026-07-08 15:59 +0000
1from typing import Optional
3import pandas as pd
5from .axis import AxisId
6from .common import PerMember
7from .tensor import Tensor
10def restore_batch_multi_index(
11 inputs: PerMember[Optional[Tensor]], outputs: PerMember[Optional[Tensor]]
12) -> PerMember[Optional[Tensor]]:
13 """Restore the first batch multi-index found in the inputs to all outputs with batch dimension."""
14 for tensor in inputs.values():
15 if tensor is None:
16 continue
18 idx = tensor.data.indexes.get(AxisId("batch")) # pyright: ignore[reportUnknownVariableType]
19 if isinstance(idx, pd.MultiIndex):
20 outputs = {
21 k: v.assign_batch_multi_index(idx)
22 if v is not None and AxisId("batch") in v.dims
23 else v
24 for k, v in outputs.items()
25 }
26 break
28 return outputs