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 CHANGED
@@ -10,7 +10,7 @@ attribute access instead of at import time. Declaration-only consumers
10
10
  NumPy/numcodecs import cost.
11
11
  """
12
12
 
13
- __version__ = "0.3.4"
13
+ __version__ = "0.3.7"
14
14
 
15
15
  _LAZY_EXPORTS: dict[str, str] = {
16
16
  "MemoryType": ".types",
@@ -1,19 +1,11 @@
1
- """Typed array-operation leaves carried by ``MemoryType`` declarations."""
1
+ """Native array operations carried by MemoryType declarations."""
2
2
 
3
- from collections.abc import Callable, Sequence
4
- from dataclasses import dataclass
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
- def _torch_scale(result: Any, target_dtype: Any, module: Any) -> Any:
177
- if not hasattr(result, "dtype"):
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 _tensorflow_stack(values: Sequence[Any], module: Any) -> Any:
205
- return module.stack(tuple(values), axis=0)
206
-
207
-
208
- def _tensorflow_cast(data: Any, dtype: Any, module: Any) -> Any:
209
- return module.cast(data, _mapped_dtype(dtype, module))
210
-
211
-
212
- def _tensorflow_scale(result: Any, target_dtype: Any, module: Any) -> Any:
213
- if not hasattr(result, "dtype"):
214
- return result
215
- mapped = _mapped_dtype(target_dtype, module)
216
- floats = (module.float16, module.float32, module.float64)
217
- if not (result.dtype in floats and np.issubdtype(np.dtype(target_dtype), np.integer)):
218
- return module.cast(result, mapped)
219
- result_min = module.reduce_min(result)
220
- result_max = module.reduce_max(result)
221
- if result_max <= result_min:
222
- return module.cast(result, mapped)
223
- scaled = _scaled_values(result, result_min, result_max, target_dtype)
224
- bounds = _clamp_bounds(target_dtype)
225
- if bounds is not None:
226
- scaled = module.clip_by_value(scaled, *bounds)
227
- return module.cast(scaled, mapped)
228
-
229
-
230
- def _jax_to_numpy(data: Any, module: Any) -> Any:
231
- del module
232
- return np.asarray(data)
233
-
234
-
235
- def _jax_from_numpy(data: Any, module: Any, device_id: int) -> Any:
236
- devices = tuple(device for device in module.devices() if device.platform == "gpu")
237
- return module.device_put(data, devices[device_id])
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
- jnp = module.numpy
264
- mapped = _mapped_dtype(target_dtype, jnp)
265
- floats = (jnp.float16, jnp.float32, jnp.float64)
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
- def _pyclesperanto_scale(result: Any, target_dtype: Any, module: Any) -> Any:
310
- if not hasattr(result, "dtype"):
311
- return result
312
- target_is_int = np.issubdtype(np.dtype(target_dtype), np.integer)
313
- if not (np.issubdtype(result.dtype, np.floating) and target_is_int):
314
- return module.push(module.pull(result).astype(target_dtype))
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(cls) -> inspect.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
- # Thread-local storage for GPU streams and contexts
268
- _thread_gpu_contexts = threading.local()
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
- """Thread-local streams keyed by framework-local device identity."""
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 default explicitly.
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
- DtypeConversionConfig.default_value(),
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 = _get_thread_gpu_context()
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
- device_id = self.device_id_of(data, framework)
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
- device_id = self.device_id_of(left, framework)
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.4
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=a27ES1ItSBsT77vP3f12ILNu0MzrvQfcpCBVDXwi-Sg,3184
1
+ arraybridge/__init__.py,sha256=lmu9luoyMe9tpvuDlxSvwsYhRcFKkLKP7Smm93RndS8,3184
2
2
  arraybridge/array_geometry.py,sha256=CopfGmFVKwV5l-lOTImhm1YVh18weUng0ReEF58qlPI,1689
3
- arraybridge/array_operations.py,sha256=fp69o9BYcCzfhmLcWrH68SVGSfPL3nb6Cty-Amd_oQM,13082
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=mYIxBbCuWQVN8dhAE2M7i2Yo0KbcMyY4MZmsfGXGm-0,19975
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=e5R2DVflQAjXrfTk033KRdiiwoip3IQAvgDCYUjPdiQ,35807
16
+ arraybridge/types.py,sha256=_RrqtA2q4ocy1huxLIDiDhoZySCQpNhorV5_7HY0S-Y,37287
17
17
  arraybridge/utils.py,sha256=RqThtYScPEOEQG491dA8b_VvUTjgcFhwLbh0erQu5Ns,4833
18
- arraybridge-0.3.4.dist-info/METADATA,sha256=jq9-ULcNwbZo7cMtZlhhqj4F2QlV-8PrpdM1Tv0tLCE,5275
19
- arraybridge-0.3.4.dist-info/WHEEL,sha256=lCkmxWfQsSc9CfIClYeavTdQeEX2toPqufh9gI35EQA,87
20
- arraybridge-0.3.4.dist-info/licenses/LICENSE,sha256=xagEoeTAj1WT64RmyR3E6HH-eTGdgXN6gqPMUUt7L_Y,1070
21
- arraybridge-0.3.4.dist-info/RECORD,,
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,,