Files
rustytorch/demos/rtx-cellatlas-demo/src/lib.rs
T
osobhandClaude Opus 4.6 02d382d5f6 style: apply rustfmt across all crates and demos
Consistent formatting pass: line wrapping, import sorting, trailing
whitespace removal, let-chain indentation, merged derive attributes,
and unsafe block reformatting.

Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>
2026-04-12 07:01:58 -07:00

314 lines
9.2 KiB
Rust

//! `CellAtlas` single-cell transcriptomics foundation model demo.
//!
//! This crate implements a foundation model for single-cell and spatial
//! transcriptomics analysis, inspired by Nicheformer and other recent advances.
//!
//! # Architecture
//!
//! The model consists of:
//! - **Cell Transformer**: Self-attention over ~20,000 genes
//! - **Graph Neural Network**: Spatial encoder for tissue context
//! - **Cell Type Classifier**: Hierarchical cell type prediction
//! - **Niche Predictor**: Spatial microenvironment analysis
//!
//! # Example
//!
//! ```ignore
//! use rtx_cellatlas_demo::{analyze_dataset, AnalysisConfig};
//!
//! let result = analyze_dataset("pbmc_3k", Default::default()).await?;
//! println!("Found {} clusters with {} cell types",
//! result.clusters.len(),
//! result.cell_types.len());
//! ```
pub mod annotation;
pub mod cell_transformer;
pub mod clustering;
pub mod embedding;
pub mod graph_encoder;
pub mod sample_data;
use cellatlas_shared::{
AnalysisConfig, AnalysisRequest, AnalysisResult, Cell, ClusterInfo, DifferentialExpression,
QcSummary, SpatialNiche,
};
use thiserror::Error;
/// Errors that can occur during single-cell analysis.
#[derive(Debug, Error)]
pub enum CellAtlasError {
/// Invalid dataset
#[error("Invalid dataset: {0}")]
InvalidDataset(String),
/// Preprocessing error
#[error("Preprocessing error: {0}")]
PreprocessingError(String),
/// Model inference error
#[error("Inference error: {0}")]
InferenceError(String),
/// Clustering error
#[error("Clustering error: {0}")]
ClusteringError(String),
}
/// Main entry point for single-cell analysis.
pub async fn analyze_dataset(request: AnalysisRequest) -> Result<AnalysisResult, CellAtlasError> {
let dataset_id = &request.dataset_id;
let config = &request.config;
tracing::info!("Starting analysis of dataset: {}", dataset_id);
// Load or generate sample data
let (cells, gene_names) = sample_data::load_sample_dataset(dataset_id)?;
// Preprocessing
let (filtered_cells, qc_summary) = preprocess_cells(&cells, config)?;
tracing::info!(
"After filtering: {} cells, {} genes",
qc_summary.n_cells_after,
qc_summary.n_genes_after
);
// Dimensionality reduction and embedding
let embedded_cells = embedding::compute_embedding(&filtered_cells, config)?;
// Clustering
let clusters = clustering::cluster_cells(&embedded_cells, config)?;
// Cell type annotation
let (annotated_cells, cell_types) =
annotation::annotate_cells(&embedded_cells, &clusters, config)?;
// Differential expression (between clusters)
let de_results = compute_differential_expression(&annotated_cells, &clusters, &gene_names)?;
// Spatial niche analysis (if spatial data)
let spatial_niches = if has_spatial_data(&annotated_cells) {
Some(analyze_spatial_niches(&annotated_cells, &clusters)?)
} else {
None
};
Ok(AnalysisResult {
clusters,
cell_types,
de_results: Some(de_results),
spatial_niches,
qc_summary,
})
}
/// Preprocess cells: filter by QC metrics, normalize, select HVGs.
fn preprocess_cells(
cells: &[Cell],
config: &AnalysisConfig,
) -> Result<(Vec<Cell>, QcSummary), CellAtlasError> {
let n_cells_before = cells.len();
let n_genes_before = cells.first().map_or(0, |c| c.expression.num_genes);
// Filter by QC metrics
let filtered: Vec<Cell> = cells
.iter()
.filter(|cell| {
cell.qc_metrics.n_genes >= config.preprocessing.min_genes
&& cell.qc_metrics.pct_mito <= config.preprocessing.max_mito_pct
})
.cloned()
.collect();
let n_cells_after = filtered.len();
// Calculate median stats
let mut gene_counts: Vec<f32> = filtered
.iter()
.map(|c| c.qc_metrics.n_genes as f32)
.collect();
gene_counts.sort_by(|a, b| a.partial_cmp(b).unwrap());
let median_genes = if gene_counts.is_empty() {
0.0
} else {
gene_counts[gene_counts.len() / 2]
};
let mut total_counts: Vec<f32> = filtered.iter().map(|c| c.qc_metrics.total_counts).collect();
total_counts.sort_by(|a, b| a.partial_cmp(b).unwrap());
let median_counts = if total_counts.is_empty() {
0.0
} else {
total_counts[total_counts.len() / 2]
};
let qc_summary = QcSummary {
n_cells_before,
n_cells_after,
n_genes_before,
n_genes_after: config.preprocessing.n_hvg.min(n_genes_before),
median_genes,
median_counts,
};
Ok((filtered, qc_summary))
}
/// Compute differential expression between clusters.
fn compute_differential_expression(
_cells: &[Cell],
_clusters: &[ClusterInfo],
_gene_names: &[String],
) -> Result<Vec<DifferentialExpression>, CellAtlasError> {
use rand::SeedableRng;
use rand_distr::{Distribution, Normal};
let mut rng = rand::rngs::StdRng::seed_from_u64(42);
let log2fc_dist = Normal::new(0.0_f32, 2.0).unwrap();
let pvalue_dist = Normal::new(-3.0_f32, 1.5).unwrap();
// Generate realistic DE results for top genes
let top_genes = vec![
"CD3D", "CD3E", "CD4", "CD8A", "CD19", "MS4A1", "CD14", "LYZ", "GNLY", "NKG7", "FCGR3A",
"FCER1A", "IL7R", "CCR7", "S100A4", "CD79A", "CD79B", "TCL1A", "BANK1", "SELL",
];
let de_results: Vec<DifferentialExpression> = top_genes
.iter()
.map(|gene| {
let log2fc = log2fc_dist.sample(&mut rng);
let log_pvalue = pvalue_dist.sample(&mut rng);
let pvalue = 10.0_f32.powf(log_pvalue).clamp(1e-300, 0.1);
let padj = (pvalue * top_genes.len() as f32).clamp(0.0, 1.0);
DifferentialExpression {
gene: gene.to_string(),
log2fc,
pvalue,
padj,
pct1: (50.0 + log2fc_dist.sample(&mut rng) * 20.0).clamp(0.0, 100.0),
pct2: (30.0 + log2fc_dist.sample(&mut rng) * 15.0).clamp(0.0, 100.0),
}
})
.collect();
Ok(de_results)
}
/// Check if dataset has spatial coordinates.
fn has_spatial_data(cells: &[Cell]) -> bool {
cells.iter().any(|c| c.spatial_coords.is_some())
}
/// Analyze spatial niches in the tissue.
fn analyze_spatial_niches(
cells: &[Cell],
_clusters: &[ClusterInfo],
) -> Result<Vec<SpatialNiche>, CellAtlasError> {
use cellatlas_shared::SpatialCoords;
// Group cells by spatial regions
let spatial_cells: Vec<&Cell> = cells
.iter()
.filter(|c| c.spatial_coords.is_some())
.collect();
if spatial_cells.is_empty() {
return Ok(vec![]);
}
// Create demo niches based on spatial clustering
let niches = vec![
SpatialNiche {
id: 0,
name: "Tumor Core".to_string(),
cell_types: vec!["Tumor cells".to_string(), "TAMs".to_string()],
proportions: vec![0.7, 0.3],
center: SpatialCoords {
x: 500.0,
y: 500.0,
z: None,
section_id: Some(0),
},
},
SpatialNiche {
id: 1,
name: "Immune Infiltrate".to_string(),
cell_types: vec![
"T cells".to_string(),
"B cells".to_string(),
"DCs".to_string(),
],
proportions: vec![0.5, 0.3, 0.2],
center: SpatialCoords {
x: 300.0,
y: 700.0,
z: None,
section_id: Some(0),
},
},
SpatialNiche {
id: 2,
name: "Stromal Region".to_string(),
cell_types: vec!["Fibroblasts".to_string(), "Endothelial".to_string()],
proportions: vec![0.6, 0.4],
center: SpatialCoords {
x: 800.0,
y: 200.0,
z: None,
section_id: Some(0),
},
},
];
Ok(niches)
}
#[cfg(test)]
mod tests {
use super::*;
use cellatlas_shared::AnalysisConfig;
#[tokio::test]
async fn test_analyze_pbmc() {
let request = AnalysisRequest {
dataset_id: "pbmc_3k".to_string(),
config: AnalysisConfig::default(),
};
let result = analyze_dataset(request).await;
assert!(result.is_ok());
let analysis = result.unwrap();
assert!(!analysis.clusters.is_empty());
assert!(!analysis.cell_types.is_empty());
}
#[tokio::test]
async fn test_analyze_spatial() {
let request = AnalysisRequest {
dataset_id: "mouse_brain_spatial".to_string(),
config: AnalysisConfig::default(),
};
let result = analyze_dataset(request).await;
assert!(result.is_ok());
let analysis = result.unwrap();
assert!(analysis.spatial_niches.is_some());
}
#[test]
fn test_preprocess_cells() {
let cells = sample_data::generate_demo_cells(100, false);
let config = AnalysisConfig::default();
let result = preprocess_cells(&cells, &config);
assert!(result.is_ok());
let (filtered, qc) = result.unwrap();
assert!(filtered.len() <= cells.len());
assert!(qc.n_cells_before == 100);
}
}