Skip to content

Commit cd1e5b3

Browse files
authored
fix(indexing): avoid full-grid allocations for sparse selections (#4356)
* fix(indexing): make sparse selections O(npoints) instead of O(nchunks) CoordinateIndexer (general path and sorted-1D fast path) and IntArrayDimIndexer built a dense per-chunk histogram via np.bincount(..., minlength=nchunks) / np.zeros(nchunks) plus a full-length cumsum. For arrays with very many chunks this allocates memory proportional to the total chunk count regardless of how few points are selected — a 3-point write to an array with 1.4e11 chunks tried to allocate 1 TiB and raised MemoryError. Replace the dense histogram with run boundaries computed directly on the sorted raveled chunk ids (sorted_run_ends), storing a compressed cumsum aligned with the occupied chunks, and index it positionally in __iter__. Memory and time now scale with the number of selected points. Fixes #4174 Assisted-by: ClaudeCode:claude-fable-5 * refactor(indexing): keep dense indexer attributes as deprecated properties The compressed per-occupied-chunk cumsum introduced for gh-4174 changed the observable semantics of chunk_nitems_cumsum (and removed chunk_nitems on IntArrayDimIndexer). Although zarr.core is documented as private API, external code is known to introspect these indexers, so be conservative: store the compressed offsets under a new name (chunk_run_ends, aligned with chunk_rixs / dim_chunk_ixs) and restore chunk_nitems / chunk_nitems_cumsum as properties that lazily rebuild the original dense arrays, warning with ZarrDeprecationWarning. The O(nchunks) cost is now only paid if someone actually accesses them — which was the status quo before the fix. Assisted-by: ClaudeCode:claude-fable-5 * fix(indexing): exact ceildiv everywhere and sparse boolean-axis selections Review follow-ups for the sparse-selection change: - Fix ceildiv itself rather than bypassing it at one call site. It went through float division, so ceildiv(2**62 - 1, 1) returned 2**62; four other callers (slice arithmetic, the regular-grid check, dask-style chunk sizes) had the same latent error. Integers now divide exactly; floats keep the ceil-of-quotient path. FixedDimension goes back to calling it. - BoolArrayDimIndexer still allocated a dense per-chunk count array and ran a Python loop over every chunk, so an orthogonal boolean selection on a finely chunked axis paid O(nchunks) in time and memory on top of the mask. It now derives the occupied chunks and run ends from the selected positions with sorted_run_ends, like the other two indexers, and exposes the same deprecated dense properties. A boolean axis with 2**22 chunks goes from a multi-second loop to ~2 ms. - Changelog: state precisely which selection kinds are covered, that reading the deprecated properties rebuilds the dense array, and the ceildiv correction. Assisted-by: ClaudeCode:claude-fable-5-1 * fix(indexing): defer boolean-axis allocation changes Retain sparse coordinate and integer indexing plus exact ceildiv, while restoring the existing boolean-axis implementation to avoid dense-mask regressions. Assisted-by: Codex:GPT-6 * docs: describe current ceildiv behavior Assisted-by: Codex:GPT-6 * fix(indexing): lazily cache legacy dense attributes Preserve dense attribute identity, mutation behavior, and dataclass fields without deprecation warnings while keeping normal indexing sparse. Assisted-by: Codex:GPT-6 * fix(indexing): use a dedicated integer ceiling division helper Restore the original ceildiv behavior and exports, and use ceildiv_int for integer-only chunk and slice calculations without type-based dispatch. Assisted-by: Codex:GPT-6 * refactor: remove unnecessary ceildiv re-exports Assisted-by: Codex:GPT-6 * test: use Expect cases for integer ceiling division Assisted-by: Codex:GPT-6 * test(indexing): restore the sparse projection property test The test was added in a merge commit, which the linear rebase onto main dropped. Assisted-by: ClaudeCode:claude-opus-5-5 * docs: number the changelog fragment after the merging PR Towncrier renders fragment names as pull-request links, so the fragment takes #4356 rather than issue #4174. The paragraph describing #4218 is dropped: it shipped in the 3.4.0 release notes. Assisted-by: ClaudeCode:claude-opus-5-5
1 parent 3799b64 commit cd1e5b3

7 files changed

Lines changed: 364 additions & 50 deletions

File tree

‎changes/4356.bugfix.md‎

Lines changed: 16 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,16 @@
1+
Coordinate selections (including `vindex` masks) and orthogonal integer-array
2+
selections no longer allocate counts proportional to the total number of chunks.
3+
Previously, selecting a few points on an array with billions of chunks could raise
4+
`MemoryError`. The indexers now store occupied chunk IDs and run-end offsets
5+
(`chunk_run_ends`); unsorted selections can still require sorting. The dense
6+
`chunk_nitems` / `chunk_nitems_cumsum` attributes of `CoordinateIndexer` and
7+
`IntArrayDimIndexer` remain available without deprecation warnings. Each array is
8+
materialized on first access and cached, preserving its identity and in-place
9+
mutations on subsequent reads. These attributes remain dataclass fields, so
10+
`dataclasses.asdict()` materializes them; `repr` omits them to avoid dense allocation
11+
during inspection. Accessing a dense attribute still allocates memory proportional
12+
to the full chunk grid.
13+
14+
The new `zarr.core.common.ceildiv_int` helper uses exact integer arithmetic for
15+
chunk counts and slice calculations. For example, `ceildiv_int(2**62 - 1, 1)`
16+
returns `2**62 - 1`. The existing `ceildiv` helper retains its floating-point behavior.

‎src/zarr/core/array.py‎

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -69,7 +69,7 @@
6969
ZarrFormat,
7070
_default_zarr_format,
7171
_warn_order_kwarg,
72-
ceildiv,
72+
ceildiv_int,
7373
concurrent_map,
7474
parse_shapelike,
7575
product,
@@ -195,7 +195,7 @@ def _chunk_sizes_from_shape(
195195
"""Compute dask-style chunk sizes from an array shape and uniform chunk shape."""
196196
result: list[tuple[int, ...]] = []
197197
for s, c in zip(array_shape, chunk_shape, strict=True):
198-
nchunks = ceildiv(s, c)
198+
nchunks = ceildiv_int(s, c)
199199
sizes = tuple(min(c, s - i * c) for i in range(nchunks))
200200
result.append(sizes)
201201
return tuple(result)
@@ -1147,7 +1147,7 @@ def _chunk_grid_shape(self) -> tuple[int, ...]:
11471147
if (sharding_codec := _sharding_codec(self.metadata)) is not None:
11481148
# When sharding, count inner chunks across the whole array
11491149
chunk_shape = sharding_codec.chunk_shape
1150-
return tuple(starmap(ceildiv, zip(self.shape, chunk_shape, strict=True)))
1150+
return tuple(starmap(ceildiv_int, zip(self.shape, chunk_shape, strict=True)))
11511151
return self._chunk_grid.grid_shape
11521152

11531153
@property

‎src/zarr/core/chunk_grids.py‎

Lines changed: 5 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -25,7 +25,7 @@
2525
import zarr
2626
from zarr.core.common import (
2727
ShapeLike,
28-
ceildiv,
28+
ceildiv_int,
2929
parse_shapelike,
3030
)
3131
from zarr.errors import ZarrUserWarning
@@ -62,7 +62,7 @@ def __post_init__(self) -> None:
6262
if self.size == 0:
6363
n = 0
6464
else:
65-
n = ceildiv(self.extent, self.size)
65+
n = ceildiv_int(self.extent, self.size)
6666
object.__setattr__(self, "nchunks", n)
6767
object.__setattr__(self, "ngridcells", n)
6868

@@ -465,7 +465,9 @@ def from_sizes(
465465
if (
466466
edges_list[0] > 0
467467
and all(e == edges_list[0] for e in edges_list)
468-
and (extent == edge_sum or len(edges_list) == ceildiv(extent, edges_list[0]))
468+
and (
469+
extent == edge_sum or len(edges_list) == ceildiv_int(extent, edges_list[0])
470+
)
469471
):
470472
dims.append(FixedDimension(size=edges_list[0], extent=extent))
471473
else:

‎src/zarr/core/common.py‎

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -92,11 +92,17 @@ def product(tup: tuple[int, ...]) -> int:
9292

9393

9494
def ceildiv(a: float, b: float) -> int:
95+
"""Ceiling of ``a / b`` using floating-point division; zero when ``a`` is zero."""
9596
if a == 0:
9697
return 0
9798
return math.ceil(a / b)
9899

99100

101+
def ceildiv_int(a: int, b: int) -> int:
102+
"""Ceiling of integer division using exact Python integer arithmetic."""
103+
return -(-int(a) // int(b))
104+
105+
100106
def concurrent_iter[T: tuple[Any, ...], V](
101107
items: Iterable[T],
102108
func: Callable[..., Awaitable[V]],

‎src/zarr/core/indexing.py‎

Lines changed: 85 additions & 42 deletions
Original file line numberDiff line numberDiff line change
@@ -1,10 +1,9 @@
11
from __future__ import annotations
22

33
import itertools
4-
import math
54
import numbers
65
from collections.abc import Iterator, Sequence
7-
from dataclasses import dataclass
6+
from dataclasses import dataclass, field
87
from enum import Enum
98
from functools import lru_cache
109
from types import EllipsisType
@@ -23,7 +22,7 @@
2322
import numpy.typing as npt
2423

2524
from zarr.core.chunk_grids import FixedDimension
26-
from zarr.core.common import ceildiv, product
25+
from zarr.core.common import ceildiv_int, product
2726
from zarr.core.metadata.v2 import ArrayV2Metadata
2827
from zarr.core.metadata.v3 import ArrayV3Metadata
2928
from zarr.errors import (
@@ -207,7 +206,7 @@ def _iter_regions(
207206
# ((slice(0, 1, 1), slice(0, 2, 1)), (slice(1, 2, 1), slice(0, 2, 1)))
208207
```
209208
"""
210-
grid_shape = tuple(itertools.starmap(ceildiv, zip(domain_shape, region_shape, strict=True)))
209+
grid_shape = tuple(itertools.starmap(ceildiv_int, zip(domain_shape, region_shape, strict=True)))
211210
for grid_position in _iter_grid(
212211
grid_shape=grid_shape, origin=origin, selection_shape=selection_shape, order=order
213212
):
@@ -415,7 +414,7 @@ def __init__(
415414

416415
object.__setattr__(self, "dim_len", dim_len)
417416
object.__setattr__(self, "dim_grid", dim_grid)
418-
object.__setattr__(self, "nitems", max(0, ceildiv((stop - start), step)))
417+
object.__setattr__(self, "nitems", max(0, ceildiv_int((stop - start), step)))
419418
object.__setattr__(self, "nchunks", dim_grid.nchunks)
420419

421420
def __iter__(self) -> Iterator[ChunkDimProjection]:
@@ -441,7 +440,7 @@ def __iter__(self) -> Iterator[ChunkDimProjection]:
441440
if remainder:
442441
dim_chunk_sel_start += self.step - remainder
443442
# compute number of previous items, provides offset into output array
444-
dim_out_offset = ceildiv((dim_offset - self.start), self.step)
443+
dim_out_offset = ceildiv_int((dim_offset - self.start), self.step)
445444
else:
446445
# selection starts within current chunk
447446
dim_chunk_sel_start = self.start - dim_offset
@@ -455,7 +454,7 @@ def __iter__(self) -> Iterator[ChunkDimProjection]:
455454
dim_chunk_sel_stop = self.stop - dim_offset
456455

457456
dim_chunk_sel = slice(dim_chunk_sel_start, dim_chunk_sel_stop, self.step)
458-
dim_chunk_nitems = ceildiv((dim_chunk_sel_stop - dim_chunk_sel_start), self.step)
457+
dim_chunk_nitems = ceildiv_int((dim_chunk_sel_stop - dim_chunk_sel_start), self.step)
459458

460459
# If there are no elements on the selection within this chunk, then skip
461460
if dim_chunk_nitems == 0:
@@ -744,6 +743,23 @@ def boundscheck_indices(x: npt.NDArray[Any], dim_len: int) -> None:
744743
raise BoundsCheckError(msg)
745744

746745

746+
def sorted_run_ends(
747+
a: npt.NDArray[Any],
748+
) -> tuple[npt.NDArray[np.intp], npt.NDArray[np.intp]]:
749+
"""Group a sorted 1-D integer array into runs of equal values.
750+
751+
Returns `(values, run_ends)` where `values` holds the distinct values in order and
752+
`run_ends[i]` is the exclusive end offset of run `i` in `a`. Cost is O(len(a)),
753+
independent of the range of values — unlike a dense `np.bincount` histogram, which
754+
allocates O(max value) memory (see gh-4174).
755+
"""
756+
if a.size == 0:
757+
return np.empty(0, dtype=np.intp), np.empty(0, dtype=np.intp)
758+
run_starts = np.concatenate(([0], np.nonzero(np.diff(a))[0] + 1))
759+
run_ends = np.append(run_starts[1:], a.size).astype(np.intp, copy=False)
760+
return a[run_starts].astype(np.intp, copy=False), run_ends
761+
762+
747763
@dataclass(frozen=True)
748764
class IntArrayDimIndexer:
749765
"""Integer array selection against a single dimension."""
@@ -755,9 +771,12 @@ class IntArrayDimIndexer:
755771
order: Order
756772
dim_sel: npt.NDArray[np.intp]
757773
dim_out_sel: npt.NDArray[np.intp]
758-
chunk_nitems: int
774+
# Dense compatibility arrays are populated on first access, not during indexing.
775+
chunk_nitems: npt.NDArray[np.intp] = field(repr=False)
759776
dim_chunk_ixs: npt.NDArray[np.intp]
760-
chunk_nitems_cumsum: npt.NDArray[np.intp]
777+
chunk_nitems_cumsum: npt.NDArray[np.intp] = field(repr=False)
778+
# end offset of each occupied chunk's run of selected items, aligned with dim_chunk_ixs
779+
chunk_run_ends: npt.NDArray[np.intp]
761780

762781
def __init__(
763782
self,
@@ -802,23 +821,21 @@ def __init__(
802821

803822
if order == Order.INCREASING:
804823
dim_out_sel = None
824+
dim_sel_chunk_sorted = dim_sel_chunk
805825
elif order == Order.DECREASING:
806826
dim_sel = dim_sel[::-1]
807827
# TODO should be possible to do this without creating an arange
808828
dim_out_sel = np.arange(nitems - 1, -1, -1)
829+
dim_sel_chunk_sorted = dim_sel_chunk[::-1]
809830
else:
810831
# sort indices to group by chunk
811832
dim_out_sel = np.argsort(dim_sel_chunk)
812833
dim_sel = np.take(dim_sel, dim_out_sel)
834+
dim_sel_chunk_sorted = dim_sel_chunk[dim_out_sel]
813835

814-
# precompute number of selected items for each chunk
815-
chunk_nitems = np.bincount(dim_sel_chunk, minlength=nchunks)
816-
817-
# find chunks that we need to visit
818-
dim_chunk_ixs = np.nonzero(chunk_nitems)[0]
819-
820-
# compute offsets into the output array
821-
chunk_nitems_cumsum = np.cumsum(chunk_nitems)
836+
# the chunks to visit and, per occupied chunk, the end offset of its run of
837+
# selected items — O(nitems), never O(nchunks)
838+
dim_chunk_ixs, chunk_run_ends = sorted_run_ends(dim_sel_chunk_sorted)
822839

823840
# store attributes
824841
object.__setattr__(self, "dim_len", dim_len)
@@ -828,21 +845,32 @@ def __init__(
828845
object.__setattr__(self, "order", order)
829846
object.__setattr__(self, "dim_sel", dim_sel)
830847
object.__setattr__(self, "dim_out_sel", dim_out_sel)
831-
object.__setattr__(self, "chunk_nitems", chunk_nitems)
832848
object.__setattr__(self, "dim_chunk_ixs", dim_chunk_ixs)
833-
object.__setattr__(self, "chunk_nitems_cumsum", chunk_nitems_cumsum)
849+
object.__setattr__(self, "chunk_run_ends", chunk_run_ends)
850+
851+
def __getattr__(self, name: str) -> npt.NDArray[np.intp]:
852+
if name not in ("chunk_nitems", "chunk_nitems_cumsum"):
853+
raise AttributeError(f"{type(self).__name__!r} object has no attribute {name!r}")
854+
dense = np.zeros(self.nchunks, dtype=np.intp)
855+
dense[self.dim_chunk_ixs] = np.diff(self.chunk_run_ends, prepend=0)
856+
if name == "chunk_nitems_cumsum":
857+
np.cumsum(dense, out=dense)
858+
object.__setattr__(self, name, dense)
859+
return dense
834860

835861
def __iter__(self) -> Iterator[ChunkDimProjection]:
836862
g = self.dim_grid
863+
dense_cumsum = self.__dict__.get("chunk_nitems_cumsum")
837864

838-
for dim_chunk_ix in self.dim_chunk_ixs:
865+
for i, dim_chunk_ix in enumerate(self.dim_chunk_ixs):
839866
dim_out_sel: slice | npt.NDArray[np.intp]
840867
# find region in output
841-
if dim_chunk_ix == 0:
842-
start = 0
868+
if dense_cumsum is None:
869+
start = 0 if i == 0 else self.chunk_run_ends[i - 1]
870+
stop = self.chunk_run_ends[i]
843871
else:
844-
start = self.chunk_nitems_cumsum[dim_chunk_ix - 1]
845-
stop = self.chunk_nitems_cumsum[dim_chunk_ix]
872+
start = 0 if dim_chunk_ix == 0 else dense_cumsum[dim_chunk_ix - 1]
873+
stop = dense_cumsum[dim_chunk_ix]
846874
if self.order == Order.INCREASING:
847875
dim_out_sel = slice(start, stop)
848876
else:
@@ -1180,7 +1208,11 @@ class CoordinateIndexer(Indexer):
11801208
sel_shape: tuple[int, ...]
11811209
selection: CoordinateSelectionNormalized
11821210
sel_sort: npt.NDArray[np.intp] | None
1183-
chunk_nitems_cumsum: npt.NDArray[np.intp]
1211+
# Exclude the lazy dense field from repr to keep inspection sparse.
1212+
chunk_nitems_cumsum: npt.NDArray[np.intp] = field(repr=False)
1213+
cdata_shape: tuple[int, ...]
1214+
# end offset of each occupied chunk's run of selected points, aligned with chunk_rixs
1215+
chunk_run_ends: npt.NDArray[np.intp]
11841216
chunk_rixs: npt.NDArray[np.intp]
11851217
chunk_mixs: tuple[npt.NDArray[np.intp], ...]
11861218
shape: tuple[int, ...]
@@ -1197,7 +1229,6 @@ def __init__(
11971229
cdata_shape = (1,)
11981230
else:
11991231
cdata_shape = tuple(g.nchunks for g in dim_grids)
1200-
nchunks = math.prod(cdata_shape)
12011232

12021233
# some initial normalization
12031234
selection_normalized = cast("CoordinateSelectionNormalized", ensure_tuple(selection))
@@ -1262,15 +1293,15 @@ def __init__(
12621293
edges = np.arange(first + 1, last + 1, dtype=coords.dtype) * size
12631294
cuts = np.searchsorted(coords, edges)
12641295
counts = np.diff(cuts, prepend=0, append=coords.size)
1265-
chunk_rixs = (first + np.nonzero(counts)[0]).astype(np.intp)
1266-
chunk_nitems = np.zeros(nchunks, dtype=np.intp)
1267-
chunk_nitems[first : last + 1] = counts
1268-
chunk_nitems_cumsum = np.cumsum(chunk_nitems)
1296+
occupied = np.nonzero(counts)[0]
1297+
chunk_rixs = (first + occupied).astype(np.intp)
1298+
chunk_run_ends = np.cumsum(counts[occupied])
12691299

12701300
object.__setattr__(self, "sel_shape", coords.shape)
12711301
object.__setattr__(self, "selection", (coords,))
12721302
object.__setattr__(self, "sel_sort", None)
1273-
object.__setattr__(self, "chunk_nitems_cumsum", chunk_nitems_cumsum)
1303+
object.__setattr__(self, "cdata_shape", cdata_shape)
1304+
object.__setattr__(self, "chunk_run_ends", chunk_run_ends)
12741305
object.__setattr__(self, "chunk_rixs", chunk_rixs)
12751306
object.__setattr__(self, "chunk_mixs", (chunk_rixs,))
12761307
object.__setattr__(self, "dim_grids", dim_grids)
@@ -1315,39 +1346,51 @@ def __init__(
13151346
# optimisation, only sort if needed
13161347
sel_sort = np.argsort(chunks_raveled_indices)
13171348
selection_broadcast = tuple(dim_sel[sel_sort] for dim_sel in selection_broadcast)
1349+
chunks_raveled_indices = chunks_raveled_indices[sel_sort]
13181350
else:
13191351
sel_sort = None
13201352

13211353
shape = selection_broadcast[0].shape or (1,)
13221354

1323-
# precompute number of selected items for each chunk
1324-
chunk_nitems = np.bincount(chunks_raveled_indices, minlength=nchunks)
1325-
chunk_nitems_cumsum = np.cumsum(chunk_nitems)
1326-
# locate the chunks we need to process
1327-
chunk_rixs = np.nonzero(chunk_nitems)[0]
1355+
# the chunks to visit and, per occupied chunk, the end offset of its run of
1356+
# selected points — O(npoints), never O(nchunks)
1357+
chunk_rixs, chunk_run_ends = sorted_run_ends(chunks_raveled_indices)
13281358

13291359
# unravel chunk indices
13301360
chunk_mixs = np.unravel_index(chunk_rixs, cdata_shape)
13311361

13321362
object.__setattr__(self, "sel_shape", sel_shape)
13331363
object.__setattr__(self, "selection", selection_broadcast)
13341364
object.__setattr__(self, "sel_sort", sel_sort)
1335-
object.__setattr__(self, "chunk_nitems_cumsum", chunk_nitems_cumsum)
1365+
object.__setattr__(self, "cdata_shape", cdata_shape)
1366+
object.__setattr__(self, "chunk_run_ends", chunk_run_ends)
13361367
object.__setattr__(self, "chunk_rixs", chunk_rixs)
13371368
object.__setattr__(self, "chunk_mixs", chunk_mixs)
13381369
object.__setattr__(self, "dim_grids", dim_grids)
13391370
object.__setattr__(self, "shape", shape)
13401371
object.__setattr__(self, "drop_axes", ())
13411372

1373+
def __getattr__(self, name: str) -> npt.NDArray[np.intp]:
1374+
if name != "chunk_nitems_cumsum":
1375+
raise AttributeError(f"{type(self).__name__!r} object has no attribute {name!r}")
1376+
dense = np.zeros(product(self.cdata_shape), dtype=np.intp)
1377+
dense[self.chunk_rixs] = np.diff(self.chunk_run_ends, prepend=0)
1378+
np.cumsum(dense, out=dense)
1379+
object.__setattr__(self, name, dense)
1380+
return dense
1381+
13421382
def __iter__(self) -> Iterator[ChunkProjection]:
1383+
dense_cumsum = self.__dict__.get("chunk_nitems_cumsum")
13431384
# iterate over chunks
1344-
for i, chunk_rix in enumerate(self.chunk_rixs):
1385+
for i in range(len(self.chunk_rixs)):
13451386
chunk_coords = tuple(m[i] for m in self.chunk_mixs)
1346-
if chunk_rix == 0:
1347-
start = 0
1387+
if dense_cumsum is None:
1388+
start = 0 if i == 0 else self.chunk_run_ends[i - 1]
1389+
stop = self.chunk_run_ends[i]
13481390
else:
1349-
start = self.chunk_nitems_cumsum[chunk_rix - 1]
1350-
stop = self.chunk_nitems_cumsum[chunk_rix]
1391+
chunk_rix = self.chunk_rixs[i]
1392+
start = 0 if chunk_rix == 0 else dense_cumsum[chunk_rix - 1]
1393+
stop = dense_cumsum[chunk_rix]
13511394
out_selection: slice | npt.NDArray[np.intp]
13521395
if self.sel_sort is None:
13531396
out_selection = slice(start, stop)

0 commit comments

Comments
 (0)