Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
30 changes: 21 additions & 9 deletions src/methods/babel/babel_predict/script.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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...")
Expand Down
49 changes: 32 additions & 17 deletions src/methods/babel/babel_train/script.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -38,48 +38,60 @@
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):
return self.x1.shape[0]

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


Expand Down Expand Up @@ -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]
Expand Down Expand Up @@ -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:
Expand Down
Loading