From 7f94d27925348f0cf0e4cc200692a4214c3722b2 Mon Sep 17 00:00:00 2001 From: Tristan Simas Date: Tue, 29 Sep 2026 15:58:07 -0400 Subject: [PATCH] Let callable owners declare typed direct-call dtype defaults --- src/arraybridge/decorators.py | 24 ++++++++--- tests/test_callable_dtype_default.py | 61 ++++++++++++++++++++++++++++ 2 files changed, 80 insertions(+), 5 deletions(-) create mode 100644 tests/test_callable_dtype_default.py diff --git a/src/arraybridge/decorators.py b/src/arraybridge/decorators.py index 9612137..b3c4dd0 100644 --- a/src/arraybridge/decorators.py +++ b/src/arraybridge/decorators.py @@ -82,11 +82,13 @@ def annotation_type(cls): return DtypeConversionConfig @classmethod - def parameter(cls) -> inspect.Parameter: + def parameter( + cls, *, default_value: "DtypeConversionConfig | None" = None + ) -> inspect.Parameter: return inspect.Parameter( cls.require_parameter_name(), inspect.Parameter.KEYWORD_ONLY, - default=cls.default_value(), + default=cls.default_value() if default_value is None else default_value, annotation=cls.annotation_type(), ) @@ -355,6 +357,7 @@ def wrap_dtype_preserving_callable( mem_type: MemoryType, *, slice_by_slice_default: bool = False, + dtype_config_default: DtypeConversionConfig | None = None, ): """ Return a callable with ArrayBridge dtype and slice controls. @@ -366,18 +369,25 @@ def wrap_dtype_preserving_callable( input_memory_type = MemoryType(MemoryContractAttribute.INPUT.read(func, mem_type.value)) output_memory_type = MemoryType(MemoryContractAttribute.OUTPUT.read(func, mem_type.value)) scale_func = output_memory_type.scale_dtype + default_dtype_config = ( + DtypeConversionConfig.default_value() + if dtype_config_default is None + else dtype_config_default + ) + if not isinstance(default_dtype_config, DtypeConversionConfig): + raise TypeError("Callable dtype default must be a DtypeConversionConfig.") @functools.wraps(func) def dtype_wrapper(image, *args, **kwargs): # Pipeline runtimes may inject dtype_config; direct calls use the same - # preserve-input default explicitly. + # callable-owned default explicitly (preserve-input when undeclared). slice_by_slice = kwargs.pop( SliceBySliceRuntimeParameter.require_parameter_name(), slice_by_slice_default, ) dtype_config: DtypeConversionConfig = kwargs.pop( DtypeConversionConfig.require_parameter_name(), - DtypeConversionConfig.default_value(), + default_dtype_config, ) dtype_conversion = dtype_config.default_dtype_conversion @@ -427,7 +437,7 @@ def _apply_dtype_conversion(array): ) ) dtype_signature = KeywordOnlySignatureExtension(dtype_signature).with_parameter( - DtypeConversionConfig.parameter() + DtypeConversionConfig.parameter(default_value=default_dtype_config) ) setattr(dtype_wrapper, "__signature__", dtype_signature) @@ -513,6 +523,7 @@ def decorator( oom_recovery=True, contract=None, slice_by_slice_default=False, + dtype_config_default: DtypeConversionConfig | None = None, ): """ Decorator for {mem_type} memory type functions. @@ -524,6 +535,8 @@ def decorator( oom_recovery: Enable automatic OOM recovery (default: True) contract: Optional validation function for outputs slice_by_slice_default: Default for the decorator-owned slice control + dtype_config_default: Callable-owned dtype policy for direct calls; + explicit runtime dtype_config still overrides this default. Returns: Decorated function with memory type metadata and dtype preservation @@ -541,6 +554,7 @@ def inner_decorator(func): func, mem_type, slice_by_slice_default=slice_by_slice_default, + dtype_config_default=dtype_config_default, ) # Apply GPU wrapper if this is a GPU memory type diff --git a/tests/test_callable_dtype_default.py b/tests/test_callable_dtype_default.py new file mode 100644 index 0000000..e3bec06 --- /dev/null +++ b/tests/test_callable_dtype_default.py @@ -0,0 +1,61 @@ +"""Callable defaults use the existing typed dtype policy and real conversion.""" + +from dataclasses import dataclass +import inspect + +import numpy as np +import pytest + +from arraybridge.decorators import ( + DtypeConversion, + DtypeConversionConfig, + PreserveInputDtypeConfig, + numpy as numpy_func, +) + + +@dataclass(frozen=True) +class NativeConfig(DtypeConversionConfig): + default_dtype_conversion: DtypeConversion = DtypeConversion.NATIVE_OUTPUT + + +def test_declared_native_output_preserves_fraction_negative_and_large_values(): + native = NativeConfig() + + @numpy_func(dtype_config_default=native) + def correct(image): + return np.array([-1.25, 12.5, 70000.5], dtype=np.float32) + + output = correct(np.array([0, 1, 65535], dtype=np.uint16)) + np.testing.assert_array_equal(output, [-1.25, 12.5, 70000.5]) + assert output.dtype == np.float32 + assert inspect.signature(correct).parameters["dtype_config"].default is native + + +def test_explicit_runtime_dtype_policy_overrides_callable_default(): + @numpy_func(dtype_config_default=NativeConfig()) + def correct(image): + return image.astype(np.float32) + 0.25, np.array([2.5], dtype=np.float32) + + output, sidecar = correct( + np.array([1, 2], dtype=np.uint16), + dtype_config=PreserveInputDtypeConfig(), + ) + assert output.dtype == np.uint16 + assert sidecar.dtype == np.float32 + np.testing.assert_array_equal(sidecar, [2.5]) + + +def test_undeclared_defaults_still_preserve_input_dtype(): + @numpy_func + def correct(image): + return image.astype(np.float32) + 0.25 + + assert correct(np.array([1, 2], dtype=np.uint16)).dtype == np.uint16 + + +def test_untyped_default_is_rejected_at_decoration(): + with pytest.raises(TypeError, match="DtypeConversionConfig"): + numpy_func(dtype_config_default={"default_dtype_conversion": "native"})( + lambda image: image + )