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

1from typing import Optional 

2 

3import pandas as pd 

4 

5from .axis import AxisId 

6from .common import PerMember 

7from .tensor import Tensor 

8 

9 

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 

17 

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 

27 

28 return outputs