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
1 change: 1 addition & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -75,6 +75,7 @@ all = [
]
dev = [
"pytest>=7.0",
"dill>=0.3.8",
"pytest-cov>=4.0",
"build>=1.0",
"twine>=5.0",
Expand Down
27 changes: 15 additions & 12 deletions src/arraybridge/decorators.py
Original file line number Diff line number Diff line change
Expand Up @@ -264,12 +264,22 @@ def _insertion_index(parameters: list[inspect.Parameter]) -> int:
return len(parameters)


# Thread-local storage for GPU streams and contexts
_thread_gpu_contexts = threading.local()
class ThreadGPUContext:
"""Runtime-owned thread-local streams keyed by framework/device identity.

Keep the runtime handle on its importable owner, not in decorator function
globals. A retained, unpublished decorated callable is serialized by value;
its durable closure must not pull a ``threading.local`` into history.
"""

class ThreadGPUContext:
"""Thread-local streams keyed by framework-local device identity."""
_contexts: ClassVar[threading.local] = threading.local()

@classmethod
def current(cls) -> "ThreadGPUContext":
"""Return this thread's runtime context without serializing its handle."""
if not hasattr(cls._contexts, "context"):
cls._contexts.context = cls()
return cls._contexts.context

def __init__(self):
self._streams: dict[tuple[MemoryType, int], Any] = {}
Expand Down Expand Up @@ -301,13 +311,6 @@ def stream_for(
return device_id, self._streams[key]


def _get_thread_gpu_context():
"""Get or create thread-local GPU context."""
if not hasattr(_thread_gpu_contexts, "context"):
_thread_gpu_contexts.context = ThreadGPUContext()
return _thread_gpu_contexts.context


def memory_types(
input_type: str | MemoryType,
output_type: str | MemoryType,
Expand Down Expand Up @@ -461,7 +464,7 @@ def gpu_wrapper(*args, **kwargs):
# Check if GPU is available for this framework
if framework is not None and mem_type.available_device_ids(framework):
# Get thread-local context
ctx = _get_thread_gpu_context()
ctx = ThreadGPUContext.current()

device_id, stream = ctx.stream_for(mem_type, framework)

Expand Down
50 changes: 50 additions & 0 deletions tests/test_durable_decorator_context.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,50 @@
"""Retained callable declarations exclude thread-local GPU lifecycle handles."""

import threading
from concurrent.futures import ThreadPoolExecutor

import dill
import numpy as np
import pytest

from arraybridge.decorators import ThreadGPUContext, cupy, numpy
from arraybridge.types import MemoryType


def test_context_is_stable_per_thread_and_never_shared_between_threads():
current = ThreadGPUContext.current()
with ThreadPoolExecutor(max_workers=1) as executor:
other, repeated = executor.submit(
lambda: (ThreadGPUContext.current(), ThreadGPUContext.current())
).result()
assert ThreadGPUContext.current() is current
assert other is repeated
assert other is not current


@pytest.mark.parametrize("decorator", [numpy, cupy])
def test_unpublished_decorated_callable_serializes_without_runtime_context(decorator):
# Like a replaced custom function retained by an undo snapshot, this
# declaration has no importable public alias. Dill must persist its body.
@decorator
def declared(image, scale=3):
return image * scale

context = ThreadGPUContext.current()
key = (MemoryType.CUPY, -1)
handle = threading.local()
context._streams[key] = handle
try:
restored = dill.loads(dill.dumps(declared))
assert restored.__name__ == declared.__name__
assert restored.__wrapped__.__name__ == declared.__wrapped__.__name__
# CPU NumPy invocation checks actual restored behavior. GPU behavior
# keeps its existing focused tests; this receipt starts no GPU runtime.
if decorator is numpy:
np.testing.assert_array_equal(
restored(np.array([1, 2], dtype=np.uint16), scale=4), [4, 8]
)
assert ThreadGPUContext.current() is context
assert context._streams[key] is handle
finally:
del context._streams[key]
Loading