Skip to content
Open
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
64 changes: 49 additions & 15 deletions DESIGN.md
Original file line number Diff line number Diff line change
Expand Up @@ -55,6 +55,7 @@ arrive DLPack-ready with no labeled-array machinery in the loop. (xarray is stil
| Neighbor | Why insitubatch is different |
|---|-------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------|
| **MosaicML Streaming / WebDataset** | They require **resharding** into a sample-oriented format (MDS/tar) — a full ETL copy, and a "sample" becomes an opaque blob. insitubatch trains **in place** on the existing ndim Zarr; splits/shuffle/batches live in **coordinate space**. |
| **annbatch (scverse)** | The nearest neighbor and the sharpest mirror — same premise (loaders are the bottleneck; large contiguous reads feed the GPU), **opposite bet on where randomness comes from**. annbatch runs a **preshuffler** that rewrites a randomized Zarr copy, then streams big contiguous slices through an in-memory shuffle buffer that **breaks chunk structure** (a chunk's rows scatter across batches). Clean global shuffle, paid for with a **rewrite** — which also freezes the sampling policy into byte order (their noted limit: weighted sampling needs re-writing). insitubatch trains **in place**: bounded block-shuffle over windowed anchors on the native grid, chunks kept **resident and reused**. Two points on one spectrum — annbatch owns *rewrite-is-fine, 1-D obs rows, local disk* (single-cell anndata); insitubatch owns *immutable, remote, PB-scale, n-D-windowed, multi-variable*. Both carry one bounded-randomness knob (their buffer `m` ≈ our `block_chunks`); the difference is the currency spent — a one-time rewrite vs. read locality. |
| **xbatcher + DataLoader (worker stack)** | xbatcher (an xarray-contrib community project) is the standard, elegant way to *define* ndim batches — **xarray-native**, yielding `xr.DataArray`. Its torch-worker engine (N **processes**) is strongest at the GRIB / one-sample-per-chunk end, and elsewhere pays worker cold start, memory ∝ workers, and per-sample decode on the uncached path. insitubatch keeps the same ndim batch semantics but stays at the **numpy/tensor** level on a *different engine* — one async loop — winning **cold start** (inference) and **memory** across the chunk spectrum (training). The caches differ in kind: insitu's in-place decoded-chunk pool (deduped, no second copy) vs xbatcher's opt-in **materialized-batch** zarr store; both now persist across runs (insitu via `persist=True`). Complementary tools; pick by regime (and we tune xbatcher well before any comparison). |
| **Earth2Studio (NVIDIA)** | An **xarray-centric** inference framework: its `DataSource` yields `xr.DataArray`, and xarray is load-bearing down to `prep_data_array`. insitubatch doesn't build xarray — *inside* their loop the win is an obstore store-swap (an obstore contribution, not ours); *around* their models it feeds `(tensor, coords)` batches straight to `model.create_iterator`, where `coords` is a light metadata dict, not the xarray machinery. |
| **zarrs / tensorstore / zarr-python** | **Substrate, not peers.** These are zarr *implementations* — chunk-granular read + codec pipelines. We consume one (zarr-python's, over an obstore/arraylake/fsspec `Store`) and build the sample-oriented layer above it: dedup, splits, shuffle, residency, batch assembly. A faster implementation underneath is a *win we inherit*, not a competitor — see the zarrs codec-pipeline spike in Open questions. |
Expand All @@ -81,8 +82,8 @@ a bounded decode/shuffle buffer. Torch runs `num_workers=0`, `batch_size=None`.
## The central abstraction: the read plan

The unit of work is neither *sample* nor *chunk* — it's a **read plan**:
required samples → **deduplicated** set of chunk reads + a gather map back to
samples. This makes the whole spectrum one code path:
required samples → **deduplicated** set of chunk reads, gathered back to samples
(the index recomputed at gather, not stored). This makes the whole spectrum one code path:

```
fat chunks ──────────────────────────────► GRIB-per-timestep (degenerate)
Expand All @@ -94,6 +95,13 @@ shared-cache + intra-chunk shuffle win async fan-out is the whole game
`build_stored_chunk_reads()` is vectorized: Python touches **O(reads)**, never
O(samples).

The plan buys two orthogonal guarantees — **read-once** (a stored tile fetched and
decoded once, however many samples touch it) and **sample-once** (each valid sample in
exactly one batch, tracked by the `order` ledger, with the fancy index recomputed in
`gather` and never stored) — both independent of batch size. The mechanism, the edge
cases (short final chunk, dropped windowed edges), and the ragged per-block tail are the
live contract: see [Read-once and sample-once](docs/architecture.md#read-once-and-sample-once).

## Sample geometry — how the ladder evolved

> The **live contract** is in [docs/architecture.md](docs/architecture.md) ("the
Expand Down Expand Up @@ -220,19 +228,21 @@ blocks (safest for time series); optional chunk-shuffle for exchangeable samples

## Shuffle (the interesting compromise)

Global shuffle ⊥ chunk-aligned reads. Two-level approximation, after MosaicML
Streaming's `py1e`/`py1br`:

1. **Chunk permutation** — shuffle the order chunks are scheduled per epoch,
keyed on `(seed, epoch)` only (canonical: hardware-independent, resumable).
2. **Shuffle-block buffer** — hold a window of `block_chunks` decoded chunks and
draw batches across the window.

`block_chunks ≳ 10×` samples-per-chunk ≈ global quality; `block_chunks` is the
single **quality ↔ memory** knob. `shuffle_quality()` scores a draw order so the
knob can be tuned empirically. `shuffle=False` (eval / inference / reconstruction)
swaps in `sequential_order` — chunks and samples in order, no permutation. Both
order functions size a short final chunk correctly (no out-of-range draws).
Global shuffle ⊥ chunk-aligned reads: an *exact* global shuffle wants a random chunk
per sample, which defeats the whole point of reading a chunk once and gathering every
sample in it. The compromise, after MosaicML Streaming's `py1e`/`py1br`, is two-level —
permute *which* chunks are scheduled each epoch (keyed on `(seed, epoch)`:
hardware-independent, resumable), then shuffle all samples within a window of
`block_chunks` chunks. Over epochs the permutation re-pairs chunks, so it converges
toward a global shuffle at `O(block_chunks)` memory, not `O(dataset)`.

That is the road taken; the **mechanism, the convergence argument, and the
archive → batch figure are the live contract** in
[docs/architecture.md](docs/architecture.md#why-the-block-local-shuffle-is-enough). The
bet it encodes — approximate-but-in-place over exact-but-resharded — is the same one the
[annbatch contrast](#what-it-is-by-contrast) turns on: they pre-shuffle to disk and
dissolve chunks into scattered rows; we keep chunks whole and resident and pay with an
approximate shuffle instead.

## Memory model

Expand Down Expand Up @@ -390,6 +400,30 @@ The shape above wasn't the first cut. The pivots that got here, and the roads no
Things wrong or missing in *our* code today, with the reasoning that sets their priority.
(Things simply *not built yet* are milestones — see Roadmap.)

- **Per-block ragged batches (possible wart).** Batches never span shuffle blocks, so a
block's sample count not dividing `batch_size` yields a short final batch *per block*, not
one per epoch — behaviour and step-count impact in
[Read-once and sample-once](docs/architecture.md#read-once-and-sample-once). Correct
(sample-once holds), but it can surprise consumers assuming a uniform batch size. No
`drop_last` today. Options if users want uniformity: (a) opt-in `drop_last`; (b) carry a
block's remainder into the next block's first batch (extends that block's pin lifetime past
its own batches); (c) a whole-epoch re-batch (cheap — `order` is already materialized — but
mixes chunks across block boundaries, diluting the block-local residency guarantee).
**Priority pending user feedback** — no correctness impact, so we hold until someone hits it.

- **Batch buffers are fresh per batch, and host memory is not pinned.** `pool.gather`
allocates a new array per batch rather than reusing one — the DLPack export aliases the
buffer into the consumer's tensor, so reuse needs lifetime tracking against the
`prefetch_depth` window — and nothing pins host memory for GPU transfer, so H2D copies
are pageable. The unified fix is a small **ring of pre-pinned buffers** (depth
≈ `prefetch_depth`+1) handed out round-robin (a buffer is safe to reuse once its batch
has left the prefetch window): it removes the per-batch allocation *and* enables
`non_blocking` H2D transfer in one change — pinning fresh per batch would be worse
(`cudaHostAlloc` is expensive). Value scales with payload: negligible for small weather
fields (µs transfers, already hidden by prefetch), real for ViT/microscopy batches — so
**profile the H2D ceiling first** (pinned vs pageable bandwidth; is the copy on the
critical path or already overlapped?). Ties into the GPU `device_transform` stage (M2).

- **`window_factor` sizes residency by span, not by the offset set.** `source.py` sizes the
shuffle/residency working set with `window_factor = 2 + ceil(span/spc)` where
`span = max(offset) - min(offset)`. For a **sparse** offset set (leads `{24h, 240h}`:
Expand Down
104 changes: 102 additions & 2 deletions docs/architecture.md
Original file line number Diff line number Diff line change
Expand Up @@ -121,6 +121,15 @@ and amortized across every sample that touches it; **read concurrency
total memory is the budget + the prefetch queue (depth `d`) + the in-flight tiles —
every term a tunable cap, none scaling with batch size or epoch length.

The batch buffers themselves are **not** pooled: `gather` allocates a fresh array per
batch (the DLPack export aliases it into the consumer's tensor, so safe reuse would
need lifetime tracking against the prefetch window), and host memory is **not pinned**
yet, so host→device copies are pageable. Both are one deferred fix — a small ring of
pre-pinned buffers (depth ≈ `prefetch_depth`) reused round-robin, which kills the
per-batch allocation *and* enables `non_blocking` transfer together. See Known
limitations in
[DESIGN.md](https://github.com/emfdavid/insitubatch/blob/main/DESIGN.md#known-limitations--defects).

### Event-loop ownership (fsspec backends)

A genuinely-async fsspec backend (gcsfs, s3fs) binds its aiohttp session to the *first
Expand Down Expand Up @@ -354,7 +363,95 @@ a full shuffle in practice:
So even a modest window asymptotes quickly toward a global shuffle over the many epochs
training actually runs — at memory cost `O(block_chunks)`, not `O(dataset)`.
[`shuffle_quality`](api.md) scores an emitted order against a perfect global shuffle if you
want to see it on your own data.
want to see it on your own data. (`shuffle=False` swaps in a sequential order — chunks and
samples in order — for eval / inference / reconstruction.)

### Read-once and sample-once

The pipeline holds two guarantees, and they are **orthogonal** — batch size touches neither:

- **read-once** — a stored tile is fetched and decoded exactly once, however many samples,
batches, or epochs reference it (the read plan dedups; the `ChunkPool` keeps it resident;
gather reads from the slot).
- **sample-once** — each valid sample lands in exactly one batch.

`order` is the ledger for sample-once: an `(N, 2)` array of `[chunk_id, within]`, one row per
drawn sample. It is a permutation of every valid anchor, so each sample appears once. The
**sample-level fancy index is never stored** — `gather` recomputes it per batch from the rows
(`anchor = chunk_id·spc + within`; `sample = anchor + offset`; scatter `slot.data[within[mask]]`
per unique read chunk), Python work `O(chunks-in-batch)`, never `O(samples)`.

The figure traces one epoch, archive → batch. Colour tracks the *bytes* — a chunk keeps its
colour from the archive ① through the resident pool ③; the **read plan** ② in between is
metadata (tile keys in shuffle-block order, *not* bytes), so the plan (dashed) stays distinct
from the movement (thick). In stage ④ a batch is a row of colour-coded slices drawn across
several resident chunks: one chunk feeds many batches yet is **never broken up** — it stays
whole and reused, which is what read-once buys. (Loaders that pre-shuffle to disk instead
dissolve the chunk into scattered rows; see the annbatch contrast in
[DESIGN.md](https://github.com/emfdavid/insitubatch/blob/main/DESIGN.md#what-it-is-by-contrast).)

```mermaid
flowchart TB
subgraph S1["① immutable archive — chunks along the sample axis"]
direction LR
A0["chunk 0"]:::cA
A1["chunk 1"]:::cB
A2["chunk 2"]:::cC
A3["chunk 3"]:::cD
end
subgraph S2["② read plan — deduped, shuffle-block order"]
direction LR
P["what to fetch, in what order · read 2 · 0 · 3 · 1 · … · <i>tile keys — not bytes</i>"]:::plan
end
subgraph S3["③ ChunkPool — block_chunks tiles resident &amp; intact"]
direction LR
R2["chunk 2"]:::cC
R0["chunk 0"]:::cA
R3["chunk 3"]:::cD
end
subgraph S4["④ windowed anchor draw — each batch interleaves slices from several resident chunks (which stay whole, reused)"]
direction TB
subgraph B1["batch #1"]
direction LR
b1a["c2"]:::cC
b1b["c0"]:::cA
b1c["c3"]:::cD
end
subgraph B2["batch #2"]
direction LR
b2a["c0"]:::cA
b2b["c3"]:::cD
b2c["c2"]:::cC
end
Bn["… batch #n"]:::batch
end
S1 -.->|"plan the reads<br/>dedup + draw order"| S2
S2 ==>|"scheduler fetches + decodes<br/>one per tile — the movement"| S3
S3 -->|"gather by (chunk_id, within)<br/>draw rows"| S4
classDef cA fill:#cfe0f7,stroke:#5b7aa6;
classDef cB fill:#f7cfd8,stroke:#a65b6e;
classDef cC fill:#f7e0cf,stroke:#a6805b;
classDef cD fill:#d3f0d8,stroke:#5ba66e;
classDef plan fill:#ffffff,stroke:#666,stroke-dasharray:5 3,color:#333;
classDef batch fill:#ededed,stroke:#888;
```

**"Exactly once" means every *valid* anchor.** `order` is built to handle the edges: a **short
final chunk** emits `within` only up to its real length, and windowed sampling drops anchors
whose `anchor+offset` would read off the array (incomplete windows — correct to skip for
training; the inference path validates the range and raises instead). Guaranteed by
`test_order_covers_every_sample_exactly_once`, `test_order_handles_partial_final_chunk`, and
the decode-once suite.

**The tail is ragged per block, not per epoch.** Batches do not span shuffle blocks: the
producer batches `order` within each block's row range and restarts at the boundary. So when a
block's sample count (`≈ block_chunks × spc`, minus edges) is not a multiple of `batch_size` —
the common case — the **last batch of every block is short**. No sample is lost or duplicated
(sample-once holds), but an epoch yields *several* short batches, and steps-per-epoch is
`Σ ⌈block_samples / bs⌉`, not `⌈N / bs⌉` — worth knowing for BatchNorm on a small tail or
step-count math. There is deliberately no `drop_last` today; whether to add one is an open
question (see Known limitations in
[DESIGN.md](https://github.com/emfdavid/insitubatch/blob/main/DESIGN.md#known-limitations--defects)).

## Transforms — three stages, placed by cost

Expand Down Expand Up @@ -785,7 +882,10 @@ pretending to be a general compute graph.
non-sample axis; there is no spatial patching/cropping in the read yet, so a giant
field is assembled whole (crop it in a `batch_transform`). Native patch geometry —
and memory-optimal partial-field residency — is a reserved extension (see the
axis-role contract).
axis-role contract). The shipped random-crop example (`examples/wb2_dataloader.py`)
does exactly this — crops in a `batch_transform` over the whole-field cache, so the
crop re-randomizes each epoch for free — but its crop is a per-sample Python loop,
illustrative rather than vectorized.
- **Not a compute framework.** No general task graph, no cross-chunk reductions on
the hot path, no lazy dask-style evaluation — by design (dask on the hot path is
the thing we route around). Reductions like fitting a scaler run *over the loader*
Expand Down