Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
24 changes: 19 additions & 5 deletions src/arraybridge/decorators.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(),
)

Expand Down Expand Up @@ -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.
Expand All @@ -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

Expand Down Expand Up @@ -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)

Expand Down Expand Up @@ -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.
Expand All @@ -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
Expand All @@ -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
Expand Down
61 changes: 61 additions & 0 deletions tests/test_callable_dtype_default.py
Original file line number Diff line number Diff line change
@@ -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
)
Loading