GID-Flow / PDGrapher /scripts /preprocess_sciplex3.py
Boom5426's picture
Upload GID-Flow project snapshot (deduped: code + key artifacts)
07fcdfe verified
Raw
History Blame Contribute Delete
6.56 kB
"""Preprocess SciPlex3 data for training.
Memory-efficient: uses backed mode for initial load, filters metadata first,
keeps sparse matrix through HVG selection, only extracts dense for final 2000 HVGs.
Usage:
python scripts/preprocess_sciplex3.py \
--h5ad /data/boom/Protein/regulatory_field/data/raw/scPerturb/rna_protein/SrivatsanTrapnell2020_sciplex3.h5ad \
--output data/processed/sciplex3_k562_24h.pt \
--cell_lines K562 \
--doses 10 100 1000 10000 \
--times 24 \
--n_hvg 2000
"""
import argparse
import gc
import os
import sys
import numpy as np
import torch
def preprocess(args):
"""Preprocess SciPlex3 and save compact representation."""
try:
import anndata as ad
import scipy.sparse as sp
except ImportError:
print("Need anndata and scipy")
sys.exit(1)
import scanpy as sc
print(f"Loading {args.h5ad}...")
# Use backed mode to avoid loading full matrix into memory
# We only need obs (metadata) for filtering first
try:
adata = ad.read_h5ad(args.h5ad, backed='r')
print(f" Loaded in backed mode: {adata.shape}")
# Materialize to memory (backed mode → in-memory)
adata = adata.to_memory()
print(f" Materialized to memory")
except Exception as e:
print(f" Backed mode failed ({e}), loading directly")
adata = ad.read_h5ad(args.h5ad)
print(f" Raw shape: {adata.shape}")
# Filter by metadata BEFORE any dense operations
obs = adata.obs.copy()
if args.cell_lines:
obs = obs[obs['cell_line'].isin(args.cell_lines)]
if args.doses:
obs = obs[obs['dose_value'].isin(args.doses)]
elif not args.include_vehicle:
obs = obs[obs['dose_value'] > 0]
if args.times:
obs = obs[obs['time'].isin(args.times)]
keep_idx = obs.index.values
n_keep = len(keep_idx)
print(f" After filtering: {n_keep} cells (from {adata.shape[0]})")
# Check if X is sparse before subsetting
is_sparse = sp.issparse(adata.X)
if is_sparse:
print(f" Source matrix is sparse ({adata.X.nnz / 1e6:.1f}M non-zeros)")
# Subset to filtered cells
adata = adata[keep_idx].copy()
print(f" Subset shape: {adata.shape}")
# Check memory usage after subset
if sp.issparse(adata.X):
print(f" Sparse matrix: {adata.X.nnz / 1e6:.1f}M non-zeros, "
f"{adata.X.data.nbytes / 1e6:.1f}MB data")
else:
print(f" Dense matrix: {adata.X.nbytes / 1e9:.2f} GB — WARNING: this will use lots of RAM!")
# Normalize + log1p (operates on sparse matrix)
print(" Normalizing...")
sc.pp.normalize_total(adata, target_sum=1e4)
sc.pp.log1p(adata)
# HVG selection on SPARSE matrix — memory efficient
print(f" Selecting {args.n_hvg} HVGs (sparse mode)...")
try:
# Try scanpy's HVG (works on sparse, uses seurat_v3)
sc.pp.highly_variable_genes(adata, n_top_genes=args.n_hvg, flavor="seurat_v3")
hvg_mask = adata.var['highly_variable'].values
hvg_idx = np.where(hvg_mask)[0]
print(f" Scanpy HVG selected: {len(hvg_idx)} genes")
except Exception as e:
print(f" Scanpy HVG failed ({e}), using variance-based sampling")
# For variance-based: sample cells to estimate variance
n_sample = min(10000, adata.shape[0])
sample_idx = np.random.choice(adata.shape[0], n_sample, replace=False)
if sp.issparse(adata.X):
X_sample = adata.X[sample_idx].toarray()
else:
X_sample = adata.X[sample_idx]
gene_var = np.var(X_sample, axis=0)
del X_sample
gc.collect()
top_idx = np.argsort(gene_var)[-args.n_hvg:]
hvg_idx = np.sort(top_idx)
print(f" Variance-based HVG: {len(hvg_idx)} genes")
# Extract ONLY the HVG columns as dense (2000 genes × N cells)
print(f" Extracting HVG expression matrix...")
n_cells = adata.shape[0]
n_hvg = len(hvg_idx)
print(f" Extracting {n_hvg} genes × {n_cells} cells...")
if sp.issparse(adata.X):
X_hvg = adata.X[:, hvg_idx].toarray().astype(np.float32)
else:
X_hvg = np.array(adata.X[:, hvg_idx], dtype=np.float32)
gene_names = list(adata.var.index[hvg_idx])
# Fix: use ensembl_id if var.index is non-standard (e.g. "nan:X")
if hasattr(adata.var, 'columns') and 'ensembl_id' in adata.var.columns:
ensembl_ids = adata.var['ensembl_id'].values.copy()
# Skip header row if present (CSV parsing artifact)
if len(ensembl_ids) > 0 and str(ensembl_ids[0]) == 'id gene_short_name':
ensembl_ids = ensembl_ids[1:]
# Pad to original length if header was removed
if len(ensembl_ids) < len(adata.var):
ensembl_ids = np.array(list(ensembl_ids) + ['unknown'] * (len(adata.var) - len(ensembl_ids)))
gene_names = [str(ensembl_ids[i]) for i in hvg_idx]
print(f" HVG matrix: {X_hvg.shape}, {X_hvg.nbytes / 1e9:.2f} GB")
del adata # free memory
gc.collect()
# Extract metadata (obs was already filtered)
obs['_cell_idx'] = np.arange(len(obs))
# Save
os.makedirs(os.path.dirname(args.output) or '.', exist_ok=True)
save_dict = {
'X': X_hvg,
'gene_names': gene_names,
'obs': obs,
'cell_lines': args.cell_lines,
'doses': args.doses,
'times': args.times,
}
torch.save(save_dict, args.output)
print(f" Saved to {args.output}")
file_size_gb = os.path.getsize(args.output) / 1e9
print(f" File size: {file_size_gb:.2f} GB")
# Summary
print(f"\nSummary:")
print(f" Cells: {len(obs)}")
print(f" Genes: {len(gene_names)}")
print(f" Cell lines: {obs['cell_line'].value_counts().to_dict()}")
print(f" Doses: {sorted(obs['dose_value'].unique())}")
print(f" Drugs: {obs['perturbation'].nunique()}")
print(f" Perturbations: {obs['perturbation'].value_counts().head(10).to_dict()}")
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--h5ad", required=True)
parser.add_argument("--output", required=True)
parser.add_argument("--cell_lines", nargs="*")
parser.add_argument("--doses", nargs="*", type=float)
parser.add_argument("--times", nargs="*", type=float)
parser.add_argument("--n_hvg", type=int, default=2000)
parser.add_argument("--include_vehicle", action="store_true")
args = parser.parse_args()
preprocess(args)
if __name__ == "__main__":
main()