arraybridge 0.3.4__py3-none-any.whl → 0.3.7__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- arraybridge/__init__.py +1 -1
- arraybridge/array_operations.py +378 -324
- arraybridge/decorators.py +41 -16
- arraybridge/types.py +38 -7
- {arraybridge-0.3.4.dist-info → arraybridge-0.3.7.dist-info}/METADATA +16 -1
- {arraybridge-0.3.4.dist-info → arraybridge-0.3.7.dist-info}/RECORD +8 -8
- {arraybridge-0.3.4.dist-info → arraybridge-0.3.7.dist-info}/WHEEL +0 -0
- {arraybridge-0.3.4.dist-info → arraybridge-0.3.7.dist-info}/licenses/LICENSE +0 -0
arraybridge/__init__.py
CHANGED
arraybridge/array_operations.py
CHANGED
|
@@ -1,19 +1,11 @@
|
|
|
1
|
-
"""
|
|
1
|
+
"""Native array operations carried by MemoryType declarations."""
|
|
2
2
|
|
|
3
|
-
from
|
|
4
|
-
from
|
|
3
|
+
from abc import ABC, abstractmethod
|
|
4
|
+
from collections.abc import Sequence
|
|
5
5
|
from typing import Any
|
|
6
6
|
|
|
7
7
|
import numpy as np
|
|
8
8
|
|
|
9
|
-
ToNumpy = Callable[[Any, Any], Any]
|
|
10
|
-
FromNumpy = Callable[[Any, Any, int], Any]
|
|
11
|
-
StackArrays = Callable[[Sequence[Any], Any], Any]
|
|
12
|
-
ScaleDtype = Callable[[Any, Any, Any], Any]
|
|
13
|
-
DtypeName = Callable[[Any], str]
|
|
14
|
-
CastArray = Callable[[Any, Any, Any], Any]
|
|
15
|
-
LogicalAnd = Callable[[Any, Any, Any], Any]
|
|
16
|
-
|
|
17
9
|
_SCALING_RANGES: dict[str, float | tuple[float, float]] = {
|
|
18
10
|
"uint8": 255.0,
|
|
19
11
|
"uint16": 65535.0,
|
|
@@ -64,103 +56,6 @@ def _clamp_bounds(target_dtype: Any) -> tuple[float, float] | None:
|
|
|
64
56
|
return 0, range_info
|
|
65
57
|
|
|
66
58
|
|
|
67
|
-
def _identity_to_numpy(data: Any, module: Any) -> Any:
|
|
68
|
-
del module
|
|
69
|
-
return data
|
|
70
|
-
|
|
71
|
-
|
|
72
|
-
def _identity_from_numpy(data: Any, module: Any, device_id: int) -> Any:
|
|
73
|
-
del module, device_id
|
|
74
|
-
return data
|
|
75
|
-
|
|
76
|
-
|
|
77
|
-
def _numpy_stack(values: Sequence[Any], module: Any) -> Any:
|
|
78
|
-
return module.stack(values, axis=0)
|
|
79
|
-
|
|
80
|
-
|
|
81
|
-
def _numpy_cast(data: Any, dtype: Any, module: Any) -> Any:
|
|
82
|
-
del module
|
|
83
|
-
return data.astype(dtype, copy=False)
|
|
84
|
-
|
|
85
|
-
|
|
86
|
-
def _module_logical_and(left: Any, right: Any, module: Any) -> Any:
|
|
87
|
-
return module.logical_and(left, right)
|
|
88
|
-
|
|
89
|
-
|
|
90
|
-
def _numpy_scale(result: Any, target_dtype: Any, module: Any) -> Any:
|
|
91
|
-
if not hasattr(result, "dtype"):
|
|
92
|
-
return result
|
|
93
|
-
if not (
|
|
94
|
-
module.issubdtype(result.dtype, module.floating)
|
|
95
|
-
and module.issubdtype(target_dtype, module.integer)
|
|
96
|
-
):
|
|
97
|
-
return result.astype(target_dtype)
|
|
98
|
-
result_min = result.min()
|
|
99
|
-
result_max = result.max()
|
|
100
|
-
if result_max <= result_min:
|
|
101
|
-
return result.astype(target_dtype)
|
|
102
|
-
scaled = _scaled_values(result, result_min, result_max, target_dtype)
|
|
103
|
-
bounds = _clamp_bounds(target_dtype)
|
|
104
|
-
if bounds is not None:
|
|
105
|
-
scaled = module.clip(scaled, *bounds)
|
|
106
|
-
return scaled.astype(target_dtype)
|
|
107
|
-
|
|
108
|
-
|
|
109
|
-
def _cupy_to_numpy(data: Any, module: Any) -> Any:
|
|
110
|
-
del module
|
|
111
|
-
return data.get()
|
|
112
|
-
|
|
113
|
-
|
|
114
|
-
def _cupy_from_numpy(data: Any, module: Any, device_id: int) -> Any:
|
|
115
|
-
del device_id
|
|
116
|
-
return module.array(data)
|
|
117
|
-
|
|
118
|
-
|
|
119
|
-
def _cupy_stack(values: Sequence[Any], module: Any) -> Any:
|
|
120
|
-
return module.stack(values, axis=0)
|
|
121
|
-
|
|
122
|
-
|
|
123
|
-
def _cupy_scale(result: Any, target_dtype: Any, module: Any) -> Any:
|
|
124
|
-
if not hasattr(result, "dtype"):
|
|
125
|
-
return result
|
|
126
|
-
if not (
|
|
127
|
-
module.issubdtype(result.dtype, module.floating)
|
|
128
|
-
and not module.issubdtype(target_dtype, module.floating)
|
|
129
|
-
):
|
|
130
|
-
return result.astype(target_dtype)
|
|
131
|
-
result_min = module.min(result)
|
|
132
|
-
result_max = module.max(result)
|
|
133
|
-
if result_max <= result_min:
|
|
134
|
-
return result.astype(target_dtype)
|
|
135
|
-
scaled = _scaled_values(result, result_min, result_max, target_dtype)
|
|
136
|
-
bounds = _clamp_bounds(target_dtype)
|
|
137
|
-
if bounds is not None:
|
|
138
|
-
scaled = module.clip(scaled, *bounds)
|
|
139
|
-
return scaled.astype(target_dtype)
|
|
140
|
-
|
|
141
|
-
|
|
142
|
-
def _torch_to_numpy(data: Any, module: Any) -> Any:
|
|
143
|
-
del module
|
|
144
|
-
return data.cpu().numpy()
|
|
145
|
-
|
|
146
|
-
|
|
147
|
-
def _torch_from_numpy(data: Any, module: Any, device_id: int) -> Any:
|
|
148
|
-
host_data = (
|
|
149
|
-
np.ascontiguousarray(data)
|
|
150
|
-
if any(stride < 0 for stride in getattr(data, "strides", ()))
|
|
151
|
-
else data
|
|
152
|
-
)
|
|
153
|
-
return module.from_numpy(host_data).to(f"cuda:{device_id}")
|
|
154
|
-
|
|
155
|
-
|
|
156
|
-
def _torch_stack(values: Sequence[Any], module: Any) -> Any:
|
|
157
|
-
return module.stack(tuple(values), dim=0)
|
|
158
|
-
|
|
159
|
-
|
|
160
|
-
def _torch_cast(data: Any, dtype: Any, module: Any) -> Any:
|
|
161
|
-
return data.to(dtype=_mapped_dtype(dtype, module))
|
|
162
|
-
|
|
163
|
-
|
|
164
59
|
def _mapped_dtype(target_dtype: Any, module: Any) -> Any:
|
|
165
60
|
try:
|
|
166
61
|
dtype_name = np.dtype(target_dtype).name
|
|
@@ -173,224 +68,383 @@ def _mapped_dtype(target_dtype: Any, module: Any) -> Any:
|
|
|
173
68
|
return mapped
|
|
174
69
|
|
|
175
70
|
|
|
176
|
-
|
|
177
|
-
|
|
71
|
+
class ArrayOperations(ABC):
|
|
72
|
+
"""Native array semantics selected by the existing MemoryType declaration."""
|
|
73
|
+
|
|
74
|
+
@staticmethod
|
|
75
|
+
@abstractmethod
|
|
76
|
+
def to_numpy(data: Any, module: Any) -> Any:
|
|
77
|
+
"""Project pixels to the explicit host boundary."""
|
|
78
|
+
|
|
79
|
+
@staticmethod
|
|
80
|
+
@abstractmethod
|
|
81
|
+
def from_numpy(data: Any, module: Any, device_id: int) -> Any:
|
|
82
|
+
"""Admit host pixels on the selected framework device."""
|
|
83
|
+
|
|
84
|
+
@staticmethod
|
|
85
|
+
@abstractmethod
|
|
86
|
+
def scale_dtype(result: Any, target_dtype: Any, module: Any) -> Any:
|
|
87
|
+
"""Apply the framework's existing intensity conversion law."""
|
|
88
|
+
|
|
89
|
+
@staticmethod
|
|
90
|
+
def dtype_name(dtype: Any) -> str:
|
|
91
|
+
return _numpy_dtype_name(dtype)
|
|
92
|
+
|
|
93
|
+
@staticmethod
|
|
94
|
+
def stack(values: Sequence[Any], module: Any) -> Any:
|
|
95
|
+
return module.stack(tuple(values), axis=0)
|
|
96
|
+
|
|
97
|
+
@staticmethod
|
|
98
|
+
def cast(data: Any, dtype: Any, module: Any) -> Any:
|
|
99
|
+
return data.astype(dtype, copy=False)
|
|
100
|
+
|
|
101
|
+
@staticmethod
|
|
102
|
+
def logical_and(left: Any, right: Any, module: Any) -> Any:
|
|
103
|
+
return module.logical_and(left, right)
|
|
104
|
+
|
|
105
|
+
@staticmethod
|
|
106
|
+
def reshape(data: Any, shape: tuple[int, ...], module: Any) -> Any:
|
|
107
|
+
return data.reshape(shape)
|
|
108
|
+
|
|
109
|
+
@staticmethod
|
|
110
|
+
def broadcast_to(data: Any, shape: tuple[int, ...], module: Any) -> Any:
|
|
111
|
+
return module.broadcast_to(data, shape)
|
|
112
|
+
|
|
113
|
+
@staticmethod
|
|
114
|
+
def ones_like(reference: Any, shape: tuple[int, ...], dtype: Any, module: Any) -> Any:
|
|
115
|
+
return module.ones(shape, dtype=dtype)
|
|
116
|
+
|
|
117
|
+
|
|
118
|
+
class NumpyArrayOperations(ArrayOperations):
|
|
119
|
+
"""Numpy native operation leaves."""
|
|
120
|
+
|
|
121
|
+
@staticmethod
|
|
122
|
+
def to_numpy(data: Any, module: Any) -> Any:
|
|
123
|
+
del module
|
|
124
|
+
return data
|
|
125
|
+
|
|
126
|
+
@staticmethod
|
|
127
|
+
def from_numpy(data: Any, module: Any, device_id: int) -> Any:
|
|
128
|
+
del module, device_id
|
|
129
|
+
return data
|
|
130
|
+
|
|
131
|
+
@staticmethod
|
|
132
|
+
def scale_dtype(result: Any, target_dtype: Any, module: Any) -> Any:
|
|
133
|
+
if not hasattr(result, "dtype"):
|
|
134
|
+
return result
|
|
135
|
+
if not (
|
|
136
|
+
module.issubdtype(result.dtype, module.floating)
|
|
137
|
+
and module.issubdtype(target_dtype, module.integer)
|
|
138
|
+
):
|
|
139
|
+
return result.astype(target_dtype)
|
|
140
|
+
result_min = result.min()
|
|
141
|
+
result_max = result.max()
|
|
142
|
+
if result_max <= result_min:
|
|
143
|
+
return result.astype(target_dtype)
|
|
144
|
+
scaled = _scaled_values(result, result_min, result_max, target_dtype)
|
|
145
|
+
bounds = _clamp_bounds(target_dtype)
|
|
146
|
+
if bounds is not None:
|
|
147
|
+
scaled = module.clip(scaled, *bounds)
|
|
148
|
+
return scaled.astype(target_dtype)
|
|
149
|
+
|
|
150
|
+
|
|
151
|
+
class CupyArrayOperations(ArrayOperations):
|
|
152
|
+
"""Cupy native operation leaves."""
|
|
153
|
+
|
|
154
|
+
@staticmethod
|
|
155
|
+
def to_numpy(data: Any, module: Any) -> Any:
|
|
156
|
+
del module
|
|
157
|
+
return data.get()
|
|
158
|
+
|
|
159
|
+
@staticmethod
|
|
160
|
+
def from_numpy(data: Any, module: Any, device_id: int) -> Any:
|
|
161
|
+
del device_id
|
|
162
|
+
return module.array(data)
|
|
163
|
+
|
|
164
|
+
@staticmethod
|
|
165
|
+
def scale_dtype(result: Any, target_dtype: Any, module: Any) -> Any:
|
|
166
|
+
if not hasattr(result, "dtype"):
|
|
167
|
+
return result
|
|
168
|
+
if not (
|
|
169
|
+
module.issubdtype(result.dtype, module.floating)
|
|
170
|
+
and not module.issubdtype(target_dtype, module.floating)
|
|
171
|
+
):
|
|
172
|
+
return result.astype(target_dtype)
|
|
173
|
+
result_min = module.min(result)
|
|
174
|
+
result_max = module.max(result)
|
|
175
|
+
if result_max <= result_min:
|
|
176
|
+
return result.astype(target_dtype)
|
|
177
|
+
scaled = _scaled_values(result, result_min, result_max, target_dtype)
|
|
178
|
+
bounds = _clamp_bounds(target_dtype)
|
|
179
|
+
if bounds is not None:
|
|
180
|
+
scaled = module.clip(scaled, *bounds)
|
|
181
|
+
return scaled.astype(target_dtype)
|
|
182
|
+
|
|
183
|
+
|
|
184
|
+
class TorchArrayOperations(ArrayOperations):
|
|
185
|
+
"""Torch native operation leaves."""
|
|
186
|
+
|
|
187
|
+
@staticmethod
|
|
188
|
+
def to_numpy(data: Any, module: Any) -> Any:
|
|
189
|
+
del module
|
|
190
|
+
return data.cpu().numpy()
|
|
191
|
+
|
|
192
|
+
@staticmethod
|
|
193
|
+
def from_numpy(data: Any, module: Any, device_id: int) -> Any:
|
|
194
|
+
host_data = (
|
|
195
|
+
np.ascontiguousarray(data)
|
|
196
|
+
if any(stride < 0 for stride in getattr(data, "strides", ()))
|
|
197
|
+
else data
|
|
198
|
+
)
|
|
199
|
+
return module.from_numpy(host_data).to(f"cuda:{device_id}")
|
|
200
|
+
|
|
201
|
+
@staticmethod
|
|
202
|
+
def scale_dtype(result: Any, target_dtype: Any, module: Any) -> Any:
|
|
203
|
+
if not hasattr(result, "dtype"):
|
|
204
|
+
return result
|
|
205
|
+
mapped = _mapped_dtype(target_dtype, module)
|
|
206
|
+
floats = (module.float16, module.float32, module.float64)
|
|
207
|
+
if not (result.dtype in floats and np.issubdtype(np.dtype(target_dtype), np.integer)):
|
|
208
|
+
return result.to(mapped)
|
|
209
|
+
result_min = result.min()
|
|
210
|
+
result_max = result.max()
|
|
211
|
+
if result_max <= result_min:
|
|
212
|
+
return result.to(mapped)
|
|
213
|
+
scaled = _scaled_values(result, result_min, result_max, target_dtype)
|
|
214
|
+
bounds = _clamp_bounds(target_dtype)
|
|
215
|
+
if bounds is not None:
|
|
216
|
+
scaled = module.clamp(scaled, min=bounds[0], max=bounds[1])
|
|
217
|
+
return scaled.to(mapped)
|
|
218
|
+
|
|
219
|
+
@staticmethod
|
|
220
|
+
def stack(values: Sequence[Any], module: Any) -> Any:
|
|
221
|
+
return module.stack(tuple(values), dim=0)
|
|
222
|
+
|
|
223
|
+
@staticmethod
|
|
224
|
+
def cast(data: Any, dtype: Any, module: Any) -> Any:
|
|
225
|
+
return data.to(dtype=_mapped_dtype(dtype, module))
|
|
226
|
+
|
|
227
|
+
@staticmethod
|
|
228
|
+
def dtype_name(dtype: Any) -> str:
|
|
229
|
+
return str(dtype).rsplit(".", maxsplit=1)[-1]
|
|
230
|
+
|
|
231
|
+
@staticmethod
|
|
232
|
+
def broadcast_to(data: Any, shape: tuple[int, ...], module: Any) -> Any:
|
|
233
|
+
return data.expand(shape)
|
|
234
|
+
|
|
235
|
+
@staticmethod
|
|
236
|
+
def ones_like(reference: Any, shape: tuple[int, ...], dtype: Any, module: Any) -> Any:
|
|
237
|
+
return module.ones(shape, dtype=_mapped_dtype(dtype, module), device=reference.device)
|
|
238
|
+
|
|
239
|
+
|
|
240
|
+
class TensorflowArrayOperations(ArrayOperations):
|
|
241
|
+
"""Tensorflow native operation leaves."""
|
|
242
|
+
|
|
243
|
+
@staticmethod
|
|
244
|
+
def to_numpy(data: Any, module: Any) -> Any:
|
|
245
|
+
del module
|
|
246
|
+
return data.numpy()
|
|
247
|
+
|
|
248
|
+
@staticmethod
|
|
249
|
+
def from_numpy(data: Any, module: Any, device_id: int) -> Any:
|
|
250
|
+
del device_id
|
|
251
|
+
return module.convert_to_tensor(data)
|
|
252
|
+
|
|
253
|
+
@staticmethod
|
|
254
|
+
def scale_dtype(result: Any, target_dtype: Any, module: Any) -> Any:
|
|
255
|
+
if not hasattr(result, "dtype"):
|
|
256
|
+
return result
|
|
257
|
+
mapped = _mapped_dtype(target_dtype, module)
|
|
258
|
+
floats = (module.float16, module.float32, module.float64)
|
|
259
|
+
if not (result.dtype in floats and np.issubdtype(np.dtype(target_dtype), np.integer)):
|
|
260
|
+
return module.cast(result, mapped)
|
|
261
|
+
result_min = module.reduce_min(result)
|
|
262
|
+
result_max = module.reduce_max(result)
|
|
263
|
+
if result_max <= result_min:
|
|
264
|
+
return module.cast(result, mapped)
|
|
265
|
+
scaled = _scaled_values(result, result_min, result_max, target_dtype)
|
|
266
|
+
bounds = _clamp_bounds(target_dtype)
|
|
267
|
+
if bounds is not None:
|
|
268
|
+
scaled = module.clip_by_value(scaled, *bounds)
|
|
269
|
+
return module.cast(scaled, mapped)
|
|
270
|
+
|
|
271
|
+
@staticmethod
|
|
272
|
+
def cast(data: Any, dtype: Any, module: Any) -> Any:
|
|
273
|
+
return module.cast(data, _mapped_dtype(dtype, module))
|
|
274
|
+
|
|
275
|
+
@staticmethod
|
|
276
|
+
def reshape(data: Any, shape: tuple[int, ...], module: Any) -> Any:
|
|
277
|
+
return module.reshape(data, shape)
|
|
278
|
+
|
|
279
|
+
@staticmethod
|
|
280
|
+
def ones_like(reference: Any, shape: tuple[int, ...], dtype: Any, module: Any) -> Any:
|
|
281
|
+
with module.device(reference.device):
|
|
282
|
+
return module.ones(shape, dtype=_mapped_dtype(dtype, module))
|
|
283
|
+
|
|
284
|
+
|
|
285
|
+
class JaxArrayOperations(ArrayOperations):
|
|
286
|
+
"""Jax native operation leaves."""
|
|
287
|
+
|
|
288
|
+
@staticmethod
|
|
289
|
+
def to_numpy(data: Any, module: Any) -> Any:
|
|
290
|
+
del module
|
|
291
|
+
return np.asarray(data)
|
|
292
|
+
|
|
293
|
+
@staticmethod
|
|
294
|
+
def from_numpy(data: Any, module: Any, device_id: int) -> Any:
|
|
295
|
+
devices = tuple(device for device in module.devices() if device.platform == "gpu")
|
|
296
|
+
return module.device_put(data, devices[device_id])
|
|
297
|
+
|
|
298
|
+
@staticmethod
|
|
299
|
+
def scale_dtype(result: Any, target_dtype: Any, module: Any) -> Any:
|
|
300
|
+
if not hasattr(result, "dtype"):
|
|
301
|
+
return result
|
|
302
|
+
if np.dtype(target_dtype) == np.dtype(np.float64):
|
|
303
|
+
x64_enabled = getattr(module.config, "x64_enabled", None)
|
|
304
|
+
if x64_enabled is None:
|
|
305
|
+
x64_enabled = module.config.read("jax_enable_x64")
|
|
306
|
+
if not x64_enabled:
|
|
307
|
+
raise ValueError(
|
|
308
|
+
"JAX float64 output requires x64 mode; set JAX_ENABLE_X64=true before import"
|
|
309
|
+
)
|
|
310
|
+
jnp = module.numpy
|
|
311
|
+
mapped = _mapped_dtype(target_dtype, jnp)
|
|
312
|
+
floats = (jnp.float16, jnp.float32, jnp.float64)
|
|
313
|
+
if not (result.dtype in floats and np.issubdtype(np.dtype(target_dtype), np.integer)):
|
|
314
|
+
return result.astype(mapped)
|
|
315
|
+
result_min = jnp.min(result)
|
|
316
|
+
result_max = jnp.max(result)
|
|
317
|
+
if result_max <= result_min:
|
|
318
|
+
return result.astype(mapped)
|
|
319
|
+
scaled = _scaled_values(result, result_min, result_max, target_dtype)
|
|
320
|
+
bounds = _clamp_bounds(target_dtype)
|
|
321
|
+
if bounds is not None:
|
|
322
|
+
scaled = jnp.clip(scaled, *bounds)
|
|
323
|
+
return scaled.astype(mapped)
|
|
324
|
+
|
|
325
|
+
@staticmethod
|
|
326
|
+
def stack(values: Sequence[Any], module: Any) -> Any:
|
|
327
|
+
return module.numpy.stack(tuple(values), axis=0)
|
|
328
|
+
|
|
329
|
+
@staticmethod
|
|
330
|
+
def cast(data: Any, dtype: Any, module: Any) -> Any:
|
|
331
|
+
return data.astype(_mapped_dtype(dtype, module.numpy))
|
|
332
|
+
|
|
333
|
+
@staticmethod
|
|
334
|
+
def logical_and(left: Any, right: Any, module: Any) -> Any:
|
|
335
|
+
return module.numpy.logical_and(left, right)
|
|
336
|
+
|
|
337
|
+
@staticmethod
|
|
338
|
+
def broadcast_to(data: Any, shape: tuple[int, ...], module: Any) -> Any:
|
|
339
|
+
return module.numpy.broadcast_to(data, shape)
|
|
340
|
+
|
|
341
|
+
@staticmethod
|
|
342
|
+
def ones_like(reference: Any, shape: tuple[int, ...], dtype: Any, module: Any) -> Any:
|
|
343
|
+
return module.numpy.ones(shape, dtype=dtype, device=reference.device)
|
|
344
|
+
|
|
345
|
+
|
|
346
|
+
class PyclesperantoArrayOperations(ArrayOperations):
|
|
347
|
+
"""Pyclesperanto native operation leaves."""
|
|
348
|
+
|
|
349
|
+
@staticmethod
|
|
350
|
+
def to_numpy(data: Any, module: Any) -> Any:
|
|
351
|
+
return module.pull(data)
|
|
352
|
+
|
|
353
|
+
@staticmethod
|
|
354
|
+
def from_numpy(data: Any, module: Any, device_id: int) -> Any:
|
|
355
|
+
del device_id
|
|
356
|
+
return module.push(data)
|
|
357
|
+
|
|
358
|
+
@staticmethod
|
|
359
|
+
def scale_dtype(result: Any, target_dtype: Any, module: Any) -> Any:
|
|
360
|
+
if not hasattr(result, "dtype"):
|
|
361
|
+
return result
|
|
362
|
+
target_is_int = np.issubdtype(np.dtype(target_dtype), np.integer)
|
|
363
|
+
if not (np.issubdtype(result.dtype, np.floating) and target_is_int):
|
|
364
|
+
return module.push(module.pull(result).astype(target_dtype))
|
|
365
|
+
result_min = float(module.minimum_of_all_pixels(result))
|
|
366
|
+
result_max = float(module.maximum_of_all_pixels(result))
|
|
367
|
+
if result_max <= result_min:
|
|
368
|
+
return module.push(module.pull(result).astype(target_dtype))
|
|
369
|
+
normalized = module.subtract_image_from_scalar(result, scalar=result_min)
|
|
370
|
+
normalized = module.multiply_image_and_scalar(
|
|
371
|
+
normalized,
|
|
372
|
+
scalar=1.0 / (result_max - result_min),
|
|
373
|
+
)
|
|
374
|
+
range_info = _SCALING_RANGES.get(_dtype_name(target_dtype))
|
|
375
|
+
if isinstance(range_info, tuple):
|
|
376
|
+
scale, offset = range_info
|
|
377
|
+
scaled = module.multiply_image_and_scalar(normalized, scalar=scale)
|
|
378
|
+
scaled = module.subtract_image_from_scalar(scaled, scalar=offset)
|
|
379
|
+
elif range_info is not None:
|
|
380
|
+
scaled = module.multiply_image_and_scalar(normalized, scalar=range_info)
|
|
381
|
+
else:
|
|
382
|
+
scaled = normalized
|
|
383
|
+
host_values = module.pull(scaled)
|
|
384
|
+
bounds = _clamp_bounds(target_dtype)
|
|
385
|
+
if bounds is not None:
|
|
386
|
+
host_values = np.clip(host_values, *bounds)
|
|
387
|
+
return module.push(host_values.astype(target_dtype))
|
|
388
|
+
|
|
389
|
+
@staticmethod
|
|
390
|
+
def stack(values: Sequence[Any], module: Any) -> Any:
|
|
391
|
+
if not values:
|
|
392
|
+
raise ValueError("Cannot stack an empty pyclesperanto sequence")
|
|
393
|
+
if len(values) == 1:
|
|
394
|
+
source = values[0]
|
|
395
|
+
result = module.create((1, *source.shape), dtype=source.dtype)
|
|
396
|
+
return module.copy_slice(source, result, 0)
|
|
397
|
+
result = values[0]
|
|
398
|
+
for value in values[1:]:
|
|
399
|
+
result = module.concatenate_along_z(result, value)
|
|
178
400
|
return result
|
|
179
|
-
mapped = _mapped_dtype(target_dtype, module)
|
|
180
|
-
floats = (module.float16, module.float32, module.float64)
|
|
181
|
-
if not (result.dtype in floats and np.issubdtype(np.dtype(target_dtype), np.integer)):
|
|
182
|
-
return result.to(mapped)
|
|
183
|
-
result_min = result.min()
|
|
184
|
-
result_max = result.max()
|
|
185
|
-
if result_max <= result_min:
|
|
186
|
-
return result.to(mapped)
|
|
187
|
-
scaled = _scaled_values(result, result_min, result_max, target_dtype)
|
|
188
|
-
bounds = _clamp_bounds(target_dtype)
|
|
189
|
-
if bounds is not None:
|
|
190
|
-
scaled = module.clamp(scaled, min=bounds[0], max=bounds[1])
|
|
191
|
-
return scaled.to(mapped)
|
|
192
|
-
|
|
193
|
-
|
|
194
|
-
def _tensorflow_to_numpy(data: Any, module: Any) -> Any:
|
|
195
|
-
del module
|
|
196
|
-
return data.numpy()
|
|
197
|
-
|
|
198
|
-
|
|
199
|
-
def _tensorflow_from_numpy(data: Any, module: Any, device_id: int) -> Any:
|
|
200
|
-
del device_id
|
|
201
|
-
return module.convert_to_tensor(data)
|
|
202
401
|
|
|
203
|
-
|
|
204
|
-
def
|
|
205
|
-
|
|
206
|
-
|
|
207
|
-
|
|
208
|
-
def
|
|
209
|
-
|
|
210
|
-
|
|
211
|
-
|
|
212
|
-
def
|
|
213
|
-
|
|
214
|
-
|
|
215
|
-
|
|
216
|
-
|
|
217
|
-
|
|
218
|
-
|
|
219
|
-
|
|
220
|
-
|
|
221
|
-
|
|
222
|
-
|
|
223
|
-
|
|
224
|
-
|
|
225
|
-
|
|
226
|
-
|
|
227
|
-
|
|
228
|
-
|
|
229
|
-
|
|
230
|
-
|
|
231
|
-
|
|
232
|
-
|
|
233
|
-
|
|
234
|
-
|
|
235
|
-
|
|
236
|
-
|
|
237
|
-
|
|
238
|
-
|
|
239
|
-
|
|
240
|
-
def _jax_stack(values: Sequence[Any], module: Any) -> Any:
|
|
241
|
-
return module.numpy.stack(tuple(values), axis=0)
|
|
242
|
-
|
|
243
|
-
|
|
244
|
-
def _jax_cast(data: Any, dtype: Any, module: Any) -> Any:
|
|
245
|
-
return data.astype(_mapped_dtype(dtype, module.numpy))
|
|
246
|
-
|
|
247
|
-
|
|
248
|
-
def _jax_logical_and(left: Any, right: Any, module: Any) -> Any:
|
|
249
|
-
return module.numpy.logical_and(left, right)
|
|
250
|
-
|
|
251
|
-
|
|
252
|
-
def _jax_scale(result: Any, target_dtype: Any, module: Any) -> Any:
|
|
253
|
-
if not hasattr(result, "dtype"):
|
|
254
|
-
return result
|
|
255
|
-
if np.dtype(target_dtype) == np.dtype(np.float64):
|
|
256
|
-
x64_enabled = getattr(module.config, "x64_enabled", None)
|
|
257
|
-
if x64_enabled is None:
|
|
258
|
-
x64_enabled = module.config.read("jax_enable_x64")
|
|
259
|
-
if not x64_enabled:
|
|
260
|
-
raise ValueError(
|
|
261
|
-
"JAX float64 output requires x64 mode; set JAX_ENABLE_X64=true before import"
|
|
402
|
+
@staticmethod
|
|
403
|
+
def cast(data: Any, dtype: Any, module: Any) -> Any:
|
|
404
|
+
return module.push(module.pull(data).astype(dtype, copy=False))
|
|
405
|
+
|
|
406
|
+
@staticmethod
|
|
407
|
+
def logical_and(left: Any, right: Any, module: Any) -> Any:
|
|
408
|
+
return module.push(np.logical_and(module.pull(left), module.pull(right)))
|
|
409
|
+
|
|
410
|
+
@staticmethod
|
|
411
|
+
def reshape(data: Any, shape: tuple[int, ...], module: Any) -> Any:
|
|
412
|
+
if tuple(data.shape) == shape:
|
|
413
|
+
return data
|
|
414
|
+
raise NotImplementedError(
|
|
415
|
+
"pyclesperanto does not provide native array reshape; its reshape downloads pixels"
|
|
416
|
+
)
|
|
417
|
+
|
|
418
|
+
@staticmethod
|
|
419
|
+
def broadcast_to(data: Any, shape: tuple[int, ...], module: Any) -> Any:
|
|
420
|
+
if tuple(data.shape) == shape:
|
|
421
|
+
return data
|
|
422
|
+
raise NotImplementedError("pyclesperanto does not provide native array broadcasting")
|
|
423
|
+
|
|
424
|
+
@staticmethod
|
|
425
|
+
def ones_like(reference: Any, shape: tuple[int, ...], dtype: Any, module: Any) -> Any:
|
|
426
|
+
if np.dtype(dtype).name not in {
|
|
427
|
+
"float32",
|
|
428
|
+
"int8",
|
|
429
|
+
"int16",
|
|
430
|
+
"int32",
|
|
431
|
+
"uint8",
|
|
432
|
+
"uint16",
|
|
433
|
+
"uint32",
|
|
434
|
+
}:
|
|
435
|
+
raise NotImplementedError(f"pyclesperanto cannot allocate native dtype {dtype!r}")
|
|
436
|
+
if not 1 <= len(shape) <= 3:
|
|
437
|
+
raise NotImplementedError(
|
|
438
|
+
"pyclesperanto native allocation requires one to three dimensions"
|
|
262
439
|
)
|
|
263
|
-
|
|
264
|
-
|
|
265
|
-
|
|
266
|
-
if not (result.dtype in floats and np.issubdtype(np.dtype(target_dtype), np.integer)):
|
|
267
|
-
return result.astype(mapped)
|
|
268
|
-
result_min = jnp.min(result)
|
|
269
|
-
result_max = jnp.max(result)
|
|
270
|
-
if result_max <= result_min:
|
|
271
|
-
return result.astype(mapped)
|
|
272
|
-
scaled = _scaled_values(result, result_min, result_max, target_dtype)
|
|
273
|
-
bounds = _clamp_bounds(target_dtype)
|
|
274
|
-
if bounds is not None:
|
|
275
|
-
scaled = jnp.clip(scaled, *bounds)
|
|
276
|
-
return scaled.astype(mapped)
|
|
277
|
-
|
|
278
|
-
|
|
279
|
-
def _pyclesperanto_to_numpy(data: Any, module: Any) -> Any:
|
|
280
|
-
return module.pull(data)
|
|
281
|
-
|
|
282
|
-
|
|
283
|
-
def _pyclesperanto_from_numpy(data: Any, module: Any, device_id: int) -> Any:
|
|
284
|
-
del device_id
|
|
285
|
-
return module.push(data)
|
|
286
|
-
|
|
287
|
-
|
|
288
|
-
def _pyclesperanto_stack(values: Sequence[Any], module: Any) -> Any:
|
|
289
|
-
if not values:
|
|
290
|
-
raise ValueError("Cannot stack an empty pyclesperanto sequence")
|
|
291
|
-
if len(values) == 1:
|
|
292
|
-
source = values[0]
|
|
293
|
-
result = module.create((1, *source.shape), dtype=source.dtype)
|
|
294
|
-
return module.copy_slice(source, result, 0)
|
|
295
|
-
result = values[0]
|
|
296
|
-
for value in values[1:]:
|
|
297
|
-
result = module.concatenate_along_z(result, value)
|
|
298
|
-
return result
|
|
299
|
-
|
|
300
|
-
|
|
301
|
-
def _pyclesperanto_cast(data: Any, dtype: Any, module: Any) -> Any:
|
|
302
|
-
return module.push(module.pull(data).astype(dtype, copy=False))
|
|
303
|
-
|
|
304
|
-
|
|
305
|
-
def _pyclesperanto_logical_and(left: Any, right: Any, module: Any) -> Any:
|
|
306
|
-
return module.push(np.logical_and(module.pull(left), module.pull(right)))
|
|
440
|
+
result = module.create(shape, dtype=dtype, device=reference.device)
|
|
441
|
+
module.set(result, scalar=1, device=reference.device)
|
|
442
|
+
return result
|
|
307
443
|
|
|
308
444
|
|
|
309
|
-
|
|
310
|
-
|
|
311
|
-
|
|
312
|
-
|
|
313
|
-
|
|
314
|
-
|
|
315
|
-
result_min = float(module.minimum_of_all_pixels(result))
|
|
316
|
-
result_max = float(module.maximum_of_all_pixels(result))
|
|
317
|
-
if result_max <= result_min:
|
|
318
|
-
return module.push(module.pull(result).astype(target_dtype))
|
|
319
|
-
normalized = module.subtract_image_from_scalar(result, scalar=result_min)
|
|
320
|
-
normalized = module.multiply_image_and_scalar(
|
|
321
|
-
normalized,
|
|
322
|
-
scalar=1.0 / (result_max - result_min),
|
|
323
|
-
)
|
|
324
|
-
range_info = _SCALING_RANGES.get(_dtype_name(target_dtype))
|
|
325
|
-
if isinstance(range_info, tuple):
|
|
326
|
-
scale, offset = range_info
|
|
327
|
-
scaled = module.multiply_image_and_scalar(normalized, scalar=scale)
|
|
328
|
-
scaled = module.subtract_image_from_scalar(scaled, scalar=offset)
|
|
329
|
-
elif range_info is not None:
|
|
330
|
-
scaled = module.multiply_image_and_scalar(normalized, scalar=range_info)
|
|
331
|
-
else:
|
|
332
|
-
scaled = normalized
|
|
333
|
-
host_values = module.pull(scaled)
|
|
334
|
-
bounds = _clamp_bounds(target_dtype)
|
|
335
|
-
if bounds is not None:
|
|
336
|
-
host_values = np.clip(host_values, *bounds)
|
|
337
|
-
return module.push(host_values.astype(target_dtype))
|
|
338
|
-
|
|
339
|
-
|
|
340
|
-
@dataclass(frozen=True, slots=True)
|
|
341
|
-
class ArrayOperations:
|
|
342
|
-
"""Framework-specific array leaves referenced by one declaration."""
|
|
343
|
-
|
|
344
|
-
to_numpy: ToNumpy
|
|
345
|
-
from_numpy: FromNumpy
|
|
346
|
-
stack: StackArrays
|
|
347
|
-
scale_dtype: ScaleDtype
|
|
348
|
-
dtype_name: DtypeName = _numpy_dtype_name
|
|
349
|
-
cast: CastArray = _numpy_cast
|
|
350
|
-
logical_and: LogicalAnd = _module_logical_and
|
|
351
|
-
|
|
352
|
-
|
|
353
|
-
NUMPY_OPERATIONS = ArrayOperations(
|
|
354
|
-
to_numpy=_identity_to_numpy,
|
|
355
|
-
from_numpy=_identity_from_numpy,
|
|
356
|
-
stack=_numpy_stack,
|
|
357
|
-
scale_dtype=_numpy_scale,
|
|
358
|
-
)
|
|
359
|
-
CUPY_OPERATIONS = ArrayOperations(
|
|
360
|
-
to_numpy=_cupy_to_numpy,
|
|
361
|
-
from_numpy=_cupy_from_numpy,
|
|
362
|
-
stack=_cupy_stack,
|
|
363
|
-
scale_dtype=_cupy_scale,
|
|
364
|
-
)
|
|
365
|
-
TORCH_OPERATIONS = ArrayOperations(
|
|
366
|
-
to_numpy=_torch_to_numpy,
|
|
367
|
-
from_numpy=_torch_from_numpy,
|
|
368
|
-
stack=_torch_stack,
|
|
369
|
-
scale_dtype=_torch_scale,
|
|
370
|
-
dtype_name=_torch_dtype_name,
|
|
371
|
-
cast=_torch_cast,
|
|
372
|
-
)
|
|
373
|
-
TENSORFLOW_OPERATIONS = ArrayOperations(
|
|
374
|
-
to_numpy=_tensorflow_to_numpy,
|
|
375
|
-
from_numpy=_tensorflow_from_numpy,
|
|
376
|
-
stack=_tensorflow_stack,
|
|
377
|
-
scale_dtype=_tensorflow_scale,
|
|
378
|
-
dtype_name=_tensorflow_dtype_name,
|
|
379
|
-
cast=_tensorflow_cast,
|
|
380
|
-
)
|
|
381
|
-
JAX_OPERATIONS = ArrayOperations(
|
|
382
|
-
to_numpy=_jax_to_numpy,
|
|
383
|
-
from_numpy=_jax_from_numpy,
|
|
384
|
-
stack=_jax_stack,
|
|
385
|
-
scale_dtype=_jax_scale,
|
|
386
|
-
cast=_jax_cast,
|
|
387
|
-
logical_and=_jax_logical_and,
|
|
388
|
-
)
|
|
389
|
-
PYCLESPERANTO_OPERATIONS = ArrayOperations(
|
|
390
|
-
to_numpy=_pyclesperanto_to_numpy,
|
|
391
|
-
from_numpy=_pyclesperanto_from_numpy,
|
|
392
|
-
stack=_pyclesperanto_stack,
|
|
393
|
-
scale_dtype=_pyclesperanto_scale,
|
|
394
|
-
cast=_pyclesperanto_cast,
|
|
395
|
-
logical_and=_pyclesperanto_logical_and,
|
|
396
|
-
)
|
|
445
|
+
NUMPY_OPERATIONS = NumpyArrayOperations()
|
|
446
|
+
CUPY_OPERATIONS = CupyArrayOperations()
|
|
447
|
+
TORCH_OPERATIONS = TorchArrayOperations()
|
|
448
|
+
TENSORFLOW_OPERATIONS = TensorflowArrayOperations()
|
|
449
|
+
JAX_OPERATIONS = JaxArrayOperations()
|
|
450
|
+
PYCLESPERANTO_OPERATIONS = PyclesperantoArrayOperations()
|
arraybridge/decorators.py
CHANGED
|
@@ -82,11 +82,13 @@ class DtypeConversionConfig(RuntimeParameterDeclarationABC):
|
|
|
82
82
|
return DtypeConversionConfig
|
|
83
83
|
|
|
84
84
|
@classmethod
|
|
85
|
-
def parameter(
|
|
85
|
+
def parameter(
|
|
86
|
+
cls, *, default_value: "DtypeConversionConfig | None" = None
|
|
87
|
+
) -> inspect.Parameter:
|
|
86
88
|
return inspect.Parameter(
|
|
87
89
|
cls.require_parameter_name(),
|
|
88
90
|
inspect.Parameter.KEYWORD_ONLY,
|
|
89
|
-
default=cls.default_value(),
|
|
91
|
+
default=cls.default_value() if default_value is None else default_value,
|
|
90
92
|
annotation=cls.annotation_type(),
|
|
91
93
|
)
|
|
92
94
|
|
|
@@ -264,12 +266,30 @@ class KeywordOnlySignatureExtension:
|
|
|
264
266
|
return len(parameters)
|
|
265
267
|
|
|
266
268
|
|
|
267
|
-
|
|
268
|
-
|
|
269
|
+
class ThreadGPUContextStorage(threading.local):
|
|
270
|
+
"""Declare the existing native thread-local context slot and its type."""
|
|
271
|
+
|
|
272
|
+
context: "ThreadGPUContext | None" = None
|
|
269
273
|
|
|
270
274
|
|
|
271
275
|
class ThreadGPUContext:
|
|
272
|
-
"""
|
|
276
|
+
"""Runtime-owned thread-local streams keyed by framework/device identity.
|
|
277
|
+
|
|
278
|
+
Keep the runtime handle on its importable owner, not in decorator function
|
|
279
|
+
globals. A retained, unpublished decorated callable is serialized by value;
|
|
280
|
+
its durable closure must not pull a ``threading.local`` into history.
|
|
281
|
+
"""
|
|
282
|
+
|
|
283
|
+
_contexts: ClassVar[ThreadGPUContextStorage] = ThreadGPUContextStorage()
|
|
284
|
+
|
|
285
|
+
@classmethod
|
|
286
|
+
def current(cls) -> "ThreadGPUContext":
|
|
287
|
+
"""Return this thread's runtime context without serializing its handle."""
|
|
288
|
+
context = cls._contexts.context
|
|
289
|
+
if context is None:
|
|
290
|
+
context = cls()
|
|
291
|
+
cls._contexts.context = context
|
|
292
|
+
return context
|
|
273
293
|
|
|
274
294
|
def __init__(self):
|
|
275
295
|
self._streams: dict[tuple[MemoryType, int], Any] = {}
|
|
@@ -301,13 +321,6 @@ class ThreadGPUContext:
|
|
|
301
321
|
return device_id, self._streams[key]
|
|
302
322
|
|
|
303
323
|
|
|
304
|
-
def _get_thread_gpu_context():
|
|
305
|
-
"""Get or create thread-local GPU context."""
|
|
306
|
-
if not hasattr(_thread_gpu_contexts, "context"):
|
|
307
|
-
_thread_gpu_contexts.context = ThreadGPUContext()
|
|
308
|
-
return _thread_gpu_contexts.context
|
|
309
|
-
|
|
310
|
-
|
|
311
324
|
def memory_types(
|
|
312
325
|
input_type: str | MemoryType,
|
|
313
326
|
output_type: str | MemoryType,
|
|
@@ -352,6 +365,7 @@ def wrap_dtype_preserving_callable(
|
|
|
352
365
|
mem_type: MemoryType,
|
|
353
366
|
*,
|
|
354
367
|
slice_by_slice_default: bool = False,
|
|
368
|
+
dtype_config_default: DtypeConversionConfig | None = None,
|
|
355
369
|
):
|
|
356
370
|
"""
|
|
357
371
|
Return a callable with ArrayBridge dtype and slice controls.
|
|
@@ -363,18 +377,25 @@ def wrap_dtype_preserving_callable(
|
|
|
363
377
|
input_memory_type = MemoryType(MemoryContractAttribute.INPUT.read(func, mem_type.value))
|
|
364
378
|
output_memory_type = MemoryType(MemoryContractAttribute.OUTPUT.read(func, mem_type.value))
|
|
365
379
|
scale_func = output_memory_type.scale_dtype
|
|
380
|
+
default_dtype_config = (
|
|
381
|
+
DtypeConversionConfig.default_value()
|
|
382
|
+
if dtype_config_default is None
|
|
383
|
+
else dtype_config_default
|
|
384
|
+
)
|
|
385
|
+
if not isinstance(default_dtype_config, DtypeConversionConfig):
|
|
386
|
+
raise TypeError("Callable dtype default must be a DtypeConversionConfig.")
|
|
366
387
|
|
|
367
388
|
@functools.wraps(func)
|
|
368
389
|
def dtype_wrapper(image, *args, **kwargs):
|
|
369
390
|
# Pipeline runtimes may inject dtype_config; direct calls use the same
|
|
370
|
-
# preserve-input
|
|
391
|
+
# callable-owned default explicitly (preserve-input when undeclared).
|
|
371
392
|
slice_by_slice = kwargs.pop(
|
|
372
393
|
SliceBySliceRuntimeParameter.require_parameter_name(),
|
|
373
394
|
slice_by_slice_default,
|
|
374
395
|
)
|
|
375
396
|
dtype_config: DtypeConversionConfig = kwargs.pop(
|
|
376
397
|
DtypeConversionConfig.require_parameter_name(),
|
|
377
|
-
|
|
398
|
+
default_dtype_config,
|
|
378
399
|
)
|
|
379
400
|
dtype_conversion = dtype_config.default_dtype_conversion
|
|
380
401
|
|
|
@@ -424,7 +445,7 @@ def wrap_dtype_preserving_callable(
|
|
|
424
445
|
)
|
|
425
446
|
)
|
|
426
447
|
dtype_signature = KeywordOnlySignatureExtension(dtype_signature).with_parameter(
|
|
427
|
-
DtypeConversionConfig.parameter()
|
|
448
|
+
DtypeConversionConfig.parameter(default_value=default_dtype_config)
|
|
428
449
|
)
|
|
429
450
|
setattr(dtype_wrapper, "__signature__", dtype_signature)
|
|
430
451
|
|
|
@@ -461,7 +482,7 @@ def _create_gpu_wrapper(func, mem_type: MemoryType, oom_recovery: bool):
|
|
|
461
482
|
# Check if GPU is available for this framework
|
|
462
483
|
if framework is not None and mem_type.available_device_ids(framework):
|
|
463
484
|
# Get thread-local context
|
|
464
|
-
ctx =
|
|
485
|
+
ctx = ThreadGPUContext.current()
|
|
465
486
|
|
|
466
487
|
device_id, stream = ctx.stream_for(mem_type, framework)
|
|
467
488
|
|
|
@@ -510,6 +531,7 @@ def _create_memory_decorator(mem_type: MemoryType):
|
|
|
510
531
|
oom_recovery=True,
|
|
511
532
|
contract=None,
|
|
512
533
|
slice_by_slice_default=False,
|
|
534
|
+
dtype_config_default: DtypeConversionConfig | None = None,
|
|
513
535
|
):
|
|
514
536
|
"""
|
|
515
537
|
Decorator for {mem_type} memory type functions.
|
|
@@ -521,6 +543,8 @@ def _create_memory_decorator(mem_type: MemoryType):
|
|
|
521
543
|
oom_recovery: Enable automatic OOM recovery (default: True)
|
|
522
544
|
contract: Optional validation function for outputs
|
|
523
545
|
slice_by_slice_default: Default for the decorator-owned slice control
|
|
546
|
+
dtype_config_default: Callable-owned dtype policy for direct calls;
|
|
547
|
+
explicit runtime dtype_config still overrides this default.
|
|
524
548
|
|
|
525
549
|
Returns:
|
|
526
550
|
Decorated function with memory type metadata and dtype preservation
|
|
@@ -538,6 +562,7 @@ def _create_memory_decorator(mem_type: MemoryType):
|
|
|
538
562
|
func,
|
|
539
563
|
mem_type,
|
|
540
564
|
slice_by_slice_default=slice_by_slice_default,
|
|
565
|
+
dtype_config_default=dtype_config_default,
|
|
541
566
|
)
|
|
542
567
|
|
|
543
568
|
# Apply GPU wrapper if this is a GPU memory type
|
arraybridge/types.py
CHANGED
|
@@ -12,7 +12,7 @@ import importlib.util
|
|
|
12
12
|
import logging
|
|
13
13
|
import os
|
|
14
14
|
import sys
|
|
15
|
-
from collections.abc import Callable, Iterator, Mapping, MutableMapping
|
|
15
|
+
from collections.abc import Callable, Iterator, Mapping, MutableMapping, Sequence
|
|
16
16
|
from contextlib import AbstractContextManager, contextmanager, nullcontext
|
|
17
17
|
from dataclasses import dataclass
|
|
18
18
|
from enum import Enum
|
|
@@ -834,11 +834,44 @@ class MemoryType(_MemoryTypeFields, Enum):
|
|
|
834
834
|
"""Cast one array through this framework member's operation leaf."""
|
|
835
835
|
|
|
836
836
|
framework = module if module is not None else self.import_module()
|
|
837
|
-
|
|
838
|
-
scope = nullcontext() if device_id is None else self.device_scope(device_id, framework)
|
|
839
|
-
with scope:
|
|
837
|
+
with self._array_device_scope(data, framework):
|
|
840
838
|
return self._operations.cast(data, dtype, framework)
|
|
841
839
|
|
|
840
|
+
def _array_device_scope(self, reference: Any, module: Any) -> AbstractContextManager[None]:
|
|
841
|
+
"""Derive every native operation scope from its actual input array."""
|
|
842
|
+
device_id = self.device_id_of(reference, module)
|
|
843
|
+
return nullcontext() if device_id is None else self.device_scope(device_id, module)
|
|
844
|
+
|
|
845
|
+
def reshape(self, data: Any, shape: Sequence[int], module: Any | None = None) -> Any:
|
|
846
|
+
"""Reshape on the input device without an implicit host projection."""
|
|
847
|
+
framework = module if module is not None else self.import_module()
|
|
848
|
+
with self._array_device_scope(data, framework):
|
|
849
|
+
return self._operations.reshape(data, tuple(shape), framework)
|
|
850
|
+
|
|
851
|
+
def broadcast_to(self, data: Any, shape: Sequence[int], module: Any | None = None) -> Any:
|
|
852
|
+
"""Broadcast on the input device without an implicit host projection."""
|
|
853
|
+
framework = module if module is not None else self.import_module()
|
|
854
|
+
with self._array_device_scope(data, framework):
|
|
855
|
+
return self._operations.broadcast_to(data, tuple(shape), framework)
|
|
856
|
+
|
|
857
|
+
def ones_like(
|
|
858
|
+
self,
|
|
859
|
+
reference: Any,
|
|
860
|
+
*,
|
|
861
|
+
shape: Sequence[int] | None = None,
|
|
862
|
+
dtype: Any = bool,
|
|
863
|
+
module: Any | None = None,
|
|
864
|
+
) -> Any:
|
|
865
|
+
"""Allocate ones on the supplied reference's framework-local device."""
|
|
866
|
+
framework = module if module is not None else self.import_module()
|
|
867
|
+
with self._array_device_scope(reference, framework):
|
|
868
|
+
return self._operations.ones_like(
|
|
869
|
+
reference,
|
|
870
|
+
tuple(reference.shape if shape is None else shape),
|
|
871
|
+
dtype,
|
|
872
|
+
framework,
|
|
873
|
+
)
|
|
874
|
+
|
|
842
875
|
def logical_and(
|
|
843
876
|
self,
|
|
844
877
|
left: Any,
|
|
@@ -848,9 +881,7 @@ class MemoryType(_MemoryTypeFields, Enum):
|
|
|
848
881
|
"""Intersect two arrays through this framework member's operation leaf."""
|
|
849
882
|
|
|
850
883
|
framework = module if module is not None else self.import_module()
|
|
851
|
-
|
|
852
|
-
scope = nullcontext() if device_id is None else self.device_scope(device_id, framework)
|
|
853
|
-
with scope:
|
|
884
|
+
with self._array_device_scope(left, framework):
|
|
854
885
|
return self._operations.logical_and(left, right, framework)
|
|
855
886
|
|
|
856
887
|
def available_device_ids(self, module: Any | None = None) -> tuple[int, ...]:
|
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: arraybridge
|
|
3
|
-
Version: 0.3.
|
|
3
|
+
Version: 0.3.7
|
|
4
4
|
Summary: Unified API for NumPy, CuPy, PyTorch, TensorFlow, JAX, and pyclesperanto with automatic memory type conversion
|
|
5
5
|
Project-URL: Homepage, https://github.com/OpenHCSDev/arraybridge
|
|
6
6
|
Project-URL: Documentation, https://arraybridge.readthedocs.io
|
|
@@ -36,6 +36,7 @@ Requires-Dist: cupy>=10.0; extra == 'cupy'
|
|
|
36
36
|
Provides-Extra: dev
|
|
37
37
|
Requires-Dist: black>=23.0; extra == 'dev'
|
|
38
38
|
Requires-Dist: build>=1.0; extra == 'dev'
|
|
39
|
+
Requires-Dist: dill>=0.3.8; extra == 'dev'
|
|
39
40
|
Requires-Dist: mypy>=1.0; extra == 'dev'
|
|
40
41
|
Requires-Dist: packaging>=23.0; extra == 'dev'
|
|
41
42
|
Requires-Dist: pytest-cov>=4.0; extra == 'dev'
|
|
@@ -142,3 +143,17 @@ pip install "arraybridge[cupy]"
|
|
|
142
143
|
```
|
|
143
144
|
|
|
144
145
|
Documentation: <https://arraybridge.readthedocs.io>
|
|
146
|
+
|
|
147
|
+
### Native array geometry
|
|
148
|
+
|
|
149
|
+
`MemoryType.reshape(array, shape)` and `MemoryType.broadcast_to(array, shape)`
|
|
150
|
+
operate in the array's framework and device. `MemoryType.ones_like(reference,
|
|
151
|
+
shape=..., dtype=...)` allocates on the reference's device; omitting `shape`
|
|
152
|
+
retains its shape. Reshape and broadcast preserve views where the backend can
|
|
153
|
+
represent them. Broadcasting is not a promise of a writable, independent buffer.
|
|
154
|
+
|
|
155
|
+
These operations never silently project pixels to NumPy. Pyclesperanto cannot
|
|
156
|
+
provide native changed-shape reshape or broadcasting, so those requests raise
|
|
157
|
+
`NotImplementedError`. Its native ones allocation supports one to three
|
|
158
|
+
dimensions and its supported integer/float32 dtypes; boolean allocation is
|
|
159
|
+
explicitly unsupported. Existing conversion and disk boundaries remain separate.
|
|
@@ -1,10 +1,10 @@
|
|
|
1
|
-
arraybridge/__init__.py,sha256=
|
|
1
|
+
arraybridge/__init__.py,sha256=lmu9luoyMe9tpvuDlxSvwsYhRcFKkLKP7Smm93RndS8,3184
|
|
2
2
|
arraybridge/array_geometry.py,sha256=CopfGmFVKwV5l-lOTImhm1YVh18weUng0ReEF58qlPI,1689
|
|
3
|
-
arraybridge/array_operations.py,sha256=
|
|
3
|
+
arraybridge/array_operations.py,sha256=3cT17uSpq9tQPj_HkxHfR8_q6C8tzygI2HiM0yuAaAk,16381
|
|
4
4
|
arraybridge/array_payload.py,sha256=6_jJwPRya0MHrUNSA9kXPlTISpLa0G74rqDValGXKx4,1079
|
|
5
5
|
arraybridge/converters.py,sha256=q44E6Pg1rJz8yq6WQIPLOZPv1lqqcUcEG4Gh4D5M1_o,2649
|
|
6
6
|
arraybridge/converters_registry.py,sha256=_hHtbd-3SMIKUw8fvNu3twPWJY7qHyNyQ_NhVp3eSHQ,5647
|
|
7
|
-
arraybridge/decorators.py,sha256=
|
|
7
|
+
arraybridge/decorators.py,sha256=KPm_R3C9TMHvNARf3vKAZQY-3gKSlvYrimxBQMyI5A4,21271
|
|
8
8
|
arraybridge/dtype_scaling.py,sha256=KDgYyCvxgsQXLJJl_bs5nH46zQimsA7c1Olv6xJ2z1o,831
|
|
9
9
|
arraybridge/exceptions.py,sha256=fhej5ZS39QcBFWzNeNYjoqeLVnOgwwp-cWIme25jyuo,752
|
|
10
10
|
arraybridge/framework_config.py,sha256=5yTTt9NWcwzKDMXbc-fqCdiFl-piKOt3SUnxmRXXCQk,744
|
|
@@ -13,9 +13,9 @@ arraybridge/gpu_cleanup.py,sha256=ADI4wMLk9GIO9rJ6BJuNk3uZIyANoAqJ7VFTnLJzrbE,25
|
|
|
13
13
|
arraybridge/oom_recovery.py,sha256=hGvOuJHVSopiIqblzoMFlqW8_oK6ZlAGF_l4pfEljeI,2403
|
|
14
14
|
arraybridge/slice_processing.py,sha256=5R10OmHW8BSbJFwT1URd2Sz3P15o5EhaI2VnVUUIvkQ,4450
|
|
15
15
|
arraybridge/stack_utils.py,sha256=FVikN0XG-Xryw-1p9MaSRUTZet7kLcy0gTEERSjc9r4,6915
|
|
16
|
-
arraybridge/types.py,sha256=
|
|
16
|
+
arraybridge/types.py,sha256=_RrqtA2q4ocy1huxLIDiDhoZySCQpNhorV5_7HY0S-Y,37287
|
|
17
17
|
arraybridge/utils.py,sha256=RqThtYScPEOEQG491dA8b_VvUTjgcFhwLbh0erQu5Ns,4833
|
|
18
|
-
arraybridge-0.3.
|
|
19
|
-
arraybridge-0.3.
|
|
20
|
-
arraybridge-0.3.
|
|
21
|
-
arraybridge-0.3.
|
|
18
|
+
arraybridge-0.3.7.dist-info/METADATA,sha256=Rpr8OMPyw7t9BXgM3EWf2jAm03vxUdVIPeGN3L29FmU,6124
|
|
19
|
+
arraybridge-0.3.7.dist-info/WHEEL,sha256=lCkmxWfQsSc9CfIClYeavTdQeEX2toPqufh9gI35EQA,87
|
|
20
|
+
arraybridge-0.3.7.dist-info/licenses/LICENSE,sha256=xagEoeTAj1WT64RmyR3E6HH-eTGdgXN6gqPMUUt7L_Y,1070
|
|
21
|
+
arraybridge-0.3.7.dist-info/RECORD,,
|
|
File without changes
|
|
File without changes
|