Coverage for src/evutils/transforms/compose.py: 77%

71 statements  

« prev     ^ index     » next       coverage.py v7.15.1, created at 2026-07-18 05:24 +0000

1import numpy as np 

2from typing import Iterable, Callable, Union 

3from evutils.types import SoaArray 

4 

5def unwrap_events(events: Union[np.ndarray, SoaArray]): 

6 """Unwraps events into a tuple of (t, x, y, p) arrays.""" 

7 if isinstance(events, SoaArray): 

8 return events.t, events.x, events.y, events.p 

9 elif isinstance(events, np.ndarray) and events.dtype.names is not None: 

10 return events['t'], events['x'], events['y'], events['p'] 

11 else: 

12 raise TypeError(f"Unsupported event format: {type(events)}") 

13 

14def repack_events(original_events: Union[np.ndarray, SoaArray], t: np.ndarray, x: np.ndarray, y: np.ndarray, p: np.ndarray) -> Union[np.ndarray, SoaArray]: 

15 """Repacks raw arrays back into the user's original format.""" 

16 if isinstance(original_events, SoaArray): 

17 # Carry any metadata (e.g. sensor_size) through the transform. Ops that 

18 # invalidate it (spatial cropping) must update it themselves. 

19 return original_events.__class__(t=t, x=x, y=y, p=p, 

20 metadata=original_events.metadata) 

21 elif isinstance(original_events, np.ndarray) and original_events.dtype.names is not None: 

22 new_events = np.empty(len(t), dtype=original_events.dtype) 

23 new_events['t'] = t 

24 new_events['x'] = x 

25 new_events['y'] = y 

26 new_events['p'] = p 

27 return new_events 

28 else: 

29 raise TypeError(f"Unsupported event format: {type(original_events)}") 

30 

31class Compose: 

32 """Composes several transforms together. 

33  

34 Groups contiguous blocks of evutils `Transform` objects to execute them purely in  

35 C-space via Numba JIT without intermediate unpacking/repacking overhead. 

36 Freely accepts standard Callables (e.g. PyTorch/Tonic transforms). 

37 """ 

38 

39 def __init__(self, transforms: Iterable[Callable]): 

40 import inspect 

41 self.transforms = list(transforms) 

42 self._execution_plan = [] 

43 

44 # Pre-compute the JIT blocks 

45 current_jit_block = [] 

46 for t in self.transforms: 

47 if hasattr(t, "_forward_jit"): 

48 current_jit_block.append(t) 

49 else: 

50 if current_jit_block: 

51 self._execution_plan.append(("jit", current_jit_block, False)) 

52 current_jit_block = [] 

53 

54 try: 

55 sig = inspect.signature(t) 

56 accepts_target = len(sig.parameters) > 1 

57 except ValueError: 

58 accepts_target = False 

59 

60 self._execution_plan.append(("standard", t, accepts_target)) 

61 

62 if current_jit_block: 

63 self._execution_plan.append(("jit", current_jit_block, False)) 

64 

65 def __call__(self, events, target=None): 

66 for step_type, block, accepts_target in self._execution_plan: 

67 if len(events) == 0: 

68 break 

69 

70 if step_type == "jit": 

71 t, x, y, p = unwrap_events(events) 

72 for transform in block: 

73 # Let the transform fall back to the container's sensor_size. 

74 transform.bind_context(events) 

75 # A transform earlier in the block may have dropped every 

76 # event; skip kernels that assume non-empty input (e.g. 

77 # deriving a sensor extent from x.max()), but still let the 

78 # transform update the target -- matching standalone 

79 # Transform.__call__, which applies both regardless of count. 

80 if len(t) > 0: 

81 t, x, y, p = transform._forward_jit(t, x, y, p) 

82 if target is not None: 

83 target = transform._transform_target(target) 

84 events = repack_events(events, t, x, y, p) 

85 else: 

86 if target is not None: 

87 if accepts_target: 

88 res = block(events, target) 

89 if isinstance(res, tuple) and len(res) == 2: 

90 events, target = res 

91 else: 

92 events = res 

93 else: 

94 events = block(events) 

95 else: 

96 events = block(events) 

97 

98 if target is not None: 

99 return events, target 

100 return events 

101 

102 def __repr__(self): 

103 format_string = self.__class__.__name__ + "(" 

104 for t in self.transforms: 

105 format_string += "\n" 

106 format_string += f" {t}" 

107 format_string += "\n)" 

108 return format_string