|
4 | 4 | from enum import IntEnum |
5 | 5 | from typing import TYPE_CHECKING |
6 | 6 |
|
| 7 | +import numpy as np |
| 8 | + |
7 | 9 | if TYPE_CHECKING: |
8 | 10 | import numpy as np |
9 | 11 |
|
@@ -67,7 +69,6 @@ def search(self, z: float, y: float, x: float, ei=None) -> dict[str, tuple[int, |
67 | 69 | """ |
68 | 70 | ... |
69 | 71 |
|
70 | | - @abstractmethod |
71 | 72 | def ravel_index(self, axis_indices: dict[str, int]) -> int: |
72 | 73 | """ |
73 | 74 | Convert a dictionary of axis indices to a single encoded index (ei). |
@@ -101,9 +102,10 @@ def ravel_index(self, axis_indices: dict[str, int]) -> int: |
101 | 102 | NotImplementedError |
102 | 103 | Raised if the method is not implemented for the current grid type. |
103 | 104 | """ |
104 | | - ... |
| 105 | + dims = np.array([self.get_axis_dim(axis) for axis in self.axes], dtype=int) |
| 106 | + indices = np.array([axis_indices[axis] for axis in self.axes], dtype=int) |
| 107 | + return _ravel(dims, indices) |
105 | 108 |
|
106 | | - @abstractmethod |
107 | 109 | def unravel_index(self, ei: int) -> dict[str, int]: |
108 | 110 | """ |
109 | 111 | Convert a single encoded index (ei) back to a dictionary of axis indices. |
@@ -134,14 +136,19 @@ def unravel_index(self, ei: int) -> dict[str, int]: |
134 | 136 | NotImplementedError |
135 | 137 | Raised if the method is not implemented for the current grid type. |
136 | 138 | """ |
137 | | - ... |
| 139 | + dims = np.array([self.get_axis_dim(axis) for axis in self.axes], dtype=int) |
| 140 | + indices = _unravel(dims, ei) |
| 141 | + return dict(zip(self.axes, indices, strict=True)) |
138 | 142 |
|
139 | 143 | @property |
140 | 144 | @abstractmethod |
141 | 145 | def axes(self) -> list[str]: |
142 | 146 | """ |
143 | 147 | Return a list of axis names that are part of this grid. |
144 | 148 |
|
| 149 | + This list must at least be of length 1, and `get_axis_dim` should |
| 150 | + return a valid integer for each axis name in the list. |
| 151 | +
|
145 | 152 | Returns |
146 | 153 | ------- |
147 | 154 | list[str] |
@@ -170,3 +177,65 @@ def get_axis_dim(self, axis: str) -> int: |
170 | 177 | If the specified axis is not part of this grid. |
171 | 178 | """ |
172 | 179 | ... |
| 180 | + |
| 181 | + |
| 182 | +def _unravel(dims, ei): |
| 183 | + """ |
| 184 | + Converts a flattened (raveled) index back to multi-dimensional indices. |
| 185 | +
|
| 186 | + Args: |
| 187 | + dims (1d-array-like): The dimensions along each axis |
| 188 | + ei (int): The flattened index to convert |
| 189 | +
|
| 190 | + Returns |
| 191 | + ------- |
| 192 | + array-like: Indices along each axis corresponding to the given flattened index |
| 193 | +
|
| 194 | + Example: |
| 195 | + >>> dims = [2, 3, 4] |
| 196 | + >>> ei = 9 |
| 197 | + >>> unravel(dims, ei) |
| 198 | + array([0, 2, 1]) |
| 199 | + # Calculation: |
| 200 | + # i0 = 9 // (3*4) = 9 // 12 = 0 |
| 201 | + # remainder = 9 % 12 = 9 |
| 202 | + # i1 = 9 // 4 = 2 |
| 203 | + # i2 = 9 % 4 = 1 |
| 204 | + """ |
| 205 | + strides = np.cumprod(dims[::-1])[::-1] |
| 206 | + |
| 207 | + indices = np.empty(len(dims), dtype=int) |
| 208 | + |
| 209 | + for i in range(len(dims) - 1): |
| 210 | + indices[i] = ei // strides[i + 1] |
| 211 | + ei = ei % strides[i + 1] |
| 212 | + |
| 213 | + indices[-1] = ei |
| 214 | + return indices |
| 215 | + |
| 216 | + |
| 217 | +def _ravel(dims, indices): |
| 218 | + """ |
| 219 | + Converts indices to a flattened (raveled) index. |
| 220 | +
|
| 221 | + Args: |
| 222 | + dims (1d-array-like): The dimensions along each axis |
| 223 | + indices (array-like): Indices along each axis to convert |
| 224 | +
|
| 225 | + Returns |
| 226 | + ------- |
| 227 | + int: The flattened index corresponding to the given indices |
| 228 | +
|
| 229 | + Example: |
| 230 | + >>> dims = [2, 3, 4] |
| 231 | + >>> indices = [0, 2, 1] |
| 232 | + >>> ravel(dims, indices) |
| 233 | + 9 |
| 234 | + # Calculation: 0 * (3 * 4) + 2 * (4) + 1 = 0 + 8 + 1 = 9 |
| 235 | + """ |
| 236 | + strides = np.cumprod(dims[::-1])[::-1] |
| 237 | + ei = 0 |
| 238 | + for i in range(len(dims) - 1): |
| 239 | + ei += indices[i] * strides[i + 1] |
| 240 | + |
| 241 | + return ei + indices[-1] |
0 commit comments