From 309a89a6cb5f6405cea11f05f83a72fff09c70e3 Mon Sep 17 00:00:00 2001 From: benjaminfreyuu Date: Fri, 7 Aug 2026 15:21:36 +0200 Subject: [PATCH] babel: keep ATAC/RNA sparse instead of densifying up front MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit babel_train densified the full ATAC (and RNA) matrix via a single _to_dense() before building the Dataset. On real ATAC data this OOMs — the 2022 pbmc_multiome run tried to allocate 222 GiB for a 130095 x 228942 float64 array and crashed (numpy ArrayMemoryError). The tiny 2021 test fixture (1500 peaks) hid it. Restore the original BABEL approach (wukevin/babel sc_data_loaders): keep the matrices sparse (CSR) and densify one cell at a time in PairedDataset.__getitem__ (mirrors BABEL's per-cell `utils.ensure_arr(X[i]).flatten()`). Peaks are sliced per chromosome on a CSC view. babel_predict likewise binarizes sparsely and runs inference in cell chunks instead of loading the whole dense matrix onto the GPU. Peak memory now scales with a minibatch, not the full matrix: on a synthetic 20k x 60k ATAC (9.6 GB dense) the sparse path peaks at 1.7 GB while the old toarray() path is OOM-killed under a 6 GB cap. Verified: viash test passes for babel_train (1/1) and babel_predict (1/1). Co-Authored-By: Claude Opus 4.8 --- src/methods/babel/babel_predict/script.py | 30 +++++++++----- src/methods/babel/babel_train/script.py | 49 +++++++++++++++-------- 2 files changed, 53 insertions(+), 26 deletions(-) diff --git a/src/methods/babel/babel_predict/script.py b/src/methods/babel/babel_predict/script.py index 837d6b72..f9294a2f 100644 --- a/src/methods/babel/babel_predict/script.py +++ b/src/methods/babel/babel_predict/script.py @@ -25,10 +25,6 @@ from model import AssymSplicedAutoEncoder -def _to_dense(X): - return X.toarray() if issparse(X) else np.asarray(X) - - def _lognorm_per_cell(pred_counts, target_sum=1e4): """Convert the NB decoder's raw-count-scale mean prediction into the same log1p(normalized-to-target_sum) space used for the "normalized" layer @@ -74,12 +70,28 @@ def _lognorm_per_cell(pred_counts, target_sum=1e4): "Test ATAC var_names do not match the peak order the model was trained with; " "reindexing across mismatched peak sets is not supported." ) -X_bin = (_to_dense(adata_test_mod1.layers.get("counts", adata_test_mod1.X)) > 0).astype(np.float32) -X_per_chrom = [torch.from_numpy(X_bin[:, idxs]).to(device) for idxs in chrom_groups.values()] +# Binarize the test ATAC while keeping it sparse, pre-slice peaks per chromosome, and +# run inference in cell chunks so the full dense matrix (and a whole-matrix GPU tensor) +# is never materialized -- the same reason babel_train stays sparse. +counts = adata_test_mod1.layers.get("counts", adata_test_mod1.X) +counts = counts.tocsc() if issparse(counts) else csc_matrix(counts) +X_bin = (counts > 0).astype(np.float32) +per_chrom = [X_bin[:, idxs].tocsr() for idxs in chrom_groups.values()] + +n_cells = X_bin.shape[0] +chunk_size = 4096 +pred_chunks = [] with torch.no_grad(): - encoded = model.encoder2(X_per_chrom) - pred_mean, _, _ = model.decoder1(encoded) -pred = _lognorm_per_cell(pred_mean.cpu().numpy()) + for start in range(0, n_cells, chunk_size): + end = min(start + chunk_size, n_cells) + x_pc = [ + torch.from_numpy(c[start:end].toarray().astype(np.float32)).to(device) + for c in per_chrom + ] + encoded = model.encoder2(x_pc) + pred_mean, _, _ = model.decoder1(encoded) + pred_chunks.append(pred_mean.cpu().numpy()) +pred = _lognorm_per_cell(np.concatenate(pred_chunks, axis=0)) out_var = adata_train_mod2.var logger.info("Writing predictions...") diff --git a/src/methods/babel/babel_train/script.py b/src/methods/babel/babel_train/script.py index a0576d87..eca8bac1 100644 --- a/src/methods/babel/babel_train/script.py +++ b/src/methods/babel/babel_train/script.py @@ -5,7 +5,7 @@ import anndata as ad import numpy as np import scanpy as sc -from scipy.sparse import issparse +from scipy.sparse import csc_matrix, csr_matrix, issparse import skorch import torch from torch import nn @@ -38,36 +38,48 @@ from losses import QuadLoss -def _to_dense(X): - return X.toarray() if issparse(X) else np.asarray(X) +def _as_csr(X): + return X.tocsr() if issparse(X) else csr_matrix(X) def _rna_matrix(adata): - counts = _to_dense(adata.layers.get("counts", adata.X)).astype(np.float32) + # Keep everything sparse; the full ATAC/RNA matrices are only densified one cell + # at a time in PairedDataset.__getitem__ (mirrors the original BABEL data loader). + # Densifying up front OOMs on real ATAC data (e.g. 130k cells x 229k peaks would + # need ~222 GiB dense). + counts = _as_csr(adata.layers.get("counts", adata.X)).astype(np.float32) if "normalized" in adata.layers: - X = _to_dense(adata.layers["normalized"]).astype(np.float32) + X = _as_csr(adata.layers["normalized"]).astype(np.float32) else: tmp = ad.AnnData(counts.copy()) sc.pp.normalize_total(tmp) sc.pp.log1p(tmp) - X = tmp.X.astype(np.float32) - size_factors = counts.sum(axis=1, keepdims=True) + X = _as_csr(tmp.X).astype(np.float32) + size_factors = np.asarray(counts.sum(axis=1)).reshape(-1, 1) med = np.median(size_factors[size_factors > 0]) if np.any(size_factors > 0) else 1.0 size_factors = size_factors / (med if med > 0 else 1.0) return X, counts, size_factors.astype(np.float32) def _atac_binarized(adata): - counts = _to_dense(adata.layers.get("counts", adata.X)) + counts = _as_csr(adata.layers.get("counts", adata.X)) return (counts > 0).astype(np.float32) +def _row_tensor(mat, idx): + """Densify a single row of a sparse matrix into a 1-D float32 tensor + (mirrors BABEL's per-cell `utils.ensure_arr(X[i]).flatten()`).""" + return torch.from_numpy(mat[idx].toarray().ravel().astype(np.float32)) + + class PairedDataset(Dataset): def __init__(self, x1, x2_per_chrom, y1, y2, size_factors1): - self.x1 = torch.from_numpy(x1) - self.x2_per_chrom = [torch.from_numpy(c) for c in x2_per_chrom] - self.y1 = torch.from_numpy(y1) - self.y2 = torch.from_numpy(y2) + # Store the RNA/ATAC matrices sparse (CSR); rows are densified on demand in + # __getitem__ so the full dense matrix is never materialized. + self.x1 = _as_csr(x1) + self.x2_per_chrom = [_as_csr(c) for c in x2_per_chrom] + self.y1 = _as_csr(y1) + self.y2 = _as_csr(y2) self.size_factors1 = torch.from_numpy(size_factors1) def __len__(self): @@ -75,11 +87,11 @@ def __len__(self): def __getitem__(self, idx): X = { - "x1": self.x1[idx], - "x2_per_chrom": [c[idx] for c in self.x2_per_chrom], + "x1": _row_tensor(self.x1, idx), + "x2_per_chrom": [_row_tensor(c, idx) for c in self.x2_per_chrom], "size_factors1": self.size_factors1[idx], } - y = torch.cat([self.y1[idx], self.y2[idx]]) # dummy combined target, unused directly + y = torch.cat([_row_tensor(self.y1, idx), _row_tensor(self.y2, idx)]) # dummy combined target, unused directly return X, y @@ -152,7 +164,10 @@ def get_loss(self, y_pred, y_true, X=None, training=False): X_atac_bin = _atac_binarized(adata_atac) chrom_counts, chrom_groups = parse_chrom_groups(adata_atac.var_names) -X_atac_per_chrom = [X_atac_bin[:, idxs] for idxs in chrom_groups.values()] +# Slice peaks per chromosome on a CSC view (fast column indexing); each group stays +# sparse and is densified per-cell later. +X_atac_bin_csc = X_atac_bin.tocsc() +X_atac_per_chrom = [X_atac_bin_csc[:, idxs] for idxs in chrom_groups.values()] n_genes = X_rna.shape[1] n_peaks = X_atac_bin.shape[1] @@ -206,7 +221,7 @@ def get_loss(self, y_pred, y_true, X=None, training=False): "direction": direction, "rna_var_names": list(adata_rna.var_names), "atac_var_names": list(adata_atac.var_names), - "size_factor_median": float(np.median(Y_rna_counts.sum(axis=1))), + "size_factor_median": float(np.median(np.asarray(Y_rna_counts.sum(axis=1)).ravel())), } with open(par["output"], "wb") as f: