auroraomics 0.1.0.dev0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
auroraomics/h5ad.py ADDED
@@ -0,0 +1,398 @@
1
+ """Write the standard result file with ``h5py`` alone.
2
+
3
+ The result of a prediction is an ``.h5ad``: the expression matrix in ``X``, one
4
+ row per spot and one column per gene, with the per-spot table in ``obs``, the
5
+ per-gene table in ``var``, the spot coordinates in ``obsm["spatial"]``, the run
6
+ provenance in ``uns`` and any derived matrix in ``layers``. It reads back as an
7
+ ordinary ``AnnData``.
8
+
9
+ **Why this module does not import anndata.** Two reasons, and the second is the
10
+ load-bearing one:
11
+
12
+ * anndata pulls pandas and its stack. This code runs inside GPU images whose
13
+ compiled dependencies are ABI-pinned, and beside a model that wants the
14
+ memory; the format is a handful of HDF5 attributes, so paying for that stack
15
+ to write them is a poor trade.
16
+ * a matrix bigger than memory cannot be handed to a library that takes an
17
+ in-memory array. Accumulating a (300k x 19k) float32 matrix to write it in
18
+ one call needs ~23 GB, which is simply not available — and the model that
19
+ produced it is holding several GB of its own. So ``X`` is created empty,
20
+ chunked and compressed, and filled batch by batch as the batches arrive:
21
+ peak memory is one batch, whatever the slide.
22
+
23
+ The format is the on-disk contract anndata 0.10 and 0.11 read, expressed as
24
+ ``encoding-type``/``encoding-version`` attributes on every element. The
25
+ :func:`write_result` docstring names each one; the test suite holds the actual
26
+ guarantee, by opening what this module writes with anndata itself.
27
+ """
28
+
29
+ from __future__ import annotations
30
+
31
+ from collections.abc import Iterable, Mapping, Sequence
32
+ from pathlib import Path
33
+ from typing import Any
34
+
35
+ import h5py
36
+ import numpy as np
37
+
38
+ from . import contracts
39
+
40
+ # ── the on-disk encoding vocabulary ─────────────────────────────────────────
41
+ _ANNDATA = {"encoding-type": "anndata", "encoding-version": "0.1.0"}
42
+ _DICT = {"encoding-type": "dict", "encoding-version": "0.1.0"}
43
+ _DATAFRAME = {"encoding-type": "dataframe", "encoding-version": "0.2.0"}
44
+ _ARRAY = {"encoding-type": "array", "encoding-version": "0.2.0"}
45
+ _STRING_ARRAY = {"encoding-type": "string-array", "encoding-version": "0.2.0"}
46
+ _STRING = {"encoding-type": "string", "encoding-version": "0.2.0"}
47
+ _SCALAR = {"encoding-type": "numeric-scalar", "encoding-version": "0.2.0"}
48
+
49
+ #: The name a dataframe group gives its index when the caller has none. This is
50
+ #: anndata's own convention for an unnamed index, and it round-trips as one.
51
+ _DEFAULT_INDEX_NAME = "_index"
52
+
53
+ #: Rows per chunk of a matrix. Chunks are whole rows across every column, so a
54
+ #: chunk is one compression unit and one disk write; 1024 rows of 19k float32
55
+ #: genes is ~78 MB, which is why the cache below matters.
56
+ _CHUNK_ROWS = 1024
57
+
58
+ #: Floor and ceiling for the HDF5 raw-chunk cache. See :func:`_cache_bytes`.
59
+ _CACHE_FLOOR = 32 * 1024 * 1024
60
+ _CACHE_CEILING = 256 * 1024 * 1024
61
+
62
+ _STRING_DTYPE = h5py.string_dtype(encoding="utf-8")
63
+
64
+
65
+ class ResultError(ValueError):
66
+ """The pieces handed to :func:`write_result` cannot form a valid result."""
67
+
68
+
69
+ def _cache_bytes(n_cols: int, itemsize: int, n_rows: int) -> int:
70
+ """How much raw-chunk cache to give the file, and why it is not the default.
71
+
72
+ Batches advance monotonically through the matrix and are far smaller than a
73
+ chunk — a model predicts a few dozen rows at a time while a chunk is
74
+ :data:`_CHUNK_ROWS` of them. With HDF5's default 1 MB cache the active
75
+ chunk does not fit, so every batch evicts it, and every eviction
76
+ re-compresses and rewrites the *whole* chunk: a read-modify-write storm
77
+ that stalls the producer whenever the target is a slow disk. Sized to hold
78
+ several whole chunks, each chunk is compressed and written exactly once, as
79
+ it fills.
80
+
81
+ The ceiling keeps the process modest, and the floor keeps a narrow matrix
82
+ (few genes, so small chunks) from inheriting a cache too small to help.
83
+ """
84
+ chunk_rows = max(1, min(_CHUNK_ROWS, n_rows))
85
+ chunk_bytes = chunk_rows * max(1, n_cols) * itemsize
86
+ return int(min(max(chunk_bytes * 4, _CACHE_FLOOR), _CACHE_CEILING))
87
+
88
+
89
+ def _is_string_data(values: Any) -> bool:
90
+ """Does this column hold text rather than numbers?"""
91
+ arr = np.asarray(values)
92
+ if arr.dtype.kind in "US":
93
+ return True
94
+ if arr.dtype.kind == "O":
95
+ return all(isinstance(v, (str, bytes)) for v in arr.ravel())
96
+ return False
97
+
98
+
99
+ def _write_string_array(
100
+ parent: h5py.Group, name: str, values: Any, compression: str | None
101
+ ) -> h5py.Dataset:
102
+ data = np.asarray(
103
+ [v.decode() if isinstance(v, bytes) else str(v) for v in np.asarray(values)],
104
+ dtype=object,
105
+ )
106
+ dataset = parent.create_dataset(
107
+ name, data=data, dtype=_STRING_DTYPE, compression=compression
108
+ )
109
+ dataset.attrs.update(_STRING_ARRAY)
110
+ return dataset
111
+
112
+
113
+ def _write_column(
114
+ parent: h5py.Group, name: str, values: Any, compression: str | None
115
+ ) -> None:
116
+ if _is_string_data(values):
117
+ _write_string_array(parent, name, values, compression)
118
+ return
119
+ arr = np.asarray(values)
120
+ if arr.ndim != 1:
121
+ raise ResultError(
122
+ f"column {name!r} must be one-dimensional, got shape {arr.shape}"
123
+ )
124
+ dataset = parent.create_dataset(name, data=arr, compression=compression)
125
+ dataset.attrs.update(_ARRAY)
126
+
127
+
128
+ def _write_dataframe(
129
+ parent: h5py.Group,
130
+ name: str,
131
+ columns: Mapping[str, Any],
132
+ index: Sequence[Any],
133
+ index_name: str,
134
+ compression: str | None,
135
+ ) -> None:
136
+ """Write one table as a ``dataframe`` group.
137
+
138
+ The group carries the column order and the index's name as attributes: the
139
+ index is a dataset like any other, and ``_index`` is what says which one it
140
+ is. Order is an attribute rather than the file's own key order because HDF5
141
+ sorts its keys alphabetically, so ``column-order`` is the only place the
142
+ caller's column order survives.
143
+ """
144
+ if index_name in columns:
145
+ raise ResultError(
146
+ f"{name}: {index_name!r} is the index name, so it cannot also be a "
147
+ "column — rename one of them"
148
+ )
149
+ group = parent.create_group(name)
150
+ group.attrs.update(_DATAFRAME)
151
+ group.attrs["_index"] = index_name
152
+ group.attrs.create(
153
+ "column-order", data=list(columns), shape=(len(columns),), dtype=_STRING_DTYPE
154
+ )
155
+ _write_string_array(group, index_name, index, compression)
156
+ for column, values in columns.items():
157
+ arr = np.asarray(values)
158
+ if arr.shape[0] != len(index):
159
+ raise ResultError(
160
+ f"{name}: column {column!r} has {arr.shape[0]} values for "
161
+ f"{len(index)} rows"
162
+ )
163
+ _write_column(group, column, values, compression)
164
+
165
+
166
+ def _write_mapping(
167
+ parent: h5py.Group, name: str, values: Mapping[str, Any], compression: str | None
168
+ ) -> h5py.Group:
169
+ """Write a nested mapping as a ``dict`` group, recursing into sub-mappings.
170
+
171
+ Scalars, strings, numeric arrays and string arrays each get the encoding
172
+ anndata expects for them; anything else is refused by name rather than
173
+ written as something a reader will mis-parse.
174
+ """
175
+ group = parent.create_group(name)
176
+ group.attrs.update(_DICT)
177
+ for key, value in values.items():
178
+ if not isinstance(key, str):
179
+ raise ResultError(f"{name}: keys must be strings, got {key!r}")
180
+ if isinstance(value, Mapping):
181
+ _write_mapping(group, key, value, compression)
182
+ elif isinstance(value, (str, bytes)):
183
+ dataset = group.create_dataset(
184
+ key,
185
+ data=value.decode() if isinstance(value, bytes) else value,
186
+ dtype=_STRING_DTYPE,
187
+ )
188
+ dataset.attrs.update(_STRING)
189
+ elif isinstance(value, (bool, int, float, np.bool_, np.number)):
190
+ dataset = group.create_dataset(key, data=value)
191
+ dataset.attrs.update(_SCALAR)
192
+ elif isinstance(value, (Sequence, np.ndarray)):
193
+ if _is_string_data(value):
194
+ _write_string_array(group, key, value, compression)
195
+ else:
196
+ dataset = group.create_dataset(
197
+ key, data=np.asarray(value), compression=compression
198
+ )
199
+ dataset.attrs.update(_ARRAY)
200
+ else:
201
+ raise ResultError(
202
+ f"{name}.{key} is a {type(value).__name__}, which has no "
203
+ "representation in this file format; convert it to a mapping, "
204
+ "a string, a number or an array first"
205
+ )
206
+ return group
207
+
208
+
209
+ def _write_matrix(
210
+ parent: h5py.Group,
211
+ name: str,
212
+ source: Any,
213
+ n_rows: int,
214
+ n_cols: int,
215
+ dtype: Any,
216
+ compression: str | None,
217
+ ) -> None:
218
+ """Create a chunked matrix and fill it from an array or a stream of batches.
219
+
220
+ ``source`` is either a two-dimensional array, or an iterable of row batches
221
+ that together cover exactly ``n_rows``. Batches are written where they
222
+ land and then dropped, so nothing here ever holds more than one of them —
223
+ that is the whole point of taking an iterable, and a test pins it by
224
+ checking that earlier batches have been collected before the last arrives.
225
+ """
226
+ dataset = parent.create_dataset(
227
+ name,
228
+ shape=(n_rows, n_cols),
229
+ dtype=dtype,
230
+ chunks=(max(1, min(_CHUNK_ROWS, n_rows)), max(1, n_cols)),
231
+ compression=compression,
232
+ )
233
+ dataset.attrs.update(_ARRAY)
234
+
235
+ # A whole matrix, or a stream of row batches. Anything that already knows
236
+ # it is two-dimensional (a numpy array, an open dataset from another file)
237
+ # is the whole matrix; everything else is iterated.
238
+ batches: Iterable[Any]
239
+ if getattr(source, "ndim", None) == 2:
240
+ batches = (source,)
241
+ else:
242
+ batches = source
243
+
244
+ written = 0
245
+ for index, batch in enumerate(batches):
246
+ arr = np.asarray(batch)
247
+ if arr.ndim != 2 or arr.shape[1] != n_cols:
248
+ raise ResultError(
249
+ f"{name}: batch {index} has shape {arr.shape}, expected "
250
+ f"(rows, {n_cols})"
251
+ )
252
+ end = written + arr.shape[0]
253
+ if end > n_rows:
254
+ raise ResultError(
255
+ f"{name}: batches carry more than {n_rows} rows (batch {index} "
256
+ f"would reach row {end})"
257
+ )
258
+ dataset[written:end, :] = arr.astype(dtype, copy=False)
259
+ written = end
260
+ if written != n_rows:
261
+ raise ResultError(
262
+ f"{name}: batches carried {written} rows for {n_rows} observations "
263
+ "— a result whose matrix is part zero-filled is worse than no file"
264
+ )
265
+
266
+
267
+ def _row_count(columns: Mapping[str, Any] | None, index: Sequence[Any] | None) -> int:
268
+ if index is not None:
269
+ return len(index)
270
+ if columns:
271
+ first = next(iter(columns.values()))
272
+ return int(np.asarray(first).shape[0])
273
+ raise ResultError(
274
+ "cannot tell how many rows this result has: pass an index, or at least "
275
+ "one column. A streamed matrix is created before its first batch "
276
+ "arrives, so the shape has to be known in advance"
277
+ )
278
+
279
+
280
+ def write_result(
281
+ path: str | Path,
282
+ *,
283
+ obs: Mapping[str, Any] | None = None,
284
+ var: Mapping[str, Any] | None = None,
285
+ x: Any = None,
286
+ obs_index: Sequence[Any] | None = None,
287
+ var_index: Sequence[Any] | None = None,
288
+ spatial: Any = None,
289
+ uns: Mapping[str, Any] | None = None,
290
+ layers: Mapping[str, Any] | None = None,
291
+ obs_index_name: str = _DEFAULT_INDEX_NAME,
292
+ var_index_name: str | None = None,
293
+ dtype: Any = np.float32,
294
+ compression: str | None = "gzip",
295
+ ) -> Path:
296
+ """Write one result file and return its path.
297
+
298
+ Parameters
299
+ ----------
300
+ path:
301
+ Where to write. An existing file is replaced.
302
+ obs, var:
303
+ The per-spot and per-gene tables, as mappings of column name to a
304
+ one-dimensional sequence. Text columns are stored as text; everything
305
+ else keeps its numpy dtype.
306
+ x:
307
+ The expression matrix: either a two-dimensional array, or an **iterable
308
+ of row batches** covering every observation in order. The iterable form
309
+ is what makes a matrix larger than memory writable — see
310
+ :func:`_write_matrix`.
311
+ obs_index, var_index:
312
+ Row and column names. Default to the positions as strings.
313
+ var_index_name:
314
+ The name ``var``'s index carries in the file. Defaults to the gene
315
+ identifier the shared result contract declares, so a reader knows what
316
+ kind of identifier it is looking at without guessing.
317
+ spatial:
318
+ An ``(n_obs, 2)`` array of spot coordinates, written to
319
+ ``obsm["spatial"]`` — where every viewer looks for them.
320
+ uns:
321
+ Free-form provenance, nested mappings allowed.
322
+ layers:
323
+ Extra matrices of the same shape as ``x``, each an array or an iterable
324
+ of batches, written under ``layers/<name>``.
325
+ dtype:
326
+ The matrix dtype. float32 by default: float16 halves the file but is
327
+ not enough precision for the statistics these matrices get fed to, and
328
+ several tools refuse it outright.
329
+ compression:
330
+ Passed to HDF5. ``"gzip"`` by default; ``None`` writes uncompressed.
331
+ """
332
+ obs = dict(obs or {})
333
+ var = dict(var or {})
334
+ n_obs = _row_count(obs, obs_index)
335
+ n_vars = _row_count(var, var_index)
336
+ if obs_index is None:
337
+ obs_index = [str(i) for i in range(n_obs)]
338
+ if var_index is None:
339
+ var_index = [str(i) for i in range(n_vars)]
340
+ if var_index_name is None:
341
+ var_index_name = str(contracts.result_layout()["var_index"])
342
+ if x is None:
343
+ raise ResultError("a result needs a matrix; pass x=")
344
+
345
+ if layers and "X" in layers:
346
+ raise ResultError(
347
+ "a layer cannot be called 'X' — that is the matrix itself, and one "
348
+ "of the two would silently win"
349
+ )
350
+ matrices: dict[str, Any] = {"X": x}
351
+ matrices.update(layers or {})
352
+ itemsize = max(np.dtype(dtype).itemsize, 1)
353
+
354
+ path = Path(path)
355
+ # The cache is a property of the OPEN FILE, so it has to be sized before
356
+ # anything is created in it — which is why the shape is settled above.
357
+ with h5py.File(
358
+ path,
359
+ "w",
360
+ rdcc_nbytes=_cache_bytes(n_vars, itemsize, n_obs),
361
+ rdcc_nslots=1009,
362
+ # Prefer evicting chunks that are fully written, which is every chunk
363
+ # behind the advancing write position.
364
+ rdcc_w0=1.0,
365
+ ) as handle:
366
+ handle.attrs.update(_ANNDATA)
367
+ _write_dataframe(handle, "obs", obs, obs_index, obs_index_name, compression)
368
+ _write_dataframe(handle, "var", var, var_index, var_index_name, compression)
369
+
370
+ obsm = handle.create_group("obsm")
371
+ obsm.attrs.update(_DICT)
372
+ if spatial is not None:
373
+ coords = np.asarray(spatial)
374
+ if coords.ndim != 2 or coords.shape[0] != n_obs:
375
+ raise ResultError(
376
+ f"spatial must be (n_obs, dims) = ({n_obs}, ...), got shape "
377
+ f"{coords.shape}"
378
+ )
379
+ dataset = obsm.create_dataset(
380
+ "spatial", data=coords, compression=compression
381
+ )
382
+ dataset.attrs.update(_ARRAY)
383
+
384
+ _write_mapping(handle, "uns", dict(uns or {}), compression)
385
+
386
+ layers_group = handle.create_group("layers")
387
+ layers_group.attrs.update(_DICT)
388
+ for name, source in matrices.items():
389
+ parent = handle if name == "X" else layers_group
390
+ _write_matrix(
391
+ parent, name, source, n_obs, n_vars, dtype, compression
392
+ )
393
+
394
+ # anndata writes these three even when empty, and a reader that expects
395
+ # them should not have to special-case a file this module produced.
396
+ for name in ("varm", "obsp", "varp"):
397
+ handle.create_group(name).attrs.update(_DICT)
398
+ return path