diff --git a/.cursor/rules/git-commit.mdc b/.cursor/rules/git-commit.mdc
new file mode 100644
index 0000000..16ca1f8
--- /dev/null
+++ b/.cursor/rules/git-commit.mdc
@@ -0,0 +1,11 @@
+---
+description: Git commit message preferences for this project
+alwaysApply: true
+---
+
+# Git Commits
+
+When creating git commits in this project:
+
+- Do **not** add `Co-authored-by: Cursor` (or any Cursor co-author trailer) to commit messages.
+- Keep commit messages focused on the change itself, following the repository's existing style.
diff --git a/CHANGELOG.md b/CHANGELOG.md
index 6b692da..44e190a 100644
--- a/CHANGELOG.md
+++ b/CHANGELOG.md
@@ -21,6 +21,81 @@ removed no sooner than the next major (see `docs/API_STABILITY.md`).
### Added
+- **Canonical 2D → AI → 3D roadmap and image IO (Epic 83)**: `docs/ROADMAP.md`
+ reserves Epics 83–90; new `spatialrust-image-io` provides bounded path,
+ reader/writer, and memory PNG/JPEG/PNM codecs, independently gated TIFF and
+ OpenEXR, typed pixels, Exif orientation handling, Python/NumPy bindings,
+ property tests, and 640p/1080p/4K decode benchmarks.
+
+- **Shared CPU filters (Epic 84A–84B)**: `spatialrust-vision` now exposes a
+ common `BorderMode`, validated 1D/2D kernels, OpenCV-style filter2D
+ correlation, explicit convolution, f32-output and separable filters, and
+ normalized box/Gaussian blur, median and bilateral filters, signed
+ Sobel/Scharr/Laplacian derivatives, and Gaussian pyramids. The feature
+ includes strided-view property coverage, Python bindings/stubs, OpenCV
+ comparison, and 640p/1080p/4K benchmarks.
+
+- **CPU morphology (Epic 84C)**: validated rectangular, cross, elliptical,
+ diamond, and custom structuring elements; explicit-anchor erode/dilate;
+ open/close/gradient/top-hat/black-hat operations; additive feature and meta
+ feature; u8/u16/f32 and strided-view tests; Python bindings; exact OpenCV
+ comparisons; and 640p/1080p/4K benchmarks.
+
+- **CPU image analysis (Epic 84D)**: fixed and adaptive thresholds, u8/u16
+ Otsu selection, masked configurable histograms, exact u8 equalization,
+ contrast-limited adaptive equalization, and checked summed-area tables;
+ additive Rust/meta features, Python bindings/stubs, strided properties,
+ OpenCV comparisons, and representative-resolution benchmarks.
+
+- **Canny edge detection (Epic 84E)**: configurable 3/5/7 Sobel apertures,
+ L1/L2 gradient magnitude, directional non-maximum suppression, 8-neighbor
+ hysteresis, and inspectable intermediate stages; additive Rust/meta features,
+ strided and property tests, Python binding/stub, 640p/1080p/4K benchmark, and
+ exact OpenCV comparison across all six aperture/magnitude combinations.
+
+- **Tensor foundation (Epic 85A)**: new dependency-light `spatialrust-tensor`
+ crate with byte-addressable dtype, arbitrary-rank shape, signed element
+ strides, checked byte offsets/spans, explicit device identity, safe borrowed
+ CPU views, and named owned copies. Non-host device memory cannot be exposed as
+ a Rust byte slice, and the meta-crate integration is opt-in through `tensor`.
+
+- **Image/spatial tensor bridges (Epic 85B)**: packed interleaved images expose
+ zero-copy HWC views, packed planar images expose zero-copy CHW views, and
+ Schema-SoA `f32` point fields expose zero-copy one-dimensional views. Explicit
+ `pack_*` operations handle padded/ROI images, with feature-alone tests and
+ 640p/1080p/4K packing benchmarks.
+
+- **DLPack and Python tensor interoperability (Epic 85C–85D)**: audited
+ `DLManagedTensorVersioned` major-version 1 CPU import/export, explicit deleter
+ transfer, read-only/copy flags, signed strides and byte offsets, malformed ABI
+ rejection, and zero-copy Python `__dlpack__`/`__dlpack_device__`. NumPy and
+ PyTorch round trips preserve allocations and producer lifetimes; device or
+ host copy requests remain explicit.
+
+- **Inference contracts and ONNX Runtime CPU (Epic 86)**: new optional
+ `spatialrust-ai` crate with named dynamic model metadata, stable backend and
+ session traits, explicit input/output copy permissions, CPU ONNX Runtime,
+ separately gated CUDA/TensorRT/DirectML providers, typed zero-copy I/O
+ Binding, caller-preallocated u8/u16/f32 outputs, runtime-allocation retention,
+ output-to-input chaining, Python bindings/stubs, reference-runtime comparison,
+ and 640p/1080p/4K Criterion coverage. Multi-byte raw storage is rejected at
+ zero-copy boundaries instead of being cast from an under-aligned byte buffer.
+
+- **Feature2D and ORB matching (Epic 87)**: checked keypoint, binary/float
+ descriptor, feature-set, and match contracts; Harris and Shi–Tomasi corners;
+ OpenCV-exact FAST-9/16 detection and scores; deterministic multi-scale ORB
+ with 256-bit rotated BRIEF; Hamming/L2 brute-force matching with ratio,
+ cross-check, and distance filters; Python/NumPy bindings and stubs; property
+ tests, OpenCV comparison, and 640p/1080p/4K Criterion baselines.
+
+- **Camera geometry, motion, and stereo (Epic 88)**: checked correspondence and
+ projective contracts; normalized DLT and deterministic RANSAC for homography,
+ fundamental, and essential matrices; triangulation and essential pose
+ disambiguation; EPnP-class PnP with iterative refine and RANSAC; sparse
+ pyramidal Lucas–Kanade tracking; stereo rig, rectify remap grids, SAD block
+ matching, and disparity-to-depth/XYZ reproject; Python bindings; OpenCV
+ comparison with documented tolerances; property tests; and Criterion coverage.
+
- **AI-ready image and vision foundation (Epics 75–79)**: mutable ROI views,
planar/interleaved layouts and color metadata in `spatialrust-image`; new
feature-gated `spatialrust-vision` preprocessing, warp, detection, mask/RLE,
diff --git a/Cargo.toml b/Cargo.toml
index f58f0bd..f384303 100644
--- a/Cargo.toml
+++ b/Cargo.toml
@@ -15,6 +15,9 @@ members = [
"crates/spatialrust-transform",
"crates/spatialrust-voxelize",
"crates/spatialrust-image",
+ "crates/spatialrust-image-io",
+ "crates/spatialrust-tensor",
+ "crates/spatialrust-ai",
"crates/spatialrust-camera",
"crates/spatialrust-vision",
]
@@ -47,6 +50,9 @@ spatialrust-metrics = { path = "crates/spatialrust-metrics", version = "1.0.0",
spatialrust-transform = { path = "crates/spatialrust-transform", version = "1.0.0", default-features = false }
spatialrust-voxelize = { path = "crates/spatialrust-voxelize", version = "1.0.0", default-features = false }
spatialrust-image = { path = "crates/spatialrust-image", version = "1.0.0" }
+spatialrust-image-io = { path = "crates/spatialrust-image-io", version = "1.0.0", default-features = false }
+spatialrust-tensor = { path = "crates/spatialrust-tensor", version = "1.0.0", default-features = false }
+spatialrust-ai = { path = "crates/spatialrust-ai", version = "1.0.0", default-features = false }
spatialrust-camera = { path = "crates/spatialrust-camera", version = "1.0.0" }
spatialrust-vision = { path = "crates/spatialrust-vision", version = "1.0.0", default-features = false }
@@ -63,6 +69,10 @@ thiserror = "2"
wgpu = "24"
criterion = { version = "0.5", features = ["html_reports"] }
proptest = "1"
+image = { version = "0.24.9", default-features = false }
+kamadak-exif = "0.6.1"
+exr = { version = "=1.72.0", default-features = false }
+tempfile = "3"
[profile.release]
lto = "thin"
diff --git a/README.md b/README.md
index f7e9a02..d8a2a54 100644
--- a/README.md
+++ b/README.md
@@ -157,8 +157,11 @@ One dataflow, focused crates — each pipeline stage maps to the crate that impl
| `spatialrust-core` | Point schema, metadata, execution traits |
| `spatialrust-math` | Vec/Mat/Pose math primitives |
| `spatialrust-image` | Typed image buffers and zero-copy strided views |
+| `spatialrust-image-io` | Bounded PNG/JPEG/PNM codecs; opt-in TIFF/OpenEXR |
+| `spatialrust-tensor` | Runtime-independent dtype/shape/stride/device ownership and DLPack |
+| `spatialrust-ai` | Explicit-copy inference contracts and opt-in ONNX Runtime providers |
| `spatialrust-camera` | Pinhole/Brown–Conrady camera models and RGB-D conversion |
-| `spatialrust-vision` | Resize/preprocess, warps, detection postprocess, masks, and dense spatial maps |
+| `spatialrust-vision` | CPU filters, Feature2D/ORB matching, resize/preprocess, warps, detection postprocess, masks, and dense spatial maps |
| `spatialrust-io` | Point cloud readers/writers (PCD, PLY, LAS, COPC) |
| `spatialrust-search` | KD-tree search, k-NN / radius graphs |
| `spatialrust-filtering` | Voxel / FPS downsample, outlier removal, crop, MLS |
@@ -220,6 +223,26 @@ The reproducible algorithm comparison is in
`bench/opencv_vision_comparison/`; the complete synthetic demo is
`crates/spatialrust-py/examples/vision_ai_pipeline.py`.
+The same feature includes Harris, Shi–Tomasi, exact FAST-9/16, multi-scale ORB,
+and checked Hamming/L2 descriptor matching. Python exposes `orb_features` and
+NumPy matcher functions; OpenCV is used only by the numerical comparison suite.
+
+An ONNX Runtime wheel is opt-in (`maturin develop --features onnxruntime`). Its
+Python API uses named CPU I/O Binding by default; `copy=True` is the explicit
+fallback for inputs that must be repacked:
+
+```python
+session = sr.OnnxRuntimeSession("model.onnx", deterministic=True)
+input_tensor = sr.tensor_copy_from_numpy(chw)
+outputs = session.run({"images": input_tensor})
+scores = np.from_dlpack(outputs["scores"])
+```
+
+The Rust features are `ai`, `ai-onnxruntime`, and separate
+`ai-onnxruntime-{cuda,tensorrt,directml}` provider gates. The optional ONNX
+Runtime adapter currently has a feature-specific Rust 1.88 MSRV; it does not
+raise the default workspace MSRV.
+
diff --git a/bench/opencv_vision_comparison/README.md b/bench/opencv_vision_comparison/README.md
index 8682aa4..1a47998 100644
--- a/bench/opencv_vision_comparison/README.md
+++ b/bench/opencv_vision_comparison/README.md
@@ -1,8 +1,14 @@
# OpenCV vision comparison
This deterministic harness compares SpatialRust's Python-visible CPU vision
-primitives with OpenCV: four resize filters, RGB-to-gray/HSV conversion,
-bilinear remap, NMS, and connected-component areas.
+primitives with OpenCV: linear, median, and bilateral filters; Sobel, Scharr,
+Laplacian, Gaussian pyramids, morphology, thresholds, histograms, CLAHE, integral images,
+and Canny across 3/5/7 Sobel apertures and L1/L2 gradients;
+Harris, Shi–Tomasi, and FAST-9/16 keypoint coordinates/order (plus exact FAST scores);
+Hamming/L2 brute-force nearest matches; ORB keypoint repeatability and descriptor layout;
+homography transfer residuals, PnP translation vs OpenCV, and StereoBM center disparity
+on a synthetic textured pair; four resize filters; RGB-to-gray/HSV conversion;
+bilinear remap; NMS; and connected-component areas.
From the repository root, after installing the editable Python extension:
diff --git a/bench/opencv_vision_comparison/run.py b/bench/opencv_vision_comparison/run.py
index 9c588ae..81b0b95 100644
--- a/bench/opencv_vision_comparison/run.py
+++ b/bench/opencv_vision_comparison/run.py
@@ -24,6 +24,287 @@ def main() -> None:
size = (61, 43)
results: dict[str, object] = {"opencv_version": cv2.__version__}
+ kernel = np.array([[0.0, -0.25, 0.0], [-0.25, 2.0, -0.25], [0.0, -0.25, 0.0]])
+ filtered = sr.filter2d_image(image, kernel)
+ filtered_cv = cv2.filter2D(image, -1, kernel, borderType=cv2.BORDER_REFLECT_101)
+ filter_error = max_abs(filtered, filtered_cv)
+ results["filter2d_max_u8_error"] = filter_error
+ if filter_error > 1:
+ raise AssertionError(f"filter2D error {filter_error} > 1")
+
+ gaussian = sr.gaussian_blur_image(image, 5, 3, 1.2, 0.8)
+ gaussian_cv = cv2.GaussianBlur(
+ image, (5, 3), 1.2, sigmaY=0.8, borderType=cv2.BORDER_REFLECT_101
+ )
+ gaussian_error = max_abs(gaussian, gaussian_cv)
+ results["gaussian_blur_max_u8_error"] = gaussian_error
+ if gaussian_error > 1:
+ raise AssertionError(f"Gaussian blur error {gaussian_error} > 1")
+
+ median = sr.median_blur_image(image, 5)
+ median_cv = cv2.medianBlur(image, 5)
+ median_error = max_abs(median, median_cv)
+ results["median_blur_max_u8_error"] = median_error
+ if median_error != 0:
+ raise AssertionError(f"median blur error {median_error} != 0")
+
+ bilateral = sr.bilateral_filter_image(image, 5, 40.0, 3.0)
+ bilateral_cv = cv2.bilateralFilter(
+ image, 5, 40.0, 3.0, borderType=cv2.BORDER_REFLECT_101
+ )
+ bilateral_error = max_abs(bilateral, bilateral_cv)
+ results["bilateral_filter_max_u8_error"] = bilateral_error
+ if bilateral_error > 2:
+ raise AssertionError(f"bilateral filter error {bilateral_error} > 2")
+
+ gray_derivative = cv2.cvtColor(image, cv2.COLOR_RGB2GRAY)
+ derivative_cases = {
+ "sobel_x": (
+ sr.sobel_image(gray_derivative, 1, 0, 5),
+ cv2.Sobel(gray_derivative, cv2.CV_32F, 1, 0, ksize=5),
+ ),
+ "scharr_y": (
+ sr.scharr_image(gray_derivative, 0, 1),
+ cv2.Scharr(gray_derivative, cv2.CV_32F, 0, 1),
+ ),
+ "laplacian": (
+ sr.laplacian_image(gray_derivative, 3),
+ cv2.Laplacian(gray_derivative, cv2.CV_32F, ksize=3),
+ ),
+ }
+ for name, (actual, expected) in derivative_cases.items():
+ error = float(np.max(np.abs(actual - expected)))
+ results[f"{name}_max_f32_error"] = error
+ if error > 1e-4:
+ raise AssertionError(f"{name} error {error} > 1e-4")
+
+ pyramid = sr.pyr_down_image(image)
+ pyramid_cv = cv2.pyrDown(image)
+ pyramid_error = max_abs(pyramid, pyramid_cv)
+ results["pyr_down_max_u8_error"] = pyramid_error
+ if pyramid_error > 1:
+ raise AssertionError(f"pyrDown error {pyramid_error} > 1")
+ pyramid_up = sr.pyr_up_image(pyramid)
+ pyramid_up_cv = cv2.pyrUp(pyramid_cv)
+ pyramid_up_error = max_abs(pyramid_up, pyramid_up_cv)
+ results["pyr_up_max_u8_error"] = pyramid_up_error
+ if pyramid_up_error > 1:
+ raise AssertionError(f"pyrUp error {pyramid_up_error} > 1")
+
+ morphology_source = gray_derivative
+ morphology_cases = {
+ "erode": cv2.MORPH_ERODE,
+ "dilate": cv2.MORPH_DILATE,
+ "open": cv2.MORPH_OPEN,
+ "close": cv2.MORPH_CLOSE,
+ "gradient": cv2.MORPH_GRADIENT,
+ "tophat": cv2.MORPH_TOPHAT,
+ "blackhat": cv2.MORPH_BLACKHAT,
+ }
+ shape_cases = {
+ "rect": cv2.MORPH_RECT,
+ "cross": cv2.MORPH_CROSS,
+ "ellipse": cv2.MORPH_ELLIPSE,
+ }
+ for shape_name, shape_code in shape_cases.items():
+ element = cv2.getStructuringElement(shape_code, (5, 3))
+ for operation_name, operation_code in morphology_cases.items():
+ actual = sr.morphology_image(
+ morphology_source, operation_name, 5, 3, shape_name, 2
+ )
+ expected = cv2.morphologyEx(
+ morphology_source,
+ operation_code,
+ element,
+ iterations=2,
+ borderType=cv2.BORDER_REPLICATE,
+ )
+ error = max_abs(actual, expected)
+ results[f"morphology_{shape_name}_{operation_name}_max_u8_error"] = error
+ if error != 0:
+ raise AssertionError(
+ f"{shape_name}/{operation_name} morphology error {error} != 0"
+ )
+
+ threshold_actual = sr.threshold_image(gray_derivative, 117.0)
+ _, threshold_expected = cv2.threshold(gray_derivative, 117.0, 255, cv2.THRESH_BINARY)
+ results["threshold_max_u8_error"] = max_abs(threshold_actual, threshold_expected)
+ if results["threshold_max_u8_error"] != 0:
+ raise AssertionError("fixed threshold mismatch")
+
+ otsu_value, otsu_actual = sr.otsu_threshold_image(gray_derivative)
+ otsu_cv, otsu_expected = cv2.threshold(
+ gray_derivative, 0, 255, cv2.THRESH_BINARY | cv2.THRESH_OTSU
+ )
+ results["otsu_threshold"] = otsu_value
+ results["otsu_max_u8_error"] = max_abs(otsu_actual, otsu_expected)
+ if otsu_value != int(otsu_cv) or results["otsu_max_u8_error"] != 0:
+ raise AssertionError(f"Otsu mismatch: {otsu_value} != {otsu_cv}")
+
+ for method_name, method_code in {
+ "mean": cv2.ADAPTIVE_THRESH_MEAN_C,
+ "gaussian": cv2.ADAPTIVE_THRESH_GAUSSIAN_C,
+ }.items():
+ actual = sr.adaptive_threshold_image(gray_derivative, 7, 3.0, method_name)
+ expected = cv2.adaptiveThreshold(
+ gray_derivative, 255, method_code, cv2.THRESH_BINARY, 7, 3.0
+ )
+ error = max_abs(actual, expected)
+ results[f"adaptive_{method_name}_max_u8_error"] = error
+ if error != 0:
+ raise AssertionError(f"adaptive {method_name} mismatch: {error}")
+
+ histogram_actual = sr.histogram_image(gray_derivative)
+ histogram_expected = cv2.calcHist([gray_derivative], [0], None, [256], [0, 256]).reshape(-1)
+ histogram_error = int(np.max(np.abs(histogram_actual.astype(np.int64) - histogram_expected.astype(np.int64))))
+ results["histogram_max_count_error"] = histogram_error
+ if histogram_error != 0:
+ raise AssertionError(f"histogram mismatch: {histogram_error}")
+
+ equalized = sr.equalize_histogram_image(gray_derivative)
+ equalized_cv = cv2.equalizeHist(gray_derivative)
+ results["equalize_hist_max_u8_error"] = max_abs(equalized, equalized_cv)
+ if results["equalize_hist_max_u8_error"] != 0:
+ raise AssertionError("histogram equalization mismatch")
+
+ clahe_actual = sr.clahe_image(gray_derivative, 2.0, 8, 8)
+ clahe_expected = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8)).apply(gray_derivative)
+ clahe_error = max_abs(clahe_actual, clahe_expected)
+ results["clahe_max_u8_error"] = clahe_error
+ if clahe_error > 1:
+ raise AssertionError(f"CLAHE mismatch: {clahe_error}")
+
+ integral_actual = sr.integral_image_u8(gray_derivative)
+ integral_expected = cv2.integral(gray_derivative, sdepth=cv2.CV_64F)
+ integral_error = float(np.max(np.abs(integral_actual - integral_expected)))
+ results["integral_max_f64_error"] = integral_error
+ if integral_error != 0.0:
+ raise AssertionError(f"integral image mismatch: {integral_error}")
+
+ for aperture_size in (3, 5, 7):
+ for l2_gradient in (False, True):
+ actual = sr.canny_image(
+ gray_derivative, 50.0, 100.0, aperture_size, l2_gradient
+ )
+ expected = cv2.Canny(
+ gray_derivative,
+ 50.0,
+ 100.0,
+ apertureSize=aperture_size,
+ L2gradient=l2_gradient,
+ )
+ mismatch = int(np.count_nonzero(actual != expected))
+ name = f"canny_aperture_{aperture_size}_{'l2' if l2_gradient else 'l1'}"
+ results[f"{name}_mismatch_pixels"] = mismatch
+ if mismatch != 0:
+ raise AssertionError(f"{name} mismatch: {mismatch} pixels")
+
+ for nonmax_suppression in (False, True):
+ actual = sr.fast_keypoints(gray_derivative, 20, nonmax_suppression)
+ detector = cv2.FastFeatureDetector_create(
+ 20, nonmax_suppression, cv2.FAST_FEATURE_DETECTOR_TYPE_9_16
+ )
+ expected = detector.detect(gray_derivative, None)
+ actual_rows = [
+ (round(point.x), round(point.y), round(point.response)) for point in actual
+ ]
+ expected_rows = [
+ (round(point.pt[0]), round(point.pt[1]), round(point.response))
+ for point in expected
+ ]
+ name = f"fast_9_16_{'nms' if nonmax_suppression else 'raw'}"
+ results[f"{name}_keypoints"] = len(actual_rows)
+ if actual_rows != expected_rows:
+ raise AssertionError(f"{name} keypoints or scores differ from OpenCV")
+
+ for use_harris in (False, True):
+ if use_harris:
+ actual = sr.harris_keypoints(gray_derivative, 100, 0.01, 1.0, 3, 3, 0.04)
+ name = "harris"
+ else:
+ actual = sr.shi_tomasi_keypoints(gray_derivative, 100, 0.01, 1.0, 3, 3)
+ name = "shi_tomasi"
+ expected = cv2.goodFeaturesToTrack(
+ gray_derivative,
+ maxCorners=100,
+ qualityLevel=0.01,
+ minDistance=1.0,
+ mask=None,
+ blockSize=3,
+ useHarrisDetector=use_harris,
+ k=0.04,
+ )
+ actual_points = [(round(point.x), round(point.y)) for point in actual]
+ expected_points = (
+ []
+ if expected is None
+ else [(round(point[0][0]), round(point[0][1])) for point in expected]
+ )
+ results[f"{name}_keypoints"] = len(actual_points)
+ if actual_points != expected_points:
+ raise AssertionError(f"{name} ordering or coordinates differ from OpenCV")
+
+ binary_query = rng.integers(0, 256, size=(23, 32), dtype=np.uint8)
+ binary_train = rng.integers(0, 256, size=(31, 32), dtype=np.uint8)
+ actual_binary = sr.match_binary_descriptors(binary_query, binary_train)
+ expected_binary = cv2.BFMatcher(cv2.NORM_HAMMING).match(binary_query, binary_train)
+ actual_binary_rows = [(query, train, distance) for query, train, distance in actual_binary]
+ expected_binary_rows = [
+ (match.queryIdx, match.trainIdx, match.distance) for match in expected_binary
+ ]
+ results["hamming_matches"] = len(actual_binary_rows)
+ if actual_binary_rows != expected_binary_rows:
+ raise AssertionError("Hamming nearest matches differ from OpenCV BFMatcher")
+
+ float_query = rng.normal(size=(19, 17)).astype(np.float32)
+ float_train = rng.normal(size=(29, 17)).astype(np.float32)
+ actual_float = sr.match_float_descriptors(float_query, float_train)
+ expected_float = cv2.BFMatcher(cv2.NORM_L2).match(float_query, float_train)
+ float_index_mismatches = sum(
+ (query, train) != (expected.queryIdx, expected.trainIdx)
+ for (query, train, _), expected in zip(actual_float, expected_float)
+ )
+ float_distance_error = max(
+ abs(distance - expected.distance)
+ for (_, _, distance), expected in zip(actual_float, expected_float)
+ )
+ results["l2_match_index_mismatches"] = float_index_mismatches
+ results["l2_match_max_distance_error"] = float_distance_error
+ if float_index_mismatches != 0 or float_distance_error > 1e-5:
+ raise AssertionError(
+ f"L2 BFMatcher mismatch: indices={float_index_mismatches}, distance={float_distance_error}"
+ )
+
+ orb_image = cv2.resize(gray_derivative, (320, 240), interpolation=cv2.INTER_CUBIC)
+ actual_orb, actual_descriptors = sr.orb_features(
+ orb_image, max_features=200, edge_threshold=16
+ )
+ cv_orb = cv2.ORB_create(nfeatures=200, edgeThreshold=16)
+ expected_orb, expected_descriptors = cv_orb.detectAndCompute(orb_image, None)
+ actual_coordinates = np.array([(point.x, point.y) for point in actual_orb], dtype=np.float32)
+ expected_coordinates = np.array([point.pt for point in expected_orb], dtype=np.float32)
+ repeatable = 0
+ if len(actual_coordinates) and len(expected_coordinates):
+ nearest_distances = np.sqrt(
+ np.min(
+ np.sum(
+ (actual_coordinates[:, None, :] - expected_coordinates[None, :, :]) ** 2,
+ axis=2,
+ ),
+ axis=1,
+ )
+ )
+ repeatable = int(np.count_nonzero(nearest_distances <= 2.0))
+ repeatability = repeatable / max(1, min(len(actual_orb), len(expected_orb)))
+ results["orb_spatialrust_keypoints"] = len(actual_orb)
+ results["orb_opencv_keypoints"] = len(expected_orb)
+ results["orb_coordinate_repeatability_2px"] = repeatability
+ results["orb_descriptor_width"] = actual_descriptors.shape[1]
+ if actual_descriptors.shape != (len(actual_orb), 32):
+ raise AssertionError("SpatialRust ORB descriptor layout is not N x 32")
+ if expected_descriptors is None or repeatability < 0.30:
+ raise AssertionError(f"ORB repeatability against OpenCV is too low: {repeatability}")
+
resize_cases = {
"nearest": getattr(cv2, "INTER_NEAREST_EXACT", cv2.INTER_NEAREST),
"bilinear": cv2.INTER_LINEAR,
@@ -101,6 +382,103 @@ def main() -> None:
if areas != areas_cv:
raise AssertionError(f"component areas mismatch: {areas} != {areas_cv}")
+ # Geometry: planar homography residual agreement (not scale-normalized identity).
+ source = np.array(
+ [[10.0, 12.0], [70.0, 14.0], [18.0, 55.0], [66.0, 60.0], [40.0, 34.0], [28.0, 22.0]],
+ dtype=np.float64,
+ )
+ homography = np.array(
+ [[1.04, 0.015, 2.5], [-0.02, 0.97, -1.25], [0.0002, -0.0001, 1.0]],
+ dtype=np.float64,
+ )
+ target = []
+ for point in source:
+ projected = homography @ np.array([point[0], point[1], 1.0])
+ target.append([projected[0] / projected[2], projected[1] / projected[2]])
+ target = np.asarray(target, dtype=np.float64)
+ estimated, inliers, residuals = sr.estimate_homography_ransac(
+ source, target, threshold=1.0, seed=3
+ )
+ estimated_cv, mask_cv = cv2.findHomography(source, target, method=0)
+ max_residual = float(np.max(residuals))
+ results["homography_max_residual"] = max_residual
+ results["homography_inliers"] = int(np.sum(inliers))
+ if max_residual > 1e-6 or estimated_cv is None:
+ raise AssertionError(f"homography residual {max_residual} too large")
+ # Compare transfer error of both models; allow tiny numeric disagreement.
+ def transfer_error(matrix: np.ndarray) -> float:
+ errors = []
+ for src, dst in zip(source, target):
+ projected = matrix @ np.array([src[0], src[1], 1.0])
+ errors.append(
+ np.hypot(projected[0] / projected[2] - dst[0], projected[1] / projected[2] - dst[1])
+ )
+ return float(np.max(errors))
+
+ transfer_sr = transfer_error(estimated)
+ transfer_cv = transfer_error(estimated_cv)
+ results["homography_transfer_sr"] = transfer_sr
+ results["homography_transfer_cv"] = transfer_cv
+ if transfer_sr > 1e-5 or transfer_cv > 1e-5:
+ raise AssertionError("homography transfer error exceeds tolerance")
+
+ objects = np.array(
+ [
+ [0.0, 0.0, 0.0],
+ [0.25, 0.0, 0.0],
+ [0.0, 0.2, 0.0],
+ [0.0, 0.0, 0.15],
+ [0.1, 0.1, 0.05],
+ [0.05, -0.08, 0.02],
+ [-0.1, 0.05, 0.08],
+ [0.12, 0.04, -0.03],
+ ],
+ dtype=np.float64,
+ )
+ fx = fy = 500.0
+ cx = cy = 240.0
+ true_r = np.eye(3, dtype=np.float64)
+ true_t = np.array([0.12, -0.04, 2.4], dtype=np.float64)
+ images = []
+ for point in objects:
+ camera = true_r @ point + true_t
+ images.append([fx * camera[0] / camera[2] + cx, fy * camera[1] / camera[2] + cy])
+ images = np.asarray(images, dtype=np.float64)
+ rotation, translation = sr.solve_pnp(objects, images, fx, fy, cx, cy, 480, 480)
+ ok_cv, rvec, tvec = cv2.solvePnP(
+ objects.astype(np.float64),
+ images.astype(np.float64),
+ np.array([[fx, 0.0, cx], [0.0, fy, cy], [0.0, 0.0, 1.0]]),
+ None,
+ flags=cv2.SOLVEPNP_ITERATIVE,
+ )
+ if not ok_cv:
+ raise AssertionError("OpenCV solvePnP failed")
+ t_error = float(np.linalg.norm(translation - tvec.reshape(3)))
+ results["pnp_translation_l2_vs_opencv"] = t_error
+ if abs(translation[2] - true_t[2]) > 0.05 or t_error > 0.05:
+ raise AssertionError(f"PnP translation disagreement {t_error}")
+ _ = rotation # shape already checked by consumer use
+
+ width, height = 128, 96
+ disparity = 16
+ yy, xx = np.indices((height, width), dtype=np.int32)
+ left = ((xx * 17 + yy * 29) % 200 + 20).astype(np.uint8)
+ right = np.zeros_like(left)
+ right[:, : width - disparity] = left[:, disparity:]
+ disparity_sr = sr.stereo_block_match(
+ left, right, window_size=11, min_disparity=1, num_disparities=32, uniqueness_ratio=5.0
+ )
+ matcher = cv2.StereoBM_create(numDisparities=32, blockSize=11)
+ disparity_cv = matcher.compute(left, right).astype(np.float32) / 16.0
+ center = (height // 2, width // 2)
+ results["stereo_bm_sr_center"] = float(disparity_sr[center])
+ results["stereo_bm_cv_center"] = float(disparity_cv[center])
+ if abs(float(disparity_sr[center]) - float(disparity)) > 1.0:
+ raise AssertionError("SpatialRust StereoBM center disparity incorrect")
+ if abs(float(disparity_cv[center]) - float(disparity)) > 1.5:
+ raise AssertionError("OpenCV StereoBM center disparity unexpected for synthetic pair")
+
results["status"] = "pass"
print(json.dumps(results, indent=2, sort_keys=True))
diff --git a/crates/spatialrust-ai/Cargo.toml b/crates/spatialrust-ai/Cargo.toml
new file mode 100644
index 0000000..500e756
--- /dev/null
+++ b/crates/spatialrust-ai/Cargo.toml
@@ -0,0 +1,30 @@
+[package]
+name = "spatialrust-ai"
+version.workspace = true
+edition.workspace = true
+license.workspace = true
+authors.workspace = true
+repository.workspace = true
+rust-version.workspace = true
+description = "Backend-independent inference contracts for SpatialRust"
+
+[features]
+default = []
+onnxruntime = ["dep:ort"]
+onnxruntime-cuda = ["onnxruntime", "ort/cuda"]
+onnxruntime-tensorrt = ["onnxruntime-cuda", "ort/tensorrt"]
+onnxruntime-directml = ["onnxruntime", "ort/directml"]
+
+[dependencies]
+bytemuck.workspace = true
+spatialrust-tensor.workspace = true
+thiserror.workspace = true
+ort = { version = "=2.0.0-rc.12", optional = true, default-features = false, features = ["std", "ndarray", "download-binaries", "tls-native", "copy-dylibs", "api-24"] }
+
+[dev-dependencies]
+criterion.workspace = true
+
+[[bench]]
+name = "onnxruntime"
+harness = false
+required-features = ["onnxruntime"]
diff --git a/crates/spatialrust-ai/benches/onnxruntime.rs b/crates/spatialrust-ai/benches/onnxruntime.rs
new file mode 100644
index 0000000..05476f9
--- /dev/null
+++ b/crates/spatialrust-ai/benches/onnxruntime.rs
@@ -0,0 +1,71 @@
+use std::{hint::black_box, sync::Arc, time::Duration};
+
+use criterion::{criterion_group, criterion_main, BenchmarkId, Criterion, Throughput};
+use spatialrust_ai::{
+ CopyPolicy, InferenceBackend, IoBinding, ModelSource, NamedTensors, OnnxRuntimeBackend,
+ OutputBinding, RunOptions, SessionOptions,
+};
+use spatialrust_tensor::{DataType, Device, TensorBuffer, TensorDescriptor};
+
+const DOUBLE_DYNAMIC: &[u8] = &[
+ 8, 8, 18, 16, 115, 112, 97, 116, 105, 97, 108, 114, 117, 115, 116, 45, 116, 101, 115, 116, 58,
+ 106, 10, 27, 10, 5, 105, 110, 112, 117, 116, 10, 5, 105, 110, 112, 117, 116, 18, 6, 111, 117,
+ 116, 112, 117, 116, 34, 3, 65, 100, 100, 18, 14, 100, 111, 117, 98, 108, 101, 95, 100, 121,
+ 110, 97, 109, 105, 99, 90, 28, 10, 5, 105, 110, 112, 117, 116, 18, 19, 10, 17, 8, 1, 18, 13,
+ 10, 7, 18, 5, 98, 97, 116, 99, 104, 10, 2, 8, 3, 98, 29, 10, 6, 111, 117, 116, 112, 117, 116,
+ 18, 19, 10, 17, 8, 1, 18, 13, 10, 7, 18, 5, 98, 97, 116, 99, 104, 10, 2, 8, 3, 66, 4, 10, 0,
+ 16, 13,
+];
+
+fn benchmark_onnxruntime(c: &mut Criterion) {
+ let backend = OnnxRuntimeBackend;
+ let mut session = backend
+ .create_session(&ModelSource::Bytes(Arc::from(DOUBLE_DYNAMIC)), &SessionOptions::default())
+ .expect("embedded model");
+ let mut group = c.benchmark_group("onnxruntime_cpu_dynamic_rgb_f32");
+ group.sample_size(10);
+ group.warm_up_time(Duration::from_millis(500));
+ group.measurement_time(Duration::from_secs(2));
+
+ for (label, width, height) in
+ [("640p", 640_usize, 480_usize), ("1080p", 1920, 1080), ("4k", 3840, 2160)]
+ {
+ let pixels = width * height;
+ let descriptor = TensorDescriptor::contiguous(DataType::F32, vec![pixels, 3], Device::CPU);
+ let input = TensorBuffer::try_from_f32(vec![0.5; pixels * 3], descriptor).unwrap();
+ let mut inputs = NamedTensors::new();
+ inputs.insert("input", input).unwrap();
+ group.throughput(Throughput::Bytes((pixels * 3 * 4 * 2) as u64));
+
+ group.bench_with_input(BenchmarkId::new("copy_run", label), &inputs, |b, inputs| {
+ b.iter(|| {
+ black_box(
+ session
+ .run_with_options(
+ inputs.clone(),
+ RunOptions {
+ input_copy: CopyPolicy::Allow,
+ output_copy: CopyPolicy::Allow,
+ },
+ )
+ .unwrap(),
+ )
+ });
+ });
+ group.bench_with_input(BenchmarkId::new("io_binding", label), &inputs, |b, inputs| {
+ b.iter(|| {
+ let mut binding = IoBinding::try_new(
+ inputs.clone(),
+ vec![OutputBinding::Allocate { name: "output".into(), device: Device::CPU }],
+ )
+ .unwrap();
+ session.run_with_binding(&mut binding).unwrap();
+ black_box(binding.into_results())
+ });
+ });
+ }
+ group.finish();
+}
+
+criterion_group!(benches, benchmark_onnxruntime);
+criterion_main!(benches);
diff --git a/crates/spatialrust-ai/src/lib.rs b/crates/spatialrust-ai/src/lib.rs
new file mode 100644
index 0000000..4f64f4e
--- /dev/null
+++ b/crates/spatialrust-ai/src/lib.rs
@@ -0,0 +1,648 @@
+//! Backend-independent model, session, named-I/O, and explicit binding contracts.
+//!
+//! The default build contains no inference runtime. ONNX Runtime and hardware
+//! execution providers are additive features and must preserve the copy/device
+//! choices represented by these types.
+
+#![deny(unsafe_code)]
+#![warn(missing_docs)]
+
+use std::{path::PathBuf, sync::Arc};
+
+use spatialrust_tensor::{DataType, Device, TensorBuffer, TensorDescriptor};
+
+#[cfg(feature = "onnxruntime")]
+mod onnxruntime;
+#[cfg(feature = "onnxruntime")]
+pub use onnxruntime::{OnnxRuntimeBackend, OnnxRuntimeSession};
+
+/// Result type for inference operations.
+pub type AiResult = Result;
+
+/// Errors shared by inference contracts and backend adapters.
+#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
+pub enum AiError {
+ /// A named input or output appears more than once.
+ #[error("duplicate tensor name `{0}`")]
+ DuplicateName(String),
+ /// A required model input was not supplied.
+ #[error("missing required input `{0}`")]
+ MissingInput(String),
+ /// A tensor name is not declared by the model.
+ #[error("unexpected tensor `{0}`")]
+ UnexpectedTensor(String),
+ /// Actual dtype differs from the model contract.
+ #[error("tensor `{name}` requires {expected:?}, found {actual:?}")]
+ DataTypeMismatch {
+ /// Tensor name.
+ name: String,
+ /// Model dtype.
+ expected: DataType,
+ /// Supplied dtype.
+ actual: DataType,
+ },
+ /// Actual rank or dimension differs from the model contract.
+ #[error("tensor `{name}` shape {actual:?} does not match {expected:?}")]
+ ShapeMismatch {
+ /// Tensor name.
+ name: String,
+ /// Model dimensions.
+ expected: Vec,
+ /// Supplied dimensions.
+ actual: Vec,
+ },
+ /// A preallocated output does not have enough bytes.
+ #[error("preallocated output `{name}` needs {required} bytes, found {found}")]
+ OutputBufferTooSmall {
+ /// Output name.
+ name: String,
+ /// Required allocation bytes.
+ required: usize,
+ /// Available allocation bytes.
+ found: usize,
+ },
+ /// The selected backend cannot honor an operation or option.
+ #[error("backend `{backend}` does not support {operation}")]
+ Unsupported {
+ /// Backend identifier.
+ backend: String,
+ /// Unsupported operation.
+ operation: String,
+ },
+ /// Backend-specific failure with stable outer context.
+ #[error("backend `{backend}` failed: {message}")]
+ Backend {
+ /// Backend identifier.
+ backend: String,
+ /// Backend error message.
+ message: String,
+ },
+ /// Invalid public configuration.
+ #[error("invalid inference configuration: {0}")]
+ InvalidConfiguration(String),
+ /// A backend would need a copy that the caller did not authorize.
+ #[error(
+ "{direction} copy is required for tensor `{name}`; opt in explicitly or use I/O binding"
+ )]
+ CopyRequired {
+ /// Input or output transfer direction.
+ direction: &'static str,
+ /// Model-visible tensor name.
+ name: String,
+ },
+}
+
+/// One model dimension, fixed, unconstrained, or symbolically dynamic.
+#[derive(Clone, Debug, PartialEq, Eq, Hash)]
+pub enum Dimension {
+ /// Exact dimension size.
+ Fixed(usize),
+ /// Dynamic dimension without a model-provided symbol.
+ Dynamic,
+ /// Dynamic dimension sharing a model-provided symbolic name.
+ Symbol(String),
+}
+
+impl Dimension {
+ fn accepts(&self, actual: usize) -> bool {
+ match self {
+ Self::Fixed(expected) => *expected == actual,
+ Self::Dynamic | Self::Symbol(_) => true,
+ }
+ }
+}
+
+/// One named model input or output contract.
+#[derive(Clone, Debug, PartialEq, Eq)]
+pub struct TensorSpec {
+ /// ONNX/model-visible name.
+ pub name: String,
+ /// Required scalar/vector dtype.
+ pub dtype: DataType,
+ /// Ordered fixed or dynamic dimensions.
+ pub shape: Vec,
+}
+
+impl TensorSpec {
+ /// Creates a named tensor specification.
+ pub fn new(name: impl Into, dtype: DataType, shape: Vec) -> Self {
+ Self { name: name.into(), dtype, shape }
+ }
+
+ /// Validates a concrete tensor descriptor against this specification.
+ pub fn validate(&self, descriptor: &TensorDescriptor) -> AiResult<()> {
+ if descriptor.dtype() != self.dtype {
+ return Err(AiError::DataTypeMismatch {
+ name: self.name.clone(),
+ expected: self.dtype,
+ actual: descriptor.dtype(),
+ });
+ }
+ if descriptor.shape().len() != self.shape.len()
+ || !self
+ .shape
+ .iter()
+ .zip(descriptor.shape())
+ .all(|(expected, &actual)| expected.accepts(actual))
+ {
+ return Err(AiError::ShapeMismatch {
+ name: self.name.clone(),
+ expected: self.shape.clone(),
+ actual: descriptor.shape().to_vec(),
+ });
+ }
+ Ok(())
+ }
+}
+
+/// Model-visible named input and output metadata.
+#[derive(Clone, Debug, Default, PartialEq, Eq)]
+pub struct ModelInfo {
+ /// Optional producer/model identifier.
+ pub name: Option,
+ /// Ordered model inputs.
+ pub inputs: Vec,
+ /// Ordered model outputs.
+ pub outputs: Vec,
+}
+
+impl ModelInfo {
+ /// Validates uniqueness of all input and output names.
+ pub fn validate(&self) -> AiResult<()> {
+ ensure_unique(self.inputs.iter().map(|spec| spec.name.as_str()))?;
+ ensure_unique(self.outputs.iter().map(|spec| spec.name.as_str()))?;
+ Ok(())
+ }
+
+ /// Validates required names, dtype, and dynamic/fixed shapes for one request.
+ pub fn validate_inputs(&self, inputs: &NamedTensors) -> AiResult<()> {
+ for spec in &self.inputs {
+ let tensor =
+ inputs.get(&spec.name).ok_or_else(|| AiError::MissingInput(spec.name.clone()))?;
+ spec.validate(tensor.descriptor())?;
+ }
+ for (name, _) in inputs.iter() {
+ if !self.inputs.iter().any(|spec| spec.name == name) {
+ return Err(AiError::UnexpectedTensor(name.to_owned()));
+ }
+ }
+ Ok(())
+ }
+}
+
+fn ensure_unique<'a>(names: impl IntoIterator- ) -> AiResult<()> {
+ let mut seen = Vec::<&str>::new();
+ for name in names {
+ if seen.contains(&name) {
+ return Err(AiError::DuplicateName(name.to_owned()));
+ }
+ seen.push(name);
+ }
+ Ok(())
+}
+
+/// Ordered, uniquely named tensor collection.
+#[derive(Clone, Debug, Default, PartialEq, Eq)]
+pub struct NamedTensors {
+ values: Vec<(String, TensorBuffer)>,
+}
+
+impl NamedTensors {
+ /// Creates an empty collection.
+ pub const fn new() -> Self {
+ Self { values: Vec::new() }
+ }
+
+ /// Inserts a unique named tensor while retaining insertion order.
+ pub fn insert(&mut self, name: impl Into, tensor: TensorBuffer) -> AiResult<()> {
+ let name = name.into();
+ if self.values.iter().any(|(existing, _)| existing == &name) {
+ return Err(AiError::DuplicateName(name));
+ }
+ self.values.push((name, tensor));
+ Ok(())
+ }
+
+ /// Returns a tensor by model-visible name.
+ pub fn get(&self, name: &str) -> Option<&TensorBuffer> {
+ self.values.iter().find(|(candidate, _)| candidate == name).map(|(_, value)| value)
+ }
+
+ /// Returns ordered `(name, tensor)` pairs.
+ pub fn iter(&self) -> impl Iterator
- {
+ self.values.iter().map(|(name, tensor)| (name.as_str(), tensor))
+ }
+
+ /// Returns the tensor count.
+ pub fn len(&self) -> usize {
+ self.values.len()
+ }
+
+ /// Returns whether no tensors are present.
+ pub fn is_empty(&self) -> bool {
+ self.values.is_empty()
+ }
+
+ /// Consumes the collection into ordered pairs.
+ pub fn into_values(self) -> Vec<(String, TensorBuffer)> {
+ self.values
+ }
+}
+
+/// Model bytes or a filesystem path supplied explicitly at session creation.
+#[derive(Clone, Debug)]
+pub enum ModelSource {
+ /// Read model bytes from this path during session creation.
+ Path(PathBuf),
+ /// Immutable in-memory model bytes.
+ Bytes(Arc<[u8]>),
+}
+
+/// Graph optimization level requested from a backend.
+#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Hash)]
+pub enum GraphOptimization {
+ /// Disable graph rewrites.
+ Disabled,
+ /// Apply safe basic rewrites.
+ Basic,
+ /// Apply extended rewrites.
+ Extended,
+ /// Apply all backend-supported rewrites.
+ #[default]
+ All,
+}
+
+/// Runtime-independent session configuration.
+#[derive(Clone, Debug, PartialEq, Eq)]
+pub struct SessionOptions {
+ /// Intra-operator worker count; `None` delegates to the backend.
+ pub intra_threads: Option,
+ /// Inter-operator worker count; `None` delegates to the backend.
+ pub inter_threads: Option,
+ /// Graph rewrite level.
+ pub graph_optimization: GraphOptimization,
+ /// Request deterministic kernels where supported.
+ pub deterministic: bool,
+}
+
+/// Whether a run may allocate and copy host tensor data.
+#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Hash)]
+pub enum CopyPolicy {
+ /// Fail instead of silently copying.
+ #[default]
+ Forbid,
+ /// Permit a documented host-to-host copy for this run.
+ Allow,
+}
+
+/// Per-run host copy permissions.
+#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Hash)]
+pub struct RunOptions {
+ /// Permission to repack/copy named inputs into backend values.
+ pub input_copy: CopyPolicy,
+ /// Permission to copy backend-owned outputs into `TensorBuffer`.
+ pub output_copy: CopyPolicy,
+}
+
+impl Default for SessionOptions {
+ fn default() -> Self {
+ Self {
+ intra_threads: None,
+ inter_threads: None,
+ graph_optimization: GraphOptimization::All,
+ deterministic: false,
+ }
+ }
+}
+
+impl SessionOptions {
+ /// Rejects zero thread counts.
+ pub fn validate(&self) -> AiResult<()> {
+ if self.intra_threads == Some(0) || self.inter_threads == Some(0) {
+ return Err(AiError::InvalidConfiguration("thread counts must be positive".into()));
+ }
+ Ok(())
+ }
+}
+
+/// Explicit output destination for an I/O-bound run.
+#[derive(Clone, Debug, PartialEq, Eq)]
+pub enum OutputBinding {
+ /// Ask the backend to allocate this output on the named device.
+ Allocate {
+ /// Model output name.
+ name: String,
+ /// Required allocation device.
+ device: Device,
+ },
+ /// Write into this caller-owned CPU allocation.
+ PreallocatedCpu(PreallocatedOutput),
+}
+
+/// Caller-owned mutable bytes for an explicitly bound CPU output.
+#[derive(Clone, Debug, PartialEq, Eq)]
+pub struct PreallocatedOutput {
+ name: String,
+ descriptor: TensorDescriptor,
+ storage: Option,
+}
+
+#[derive(Clone, Debug)]
+pub(crate) enum PreallocatedStorage {
+ Bytes(Vec),
+ U16(Vec),
+ F32(Vec),
+}
+
+impl PreallocatedStorage {
+ fn allocation_bytes(&self) -> &[u8] {
+ match self {
+ Self::Bytes(values) => values,
+ Self::U16(values) => bytemuck::cast_slice(values),
+ Self::F32(values) => bytemuck::cast_slice(values),
+ }
+ }
+
+ fn allocation_bytes_mut(&mut self) -> &mut [u8] {
+ match self {
+ Self::Bytes(values) => values,
+ Self::U16(values) => bytemuck::cast_slice_mut(values),
+ Self::F32(values) => bytemuck::cast_slice_mut(values),
+ }
+ }
+}
+
+impl PartialEq for PreallocatedStorage {
+ fn eq(&self, other: &Self) -> bool {
+ self.allocation_bytes() == other.allocation_bytes()
+ }
+}
+
+impl Eq for PreallocatedStorage {}
+
+impl PreallocatedOutput {
+ /// Creates a checked preallocated CPU output.
+ pub fn try_new(
+ name: impl Into,
+ descriptor: TensorDescriptor,
+ bytes: Vec,
+ ) -> AiResult {
+ let name = name.into();
+ let required = descriptor
+ .required_byte_range()
+ .map_err(|error| AiError::InvalidConfiguration(error.to_string()))?
+ .end;
+ if !descriptor.device().is_host_accessible() {
+ return Err(AiError::InvalidConfiguration(
+ "PreallocatedCpu requires host-accessible storage".into(),
+ ));
+ }
+ if bytes.len() < required {
+ return Err(AiError::OutputBufferTooSmall { name, required, found: bytes.len() });
+ }
+ Ok(Self { name, descriptor, storage: Some(PreallocatedStorage::Bytes(bytes)) })
+ }
+
+ /// Allocates aligned storage for a compact `u8`, `u16`, or `f32` CPU output.
+ pub fn allocate(name: impl Into, descriptor: TensorDescriptor) -> AiResult {
+ if !descriptor.is_c_contiguous() || descriptor.byte_offset() != 0 {
+ return Err(AiError::InvalidConfiguration(
+ "preallocated output must be compact with byte_offset=0".into(),
+ ));
+ }
+ if descriptor.device() != Device::CPU {
+ return Err(AiError::InvalidConfiguration(
+ "preallocated output allocation currently requires Device::CPU".into(),
+ ));
+ }
+ let elements = descriptor
+ .element_count()
+ .map_err(|error| AiError::InvalidConfiguration(error.to_string()))?;
+ let storage = match descriptor.dtype() {
+ DataType::U8 => PreallocatedStorage::Bytes(vec![0; elements]),
+ DataType::U16 => PreallocatedStorage::U16(vec![0; elements]),
+ DataType::F32 => PreallocatedStorage::F32(vec![0.0; elements]),
+ dtype => {
+ return Err(AiError::InvalidConfiguration(format!(
+ "preallocated output dtype {dtype:?} is not supported; use u8, u16, or f32"
+ )))
+ }
+ };
+ Ok(Self { name: name.into(), descriptor, storage: Some(storage) })
+ }
+
+ /// Returns the model output name.
+ pub fn name(&self) -> &str {
+ &self.name
+ }
+
+ /// Returns output metadata.
+ pub const fn descriptor(&self) -> &TensorDescriptor {
+ &self.descriptor
+ }
+
+ /// Returns mutable caller-owned output bytes for a backend binding.
+ pub fn allocation_bytes_mut(&mut self) -> AiResult<&mut [u8]> {
+ self.storage.as_mut().map(PreallocatedStorage::allocation_bytes_mut).ok_or_else(|| {
+ AiError::InvalidConfiguration(format!(
+ "preallocated output `{}` has already been consumed",
+ self.name
+ ))
+ })
+ }
+
+ /// Returns whether a bound run has taken ownership of this allocation.
+ pub fn is_consumed(&self) -> bool {
+ self.storage.is_none()
+ }
+
+ /// Converts a completed binding into generic owned tensor storage.
+ pub fn into_tensor(mut self) -> AiResult {
+ let storage = self.take_storage()?;
+ tensor_from_preallocated(storage, self.descriptor)
+ }
+
+ pub(crate) fn take_storage(&mut self) -> AiResult {
+ self.storage.take().ok_or_else(|| {
+ AiError::InvalidConfiguration(format!(
+ "preallocated output `{}` has already been consumed",
+ self.name
+ ))
+ })
+ }
+}
+
+fn tensor_from_preallocated(
+ storage: PreallocatedStorage,
+ descriptor: TensorDescriptor,
+) -> AiResult {
+ let result = match storage {
+ PreallocatedStorage::Bytes(values) => TensorBuffer::try_new(values, descriptor),
+ PreallocatedStorage::U16(values) => TensorBuffer::try_from_u16(values, descriptor),
+ PreallocatedStorage::F32(values) => TensorBuffer::try_from_f32(values, descriptor),
+ };
+ result.map_err(|error| AiError::InvalidConfiguration(error.to_string()))
+}
+
+/// Inputs, requested output destinations, and completed named results.
+#[derive(Clone, Debug, PartialEq, Eq)]
+pub struct IoBinding {
+ inputs: NamedTensors,
+ outputs: Vec,
+ results: Option,
+}
+
+impl IoBinding {
+ /// Creates an explicit binding request.
+ pub fn try_new(inputs: NamedTensors, outputs: Vec) -> AiResult {
+ ensure_unique(outputs.iter().map(|output| match output {
+ OutputBinding::Allocate { name, .. } => name.as_str(),
+ OutputBinding::PreallocatedCpu(output) => output.name(),
+ }))?;
+ Ok(Self { inputs, outputs, results: None })
+ }
+
+ /// Returns bound named inputs.
+ pub const fn inputs(&self) -> &NamedTensors {
+ &self.inputs
+ }
+
+ /// Returns requested output destinations.
+ pub fn outputs(&self) -> &[OutputBinding] {
+ &self.outputs
+ }
+
+ /// Returns mutable requested output destinations.
+ pub fn outputs_mut(&mut self) -> &mut [OutputBinding] {
+ &mut self.outputs
+ }
+
+ /// Stores completed named results. Intended for backend implementers.
+ pub fn set_results(&mut self, results: NamedTensors) {
+ self.results = Some(results);
+ }
+
+ /// Clears results before a backend starts another run.
+ pub fn clear_results(&mut self) {
+ self.results = None;
+ }
+
+ /// Returns completed results after a successful bound run.
+ pub fn results(&self) -> Option<&NamedTensors> {
+ self.results.as_ref()
+ }
+
+ /// Consumes completed results, if present.
+ pub fn into_results(self) -> Option {
+ self.results
+ }
+}
+
+/// Stable interface implemented by inference engines.
+pub trait InferenceBackend: Send + Sync {
+ /// Stable backend identifier such as `onnxruntime-cpu`.
+ fn name(&self) -> &str;
+
+ /// Loads one model and returns an independent mutable session.
+ fn create_session(
+ &self,
+ source: &ModelSource,
+ options: &SessionOptions,
+ ) -> AiResult>;
+}
+
+/// Loaded model session with named dynamic I/O.
+pub trait ModelSession: Send {
+ /// Stable backend identifier for diagnostics and capability checks.
+ fn backend_name(&self) -> &str;
+
+ /// Returns model input/output metadata captured at load time.
+ fn model_info(&self) -> &ModelInfo;
+
+ /// Runs without authorizing hidden host copies.
+ fn run(&mut self, inputs: NamedTensors) -> AiResult {
+ self.run_with_options(inputs, RunOptions::default())
+ }
+
+ /// Runs named I/O with explicit host copy permissions.
+ fn run_with_options(
+ &mut self,
+ inputs: NamedTensors,
+ options: RunOptions,
+ ) -> AiResult;
+
+ /// Runs with explicit input/output device or allocation bindings.
+ fn run_with_binding(&mut self, _binding: &mut IoBinding) -> AiResult<()> {
+ Err(AiError::Unsupported {
+ backend: self.backend_name().into(),
+ operation: "explicit I/O binding".into(),
+ })
+ }
+}
+
+#[cfg(test)]
+mod tests {
+ use super::{
+ AiError, Dimension, IoBinding, ModelInfo, NamedTensors, OutputBinding, PreallocatedOutput,
+ TensorSpec,
+ };
+ use spatialrust_tensor::{DataType, Device, TensorBuffer, TensorDescriptor};
+
+ fn tensor(shape: Vec) -> TensorBuffer {
+ let len = shape.iter().product::() * 4;
+ TensorBuffer::try_new(
+ vec![0; len],
+ TensorDescriptor::contiguous(DataType::F32, shape, Device::CPU),
+ )
+ .unwrap()
+ }
+
+ #[test]
+ fn validates_named_dynamic_inputs() {
+ let info = ModelInfo {
+ name: Some("dynamic-batch".into()),
+ inputs: vec![TensorSpec::new(
+ "images",
+ DataType::F32,
+ vec![Dimension::Symbol("batch".into()), Dimension::Fixed(3), Dimension::Dynamic],
+ )],
+ outputs: vec![],
+ };
+ let mut inputs = NamedTensors::new();
+ inputs.insert("images", tensor(vec![4, 3, 224])).unwrap();
+ info.validate_inputs(&inputs).unwrap();
+ let wrong = TensorDescriptor::contiguous(DataType::F32, vec![4, 1, 224], Device::CPU);
+ assert!(matches!(info.inputs[0].validate(&wrong), Err(AiError::ShapeMismatch { .. })));
+ }
+
+ #[test]
+ fn named_tensors_reject_duplicates_and_missing_inputs() {
+ let mut inputs = NamedTensors::new();
+ inputs.insert("x", tensor(vec![2])).unwrap();
+ assert!(matches!(inputs.insert("x", tensor(vec![2])), Err(AiError::DuplicateName(_))));
+ let info = ModelInfo {
+ name: None,
+ inputs: vec![TensorSpec::new("y", DataType::F32, vec![Dimension::Fixed(2)])],
+ outputs: vec![],
+ };
+ assert!(matches!(info.validate_inputs(&inputs), Err(AiError::MissingInput(_))));
+ }
+
+ #[test]
+ fn io_binding_requires_explicit_unique_outputs() {
+ let descriptor = TensorDescriptor::contiguous(DataType::F32, vec![2], Device::CPU);
+ let output = PreallocatedOutput::allocate("scores", descriptor).unwrap();
+ let binding =
+ IoBinding::try_new(NamedTensors::new(), vec![OutputBinding::PreallocatedCpu(output)])
+ .unwrap();
+ assert_eq!(binding.outputs().len(), 1);
+ assert!(IoBinding::try_new(
+ NamedTensors::new(),
+ vec![
+ OutputBinding::Allocate { name: "y".into(), device: Device::CPU },
+ OutputBinding::Allocate { name: "y".into(), device: Device::CPU },
+ ],
+ )
+ .is_err());
+ }
+}
diff --git a/crates/spatialrust-ai/src/onnxruntime.rs b/crates/spatialrust-ai/src/onnxruntime.rs
new file mode 100644
index 0000000..f414f02
--- /dev/null
+++ b/crates/spatialrust-ai/src/onnxruntime.rs
@@ -0,0 +1,759 @@
+//! ONNX Runtime CPU execution provider adapter.
+
+use std::sync::Arc;
+
+use ort::{
+ ep::CPU,
+ memory::MemoryInfo,
+ session::{builder::GraphOptimizationLevel, Session},
+ value::{DynValue, Tensor, TensorElementType, TensorRef, TensorValueType, ValueType},
+};
+use spatialrust_tensor::{DataType, Device, HostTensorStorage, TensorBuffer, TensorDescriptor};
+
+use crate::{
+ AiError, AiResult, CopyPolicy, Dimension, GraphOptimization, InferenceBackend,
+ IoBinding as SpatialIoBinding, ModelInfo, ModelSession, ModelSource, NamedTensors,
+ OutputBinding, PreallocatedStorage, RunOptions, SessionOptions, TensorSpec,
+};
+
+const BACKEND_NAME: &str = "onnxruntime-cpu";
+
+/// ONNX Runtime backend configured explicitly for the CPU execution provider.
+#[derive(Clone, Copy, Debug, Default)]
+pub struct OnnxRuntimeBackend;
+
+impl InferenceBackend for OnnxRuntimeBackend {
+ fn name(&self) -> &str {
+ BACKEND_NAME
+ }
+
+ fn create_session(
+ &self,
+ source: &ModelSource,
+ options: &SessionOptions,
+ ) -> AiResult> {
+ options.validate()?;
+ let mut builder = Session::builder().map_err(ort_error)?;
+ builder = builder.with_execution_providers([CPU::default().build()]).map_err(ort_error)?;
+ if let Some(threads) = options.intra_threads {
+ builder = builder.with_intra_threads(threads).map_err(ort_error)?;
+ }
+ if let Some(threads) = options.inter_threads {
+ builder = builder.with_inter_threads(threads).map_err(ort_error)?;
+ }
+ builder = builder
+ .with_optimization_level(match options.graph_optimization {
+ GraphOptimization::Disabled => GraphOptimizationLevel::Disable,
+ GraphOptimization::Basic => GraphOptimizationLevel::Level1,
+ GraphOptimization::Extended => GraphOptimizationLevel::Level2,
+ GraphOptimization::All => GraphOptimizationLevel::Level3,
+ })
+ .map_err(ort_error)?;
+ builder = builder.with_deterministic_compute(options.deterministic).map_err(ort_error)?;
+ let session = match source {
+ ModelSource::Path(path) => builder.commit_from_file(path).map_err(ort_error)?,
+ ModelSource::Bytes(bytes) => builder.commit_from_memory(bytes).map_err(ort_error)?,
+ };
+ let info = model_info(&session)?;
+ Ok(Box::new(OnnxRuntimeSession { session, info }))
+ }
+}
+
+/// Loaded ONNX Runtime CPU session.
+#[derive(Debug)]
+pub struct OnnxRuntimeSession {
+ session: Session,
+ info: ModelInfo,
+}
+
+impl ModelSession for OnnxRuntimeSession {
+ fn backend_name(&self) -> &str {
+ BACKEND_NAME
+ }
+
+ fn model_info(&self) -> &ModelInfo {
+ &self.info
+ }
+
+ fn run_with_options(
+ &mut self,
+ inputs: NamedTensors,
+ options: RunOptions,
+ ) -> AiResult {
+ self.info.validate_inputs(&inputs)?;
+ if options.input_copy == CopyPolicy::Forbid {
+ return Err(AiError::CopyRequired {
+ direction: "input host",
+ name: self.info.inputs.first().map_or_else(String::new, |spec| spec.name.clone()),
+ });
+ }
+ if options.output_copy == CopyPolicy::Forbid {
+ return Err(AiError::CopyRequired {
+ direction: "output host",
+ name: self.info.outputs.first().map_or_else(String::new, |spec| spec.name.clone()),
+ });
+ }
+ let values = inputs
+ .iter()
+ .map(|(name, tensor)| Ok((name.to_owned(), to_ort_tensor(name, tensor)?)))
+ .collect::>>()?;
+ let outputs = self.session.run(values).map_err(ort_error)?;
+ let mut named = NamedTensors::new();
+ for spec in &self.info.outputs {
+ let value = outputs.get(&spec.name).ok_or_else(|| AiError::Backend {
+ backend: BACKEND_NAME.into(),
+ message: format!("runtime omitted output `{}`", spec.name),
+ })?;
+ named.insert(spec.name.clone(), copy_ort_output(spec, value)?)?;
+ }
+ Ok(named)
+ }
+
+ fn run_with_binding(&mut self, binding: &mut SpatialIoBinding) -> AiResult<()> {
+ binding.clear_results();
+ self.info.validate_inputs(binding.inputs())?;
+ let mut ort_binding = self.session.create_binding().map_err(ort_error)?;
+ for (name, tensor) in binding.inputs().iter() {
+ bind_zero_copy_input(&mut ort_binding, name, tensor)?;
+ }
+
+ let requested = binding
+ .outputs()
+ .iter()
+ .map(|output| match output {
+ OutputBinding::Allocate { name, .. } => name.clone(),
+ OutputBinding::PreallocatedCpu(output) => output.name().to_owned(),
+ })
+ .collect::>();
+ for output in binding.outputs_mut() {
+ match output {
+ OutputBinding::Allocate { name, device } => {
+ output_spec(&self.info, name)?;
+ if *device != Device::CPU {
+ return Err(AiError::Unsupported {
+ backend: BACKEND_NAME.into(),
+ operation: format!("bound output device {device:?}"),
+ });
+ }
+ ort_binding
+ .bind_output_to_device(name.clone(), &MemoryInfo::default())
+ .map_err(ort_error)?;
+ }
+ OutputBinding::PreallocatedCpu(output) => {
+ let spec = output_spec(&self.info, output.name())?;
+ spec.validate(output.descriptor())?;
+ if !output.descriptor().is_c_contiguous()
+ || output.descriptor().byte_offset() != 0
+ {
+ return Err(AiError::Unsupported {
+ backend: BACKEND_NAME.into(),
+ operation: format!(
+ "strided preallocated output `{}`; allocate a compact buffer",
+ output.name()
+ ),
+ });
+ }
+ let name = output.name().to_owned();
+ let shape = output.descriptor().shape().to_vec();
+ bind_preallocated_output(
+ &mut ort_binding,
+ &name,
+ shape,
+ output.take_storage()?,
+ spec.dtype,
+ )?;
+ }
+ }
+ }
+
+ let mut outputs = self.session.run_binding(&ort_binding).map_err(ort_error)?;
+ let mut results = NamedTensors::new();
+ for name in requested {
+ let spec = output_spec(&self.info, &name)?;
+ let value = outputs.remove(&name).ok_or_else(|| AiError::Backend {
+ backend: BACKEND_NAME.into(),
+ message: format!("runtime omitted bound output `{name}`"),
+ })?;
+ results.insert(name, retain_ort_output(spec, value)?)?;
+ }
+ drop(outputs);
+ binding.set_results(results);
+ Ok(())
+ }
+}
+
+fn output_spec<'a>(info: &'a ModelInfo, name: &str) -> AiResult<&'a TensorSpec> {
+ info.outputs
+ .iter()
+ .find(|spec| spec.name == name)
+ .ok_or_else(|| AiError::UnexpectedTensor(name.to_owned()))
+}
+
+fn bind_zero_copy_input(
+ binding: &mut ort::session::IoBinding,
+ name: &str,
+ tensor: &TensorBuffer,
+) -> AiResult<()> {
+ let descriptor = tensor.descriptor();
+ if !descriptor.is_c_contiguous() || descriptor.byte_offset() != 0 {
+ return Err(AiError::Unsupported {
+ backend: BACKEND_NAME.into(),
+ operation: format!("strided bound input `{name}`; explicitly pack it first"),
+ });
+ }
+ if let Some(storage) = tensor.host_storage() {
+ if let Some(storage) = storage.as_any().downcast_ref::() {
+ return storage.bind_input(binding, name);
+ }
+ }
+ let expected = descriptor
+ .element_count()
+ .map_err(|error| AiError::InvalidConfiguration(error.to_string()))?;
+ let shape = descriptor.shape().to_vec();
+ macro_rules! bind_arc {
+ ($getter:ident, $type:ty) => {{
+ let values = tensor.$getter().ok_or_else(|| AiError::CopyRequired {
+ direction: "input alignment",
+ name: name.to_owned(),
+ })?;
+ if values.len() != expected {
+ return Err(AiError::InvalidConfiguration(format!(
+ "bound input `{name}` allocation has {} elements, expected {expected}",
+ values.len()
+ )));
+ }
+ let value = TensorRef::<$type>::from_array_view((shape, values)).map_err(ort_error)?;
+ binding.bind_input(name, &value).map_err(ort_error)
+ }};
+ }
+ match descriptor.dtype() {
+ DataType::U8 => bind_arc!(shared_bytes, u8),
+ DataType::U16 => bind_arc!(shared_u16, u16),
+ DataType::U32 => bind_arc!(shared_u32, u32),
+ DataType::I16 => bind_arc!(shared_i16, i16),
+ DataType::I32 => bind_arc!(shared_i32, i32),
+ DataType::I64 => bind_arc!(shared_i64, i64),
+ DataType::F32 => bind_arc!(shared_f32, f32),
+ DataType::F64 => bind_arc!(shared_f64, f64),
+ dtype => Err(AiError::Unsupported {
+ backend: BACKEND_NAME.into(),
+ operation: format!("zero-copy bound input dtype {dtype:?}"),
+ }),
+ }
+}
+
+fn bind_preallocated_output(
+ binding: &mut ort::session::IoBinding,
+ name: &str,
+ shape: Vec,
+ storage: PreallocatedStorage,
+ dtype: DataType,
+) -> AiResult<()> {
+ match (storage, dtype) {
+ (PreallocatedStorage::Bytes(values), DataType::U8) => binding
+ .bind_output(name, Tensor::::from_array((shape, values)).map_err(ort_error)?)
+ .map_err(ort_error),
+ (PreallocatedStorage::U16(values), DataType::U16) => binding
+ .bind_output(name, Tensor::::from_array((shape, values)).map_err(ort_error)?)
+ .map_err(ort_error),
+ (PreallocatedStorage::F32(values), DataType::F32) => binding
+ .bind_output(name, Tensor::::from_array((shape, values)).map_err(ort_error)?)
+ .map_err(ort_error),
+ _ => Err(AiError::CopyRequired { direction: "output alignment", name: name.to_owned() }),
+ }
+}
+
+#[derive(Debug)]
+enum OrtHostStorage {
+ U8(Tensor),
+ U16(Tensor),
+ U32(Tensor),
+ I8(Tensor),
+ I16(Tensor),
+ I32(Tensor),
+ I64(Tensor),
+ F32(Tensor),
+ F64(Tensor),
+}
+
+impl OrtHostStorage {
+ fn bind_input(&self, binding: &mut ort::session::IoBinding, name: &str) -> AiResult<()> {
+ macro_rules! bind {
+ ($value:expr) => {
+ binding.bind_input(name, $value).map_err(ort_error)
+ };
+ }
+ match self {
+ Self::U8(value) => bind!(value),
+ Self::U16(value) => bind!(value),
+ Self::U32(value) => bind!(value),
+ Self::I8(value) => bind!(value),
+ Self::I16(value) => bind!(value),
+ Self::I32(value) => bind!(value),
+ Self::I64(value) => bind!(value),
+ Self::F32(value) => bind!(value),
+ Self::F64(value) => bind!(value),
+ }
+ }
+}
+
+impl HostTensorStorage for OrtHostStorage {
+ fn dtype(&self) -> DataType {
+ match self {
+ Self::U8(_) => DataType::U8,
+ Self::U16(_) => DataType::U16,
+ Self::U32(_) => DataType::U32,
+ Self::I8(_) => DataType::I8,
+ Self::I16(_) => DataType::I16,
+ Self::I32(_) => DataType::I32,
+ Self::I64(_) => DataType::I64,
+ Self::F32(_) => DataType::F32,
+ Self::F64(_) => DataType::F64,
+ }
+ }
+
+ fn allocation_bytes(&self) -> &[u8] {
+ macro_rules! bytes {
+ ($value:expr) => {
+ bytemuck::cast_slice($value.extract_tensor().1)
+ };
+ }
+ match self {
+ Self::U8(value) => value.extract_tensor().1,
+ Self::U16(value) => bytes!(value),
+ Self::U32(value) => bytes!(value),
+ Self::I8(value) => bytes!(value),
+ Self::I16(value) => bytes!(value),
+ Self::I32(value) => bytes!(value),
+ Self::I64(value) => bytes!(value),
+ Self::F32(value) => bytes!(value),
+ Self::F64(value) => bytes!(value),
+ }
+ }
+
+ fn as_any(&self) -> &dyn std::any::Any {
+ self
+ }
+}
+
+fn retain_ort_output(spec: &TensorSpec, value: DynValue) -> AiResult {
+ macro_rules! retain {
+ ($type:ty, $variant:ident) => {{
+ let tensor = value.downcast::>().map_err(ort_error)?;
+ let shape = tensor
+ .extract_tensor()
+ .0
+ .iter()
+ .map(|&dimension| usize::try_from(dimension).expect("runtime shape is concrete"))
+ .collect::>();
+ let storage: Arc = Arc::new(OrtHostStorage::$variant(tensor));
+ (shape, storage)
+ }};
+ }
+ let (shape, storage) = match spec.dtype {
+ DataType::U8 => retain!(u8, U8),
+ DataType::U16 => retain!(u16, U16),
+ DataType::U32 => retain!(u32, U32),
+ DataType::I8 => retain!(i8, I8),
+ DataType::I16 => retain!(i16, I16),
+ DataType::I32 => retain!(i32, I32),
+ DataType::I64 => retain!(i64, I64),
+ DataType::F32 => retain!(f32, F32),
+ DataType::F64 => retain!(f64, F64),
+ dtype => {
+ return Err(AiError::Unsupported {
+ backend: BACKEND_NAME.into(),
+ operation: format!("retaining bound output dtype {dtype:?}"),
+ })
+ }
+ };
+ TensorBuffer::try_from_host_storage(
+ storage,
+ TensorDescriptor::contiguous(spec.dtype, shape, Device::CPU),
+ )
+ .map_err(|error| AiError::Backend { backend: BACKEND_NAME.into(), message: error.to_string() })
+}
+
+fn model_info(session: &Session) -> AiResult {
+ let inputs = session.inputs().iter().map(outlet_spec).collect::>>()?;
+ let outputs = session.outputs().iter().map(outlet_spec).collect::>>()?;
+ let info = ModelInfo { name: None, inputs, outputs };
+ info.validate()?;
+ Ok(info)
+}
+
+fn outlet_spec(outlet: &ort::value::Outlet) -> AiResult {
+ let ValueType::Tensor { ty, shape, dimension_symbols } = outlet.dtype() else {
+ return Err(AiError::Unsupported {
+ backend: BACKEND_NAME.into(),
+ operation: format!("non-tensor ONNX value `{}`", outlet.name()),
+ });
+ };
+ let dtype = decode_ort_dtype(*ty).ok_or_else(|| AiError::Unsupported {
+ backend: BACKEND_NAME.into(),
+ operation: format!("ONNX dtype {ty} for `{}`", outlet.name()),
+ })?;
+ let dimensions = shape
+ .iter()
+ .enumerate()
+ .map(|(index, &dimension)| {
+ if dimension >= 0 {
+ Dimension::Fixed(dimension as usize)
+ } else {
+ let symbol = &dimension_symbols[index];
+ if symbol.is_empty() {
+ Dimension::Dynamic
+ } else {
+ Dimension::Symbol(symbol.clone())
+ }
+ }
+ })
+ .collect();
+ Ok(TensorSpec::new(outlet.name(), dtype, dimensions))
+}
+
+fn decode_ort_dtype(dtype: TensorElementType) -> Option {
+ Some(match dtype {
+ TensorElementType::Float32 => DataType::F32,
+ TensorElementType::Float64 => DataType::F64,
+ TensorElementType::Uint8 => DataType::U8,
+ TensorElementType::Uint16 => DataType::U16,
+ TensorElementType::Uint32 => DataType::U32,
+ TensorElementType::Int8 => DataType::I8,
+ TensorElementType::Int16 => DataType::I16,
+ TensorElementType::Int32 => DataType::I32,
+ TensorElementType::Int64 => DataType::I64,
+ TensorElementType::Float16 => DataType::F16,
+ TensorElementType::Bfloat16 => DataType::BF16,
+ TensorElementType::Bool => DataType::BOOL,
+ _ => return None,
+ })
+}
+
+fn to_ort_tensor(name: &str, tensor: &TensorBuffer) -> AiResult {
+ let descriptor = tensor.descriptor();
+ if !descriptor.is_c_contiguous() {
+ return Err(AiError::Unsupported {
+ backend: BACKEND_NAME.into(),
+ operation: format!("strided input `{name}`; explicitly pack it first"),
+ });
+ }
+ let range = descriptor.required_byte_range().map_err(|error| AiError::Backend {
+ backend: BACKEND_NAME.into(),
+ message: error.to_string(),
+ })?;
+ let bytes = &tensor.allocation_bytes()[range];
+ let shape = descriptor.shape().to_vec();
+ macro_rules! typed {
+ ($type:ty, $width:expr) => {{
+ let values = decode_ne::<$type>(bytes, $width, |chunk| {
+ <$type>::from_ne_bytes(chunk.try_into().expect("fixed chunk width"))
+ });
+ Tensor::<$type>::from_array((shape, values)).map(|value| value.into_dyn())
+ }};
+ }
+ let value = match descriptor.dtype() {
+ DataType::U8 => Tensor::::from_array((shape, bytes.to_vec())).map(|v| v.into_dyn()),
+ DataType::U16 => typed!(u16, 2),
+ DataType::U32 => typed!(u32, 4),
+ DataType::I8 => Tensor::::from_array((
+ shape,
+ bytes.iter().map(|&value| value as i8).collect::>(),
+ ))
+ .map(|v| v.into_dyn()),
+ DataType::I16 => typed!(i16, 2),
+ DataType::I32 => typed!(i32, 4),
+ DataType::I64 => typed!(i64, 8),
+ DataType::F32 => typed!(f32, 4),
+ DataType::F64 => typed!(f64, 8),
+ dtype => {
+ return Err(AiError::Unsupported {
+ backend: BACKEND_NAME.into(),
+ operation: format!("input dtype {dtype:?}"),
+ })
+ }
+ };
+ value.map_err(ort_error)
+}
+
+fn decode_ne(bytes: &[u8], width: usize, decode: impl Fn(&[u8]) -> T) -> Vec {
+ debug_assert_eq!(bytes.len() % width, 0);
+ bytes.chunks_exact(width).map(decode).collect()
+}
+
+fn copy_ort_output(spec: &TensorSpec, value: &DynValue) -> AiResult {
+ macro_rules! extract {
+ ($type:ty, $constructor:ident) => {{
+ let (shape, values) = value.try_extract_tensor::<$type>().map_err(ort_error)?;
+ let descriptor = TensorDescriptor::contiguous(
+ spec.dtype,
+ shape.iter().map(|&dimension| dimension as usize).collect(),
+ Device::CPU,
+ );
+ TensorBuffer::$constructor(values.to_vec(), descriptor)
+ }};
+ }
+ let result = match spec.dtype {
+ DataType::U8 => extract!(u8, try_new),
+ DataType::U16 => extract!(u16, try_from_u16),
+ DataType::U32 => extract!(u32, try_from_u32),
+ DataType::I8 => {
+ let (shape, values) = value.try_extract_tensor::().map_err(ort_error)?;
+ let descriptor = TensorDescriptor::contiguous(
+ DataType::I8,
+ shape.iter().map(|&dimension| dimension as usize).collect(),
+ Device::CPU,
+ );
+ TensorBuffer::try_new(values.iter().map(|&item| item as u8).collect(), descriptor)
+ }
+ DataType::I16 => extract!(i16, try_from_i16),
+ DataType::I32 => extract!(i32, try_from_i32),
+ DataType::I64 => extract!(i64, try_from_i64),
+ DataType::F32 => extract!(f32, try_from_f32),
+ DataType::F64 => extract!(f64, try_from_f64),
+ dtype => {
+ return Err(AiError::Unsupported {
+ backend: BACKEND_NAME.into(),
+ operation: format!("output dtype {dtype:?}"),
+ })
+ }
+ };
+ result.map_err(|error| AiError::Backend {
+ backend: BACKEND_NAME.into(),
+ message: error.to_string(),
+ })
+}
+
+fn ort_error(error: impl std::fmt::Display) -> AiError {
+ AiError::Backend { backend: BACKEND_NAME.into(), message: error.to_string() }
+}
+
+#[cfg(test)]
+mod tests {
+ use std::sync::Arc;
+
+ use super::OnnxRuntimeBackend;
+ use crate::{
+ AiError, CopyPolicy, Dimension, InferenceBackend, IoBinding, ModelSource, NamedTensors,
+ OutputBinding, PreallocatedOutput, RunOptions, SessionOptions,
+ };
+ use spatialrust_tensor::{DataType, Device, TensorBuffer, TensorDescriptor};
+
+ const DOUBLE_DYNAMIC: &[u8] = &[
+ 8, 8, 18, 16, 115, 112, 97, 116, 105, 97, 108, 114, 117, 115, 116, 45, 116, 101, 115, 116,
+ 58, 106, 10, 27, 10, 5, 105, 110, 112, 117, 116, 10, 5, 105, 110, 112, 117, 116, 18, 6,
+ 111, 117, 116, 112, 117, 116, 34, 3, 65, 100, 100, 18, 14, 100, 111, 117, 98, 108, 101, 95,
+ 100, 121, 110, 97, 109, 105, 99, 90, 28, 10, 5, 105, 110, 112, 117, 116, 18, 19, 10, 17, 8,
+ 1, 18, 13, 10, 7, 18, 5, 98, 97, 116, 99, 104, 10, 2, 8, 3, 98, 29, 10, 6, 111, 117, 116,
+ 112, 117, 116, 18, 19, 10, 17, 8, 1, 18, 13, 10, 7, 18, 5, 98, 97, 116, 99, 104, 10, 2, 8,
+ 3, 66, 4, 10, 0, 16, 13,
+ ];
+
+ fn f32_tensor(shape: Vec, values: &[f32]) -> TensorBuffer {
+ TensorBuffer::try_from_f32(
+ values.to_vec(),
+ TensorDescriptor::contiguous(DataType::F32, shape, Device::CPU),
+ )
+ .unwrap()
+ }
+
+ fn f32_values(tensor: &TensorBuffer) -> Vec {
+ tensor
+ .allocation_bytes()
+ .chunks_exact(4)
+ .map(|chunk| f32::from_ne_bytes(chunk.try_into().unwrap()))
+ .collect()
+ }
+
+ fn dynamic_identity_model_with_onnx_dtype(dtype: u8) -> Arc<[u8]> {
+ let mut model = DOUBLE_DYNAMIC.to_vec();
+ let graph = model
+ .windows(4)
+ .position(|window| window == [58, 106, 10, 27])
+ .expect("embedded graph and node lengths");
+ model[graph + 1] = 104;
+ let node = graph + 2;
+ let identity_node = [
+ 10, 25, 10, 5, 105, 110, 112, 117, 116, 18, 6, 111, 117, 116, 112, 117, 116, 34, 8, 73,
+ 100, 101, 110, 116, 105, 116, 121,
+ ];
+ model.splice(node..node + 29, identity_node);
+ let mut replacements = 0;
+ for index in 0..model.len().saturating_sub(3) {
+ if model[index..index + 4] == [8, 1, 18, 13] {
+ model[index + 1] = dtype;
+ replacements += 1;
+ }
+ }
+ assert_eq!(replacements, 2, "input and output tensor types must both be replaced");
+ Arc::from(model)
+ }
+
+ #[test]
+ fn cpu_session_preserves_named_dynamic_contract_and_requires_copy_opt_in() {
+ let backend = OnnxRuntimeBackend;
+ let mut session = backend
+ .create_session(
+ &ModelSource::Bytes(Arc::from(DOUBLE_DYNAMIC)),
+ &SessionOptions::default(),
+ )
+ .unwrap();
+ assert_eq!(session.model_info().inputs[0].name, "input");
+ assert_eq!(
+ session.model_info().inputs[0].shape,
+ vec![Dimension::Symbol("batch".into()), Dimension::Fixed(3)]
+ );
+
+ let mut inputs = NamedTensors::new();
+ inputs.insert("input", f32_tensor(vec![2, 3], &[1.0, 2.0, 3.0, 4.0, 5.0, 6.0])).unwrap();
+ assert!(matches!(session.run(inputs.clone()), Err(AiError::CopyRequired { .. })));
+ let outputs = session
+ .run_with_options(
+ inputs,
+ RunOptions { input_copy: CopyPolicy::Allow, output_copy: CopyPolicy::Allow },
+ )
+ .unwrap();
+ let output = outputs.get("output").unwrap();
+ assert_eq!(output.descriptor().shape(), &[2, 3]);
+ assert_eq!(f32_values(output), &[2.0, 4.0, 6.0, 8.0, 10.0, 12.0]);
+ }
+
+ #[test]
+ fn io_binding_allocates_dynamic_output_and_reuses_it_as_zero_copy_input() {
+ let backend = OnnxRuntimeBackend;
+ let mut session = backend
+ .create_session(
+ &ModelSource::Bytes(Arc::from(DOUBLE_DYNAMIC)),
+ &SessionOptions::default(),
+ )
+ .unwrap();
+ let mut inputs = NamedTensors::new();
+ inputs.insert("input", f32_tensor(vec![1, 3], &[1.0, 2.0, 3.0])).unwrap();
+ let mut binding = IoBinding::try_new(
+ inputs,
+ vec![OutputBinding::Allocate { name: "output".into(), device: Device::CPU }],
+ )
+ .unwrap();
+ session.run_with_binding(&mut binding).unwrap();
+ let first = binding.results().unwrap().get("output").unwrap();
+ assert!(first.host_storage().is_some());
+ assert_eq!(f32_values(first), &[2.0, 4.0, 6.0]);
+
+ let mut inputs = NamedTensors::new();
+ inputs.insert("input", first.clone()).unwrap();
+ let mut chained = IoBinding::try_new(
+ inputs,
+ vec![OutputBinding::Allocate { name: "output".into(), device: Device::CPU }],
+ )
+ .unwrap();
+ session.run_with_binding(&mut chained).unwrap();
+ assert_eq!(
+ f32_values(chained.results().unwrap().get("output").unwrap()),
+ &[4.0, 8.0, 12.0]
+ );
+ }
+
+ #[test]
+ fn io_binding_writes_directly_into_caller_preallocated_f32_storage() {
+ let backend = OnnxRuntimeBackend;
+ let mut session = backend
+ .create_session(
+ &ModelSource::Bytes(Arc::from(DOUBLE_DYNAMIC)),
+ &SessionOptions::default(),
+ )
+ .unwrap();
+ let mut inputs = NamedTensors::new();
+ inputs.insert("input", f32_tensor(vec![2, 3], &[1.0, 2.0, 3.0, 4.0, 5.0, 6.0])).unwrap();
+ let descriptor = TensorDescriptor::contiguous(DataType::F32, vec![2, 3], Device::CPU);
+ let mut output = PreallocatedOutput::allocate("output", descriptor).unwrap();
+ let allocation = output.allocation_bytes_mut().unwrap().as_ptr();
+ let mut binding =
+ IoBinding::try_new(inputs, vec![OutputBinding::PreallocatedCpu(output)]).unwrap();
+ session.run_with_binding(&mut binding).unwrap();
+ let result = binding.results().unwrap().get("output").unwrap();
+ assert_eq!(result.allocation_bytes().as_ptr(), allocation);
+ assert_eq!(f32_values(result), &[2.0, 4.0, 6.0, 8.0, 10.0, 12.0]);
+ }
+
+ #[test]
+ fn io_binding_rejects_unaligned_raw_f32_storage_instead_of_copying() {
+ let backend = OnnxRuntimeBackend;
+ let mut session = backend
+ .create_session(
+ &ModelSource::Bytes(Arc::from(DOUBLE_DYNAMIC)),
+ &SessionOptions::default(),
+ )
+ .unwrap();
+ let descriptor = TensorDescriptor::contiguous(DataType::F32, vec![1, 3], Device::CPU);
+ let input = TensorBuffer::try_new(vec![0; 12], descriptor).unwrap();
+ let mut inputs = NamedTensors::new();
+ inputs.insert("input", input).unwrap();
+ let mut binding = IoBinding::try_new(
+ inputs,
+ vec![OutputBinding::Allocate { name: "output".into(), device: Device::CPU }],
+ )
+ .unwrap();
+ assert!(matches!(
+ session.run_with_binding(&mut binding),
+ Err(AiError::CopyRequired { direction: "input alignment", .. })
+ ));
+ }
+
+ #[test]
+ fn io_binding_supports_aligned_u8_and_u16_storage() {
+ let backend = OnnxRuntimeBackend;
+
+ let mut u8_session = backend
+ .create_session(
+ &ModelSource::Bytes(dynamic_identity_model_with_onnx_dtype(2)),
+ &SessionOptions::default(),
+ )
+ .unwrap();
+ let mut u8_inputs = NamedTensors::new();
+ u8_inputs
+ .insert(
+ "input",
+ TensorBuffer::try_new(
+ vec![1, 2, 3],
+ TensorDescriptor::contiguous(DataType::U8, vec![1, 3], Device::CPU),
+ )
+ .unwrap(),
+ )
+ .unwrap();
+ let mut u8_binding = IoBinding::try_new(
+ u8_inputs,
+ vec![OutputBinding::Allocate { name: "output".into(), device: Device::CPU }],
+ )
+ .unwrap();
+ u8_session.run_with_binding(&mut u8_binding).unwrap();
+ assert_eq!(
+ u8_binding.results().unwrap().get("output").unwrap().allocation_bytes(),
+ &[1, 2, 3]
+ );
+
+ let mut u16_session = backend
+ .create_session(
+ &ModelSource::Bytes(dynamic_identity_model_with_onnx_dtype(4)),
+ &SessionOptions::default(),
+ )
+ .unwrap();
+ let descriptor = TensorDescriptor::contiguous(DataType::U16, vec![1, 3], Device::CPU);
+ let mut u16_inputs = NamedTensors::new();
+ u16_inputs
+ .insert(
+ "input",
+ TensorBuffer::try_from_u16(vec![10, 20, 30], descriptor.clone()).unwrap(),
+ )
+ .unwrap();
+ let output = PreallocatedOutput::allocate("output", descriptor).unwrap();
+ let mut u16_binding =
+ IoBinding::try_new(u16_inputs, vec![OutputBinding::PreallocatedCpu(output)]).unwrap();
+ u16_session.run_with_binding(&mut u16_binding).unwrap();
+ assert_eq!(
+ bytemuck::cast_slice::(
+ u16_binding.results().unwrap().get("output").unwrap().allocation_bytes()
+ ),
+ &[10, 20, 30]
+ );
+ }
+}
diff --git a/crates/spatialrust-core/src/tensor.rs b/crates/spatialrust-core/src/tensor.rs
index 0bc8f7e..dfc0f60 100644
--- a/crates/spatialrust-core/src/tensor.rs
+++ b/crates/spatialrust-core/src/tensor.rs
@@ -37,7 +37,7 @@ impl<'a> SpatialTensor<'a> {
/// Returns the underlying point cloud.
#[must_use]
- pub const fn cloud(&self) -> &PointCloud {
+ pub const fn cloud(&self) -> &'a PointCloud {
self.cloud
}
diff --git a/crates/spatialrust-image-io/Cargo.toml b/crates/spatialrust-image-io/Cargo.toml
new file mode 100644
index 0000000..093bd45
--- /dev/null
+++ b/crates/spatialrust-image-io/Cargo.toml
@@ -0,0 +1,37 @@
+[package]
+name = "spatialrust-image-io"
+version.workspace = true
+edition.workspace = true
+license.workspace = true
+authors.workspace = true
+repository.workspace = true
+rust-version.workspace = true
+description = "Feature-gated image codecs and bounded stream IO for SpatialRust"
+
+[features]
+default = ["standard"]
+standard = ["png", "jpeg", "pnm", "exif"]
+png = ["image/png"]
+jpeg = ["image/jpeg"]
+pnm = ["image/pnm"]
+tiff = ["image/tiff", "exif"]
+openexr = ["image/openexr", "dep:exr"]
+exif = ["dep:kamadak-exif"]
+full = ["standard", "tiff", "openexr"]
+
+[dependencies]
+spatialrust-image.workspace = true
+image.workspace = true
+kamadak-exif = { workspace = true, optional = true }
+exr = { workspace = true, optional = true }
+thiserror.workspace = true
+
+[dev-dependencies]
+criterion.workspace = true
+proptest.workspace = true
+tempfile.workspace = true
+
+[[bench]]
+name = "decode"
+harness = false
+required-features = ["standard"]
diff --git a/crates/spatialrust-image-io/benches/decode.rs b/crates/spatialrust-image-io/benches/decode.rs
new file mode 100644
index 0000000..c380aeb
--- /dev/null
+++ b/crates/spatialrust-image-io/benches/decode.rs
@@ -0,0 +1,35 @@
+use criterion::{black_box, criterion_group, criterion_main, BenchmarkId, Criterion, Throughput};
+use spatialrust_image::Image;
+use spatialrust_image_io::{
+ decode_bytes, encode_bytes, DecodedPixels, EncodeOptions, ImageFileFormat,
+};
+
+fn encoded_rgb(width: usize, height: usize, format: ImageFileFormat) -> Vec {
+ let data =
+ (0..width * height * 3).map(|index| ((index * 31 + index / 7) & 0xff) as u8).collect();
+ let image = Image::try_new(width, height, data).expect("valid benchmark image");
+ encode_bytes(&DecodedPixels::Rgb8(image), EncodeOptions::new(format))
+ .expect("benchmark fixture encodes")
+}
+
+fn benchmark_decode(c: &mut Criterion) {
+ let mut group = c.benchmark_group("decode_rgb8");
+ group.sample_size(10);
+ for &(name, width, height) in &[("640p", 640, 480), ("1080p", 1920, 1080), ("4k", 3840, 2160)] {
+ for format in [ImageFileFormat::Png, ImageFileFormat::Jpeg] {
+ let encoded = encoded_rgb(width, height, format);
+ group.throughput(Throughput::Elements((width * height) as u64));
+ group.bench_with_input(
+ BenchmarkId::new(format!("{format}/{name}"), encoded.len()),
+ &encoded,
+ |b, bytes| {
+ b.iter(|| decode_bytes(black_box(bytes), Default::default()).unwrap());
+ },
+ );
+ }
+ }
+ group.finish();
+}
+
+criterion_group!(benches, benchmark_decode);
+criterion_main!(benches);
diff --git a/crates/spatialrust-image-io/src/lib.rs b/crates/spatialrust-image-io/src/lib.rs
new file mode 100644
index 0000000..2573dcf
--- /dev/null
+++ b/crates/spatialrust-image-io/src/lib.rs
@@ -0,0 +1,860 @@
+//! Bounded, feature-gated image decoding and encoding for SpatialRust.
+//!
+//! All input is staged in CPU memory with an explicit compressed-input limit.
+//! This crate never uploads to or reads from a device. Codec dependencies are
+//! selected independently from the storage-only `spatialrust-image` crate.
+//!
+//! ```no_run
+//! use spatialrust_image_io::{decode_path, DecodeOptions, DecodedPixels};
+//!
+//! let decoded = decode_path("frame.png", DecodeOptions::default())?;
+//! match decoded.pixels() {
+//! DecodedPixels::Rgb8(image) => assert!(image.width() > 0),
+//! other => println!("decoded {other:?}"),
+//! }
+//! # Ok::<(), spatialrust_image_io::ImageIoError>(())
+//! ```
+
+#![deny(unsafe_code)]
+#![warn(missing_docs)]
+
+use std::fs::File;
+use std::io::{BufReader, BufWriter, Cursor, Read, Seek, Write};
+use std::path::Path;
+
+use image::{DynamicImage, GenericImageView, ImageBuffer, ImageFormat, ImageOutputFormat};
+use spatialrust_image::{AlphaMode, ColorRange, ColorSpace, Image, ImageError, ImageMetadata};
+
+/// Errors produced by bounded image IO.
+#[derive(Debug, thiserror::Error)]
+pub enum ImageIoError {
+ /// Reading or writing the stream failed.
+ #[error(transparent)]
+ Io(#[from] std::io::Error),
+ /// A codec rejected the image or encoded output.
+ #[error(transparent)]
+ Codec(#[from] image::ImageError),
+ /// Constructing a typed SpatialRust image failed.
+ #[error(transparent)]
+ Image(#[from] ImageError),
+ /// The stream exceeded the configured compressed-input limit.
+ #[error("encoded input exceeds the {maximum} byte limit")]
+ InputTooLarge {
+ /// Maximum accepted encoded byte count.
+ maximum: usize,
+ },
+ /// Decoded dimensions or pixel count exceeded the configured limits.
+ #[error("decoded image dimensions {width}x{height} exceed configured limits")]
+ DimensionsTooLarge {
+ /// Encoded image width.
+ width: u32,
+ /// Encoded image height.
+ height: u32,
+ },
+ /// The detected/requested format is not enabled in this build.
+ #[error("image format `{0}` is not enabled in this build")]
+ FormatDisabled(ImageFileFormat),
+ /// The input format is not one of SpatialRust's image-IO formats.
+ #[error("unsupported image format: {0}")]
+ UnsupportedFormat(String),
+ /// Encoding options were invalid.
+ #[error("invalid encode option: {0}")]
+ InvalidEncodeOption(String),
+ /// A decoded dimension cannot be represented by this platform.
+ #[error("decoded dimensions cannot be represented by usize")]
+ DimensionOverflow,
+}
+
+/// File formats recognized by the stable SpatialRust image-IO boundary.
+#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
+pub enum ImageFileFormat {
+ /// Portable Network Graphics.
+ Png,
+ /// JPEG/JFIF image.
+ Jpeg,
+ /// Portable anymap family (PBM/PGM/PPM/PAM).
+ Pnm,
+ /// Tagged Image File Format.
+ Tiff,
+ /// OpenEXR high-dynamic-range image.
+ OpenExr,
+}
+
+impl std::fmt::Display for ImageFileFormat {
+ fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
+ formatter.write_str(match self {
+ Self::Png => "png",
+ Self::Jpeg => "jpeg",
+ Self::Pnm => "pnm",
+ Self::Tiff => "tiff",
+ Self::OpenExr => "openexr",
+ })
+ }
+}
+
+impl ImageFileFormat {
+ /// Returns whether this codec is compiled into the current crate.
+ #[must_use]
+ pub const fn is_enabled(self) -> bool {
+ match self {
+ Self::Png => cfg!(feature = "png"),
+ Self::Jpeg => cfg!(feature = "jpeg"),
+ Self::Pnm => cfg!(feature = "pnm"),
+ Self::Tiff => cfg!(feature = "tiff"),
+ Self::OpenExr => cfg!(feature = "openexr"),
+ }
+ }
+
+ fn from_backend(format: ImageFormat) -> Result {
+ match format {
+ ImageFormat::Png => Ok(Self::Png),
+ ImageFormat::Jpeg => Ok(Self::Jpeg),
+ ImageFormat::Pnm => Ok(Self::Pnm),
+ ImageFormat::Tiff => Ok(Self::Tiff),
+ ImageFormat::OpenExr => Ok(Self::OpenExr),
+ other => Err(ImageIoError::UnsupportedFormat(format!("{other:?}"))),
+ }
+ }
+
+ fn require_enabled(self) -> Result<(), ImageIoError> {
+ if self.is_enabled() {
+ Ok(())
+ } else {
+ Err(ImageIoError::FormatDisabled(self))
+ }
+ }
+}
+
+/// Strict application-level resource limits for one decode operation.
+#[derive(Clone, Copy, Debug, PartialEq, Eq)]
+pub struct DecodeLimits {
+ /// Maximum compressed input size in bytes.
+ pub max_input_bytes: usize,
+ /// Maximum decoded width.
+ pub max_width: u32,
+ /// Maximum decoded height.
+ pub max_height: u32,
+ /// Maximum decoded pixel count.
+ pub max_pixels: u64,
+ /// Best-effort maximum simultaneous allocation performed by the backend.
+ pub max_alloc_bytes: u64,
+}
+
+impl Default for DecodeLimits {
+ fn default() -> Self {
+ Self {
+ max_input_bytes: 256 * 1024 * 1024,
+ max_width: 32_768,
+ max_height: 32_768,
+ max_pixels: 100_000_000,
+ max_alloc_bytes: 512 * 1024 * 1024,
+ }
+ }
+}
+
+impl DecodeLimits {
+ fn validate_dimensions(self, width: u32, height: u32) -> Result<(), ImageIoError> {
+ let pixels = u64::from(width).saturating_mul(u64::from(height));
+ if width > self.max_width || height > self.max_height || pixels > self.max_pixels {
+ return Err(ImageIoError::DimensionsTooLarge { width, height });
+ }
+ Ok(())
+ }
+
+ fn backend(self) -> image::io::Limits {
+ let mut limits = image::io::Limits::default();
+ limits.max_image_width = Some(self.max_width);
+ limits.max_image_height = Some(self.max_height);
+ limits.max_alloc = Some(self.max_alloc_bytes);
+ limits
+ }
+}
+
+/// Options controlling one decode operation.
+#[derive(Clone, Copy, Debug, PartialEq, Eq)]
+pub struct DecodeOptions {
+ /// Resource limits applied before and during decoding.
+ pub limits: DecodeLimits,
+ /// Apply Exif orientation to produce canonical top-left pixels.
+ pub apply_orientation: bool,
+}
+
+impl Default for DecodeOptions {
+ fn default() -> Self {
+ Self { limits: DecodeLimits::default(), apply_orientation: true }
+ }
+}
+
+/// Exif image orientation in encoded-pixel coordinates.
+#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Hash)]
+#[repr(u8)]
+pub enum Orientation {
+ /// No orientation tag was found.
+ #[default]
+ Unspecified = 0,
+ /// Top-left origin; no transform.
+ Normal = 1,
+ /// Mirror left-to-right.
+ FlipHorizontal = 2,
+ /// Rotate 180 degrees.
+ Rotate180 = 3,
+ /// Mirror top-to-bottom.
+ FlipVertical = 4,
+ /// Reflect across the top-left to bottom-right diagonal.
+ Transpose = 5,
+ /// Rotate 90 degrees clockwise.
+ Rotate90 = 6,
+ /// Reflect across the top-right to bottom-left diagonal.
+ Transverse = 7,
+ /// Rotate 270 degrees clockwise.
+ Rotate270 = 8,
+}
+
+impl Orientation {
+ #[cfg(feature = "exif")]
+ fn from_exif(value: u32) -> Self {
+ match value {
+ 1 => Self::Normal,
+ 2 => Self::FlipHorizontal,
+ 3 => Self::Rotate180,
+ 4 => Self::FlipVertical,
+ 5 => Self::Transpose,
+ 6 => Self::Rotate90,
+ 7 => Self::Transverse,
+ 8 => Self::Rotate270,
+ _ => Self::Unspecified,
+ }
+ }
+
+ fn changes_pixels(self) -> bool {
+ !matches!(self, Self::Unspecified | Self::Normal)
+ }
+}
+
+/// Encoded sample/channel type before conversion into SpatialRust storage.
+#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
+pub enum SourceColorType {
+ /// 8-bit grayscale.
+ Gray8,
+ /// 8-bit grayscale plus alpha.
+ GrayAlpha8,
+ /// 8-bit RGB.
+ Rgb8,
+ /// 8-bit RGBA.
+ Rgba8,
+ /// 16-bit grayscale.
+ Gray16,
+ /// 16-bit grayscale plus alpha.
+ GrayAlpha16,
+ /// 16-bit RGB.
+ Rgb16,
+ /// 16-bit RGBA.
+ Rgba16,
+ /// 32-bit floating-point RGB.
+ Rgb32Float,
+ /// 32-bit floating-point RGBA.
+ Rgba32Float,
+}
+
+impl SourceColorType {
+ fn from_backend(color: image::ColorType) -> Result {
+ match color {
+ image::ColorType::L8 => Ok(Self::Gray8),
+ image::ColorType::La8 => Ok(Self::GrayAlpha8),
+ image::ColorType::Rgb8 => Ok(Self::Rgb8),
+ image::ColorType::Rgba8 => Ok(Self::Rgba8),
+ image::ColorType::L16 => Ok(Self::Gray16),
+ image::ColorType::La16 => Ok(Self::GrayAlpha16),
+ image::ColorType::Rgb16 => Ok(Self::Rgb16),
+ image::ColorType::Rgba16 => Ok(Self::Rgba16),
+ image::ColorType::Rgb32F => Ok(Self::Rgb32Float),
+ image::ColorType::Rgba32F => Ok(Self::Rgba32Float),
+ other => Err(ImageIoError::UnsupportedFormat(format!(
+ "unsupported decoded color type {other:?}"
+ ))),
+ }
+ }
+}
+
+/// Owned typed pixels produced by a decoder or accepted by an encoder.
+#[derive(Clone, Debug, PartialEq)]
+pub enum DecodedPixels {
+ /// One-channel 8-bit grayscale.
+ Gray8(Image),
+ /// Two-channel 8-bit grayscale and alpha.
+ GrayAlpha8(Image),
+ /// Three-channel 8-bit RGB.
+ Rgb8(Image),
+ /// Four-channel 8-bit RGBA.
+ Rgba8(Image),
+ /// One-channel 16-bit grayscale.
+ Gray16(Image),
+ /// Two-channel 16-bit grayscale and alpha.
+ GrayAlpha16(Image),
+ /// Three-channel 16-bit RGB.
+ Rgb16(Image),
+ /// Four-channel 16-bit RGBA.
+ Rgba16(Image),
+ /// Three-channel linear-light floating-point RGB.
+ Rgb32Float(Image),
+ /// Four-channel linear-light floating-point RGBA.
+ Rgba32Float(Image),
+}
+
+impl DecodedPixels {
+ /// Returns pixel width.
+ #[must_use]
+ pub fn width(&self) -> usize {
+ match self {
+ Self::Gray8(image) => image.width(),
+ Self::GrayAlpha8(image) => image.width(),
+ Self::Rgb8(image) => image.width(),
+ Self::Rgba8(image) => image.width(),
+ Self::Gray16(image) => image.width(),
+ Self::GrayAlpha16(image) => image.width(),
+ Self::Rgb16(image) => image.width(),
+ Self::Rgba16(image) => image.width(),
+ Self::Rgb32Float(image) => image.width(),
+ Self::Rgba32Float(image) => image.width(),
+ }
+ }
+
+ /// Returns pixel height.
+ #[must_use]
+ pub fn height(&self) -> usize {
+ match self {
+ Self::Gray8(image) => image.height(),
+ Self::GrayAlpha8(image) => image.height(),
+ Self::Rgb8(image) => image.height(),
+ Self::Rgba8(image) => image.height(),
+ Self::Gray16(image) => image.height(),
+ Self::GrayAlpha16(image) => image.height(),
+ Self::Rgb16(image) => image.height(),
+ Self::Rgba16(image) => image.height(),
+ Self::Rgb32Float(image) => image.height(),
+ Self::Rgba32Float(image) => image.height(),
+ }
+ }
+
+ fn to_dynamic(&self) -> Result {
+ let width = u32::try_from(self.width()).map_err(|_| ImageIoError::DimensionOverflow)?;
+ let height = u32::try_from(self.height()).map_err(|_| ImageIoError::DimensionOverflow)?;
+ macro_rules! buffer {
+ ($image:expr, $pixel:ty, $variant:path) => {{
+ let typed = ImageBuffer::<$pixel, Vec<_>>::from_raw(
+ width,
+ height,
+ $image.as_slice().to_vec(),
+ )
+ .ok_or(ImageIoError::DimensionOverflow)?;
+ $variant(typed)
+ }};
+ }
+ Ok(match self {
+ Self::Gray8(image) => buffer!(image, image::Luma, DynamicImage::ImageLuma8),
+ Self::GrayAlpha8(image) => {
+ buffer!(image, image::LumaA, DynamicImage::ImageLumaA8)
+ }
+ Self::Rgb8(image) => buffer!(image, image::Rgb, DynamicImage::ImageRgb8),
+ Self::Rgba8(image) => buffer!(image, image::Rgba, DynamicImage::ImageRgba8),
+ Self::Gray16(image) => buffer!(image, image::Luma, DynamicImage::ImageLuma16),
+ Self::GrayAlpha16(image) => {
+ buffer!(image, image::LumaA, DynamicImage::ImageLumaA16)
+ }
+ Self::Rgb16(image) => buffer!(image, image::Rgb, DynamicImage::ImageRgb16),
+ Self::Rgba16(image) => buffer!(image, image::Rgba, DynamicImage::ImageRgba16),
+ Self::Rgb32Float(image) => {
+ buffer!(image, image::Rgb, DynamicImage::ImageRgb32F)
+ }
+ Self::Rgba32Float(image) => {
+ buffer!(image, image::Rgba, DynamicImage::ImageRgba32F)
+ }
+ })
+ }
+}
+
+/// Metadata describing encoded pixels and orientation handling.
+#[derive(Clone, Copy, Debug, PartialEq, Eq)]
+pub struct DecodedMetadata {
+ /// Detected container/codec format.
+ pub format: ImageFileFormat,
+ /// Sample and channel representation reported by the decoder.
+ pub source_color_type: SourceColorType,
+ /// Orientation tag found in the encoded container.
+ pub orientation: Orientation,
+ /// Whether a non-identity orientation was applied to the returned pixels.
+ pub orientation_applied: bool,
+}
+
+/// One decoded image with owned typed pixels and source metadata.
+#[derive(Clone, Debug, PartialEq)]
+pub struct DecodedImage {
+ pixels: DecodedPixels,
+ metadata: DecodedMetadata,
+}
+
+impl DecodedImage {
+ /// Borrows decoded pixels.
+ #[must_use]
+ pub fn pixels(&self) -> &DecodedPixels {
+ &self.pixels
+ }
+
+ /// Returns source metadata.
+ #[must_use]
+ pub const fn metadata(&self) -> DecodedMetadata {
+ self.metadata
+ }
+
+ /// Consumes the wrapper and returns pixels.
+ #[must_use]
+ pub fn into_pixels(self) -> DecodedPixels {
+ self.pixels
+ }
+
+ /// Returns decoded width after orientation handling.
+ #[must_use]
+ pub fn width(&self) -> usize {
+ self.pixels.width()
+ }
+
+ /// Returns decoded height after orientation handling.
+ #[must_use]
+ pub fn height(&self) -> usize {
+ self.pixels.height()
+ }
+}
+
+/// Options controlling encoded output.
+#[derive(Clone, Copy, Debug, PartialEq, Eq)]
+pub struct EncodeOptions {
+ /// Output codec/container.
+ pub format: ImageFileFormat,
+ /// JPEG quality in the inclusive range 1–100.
+ pub jpeg_quality: u8,
+}
+
+impl EncodeOptions {
+ /// Creates options for one format with JPEG quality 90.
+ #[must_use]
+ pub const fn new(format: ImageFileFormat) -> Self {
+ Self { format, jpeg_quality: 90 }
+ }
+
+ fn output_format(self) -> Result {
+ self.format.require_enabled()?;
+ if self.jpeg_quality == 0 || self.jpeg_quality > 100 {
+ return Err(ImageIoError::InvalidEncodeOption(
+ "jpeg_quality must be in 1..=100".to_owned(),
+ ));
+ }
+ match self.format {
+ #[cfg(feature = "png")]
+ ImageFileFormat::Png => Ok(ImageOutputFormat::Png),
+ #[cfg(not(feature = "png"))]
+ ImageFileFormat::Png => Err(ImageIoError::FormatDisabled(self.format)),
+ #[cfg(feature = "jpeg")]
+ ImageFileFormat::Jpeg => Ok(ImageOutputFormat::Jpeg(self.jpeg_quality)),
+ #[cfg(not(feature = "jpeg"))]
+ ImageFileFormat::Jpeg => Err(ImageIoError::FormatDisabled(self.format)),
+ #[cfg(feature = "pnm")]
+ ImageFileFormat::Pnm => {
+ Ok(ImageOutputFormat::Pnm(image::codecs::pnm::PnmSubtype::ArbitraryMap))
+ }
+ #[cfg(not(feature = "pnm"))]
+ ImageFileFormat::Pnm => Err(ImageIoError::FormatDisabled(self.format)),
+ #[cfg(feature = "tiff")]
+ ImageFileFormat::Tiff => Ok(ImageOutputFormat::Tiff),
+ #[cfg(not(feature = "tiff"))]
+ ImageFileFormat::Tiff => Err(ImageIoError::FormatDisabled(self.format)),
+ #[cfg(feature = "openexr")]
+ ImageFileFormat::OpenExr => Ok(ImageOutputFormat::OpenExr),
+ #[cfg(not(feature = "openexr"))]
+ ImageFileFormat::OpenExr => Err(ImageIoError::FormatDisabled(self.format)),
+ }
+ }
+}
+
+/// Decodes an image from an encoded byte slice.
+pub fn decode_bytes(bytes: &[u8], options: DecodeOptions) -> Result {
+ if bytes.len() > options.limits.max_input_bytes {
+ return Err(ImageIoError::InputTooLarge { maximum: options.limits.max_input_bytes });
+ }
+ let backend_format = image::guess_format(bytes)?;
+ let format = ImageFileFormat::from_backend(backend_format)?;
+ format.require_enabled()?;
+
+ let dimensions =
+ image::io::Reader::with_format(Cursor::new(bytes), backend_format).into_dimensions()?;
+ options.limits.validate_dimensions(dimensions.0, dimensions.1)?;
+
+ let orientation = read_orientation(bytes);
+ let mut reader = image::io::Reader::with_format(Cursor::new(bytes), backend_format);
+ reader.limits(options.limits.backend());
+ let dynamic = reader.decode()?;
+ let source_color_type = SourceColorType::from_backend(dynamic.color())?;
+ let orientation_applied = options.apply_orientation && orientation.changes_pixels();
+ let dynamic =
+ if orientation_applied { apply_orientation(dynamic, orientation) } else { dynamic };
+ let pixels = dynamic_to_pixels(dynamic)?;
+ Ok(DecodedImage {
+ pixels,
+ metadata: DecodedMetadata { format, source_color_type, orientation, orientation_applied },
+ })
+}
+
+/// Reads a bounded encoded stream and decodes it.
+pub fn decode_reader(
+ reader: R,
+ options: DecodeOptions,
+) -> Result {
+ let maximum = options.limits.max_input_bytes;
+ let take_limit = u64::try_from(maximum).unwrap_or(u64::MAX).saturating_add(1);
+ let mut bytes = Vec::new();
+ reader.take(take_limit).read_to_end(&mut bytes)?;
+ if bytes.len() > maximum {
+ return Err(ImageIoError::InputTooLarge { maximum });
+ }
+ decode_bytes(&bytes, options)
+}
+
+/// Opens and decodes one image path using content-based format detection.
+pub fn decode_path(
+ path: impl AsRef,
+ options: DecodeOptions,
+) -> Result {
+ decode_reader(BufReader::new(File::open(path)?), options)
+}
+
+/// Encodes typed pixels to a seekable writer.
+///
+/// The codec backend receives a temporary owned CPU buffer. No device transfer
+/// occurs, and the source image remains unchanged.
+pub fn encode_writer(
+ writer: &mut W,
+ pixels: &DecodedPixels,
+ options: EncodeOptions,
+) -> Result<(), ImageIoError> {
+ let dynamic = pixels.to_dynamic()?;
+ dynamic.write_to(writer, options.output_format()?)?;
+ Ok(())
+}
+
+/// Encodes typed pixels into an owned byte vector.
+pub fn encode_bytes(
+ pixels: &DecodedPixels,
+ options: EncodeOptions,
+) -> Result, ImageIoError> {
+ let mut cursor = Cursor::new(Vec::new());
+ encode_writer(&mut cursor, pixels, options)?;
+ Ok(cursor.into_inner())
+}
+
+/// Creates or truncates a path and encodes typed pixels into it.
+pub fn encode_path(
+ path: impl AsRef,
+ pixels: &DecodedPixels,
+ options: EncodeOptions,
+) -> Result<(), ImageIoError> {
+ let mut writer = BufWriter::new(File::create(path)?);
+ encode_writer(&mut writer, pixels, options)?;
+ writer.flush()?;
+ Ok(())
+}
+
+fn metadata(color_space: ColorSpace, alpha_mode: AlphaMode, floating: bool) -> ImageMetadata {
+ ImageMetadata {
+ color_space,
+ color_range: if floating { ColorRange::Unspecified } else { ColorRange::Full },
+ alpha_mode,
+ }
+}
+
+fn dimensions(dynamic: &DynamicImage) -> Result<(usize, usize), ImageIoError> {
+ let (width, height) = dynamic.dimensions();
+ Ok((
+ usize::try_from(width).map_err(|_| ImageIoError::DimensionOverflow)?,
+ usize::try_from(height).map_err(|_| ImageIoError::DimensionOverflow)?,
+ ))
+}
+
+fn dynamic_to_pixels(dynamic: DynamicImage) -> Result {
+ let (width, height) = dimensions(&dynamic)?;
+ Ok(match dynamic {
+ DynamicImage::ImageLuma8(image) => DecodedPixels::Gray8(Image::try_new_with_metadata(
+ width,
+ height,
+ image.into_raw(),
+ metadata(ColorSpace::Gray, AlphaMode::None, false),
+ )?),
+ DynamicImage::ImageLumaA8(image) => {
+ DecodedPixels::GrayAlpha8(Image::try_new_with_metadata(
+ width,
+ height,
+ image.into_raw(),
+ metadata(ColorSpace::Unknown, AlphaMode::Straight, false),
+ )?)
+ }
+ DynamicImage::ImageRgb8(image) => DecodedPixels::Rgb8(Image::try_new_with_metadata(
+ width,
+ height,
+ image.into_raw(),
+ metadata(ColorSpace::Rgb, AlphaMode::None, false),
+ )?),
+ DynamicImage::ImageRgba8(image) => DecodedPixels::Rgba8(Image::try_new_with_metadata(
+ width,
+ height,
+ image.into_raw(),
+ metadata(ColorSpace::Rgba, AlphaMode::Straight, false),
+ )?),
+ DynamicImage::ImageLuma16(image) => DecodedPixels::Gray16(Image::try_new_with_metadata(
+ width,
+ height,
+ image.into_raw(),
+ metadata(ColorSpace::Gray, AlphaMode::None, false),
+ )?),
+ DynamicImage::ImageLumaA16(image) => {
+ DecodedPixels::GrayAlpha16(Image::try_new_with_metadata(
+ width,
+ height,
+ image.into_raw(),
+ metadata(ColorSpace::Unknown, AlphaMode::Straight, false),
+ )?)
+ }
+ DynamicImage::ImageRgb16(image) => DecodedPixels::Rgb16(Image::try_new_with_metadata(
+ width,
+ height,
+ image.into_raw(),
+ metadata(ColorSpace::Rgb, AlphaMode::None, false),
+ )?),
+ DynamicImage::ImageRgba16(image) => DecodedPixels::Rgba16(Image::try_new_with_metadata(
+ width,
+ height,
+ image.into_raw(),
+ metadata(ColorSpace::Rgba, AlphaMode::Straight, false),
+ )?),
+ DynamicImage::ImageRgb32F(image) => {
+ DecodedPixels::Rgb32Float(Image::try_new_with_metadata(
+ width,
+ height,
+ image.into_raw(),
+ metadata(ColorSpace::LinearRgb, AlphaMode::None, true),
+ )?)
+ }
+ DynamicImage::ImageRgba32F(image) => {
+ DecodedPixels::Rgba32Float(Image::try_new_with_metadata(
+ width,
+ height,
+ image.into_raw(),
+ metadata(ColorSpace::Unknown, AlphaMode::Straight, true),
+ )?)
+ }
+ _ => {
+ return Err(ImageIoError::UnsupportedFormat(
+ "decoder returned an unsupported dynamic image variant".to_owned(),
+ ))
+ }
+ })
+}
+
+fn apply_orientation(image: DynamicImage, orientation: Orientation) -> DynamicImage {
+ match orientation {
+ Orientation::Unspecified | Orientation::Normal => image,
+ Orientation::FlipHorizontal => image.fliph(),
+ Orientation::Rotate180 => image.rotate180(),
+ Orientation::FlipVertical => image.flipv(),
+ Orientation::Transpose => image.rotate90().fliph(),
+ Orientation::Rotate90 => image.rotate90(),
+ Orientation::Transverse => image.rotate90().flipv(),
+ Orientation::Rotate270 => image.rotate270(),
+ }
+}
+
+#[cfg(feature = "exif")]
+fn read_orientation(bytes: &[u8]) -> Orientation {
+ use exif::{In, Reader, Tag};
+
+ let mut cursor = Cursor::new(bytes);
+ Reader::new()
+ .read_from_container(&mut cursor)
+ .ok()
+ .and_then(|exif| {
+ exif.get_field(Tag::Orientation, In::PRIMARY).and_then(|field| field.value.get_uint(0))
+ })
+ .map(Orientation::from_exif)
+ .unwrap_or(Orientation::Unspecified)
+}
+
+#[cfg(not(feature = "exif"))]
+fn read_orientation(_bytes: &[u8]) -> Orientation {
+ Orientation::Unspecified
+}
+
+#[cfg(test)]
+mod tests {
+ use super::*;
+ use proptest::prelude::*;
+
+ #[cfg(any(feature = "png", feature = "jpeg", feature = "pnm"))]
+ fn rgb_fixture() -> DecodedPixels {
+ DecodedPixels::Rgb8(
+ Image::try_new_with_metadata(
+ 3,
+ 2,
+ vec![255, 0, 0, 0, 255, 0, 0, 0, 255, 10, 20, 30, 40, 50, 60, 70, 80, 90],
+ metadata(ColorSpace::Rgb, AlphaMode::None, false),
+ )
+ .unwrap(),
+ )
+ }
+
+ #[cfg(feature = "png")]
+ #[test]
+ fn png_memory_roundtrip_is_exact() {
+ let source = rgb_fixture();
+ let bytes = encode_bytes(&source, EncodeOptions::new(ImageFileFormat::Png)).unwrap();
+ let decoded = decode_bytes(&bytes, DecodeOptions::default()).unwrap();
+ assert_eq!(decoded.pixels(), &source);
+ assert_eq!(decoded.metadata().format, ImageFileFormat::Png);
+ assert_eq!(decoded.metadata().source_color_type, SourceColorType::Rgb8);
+ }
+
+ #[cfg(feature = "pnm")]
+ #[test]
+ fn pnm_reader_roundtrip_is_exact() {
+ let source = rgb_fixture();
+ let bytes = encode_bytes(&source, EncodeOptions::new(ImageFileFormat::Pnm)).unwrap();
+ let decoded = decode_reader(Cursor::new(bytes), DecodeOptions::default()).unwrap();
+ assert_eq!(decoded.pixels(), &source);
+ }
+
+ #[cfg(feature = "jpeg")]
+ #[test]
+ fn jpeg_roundtrip_preserves_shape_and_type() {
+ let bytes =
+ encode_bytes(&rgb_fixture(), EncodeOptions::new(ImageFileFormat::Jpeg)).unwrap();
+ let decoded = decode_bytes(&bytes, DecodeOptions::default()).unwrap();
+ assert_eq!((decoded.width(), decoded.height()), (3, 2));
+ assert!(matches!(decoded.pixels(), DecodedPixels::Rgb8(_)));
+ }
+
+ #[cfg(feature = "png")]
+ #[test]
+ fn dimensions_and_input_are_bounded_before_decode() {
+ let bytes = encode_bytes(&rgb_fixture(), EncodeOptions::new(ImageFileFormat::Png)).unwrap();
+ let mut options = DecodeOptions::default();
+ options.limits.max_width = 2;
+ assert!(matches!(
+ decode_bytes(&bytes, options),
+ Err(ImageIoError::DimensionsTooLarge { .. }) | Err(ImageIoError::Codec(_))
+ ));
+
+ let mut options = DecodeOptions::default();
+ options.limits.max_input_bytes = bytes.len() - 1;
+ assert!(matches!(
+ decode_reader(Cursor::new(bytes), options),
+ Err(ImageIoError::InputTooLarge { .. })
+ ));
+ }
+
+ #[test]
+ fn rotation_and_diagonal_reflection_have_expected_layout() {
+ let image = image::GrayImage::from_raw(2, 3, vec![1, 2, 3, 4, 5, 6]).unwrap();
+ let rotated =
+ apply_orientation(DynamicImage::ImageLuma8(image.clone()), Orientation::Rotate90);
+ assert_eq!(rotated.dimensions(), (3, 2));
+ assert_eq!(rotated.to_luma8().into_raw(), vec![5, 3, 1, 6, 4, 2]);
+
+ let transposed = apply_orientation(DynamicImage::ImageLuma8(image), Orientation::Transpose);
+ assert_eq!(transposed.dimensions(), (3, 2));
+ assert_eq!(transposed.to_luma8().into_raw(), vec![1, 3, 5, 2, 4, 6]);
+ }
+
+ #[cfg(feature = "exif")]
+ #[test]
+ fn reads_orientation_from_minimal_tiff_exif() {
+ let bytes = [
+ b'I', b'I', 42, 0, 8, 0, 0, 0, // TIFF header and first IFD offset
+ 1, 0, // one entry
+ 0x12, 0x01, // Orientation tag
+ 3, 0, // SHORT
+ 1, 0, 0, 0, // count
+ 6, 0, 0, 0, // Rotate 90 clockwise
+ 0, 0, 0, 0, // next IFD
+ ];
+ assert_eq!(read_orientation(&bytes), Orientation::Rotate90);
+ }
+
+ #[test]
+ fn rejects_malformed_input() {
+ assert!(decode_bytes(b"not an image", DecodeOptions::default()).is_err());
+ }
+
+ #[test]
+ fn disabled_format_reports_feature_boundary() {
+ if !ImageFileFormat::Tiff.is_enabled() {
+ assert!(matches!(
+ EncodeOptions::new(ImageFileFormat::Tiff).output_format(),
+ Err(ImageIoError::FormatDisabled(ImageFileFormat::Tiff))
+ ));
+ }
+ }
+
+ proptest! {
+ #[test]
+ fn arbitrary_small_input_never_panics(bytes in proptest::collection::vec(any::(), 0..4096)) {
+ let mut options = DecodeOptions::default();
+ options.limits.max_input_bytes = 4096;
+ options.limits.max_width = 256;
+ options.limits.max_height = 256;
+ options.limits.max_pixels = 65_536;
+ options.limits.max_alloc_bytes = 4 * 1024 * 1024;
+ let _ = decode_bytes(&bytes, options);
+ }
+ }
+
+ #[cfg(feature = "png")]
+ #[test]
+ fn path_roundtrip_uses_content_detection() {
+ let directory = tempfile::tempdir().unwrap();
+ let path = directory.path().join("image-without-extension");
+ let source = rgb_fixture();
+ encode_path(&path, &source, EncodeOptions::new(ImageFileFormat::Png)).unwrap();
+ let decoded = decode_path(path, DecodeOptions::default()).unwrap();
+ assert_eq!(decoded.pixels(), &source);
+ }
+
+ #[cfg(feature = "tiff")]
+ #[test]
+ fn tiff_preserves_sixteen_bit_grayscale() {
+ let source = DecodedPixels::Gray16(
+ Image::try_new_with_metadata(
+ 3,
+ 1,
+ vec![0, 1024, u16::MAX],
+ metadata(ColorSpace::Gray, AlphaMode::None, false),
+ )
+ .unwrap(),
+ );
+ let bytes = encode_bytes(&source, EncodeOptions::new(ImageFileFormat::Tiff)).unwrap();
+ let decoded = decode_bytes(&bytes, DecodeOptions::default()).unwrap();
+ assert_eq!(decoded.pixels(), &source);
+ }
+
+ #[cfg(feature = "openexr")]
+ #[test]
+ fn openexr_preserves_float_rgb() {
+ let source = DecodedPixels::Rgb32Float(
+ Image::try_new_with_metadata(
+ 2,
+ 1,
+ vec![0.0, 0.25, 1.0, 4.0, -0.5, 2.0],
+ metadata(ColorSpace::LinearRgb, AlphaMode::None, true),
+ )
+ .unwrap(),
+ );
+ let bytes = encode_bytes(&source, EncodeOptions::new(ImageFileFormat::OpenExr)).unwrap();
+ let decoded = decode_bytes(&bytes, DecodeOptions::default()).unwrap();
+ assert_eq!(decoded.pixels(), &source);
+ }
+}
diff --git a/crates/spatialrust-math/src/mat.rs b/crates/spatialrust-math/src/mat.rs
index 43918f7..1fc8350 100644
--- a/crates/spatialrust-math/src/mat.rs
+++ b/crates/spatialrust-math/src/mat.rs
@@ -118,6 +118,46 @@ impl Mat3 {
self.m[2][0] * v.x + self.m[2][1] * v.y + self.m[2][2] * v.z,
)
}
+
+ /// Matrix multiplication.
+ #[must_use]
+ pub fn mul_mat3(self, other: Self) -> Self {
+ Self::from_rows(
+ [
+ self.m[0][0] * other.m[0][0]
+ + self.m[0][1] * other.m[1][0]
+ + self.m[0][2] * other.m[2][0],
+ self.m[0][0] * other.m[0][1]
+ + self.m[0][1] * other.m[1][1]
+ + self.m[0][2] * other.m[2][1],
+ self.m[0][0] * other.m[0][2]
+ + self.m[0][1] * other.m[1][2]
+ + self.m[0][2] * other.m[2][2],
+ ],
+ [
+ self.m[1][0] * other.m[0][0]
+ + self.m[1][1] * other.m[1][0]
+ + self.m[1][2] * other.m[2][0],
+ self.m[1][0] * other.m[0][1]
+ + self.m[1][1] * other.m[1][1]
+ + self.m[1][2] * other.m[2][1],
+ self.m[1][0] * other.m[0][2]
+ + self.m[1][1] * other.m[1][2]
+ + self.m[1][2] * other.m[2][2],
+ ],
+ [
+ self.m[2][0] * other.m[0][0]
+ + self.m[2][1] * other.m[1][0]
+ + self.m[2][2] * other.m[2][0],
+ self.m[2][0] * other.m[0][1]
+ + self.m[2][1] * other.m[1][1]
+ + self.m[2][2] * other.m[2][1],
+ self.m[2][0] * other.m[0][2]
+ + self.m[2][1] * other.m[1][2]
+ + self.m[2][2] * other.m[2][2],
+ ],
+ )
+ }
}
impl Mat4 {
diff --git a/crates/spatialrust-py/Cargo.toml b/crates/spatialrust-py/Cargo.toml
index e0c1838..674ac56 100644
--- a/crates/spatialrust-py/Cargo.toml
+++ b/crates/spatialrust-py/Cargo.toml
@@ -9,6 +9,10 @@ authors = ["SpatialRust Contributors"]
repository = "https://github.com/rsasaki0109/SpatialRust"
description = "Python bindings for SpatialRust point cloud processing"
+[features]
+default = []
+onnxruntime = ["spatialrust/ai-onnxruntime"]
+
[lib]
# Importable module name: `import spatialrust`
name = "spatialrust"
@@ -42,6 +46,8 @@ spatialrust = { path = "../spatialrust", features = [
"register-fpfh",
"camera-rgbd",
"vision-full",
+ "image-io-standard",
+ "tensor-dlpack",
] }
# Keep this crate out of the main Rust workspace so `cargo test --workspace`
diff --git a/crates/spatialrust-py/README.md b/crates/spatialrust-py/README.md
index 0752a40..91777f9 100644
--- a/crates/spatialrust-py/README.md
+++ b/crates/spatialrust-py/README.md
@@ -24,6 +24,26 @@ maturin develop --release # builds the Rust extension into the venv
maturin build --release --out dist
```
+ONNX Runtime is intentionally not part of the default wheel. Enable its CPU
+backend explicitly when building an inference wheel:
+
+```bash
+maturin develop --release --features onnxruntime
+```
+
+`OnnxRuntimeSession.run()` uses named CPU I/O Binding by default. Pass
+`copy=True` only when an explicit host conversion is acceptable.
+
+Feature2D is available in the default wheel. `harris_keypoints`,
+`shi_tomasi_keypoints`, and `fast_keypoints` return immutable keypoint metadata;
+`orb_features` returns those keypoints with a `uint8[N, 32]` descriptor matrix.
+Use `match_binary_descriptors` for Hamming distance or
+`match_float_descriptors` for Euclidean distance, with optional ratio,
+cross-check, and maximum-distance filters.
+
+Geometry bindings expose `estimate_homography_ransac`, `solve_pnp`, and
+`stereo_block_match` for NumPy `float64` / grayscale workflows.
+
## Test
The bindings have a pytest suite (`tests/`) that exercises the NumPy ⇄ Rust
diff --git a/crates/spatialrust-py/spatialrust.pyi b/crates/spatialrust-py/spatialrust.pyi
index 68f2e24..b9af992 100644
--- a/crates/spatialrust-py/spatialrust.pyi
+++ b/crates/spatialrust-py/spatialrust.pyi
@@ -11,13 +11,68 @@ import numpy as np
from numpy.typing import NDArray
__version__: str
+__all__: list[str] = [
+ "__version__", "ImageMetadata", "Tensor", "Keypoint2", "OnnxRuntimeSession",
+ "DLPackTensorView", "PointCloud", "PipelineResult", "RegionResult",
+ "DbscanResult", "GroundResult", "MultiPlaneResult", "SphereResult",
+ "CylinderResult", "RegistrationResult", "read_image",
+ "tensor_copy_from_numpy", "tensor_view_from_dlpack", "harris_keypoints",
+ "shi_tomasi_keypoints", "fast_keypoints", "orb_features",
+ "estimate_homography_ransac", "solve_pnp", "stereo_block_match",
+ "match_binary_descriptors", "match_float_descriptors", "write_image", "read",
+ "write", "voxel_downsample", "crop_box", "pass_through", "iss_keypoints",
+ "orient_normals", "detect_boundary", "mls_smooth", "farthest_point_sampling",
+ "statistical_outlier_removal", "radius_outlier_removal", "run_pipeline",
+ "region_growing", "dbscan", "ground_segmentation", "segment_multi_plane",
+ "ransac_sphere", "ransac_cylinder", "chamfer_distance", "hausdorff_distance",
+ "apply_transform", "recenter", "scale", "normalize_unit_sphere", "merge",
+ "centroid", "bounding_box", "oriented_bounding_box", "voxelize", "range_image",
+ "rgbd_to_point_cloud", "filter2d_image", "gaussian_blur_image",
+ "median_blur_image", "bilateral_filter_image", "sobel_image", "scharr_image",
+ "laplacian_image", "pyr_down_image", "pyr_up_image", "morphology_image",
+ "threshold_image", "otsu_threshold_image", "adaptive_threshold_image",
+ "histogram_image", "equalize_histogram_image", "clahe_image",
+ "integral_image_u8", "canny_image", "resize_image", "letterbox_image",
+ "normalize_image_chw", "rgb_to_gray_image", "rgb_to_hsv_image", "remap_image",
+ "nms", "soft_nms", "connected_components_image", "find_mask_contours",
+ "encode_mask_rle", "decode_mask_rle", "point_map_to_point_cloud", "knn_graph",
+ "radius_graph", "register_icp", "register_point_to_plane", "register_gicp",
+ "register_ndt", "register_fpfh_ransac", "register_fpfh_keypoints",
+]
# Convenient aliases for the array shapes the bindings exchange.
_F32Array = NDArray[np.float32] # positions, grids, range images, transforms
+_F64Array = NDArray[np.float64]
+_BoolArray = NDArray[np.bool_]
_I32Array = NDArray[np.int32] # labels, edge_index
_U32Array = NDArray[np.uint32]
_Vec3 = tuple[float, float, float]
_U8Array = NDArray[np.uint8]
+_U16Array = NDArray[np.uint16]
+
+@final
+class ImageMetadata:
+ """Container, sample type, and Exif orientation from image decoding."""
+
+ @property
+ def format(self) -> str: ...
+ @property
+ def color_type(self) -> str: ...
+ @property
+ def orientation(self) -> int: ...
+ @property
+ def orientation_applied(self) -> bool: ...
+ def __repr__(self) -> str: ...
+
+def read_image(
+ path: str, apply_orientation: bool = ...
+) -> tuple[NDArray[np.uint8] | NDArray[np.uint16] | NDArray[np.float32], ImageMetadata]: ...
+def write_image(
+ path: str,
+ image: NDArray[np.uint8] | NDArray[np.uint16],
+ format: str,
+ jpeg_quality: int = ...,
+) -> None: ...
@final
class PointCloud:
@@ -61,6 +116,87 @@ def rgbd_to_point_cloud(
# --------------------------------------------------------------------------- #
# Image preprocessing, detection, masks, and dense spatial data
# --------------------------------------------------------------------------- #
+def filter2d_image(
+ image: _U8Array, kernel: NDArray[np.float64], delta: float = ...
+) -> _U8Array: ...
+def gaussian_blur_image(
+ image: _U8Array,
+ kernel_width: int,
+ kernel_height: int,
+ sigma_x: float,
+ sigma_y: Optional[float] = ...,
+) -> _U8Array: ...
+def median_blur_image(image: _U8Array, kernel_size: int) -> _U8Array: ...
+def bilateral_filter_image(
+ image: _U8Array,
+ diameter: int,
+ sigma_color: float,
+ sigma_space: float,
+) -> _U8Array: ...
+def sobel_image(
+ image: _U8Array,
+ dx: int,
+ dy: int,
+ kernel_size: int = ...,
+ scale: float = ...,
+ delta: float = ...,
+) -> _F32Array: ...
+def scharr_image(
+ image: _U8Array,
+ dx: int,
+ dy: int,
+ scale: float = ...,
+ delta: float = ...,
+) -> _F32Array: ...
+def laplacian_image(
+ image: _U8Array,
+ kernel_size: int = ...,
+ scale: float = ...,
+ delta: float = ...,
+) -> _F32Array: ...
+def pyr_down_image(image: _U8Array) -> _U8Array: ...
+def pyr_up_image(image: _U8Array) -> _U8Array: ...
+def morphology_image(
+ image: _U8Array,
+ operation: str,
+ kernel_width: int,
+ kernel_height: int,
+ shape: str = ...,
+ iterations: int = ...,
+) -> _U8Array: ...
+def threshold_image(
+ image: _U8Array,
+ threshold: float,
+ max_value: int = ...,
+ threshold_type: str = ...,
+) -> _U8Array: ...
+def otsu_threshold_image(
+ image: _U8Array, max_value: int = ..., threshold_type: str = ...
+) -> tuple[int, _U8Array]: ...
+def adaptive_threshold_image(
+ image: _U8Array,
+ block_size: int,
+ c: float,
+ method: str = ...,
+ max_value: int = ...,
+ threshold_type: str = ...,
+) -> _U8Array: ...
+def histogram_image(image: _U8Array) -> NDArray[np.uint64]: ...
+def equalize_histogram_image(image: _U8Array) -> _U8Array: ...
+def clahe_image(
+ image: _U8Array,
+ clip_limit: float = ...,
+ tiles_x: int = ...,
+ tiles_y: int = ...,
+) -> _U8Array: ...
+def integral_image_u8(image: _U8Array) -> NDArray[np.float64]: ...
+def canny_image(
+ image: _U8Array,
+ low_threshold: float,
+ high_threshold: float,
+ aperture_size: int = ...,
+ l2_gradient: bool = ...,
+) -> _U8Array: ...
def resize_image(
image: _U8Array,
width: int,
@@ -434,3 +570,138 @@ def register_fpfh_keypoints(
ransac_iterations: int = ...,
k_neighbors: int = ...,
) -> RegistrationResult: ...
+@final
+class Tensor:
+ @property
+ def shape(self) -> list[int]: ...
+ @property
+ def dtype(self) -> str: ...
+ def __dlpack_device__(self) -> tuple[int, int]: ...
+ def __dlpack__(
+ self,
+ stream: object | None = ...,
+ *,
+ max_version: tuple[int, int] | None = ...,
+ dl_device: tuple[int, int] | None = ...,
+ copy: bool | None = ...,
+ ) -> object: ...
+ def copy(self) -> Tensor: ...
+
+@final
+class Keypoint2:
+ @property
+ def x(self) -> float: ...
+ @property
+ def y(self) -> float: ...
+ @property
+ def size(self) -> float: ...
+ @property
+ def angle_degrees(self) -> float | None: ...
+ @property
+ def response(self) -> float: ...
+ @property
+ def octave(self) -> int: ...
+ @property
+ def class_id(self) -> int | None: ...
+
+def harris_keypoints(
+ image: _U8Array,
+ max_corners: int = ...,
+ quality_level: float = ...,
+ min_distance: float = ...,
+ block_size: int = ...,
+ gradient_size: int = ...,
+ k: float = ...,
+) -> list[Keypoint2]: ...
+def shi_tomasi_keypoints(
+ image: _U8Array,
+ max_corners: int = ...,
+ quality_level: float = ...,
+ min_distance: float = ...,
+ block_size: int = ...,
+ gradient_size: int = ...,
+) -> list[Keypoint2]: ...
+def fast_keypoints(
+ image: _U8Array,
+ threshold: int = ...,
+ nonmax_suppression: bool = ...,
+) -> list[Keypoint2]: ...
+def orb_features(
+ image: _U8Array,
+ max_features: int = ...,
+ scale_factor: float = ...,
+ levels: int = ...,
+ edge_threshold: int = ...,
+ fast_threshold: int = ...,
+ patch_size: int = ...,
+ score_type: str = ...,
+) -> tuple[list[Keypoint2], _U8Array]: ...
+def estimate_homography_ransac(
+ source: _F64Array,
+ target: _F64Array,
+ threshold: float = ...,
+ confidence: float = ...,
+ max_iterations: int = ...,
+ seed: int = ...,
+) -> tuple[_F64Array, _BoolArray, _F64Array]: ...
+def solve_pnp(
+ object_points: _F64Array,
+ image_points: _F64Array,
+ fx: float,
+ fy: float,
+ cx: float,
+ cy: float,
+ width: int = ...,
+ height: int = ...,
+) -> tuple[_F64Array, _F64Array]: ...
+def stereo_block_match(
+ left: _U8Array,
+ right: _U8Array,
+ window_size: int = ...,
+ min_disparity: int = ...,
+ num_disparities: int = ...,
+ uniqueness_ratio: float = ...,
+) -> _F32Array: ...
+def match_binary_descriptors(
+ query: _U8Array,
+ train: _U8Array,
+ cross_check: bool = ...,
+ ratio: float | None = ...,
+ max_distance: float | None = ...,
+) -> list[tuple[int, int, float]]: ...
+def match_float_descriptors(
+ query: _F32Array,
+ train: _F32Array,
+ cross_check: bool = ...,
+ ratio: float | None = ...,
+ max_distance: float | None = ...,
+) -> list[tuple[int, int, float]]: ...
+
+@final
+class OnnxRuntimeSession:
+ def __new__(
+ cls,
+ path: str,
+ *,
+ intra_threads: int | None = ...,
+ inter_threads: int | None = ...,
+ deterministic: bool = ...,
+ ) -> OnnxRuntimeSession: ...
+ @property
+ def inputs(self) -> list[tuple[str, str, list[str]]]: ...
+ @property
+ def outputs(self) -> list[tuple[str, str, list[str]]]: ...
+ def run(self, inputs: dict[str, Tensor], *, copy: bool = ...) -> dict[str, Tensor]: ...
+
+def tensor_copy_from_numpy(array: NDArray[np.generic]) -> Tensor: ...
+@final
+class DLPackTensorView:
+ @property
+ def shape(self) -> list[int]: ...
+ @property
+ def dtype(self) -> str: ...
+ @property
+ def version(self) -> tuple[int, int]: ...
+ def copy(self) -> Tensor: ...
+
+def tensor_view_from_dlpack(producer: object) -> DLPackTensorView: ...
diff --git a/crates/spatialrust-py/src/dlpack_capsule.rs b/crates/spatialrust-py/src/dlpack_capsule.rs
new file mode 100644
index 0000000..d4c114e
--- /dev/null
+++ b/crates/spatialrust-py/src/dlpack_capsule.rs
@@ -0,0 +1,66 @@
+//! Audited CPython capsule boundary for DLPack ownership transfer.
+
+use pyo3::{
+ exceptions::PyBufferError,
+ prelude::*,
+ types::{PyAny, PyDict},
+};
+use spatialrust::tensor::{release_dlpack_raw, DlpackExport, DlpackImport, TensorBuffer};
+
+const VERSIONED_NAME: &[u8] = b"dltensor_versioned\0";
+const USED_VERSIONED_NAME: &[u8] = b"used_dltensor_versioned\0";
+
+unsafe extern "C" fn capsule_destructor(capsule: *mut pyo3::ffi::PyObject) {
+ let name = VERSIONED_NAME.as_ptr().cast();
+ // SAFETY: CPython calls the destructor with the capsule object. A consumer
+ // renames consumed capsules, so only the original name retains ownership.
+ if unsafe { pyo3::ffi::PyCapsule_IsValid(capsule, name) } == 1 {
+ // SAFETY: validity above proves the name and capsule pointer contract.
+ let raw = unsafe { pyo3::ffi::PyCapsule_GetPointer(capsule, name) };
+ if !raw.is_null() {
+ // SAFETY: the unconsumed capsule uniquely owns deleter responsibility.
+ unsafe { release_dlpack_raw(raw) };
+ }
+ }
+}
+
+pub(crate) fn export_tensor(py: Python<'_>, tensor: &TensorBuffer) -> PyResult> {
+ let export = DlpackExport::from_tensor(tensor)
+ .map_err(|error| PyBufferError::new_err(error.to_string()))?;
+ let raw = export.into_raw();
+ // SAFETY: `raw` is live and the static nul-terminated name outlives the capsule.
+ let capsule = unsafe {
+ pyo3::ffi::PyCapsule_New(raw, VERSIONED_NAME.as_ptr().cast(), Some(capsule_destructor))
+ };
+ if capsule.is_null() {
+ // SAFETY: capsule construction failed, so ownership was not transferred.
+ unsafe { release_dlpack_raw(raw) };
+ return Err(PyErr::fetch(py));
+ }
+ // SAFETY: PyCapsule_New returned one new owned Python reference.
+ Ok(unsafe { Py::from_owned_ptr(py, capsule) })
+}
+
+pub(crate) fn import_tensor(producer: &Bound<'_, PyAny>) -> PyResult {
+ let kwargs = PyDict::new_bound(producer.py());
+ kwargs.set_item("max_version", (1, 0))?;
+ kwargs.set_item("copy", false)?;
+ let capsule = producer.call_method("__dlpack__", (), Some(&kwargs))?;
+ let name = VERSIONED_NAME.as_ptr().cast();
+ // SAFETY: this only asks CPython to validate the exact capsule/name pair.
+ if unsafe { pyo3::ffi::PyCapsule_IsValid(capsule.as_ptr(), name) } != 1 {
+ return Err(PyBufferError::new_err("producer did not return a dltensor_versioned capsule"));
+ }
+ // SAFETY: capsule validity proves that the payload is non-null for this name.
+ let raw = unsafe { pyo3::ffi::PyCapsule_GetPointer(capsule.as_ptr(), name) };
+ // SAFETY: both names are static nul-terminated strings and capsule is valid.
+ if unsafe {
+ pyo3::ffi::PyCapsule_SetName(capsule.as_ptr(), USED_VERSIONED_NAME.as_ptr().cast())
+ } != 0
+ {
+ return Err(PyErr::fetch(producer.py()));
+ }
+ // SAFETY: renaming transferred exclusive deleter ownership from the capsule.
+ unsafe { DlpackImport::from_raw(raw) }
+ .map_err(|error| PyBufferError::new_err(error.to_string()))
+}
diff --git a/crates/spatialrust-py/src/lib.rs b/crates/spatialrust-py/src/lib.rs
index 8fd4d69..98620b8 100644
--- a/crates/spatialrust-py/src/lib.rs
+++ b/crates/spatialrust-py/src/lib.rs
@@ -6,13 +6,26 @@
// PyO3's `#[pyfunction]` expansion emits `.into()` on already-`PyErr` results.
#![allow(clippy::useless_conversion)]
+#![deny(unsafe_code)]
+
+#[allow(unsafe_code)]
+mod dlpack_capsule;
use numpy::ndarray::{Array2, Array3};
use numpy::{
- IntoPyArray, PyArray1, PyArray2, PyArray3, PyReadonlyArray1, PyReadonlyArray2, PyReadonlyArray3,
+ IntoPyArray, PyArray1, PyArray2, PyArray3, PyReadonlyArray1, PyReadonlyArray2,
+ PyReadonlyArray3, PyReadonlyArrayDyn, PyUntypedArrayMethods,
};
-use pyo3::exceptions::PyValueError;
+use pyo3::exceptions::{PyBufferError, PyRuntimeError, PyTypeError, PyValueError};
use pyo3::prelude::*;
+use pyo3::types::PyDict;
+
+#[cfg(feature = "onnxruntime")]
+use spatialrust::ai::{
+ CopyPolicy as AiCopyPolicy, InferenceBackend, IoBinding as AiIoBinding, ModelSession,
+ ModelSource, NamedTensors as AiNamedTensors, OnnxRuntimeBackend, OutputBinding,
+ RunOptions as AiRunOptions, SessionOptions as AiSessionOptions,
+};
use spatialrust::core::{PointBuffer, PointBufferSet, SpatialMetadata};
use spatialrust::features::{
@@ -26,7 +39,11 @@ use spatialrust::filtering::{
StatisticalOutlierConfig, StatisticalOutlierRemoval, VoxelGridDownsample,
VoxelGridDownsampleConfig,
};
-use spatialrust::math::Mat4;
+use spatialrust::image_io::{
+ decode_path as decode_image_path, encode_path as encode_image_path, DecodeOptions,
+ DecodedMetadata, DecodedPixels, EncodeOptions, ImageFileFormat,
+};
+use spatialrust::math::{Mat3, Mat4, Vec2, Vec3};
use spatialrust::metrics::{chamfer_distance as chamfer, hausdorff_distance as hausdorff};
use spatialrust::pipeline::{MvpPipeline, MvpPipelineConfig};
use spatialrust::registration::{
@@ -39,19 +56,39 @@ use spatialrust::segmentation::{
MultiPlaneSegmenter, RansacCylinderSegmenter, RansacPrimitiveConfig, RansacSphereSegmenter,
RegionGrowingConfig, RegionGrowingSegmenter,
};
+use spatialrust::tensor::{
+ DataType as TensorDataType, Device as TensorDevice, DlpackImport, TensorBuffer,
+ TensorDescriptor,
+};
use spatialrust::transform::{
apply_transform as apply_tf, bounding_box as bbox, centroid as cloud_centroid, merge_clouds,
normalize_unit_sphere as normalize_unit, oriented_bounding_box as obb, recenter as recenter_op,
scale_cloud,
};
use spatialrust::vision::{
- approximate_polygon as approximate_contour, connected_components as label_components,
- decode_rle as decode_mask_runs, encode_rle as encode_mask_runs,
- find_contours as trace_contours, letterbox as letterbox_op, nms as nms_op,
- pack_chw as pack_chw_op, point_map_to_point_cloud as point_map_to_cloud, remap as remap_op,
- resize as resize_op, rgb_to_gray as rgb_to_gray_op, rgb_to_hsv as rgb_to_hsv_op,
- soft_nms as soft_nms_op, BinaryMask, BorderMode, BoundingBox2, ConfidenceMap, Connectivity,
- Interpolation, MaskRle, PointMap, RleOrder, SoftNmsMethod,
+ adaptive_threshold as adaptive_threshold_op, approximate_polygon as approximate_contour,
+ bilateral_filter as bilateral_filter_op, canny as canny_op, clahe as clahe_op,
+ connected_components as label_components, decode_rle as decode_mask_runs,
+ detect_and_describe_orb as detect_and_describe_orb_op, detect_fast as detect_fast_op,
+ detect_harris as detect_harris_op, detect_shi_tomasi as detect_shi_tomasi_op,
+ encode_rle as encode_mask_runs, equalize_histogram as equalize_histogram_op,
+ estimate_homography_ransac as estimate_homography_ransac_op, filter2d as filter2d_op,
+ find_contours as trace_contours, gaussian_blur as gaussian_blur_op,
+ histogram_u8 as histogram_u8_op, integral_image as integral_image_op,
+ laplacian as laplacian_op, letterbox as letterbox_op,
+ match_descriptors as match_descriptors_op, median_blur as median_blur_op,
+ morphology_ex as morphology_ex_op, nms as nms_op, otsu_threshold_u8 as otsu_threshold_u8_op,
+ pack_chw as pack_chw_op, point_map_to_point_cloud as point_map_to_cloud,
+ pyr_down as pyr_down_op, pyr_up as pyr_up_op, remap as remap_op, resize as resize_op,
+ rgb_to_gray as rgb_to_gray_op, rgb_to_hsv as rgb_to_hsv_op, scharr as scharr_op,
+ sobel as sobel_op, soft_nms as soft_nms_op, solve_pnp as solve_pnp_op,
+ stereo_block_match as stereo_block_match_op, threshold as threshold_op,
+ AdaptiveThresholdMethod, BinaryMask, BorderMode, BoundingBox2, CameraMatrix3, CannyOptions,
+ ConfidenceMap, Connectivity, CornerSelectionOptions, DescriptorBuffer, FastOptions,
+ HarrisOptions, Interpolation, Kernel2D, Keypoint2, MaskRle, MatchOptions,
+ MorphologyOperation, MorphologyShape, ObjectImageCorrespondence, OrbOptions, OrbScoreType,
+ PointCorrespondence2, PointMap, RobustEstimationOptions, RleOrder, ShiTomasiOptions,
+ SoftNmsMethod, StereoBmOptions, StructuringElement, ThresholdType, AbsolutePose,
};
use spatialrust::voxelize::{
range_image as range_image_proj, voxelize as voxelize_grid, RangeImageConfig, VoxelFill,
@@ -77,6 +114,516 @@ fn to_py_err(err: E) -> PyErr {
PyValueError::new_err(err.to_string())
}
+#[pyclass(name = "Tensor")]
+#[derive(Clone)]
+struct PyTensor {
+ inner: TensorBuffer,
+}
+
+#[pyclass(name = "Keypoint2", frozen)]
+#[derive(Clone, Copy)]
+struct PyKeypoint2 {
+ inner: Keypoint2,
+}
+
+#[pymethods]
+impl PyKeypoint2 {
+ #[getter]
+ fn x(&self) -> f32 {
+ self.inner.x()
+ }
+
+ #[getter]
+ fn y(&self) -> f32 {
+ self.inner.y()
+ }
+
+ #[getter]
+ fn size(&self) -> f32 {
+ self.inner.size()
+ }
+
+ #[getter]
+ fn angle_degrees(&self) -> Option {
+ self.inner.angle_degrees()
+ }
+
+ #[getter]
+ fn response(&self) -> f32 {
+ self.inner.response()
+ }
+
+ #[getter]
+ fn octave(&self) -> i32 {
+ self.inner.octave()
+ }
+
+ #[getter]
+ fn class_id(&self) -> Option {
+ self.inner.class_id()
+ }
+
+ fn __repr__(&self) -> String {
+ format!(
+ "Keypoint2(x={}, y={}, size={}, angle_degrees={:?}, response={})",
+ self.x(),
+ self.y(),
+ self.size(),
+ self.angle_degrees(),
+ self.response()
+ )
+ }
+}
+
+#[pyclass(name = "OnnxRuntimeSession", unsendable)]
+struct PyOnnxRuntimeSession {
+ #[cfg(feature = "onnxruntime")]
+ inner: Box,
+}
+
+#[pymethods]
+impl PyOnnxRuntimeSession {
+ #[new]
+ #[pyo3(signature = (path, *, intra_threads=None, inter_threads=None, deterministic=false))]
+ fn new(
+ path: String,
+ intra_threads: Option,
+ inter_threads: Option,
+ deterministic: bool,
+ ) -> PyResult {
+ #[cfg(feature = "onnxruntime")]
+ {
+ let options = AiSessionOptions {
+ intra_threads,
+ inter_threads,
+ deterministic,
+ ..AiSessionOptions::default()
+ };
+ let inner = OnnxRuntimeBackend
+ .create_session(&ModelSource::Path(path.into()), &options)
+ .map_err(to_py_err)?;
+ Ok(Self { inner })
+ }
+ #[cfg(not(feature = "onnxruntime"))]
+ {
+ let _ = (path, intra_threads, inter_threads, deterministic);
+ Err(PyRuntimeError::new_err(
+ "this SpatialRust Python module was built without the `onnxruntime` feature",
+ ))
+ }
+ }
+
+ /// Returns `(name, dtype, dimensions)` metadata for model inputs.
+ #[getter]
+ fn inputs(&self) -> Vec<(String, String, Vec)> {
+ #[cfg(feature = "onnxruntime")]
+ {
+ self.inner.model_info().inputs.iter().map(python_tensor_spec).collect()
+ }
+ #[cfg(not(feature = "onnxruntime"))]
+ {
+ Vec::new()
+ }
+ }
+
+ /// Returns `(name, dtype, dimensions)` metadata for model outputs.
+ #[getter]
+ fn outputs(&self) -> Vec<(String, String, Vec)> {
+ #[cfg(feature = "onnxruntime")]
+ {
+ self.inner.model_info().outputs.iter().map(python_tensor_spec).collect()
+ }
+ #[cfg(not(feature = "onnxruntime"))]
+ {
+ Vec::new()
+ }
+ }
+
+ /// Runs named tensors; zero-copy CPU I/O binding is the default.
+ #[pyo3(signature = (inputs, *, copy=false))]
+ fn run<'py>(
+ &mut self,
+ py: Python<'py>,
+ inputs: &Bound<'py, PyDict>,
+ copy: bool,
+ ) -> PyResult> {
+ #[cfg(feature = "onnxruntime")]
+ {
+ let mut named = AiNamedTensors::new();
+ for (name, value) in inputs.iter() {
+ let name = name.extract::()?;
+ let tensor = value.extract::>()?;
+ named.insert(name, tensor.inner.clone()).map_err(to_py_err)?;
+ }
+ let outputs = if copy {
+ self.inner
+ .run_with_options(
+ named,
+ AiRunOptions {
+ input_copy: AiCopyPolicy::Allow,
+ output_copy: AiCopyPolicy::Allow,
+ },
+ )
+ .map_err(to_py_err)?
+ } else {
+ let destinations = self
+ .inner
+ .model_info()
+ .outputs
+ .iter()
+ .map(|spec| OutputBinding::Allocate {
+ name: spec.name.clone(),
+ device: TensorDevice::CPU,
+ })
+ .collect();
+ let mut binding = AiIoBinding::try_new(named, destinations).map_err(to_py_err)?;
+ self.inner.run_with_binding(&mut binding).map_err(to_py_err)?;
+ binding.into_results().ok_or_else(|| {
+ PyRuntimeError::new_err("ONNX Runtime completed without bound results")
+ })?
+ };
+ let result = PyDict::new_bound(py);
+ for (name, tensor) in outputs.into_values() {
+ result.set_item(name, Py::new(py, PyTensor { inner: tensor })?)?;
+ }
+ Ok(result)
+ }
+ #[cfg(not(feature = "onnxruntime"))]
+ {
+ let _ = (py, inputs, copy);
+ Err(PyRuntimeError::new_err(
+ "this SpatialRust Python module was built without the `onnxruntime` feature",
+ ))
+ }
+ }
+}
+
+#[cfg(feature = "onnxruntime")]
+fn python_tensor_spec(spec: &spatialrust::ai::TensorSpec) -> (String, String, Vec) {
+ let dimensions = spec
+ .shape
+ .iter()
+ .map(|dimension| match dimension {
+ spatialrust::ai::Dimension::Fixed(value) => value.to_string(),
+ spatialrust::ai::Dimension::Dynamic => "?".into(),
+ spatialrust::ai::Dimension::Symbol(value) => value.clone(),
+ })
+ .collect();
+ (spec.name.clone(), tensor_dtype_name(spec.dtype), dimensions)
+}
+
+#[pyclass(name = "DLPackTensorView", unsendable)]
+struct PyDlpackTensorView {
+ inner: DlpackImport,
+}
+
+#[pymethods]
+impl PyDlpackTensorView {
+ #[getter]
+ fn shape(&self) -> Vec {
+ self.inner.descriptor().shape().to_vec()
+ }
+
+ #[getter]
+ fn dtype(&self) -> String {
+ tensor_dtype_name(self.inner.descriptor().dtype())
+ }
+
+ #[getter]
+ fn version(&self) -> (u32, u32) {
+ self.inner.version()
+ }
+
+ /// Makes ownership independent of the DLPack producer with an explicit copy.
+ fn copy(&self) -> PyResult {
+ Ok(PyTensor { inner: self.inner.view().map_err(to_py_err)?.to_owned_copy() })
+ }
+
+ fn __repr__(&self) -> String {
+ format!(
+ "DLPackTensorView(shape={:?}, dtype='{}', device='cpu', version={:?})",
+ self.shape(),
+ self.dtype(),
+ self.version()
+ )
+ }
+}
+
+fn tensor_dtype_name(dtype: TensorDataType) -> String {
+ match dtype {
+ TensorDataType::U8 => "uint8".into(),
+ TensorDataType::U16 => "uint16".into(),
+ TensorDataType::F32 => "float32".into(),
+ _ => format!("{:?}{}x{}", dtype.code(), dtype.bits(), dtype.lanes()),
+ }
+}
+
+#[pymethods]
+impl PyTensor {
+ #[getter]
+ fn shape(&self) -> Vec {
+ self.inner.descriptor().shape().to_vec()
+ }
+
+ #[getter]
+ fn dtype(&self) -> String {
+ tensor_dtype_name(self.inner.descriptor().dtype())
+ }
+
+ /// Returns the DLPack CPU device tuple.
+ fn __dlpack_device__(&self) -> (i32, i32) {
+ (1, 0)
+ }
+
+ /// Exports a read-only, zero-copy DLPack major-version 1 capsule.
+ #[pyo3(signature = (stream=None, *, max_version=None, dl_device=None, copy=None))]
+ fn __dlpack__(
+ &self,
+ py: Python<'_>,
+ stream: Option>,
+ max_version: Option<(u32, u32)>,
+ dl_device: Option<(i32, i32)>,
+ copy: Option,
+ ) -> PyResult> {
+ if stream.is_some() {
+ return Err(PyBufferError::new_err("CPU DLPack export requires stream=None"));
+ }
+ if max_version.is_some_and(|version| version.0 < 1) {
+ return Err(PyBufferError::new_err(
+ "consumer does not support the DLPack versioned ABI",
+ ));
+ }
+ if dl_device.is_some_and(|device| device != (1, 0)) {
+ return Err(PyBufferError::new_err(
+ "Tensor is CPU-resident; explicit device transfer is required",
+ ));
+ }
+ if copy == Some(true) {
+ return Err(PyBufferError::new_err(
+ "implicit DLPack copies are disabled; call Tensor.copy() explicitly",
+ ));
+ }
+ dlpack_capsule::export_tensor(py, &self.inner).map_err(to_py_err)
+ }
+
+ /// Makes an explicit host-to-host allocation copy.
+ fn copy(&self) -> Self {
+ Self { inner: self.inner.to_owned_copy() }
+ }
+
+ fn __repr__(&self) -> String {
+ format!("Tensor(shape={:?}, dtype='{}', device='cpu')", self.shape(), self.dtype())
+ }
+}
+
+/// Copies a NumPy uint8, uint16, or float32 array into packed CPU tensor storage.
+#[pyfunction]
+fn tensor_copy_from_numpy(array: &Bound<'_, PyAny>) -> PyResult {
+ if let Ok(array) = array.extract::>() {
+ let shape = array.shape().to_vec();
+ let bytes = array.as_array().iter().copied().collect::>();
+ let descriptor = TensorDescriptor::contiguous(TensorDataType::U8, shape, TensorDevice::CPU);
+ return Ok(PyTensor {
+ inner: TensorBuffer::try_new(bytes, descriptor).map_err(to_py_err)?,
+ });
+ }
+ if let Ok(array) = array.extract::>() {
+ let shape = array.shape().to_vec();
+ let values = array.as_array().iter().copied().collect::>();
+ let descriptor =
+ TensorDescriptor::contiguous(TensorDataType::U16, shape, TensorDevice::CPU);
+ return Ok(PyTensor {
+ inner: TensorBuffer::try_from_u16(values, descriptor).map_err(to_py_err)?,
+ });
+ }
+ if let Ok(array) = array.extract::>() {
+ let shape = array.shape().to_vec();
+ let values = array.as_array().iter().copied().collect::>();
+ let descriptor =
+ TensorDescriptor::contiguous(TensorDataType::F32, shape, TensorDevice::CPU);
+ return Ok(PyTensor {
+ inner: TensorBuffer::try_from_f32(values, descriptor).map_err(to_py_err)?,
+ });
+ }
+ Err(PyTypeError::new_err("expected a NumPy uint8, uint16, or float32 array"))
+}
+
+/// Takes ownership of a producer's CPU DLPack capsule without copying its allocation.
+#[pyfunction]
+fn tensor_view_from_dlpack(producer: &Bound<'_, PyAny>) -> PyResult {
+ let inner = dlpack_capsule::import_tensor(producer)?;
+ Ok(PyDlpackTensorView { inner })
+}
+
+fn parse_image_format(format: &str) -> PyResult {
+ match format.to_ascii_lowercase().as_str() {
+ "png" => Ok(ImageFileFormat::Png),
+ "jpg" | "jpeg" => Ok(ImageFileFormat::Jpeg),
+ "pnm" | "pbm" | "pgm" | "ppm" => Ok(ImageFileFormat::Pnm),
+ other => Err(PyValueError::new_err(format!(
+ "unsupported image format `{other}` (expected: png, jpeg, or pnm)"
+ ))),
+ }
+}
+
+/// Source metadata returned alongside a decoded NumPy image.
+#[pyclass(name = "ImageMetadata", frozen)]
+#[derive(Clone)]
+struct PyImageMetadata {
+ inner: DecodedMetadata,
+}
+
+#[pymethods]
+impl PyImageMetadata {
+ #[getter]
+ fn format(&self) -> String {
+ self.inner.format.to_string()
+ }
+
+ #[getter]
+ fn color_type(&self) -> String {
+ format!("{:?}", self.inner.source_color_type)
+ }
+
+ #[getter]
+ fn orientation(&self) -> u8 {
+ self.inner.orientation as u8
+ }
+
+ #[getter]
+ fn orientation_applied(&self) -> bool {
+ self.inner.orientation_applied
+ }
+
+ fn __repr__(&self) -> String {
+ format!(
+ "ImageMetadata(format='{}', color_type='{}', orientation={}, orientation_applied={})",
+ self.format(),
+ self.color_type(),
+ self.orientation(),
+ self.orientation_applied()
+ )
+ }
+}
+
+/// Decodes PNG, JPEG, or PNM into an owned NumPy array and source metadata.
+#[pyfunction]
+#[pyo3(signature = (path, apply_orientation=true))]
+fn read_image<'py>(
+ py: Python<'py>,
+ path: &str,
+ apply_orientation: bool,
+) -> PyResult<(Py, PyImageMetadata)> {
+ let decoded =
+ decode_image_path(path, DecodeOptions { apply_orientation, ..Default::default() })
+ .map_err(to_py_err)?;
+ let metadata = PyImageMetadata { inner: decoded.metadata() };
+ let (height, width) = (decoded.height(), decoded.width());
+ macro_rules! array2 {
+ ($image:expr) => {
+ Array2::from_shape_vec((height, width), $image.into_vec())
+ .map_err(to_py_err)?
+ .into_pyarray_bound(py)
+ .into_any()
+ .unbind()
+ };
+ }
+ macro_rules! array3 {
+ ($image:expr, $channels:expr) => {
+ Array3::from_shape_vec((height, width, $channels), $image.into_vec())
+ .map_err(to_py_err)?
+ .into_pyarray_bound(py)
+ .into_any()
+ .unbind()
+ };
+ }
+ let array = match decoded.into_pixels() {
+ DecodedPixels::Gray8(image) => array2!(image),
+ DecodedPixels::GrayAlpha8(image) => array3!(image, 2),
+ DecodedPixels::Rgb8(image) => array3!(image, 3),
+ DecodedPixels::Rgba8(image) => array3!(image, 4),
+ DecodedPixels::Gray16(image) => array2!(image),
+ DecodedPixels::GrayAlpha16(image) => array3!(image, 2),
+ DecodedPixels::Rgb16(image) => array3!(image, 3),
+ DecodedPixels::Rgba16(image) => array3!(image, 4),
+ DecodedPixels::Rgb32Float(image) => array3!(image, 3),
+ DecodedPixels::Rgba32Float(image) => array3!(image, 4),
+ };
+ Ok((array, metadata))
+}
+
+/// Encodes a uint8/uint16 NumPy image as PNG, JPEG, or PNM.
+#[pyfunction]
+#[pyo3(signature = (path, image, format, jpeg_quality=90))]
+fn write_image(
+ path: &str,
+ image: &Bound<'_, PyAny>,
+ format: &str,
+ jpeg_quality: u8,
+) -> PyResult<()> {
+ let format = parse_image_format(format)?;
+ let pixels = if let Ok(array) = image.extract::>() {
+ let view = array.as_array();
+ let shape = view.shape();
+ DecodedPixels::Gray8(
+ Image::try_new(shape[1], shape[0], view.iter().copied().collect())
+ .map_err(to_py_err)?,
+ )
+ } else if let Ok(array) = image.extract::>() {
+ let view = array.as_array();
+ let shape = view.shape();
+ DecodedPixels::Gray16(
+ Image::try_new(shape[1], shape[0], view.iter().copied().collect())
+ .map_err(to_py_err)?,
+ )
+ } else if let Ok(array) = image.extract::>() {
+ let view = array.as_array();
+ let shape = view.shape();
+ let packed = view.iter().copied().collect();
+ match shape[2] {
+ 2 => DecodedPixels::GrayAlpha8(
+ Image::try_new(shape[1], shape[0], packed).map_err(to_py_err)?,
+ ),
+ 3 => {
+ DecodedPixels::Rgb8(Image::try_new(shape[1], shape[0], packed).map_err(to_py_err)?)
+ }
+ 4 => {
+ DecodedPixels::Rgba8(Image::try_new(shape[1], shape[0], packed).map_err(to_py_err)?)
+ }
+ channels => {
+ return Err(PyValueError::new_err(format!(
+ "expected 2, 3, or 4 channels, found {channels}"
+ )))
+ }
+ }
+ } else if let Ok(array) = image.extract::>() {
+ let view = array.as_array();
+ let shape = view.shape();
+ let packed = view.iter().copied().collect();
+ match shape[2] {
+ 2 => DecodedPixels::GrayAlpha16(
+ Image::try_new(shape[1], shape[0], packed).map_err(to_py_err)?,
+ ),
+ 3 => {
+ DecodedPixels::Rgb16(Image::try_new(shape[1], shape[0], packed).map_err(to_py_err)?)
+ }
+ 4 => DecodedPixels::Rgba16(
+ Image::try_new(shape[1], shape[0], packed).map_err(to_py_err)?,
+ ),
+ channels => {
+ return Err(PyValueError::new_err(format!(
+ "expected 2, 3, or 4 channels, found {channels}"
+ )))
+ }
+ }
+ } else {
+ return Err(PyValueError::new_err(
+ "expected a uint8 or uint16 NumPy array shaped (H, W) or (H, W, C)",
+ ));
+ };
+ encode_image_path(path, &pixels, EncodeOptions { format, jpeg_quality }).map_err(to_py_err)
+}
+
fn parse_policy(policy: &str) -> PyResult {
match policy.to_lowercase().as_str() {
"auto" => Ok(ExecutionPolicy::Auto),
@@ -100,6 +647,17 @@ fn parse_interpolation(interpolation: &str) -> PyResult {
}
}
+fn parse_threshold_type(value: &str) -> PyResult {
+ match value.to_ascii_lowercase().as_str() {
+ "binary" => Ok(ThresholdType::Binary),
+ "binary_inv" | "binary-inv" => Ok(ThresholdType::BinaryInv),
+ "truncate" | "trunc" => Ok(ThresholdType::Truncate),
+ "to_zero" | "to-zero" => Ok(ThresholdType::ToZero),
+ "to_zero_inv" | "to-zero-inv" => Ok(ThresholdType::ToZeroInv),
+ other => Err(PyValueError::new_err(format!("unknown threshold type `{other}`"))),
+ }
+}
+
fn rgb_image_from_numpy(array: PyReadonlyArray3<'_, u8>) -> PyResult> {
let view = array.as_array();
let shape = view.shape();
@@ -1204,6 +1762,760 @@ fn rgbd_to_point_cloud(
Ok(PyPointCloud { inner })
}
+/// Correlates an RGB image with a 2D float64 kernel using Reflect101 borders.
+#[pyfunction]
+#[pyo3(signature = (image, kernel, delta=0.0))]
+fn filter2d_image<'py>(
+ py: Python<'py>,
+ image: PyReadonlyArray3<'_, u8>,
+ kernel: PyReadonlyArray2<'_, f64>,
+ delta: f64,
+) -> PyResult>> {
+ let image = rgb_image_from_numpy(image)?;
+ let kernel_view = kernel.as_array();
+ let shape = kernel_view.shape();
+ let kernel = Kernel2D::try_new(shape[1], shape[0], kernel_view.iter().copied().collect())
+ .map_err(to_py_err)?;
+ let output =
+ filter2d_op(image.view(), &kernel, delta, BorderMode::Reflect101).map_err(to_py_err)?;
+ let array = Array3::from_shape_vec((image.height(), image.width(), 3), output.into_vec())
+ .map_err(to_py_err)?;
+ Ok(array.into_pyarray_bound(py))
+}
+
+/// Applies a normalized Gaussian blur to an RGB image using Reflect101 borders.
+#[pyfunction]
+#[pyo3(signature = (image, kernel_width, kernel_height, sigma_x, sigma_y=None))]
+fn gaussian_blur_image<'py>(
+ py: Python<'py>,
+ image: PyReadonlyArray3<'_, u8>,
+ kernel_width: usize,
+ kernel_height: usize,
+ sigma_x: f64,
+ sigma_y: Option,
+) -> PyResult>> {
+ let image = rgb_image_from_numpy(image)?;
+ let output = gaussian_blur_op(
+ image.view(),
+ kernel_width,
+ kernel_height,
+ sigma_x,
+ sigma_y.unwrap_or(sigma_x),
+ BorderMode::Reflect101,
+ )
+ .map_err(to_py_err)?;
+ let array = Array3::from_shape_vec((image.height(), image.width(), 3), output.into_vec())
+ .map_err(to_py_err)?;
+ Ok(array.into_pyarray_bound(py))
+}
+
+/// Applies an odd-aperture median filter to an RGB image.
+#[pyfunction]
+fn median_blur_image<'py>(
+ py: Python<'py>,
+ image: PyReadonlyArray3<'_, u8>,
+ kernel_size: usize,
+) -> PyResult>> {
+ let image = rgb_image_from_numpy(image)?;
+ let output =
+ median_blur_op(image.view(), kernel_size, BorderMode::Replicate).map_err(to_py_err)?;
+ let array = Array3::from_shape_vec((image.height(), image.width(), 3), output.into_vec())
+ .map_err(to_py_err)?;
+ Ok(array.into_pyarray_bound(py))
+}
+
+/// Applies an RGB bilateral filter using Reflect101 borders.
+#[pyfunction]
+fn bilateral_filter_image<'py>(
+ py: Python<'py>,
+ image: PyReadonlyArray3<'_, u8>,
+ diameter: usize,
+ sigma_color: f64,
+ sigma_space: f64,
+) -> PyResult>> {
+ let image = rgb_image_from_numpy(image)?;
+ let output = bilateral_filter_op(
+ image.view(),
+ diameter,
+ sigma_color,
+ sigma_space,
+ BorderMode::Reflect101,
+ )
+ .map_err(to_py_err)?;
+ let array = Array3::from_shape_vec((image.height(), image.width(), 3), output.into_vec())
+ .map_err(to_py_err)?;
+ Ok(array.into_pyarray_bound(py))
+}
+
+/// Computes a signed float32 Sobel derivative from a grayscale uint8 image.
+#[pyfunction]
+#[pyo3(signature = (image, dx, dy, kernel_size=3, scale=1.0, delta=0.0))]
+fn sobel_image<'py>(
+ py: Python<'py>,
+ image: PyReadonlyArray2<'_, u8>,
+ dx: usize,
+ dy: usize,
+ kernel_size: usize,
+ scale: f64,
+ delta: f64,
+) -> PyResult>> {
+ let image = gray_u8_image_from_numpy(image)?;
+ let output = sobel_op(image.view(), dx, dy, kernel_size, scale, delta, BorderMode::Reflect101)
+ .map_err(to_py_err)?;
+ let array = Array2::from_shape_vec((image.height(), image.width()), output.into_vec())
+ .map_err(to_py_err)?;
+ Ok(array.into_pyarray_bound(py))
+}
+
+/// Computes a signed float32 Scharr derivative from a grayscale uint8 image.
+#[pyfunction]
+#[pyo3(signature = (image, dx, dy, scale=1.0, delta=0.0))]
+fn scharr_image<'py>(
+ py: Python<'py>,
+ image: PyReadonlyArray2<'_, u8>,
+ dx: usize,
+ dy: usize,
+ scale: f64,
+ delta: f64,
+) -> PyResult>> {
+ let image = gray_u8_image_from_numpy(image)?;
+ let output =
+ scharr_op(image.view(), dx, dy, scale, delta, BorderMode::Reflect101).map_err(to_py_err)?;
+ let array = Array2::from_shape_vec((image.height(), image.width()), output.into_vec())
+ .map_err(to_py_err)?;
+ Ok(array.into_pyarray_bound(py))
+}
+
+/// Computes a signed float32 Laplacian from a grayscale uint8 image.
+#[pyfunction]
+#[pyo3(signature = (image, kernel_size=1, scale=1.0, delta=0.0))]
+fn laplacian_image<'py>(
+ py: Python<'py>,
+ image: PyReadonlyArray2<'_, u8>,
+ kernel_size: usize,
+ scale: f64,
+ delta: f64,
+) -> PyResult>> {
+ let image = gray_u8_image_from_numpy(image)?;
+ let output = laplacian_op(image.view(), kernel_size, scale, delta, BorderMode::Reflect101)
+ .map_err(to_py_err)?;
+ let array = Array2::from_shape_vec((image.height(), image.width()), output.into_vec())
+ .map_err(to_py_err)?;
+ Ok(array.into_pyarray_bound(py))
+}
+
+/// Reduces an RGB image with the canonical Gaussian pyramid kernel.
+#[pyfunction]
+fn pyr_down_image<'py>(
+ py: Python<'py>,
+ image: PyReadonlyArray3<'_, u8>,
+) -> PyResult>> {
+ let image = rgb_image_from_numpy(image)?;
+ let output = pyr_down_op(image.view(), BorderMode::Reflect101).map_err(to_py_err)?;
+ let array = Array3::from_shape_vec((output.height(), output.width(), 3), output.into_vec())
+ .map_err(to_py_err)?;
+ Ok(array.into_pyarray_bound(py))
+}
+
+/// Doubles an RGB image with the canonical Gaussian pyramid kernel.
+#[pyfunction]
+fn pyr_up_image<'py>(
+ py: Python<'py>,
+ image: PyReadonlyArray3<'_, u8>,
+) -> PyResult>> {
+ let image = rgb_image_from_numpy(image)?;
+ let output = pyr_up_op(image.view(), BorderMode::Reflect101).map_err(to_py_err)?;
+ let array = Array3::from_shape_vec((output.height(), output.width(), 3), output.into_vec())
+ .map_err(to_py_err)?;
+ Ok(array.into_pyarray_bound(py))
+}
+
+/// Applies grayscale erosion, dilation, or a composite morphology operation.
+#[pyfunction]
+#[pyo3(signature = (image, operation, kernel_width, kernel_height, shape="rect", iterations=1))]
+fn morphology_image<'py>(
+ py: Python<'py>,
+ image: PyReadonlyArray2<'_, u8>,
+ operation: &str,
+ kernel_width: usize,
+ kernel_height: usize,
+ shape: &str,
+ iterations: usize,
+) -> PyResult>> {
+ let image = gray_u8_image_from_numpy(image)?;
+ let shape = match shape.to_ascii_lowercase().as_str() {
+ "rect" | "rectangle" => MorphologyShape::Rect,
+ "cross" => MorphologyShape::Cross,
+ "ellipse" | "elliptical" => MorphologyShape::Ellipse,
+ "diamond" => MorphologyShape::Diamond,
+ other => {
+ return Err(PyValueError::new_err(format!(
+ "unknown morphology shape `{other}` (expected: rect, cross, ellipse, diamond)"
+ )))
+ }
+ };
+ let element =
+ StructuringElement::try_new(shape, kernel_width, kernel_height).map_err(to_py_err)?;
+ let operation = operation.to_ascii_lowercase();
+ let output = match operation.as_str() {
+ "erode" => {
+ spatialrust::vision::erode(image.view(), &element, iterations, BorderMode::Replicate)
+ }
+ "dilate" => {
+ spatialrust::vision::dilate(image.view(), &element, iterations, BorderMode::Replicate)
+ }
+ "open" => morphology_ex_op(
+ image.view(),
+ MorphologyOperation::Open,
+ &element,
+ iterations,
+ BorderMode::Replicate,
+ ),
+ "close" => morphology_ex_op(
+ image.view(),
+ MorphologyOperation::Close,
+ &element,
+ iterations,
+ BorderMode::Replicate,
+ ),
+ "gradient" => morphology_ex_op(
+ image.view(),
+ MorphologyOperation::Gradient,
+ &element,
+ iterations,
+ BorderMode::Replicate,
+ ),
+ "tophat" | "top-hat" => morphology_ex_op(
+ image.view(),
+ MorphologyOperation::TopHat,
+ &element,
+ iterations,
+ BorderMode::Replicate,
+ ),
+ "blackhat" | "black-hat" => morphology_ex_op(
+ image.view(),
+ MorphologyOperation::BlackHat,
+ &element,
+ iterations,
+ BorderMode::Replicate,
+ ),
+ other => {
+ return Err(PyValueError::new_err(format!("unknown morphology operation `{other}`")))
+ }
+ }
+ .map_err(to_py_err)?;
+ let array = Array2::from_shape_vec((image.height(), image.width()), output.into_vec())
+ .map_err(to_py_err)?;
+ Ok(array.into_pyarray_bound(py))
+}
+
+/// Applies a fixed threshold to a grayscale uint8 image.
+#[pyfunction]
+#[pyo3(signature = (image, threshold, max_value=255, threshold_type="binary"))]
+fn threshold_image<'py>(
+ py: Python<'py>,
+ image: PyReadonlyArray2<'_, u8>,
+ threshold: f64,
+ max_value: u8,
+ threshold_type: &str,
+) -> PyResult>> {
+ let image = gray_u8_image_from_numpy(image)?;
+ let output = threshold_op(
+ image.view(),
+ threshold,
+ f64::from(max_value),
+ parse_threshold_type(threshold_type)?,
+ )
+ .map_err(to_py_err)?;
+ let array = Array2::from_shape_vec((image.height(), image.width()), output.into_vec())
+ .map_err(to_py_err)?;
+ Ok(array.into_pyarray_bound(py))
+}
+
+/// Selects and applies an Otsu threshold, returning `(threshold, image)`.
+#[pyfunction]
+#[pyo3(signature = (image, max_value=255, threshold_type="binary"))]
+fn otsu_threshold_image<'py>(
+ py: Python<'py>,
+ image: PyReadonlyArray2<'_, u8>,
+ max_value: u8,
+ threshold_type: &str,
+) -> PyResult<(u8, Bound<'py, PyArray2>)> {
+ let image = gray_u8_image_from_numpy(image)?;
+ let (selected, output) =
+ otsu_threshold_u8_op(image.view(), max_value, parse_threshold_type(threshold_type)?)
+ .map_err(to_py_err)?;
+ let array = Array2::from_shape_vec((image.height(), image.width()), output.into_vec())
+ .map_err(to_py_err)?;
+ Ok((selected, array.into_pyarray_bound(py)))
+}
+
+/// Applies mean or Gaussian adaptive thresholding.
+#[pyfunction]
+#[pyo3(signature = (image, block_size, c, method="mean", max_value=255, threshold_type="binary"))]
+fn adaptive_threshold_image<'py>(
+ py: Python<'py>,
+ image: PyReadonlyArray2<'_, u8>,
+ block_size: usize,
+ c: f64,
+ method: &str,
+ max_value: u8,
+ threshold_type: &str,
+) -> PyResult>> {
+ let image = gray_u8_image_from_numpy(image)?;
+ let method = match method.to_ascii_lowercase().as_str() {
+ "mean" => AdaptiveThresholdMethod::Mean,
+ "gaussian" => AdaptiveThresholdMethod::Gaussian,
+ other => return Err(PyValueError::new_err(format!("unknown adaptive method `{other}`"))),
+ };
+ let output = adaptive_threshold_op(
+ image.view(),
+ max_value,
+ method,
+ parse_threshold_type(threshold_type)?,
+ block_size,
+ c,
+ BorderMode::Replicate,
+ )
+ .map_err(to_py_err)?;
+ let array = Array2::from_shape_vec((image.height(), image.width()), output.into_vec())
+ .map_err(to_py_err)?;
+ Ok(array.into_pyarray_bound(py))
+}
+
+/// Returns the exact 256-bin grayscale histogram.
+#[pyfunction]
+fn histogram_image<'py>(
+ py: Python<'py>,
+ image: PyReadonlyArray2<'_, u8>,
+) -> PyResult>> {
+ let image = gray_u8_image_from_numpy(image)?;
+ Ok(histogram_u8_op(image.view()).into_pyarray_bound(py))
+}
+
+/// Equalizes a grayscale uint8 histogram.
+#[pyfunction]
+fn equalize_histogram_image<'py>(
+ py: Python<'py>,
+ image: PyReadonlyArray2<'_, u8>,
+) -> PyResult>> {
+ let image = gray_u8_image_from_numpy(image)?;
+ let output = equalize_histogram_op(image.view()).map_err(to_py_err)?;
+ let array = Array2::from_shape_vec((image.height(), image.width()), output.into_vec())
+ .map_err(to_py_err)?;
+ Ok(array.into_pyarray_bound(py))
+}
+
+/// Applies contrast-limited adaptive histogram equalization.
+#[pyfunction]
+#[pyo3(signature = (image, clip_limit=2.0, tiles_x=8, tiles_y=8))]
+fn clahe_image<'py>(
+ py: Python<'py>,
+ image: PyReadonlyArray2<'_, u8>,
+ clip_limit: f64,
+ tiles_x: usize,
+ tiles_y: usize,
+) -> PyResult>> {
+ let image = gray_u8_image_from_numpy(image)?;
+ let output = clahe_op(image.view(), clip_limit, tiles_x, tiles_y).map_err(to_py_err)?;
+ let array = Array2::from_shape_vec((image.height(), image.width()), output.into_vec())
+ .map_err(to_py_err)?;
+ Ok(array.into_pyarray_bound(py))
+}
+
+/// Computes the `(H + 1, W + 1)` float64 summed-area table.
+#[pyfunction]
+fn integral_image_u8<'py>(
+ py: Python<'py>,
+ image: PyReadonlyArray2<'_, u8>,
+) -> PyResult>> {
+ let image = gray_u8_image_from_numpy(image)?;
+ let integral = integral_image_op(image.view(), 0).map_err(to_py_err)?;
+ let array =
+ Array2::from_shape_vec((integral.height(), integral.width()), integral.as_slice().to_vec())
+ .map_err(to_py_err)?;
+ Ok(array.into_pyarray_bound(py))
+}
+
+/// Detects edges in a grayscale uint8 image with Canny hysteresis.
+#[pyfunction]
+#[pyo3(signature = (image, low_threshold, high_threshold, aperture_size=3, l2_gradient=false))]
+fn canny_image<'py>(
+ py: Python<'py>,
+ image: PyReadonlyArray2<'_, u8>,
+ low_threshold: f64,
+ high_threshold: f64,
+ aperture_size: usize,
+ l2_gradient: bool,
+) -> PyResult>> {
+ let image = gray_u8_image_from_numpy(image)?;
+ let output = canny_op(
+ image.view(),
+ CannyOptions { low_threshold, high_threshold, aperture_size, l2_gradient },
+ )
+ .map_err(to_py_err)?;
+ let array = Array2::from_shape_vec((image.height(), image.width()), output.into_vec())
+ .map_err(to_py_err)?;
+ Ok(array.into_pyarray_bound(py))
+}
+
+fn corner_selection_options(
+ max_corners: usize,
+ quality_level: f32,
+ min_distance: f32,
+ block_size: usize,
+ gradient_size: usize,
+) -> CornerSelectionOptions {
+ CornerSelectionOptions {
+ max_corners,
+ quality_level,
+ min_distance,
+ block_size,
+ gradient_size,
+ border: BorderMode::Reflect101,
+ }
+}
+
+/// Detects strongest-first Harris keypoints in a grayscale uint8 image.
+#[pyfunction]
+#[pyo3(signature = (image, max_corners=0, quality_level=0.01, min_distance=1.0, block_size=3, gradient_size=3, k=0.04))]
+#[allow(clippy::too_many_arguments)]
+fn harris_keypoints(
+ image: PyReadonlyArray2<'_, u8>,
+ max_corners: usize,
+ quality_level: f32,
+ min_distance: f32,
+ block_size: usize,
+ gradient_size: usize,
+ k: f32,
+) -> PyResult> {
+ let image = gray_u8_image_from_numpy(image)?;
+ let options = HarrisOptions {
+ selection: corner_selection_options(
+ max_corners,
+ quality_level,
+ min_distance,
+ block_size,
+ gradient_size,
+ ),
+ k,
+ };
+ Ok(detect_harris_op(image.view(), options)
+ .map_err(to_py_err)?
+ .into_iter()
+ .map(|inner| PyKeypoint2 { inner })
+ .collect())
+}
+
+/// Detects strongest-first Shi–Tomasi keypoints in a grayscale uint8 image.
+#[pyfunction]
+#[pyo3(signature = (image, max_corners=0, quality_level=0.01, min_distance=1.0, block_size=3, gradient_size=3))]
+fn shi_tomasi_keypoints(
+ image: PyReadonlyArray2<'_, u8>,
+ max_corners: usize,
+ quality_level: f32,
+ min_distance: f32,
+ block_size: usize,
+ gradient_size: usize,
+) -> PyResult> {
+ let image = gray_u8_image_from_numpy(image)?;
+ let options = ShiTomasiOptions {
+ selection: corner_selection_options(
+ max_corners,
+ quality_level,
+ min_distance,
+ block_size,
+ gradient_size,
+ ),
+ };
+ Ok(detect_shi_tomasi_op(image.view(), options)
+ .map_err(to_py_err)?
+ .into_iter()
+ .map(|inner| PyKeypoint2 { inner })
+ .collect())
+}
+
+/// Detects scan-ordered FAST-9/16 keypoints in a grayscale uint8 image.
+#[pyfunction]
+#[pyo3(signature = (image, threshold=10, nonmax_suppression=true))]
+fn fast_keypoints(
+ image: PyReadonlyArray2<'_, u8>,
+ threshold: u8,
+ nonmax_suppression: bool,
+) -> PyResult> {
+ let image = gray_u8_image_from_numpy(image)?;
+ Ok(detect_fast_op(image.view(), FastOptions { threshold, nonmax_suppression })
+ .map_err(to_py_err)?
+ .into_iter()
+ .map(|inner| PyKeypoint2 { inner })
+ .collect())
+}
+
+/// Detects ORB keypoints and returns `(keypoints, uint8[N, 32] descriptors)`.
+#[pyfunction]
+#[pyo3(signature = (image, max_features=500, scale_factor=1.2, levels=8, edge_threshold=31, fast_threshold=20, patch_size=31, score_type="harris"))]
+#[allow(clippy::too_many_arguments)]
+fn orb_features<'py>(
+ py: Python<'py>,
+ image: PyReadonlyArray2<'_, u8>,
+ max_features: usize,
+ scale_factor: f32,
+ levels: usize,
+ edge_threshold: usize,
+ fast_threshold: u8,
+ patch_size: usize,
+ score_type: &str,
+) -> PyResult<(Vec, Bound<'py, PyArray2>)> {
+ let image = gray_u8_image_from_numpy(image)?;
+ let score_type = match score_type {
+ "harris" => OrbScoreType::Harris,
+ "fast" => OrbScoreType::Fast,
+ _ => return Err(PyValueError::new_err("score_type must be 'harris' or 'fast'")),
+ };
+ let features = detect_and_describe_orb_op(
+ image.view(),
+ OrbOptions {
+ max_features,
+ scale_factor,
+ levels,
+ edge_threshold,
+ fast_threshold,
+ patch_size,
+ score_type,
+ },
+ )
+ .map_err(to_py_err)?;
+ let keypoints = features
+ .keypoints()
+ .iter()
+ .copied()
+ .map(|inner| PyKeypoint2 { inner })
+ .collect();
+ let descriptors = Array2::from_shape_vec(
+ (features.descriptors().len(), features.descriptors().width()),
+ features.descriptors().binary_data().expect("ORB descriptors are binary").to_vec(),
+ )
+ .map_err(to_py_err)?
+ .into_pyarray_bound(py);
+ Ok((keypoints, descriptors))
+}
+
+fn correspondences_from_numpy(
+ source: PyReadonlyArray2<'_, f64>,
+ target: PyReadonlyArray2<'_, f64>,
+) -> PyResult> {
+ let source = source.as_array();
+ let target = target.as_array();
+ if source.shape() != target.shape() || source.ndim() != 2 || source.shape()[1] != 2 {
+ return Err(PyValueError::new_err("source/target must be Nx2 float64 arrays"));
+ }
+ source
+ .outer_iter()
+ .zip(target.outer_iter())
+ .map(|(src, dst)| {
+ PointCorrespondence2::try_new(
+ Vec2 { x: src[0], y: src[1] },
+ Vec2 { x: dst[0], y: dst[1] },
+ )
+ .map_err(to_py_err)
+ })
+ .collect()
+}
+
+fn mat3_to_numpy<'py>(py: Python<'py>, matrix: Mat3) -> Bound<'py, PyArray2> {
+ let mut values = Vec::with_capacity(9);
+ for row in &matrix.m {
+ values.extend_from_slice(row);
+ }
+ Array2::from_shape_vec((3, 3), values)
+ .expect("3x3")
+ .into_pyarray_bound(py)
+}
+
+/// Estimates a homography with deterministic RANSAC.
+///
+/// Returns `(matrix[3,3], inliers[N], residuals[N])`.
+#[pyfunction]
+#[pyo3(signature = (source, target, threshold=1.0, confidence=0.99, max_iterations=2000, seed=0))]
+fn estimate_homography_ransac<'py>(
+ py: Python<'py>,
+ source: PyReadonlyArray2<'_, f64>,
+ target: PyReadonlyArray2<'_, f64>,
+ threshold: f64,
+ confidence: f64,
+ max_iterations: usize,
+ seed: u64,
+) -> PyResult<(Bound<'py, PyArray2>, Bound<'py, PyArray1>, Bound<'py, PyArray1>)> {
+ let pairs = correspondences_from_numpy(source, target)?;
+ let estimate = estimate_homography_ransac_op(
+ &pairs,
+ RobustEstimationOptions { threshold, confidence, max_iterations, seed },
+ )
+ .map_err(to_py_err)?;
+ Ok((
+ mat3_to_numpy(py, estimate.model().matrix()),
+ numpy::PyArray1::from_vec_bound(py, estimate.inliers().to_vec()),
+ numpy::PyArray1::from_vec_bound(py, estimate.residuals().to_vec()),
+ ))
+}
+
+/// Solves PnP for object points `Nx3` and image points `Nx2`.
+///
+/// Camera intrinsics are `fx, fy, cx, cy`. Returns `(rotation[3,3], translation[3])`.
+#[pyfunction]
+#[pyo3(signature = (object_points, image_points, fx, fy, cx, cy, width=640, height=480))]
+#[allow(clippy::too_many_arguments)]
+fn solve_pnp<'py>(
+ py: Python<'py>,
+ object_points: PyReadonlyArray2<'_, f64>,
+ image_points: PyReadonlyArray2<'_, f64>,
+ fx: f64,
+ fy: f64,
+ cx: f64,
+ cy: f64,
+ width: usize,
+ height: usize,
+) -> PyResult<(Bound<'py, PyArray2>, Bound<'py, PyArray1>)> {
+ let objects = object_points.as_array();
+ let images = image_points.as_array();
+ if objects.ndim() != 2
+ || images.ndim() != 2
+ || objects.shape()[1] != 3
+ || images.shape()[1] != 2
+ || objects.shape()[0] != images.shape()[0]
+ {
+ return Err(PyValueError::new_err(
+ "object_points must be Nx3 and image_points Nx2 with matching N",
+ ));
+ }
+ let pairs = objects
+ .outer_iter()
+ .zip(images.outer_iter())
+ .map(|(object, image)| {
+ ObjectImageCorrespondence::try_new(
+ Vec3::new(object[0], object[1], object[2]),
+ Vec2 { x: image[0], y: image[1] },
+ )
+ .map_err(to_py_err)
+ })
+ .collect::>>()?;
+ let camera = CameraMatrix3::from_intrinsics(
+ CameraIntrinsics::try_new(fx, fy, cx, cy, width, height).map_err(to_py_err)?,
+ );
+ let pose: AbsolutePose = solve_pnp_op(&pairs, camera).map_err(to_py_err)?;
+ let translation = numpy::PyArray1::from_vec_bound(
+ py,
+ vec![pose.translation().x, pose.translation().y, pose.translation().z],
+ );
+ Ok((mat3_to_numpy(py, pose.rotation()), translation))
+}
+
+/// Dense SAD stereo block matching on rectified grayscale images.
+#[pyfunction]
+#[pyo3(signature = (left, right, window_size=15, min_disparity=0, num_disparities=64, uniqueness_ratio=15.0))]
+fn stereo_block_match<'py>(
+ py: Python<'py>,
+ left: PyReadonlyArray2<'_, u8>,
+ right: PyReadonlyArray2<'_, u8>,
+ window_size: usize,
+ min_disparity: i32,
+ num_disparities: i32,
+ uniqueness_ratio: f32,
+) -> PyResult>> {
+ let left = gray_u8_image_from_numpy(left)?;
+ let right = gray_u8_image_from_numpy(right)?;
+ let disparity = stereo_block_match_op(
+ left.view(),
+ right.view(),
+ StereoBmOptions { window_size, min_disparity, num_disparities, uniqueness_ratio },
+ )
+ .map_err(to_py_err)?;
+ let array = Array2::from_shape_vec(
+ (disparity.height(), disparity.width()),
+ disparity.as_slice().to_vec(),
+ )
+ .map_err(to_py_err)?;
+ Ok(array.into_pyarray_bound(py))
+}
+
+fn descriptor_match_tuples(
+ query: DescriptorBuffer,
+ train: DescriptorBuffer,
+ cross_check: bool,
+ ratio: Option,
+ max_distance: Option,
+) -> PyResult> {
+ Ok(match_descriptors_op(
+ &query,
+ &train,
+ MatchOptions { cross_check, ratio, max_distance },
+ )
+ .map_err(to_py_err)?
+ .into_iter()
+ .map(|feature_match| {
+ (
+ feature_match.query_index(),
+ feature_match.train_index(),
+ feature_match.distance(),
+ )
+ })
+ .collect())
+}
+
+/// Brute-force Hamming matching for two `uint8[N, D]` descriptor matrices.
+#[pyfunction]
+#[pyo3(signature = (query, train, cross_check=false, ratio=None, max_distance=None))]
+fn match_binary_descriptors(
+ query: PyReadonlyArray2<'_, u8>,
+ train: PyReadonlyArray2<'_, u8>,
+ cross_check: bool,
+ ratio: Option,
+ max_distance: Option,
+) -> PyResult> {
+ let query_shape = query.shape();
+ let train_shape = train.shape();
+ let query = DescriptorBuffer::try_binary(
+ query_shape[0],
+ query_shape[1],
+ query.as_array().iter().copied().collect(),
+ )
+ .map_err(to_py_err)?;
+ let train = DescriptorBuffer::try_binary(
+ train_shape[0],
+ train_shape[1],
+ train.as_array().iter().copied().collect(),
+ )
+ .map_err(to_py_err)?;
+ descriptor_match_tuples(query, train, cross_check, ratio, max_distance)
+}
+
+/// Brute-force Euclidean matching for two `float32[N, D]` descriptor matrices.
+#[pyfunction]
+#[pyo3(signature = (query, train, cross_check=false, ratio=None, max_distance=None))]
+fn match_float_descriptors(
+ query: PyReadonlyArray2<'_, f32>,
+ train: PyReadonlyArray2<'_, f32>,
+ cross_check: bool,
+ ratio: Option,
+ max_distance: Option,
+) -> PyResult> {
+ let query_shape = query.shape();
+ let train_shape = train.shape();
+ let query = DescriptorBuffer::try_float32(
+ query_shape[0],
+ query_shape[1],
+ query.as_array().iter().copied().collect(),
+ )
+ .map_err(to_py_err)?;
+ let train = DescriptorBuffer::try_float32(
+ train_shape[0],
+ train_shape[1],
+ train.as_array().iter().copied().collect(),
+ )
+ .map_err(to_py_err)?;
+ descriptor_match_tuples(query, train, cross_check, ratio, max_distance)
+}
+
/// Resizes an `(H, W, 3)` uint8 RGB image.
#[pyfunction]
#[pyo3(signature = (image, width, height, interpolation="bilinear"))]
@@ -1533,6 +2845,11 @@ fn point_map_to_point_cloud(
#[pyo3(name = "spatialrust")]
fn spatialrust_module(m: &Bound<'_, PyModule>) -> PyResult<()> {
m.add("__version__", env!("CARGO_PKG_VERSION"))?;
+ m.add_class::()?;
+ m.add_class::()?;
+ m.add_class::()?;
+ m.add_class::()?;
+ m.add_class::()?;
m.add_class::()?;
m.add_class::()?;
m.add_class::()?;
@@ -1542,6 +2859,19 @@ fn spatialrust_module(m: &Bound<'_, PyModule>) -> PyResult<()> {
m.add_class::()?;
m.add_class::()?;
m.add_class::()?;
+ m.add_function(wrap_pyfunction!(read_image, m)?)?;
+ m.add_function(wrap_pyfunction!(tensor_copy_from_numpy, m)?)?;
+ m.add_function(wrap_pyfunction!(tensor_view_from_dlpack, m)?)?;
+ m.add_function(wrap_pyfunction!(harris_keypoints, m)?)?;
+ m.add_function(wrap_pyfunction!(shi_tomasi_keypoints, m)?)?;
+ m.add_function(wrap_pyfunction!(fast_keypoints, m)?)?;
+ m.add_function(wrap_pyfunction!(orb_features, m)?)?;
+ m.add_function(wrap_pyfunction!(estimate_homography_ransac, m)?)?;
+ m.add_function(wrap_pyfunction!(solve_pnp, m)?)?;
+ m.add_function(wrap_pyfunction!(stereo_block_match, m)?)?;
+ m.add_function(wrap_pyfunction!(match_binary_descriptors, m)?)?;
+ m.add_function(wrap_pyfunction!(match_float_descriptors, m)?)?;
+ m.add_function(wrap_pyfunction!(write_image, m)?)?;
m.add_function(wrap_pyfunction!(read, m)?)?;
m.add_function(wrap_pyfunction!(write, m)?)?;
m.add_function(wrap_pyfunction!(voxel_downsample, m)?)?;
@@ -1574,6 +2904,24 @@ fn spatialrust_module(m: &Bound<'_, PyModule>) -> PyResult<()> {
m.add_function(wrap_pyfunction!(voxelize, m)?)?;
m.add_function(wrap_pyfunction!(range_image, m)?)?;
m.add_function(wrap_pyfunction!(rgbd_to_point_cloud, m)?)?;
+ m.add_function(wrap_pyfunction!(filter2d_image, m)?)?;
+ m.add_function(wrap_pyfunction!(gaussian_blur_image, m)?)?;
+ m.add_function(wrap_pyfunction!(median_blur_image, m)?)?;
+ m.add_function(wrap_pyfunction!(bilateral_filter_image, m)?)?;
+ m.add_function(wrap_pyfunction!(sobel_image, m)?)?;
+ m.add_function(wrap_pyfunction!(scharr_image, m)?)?;
+ m.add_function(wrap_pyfunction!(laplacian_image, m)?)?;
+ m.add_function(wrap_pyfunction!(pyr_down_image, m)?)?;
+ m.add_function(wrap_pyfunction!(pyr_up_image, m)?)?;
+ m.add_function(wrap_pyfunction!(morphology_image, m)?)?;
+ m.add_function(wrap_pyfunction!(threshold_image, m)?)?;
+ m.add_function(wrap_pyfunction!(otsu_threshold_image, m)?)?;
+ m.add_function(wrap_pyfunction!(adaptive_threshold_image, m)?)?;
+ m.add_function(wrap_pyfunction!(histogram_image, m)?)?;
+ m.add_function(wrap_pyfunction!(equalize_histogram_image, m)?)?;
+ m.add_function(wrap_pyfunction!(clahe_image, m)?)?;
+ m.add_function(wrap_pyfunction!(integral_image_u8, m)?)?;
+ m.add_function(wrap_pyfunction!(canny_image, m)?)?;
m.add_function(wrap_pyfunction!(resize_image, m)?)?;
m.add_function(wrap_pyfunction!(letterbox_image, m)?)?;
m.add_function(wrap_pyfunction!(normalize_image_chw, m)?)?;
diff --git a/crates/spatialrust-py/tests/test_bindings.py b/crates/spatialrust-py/tests/test_bindings.py
index 540a66f..1440ec3 100644
--- a/crates/spatialrust-py/tests/test_bindings.py
+++ b/crates/spatialrust-py/tests/test_bindings.py
@@ -469,3 +469,246 @@ def test_fpfh_ransac_returns_transform():
res = sr.register_fpfh_ransac(src, tgt, feature_radius=0.5,
max_correspondence_distance=0.2, ransac_iterations=500)
assert res.transform().shape == (4, 4)
+def test_png_image_io_roundtrip(tmp_path):
+ image = np.arange(5 * 7 * 3, dtype=np.uint8).reshape(5, 7, 3)
+ path = tmp_path / "roundtrip.png"
+ sr.write_image(str(path), image[:, ::-1], "png")
+ decoded, metadata = sr.read_image(str(path))
+ np.testing.assert_array_equal(decoded, image[:, ::-1])
+ assert metadata.format == "png"
+ assert metadata.color_type == "Rgb8"
+ assert metadata.orientation in (0, 1)
+ assert "ImageMetadata(" in repr(metadata)
+
+
+def test_filter2d_and_gaussian_preserve_rgb_shape():
+ image = np.arange(7 * 9 * 3, dtype=np.uint8).reshape(7, 9, 3)
+ identity = sr.filter2d_image(image[:, ::-1], np.array([[1.0]], dtype=np.float64))
+ np.testing.assert_array_equal(identity, image[:, ::-1])
+ blurred = sr.gaussian_blur_image(image, 5, 3, 1.2, 0.8)
+ assert blurred.shape == image.shape
+ assert blurred.dtype == np.uint8
+
+
+def test_advanced_filters_and_pyramid_shapes():
+ image = np.arange(9 * 11 * 3, dtype=np.uint8).reshape(9, 11, 3)
+ assert sr.median_blur_image(image[:, ::-1], 3).shape == image.shape
+ assert sr.bilateral_filter_image(image, 3, 20.0, 2.0).shape == image.shape
+ gray = image[..., 0]
+ for derivative in (
+ sr.sobel_image(gray, 1, 0),
+ sr.scharr_image(gray, 0, 1),
+ sr.laplacian_image(gray),
+ ):
+ assert derivative.shape == gray.shape
+ assert derivative.dtype == np.float32
+ down = sr.pyr_down_image(image)
+ assert down.shape == (5, 6, 3)
+ assert sr.pyr_up_image(down).shape == (10, 12, 3)
+
+
+def test_morphology_operations_and_noncontiguous_input():
+ mask = np.zeros((9, 11), dtype=np.uint8)
+ mask[2:7, 3:8] = 255
+ for operation in ("erode", "dilate", "open", "close", "gradient", "tophat", "blackhat"):
+ output = sr.morphology_image(mask[:, ::-1], operation, 3, 3, "ellipse", 2)
+ assert output.shape == mask.shape
+ assert output.dtype == np.uint8
+
+
+def test_threshold_histogram_clahe_and_integral_contracts():
+ image = np.arange(9 * 11, dtype=np.uint8).reshape(9, 11)[:, ::-1]
+ assert sr.threshold_image(image, 40).shape == image.shape
+ selected, otsu = sr.otsu_threshold_image(image)
+ assert 0 <= selected <= 255 and otsu.shape == image.shape
+ for method in ("mean", "gaussian"):
+ assert sr.adaptive_threshold_image(image, 5, 2.0, method).shape == image.shape
+ histogram = sr.histogram_image(image)
+ assert histogram.shape == (256,) and int(histogram.sum()) == image.size
+ assert sr.equalize_histogram_image(image).shape == image.shape
+ assert sr.clahe_image(image, 2.0, 3, 2).shape == image.shape
+ integral = sr.integral_image_u8(image)
+ assert integral.shape == (image.shape[0] + 1, image.shape[1] + 1)
+ assert integral[-1, -1] == pytest.approx(float(image.sum()))
+
+
+def test_canny_image_binary_output_and_noncontiguous_input():
+ image = np.zeros((17, 19), dtype=np.uint8)
+ image[4:13, 6:14] = 255
+ edges = sr.canny_image(image[:, ::-1], 50.0, 100.0, 3, True)
+ assert edges.shape == image.shape
+ assert edges.dtype == np.uint8
+ assert set(np.unique(edges)).issubset({0, 255})
+ assert np.count_nonzero(edges) > 0
+
+
+def test_feature2d_corner_detectors_and_keypoint_metadata():
+ image = np.zeros((25, 29), dtype=np.uint8)
+ image[5:19, 7:22] = 255
+ harris = sr.harris_keypoints(image[:, ::-1], 20, 0.01, 1.0, 3, 3, 0.04)
+ shi = sr.shi_tomasi_keypoints(image, 20, 0.01, 1.0, 3, 3)
+ assert len(harris) >= 4 and len(shi) >= 4
+ assert all(point.size == 3.0 and point.angle_degrees is None for point in harris)
+ impulse = np.zeros((9, 9), dtype=np.uint8)
+ impulse[4, 4] = 255
+ fast = sr.fast_keypoints(impulse, 20, True)
+ assert len(fast) == 1
+ assert (fast[0].x, fast[0].y, fast[0].size) == (4.0, 4.0, 7.0)
+ assert "Keypoint2(" in repr(fast[0])
+
+
+def test_orb_features_and_descriptor_matchers():
+ yy, xx = np.indices((96, 112), dtype=np.int32)
+ image = ((xx * 37 + yy * 19) ^ (xx * yy * 3) ^ ((xx // 8 + yy // 8) * 127)).astype(np.uint8)
+ keypoints, descriptors = sr.orb_features(image[:, ::-1], max_features=60, edge_threshold=16)
+ assert 0 < len(keypoints) <= 60
+ assert descriptors.shape == (len(keypoints), 32)
+ assert descriptors.dtype == np.uint8
+ assert all(point.angle_degrees is not None for point in keypoints)
+
+ matches = sr.match_binary_descriptors(descriptors, descriptors, cross_check=True)
+ assert len(matches) == len(keypoints)
+ assert all(query == train and distance == 0.0 for query, train, distance in matches)
+
+ query = np.array([[0.0, 0.0], [10.0, 10.0]], dtype=np.float32)
+ train = np.array([[1.0, 0.0], [3.0, 0.0], [10.0, 9.0]], dtype=np.float32)
+ assert sr.match_float_descriptors(query, train, cross_check=True, ratio=0.8) == [
+ (0, 0, 1.0),
+ (1, 2, 1.0),
+ ]
+
+
+def test_geometry_homography_pnp_and_stereo():
+ source = np.array(
+ [[0.0, 0.0], [40.0, 0.0], [0.0, 30.0], [40.0, 30.0], [20.0, 15.0], [10.0, 5.0]],
+ dtype=np.float64,
+ )
+ h = np.array([[1.05, 0.01, 2.0], [-0.02, 0.98, -1.5], [0.0001, 0.0, 1.0]], dtype=np.float64)
+ target = []
+ for point in source:
+ projected = h @ np.array([point[0], point[1], 1.0])
+ target.append([projected[0] / projected[2], projected[1] / projected[2]])
+ target = np.asarray(target, dtype=np.float64)
+ matrix, inliers, residuals = sr.estimate_homography_ransac(source, target, threshold=1.0)
+ assert matrix.shape == (3, 3)
+ assert inliers.dtype == np.bool_
+ assert residuals.shape == (source.shape[0],)
+ assert int(inliers.sum()) >= 4
+
+ objects = np.array(
+ [
+ [0.0, 0.0, 0.0],
+ [0.2, 0.0, 0.0],
+ [0.0, 0.15, 0.0],
+ [0.0, 0.0, 0.1],
+ [0.1, 0.05, 0.05],
+ [0.05, -0.05, 0.02],
+ ],
+ dtype=np.float64,
+ )
+ fx = fy = 500.0
+ cx, cy = 320.0, 240.0
+ rotation = np.eye(3, dtype=np.float64)
+ translation = np.array([0.1, -0.05, 2.5], dtype=np.float64)
+ images = []
+ for point in objects:
+ camera = rotation @ point + translation
+ images.append([fx * camera[0] / camera[2] + cx, fy * camera[1] / camera[2] + cy])
+ images = np.asarray(images, dtype=np.float64)
+ recovered_r, recovered_t = sr.solve_pnp(objects, images, fx, fy, cx, cy)
+ assert recovered_r.shape == (3, 3)
+ assert recovered_t.shape == (3,)
+ assert abs(recovered_t[2] - translation[2]) < 0.05
+
+ width, height = 96, 64
+ disparity = 12
+ yy, xx = np.indices((height, width), dtype=np.int32)
+ left = ((xx * 17 + yy * 29) % 200 + 20).astype(np.uint8)
+ right = np.zeros_like(left)
+ right[:, : width - disparity] = left[:, disparity:]
+ disparity_map = sr.stereo_block_match(
+ left, right, window_size=11, min_disparity=1, num_disparities=32, uniqueness_ratio=5.0
+ )
+ assert disparity_map.shape == (height, width)
+ assert abs(float(disparity_map[height // 2, width // 2]) - float(disparity)) <= 1.0
+
+
+@pytest.mark.parametrize("dtype", [np.uint8, np.uint16, np.float32])
+def test_tensor_dlpack_zero_copy_export(dtype):
+ source = np.arange(3 * 5, dtype=dtype).reshape(3, 5)[:, ::-1]
+ tensor = sr.tensor_copy_from_numpy(source)
+ first = np.from_dlpack(tensor)
+ second = np.from_dlpack(tensor)
+ np.testing.assert_array_equal(first, source)
+ assert first.dtype == source.dtype
+ assert first.shape == source.shape
+ assert first.__array_interface__["data"][0] == second.__array_interface__["data"][0]
+ assert not first.flags.writeable
+ assert tensor.__dlpack_device__() == (1, 0)
+ assert "Tensor(shape=" in repr(tensor)
+
+
+def test_tensor_dlpack_copy_and_device_requests_are_explicit():
+ tensor = sr.tensor_copy_from_numpy(np.arange(8, dtype=np.uint8))
+ copied = tensor.copy()
+ assert np.from_dlpack(copied).__array_interface__["data"][0] != np.from_dlpack(
+ tensor
+ ).__array_interface__["data"][0]
+ with pytest.raises(BufferError):
+ tensor.__dlpack__(copy=True)
+ with pytest.raises(BufferError):
+ tensor.__dlpack__(dl_device=(2, 0))
+
+
+def test_onnxruntime_dynamic_named_binding_matches_reference(tmp_path):
+ model = bytes(
+ [
+ 8, 8, 18, 16, 115, 112, 97, 116, 105, 97, 108, 114, 117, 115, 116, 45, 116,
+ 101, 115, 116, 58, 106, 10, 27, 10, 5, 105, 110, 112, 117, 116, 10, 5, 105,
+ 110, 112, 117, 116, 18, 6, 111, 117, 116, 112, 117, 116, 34, 3, 65, 100, 100,
+ 18, 14, 100, 111, 117, 98, 108, 101, 95, 100, 121, 110, 97, 109, 105, 99, 90,
+ 28, 10, 5, 105, 110, 112, 117, 116, 18, 19, 10, 17, 8, 1, 18, 13, 10, 7,
+ 18, 5, 98, 97, 116, 99, 104, 10, 2, 8, 3, 98, 29, 10, 6, 111, 117, 116,
+ 112, 117, 116, 18, 19, 10, 17, 8, 1, 18, 13, 10, 7, 18, 5, 98, 97, 116,
+ 99, 104, 10, 2, 8, 3, 66, 4, 10, 0, 16, 13,
+ ]
+ )
+ path = tmp_path / "double_dynamic.onnx"
+ path.write_bytes(model)
+ try:
+ session = sr.OnnxRuntimeSession(str(path), deterministic=True)
+ except RuntimeError as error:
+ if "without the `onnxruntime` feature" in str(error):
+ pytest.skip("extension was intentionally built without ONNX Runtime")
+ raise
+
+ source = np.arange(12, dtype=np.float32).reshape(4, 3)
+ inputs = {"input": sr.tensor_copy_from_numpy(source)}
+ assert session.inputs == [("input", "float32", ["batch", "3"])]
+ bound = np.from_dlpack(session.run(inputs)["output"])
+ copied = np.from_dlpack(session.run(inputs, copy=True)["output"])
+ expected = source * 2.0
+ np.testing.assert_array_equal(bound, expected)
+ np.testing.assert_array_equal(copied, expected)
+
+ try:
+ import onnxruntime as reference_runtime
+ except ImportError:
+ return
+ reference = reference_runtime.InferenceSession(
+ str(path), providers=["CPUExecutionProvider"]
+ ).run(None, {"input": source})[0]
+ np.testing.assert_array_equal(bound, reference)
+
+
+@pytest.mark.parametrize("dtype", [np.uint8, np.uint16, np.float32])
+def test_tensor_zero_copy_dlpack_import_retains_producer(dtype):
+ source = np.arange(12, dtype=dtype).reshape(3, 4)
+ imported = sr.tensor_view_from_dlpack(source)
+ assert imported.shape == [3, 4]
+ assert imported.dtype == np.dtype(dtype).name
+ assert imported.version[0] == 1
+ del source
+ copied = np.from_dlpack(imported.copy())
+ np.testing.assert_array_equal(copied, np.arange(12, dtype=dtype).reshape(3, 4))
+ assert "DLPackTensorView(" in repr(imported)
diff --git a/crates/spatialrust-tensor/Cargo.toml b/crates/spatialrust-tensor/Cargo.toml
new file mode 100644
index 0000000..667ca8d
--- /dev/null
+++ b/crates/spatialrust-tensor/Cargo.toml
@@ -0,0 +1,30 @@
+[package]
+name = "spatialrust-tensor"
+version.workspace = true
+edition.workspace = true
+license.workspace = true
+authors.workspace = true
+repository.workspace = true
+rust-version.workspace = true
+description = "Small tensor metadata and explicit CPU ownership contracts for SpatialRust"
+
+[features]
+default = []
+dlpack = []
+image = ["dep:spatialrust-image"]
+spatial = ["dep:spatialrust-core"]
+
+[dependencies]
+spatialrust-image = { workspace = true, optional = true }
+spatialrust-core = { workspace = true, optional = true }
+bytemuck.workspace = true
+thiserror.workspace = true
+
+[dev-dependencies]
+proptest.workspace = true
+criterion.workspace = true
+
+[[bench]]
+name = "image_bridge"
+harness = false
+required-features = ["image"]
diff --git a/crates/spatialrust-tensor/benches/image_bridge.rs b/crates/spatialrust-tensor/benches/image_bridge.rs
new file mode 100644
index 0000000..5b5488c
--- /dev/null
+++ b/crates/spatialrust-tensor/benches/image_bridge.rs
@@ -0,0 +1,30 @@
+use criterion::{black_box, criterion_group, criterion_main, BenchmarkId, Criterion, Throughput};
+use spatialrust_image::{Image, ImageView};
+use spatialrust_tensor::{interleaved_image_view, pack_interleaved_image};
+
+fn benchmark_image_bridge(c: &mut Criterion) {
+ for &(name, width, height) in &[("640p", 640, 480), ("1080p", 1920, 1080), ("4k", 3840, 2160)] {
+ let packed = Image::::try_new(width, height, vec![127; width * height * 3]).unwrap();
+ let row_stride = width * 3 + 16;
+ let padded = vec![127_u8; row_stride * height];
+ let strided = ImageView::::new(width, height, row_stride, &padded).unwrap();
+
+ let mut zero_copy = c.benchmark_group("tensor_image_zero_copy");
+ zero_copy.throughput(Throughput::Elements((width * height) as u64));
+ zero_copy.bench_function(BenchmarkId::from_parameter(name), |b| {
+ b.iter(|| black_box(interleaved_image_view(black_box(&packed)).unwrap()));
+ });
+ zero_copy.finish();
+
+ let mut packing = c.benchmark_group("tensor_image_pack_strided");
+ packing.sample_size(10);
+ packing.throughput(Throughput::Bytes((width * height * 3) as u64));
+ packing.bench_function(BenchmarkId::from_parameter(name), |b| {
+ b.iter(|| black_box(pack_interleaved_image(black_box(strided)).unwrap()));
+ });
+ packing.finish();
+ }
+}
+
+criterion_group!(benches, benchmark_image_bridge);
+criterion_main!(benches);
diff --git a/crates/spatialrust-tensor/src/dlpack.rs b/crates/spatialrust-tensor/src/dlpack.rs
new file mode 100644
index 0000000..4bd2d71
--- /dev/null
+++ b/crates/spatialrust-tensor/src/dlpack.rs
@@ -0,0 +1,606 @@
+//! Audited DLPack major-version 1 CPU ownership boundary.
+
+use std::{
+ ffi::c_void,
+ mem::ManuallyDrop,
+ ptr::{self, NonNull},
+ slice,
+};
+
+use crate::{
+ DataType, DataTypeCode, Device, DeviceKind, TensorBuffer, TensorDescriptor, TensorError,
+ TensorStorage, TensorView,
+};
+
+/// DLPack ABI major version supported by this boundary.
+pub const DLPACK_MAJOR: u32 = 1;
+/// Baseline DLPack minor version emitted by this boundary.
+pub const DLPACK_MINOR: u32 = 0;
+
+const FLAG_READ_ONLY: u64 = 1;
+const MAX_RANK: usize = 64;
+
+/// Errors raised while validating a DLPack managed tensor.
+#[derive(Debug, thiserror::Error)]
+pub enum DlpackError {
+ /// A null managed-tensor pointer was supplied.
+ #[error("DLPack managed tensor pointer is null")]
+ NullManagedTensor,
+ /// The ABI major version is incompatible.
+ #[error("unsupported DLPack ABI version {major}.{minor}; expected major {DLPACK_MAJOR}")]
+ IncompatibleVersion {
+ /// Producer major version.
+ major: u32,
+ /// Producer minor version.
+ minor: u32,
+ },
+ /// Rank is negative or exceeds the defensive limit.
+ #[error("invalid DLPack tensor rank {0}")]
+ InvalidRank(i32),
+ /// A ranked tensor has no shape pointer.
+ #[error("DLPack tensor shape pointer is null for non-zero rank")]
+ NullShape,
+ /// A dimension is negative or cannot be represented by `usize`.
+ #[error("invalid DLPack shape dimension {0}")]
+ InvalidDimension(i64),
+ /// The data type code is unsupported or not byte-addressable.
+ #[error("unsupported DLPack dtype code={code}, bits={bits}, lanes={lanes}")]
+ UnsupportedDataType {
+ /// Raw DLPack code.
+ code: u8,
+ /// Bits per lane.
+ bits: u8,
+ /// Vector lanes.
+ lanes: u16,
+ },
+ /// The device code is unknown to this DLPack minor implementation.
+ #[error("unsupported DLPack device type {0}")]
+ UnsupportedDevice(i32),
+ /// A non-empty host tensor has no data pointer.
+ #[error("non-empty DLPack host tensor has a null data pointer")]
+ NullData,
+ /// A DLPack integer does not fit the host representation.
+ #[error("DLPack metadata does not fit the host integer representation")]
+ IntegerConversion,
+ /// The decoded tensor layout is invalid.
+ #[error(transparent)]
+ Tensor(#[from] TensorError),
+}
+
+#[repr(C)]
+#[derive(Clone, Copy)]
+struct RawVersion {
+ major: u32,
+ minor: u32,
+}
+
+#[repr(C)]
+#[derive(Clone, Copy)]
+struct RawDevice {
+ device_type: i32,
+ device_id: i32,
+}
+
+#[repr(C)]
+#[derive(Clone, Copy)]
+struct RawDataType {
+ code: u8,
+ bits: u8,
+ lanes: u16,
+}
+
+#[repr(C)]
+struct RawTensor {
+ data: *mut c_void,
+ device: RawDevice,
+ ndim: i32,
+ dtype: RawDataType,
+ shape: *mut i64,
+ strides: *mut i64,
+ byte_offset: u64,
+}
+
+#[repr(C)]
+struct RawManagedTensorVersioned {
+ version: RawVersion,
+ manager_ctx: *mut c_void,
+ deleter: Option,
+ flags: u64,
+ dl_tensor: RawTensor,
+}
+
+struct ExportContext {
+ _allocation: TensorStorage,
+ _shape: Box<[i64]>,
+ _strides: Box<[i64]>,
+}
+
+/// Owner for a versioned DLPack managed tensor exported without copying CPU data.
+///
+/// Dropping this value calls the DLPack deleter. [`Self::into_raw`] transfers
+/// that responsibility to a capsule or another consumer.
+pub struct DlpackExport {
+ raw: NonNull,
+}
+
+/// Calls the producer deleter for a raw pointer previously returned by
+/// [`DlpackExport::into_raw`].
+///
+/// # Safety
+///
+/// `raw` must still carry exclusive deleter responsibility for a live
+/// `DLManagedTensorVersioned*`. It must not be used after this call.
+pub unsafe fn release_dlpack_raw(raw: *mut c_void) {
+ // SAFETY: forwarded from the public ownership contract above.
+ unsafe { call_deleter(raw.cast()) };
+}
+
+impl DlpackExport {
+ /// Shares an owned CPU tensor allocation with a DLPack consumer without copying.
+ pub fn from_tensor(tensor: &TensorBuffer) -> Result {
+ let descriptor = tensor.descriptor();
+ if !descriptor.device().is_host_accessible() {
+ return Err(TensorError::DeviceNotHostAccessible(descriptor.device()).into());
+ }
+ let shape = descriptor
+ .shape()
+ .iter()
+ .map(|&dimension| i64::try_from(dimension).map_err(|_| DlpackError::IntegerConversion))
+ .collect::, _>>()?
+ .into_boxed_slice();
+ let strides = match descriptor.strides() {
+ Some(values) => values
+ .iter()
+ .map(|&stride| i64::try_from(stride).map_err(|_| DlpackError::IntegerConversion))
+ .collect::, _>>()?,
+ None => compact_strides(descriptor.shape())?,
+ }
+ .into_boxed_slice();
+ let ndim =
+ i32::try_from(descriptor.shape().len()).map_err(|_| DlpackError::IntegerConversion)?;
+ let byte_offset =
+ u64::try_from(descriptor.byte_offset()).map_err(|_| DlpackError::IntegerConversion)?;
+ let allocation = tensor.shared_allocation();
+ let data = if allocation.is_empty() {
+ ptr::null_mut()
+ } else {
+ allocation.as_ptr().cast_mut().cast()
+ };
+ let shape_ptr = if shape.is_empty() { ptr::null_mut() } else { shape.as_ptr().cast_mut() };
+ let strides_ptr =
+ if strides.is_empty() { ptr::null_mut() } else { strides.as_ptr().cast_mut() };
+ let context =
+ Box::new(ExportContext { _allocation: allocation, _shape: shape, _strides: strides });
+ let manager_ctx = Box::into_raw(context).cast();
+ let raw = Box::new(RawManagedTensorVersioned {
+ version: RawVersion { major: DLPACK_MAJOR, minor: DLPACK_MINOR },
+ manager_ctx,
+ deleter: Some(delete_export),
+ flags: FLAG_READ_ONLY,
+ dl_tensor: RawTensor {
+ data,
+ device: encode_device(descriptor.device()),
+ ndim,
+ dtype: encode_dtype(descriptor.dtype()),
+ shape: shape_ptr,
+ strides: strides_ptr,
+ byte_offset,
+ },
+ });
+ Ok(Self { raw: NonNull::from(Box::leak(raw)) })
+ }
+
+ /// Returns the opaque managed-tensor pointer without transferring ownership.
+ pub fn as_raw(&self) -> *mut c_void {
+ self.raw.as_ptr().cast()
+ }
+
+ /// Transfers deleter responsibility to an external DLPack consumer.
+ pub fn into_raw(self) -> *mut c_void {
+ let this = ManuallyDrop::new(self);
+ this.raw.as_ptr().cast()
+ }
+}
+
+impl Drop for DlpackExport {
+ fn drop(&mut self) {
+ // SAFETY: `DlpackExport` uniquely owns deleter responsibility until `into_raw`.
+ unsafe { call_deleter(self.raw.as_ptr()) };
+ }
+}
+
+/// Validated owner of a DLPack producer's managed tensor.
+#[derive(Debug)]
+pub struct DlpackImport {
+ raw: NonNull,
+ descriptor: TensorDescriptor,
+ allocation_len: usize,
+ data: *const u8,
+ version: (u32, u32),
+ flags: u64,
+}
+
+impl DlpackImport {
+ /// Takes deleter ownership of a DLPack managed tensor and validates its host view.
+ ///
+ /// # Safety
+ ///
+ /// `raw` must be a live, exclusively transferred `DLManagedTensorVersioned*`
+ /// produced according to DLPack. The caller must not use or delete it after
+ /// this call, whether validation succeeds or fails.
+ pub unsafe fn from_raw(raw: *mut c_void) -> Result {
+ let raw = NonNull::new(raw.cast::())
+ .ok_or(DlpackError::NullManagedTensor)?;
+ let guard = IncomingGuard { raw: Some(raw) };
+
+ // SAFETY: the caller promises a live versioned header. DLPack guarantees
+ // version and deleter positions remain accessible across major mismatch.
+ let managed = unsafe { raw.as_ref() };
+ let version = (managed.version.major, managed.version.minor);
+ if version.0 != DLPACK_MAJOR {
+ return Err(DlpackError::IncompatibleVersion { major: version.0, minor: version.1 });
+ }
+ let tensor = &managed.dl_tensor;
+ if tensor.ndim < 0 || tensor.ndim as usize > MAX_RANK {
+ return Err(DlpackError::InvalidRank(tensor.ndim));
+ }
+ let rank = tensor.ndim as usize;
+ if rank != 0 && tensor.shape.is_null() {
+ return Err(DlpackError::NullShape);
+ }
+ let shape_values = if rank == 0 {
+ &[][..]
+ } else {
+ // SAFETY: producer contract supplies `ndim` readable shape entries.
+ unsafe { slice::from_raw_parts(tensor.shape, rank) }
+ };
+ let shape = shape_values
+ .iter()
+ .map(|&dimension| {
+ usize::try_from(dimension).map_err(|_| DlpackError::InvalidDimension(dimension))
+ })
+ .collect::, _>>()?;
+ let strides = if tensor.strides.is_null() {
+ None
+ } else {
+ // SAFETY: producer contract supplies `ndim` readable stride entries.
+ let values = unsafe { slice::from_raw_parts(tensor.strides, rank) };
+ Some(
+ values
+ .iter()
+ .map(|&stride| {
+ isize::try_from(stride).map_err(|_| DlpackError::IntegerConversion)
+ })
+ .collect::, _>>()?,
+ )
+ };
+ let dtype = decode_dtype(tensor.dtype)?;
+ let device = decode_device(tensor.device)?;
+ let byte_offset =
+ usize::try_from(tensor.byte_offset).map_err(|_| DlpackError::IntegerConversion)?;
+ let descriptor = match strides {
+ Some(strides) => {
+ TensorDescriptor::try_strided(dtype, shape, strides, byte_offset, device)?
+ }
+ None => {
+ let mut descriptor = TensorDescriptor::contiguous(dtype, shape, device);
+ if byte_offset != 0 {
+ descriptor = TensorDescriptor::try_strided(
+ dtype,
+ descriptor.shape().to_vec(),
+ compact_strides_isize(descriptor.shape())?,
+ byte_offset,
+ device,
+ )?;
+ }
+ descriptor
+ }
+ };
+ let range = descriptor.required_byte_range()?;
+ if range.end != 0 && tensor.data.is_null() {
+ return Err(DlpackError::NullData);
+ }
+ let raw = guard.disarm();
+ Ok(Self {
+ raw,
+ descriptor,
+ allocation_len: range.end,
+ data: tensor.data.cast(),
+ version,
+ flags: managed.flags,
+ })
+ }
+
+ /// Returns producer ABI major and minor versions.
+ pub const fn version(&self) -> (u32, u32) {
+ self.version
+ }
+
+ /// Returns raw DLPack flags.
+ pub const fn flags(&self) -> u64 {
+ self.flags
+ }
+
+ /// Returns validated tensor metadata.
+ pub const fn descriptor(&self) -> &TensorDescriptor {
+ &self.descriptor
+ }
+
+ /// Borrows the imported host allocation without copying it.
+ pub fn view(&self) -> Result, DlpackError> {
+ let bytes = if self.allocation_len == 0 {
+ &[][..]
+ } else {
+ // SAFETY: `from_raw` validated the DLPack producer contract and the
+ // managed tensor remains alive until this owner is dropped.
+ unsafe { slice::from_raw_parts(self.data, self.allocation_len) }
+ };
+ Ok(TensorView::try_new(bytes, self.descriptor.clone())?)
+ }
+}
+
+impl Drop for DlpackImport {
+ fn drop(&mut self) {
+ // SAFETY: this owner received exclusive deleter responsibility in `from_raw`.
+ unsafe { call_deleter(self.raw.as_ptr()) };
+ }
+}
+
+struct IncomingGuard {
+ raw: Option>,
+}
+
+impl IncomingGuard {
+ fn disarm(mut self) -> NonNull {
+ self.raw.take().expect("incoming pointer is present")
+ }
+}
+
+impl Drop for IncomingGuard {
+ fn drop(&mut self) {
+ if let Some(raw) = self.raw {
+ // SAFETY: the guard owns the transferred pointer on all error paths.
+ unsafe { call_deleter(raw.as_ptr()) };
+ }
+ }
+}
+
+unsafe extern "C" fn delete_export(raw: *mut RawManagedTensorVersioned) {
+ if raw.is_null() {
+ return;
+ }
+ // SAFETY: this function is installed only for allocations built by
+ // `DlpackExport::from_tensor` and is called exactly once by ownership contract.
+ let managed = unsafe { Box::from_raw(raw) };
+ if !managed.manager_ctx.is_null() {
+ // SAFETY: manager_ctx was created with Box::into_raw for ExportContext.
+ drop(unsafe { Box::from_raw(managed.manager_ctx.cast::()) });
+ }
+}
+
+unsafe fn call_deleter(raw: *mut RawManagedTensorVersioned) {
+ if raw.is_null() {
+ return;
+ }
+ // SAFETY: caller owns a live managed-tensor pointer.
+ if let Some(deleter) = unsafe { (*raw).deleter } {
+ // SAFETY: deleter belongs to this exact managed tensor.
+ unsafe { deleter(raw) };
+ }
+}
+
+fn compact_strides(shape: &[usize]) -> Result, DlpackError> {
+ let mut output = vec![0; shape.len()];
+ let mut stride = 1_i64;
+ for (index, &dimension) in shape.iter().enumerate().rev() {
+ output[index] = stride;
+ stride = stride
+ .checked_mul(
+ i64::try_from(dimension.max(1)).map_err(|_| DlpackError::IntegerConversion)?,
+ )
+ .ok_or(DlpackError::IntegerConversion)?;
+ }
+ Ok(output)
+}
+
+fn compact_strides_isize(shape: &[usize]) -> Result, DlpackError> {
+ compact_strides(shape)?
+ .into_iter()
+ .map(|stride| isize::try_from(stride).map_err(|_| DlpackError::IntegerConversion))
+ .collect()
+}
+
+fn encode_dtype(dtype: DataType) -> RawDataType {
+ RawDataType { code: dtype.code() as u8, bits: dtype.bits(), lanes: dtype.lanes() }
+}
+
+fn decode_dtype(raw: RawDataType) -> Result {
+ let code = match raw.code {
+ 0 => DataTypeCode::Int,
+ 1 => DataTypeCode::UInt,
+ 2 => DataTypeCode::Float,
+ 4 => DataTypeCode::BFloat,
+ 5 => DataTypeCode::Complex,
+ 6 => DataTypeCode::Bool,
+ _ => {
+ return Err(DlpackError::UnsupportedDataType {
+ code: raw.code,
+ bits: raw.bits,
+ lanes: raw.lanes,
+ })
+ }
+ };
+ DataType::try_new(code, raw.bits, raw.lanes).map_err(|_| DlpackError::UnsupportedDataType {
+ code: raw.code,
+ bits: raw.bits,
+ lanes: raw.lanes,
+ })
+}
+
+fn encode_device(device: Device) -> RawDevice {
+ let device_type = match device.kind {
+ DeviceKind::Cpu => 1,
+ DeviceKind::Cuda => 2,
+ DeviceKind::CudaHost => 3,
+ DeviceKind::OpenCl => 4,
+ DeviceKind::Vulkan => 7,
+ DeviceKind::Metal => 8,
+ DeviceKind::Vpi => 9,
+ DeviceKind::Rocm => 10,
+ DeviceKind::RocmHost => 11,
+ DeviceKind::External => 12,
+ DeviceKind::CudaManaged => 13,
+ DeviceKind::OneApi => 14,
+ DeviceKind::WebGpu => 15,
+ DeviceKind::Hexagon => 16,
+ DeviceKind::Maia => 17,
+ DeviceKind::Trainium => 18,
+ DeviceKind::Tpu => 19,
+ DeviceKind::TpuHost => 20,
+ };
+ RawDevice { device_type, device_id: device.id }
+}
+
+fn decode_device(raw: RawDevice) -> Result {
+ let kind = match raw.device_type {
+ 1 => DeviceKind::Cpu,
+ 2 => DeviceKind::Cuda,
+ 3 => DeviceKind::CudaHost,
+ 4 => DeviceKind::OpenCl,
+ 7 => DeviceKind::Vulkan,
+ 8 => DeviceKind::Metal,
+ 9 => DeviceKind::Vpi,
+ 10 => DeviceKind::Rocm,
+ 11 => DeviceKind::RocmHost,
+ 12 => DeviceKind::External,
+ 13 => DeviceKind::CudaManaged,
+ 14 => DeviceKind::OneApi,
+ 15 => DeviceKind::WebGpu,
+ 16 => DeviceKind::Hexagon,
+ 17 => DeviceKind::Maia,
+ 18 => DeviceKind::Trainium,
+ 19 => DeviceKind::Tpu,
+ 20 => DeviceKind::TpuHost,
+ other => return Err(DlpackError::UnsupportedDevice(other)),
+ };
+ Ok(Device { kind, id: raw.device_id })
+}
+
+#[cfg(test)]
+mod tests {
+ use super::{DlpackError, DlpackExport, DlpackImport, DLPACK_MAJOR, DLPACK_MINOR};
+ use crate::{DataType, Device, TensorBuffer, TensorDescriptor};
+
+ #[test]
+ fn owned_cpu_roundtrip_is_zero_copy_and_versioned() {
+ let tensor = TensorBuffer::try_new(
+ (0_u8..24).collect(),
+ TensorDescriptor::contiguous(DataType::F32, vec![2, 3], Device::CPU),
+ )
+ .unwrap();
+ let original = tensor.allocation_bytes().as_ptr();
+ let allocation = tensor.shared_allocation();
+ let export = DlpackExport::from_tensor(&tensor).unwrap();
+ assert_eq!(allocation.strong_count(), 3);
+ // SAFETY: into_raw transfers the live export exactly once.
+ let imported = unsafe { DlpackImport::from_raw(export.into_raw()) }.unwrap();
+ assert_eq!(imported.version(), (DLPACK_MAJOR, DLPACK_MINOR));
+ let view = imported.view().unwrap();
+ assert_eq!(view.descriptor().shape(), &[2, 3]);
+ assert_eq!(view.allocation_bytes().as_ptr(), original);
+ drop(imported);
+ assert_eq!(allocation.strong_count(), 2);
+ }
+
+ #[test]
+ fn negative_stride_and_byte_offset_roundtrip() {
+ let descriptor =
+ TensorDescriptor::try_strided(DataType::U8, vec![4], vec![-1], 3, Device::CPU).unwrap();
+ let tensor = TensorBuffer::try_new(vec![10, 20, 30, 40], descriptor).unwrap();
+ let export = DlpackExport::from_tensor(&tensor).unwrap();
+ // SAFETY: the export pointer is transferred exactly once.
+ let imported = unsafe { DlpackImport::from_raw(export.into_raw()) }.unwrap();
+ let view = imported.view().unwrap();
+ assert_eq!(view.descriptor().strides(), Some(&[-1][..]));
+ assert_eq!(view.descriptor().byte_offset(), 3);
+ assert_eq!(view.allocation_bytes(), &[10, 20, 30, 40]);
+ }
+
+ #[test]
+ fn major_mismatch_is_rejected_and_deleted() {
+ let tensor = TensorBuffer::try_new(
+ vec![1],
+ TensorDescriptor::contiguous(DataType::U8, vec![1], Device::CPU),
+ )
+ .unwrap();
+ let allocation = tensor.shared_allocation();
+ let export = DlpackExport::from_tensor(&tensor).unwrap();
+ let raw = export.into_raw().cast::();
+ // SAFETY: test exclusively owns this live export and only changes its version header.
+ unsafe { (*raw).version.major = 99 };
+ // SAFETY: the pointer is still live and transferred exactly once.
+ let error = unsafe { DlpackImport::from_raw(raw.cast()) }.unwrap_err();
+ assert!(matches!(error, DlpackError::IncompatibleVersion { major: 99, .. }));
+ assert_eq!(allocation.strong_count(), 2);
+ }
+
+ #[test]
+ fn malformed_dtype_is_rejected_and_deleted() {
+ let tensor = TensorBuffer::try_new(
+ vec![1],
+ TensorDescriptor::contiguous(DataType::U8, vec![1], Device::CPU),
+ )
+ .unwrap();
+ let allocation = tensor.shared_allocation();
+ let raw = DlpackExport::from_tensor(&tensor)
+ .unwrap()
+ .into_raw()
+ .cast::();
+ // SAFETY: test owns the live export and corrupts only metadata under test.
+ unsafe { (*raw).dl_tensor.dtype.code = 255 };
+ // SAFETY: the corrupted but live pointer is transferred exactly once.
+ let error = unsafe { DlpackImport::from_raw(raw.cast()) }.unwrap_err();
+ assert!(matches!(error, DlpackError::UnsupportedDataType { code: 255, .. }));
+ assert_eq!(allocation.strong_count(), 2);
+ }
+
+ #[test]
+ fn null_nonempty_data_is_rejected_and_deleted() {
+ let tensor = TensorBuffer::try_new(
+ vec![1, 2],
+ TensorDescriptor::contiguous(DataType::U8, vec![2], Device::CPU),
+ )
+ .unwrap();
+ let allocation = tensor.shared_allocation();
+ let raw = DlpackExport::from_tensor(&tensor)
+ .unwrap()
+ .into_raw()
+ .cast::();
+ // SAFETY: test owns the live export and corrupts only metadata under test.
+ unsafe { (*raw).dl_tensor.data = std::ptr::null_mut() };
+ // SAFETY: the corrupted but live pointer is transferred exactly once.
+ let error = unsafe { DlpackImport::from_raw(raw.cast()) }.unwrap_err();
+ assert!(matches!(error, DlpackError::NullData));
+ assert_eq!(allocation.strong_count(), 2);
+ }
+
+ #[test]
+ fn negative_shape_is_rejected_and_deleted() {
+ let tensor = TensorBuffer::try_new(
+ vec![1],
+ TensorDescriptor::contiguous(DataType::U8, vec![1], Device::CPU),
+ )
+ .unwrap();
+ let allocation = tensor.shared_allocation();
+ let raw = DlpackExport::from_tensor(&tensor)
+ .unwrap()
+ .into_raw()
+ .cast::();
+ // SAFETY: exported rank is one and shape points to one writable context entry.
+ unsafe { *(*raw).dl_tensor.shape = -1 };
+ // SAFETY: the corrupted but live pointer is transferred exactly once.
+ let error = unsafe { DlpackImport::from_raw(raw.cast()) }.unwrap_err();
+ assert!(matches!(error, DlpackError::InvalidDimension(-1)));
+ assert_eq!(allocation.strong_count(), 2);
+ }
+}
diff --git a/crates/spatialrust-tensor/src/image.rs b/crates/spatialrust-tensor/src/image.rs
new file mode 100644
index 0000000..4ad614d
--- /dev/null
+++ b/crates/spatialrust-tensor/src/image.rs
@@ -0,0 +1,153 @@
+//! Explicit bridges between typed CPU images and generic tensors.
+
+use bytemuck::{cast_slice, Pod};
+use spatialrust_image::{Image, ImageView, PlanarImage, PlanarImageView};
+
+use crate::{DataType, Device, TensorBuffer, TensorDescriptor, TensorError, TensorView};
+
+/// A native-endian scalar that has a stable tensor dtype and no invalid bit patterns.
+pub trait TensorElement: Pod {
+ /// Tensor dtype corresponding to this Rust scalar.
+ const DTYPE: DataType;
+}
+
+macro_rules! tensor_elements {
+ ($($type:ty => $dtype:expr),+ $(,)?) => {
+ $(impl TensorElement for $type {
+ const DTYPE: DataType = $dtype;
+ })+
+ };
+}
+
+tensor_elements! {
+ u8 => DataType::U8,
+ u16 => DataType::U16,
+ i8 => DataType::I8,
+ i16 => DataType::I16,
+ i32 => DataType::I32,
+ i64 => DataType::I64,
+ f32 => DataType::F32,
+ f64 => DataType::F64,
+}
+
+/// Borrows a packed interleaved image as a zero-copy `[height, width, channels]` tensor.
+pub fn interleaved_image_view(
+ image: &Image,
+) -> Result, TensorError> {
+ let descriptor = TensorDescriptor::contiguous(
+ T::DTYPE,
+ vec![image.height(), image.width(), CHANNELS],
+ Device::CPU,
+ );
+ TensorView::try_new(cast_slice(image.as_slice()), descriptor)
+}
+
+/// Borrows a packed planar image as a zero-copy `[channels, height, width]` tensor.
+pub fn planar_image_view(
+ image: &PlanarImage,
+) -> Result, TensorError> {
+ let descriptor = TensorDescriptor::contiguous(
+ T::DTYPE,
+ vec![CHANNELS, image.height(), image.width()],
+ Device::CPU,
+ );
+ TensorView::try_new(cast_slice(image.as_slice()), descriptor)
+}
+
+/// Explicitly packs a possibly strided interleaved image view into an owned HWC tensor.
+pub fn pack_interleaved_image(
+ image: ImageView<'_, T, CHANNELS>,
+) -> Result {
+ let scalar_count = image
+ .width()
+ .checked_mul(image.height())
+ .and_then(|value| value.checked_mul(CHANNELS))
+ .ok_or(TensorError::LayoutOverflow)?;
+ let byte_count =
+ scalar_count.checked_mul(T::DTYPE.element_size()).ok_or(TensorError::LayoutOverflow)?;
+ let mut bytes = Vec::with_capacity(byte_count);
+ for y in 0..image.height() {
+ let row = image.row(y).expect("row index is within the image height");
+ bytes.extend_from_slice(cast_slice(row));
+ }
+ let descriptor = TensorDescriptor::contiguous(
+ T::DTYPE,
+ vec![image.height(), image.width(), CHANNELS],
+ Device::CPU,
+ );
+ TensorBuffer::try_new(bytes, descriptor)
+}
+
+/// Explicitly packs a possibly strided planar image view into an owned CHW tensor.
+pub fn pack_planar_image(
+ image: PlanarImageView<'_, T, CHANNELS>,
+) -> Result {
+ let scalar_count = image
+ .width()
+ .checked_mul(image.height())
+ .and_then(|value| value.checked_mul(CHANNELS))
+ .ok_or(TensorError::LayoutOverflow)?;
+ let byte_count =
+ scalar_count.checked_mul(T::DTYPE.element_size()).ok_or(TensorError::LayoutOverflow)?;
+ let mut bytes = Vec::with_capacity(byte_count);
+ for channel in 0..CHANNELS {
+ for y in 0..image.height() {
+ for x in 0..image.width() {
+ let value = image
+ .get(channel, x, y)
+ .expect("channel and pixel coordinates are within the image");
+ bytes.extend_from_slice(bytemuck::bytes_of(value));
+ }
+ }
+ }
+ let descriptor = TensorDescriptor::contiguous(
+ T::DTYPE,
+ vec![CHANNELS, image.height(), image.width()],
+ Device::CPU,
+ );
+ TensorBuffer::try_new(bytes, descriptor)
+}
+
+#[cfg(test)]
+mod tests {
+ use super::{
+ interleaved_image_view, pack_interleaved_image, pack_planar_image, planar_image_view,
+ };
+ use spatialrust_image::{Image, ImageRegion, ImageView, PlanarImage, PlanarImageView};
+
+ #[test]
+ fn packed_images_are_zero_copy_hwc_and_chw() {
+ let interleaved = Image::::try_new(2, 2, (0..12).collect()).unwrap();
+ let tensor = interleaved_image_view(&interleaved).unwrap();
+ assert_eq!(tensor.descriptor().shape(), &[2, 2, 3]);
+ assert_eq!(tensor.allocation_bytes().as_ptr(), interleaved.as_slice().as_ptr().cast());
+
+ let planar =
+ PlanarImage::::try_new(2, 2, (0..12).map(|v| v as f32).collect()).unwrap();
+ let tensor = planar_image_view(&planar).unwrap();
+ assert_eq!(tensor.descriptor().shape(), &[3, 2, 2]);
+ assert_eq!(tensor.allocation_bytes().as_ptr(), planar.as_slice().as_ptr().cast());
+ }
+
+ #[test]
+ fn strided_interleaved_roi_is_explicitly_packed() {
+ let storage = (0_u8..30).collect::>();
+ let parent = ImageView::::new(3, 3, 10, &storage).unwrap();
+ let roi = parent.subview(ImageRegion::new(1, 1, 2, 2)).unwrap();
+ let tensor = pack_interleaved_image(roi).unwrap();
+ assert_eq!(tensor.descriptor().shape(), &[2, 2, 3]);
+ assert_eq!(tensor.allocation_bytes(), &[13, 14, 15, 16, 17, 18, 23, 24, 25, 26, 27, 28]);
+ }
+
+ #[test]
+ fn padded_planar_view_is_explicitly_packed() {
+ let storage = (0_u16..24).collect::>();
+ let view = PlanarImageView::::new(2, 2, 3, 12, &storage).unwrap();
+ let tensor = pack_planar_image(view).unwrap();
+ assert_eq!(tensor.descriptor().shape(), &[2, 2, 2]);
+ assert_eq!(
+ bytemuck::cast_slice::(tensor.allocation_bytes()),
+ &[0, 1, 3, 4, 12, 13, 15, 16]
+ );
+ }
+}
diff --git a/crates/spatialrust-tensor/src/lib.rs b/crates/spatialrust-tensor/src/lib.rs
new file mode 100644
index 0000000..505b49d
--- /dev/null
+++ b/crates/spatialrust-tensor/src/lib.rs
@@ -0,0 +1,763 @@
+//! Small, runtime-independent tensor descriptors and CPU storage views.
+//!
+//! Shape, element strides, byte offset, device, and ownership are explicit.
+//! This crate never uploads, downloads, or otherwise migrates storage. DLPack
+//! FFI and image bridges are additive features built on this data model.
+
+#![deny(unsafe_code)]
+#![warn(missing_docs)]
+
+use std::{any::Any, fmt::Debug, ops::Range, sync::Arc};
+
+#[cfg(feature = "image")]
+mod image;
+#[cfg(feature = "image")]
+pub use image::{
+ interleaved_image_view, pack_interleaved_image, pack_planar_image, planar_image_view,
+ TensorElement,
+};
+#[cfg(feature = "spatial")]
+mod spatial;
+#[cfg(feature = "spatial")]
+pub use spatial::{spatial_f32_field_view, SpatialTensorBridgeError};
+#[cfg(feature = "dlpack")]
+#[allow(unsafe_code)]
+mod dlpack;
+#[cfg(feature = "dlpack")]
+pub use dlpack::{
+ release_dlpack_raw, DlpackError, DlpackExport, DlpackImport, DLPACK_MAJOR, DLPACK_MINOR,
+};
+
+/// Tensor construction and layout errors.
+#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
+pub enum TensorError {
+ /// The data type has no bits or lanes, or cannot address whole bytes.
+ #[error("invalid data type: bits and lanes must be non-zero and form whole bytes")]
+ InvalidDataType,
+ /// The stride rank differs from the shape rank.
+ #[error("stride rank {strides} differs from shape rank {shape}")]
+ RankMismatch {
+ /// Number of shape dimensions.
+ shape: usize,
+ /// Number of strides.
+ strides: usize,
+ },
+ /// Shape, stride, or byte-range arithmetic overflowed.
+ #[error("tensor layout overflows addressable memory")]
+ LayoutOverflow,
+ /// The byte offset or reachable tensor span lies outside storage.
+ #[error("tensor byte range {start}..{end} lies outside {available} bytes")]
+ StorageOutOfBounds {
+ /// First reachable byte.
+ start: i128,
+ /// Exclusive final reachable byte.
+ end: i128,
+ /// Available storage size.
+ available: usize,
+ },
+ /// A host slice was paired with a device that is not host-accessible.
+ #[error("device {0:?} is not directly host-accessible")]
+ DeviceNotHostAccessible(Device),
+ /// Typed storage does not match its descriptor dtype.
+ #[error("typed storage requires {expected:?}, descriptor declares {actual:?}")]
+ DataTypeMismatch {
+ /// Storage dtype.
+ expected: DataType,
+ /// Descriptor dtype.
+ actual: DataType,
+ },
+}
+
+/// DLPack-compatible scalar category without depending on a runtime header.
+#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
+#[repr(u8)]
+pub enum DataTypeCode {
+ /// Signed integer.
+ Int = 0,
+ /// Unsigned integer.
+ UInt = 1,
+ /// IEEE floating point.
+ Float = 2,
+ /// Brain floating point.
+ BFloat = 4,
+ /// Complex floating point.
+ Complex = 5,
+ /// Boolean value.
+ Bool = 6,
+}
+
+/// Scalar category, bit width, and vector lane count.
+#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
+pub struct DataType {
+ code: DataTypeCode,
+ bits: u8,
+ lanes: u16,
+}
+
+impl DataType {
+ /// Unsigned 8-bit scalar.
+ pub const U8: Self = Self::new_unchecked(DataTypeCode::UInt, 8, 1);
+ /// Unsigned 16-bit scalar.
+ pub const U16: Self = Self::new_unchecked(DataTypeCode::UInt, 16, 1);
+ /// Unsigned 32-bit scalar.
+ pub const U32: Self = Self::new_unchecked(DataTypeCode::UInt, 32, 1);
+ /// Signed 8-bit scalar.
+ pub const I8: Self = Self::new_unchecked(DataTypeCode::Int, 8, 1);
+ /// Signed 16-bit scalar.
+ pub const I16: Self = Self::new_unchecked(DataTypeCode::Int, 16, 1);
+ /// Signed 32-bit scalar.
+ pub const I32: Self = Self::new_unchecked(DataTypeCode::Int, 32, 1);
+ /// Signed 64-bit scalar.
+ pub const I64: Self = Self::new_unchecked(DataTypeCode::Int, 64, 1);
+ /// IEEE binary16 scalar.
+ pub const F16: Self = Self::new_unchecked(DataTypeCode::Float, 16, 1);
+ /// Brain floating-point 16-bit scalar.
+ pub const BF16: Self = Self::new_unchecked(DataTypeCode::BFloat, 16, 1);
+ /// IEEE binary32 scalar.
+ pub const F32: Self = Self::new_unchecked(DataTypeCode::Float, 32, 1);
+ /// IEEE binary64 scalar.
+ pub const F64: Self = Self::new_unchecked(DataTypeCode::Float, 64, 1);
+ /// Eight-bit boolean scalar, matching common DLPack producers.
+ pub const BOOL: Self = Self::new_unchecked(DataTypeCode::Bool, 8, 1);
+
+ const fn new_unchecked(code: DataTypeCode, bits: u8, lanes: u16) -> Self {
+ Self { code, bits, lanes }
+ }
+
+ /// Creates a byte-addressable scalar or vector data type.
+ pub fn try_new(code: DataTypeCode, bits: u8, lanes: u16) -> Result {
+ let total_bits = usize::from(bits)
+ .checked_mul(usize::from(lanes))
+ .ok_or(TensorError::InvalidDataType)?;
+ if bits == 0 || lanes == 0 || total_bits % 8 != 0 {
+ return Err(TensorError::InvalidDataType);
+ }
+ Ok(Self { code, bits, lanes })
+ }
+
+ /// Returns the scalar category.
+ pub const fn code(self) -> DataTypeCode {
+ self.code
+ }
+
+ /// Returns bits per lane.
+ pub const fn bits(self) -> u8 {
+ self.bits
+ }
+
+ /// Returns the vector lane count.
+ pub const fn lanes(self) -> u16 {
+ self.lanes
+ }
+
+ /// Returns bytes occupied by one tensor element.
+ pub fn element_size(self) -> usize {
+ usize::from(self.bits) * usize::from(self.lanes) / 8
+ }
+}
+
+/// Device category represented independently of any execution backend.
+#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
+pub enum DeviceKind {
+ /// Ordinary CPU memory.
+ Cpu,
+ /// CUDA device memory.
+ Cuda,
+ /// CUDA-pinned host memory.
+ CudaHost,
+ /// OpenCL device memory.
+ OpenCl,
+ /// Vulkan device memory.
+ Vulkan,
+ /// Metal device memory.
+ Metal,
+ /// Verilog simulator buffer.
+ Vpi,
+ /// ROCm device memory.
+ Rocm,
+ /// ROCm-pinned host memory.
+ RocmHost,
+ /// Backend-specific external memory.
+ External,
+ /// CUDA managed memory.
+ CudaManaged,
+ /// oneAPI device memory.
+ OneApi,
+ /// WebGPU device memory.
+ WebGpu,
+ /// Hexagon device memory.
+ Hexagon,
+ /// Microsoft MAIA device memory.
+ Maia,
+ /// AWS Trainium device memory.
+ Trainium,
+ /// Google TPU device memory.
+ Tpu,
+ /// Google TPU pinned host memory.
+ TpuHost,
+}
+
+/// Device category and backend-local ordinal.
+#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
+pub struct Device {
+ /// Device category.
+ pub kind: DeviceKind,
+ /// Backend-local device ordinal.
+ pub id: i32,
+}
+
+impl Device {
+ /// Main CPU device.
+ pub const CPU: Self = Self { kind: DeviceKind::Cpu, id: 0 };
+
+ /// Returns whether Rust may safely expose this memory as a host byte slice.
+ pub const fn is_host_accessible(self) -> bool {
+ matches!(
+ self.kind,
+ DeviceKind::Cpu | DeviceKind::CudaHost | DeviceKind::RocmHost | DeviceKind::TpuHost
+ )
+ }
+}
+
+/// Whether a tensor object owns or borrows its backing allocation.
+#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
+pub enum Ownership {
+ /// The object owns and drops its storage.
+ Owned,
+ /// The object cannot outlive borrowed storage.
+ Borrowed,
+}
+
+/// Shape, element strides, type, offset, and device for one tensor.
+#[derive(Clone, Debug, PartialEq, Eq)]
+pub struct TensorDescriptor {
+ dtype: DataType,
+ shape: Vec,
+ strides: Option>,
+ byte_offset: usize,
+ device: Device,
+}
+
+impl TensorDescriptor {
+ /// Creates a compact C-order tensor on a device.
+ pub fn contiguous(dtype: DataType, shape: Vec, device: Device) -> Self {
+ Self { dtype, shape, strides: None, byte_offset: 0, device }
+ }
+
+ /// Creates an explicitly strided tensor. Strides are measured in elements.
+ pub fn try_strided(
+ dtype: DataType,
+ shape: Vec,
+ strides: Vec,
+ byte_offset: usize,
+ device: Device,
+ ) -> Result