Skip to content
Closed
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
77 changes: 51 additions & 26 deletions exca/cachedict/contiguous.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,29 +24,31 @@


class ContiguousMemmap:
"""Proxy around a memmap view that reads data via file I/O.

View operations (__getitem__ with basic indexing) delegate to the
underlying memmap — they only adjust pointers/strides, no page faults.
Data is materialized via file I/O only when explicitly consumed through
``np.asarray()``.

Non-contiguous access (fancy indexing, strided slicing) raises TypeError;
the caller should materialize first via ``np.asarray()``.

An optional *cache* ``threading.local`` (e.g. ``DumpContext._resource_cache``)
stores open file handles keyed by ``("ContiguousMemmap", path)``, isolated
per thread and per process (fork-safe). When *cache* is ``None``, a
module-level fallback is used.
"""Lazy proxy for contiguous views of a memmap-backed array.

Fancy indexing, non-contiguous slicing, arithmetic, and value-consuming
ndarray methods require materialization with ``np.asarray()``.

Parameters
----------
arr : np.ndarray
Memmap-backed array.
cache : threading.local, optional
Thread-local file-handle cache. If None, a module-level fallback is used.
multiplier : np.ndarray or ContiguousMemmap, optional
Values broadcast to ``arr`` and multiplied when read.
"""

__slots__ = ("_arr", "_mm", "_cache")
__slots__ = ("_arr", "_mm", "_cache", "_multiplier")

def __init__(self, arr: np.ndarray, cache: threading.local | None = None) -> None:
# Walk the base chain to the root file-level memmap.
# Slicing a np.memmap produces another np.memmap whose .offset is
# copied (not recalculated); only the root's data pointer is
# consistent with its .offset for file seeks.
def __init__(
self,
arr: np.ndarray,
cache: threading.local | None = None,
*,
multiplier: np.ndarray | tp.Self | None = None,
) -> None:
# root memmap: sliced views retain the parent's file offset
mm: np.ndarray = arr
while isinstance(getattr(mm, "base", None), np.ndarray):
mm = mm.base # type: ignore[assignment]
Expand All @@ -58,6 +60,9 @@ def __init__(self, arr: np.ndarray, cache: threading.local | None = None) -> Non
self._arr = arr
self._mm: np.memmap = mm
self._cache = cache if cache is not None else _FILE_HANDLE_CACHE
if multiplier is not None:
multiplier = np.broadcast_to(np.asarray(multiplier), arr.shape)
self._multiplier: np.ndarray | None = multiplier

def _byte_range(self, arr: np.ndarray) -> tuple[int, int]:
"""Return (byte_offset, byte_span) of *arr* relative to the root memmap."""
Expand All @@ -71,8 +76,14 @@ def _byte_range(self, arr: np.ndarray) -> tuple[int, int]:
def __len__(self) -> int:
return len(self._arr)

@property
def dtype(self) -> np.dtype[tp.Any]:
if self._multiplier is None:
return self._arr.dtype
return np.result_type(self._arr.dtype, self._multiplier.dtype)

def __repr__(self) -> str:
return f"ContiguousMemmap(shape={self._arr.shape}, dtype={self._arr.dtype})"
return f"ContiguousMemmap(shape={self._arr.shape}, dtype={self.dtype})"

def __getitem__(self, key: tp.Any) -> tp.Any:
keys = key if isinstance(key, tuple) else (key,)
Expand All @@ -82,16 +93,19 @@ def __getitem__(self, key: tp.Any) -> tp.Any:
"— use np.asarray(arr)[key] to read data first."
)
result = self._arr[key]
multiplier = None if self._multiplier is None else self._multiplier[key]
if not isinstance(result, np.ndarray):
return result # scalar
if multiplier is None:
return result
return np.multiply(result, multiplier, dtype=self.dtype)
if result.size == 0:
return np.empty(result.shape, dtype=result.dtype)
return np.empty(result.shape, dtype=self.dtype)
if any(s < 0 for s in result.strides):
raise TypeError("Non-contiguous read — use np.asarray(arr)[key] instead.")
_, span = self._byte_range(result)
if span != result.size * result.dtype.itemsize:
raise TypeError("Non-contiguous read — use np.asarray(arr)[key] instead.")
return ContiguousMemmap(result, self._cache)
return ContiguousMemmap(result, self._cache, multiplier=multiplier)

def __array__(
self, dtype: np.dtype[tp.Any] | None = None, copy: bool | None = None
Expand Down Expand Up @@ -122,6 +136,8 @@ def __array__(
strides=self._arr.strides,
)
result = np.ascontiguousarray(view)
if self._multiplier is not None:
result = np.multiply(result, self._multiplier, dtype=self.dtype)
if dtype is not None:
result = result.astype(dtype)
return result
Expand All @@ -142,14 +158,23 @@ def __getattr__(self, name: str) -> tp.Any:
if name in _SAFE_ATTRS:
return getattr(self._arr, name)
if name in _VIEW_OPS:
if self._multiplier is not None and name not in ("T", "transpose"):
raise AttributeError(
f"ContiguousMemmap with a multiplier does not support '.{name}' "
f"directly — use np.asarray(arr).{name} instead."
)
val = getattr(self._arr, name)
multiplier = (
None if self._multiplier is None else getattr(self._multiplier, name)
)
if callable(val):

def _wrap(*a: tp.Any, **kw: tp.Any) -> "ContiguousMemmap":
return ContiguousMemmap(val(*a, **kw), self._cache)
mult = None if multiplier is None else multiplier(*a, **kw)
return ContiguousMemmap(val(*a, **kw), self._cache, multiplier=mult)

return _wrap
return ContiguousMemmap(val, self._cache)
return ContiguousMemmap(val, self._cache, multiplier=multiplier)
if hasattr(np.ndarray, name):
raise AttributeError(
f"ContiguousMemmap does not support '.{name}' directly "
Expand Down
27 changes: 25 additions & 2 deletions exca/cachedict/test_contiguous.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,12 +15,16 @@
from .dumpcontext import DumpContext


def _make_cm(tmp_path: Path, arr: np.ndarray) -> ContiguousMemmap:
def _make_cm(
tmp_path: Path,
arr: np.ndarray,
multiplier: np.ndarray | ContiguousMemmap | None = None,
) -> ContiguousMemmap:
"""Dump an array via MemmapArray and return it wrapped as ContiguousMemmap."""
ctx = DumpContext(tmp_path)
with ctx:
info = ctx.dump(arr, cache_type="MemmapArray")
return ContiguousMemmap(ctx.load(info))
return ContiguousMemmap(ctx.load(info), multiplier=multiplier)


# =============================================================================
Expand Down Expand Up @@ -59,6 +63,25 @@ def test_roundtrip(tmp_path: Path) -> None:
assert isinstance(empty, np.ndarray) and empty.shape == (0, 30)


def test_multiplier(tmp_path: Path) -> None:
stored = np.arange(12, dtype=np.float16).reshape(4, 3)
multiplier = np.array([[0.25], [2.0], [16.0], [0.5]], dtype=np.float32)
expected = stored * multiplier
cached_multiplier = _make_cm(tmp_path / "multiplier", multiplier)
cm = _make_cm(tmp_path / "data", stored, cached_multiplier)

assert cm.dtype == np.float32
assert isinstance(cm._multiplier, np.ndarray)
assert cm._multiplier.strides == (4, 0)
np.testing.assert_array_equal(np.asarray(cm), expected)
assert cm[2, 1] == expected[2, 1]

moved = np.moveaxis(cm, 0, -1) # type: ignore[type-var]
np.testing.assert_array_equal(np.asarray(moved[..., 1:3]), expected.T[..., 1:3])
with pytest.raises(AttributeError, match="multiplier"):
cm.reshape(2, 6)


# =============================================================================
# View operations
# =============================================================================
Expand Down
Loading