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
« 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
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)}")
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)}")
31class Compose:
32 """Composes several transforms together.
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 """
39 def __init__(self, transforms: Iterable[Callable]):
40 import inspect
41 self.transforms = list(transforms)
42 self._execution_plan = []
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 = []
54 try:
55 sig = inspect.signature(t)
56 accepts_target = len(sig.parameters) > 1
57 except ValueError:
58 accepts_target = False
60 self._execution_plan.append(("standard", t, accepts_target))
62 if current_jit_block:
63 self._execution_plan.append(("jit", current_jit_block, False))
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
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)
98 if target is not None:
99 return events, target
100 return events
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