API reference¶
The public surface is the top-level insitubatch package — everything in its
__all__. InSituDataset and the framework adapters are re-exported there, so
import them from the package root (from insitubatch import InSituDataset, to_torch),
not from submodules. The adapters are optional: they import torch / JAX / TF lazily,
only when called, so importing insitubatch never pulls a framework in.
insitubatch¶
insitubatch -- train in place on n-dimensional cloud tensors.
The loader-orchestration layer that sits on top of already-solved async cloud IO (obstore / zarr v3 / icechunk): turns an existing Zarr archive into a shuffled, split-aware, GPU-saturating PyTorch source with no reshard and a Python hot path that scales with chunks, not samples.
See DESIGN.md for the full rationale.
ChunkPool ¶
Byte-budgeted pool of outer-chunk slots, keyed (array, chunk_index).
The pool is the assembly buffer and the cache. A slot is pinned while the
current epoch needs it (in-flight or block-not-yet-drained) and unpinned once
its block is drained;
unpinned slots stay resident (retained for cross-epoch reuse) until budget
pressure evicts them in LRU order. budget_bytes is the single knob:
- small (~
2*block_chunksworth) -> read-once (unpinned evicted promptly); - large + persistent across epochs -> a decode-once cache (a still-resident prepped chunk is a hit, skipping fetch + decode + transform).
Eviction targets unpinned-LRU only; a slot is never unpinned before it is
ready+drained, so an in-flight or in-use chunk is never dropped. Backing is heap
or mmap (see backing_dir); chunk_transforms run once per outer chunk on
the assembled array, so a hit reflects decode + transform.
One writer per backing_dir (#42). Whenever a backing dir is set -- with or
without persist, since _alloc writes the same filenames either way -- the pool
takes an advisory lock on it for its lifetime and a second writer fails fast
(:meth:_lock). readonly_cache=True takes that lock shared instead: many such
openers coexist with each other, none with a writer, none of them writes anything, and
a miss raises (:meth:_readonly_miss) rather than fetching -- the flag is an
assertion that this cache is complete for what the run reads. Slot files are replaced,
never truncated in place (:meth:_alloc), so a reader holding a mapping keeps reading
real data even while a writer re-admits the same chunk.
assembles
property
¶
True if publishing a chunk does real work (assembly / transform / write-back).
The scheduler needs this to decide where delivery runs. On the plain tiled path a
delivery is a dict write and a counter, so it belongs inline on the loop. When a
slot must publish a whole array, _advance also runs the assembly memcpy, the
user chunk_transform and the mmap write-back -- none of which may sit on an
event loop we share with the rest of the process.
active_owners
property
¶
How many iterations are live on this pool right now.
Live means minted and not yet released -- not "currently holds a pin". An iteration that starves before it can pin anything is precisely the case the starvation diagnostic exists to name, and counting pin-holders would miss it.
More than one means several iterations share this pool (zip(ds.train,
ds.val), two DataLoaders) and each needs its own working set resident at the
same time -- the most common reason a budget auto-sized for one cannot admit.
Read-only snapshot.
new_owner ¶
Mint an opaque token identifying one iteration's references.
One iteration = one owner: its scheduler's admission pins and its producer's
block pins are the same owner, so unpin_all(owner) at the next epoch's
prologue releases exactly that iteration's state and leaves a concurrent
iteration's untouched. Deliberately not user-visible -- _iterate mints one
and threads it through; nothing in the public API names an owner.
counters ¶
This owner's counters. Read before its :meth:release_owner, which drops them.
blocked_waiters ¶
(path, chunk_index) keys a thread is blocked on in :meth:wait_ready and
cannot yet proceed, across every owner or within owner alone.
A waiter whose condition already holds is excluded, even though its thread is
still parked: it is registered until the OS reschedules it, and reporting it as
blocked asserts something false about the future. The scheduler pairs this with
"no tile in flight" to prove a stall terminal (:meth:Scheduler._starvation), and
that proof is only as sound as this list -- a waiter that is about to wake up,
gather and unpin is precisely the thing that will free budget. Read-only snapshot.
owner narrows it to one pass, and which of the two a caller wants follows from
what it is waiting on. The byte budget is shared, so an admission stall is rightly
judged against every owner: another iteration's blocked consumer is holding budget
we need. A read-ahead permit is not shared -- it belongs to one
:class:~insitubatch.scheduler.Scheduler and comes back only from that pass's own
unpin_block -- so a permit stall must ask only about its own waiters
(:meth:Scheduler._ahead_starvation).
try_admit ¶
Reserve + allocate + reference one outer-chunk slot, evicting ready-LRU for room.
Admission takes one reference (incref) and claims the slot for this epoch (so
the consumer's :meth:wait_ready won't gather it until the driver has referenced
it -- see there), so the slot stays resident from its in-flight fetch through to
the consumer's release. The driver fetches each chunk once per epoch, so an
eviction before consume could not be re-fetched and would deadlock the waiter.
The consumer releases each chunk at its last use (:meth:unpin_keys); windowed
reads let one chunk be referenced by several blocks, hence reference counts, not a
boolean. Idempotent if already resident (incref only). Returns False only when
the budget is full of in-flight or referenced slots -- the caller awaits a release.
is_ready ¶
True if the chunk is resident, fully assembled, and not failed (a hit).
pin_if_ready ¶
Incref + return True iff the chunk is resident, ready, and not failed.
A cross-epoch (or, with persist, cross-run) cache hit the driver can skip
fetching -- but it must still be referenced so it stays resident through the
consumer's use (released at last use like an admitted chunk), else it could be
evicted before the waiter gathers it and, since the driver fetches each chunk
once, deadlock. One lock so the check and the incref cannot race an eviction in
between. A persisted-on-disk chunk is revived here on first touch (see
:meth:_revive), so a cross-run hit costs no fetch.
pin_keys ¶
Reference (incref) a set of (path, chunk_index) slots for a live block.
Windows make one chunk readable by several concurrent blocks, so pins are
reference-counted: each block that needs a slot increfs it on entry and
decrefs it on drain (:meth:unpin_keys). A slot with refcount > 0 is never
evicted. Pinning a not-yet-allocated key is fine -- the count is recorded and
the slot, once admitted, inherits it.
unpin_keys ¶
Release (decref) a block's (path, chunk_index) references.
A slot dropping to refcount 0 becomes LRU-evictable (retained for cross-epoch reuse until budget pressure drops it), not dropped now. Wakes any admit parked on a full budget.
buffer_stats ¶
One consistent reading of the batch-output pool, for the per-epoch summary.
set_host_allocator ¶
Point the batch-output pool at a different host allocator.
How the torch adapter installs page-locked buffers without the core importing a
framework. A method rather than letting callers reach pool._buffers directly, so
the buffer pool stays this class's private business and there is one place to look for
who can change it.
release_owner ¶
Release everything one iteration holds: its pins, and its abandoned partials.
Called by that iteration's own teardown, so the state dies with the thing that created it. A pin is per-pass working state, not cache membership -- READY chunks stay resident (unpinned) for cross-epoch reuse.
This must happen: an abandoned pass (early break) otherwise leaves its
read-ahead and un-drained block referenced forever, shrinking every later
epoch's budget until admission can free no room and the driver deadlocks.
Scoped to owner, because one pool serves several concurrent iterations.
Releasing globally is #34 -- it strips a live iteration's pins and its in-use
chunks become eviction candidates mid-gather. For the same reason a partial is
dropped only when it is quiescent and unreferenced: a slot still being written
by a live task, or held by another owner, is not ours to reclaim.
tile_write ¶
Scope one tile task's write, guaranteeing the slot hears about its end.
The pool owns the guarantee; the caller only has to be inside the scope::
with pool.tile_write(array, cid, coord) as w:
tile = await fetch_decode(...) # cancellation here is covered
w.deliver(tile) # or w.fail(exc)
A bare decrement at the end of the task would not do: a tile task can end
without ever reaching the pool -- a fetch/decode error takes an early return,
and a cancellation (an early break closes the scheduler with
cancel_futures=True) can unwind at any await. writers is what
eviction reads, so a missed decrement is a slot that is never safe to take away
and a budget that never recovers. __exit__ is the only construct that
cannot be forgotten by a future early return.
deliver / fail / scope exit all funnel into one release, and the
release is idempotent -- so a delivered tile does not double-decrement.
deliver_tile ¶
Synchronous convenience: scope one tile write and deliver it in one call.
Sugar over :meth:tile_write, not a second lever -- it opens the same scope, so
the writer count still moves in exactly one place. Use it when there is nothing
to await between reserving the write and having the tile. The scheduler cannot:
it awaits the fetch inside the scope, so it uses :meth:tile_write directly.
fail ¶
Poison one chunk so a waiting consumer re-raises instead of hanging.
Fail-fast: a fetch/decode error on any tile poisons its outer chunk; the
consumer's wait_ready surfaces it on the main thread.
It records the error and nothing else. The old version also set
ready = True -- because "stop waiting" and "safe to evict" shared one flag,
the only way to wake a waiter was to declare the slot a finished cache entry,
while sibling tile tasks were still writing into it (#33). wait_ready keys
off error independently of state, so waiters still wake immediately; the
slot is disposed of by :meth:_advance once it quiesces.
set_error ¶
Poison the whole pool (the fetch driver died) so every waiter re-raises.
Unlike :meth:fail (one chunk), this unblocks consumers waiting on chunks
that may never be allocated -- the driver failed before reaching them. The
first error wins (later failures are usually cascade noise).
wait_ready ¶
Block until the chunk is READY for this owner (or raise).
Requiring this owner's own reference (rather than a shared claimed
bool) closes a cross-epoch race: a chunk still resident-and-READY from the
prior epoch would otherwise be gathered before the driver references it,
letting the consumer's last-use release land before the driver's pin -- a lost
release that leaks a reference and, worse, lets the driver evict a chunk
mid-gather. Owner-scoped because a single bool is satisfied by another
iteration's claim (#35), which is the same race with a concurrent producer
instead of a previous epoch.
Wakes on: READY and referenced by owner; the chunk failed (:meth:fail,
which no longer has to fake readiness to get here); or the pool was poisoned
(:meth:set_error, covering a driver death before this chunk was allocated).
A key that is simply absent means the driver has not admitted it yet -- keep
waiting, as before.
gather ¶
Assemble one batch from [chunk_id, within] anchor draw rows.
Each row is one sample anchor t = chunk_id*ref_spc + within in the reference
(manifest) grid; each variable reads its array at t + offset (offset 0 is the
plain non-windowed case). Output is in anchor-row order: row i of every variable
is the same anchor, sample_indices[i] == t_i. Per variable the reads are grouped
by the variable's own (offset-shifted) chunk -- computed with that variable's
chunk size, so variables may chunk the sample axis differently -- one coalesced
fancy-index per chunk, never a Python per-sample loop. The caller must have waited
every referenced (path, offset-shifted chunk) ready.
close ¶
Free every remaining slot; release the log handle and the cache lock. Persist keeps ready cache files (each already recorded in the log at completion, so there is nothing to rewrite -- just flush + close the handle); heap/spill mmap files are unlinked. Idempotent.
Depths
dataclass
¶
PassStats
dataclass
¶
One pass over one split: counters, depths, times, and the stage they accuse.
consumer_idle
property
¶
True when the caller's loop body did essentially nothing.
A for batch in ds.train: pass loop takes batches as fast as they are produced, so
the queue is empty on every sample no matter how fast the loader is. That is a
throughput measurement, not a starved pipeline, and saying "the queue ran empty" of
it would be true and useless.
StageTimes
dataclass
¶
Cumulative seconds by stage. See the module docstring for the units rule.
producer_s
property
¶
Accounted producer-side cost, the denominator for a stage's share.
Excludes consumer_s, which is the caller's own work and not a stage we can tune.
Scheduler ¶
Owns one event loop + a decode pool; streams tiles into a caller-owned pool.
The :class:ChunkPool is passed in (dataset-owned, so it persists across epochs
as the cache). :meth:start streams the stored chunks of an ordered chunk list;
the consumer reads assembled chunks via :attr:pool and releases drained ones
via :meth:unpin. Per chunk the scheduler skips fetch if the pool already holds
it (a cross-epoch hit); misses are admitted against the pool's byte budget,
awaiting an unpin when the working set fills it.
decode_threads
property
¶
Effective decode-pool width; see :func:decode_pool_workers.
owner
property
¶
This scheduler's reference token, for the consumer's pool.wait_ready.
One iteration = one owner: whoever drains this scheduler must wait under the
same token the scheduler admits under, or its wait can never be satisfied.
InSituDataset._iterate mints the token and hands it to both sides; a
caller driving a Scheduler directly reads it here.
close ¶
Cancel this scheduler's in-flight driver. Nothing else.
Graceful: a consumer may close mid-epoch (early break) while _drive is
still streaming, so we cancel our outstanding tasks and let them unwind.
A scheduler owns no shared resource, so it tears down no shared resource. The
loop is zarr's and lives for the process; the decode pool is process-wide and
outlives every scheduler; the chunk pool is caller-owned and persists across epochs
as the cache. Closing any of them here would break unrelated work elsewhere in the
process -- which is not hypothetical: stopping the loop hangs every later
zarr.core.sync.sync() call, and shutting the decode pool down makes the next
scheduler raise cannot schedule new futures after shutdown. Both were observed.
start ¶
Begin streaming the stored chunks of chunk_ids (priority order).
chunk_ids are in the reference (manifest) grid; ref_spc is that grid's
sample-chunk size, used to map anchor chunks onto each variable's own chunks.
Returns the driver future; a failure there poisons the pool so consumers
re-raise. The consumer drives demand independently via :attr:pool.
unpin_block ¶
Release references on a set of drained (path, chunk_index) slots
(thread-safe): the slots that hit refcount 0 become LRU-evictable; wake any
admit parked on a full budget so it can evict them and proceed.
SchedulerConfig
dataclass
¶
max_inflight
class-attribute
instance-attribute
¶
Tiles in flight at once -- the single concurrency dial. Memory in flight ~= max_inflight * stored_chunk_nbytes (+ transform scratch). Residency is bounded separately by the pool's byte budget (admission evicts unpinned-LRU).
read_ahead_chunks
class-attribute
instance-attribute
¶
How many chunks this iteration may hold ahead of its own consumer.
0 leaves admission bounded only by the pool's byte budget, so one driver claims
the whole pool however large it is -- which starves any other iteration sharing it
(#64) and, on a generous budget, fetches the resident set before the consumer's first
batch. :class:~insitubatch.source.InSituDataset always supplies a value, computed by
:func:~insitubatch.source.read_ahead_bound; 0 is what a caller driving a
:class:Scheduler directly gets.
It must be at least that computed bound. Below it the consumer waits for a chunk the driver is not permitted to admit -- a deadlock of our own making rather than the budget's -- which is why a non-zero value here is a testing seam, not a tuning knob.
decode_threads
class-attribute
instance-attribute
¶
Size of the decode pool (the GIL-releasing codec decode runs here). 0 = auto =
min(32, cpu+4).
The pool is process-wide (see :func:decode_pool), so this sizes it once, on the
first dataset built in the process; a later, different value is ignored with a warning.
Thread count is a property of the machine rather than of the data, so one number per
process is the honest shape.
on_bad_chunk
class-attribute
instance-attribute
¶
What to do when a stored chunk fails to fetch/decode (truncated/corrupt --
common in GRIB-under-zarr archives like HRRR). "raise" (default) fails fast;
"nan" fills that tile with NaN (float dtypes) or the fill value, so the chunk
assembles with a hole instead of poisoning the epoch -- the caller then handles
NaN with a chunk_transform (interpolate / drop). Bad reads are recorded in
Scheduler.bad_chunks.
DrawOrder
dataclass
¶
One epoch's [chunk_id, within] draw rows, and the shuffle-blocks holding them.
The blocks are laid down by the builder that creates them and travel with the rows, so every consumer reads one account of where a block begins and ends. Dropping rows narrows the block that held them and leaves its neighbours' membership untouched.
bounds holds n_blocks + 1 row offsets, so block i is
rows[bounds[i] : bounds[i + 1]]. Blocks are contiguous, gapless, and cover every
row.
keep ¶
Drop the rows mask excludes, narrowing the blocks that held them.
A block emptied outright is dropped: it names no chunks and does no work. Every other block keeps its membership, its row range shifted by the number of rows removed ahead of it.
InSituDataset ¶
A framework-neutral source of shuffled numpy batches from Zarr, split-aware.
The dataset is not itself iterated -- you iterate one of its split views:
:attr:train (shuffled), :attr:val, :attr:test, :attr:all (deterministic).
All four share one :class:ChunkPool, so a chunk that two splits both read --
e.g. a windowed read spilling across a split boundary -- is decoded once::
ds = InSituDataset(store, manifest, geometries=geoms, batch_size=32)
for batch in ds.train: ... # one epoch; ds.set_epoch(e) reshuffles
for batch in ds.val: ...
One epoch over a view = permute the split's chunks -> walk shuffle-blocks, stream-fetching
each block's stored chunks into the pool -> gather coalesced batches cut over the whole
epoch (so only the last is short, and one may span a block boundary) -> evict.
Batches are numpy :class:Batch; convert to a framework with
:mod:insitubatch.frameworks (as_torch / to_jax / as_tf_dataset). A
different per-split configuration (e.g. train-only augmentation) is a separate dataset.
Two preprocessing hooks, placed by cost (full model in the docs, "Transforms"):
chunk_transforms--(DecodedChunk) -> DecodedChunk, run per chunk before shuffle, seeing one variable. The cacheable home for elementwise, per-variable, deterministic work (scaling, unit conversion, dtype cast); amortized over every sample in the chunk and reused across epochs. Restrict one to particular variables with :func:~insitubatch.transforms.applies--applies(["2m_temperature"], k_to_c)-- rather than testingchunk.read.arrayinside the transform: a name test in the body is invisible to the engine, which would fold a reshaping transform's declared output into every variable's geometry and truncate the ones it does not touch.batch_transforms--(Batch) -> Batch, run per assembled batch, seeing all variables aligned on the sample axis. For cross-variable derived fields and per-sample random augmentation; runs after the cache, so it is never cached.
Runnable side-by-side example: examples/transforms.py.
One writer per cache_dir. The cache directory (and persist=True on top of
it) is arbitrated across processes by an advisory lock held for the dataset's
lifetime: one writer at a time, and a second fails fast rather than silently
corrupting the first's chunks. Several jobs may read one warm cache at once with
readonly_cache=True, which takes the lock shared, never writes, and raises on a
miss -- it asserts the cache is complete for what this run reads, rather than
quietly falling back to fetching. Put cache_dir on local NVMe: over a network
filesystem the lock may be emulated per client and we can only warn. See
docs/tuning.md, "Sharing a cache_dir between processes".
all
property
¶
Iterable over every split's chunks (deterministic) -- e.g. full-archive inference.
set_epoch ¶
Call from the training loop so each epoch reshuffles deterministically.
describe ¶
What this dataset will do, from geometry and configuration -- no store access.
iterations is how many passes will share the pool at once (zip(ds.train,
ds.val) is two), since each holds its own chunk references and the automatic
budget covers one. See :meth:print_summary for the formatted view, and
:mod:insitubatch.summary for what each number means.
print_summary ¶
Print :meth:describe for a human. On demand -- construction stays quiet.
close ¶
Release the cache pool's backing (mmap handles, cached chunks) and any async store session.
The pool persists across epochs, so close it when done training -- not per
epoch. With persist=True the cache files + manifest are kept on disk for a
future run (only the in-memory handles are released); otherwise the mmap spill
files are unlinked. An fsspec/gcsfs store's aiohttp session is closed on its own
loop here (a no-op for obstore) so it does not leak or spew a teardown traceback
at GC; gcsfs recreates it lazily if the store is reused. Idempotent; also called
on GC.
SplitManifest
dataclass
¶
Which sample-axis chunk indices belong to each split.
sample_indices ¶
Expand a split's chunks into the global sample indices they contain.
DatasetReport ¶
Bases: TypedDict
The whole static picture. Returned by :meth:InSituDataset.describe.
BatchTransform ¶
Bases: Protocol
Per-batch transform applied after gather (not cached).
ChunkTransform ¶
Bases: Protocol
Per-chunk transform applied before shuffle/gather (cacheable).
StandardScaler
dataclass
¶
Global per-variable standardization with PRE-FIT, FIXED statistics.
mean/std are keyed by variable and shaped to broadcast over a chunk's
(n_samples, *inner) array WITHOUT the sample axis: a surface variable uses
shape (1, 1); per-level stats use (level, 1, 1). The same stats are
applied to every chunk of that variable -- never recomputed per chunk.
Pre-fit the stats however you like and pass them in. The recommended path is to
fit over the loader with sklearn's incremental StandardScaler.partial_fit
(which also warms the cache) and scale at the batch stage -- see
examples/fit_scaler.py; this class is the chunk-stage applier for when you
want the normalization cached with the decoded chunk.
Give a re-fitted scaler a new cache_key if you persist the cache without
cloudpickle. The fallback fingerprint hashes this class's source plus its repr, and
numpy summarizes an array over 1000 elements in repr -- so per-gridpoint statistics
(say (721, 1440)) repr identically whatever their values, and a re-fit reopens a
persisted cache as a hit, serving chunks normalized with the old numbers. Any of three
closes the hole: install insitubatch[cache] (cloudpickle hashes the values), pass
cache_key="<stats version>" and bump it on every fit, or keep stats small enough to
repr in full.
cache_key
class-attribute
instance-attribute
¶
The identity this scaler declares to the cache fingerprint, and the only one that is
exact: the fallback hashes repr, which numpy summarizes for large statistics. None
(the default) leaves identity to the fingerprint's own resolution.
validate_scope ¶
Reject a scope this scaler has no statistics for (called by :func:applies).
applies(["u10"], scaler_fitted_on_t2m) would otherwise KeyError inside a
decode thread, one chunk into training. The stats dict may be wider than the
scope -- one fitted scaler shared by several scoped uses is normal.
ArrayGeometry
dataclass
¶
The minimal geometry the engine needs about one zarr array.
We only model the sample axis explicitly, because that is the axis we split,
shuffle, and batch along; the remaining dims are carried opaquely as
inner_shape and kept contiguous to preserve partial zero-copy. shape and
chunks are in physical (zarr) axis order -- they mirror the array's own
metadata -- and sample_axis names which physical axis is the sample axis
(0 by convention: time for ERA5/HRRR; e.g. 2 for the Z of an OME-NGFF
(T,C,Z,Y,X) microscopy stack sampled slice-by-slice). The engine works in a
logical view where the sample axis leads and the inner axes follow in physical
order; the one physical<->logical permutation is confined to the scheduler
(:meth:physical_chunk_coord for read addressing; a moveaxis on the decoded
tile). Everything downstream -- planning, pooling, gather -- is sample-first.
offset makes a variable a windowed view: it reads array[anchor + offset]
along the sample axis around a shared anchor. Two geometries with the same path
and different offset (e.g. g and g.shift(1)) are two views of one array --
they decode once and share slots. Offset 0 is not special; everything is relative to
the anchor.
sample_chunk_size
property
¶
How many samples live in one chunk along the sample axis.
inner_shape
property
¶
Shape of a single sample (every axis but the sample axis, physical order).
inner_chunks
property
¶
Stored-chunk shape on the inner (non-sample) axes (physical order).
chunk_bytes
property
¶
Bytes one sample-axis chunk occupies once resident — the number to size a budget against, and the one that is easy to get wrong by hand.
Residency is the array's stored tiles, kept whole: a chunk grid that does not
divide the array evenly still stores full-size edge chunks (721 rows chunked at 180
occupy 900) and the pool does not clip them. So this is
n_inner_chunks * prod(tile_shape) * itemsize, which on a gridded inner axis is
larger than sample_chunk_size * prod(inner_shape) * itemsize — the product a
reader reaches for first, and which under-provisions exactly the deeply chunked
stores where the difference costs gigabytes.
This is the tiles term of :func:~insitubatch.pool.slot_charge_bytes, which charges
it; the two are the same expression, so they cannot drift. A configured
chunk_transform can make the assembled output the binding term instead — ask
:meth:InSituDataset.describe for the whole picture rather than summing this by
hand across variables.
shift ¶
A view of the same array read k samples later (composes: shift(1).shift(1)
is offset += 2). Declare a forecast target as g.shift(horizon).
physical_chunk_coord ¶
Full physical zarr stored-chunk coordinate for a logical read.
Reinsert the sample-axis chunk index chunk_index at sample_axis among the
inner-axis chunk coords (which are in physical inner order). With sample_axis == 0
this is exactly (chunk_index, *inner_coord) -- the identity the old code assumed.
samples_in_chunk ¶
The half-open range of global sample indices in chunk_index.
n_inner_chunks ¶
How many stored tiles compose one outer chunk (the inner-grid size).
Independent of chunk_index (a short final outer chunk is still one
axis-0 stored chunk), but kept index-keyed so the pool's completion count
reads naturally and the API survives a future per-axis sample chunking.
inner_index ¶
Row-major position of one inner stored-chunk coord in :meth:inner_coords.
The inverse of iterating inner_coords(). It is what gives the persisted cache a
stable tile order: a chunk's .npy stores its tiles tile-major at this index,
so a revived file maps back to coordinates without recording them.
slot_shape ¶
Shape of the assembled outer chunk: (n_samples_in_chunk, *inner_shape).
Axis 0 uses the actual sample count so the final short chunk is sized exactly (no over-allocation, no out-of-range scatter).
tile_shape ¶
One stored chunk's shape, sample-first -- the shape a decoded tile arrives in.
The scheduler moves the sample axis to the front on decode, so this is
(sample_chunk_size, *inner_chunks) rather than the array's physical chunk order.
Full size always: an edge chunk is stored whole and its padding is kept.
tile_placement ¶
Where one stored tile lands inside its outer chunk, as zarr's own projection.
Returns a :class:zarr.core.indexing.ChunkProjection -- the vocabulary zarr uses
for exactly this ("a mapping of items from chunk to output array"), rather than a
private (dst, src) pair. out_selection indexes the assembled outer chunk;
chunk_selection clips the full chunk-shaped decoded tile to the (possibly
partial) edge region, on axis 0 (short final outer chunk) and the inner edges
alike. is_complete_chunk says the tile is used whole -- i.e. it is not an edge
tile -- which is what tells a reader whether the stored chunk carries padding.
Both selections are sample-first, matching the slot and the sample-first tile the scheduler delivers (it moves the sample axis to the front on decode), not the array's physical axis order.
Batch
dataclass
¶
A model-ready batch.
arrays maps variable label -> stacked array of shape (batch, *inner).
sample_indices is the per-row anchor sample index t (provenance for
determinism / resumption). offsets maps each label to its sample-axis read
offset, so label v of row i was read from global sample sample_indices[i]
+ offsets[v]. A plain (non-windowed) batch has every offset 0; a forecast batch
pairs e.g. an input at offset 0 with a target at offset horizon.
The batch stays a flat {label: array} dict -- there is no lead/role axis in the
engine. Use :meth:stack to assemble a multi-step window into one array and
:meth:read_indices for a label's true provenance.
__len__ ¶
Rows in the batch -- how a caller implements drop_last for itself.
Every batch is full except the epoch's last, which is short whenever the split's
sample count does not divide batch_size. That is the ordinary data-loader
contract, and dropping it is the caller's choice, so the check has to be a one-liner
that does not require naming a variable::
for batch in ds.train:
if len(batch) < ds.batch_size:
continue
Counted off sample_indices rather than an array's leading axis: it is the anchor
row count the engine sets on every batch and the one :meth:read_indices already
builds on, so this adds a spelling rather than a second notion of how long a batch is.
read_indices ¶
Global sample index each row of label was read from: anchor + offset.
Provenance for a windowed view (e.g. to confirm a target leads its input by the intended horizon). Defaults the offset to 0 for a label without one recorded.
stack ¶
Stack several labels into one array along a new axis (default 1).
The obvious way to build a multi-step input window from a set of time-shifted
views, e.g. batch.stack(["t_m2", "t_m1", "t_0"]) -> (batch, 3, *inner).
Order follows labels; row i of every label shares anchor
sample_indices[i], so the stacked steps stay aligned. The caller chooses the
labels and their order -- the engine does not impose a window layout.
ChunkRead
dataclass
¶
A single chunk to fetch, addressed along the sample axis.
array names which zarr array (variable) this read belongs to; a training
sample that concatenates several variables produces one ChunkRead per
variable that must be co-scheduled.
DecodedChunk
dataclass
¶
A decoded, in-memory chunk, keyed by its read.
data has shape (n_samples_in_chunk, *inner_shape) -- the logical chunk, and
a real contiguous array, never a view over stored tiles.
In particular it carries no stored-chunk padding. A zarr chunk grid that does not
divide the array evenly still stores full-size chunks at the edges (721 rows chunked at
180 occupy 900), and the pool holds those tiles whole. Assembly clips every tile to its
in-bounds region before a transform sees it, so a transform never has to know the chunk
grid, mask an edge, or handle padding values. n_samples_in_chunk is likewise the
real count, so a short final chunk arrives short rather than padded.
The buffer holds a bounded number of these; memory overhead is O(in-flight chunks), independent of batch size.
StoredChunkRead
dataclass
¶
One stored chunk to fetch: a single tile of the chunk grid.
Reading a whole outer chunk per getitem lets zarr stitch the inner grid
under a second concurrency cap; fetching at stored-chunk granularity instead
-- (chunk_index, *inner_coord) -- lets a single max_inflight budget
span inner and outer reads, with no nested caps. chunk_index is the
sample-axis (outer) stored-chunk index; inner_coord is the stored-chunk
index on each inner axis (empty tuple when the inner dims are single-chunk --
the degenerate GRIB-per-timestep case).
Frozen + hashable so a plan can dedup tiles and key the in-flight set.
debug_info ¶
Return the environment report as an ordered label -> value mapping.
The structured half of :func:print_debug_info; useful if you want to attach the same
facts to a benchmark record rather than paste them into an issue.
print_debug_info ¶
Print the environment report. Paste the output into a bug or performance report.
as_tf_dataset ¶
Wrap a split view (e.g. ds.val) as a tf.data.Dataset via from_generator.
output_signature is inferred from the view's geometries: each variable is
(None, *inner) (None = the variable last-batch size) with the variable's dtype.
Both from_generator here and :func:to_tf copy into the TF runtime -- TF has no
reliable zero-copy path from insitu's buffers (its experimental DLPack mishandles buffer
ownership; see :func:to_tf). Call :func:to_tf on the raw stream when you want plain
dict[str, tf.Tensor] batches instead of a tf.data.Dataset.
as_torch ¶
Wrap a split view (e.g. ds.train) as a torch IterableDataset for DataLoader.
Each yielded item is a dict[str, torch.Tensor] (via :func:to_torch). Use
DataLoader(as_torch(ds.train), batch_size=None, num_workers=0).
Passing device moves batches here instead of in the training loop, and switches the
batch buffers to page-locked memory. The two are one option rather than two on
purpose: pinning is what makes an H2D copy genuinely asynchronous, so it is only safe
when whoever issues the copy also knows when it landed. A pin_memory=True flag the
caller could set without device would hand out pinned buffers under the caller's own
non_blocking copy and reintroduce a use-after-recycle we would then have to document
instead of prevent.
What it is worth runs between two measured bounds, and how compute-bound the loop is
decides where you land. In isolation (bench/probe_batch_buffers.py --arms overlap) a
pinned copy below ~37 MiB is hidden completely -- the pinned arm sits on the compute
floor, while a pageable one costs more than its own transfer time because it stages
through a driver bounce buffer and cannot overlap at all. Against a compute step that
already hides the copy, it is worth nothing. A real GPU training loop lands in between:
+1 to +2%, roughly flat over a 4x payload range (32 to 128 MiB per step) rather than
switching on above a threshold.
Expect no more precision than that range. An A/B/A bracket on an unchanged checkout drifts
~0.7% per block -- half the signal -- so the bound is +0.8% to +2.2% depending on how the
drift is modelled. See docs/benchmarks.md, which also explains why % of ceiling cannot
measure this at all.
pin_budget_bytes caps the page-locked total (default: an eighth of RAM, since the
kernel cannot reclaim it); past that buffers are pageable again, with one warning.
to_jax ¶
Convert a numpy Batch to a dict of jax.Array (DLPack), on JAX's default device.
Zero-copy only when the batch buffer sits on a 128-byte boundary: XLA:CPU requires that
alignment and silently falls back to a copy otherwise. numpy guarantees only 16, so with
today's per-batch np.empty roughly half of batches are copied (measured 20/40 —
bench/probe_batch_buffers.py --arms jax). Correct either way, just not free.
jnp.from_dlpack imports a host buffer, so its result is committed to cpu:0 --
and jax.jit does not move a committed array. Returning that directly meant a GPU run
trained on the CPU at full correctness with nothing raised to say so, so the batch is
placed on jax.devices()[0] here, which is where every other jax array-creation path
puts it. Pass device to choose another. On a CPU-only build the put is a
pass-through that still aliases the exported buffer, so it costs nothing there.
The transfer is asynchronous -- JAX documents it so -- but JAX keeps the source
referenced across it, so the pool's liveness poll cannot reclaim a buffer mid-DMA
(bench/probe_batch_buffers.py --arms all is what keeps that honest; see DESIGN.md).
Unlike :func:to_torch this adapter does not own the transfer: JAX cannot consume
pinned host memory (jax#22346) and exposes no event to gate recycling on.
to_tf ¶
Convert a numpy Batch to a dict of tf.Tensor (one CPU copy per variable).
Unlike torch/JAX -- whose array-accepting from_dlpack manages the exported buffer's
lifetime correctly -- TensorFlow only exposes the experimental from_dlpack(capsule),
which mishandles ownership of the exported numpy buffer: under the concurrent allocation
of the prefetch decode threads it double-frees that buffer and aborts the process
(SIGABRT, no message). convert_to_tensor copies into a TF-owned tensor instead, so TF
never touches insitu-managed memory; the batch is already an owned array, so this is a
single CPU copy. (torch stays zero-copy, JAX when aligned -- see :func:to_jax; TF's is a
DLPack limitation, not ours.)
to_torch ¶
Convert a numpy Batch to a dict of torch tensors (DLPack; zero-copy on CPU).
With device set, the copy is issued here rather than by the caller, and the batch's
source arrays are held until it lands. That is what makes pinned batch buffers safe to
recycle -- see :class:_InFlight. Without it the behaviour is unchanged: CPU tensors
aliasing the batch, and the caller's own .to(...) afterwards.
build_stored_chunk_reads ¶
Expand outer chunk ids into deduped stored-chunk reads, in priority order.
There is no gather map: the scheduler delivers tiles into per-outer-chunk slots
in a :class:~insitubatch.pool.ChunkPool, and batches are gathered straight
from those tiles by (chunk_id, within) draw rows -- the same
coordinates the shuffle order already produces. So the result is just what to
fetch, in what order; the scheduler keeps max_inflight tiles in flight
across the list.
chunk_ids are outer (sample-axis) chunk indices in draw/priority order
(e.g. the next shuffle-block's chunks first), so the soonest-needed tiles go
first. Each outer chunk expands to its inner grid; every variable contributes
its own grid (variables may chunk the inner dims differently). Order is
chunk -> variable -> inner so a whole outer chunk's tiles are scheduled
together (it can be assembled and drained promptly). Reads are keyed by the
array path (not the dict label), so several windowed views of one array
(same path, different offset) collapse to a single fetch -- decode-once.
Dedup also makes the function safe to call with repeated ids.
chunk_ids are anchor chunks in the reference grid (ref_spc = the
manifest's sample-chunk size, which defines the shuffle/split anchor grid). A windowed
variable reads array[anchor + offset], and a variable that chunks the sample axis
differently from the reference maps those anchor samples onto its own chunks -- so one
anchor chunk expands to the (offset-shifted) chunks each variable needs. With every
offset == 0 and a uniform chunk size this is exactly anchor chunk -> itself.
groups scopes that dedup to consecutive runs of chunk_ids -- the consumer's
shuffle-blocks -- rather than the whole pass, and the returned block_starts say
where each block's reads begin. A chunk that two blocks read is then admitted once per
block, which is what lets the consumer release it when its block drains and have it
re-admitted later; the pool serves that second admission from residency or from
cache_dir, and re-reads it only if neither holds it.
Per block rather than per anchor, deliberately: the consumer releases each of a block's keys once, so admitting a key twice inside one block would leave a reference nobody returns -- and a slot pinned forever is a starved pass, not a slow one.
The boundaries are returned rather than left to be recovered from the reads, because they cannot be: the driver decides admission once per run of tiles belonging to one chunk, and a chunk that ends one block and opens the next produces two adjacent runs that are indistinguishable from one. Merging them takes a single reference where the consumer will release two, and the first release unpins the chunk out from under the second block.
bottleneck ¶
Which stage limited this pass, and what to do about it.
A caller whose loop body does essentially nothing is answered first: with no training step there is no "the consumer is the constraint" to reach, whatever the queue looks like. Otherwise the rule reads the batch queue, because a full queue settles the question: every producer stage kept up and the consumer is the constraint, which is the state you want. Only once the queue is repeatedly empty is there a producer stage to accuse, and then the accusation comes from which one spent the time -- not from which one feels slow.
Returns ("unknown", ...) rather than guessing when the evidence does not separate
the candidates: no samples, or no stage clearly dominant. A confident wrong answer here
costs more than an admission, because it sends someone to tune the wrong knob.
block_shuffled_order ¶
Produce a shuffle-block-ordered list of [chunk_id, within] draws.
Chunks are permuted per epoch; within each window of block_chunks chunks
all samples are shuffled together. n_samples is the global sample-axis
length, used to size a short final chunk correctly.
Returns a :class:DrawOrder: rows of shape (N, 2) covering every sample in
chunk_ids, and the bounds of the blocks this function just laid down. The
caller needs those bounds and cannot recover them from rows -- see
:class:DrawOrder -- so they are returned rather than left to be inferred.
chunk_permutation ¶
Deterministically permute chunk ids for one epoch.
Determinism is keyed on (seed, epoch) only -- not on world size or worker count -- so a run is reproducible and resumable across hardware (the "canonical" property from MosaicML).
sequential_order ¶
In-order [chunk_id, within] draws (no permutation, no shuffle).
Used when shuffle=False (eval / inference / reconstruction): chunks in the
given order, samples in order within each. Honours a short final chunk.
block_chunks changes no row -- nothing is shuffled -- but an unshuffled pass still
streams and releases in blocks, so it is grouped the same way. One block definition
across both orders is what lets the rest of the engine stay ignorant of which it walks.
No chunks is an empty order of the same shape, not an error: a split can hold nothing
because it was asked to (fractions=(1.0, 0.0, 0.0)) or because a small store
rounded it to zero, and iterating it should yield nothing the way an empty list does.
shuffle_quality ¶
A 0..1 score for how well an emitted order mixes the source.
Heuristic: the mean absolute source-rank gap between consecutive emitted
samples, normalised by the gap a perfect global shuffle would give. 1.0 ~=
global; values near 0 mean adjacent samples still come out near each other
(poor mixing). Cheap to compute, good enough to tune block_chunks.
split_by_chunk ¶
Partition a variable's sample-axis chunks into train/val/test.
Parameters¶
fractions:
(train, val, test) fractions of chunks (not samples). Must sum to ~1.
contiguous:
If True (default), assign contiguous blocks of chunks to each split --
the safest choice for time series, where a randomly interleaved split
still risks leakage through autocorrelation across chunk boundaries. If
False, chunks are shuffled before partitioning (acceptable when samples
are exchangeable, e.g. independent scenes).
sample_range:
Optional half-open (start, stop) window of sample (outer-axis) indices
to restrict the split to before partitioning -- e.g. train on one date
range of a long archive. The selection is chunk-aligned and contiguous:
every chunk that overlaps [start, stop) is kept whole, so a window
starting or ending mid-chunk pulls in that partial edge chunk (splits are
chunk-granular -- you subset whole chunks, never individual samples). Use it
for a single contiguous window; it is not a tool for scattered/boolean
selections (those would drag in straddling chunks and silently add samples).
valid_anchor_range ¶
Half-open [lo, hi) of anchor sample-indices whose every windowed read
anchor + offset stays in [0, n_samples) -- the anchors a windowed dataset
may draw, with array-edge anchors dropped.
Offsets {-1, 0, 1} over T samples -> anchors [1, T-1). Range is the
only validity the engine enforces; whether the user's offset choices define a
meaningful (non-leaky) task is theirs to decide (DESIGN, M-W). Empty/too-wide
windows return an empty range (lo, lo).
arraylake_store ¶
Open an Arraylake repo and return its read-only Icechunk session store.
Auth comes from a cached al auth login or ARRAYLAKE_TOKEN; the client
vends the bucket credentials for the repo. The returned object is a zarr-v3
Store bound to the branch snapshot -- exactly what the engine accepts.
Requires insitubatch[arraylake].
close_store ¶
Best-effort teardown for a store that holds an async fsspec session (gcsfs, s3fs).
Such a backend creates an aiohttp session on the first event loop that awaits it --
for a zarr store, that is zarr's loop, not fsspec's -- but gcsfs's finalizer captures
fs.loop (which is None here) and closes the session on the wrong loop at GC,
spewing a harmless-looking "Task was destroyed / attached to a different loop"
traceback and leaking the connection. Closing the session here on the loop it actually
lives on makes that finalizer a no-op.
A no-op for stores with no such session (obstore's ObjectStore has no .fs) and
for already-closed or not-running loops. gcsfs recreates the session lazily, so a
store closed here still works if reused -- but call this only when done with it.
ensure_local_dir ¶
For a file:// URL, create the target directory so writes can land.
obstore's LocalStore will not create the prefix for you. No-op for non-file schemes. Returns the URL unchanged for chaining.
fsspec_store ¶
Return a zarr FsspecStore for url (any fsspec-supported backend).
Reaches stores via a backend fsspec filesystem -- notably GCS Rapid/zonal
buckets (gRPC) and GCS requester-pays, which obstore does not currently
support. **storage_options pass straight through to
FsspecStore.from_url (credentials, project, endpoint, Rapid config, ...).
Requires an fsspec backend for the URL scheme: insitubatch[gcsfs] for
gs://, or bring your own (s3fs, ...). A sync backend (e.g. local
file://) is auto-wrapped as async by zarr; gs:// via gcsfs is
natively async. See :func:obstore_store for the obstore-backed constructor.
icechunk_store ¶
Return the read-only session Store of an Icechunk repository at url.
s3://bucket/prefix (the common case for public archives), gs://, or a local
file:///path. Public buckets need anonymous=True; otherwise credentials come
from the environment the way every other AWS/GCS tool finds them (profile, instance
role, AWS_* variables, an SSO login). Extra kwargs pass through to Icechunk's
storage constructor.
Icechunk is a versioned format, so a store is a snapshot: branch picks which one,
and the returned session is read-only, which is what the engine wants -- training reads
a fixed view of the data even while a writer appends to the repo.
Requires insitubatch[icechunk]::
store = icechunk_store(
"s3://dynamical-noaa-gfs/noaa-gfs-analysis/v0.1.0.icechunk",
anonymous=True, region="us-west-2",
)
For an Arraylake-hosted repo use :func:arraylake_store instead: it is addressed by
catalog name rather than URL and authenticates with an Arraylake login.
obstore_store ¶
Return an obstore-backed zarr Store for url (any obstore scheme).
file:///abs/path.zarr for local; s3://bucket/path.zarr for cloud.
Extra kwargs pass through to obstore.store.from_url (region,
credentials, client options, ...). The read path stays pure Rust -- no fsspec
Python layer.
open_geometries ¶
Introspect a zarr group Store into {name: ArrayGeometry}.
Lets InSituDataset be built from a store alone -- geometry (shape,
chunks, dtype) is read from the array metadata rather than hand-specified.
Build the store with :func:obstore_store / :func:fsspec_store /
:func:arraylake_store, or pass any prebuilt zarr Store.
sample_axis names which physical axis is the outer (sample) axis for
every returned variable -- 0 (default: time for ERA5/HRRR) or, e.g., the
Z of an OME-NGFF (T,C,Z,Y,X) stack sampled slice-by-slice (sample_axis=2).
Variables that need different sample axes are built individually (construct
:class:ArrayGeometry per array); the shape/chunks stay in physical order.
applies ¶
Restrict a chunk_transform to the named arrays; every other array passes through.
arrays are zarr array paths (what open_geometries keys on, and what
chunk.read.array carries) -- not dict labels, since two labels may alias one array
(t2m_now / t2m_next). Unknown names raise at dataset construction::
InSituDataset(..., chunk_transforms=[
applies(["2m_temperature"], kelvin_to_celsius),
Coarsen(2), # bare = every array
])
Declaring scope here rather than with an if chunk.read.array == ... inside the
transform is not a style preference: an in-body gate is invisible to the engine, which
folds a reshaping transform's output_inner into every array's declared geometry.
The unaffected arrays are then gathered as truncated prefixes of themselves and can
never revive from cache -- with no exception raised. Scope also enters the cache
fingerprint per array, so editing one variable's transform leaves the rest valid.
Scope is matched against the zarr array path. Two labels that alias one array
(t2m_now / t2m_next) therefore share one chunk pipeline; per-label work belongs in
a batch_transform.
A transform may be parameterized by variable (a fitted scaler's statistics
legitimately are) but must not decide whether it runs. If it can validate a declared
scope -- :class:StandardScaler checks the names against its own stats dict -- it
defines validate_scope(arrays), which this calls at construction rather than leaving
it to fail in a decode thread.
InSituDataset¶
insitubatch.InSituDataset ¶
A framework-neutral source of shuffled numpy batches from Zarr, split-aware.
The dataset is not itself iterated -- you iterate one of its split views:
:attr:train (shuffled), :attr:val, :attr:test, :attr:all (deterministic).
All four share one :class:ChunkPool, so a chunk that two splits both read --
e.g. a windowed read spilling across a split boundary -- is decoded once::
ds = InSituDataset(store, manifest, geometries=geoms, batch_size=32)
for batch in ds.train: ... # one epoch; ds.set_epoch(e) reshuffles
for batch in ds.val: ...
One epoch over a view = permute the split's chunks -> walk shuffle-blocks, stream-fetching
each block's stored chunks into the pool -> gather coalesced batches cut over the whole
epoch (so only the last is short, and one may span a block boundary) -> evict.
Batches are numpy :class:Batch; convert to a framework with
:mod:insitubatch.frameworks (as_torch / to_jax / as_tf_dataset). A
different per-split configuration (e.g. train-only augmentation) is a separate dataset.
Two preprocessing hooks, placed by cost (full model in the docs, "Transforms"):
chunk_transforms--(DecodedChunk) -> DecodedChunk, run per chunk before shuffle, seeing one variable. The cacheable home for elementwise, per-variable, deterministic work (scaling, unit conversion, dtype cast); amortized over every sample in the chunk and reused across epochs. Restrict one to particular variables with :func:~insitubatch.transforms.applies--applies(["2m_temperature"], k_to_c)-- rather than testingchunk.read.arrayinside the transform: a name test in the body is invisible to the engine, which would fold a reshaping transform's declared output into every variable's geometry and truncate the ones it does not touch.batch_transforms--(Batch) -> Batch, run per assembled batch, seeing all variables aligned on the sample axis. For cross-variable derived fields and per-sample random augmentation; runs after the cache, so it is never cached.
Runnable side-by-side example: examples/transforms.py.
One writer per cache_dir. The cache directory (and persist=True on top of
it) is arbitrated across processes by an advisory lock held for the dataset's
lifetime: one writer at a time, and a second fails fast rather than silently
corrupting the first's chunks. Several jobs may read one warm cache at once with
readonly_cache=True, which takes the lock shared, never writes, and raises on a
miss -- it asserts the cache is complete for what this run reads, rather than
quietly falling back to fetching. Put cache_dir on local NVMe: over a network
filesystem the lock may be emulated per client and we can only warn. See
docs/tuning.md, "Sharing a cache_dir between processes".
all
property
¶
Iterable over every split's chunks (deterministic) -- e.g. full-archive inference.
set_epoch ¶
Call from the training loop so each epoch reshuffles deterministically.
describe ¶
What this dataset will do, from geometry and configuration -- no store access.
iterations is how many passes will share the pool at once (zip(ds.train,
ds.val) is two), since each holds its own chunk references and the automatic
budget covers one. See :meth:print_summary for the formatted view, and
:mod:insitubatch.summary for what each number means.
print_summary ¶
Print :meth:describe for a human. On demand -- construction stays quiet.
close ¶
Release the cache pool's backing (mmap handles, cached chunks) and any async store session.
The pool persists across epochs, so close it when done training -- not per
epoch. With persist=True the cache files + manifest are kept on disk for a
future run (only the in-memory handles are released); otherwise the mmap spill
files are unlinked. An fsspec/gcsfs store's aiohttp session is closed on its own
loop here (a no-op for obstore) so it does not leak or spew a teardown traceback
at GC; gcsfs recreates it lazily if the store is reused. Idempotent; also called
on GC.
Debugging¶
Environment report for bug reports -- one paste instead of a dozen questions.
print_debug_info() dumps the versions and interpreter facts that actually change
insitubatch's behavior: the storage stack (zarr / obstore / numpy), whichever framework
adapter is installed, and the free-threading state. That last one matters more here
than in most libraries -- the ChunkPool's lock discipline and its cross-thread readiness
signalling are exercised very differently on a 3.13t build with the GIL genuinely off, and
"works for me" reports have turned on exactly that difference::
python -c "import insitubatch; insitubatch.print_debug_info()"
Nothing here imports an optional dependency: versions come from installed distribution metadata. Importing torch and JAX in one process crashes (duplicate OpenMP/XLA runtimes), so a debug helper that imported what it reports on would take the process down precisely when someone is trying to report a bug.
print_debug_info ¶
Print the environment report. Paste the output into a bug or performance report.
debug_info ¶
Return the environment report as an ordered label -> value mapping.
The structured half of :func:print_debug_info; useful if you want to attach the same
facts to a benchmark record rather than paste them into an issue.
The static report¶
The entry points are the two InSituDataset methods above — describe() for the dict,
print_summary() for the rendering. These are the shapes describe() returns.
What will this dataset actually do? -- a static report, before anything runs.
InSituDataset.describe() answers from geometry and configuration only: it opens no
store, fetches nothing, and runs no pass. That is the point. The facts it reports -- a
360-byte gather run, a 2x ragged residency multiplier, a budget sized for one iteration
when you meant to run two -- are exactly the ones you want before waiting an hour to
discover them, and a report that had to touch the store would be unusable in the situation
it exists for.
It is deliberately on demand. Construction stays quiet (bar the two warnings that already fire), because layout advice on every dataset you build teaches people to skim our logs.
Runtime facts -- queue depths, per-stage timing, cache hits, peak residency -- are a different surface with a different cost model, and live in #45.
:func:working_set_bytes is shared with :class:~insitubatch.source.InSituDataset, which
sizes its automatic budget with it. One formula, called twice: a report that predicted a
different number from the one the engine uses would be worse than no report.
DatasetReport ¶
Bases: TypedDict
The whole static picture. Returned by :meth:InSituDataset.describe.
VariableReport ¶
Bases: TypedDict
One variable's geometry, and the per-chunk bytes derived from it.
ConfigReport ¶
Bases: TypedDict
What the dataset resolved to -- derived values, not the arguments passed in.
MemoryReport ¶
Bases: TypedDict
The accounted pieces of the memory model, evaluated for this configuration.
Note ¶
Bases: TypedDict
One finding. severity separates "this will hurt" from "worth knowing".
The runtime report¶
Where describe() predicts before a pass, ds.last_pass reports after one. It is a
PassStats, and last_pass.limiting_stage names the one stage worth going to fix.
What did this pass actually do? -- runtime counters, and the stage they accuse.
:mod:~insitubatch.summary predicts from geometry before anything runs; this module
reports what happened after it did. The two are deliberately separate surfaces: a
prediction that had to touch the store would be useless in the situation it exists for,
and a measurement that could be computed from configuration would not be a measurement.
The payload is :class:PassStats, and the part worth having is :func:bottleneck -- a
rule mapping observed pressure to the one stage to go fix. Depth counters alone tell you
the batch queue was empty; they do not tell you whether that was the store, the decode
pool, or a residency budget too small to admit the next chunk. Those three want opposite
responses, and telling them apart is the whole job.
The measurement rule, which is sharper than the obvious one. Wall-clock around a thread hop measures GIL wait, not work. That is how our scatter memcpy once read as 51% of the hot path when its real share is 7.7-10.1%. So:
- in-thread cost uses :func:
time.thread_time-- CPU actually burned by that thread; - waiting uses :func:
time.perf_counter, because there the waiting is the quantity.
A stage timer that over-attributes is worse than no timer: it sends people to optimize a
stage that was never the problem. The same rule bites inside a coroutine: x += await f()
loads x before evaluating the right-hand side, so the await suspends between the load
and the store and concurrent tasks lose each other's updates. Bind the cost to a local
first. Single-writer means no await between load and store, not merely one thread.
Why residency rarely wins on its own. Admission parking is backpressure: any slow
stage downstream keeps chunks pinned, the budget fills, and admission parks behind it. So a
budget that is genuinely too small sits in a narrow band between two neighbours -- a tie
with the stage actually causing the pressure (where :func:bottleneck names that stage),
and the provably-terminal stall the pool raises residency budget exhausted on. The value
of admission_parked_s is mostly in telling those two apart, which nothing else can.
Wait totals are summed across concurrent tasks and will exceed wall time. Many tiles
are in flight at once, so fetch_wait_s is task-seconds, not a share of the pass. Divide
by :attr:Depths.inflight_peak for an order-of-magnitude per-tile feel, and read the
ratios between stages rather than any absolute against the clock.
PassStats
dataclass
¶
One pass over one split: counters, depths, times, and the stage they accuse.
consumer_idle
property
¶
True when the caller's loop body did essentially nothing.
A for batch in ds.train: pass loop takes batches as fast as they are produced, so
the queue is empty on every sample no matter how fast the loader is. That is a
throughput measurement, not a starved pipeline, and saying "the queue ran empty" of
it would be true and useless.
StageTimes
dataclass
¶
Cumulative seconds by stage. See the module docstring for the units rule.
producer_s
property
¶
Accounted producer-side cost, the denominator for a stage's share.
Excludes consumer_s, which is the caller's own work and not a stage we can tune.
Depths
dataclass
¶
bottleneck ¶
Which stage limited this pass, and what to do about it.
A caller whose loop body does essentially nothing is answered first: with no training step there is no "the consumer is the constraint" to reach, whatever the queue looks like. Otherwise the rule reads the batch queue, because a full queue settles the question: every producer stage kept up and the consumer is the constraint, which is the state you want. Only once the queue is repeatedly empty is there a producer stage to accuse, and then the accusation comes from which one spent the time -- not from which one feels slow.
Returns ("unknown", ...) rather than guessing when the evidence does not separate
the candidates: no samples, or no stage clearly dominant. A confident wrong answer here
costs more than an admission, because it sends someone to tune the wrong knob.
Framework adapters¶
Thin, optional framework adapters: numpy Batch -> torch / JAX / TF via DLPack.
The core (:mod:insitubatch.source) yields numpy :class:Batch objects and imports
no framework. These adapters convert a batch's arrays to a framework's tensors with
DLPack (zero-copy on CPU where the framework supports it). The wrapping differs per
ecosystem -- there is no single cross-framework "dataset" base class:
-
torch has one.
DataLoaderrequires aDataset/IterableDatasetsubclass (itisinstance-checks), so :func:as_torchwraps the stream in one::DataLoader(as_torch(ds), batch_size=None, num_workers=0)batch_size=None(the stream already yields assembled batches) andnum_workers=0(parallelism is in our event loop; forking re-introduces the redundant-read problem). * JAX has none -- it is loader-agnostic. Iterate the dataset and call :func:to_jaxper batch. * TF adapts via a factory, not a base class: :func:as_tf_datasetwraps the stream intf.data.Dataset.from_generator.
Each framework is imported lazily inside its function, so importing this module costs
nothing and a missing framework raises a clear, actionable error. sample_indices
(provenance) stays on the numpy Batch; only the model-input arrays are converted.
to_torch ¶
Convert a numpy Batch to a dict of torch tensors (DLPack; zero-copy on CPU).
With device set, the copy is issued here rather than by the caller, and the batch's
source arrays are held until it lands. That is what makes pinned batch buffers safe to
recycle -- see :class:_InFlight. Without it the behaviour is unchanged: CPU tensors
aliasing the batch, and the caller's own .to(...) afterwards.
as_torch ¶
Wrap a split view (e.g. ds.train) as a torch IterableDataset for DataLoader.
Each yielded item is a dict[str, torch.Tensor] (via :func:to_torch). Use
DataLoader(as_torch(ds.train), batch_size=None, num_workers=0).
Passing device moves batches here instead of in the training loop, and switches the
batch buffers to page-locked memory. The two are one option rather than two on
purpose: pinning is what makes an H2D copy genuinely asynchronous, so it is only safe
when whoever issues the copy also knows when it landed. A pin_memory=True flag the
caller could set without device would hand out pinned buffers under the caller's own
non_blocking copy and reintroduce a use-after-recycle we would then have to document
instead of prevent.
What it is worth runs between two measured bounds, and how compute-bound the loop is
decides where you land. In isolation (bench/probe_batch_buffers.py --arms overlap) a
pinned copy below ~37 MiB is hidden completely -- the pinned arm sits on the compute
floor, while a pageable one costs more than its own transfer time because it stages
through a driver bounce buffer and cannot overlap at all. Against a compute step that
already hides the copy, it is worth nothing. A real GPU training loop lands in between:
+1 to +2%, roughly flat over a 4x payload range (32 to 128 MiB per step) rather than
switching on above a threshold.
Expect no more precision than that range. An A/B/A bracket on an unchanged checkout drifts
~0.7% per block -- half the signal -- so the bound is +0.8% to +2.2% depending on how the
drift is modelled. See docs/benchmarks.md, which also explains why % of ceiling cannot
measure this at all.
pin_budget_bytes caps the page-locked total (default: an eighth of RAM, since the
kernel cannot reclaim it); past that buffers are pageable again, with one warning.
to_jax ¶
Convert a numpy Batch to a dict of jax.Array (DLPack), on JAX's default device.
Zero-copy only when the batch buffer sits on a 128-byte boundary: XLA:CPU requires that
alignment and silently falls back to a copy otherwise. numpy guarantees only 16, so with
today's per-batch np.empty roughly half of batches are copied (measured 20/40 —
bench/probe_batch_buffers.py --arms jax). Correct either way, just not free.
jnp.from_dlpack imports a host buffer, so its result is committed to cpu:0 --
and jax.jit does not move a committed array. Returning that directly meant a GPU run
trained on the CPU at full correctness with nothing raised to say so, so the batch is
placed on jax.devices()[0] here, which is where every other jax array-creation path
puts it. Pass device to choose another. On a CPU-only build the put is a
pass-through that still aliases the exported buffer, so it costs nothing there.
The transfer is asynchronous -- JAX documents it so -- but JAX keeps the source
referenced across it, so the pool's liveness poll cannot reclaim a buffer mid-DMA
(bench/probe_batch_buffers.py --arms all is what keeps that honest; see DESIGN.md).
Unlike :func:to_torch this adapter does not own the transfer: JAX cannot consume
pinned host memory (jax#22346) and exposes no event to gate recycling on.
to_tf ¶
Convert a numpy Batch to a dict of tf.Tensor (one CPU copy per variable).
Unlike torch/JAX -- whose array-accepting from_dlpack manages the exported buffer's
lifetime correctly -- TensorFlow only exposes the experimental from_dlpack(capsule),
which mishandles ownership of the exported numpy buffer: under the concurrent allocation
of the prefetch decode threads it double-frees that buffer and aborts the process
(SIGABRT, no message). convert_to_tensor copies into a TF-owned tensor instead, so TF
never touches insitu-managed memory; the batch is already an owned array, so this is a
single CPU copy. (torch stays zero-copy, JAX when aligned -- see :func:to_jax; TF's is a
DLPack limitation, not ours.)
as_tf_dataset ¶
Wrap a split view (e.g. ds.val) as a tf.data.Dataset via from_generator.
output_signature is inferred from the view's geometries: each variable is
(None, *inner) (None = the variable last-batch size) with the variable's dtype.
Both from_generator here and :func:to_tf copy into the TF runtime -- TF has no
reliable zero-copy path from insitu's buffers (its experimental DLPack mishandles buffer
ownership; see :func:to_tf). Call :func:to_tf on the raw stream when you want plain
dict[str, tf.Tensor] batches instead of a tf.data.Dataset.