CNN-based regression for cell type deconvolution of bulk RNA-seq using scRNA-seq reference
Project description
CNNreg: CNN-based Cell Type Deconvolution
A deep learning approach for cell type deconvolution of bulk RNA-seq data using single-cell RNA-seq reference data. CNNreg employs a custom CNN-based regression model to estimate cell type proportions in complex tissue samples.
We recommend using CNNregR package in R for easier data handling and preprocessing. Alternatively, users can also preprocess the data by yourself and use the Python API or CLI provided in this package.
Installation
Option 1: Using pip (Recommended)
CNNreg requires PyTorch. Install PyTorch first according to your system configuration:
# For CUDA 12.6 (check your CUDA version: nvidia-smi)
pip3 install torch --index-url https://download.pytorch.org/whl/cu126
# For CUDA 11.8
pip3 install torch --index-url https://download.pytorch.org/whl/cu118
# For CPU only
pip3 install torch --index-url https://download.pytorch.org/whl/cpu
Then install CNNreg:
pip install CNNreg
Option 2: From Source (Development)
git clone https://github.com/mwang159/CNNreg.git
cd CNNreg
pip install -e .
Option 3: Using Conda Environment
# Create environment
conda create -n cnnreg_env python=3.10
conda activate cnnreg_env
# Install PyTorch (choose appropriate CUDA version)
pip3 install torch --index-url https://download.pytorch.org/whl/cu126
# Install CNNreg
pip install CNNreg
Verify Installation
# Check CLI command
cnnreg --help
# Test in Python
python -c "import CNNreg; print('CNNreg installed successfully!')"
Quick Start
Command Line Interface
cnnreg \
-bulk data/bulk.csv \
-ref data/sc_ref.csv \
-o output/ \
-C 7 \
-EP 50000 \
-pre GBM_analysis
Parameters:
-bulk: Path to bulk RNA-seq CSV file-ref: Path to reference cell type specific expression (CSE) CSV-o: Output directory-C: Number of cell types (kernel size)-EP: Maximum training epochs-pre: Output file prefix
Note: Training automatically generates cell proportion predictions for the input bulk samples. Predictions are saved at checkpoints (every 10,000 epochs) and at completion (or early stop). Both raw and row-sum-normalized proportions are saved.
Python API
import torch
import pandas as pd
from CNNreg.data import data_CSE, reformat_ref
from CNNreg.train import trainProp
from CNNreg.layers import DeconvProp
# Setup
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
# Load data
df_bulk = pd.read_csv("bulk.csv")
bulk_data = torch.tensor(
df_bulk.iloc[:, 1:].values.transpose(),
dtype=torch.float32
).to(device)
# Configure parameters
pHash = {
"bulk": "bulk.csv",
"reference": "sc_ref.csv",
"data_out_dir": "output/",
"max_epoch_cellprop": 50000,
"prefix": "GBM",
"device": device,
"n_kernel": 7,
"n_gene": df_bulk.shape[0],
"n_sample": df_bulk.shape[1] - 1,
"n_celltype": 7
}
# Initialize data dictionary
dHash = {
"bulk": bulk_data,
"sample": df_bulk.columns.values[1:],
"celltype": ["AClike", "MESlike", "NPClike", "OPClike", "OL", "Myeloid", "T"],
"CSE": data_CSE(pHash["reference"], device=device)
}
# Add required indices
k = pHash["n_gene"] * pHash["n_celltype"]
dHash["idx_feature_celltype"] = [
[int(y+x) for y in range(0, k, pHash["n_celltype"])]
for x in range(pHash["n_celltype"])
]
dHash["CSE_reformat"] = reformat_ref(dHash["CSE"].expr_cse)
# Train model
trainProp(dHash, pHash)
Input Data Format
Bulk RNA-seq Data
Rows are samples, columns are genes:
Sample,Gene1,Gene2,Gene3,...
Sample1,0.5,1.2,0.8,...
Sample2,1.1,0.9,1.5,...
Sample3,0.7,1.3,0.6,...
Reference scRNA-seq Data
Cell Type Specific Expression profiles from scRNA-seq:
CellType,Gene1,Gene2,Gene3,...
AClike_ref1,0.3,0.8,0.5,...
AClike_ref2,0.4,0.7,0.6,...
MESlike_ref1,0.9,0.2,0.4,...
Output Files
Estimation_{prefix}_epoch_{N}.csv: Raw (unnormalized) cell proportion weights at checkpoint epochs (every 10,000) and at early stop / completionEstimation_{prefix}_epoch_{N}_normalized.csv: Row-sum-normalized cell proportions (sum to 1 per sample) at the same checkpoints
Output format:
Sample,celltype_1,celltype_2,celltype_3,...
Sample1,0.23,0.15,0.31,...
Sample2,0.19,0.28,0.22,...
Model Architecture
CNNreg uses a custom 4-layer CNN pipeline specifically designed for biological deconvolution:
- RefCombLayer: Combines multiple reference samples per cell type via learned weighted combination
- SliceSumLayer: Aggregates weighted reference combinations into a single profile per cell type
- CelltypeScaleLayer: Scales the expression profile of each cell type independently (constrained to [0.25, 4])
- Conv1D Layer: Estimates cell proportions via 1D convolution (kernel size = number of cell types)
Training Strategy
Training uses a 3-phase manual SGD cycle:
- Phase 0 (
epoch % 3 == 0): Tuneconv1kernel weights (cell proportions) using Pearson correlation + expression fit loss - Phase 1 (
epoch % 3 == 1, only forepoch < 30,000): TunerefLayerweights (reference combination) - Phase 2 (
epoch % 3 == 2): TunecelltypeScaleLayerweights (per-cell-type scaling)
Early stopping activates after epoch 20,000: training halts when correlation loss begins increasing and the model has already achieved Pearson r > 0.75 on high-variability genes.
Citation
If you use CNNreg in your research, please cite:
License
This project is licensed under the MIT License - see the LICENSE file for details.
Links
- Python Package: CNNreg on PyPI
- CNNreg: CNNreg on GitHub
- CNNregR: CNNregR on GitHub
- Issues: Report bugs
Contact
Project details
Release history Release notifications | RSS feed
Download files
Download the file for your platform. If you're not sure which to choose, learn more about installing packages.
Source Distributions
Built Distribution
Filter files by name, interpreter, ABI, and platform.
If you're not sure about the file name format, learn more about wheel file names.
Copy a direct link to the current filters
File details
Details for the file cnnreg-0.1.1-py3-none-any.whl.
File metadata
- Download URL: cnnreg-0.1.1-py3-none-any.whl
- Upload date:
- Size: 11.2 kB
- Tags: Python 3
- Uploaded using Trusted Publishing? No
- Uploaded via: twine/6.2.0 CPython/3.10.18
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
aa19591df7139170c61fd3fe022153d41e1d02d6aa93803cd12021d5657e857e
|
|
| MD5 |
0561fcccc1dc7a08d4e5455a78e6d507
|
|
| BLAKE2b-256 |
f4cd493fc615d2beeadf590b1171be2b97c0bf215c0f3d4a044d30e7887d756e
|