Tissueformer: extending single-cell foundation models to predict population-level phenotypes.
The 13 matches
- [1] § Methods › Pseudobulk and type composition benchmarks ↔ applications/brain_annotation/benchmarks.py, lines 1–52 · score 0.88 · random forest classifier, max depth, logistic regression, scikit-learn, lbfgs, leaf
- [2] § Methods › Deep learning benchmarks ↔ applications/brain_annotation/benchmarks.py, lines 740–824 · score 0.84 · weight decay, ScRAT, cosine, schedule, warmup, dropout
- [3] § Methods › Data and tokenization › Data description: mouse brains ↔ colormycells/colormap.py, lines 48–141 · score 0.80 · perceptual distance, perceptually uniform, pseudobulk expression, ColorMyCells, LUV, MDS
- [4] § Methods › Brain map visualization ↔ applications/brain_annotation/paper_figures/fig3/compute_visp_area.py, lines 443–531 · score 0.74 · flatmap pixel, cuML, smooth, GPU, kernel, SVC
- [5] § Methods › Data and tokenization › Data description: mouse brains ↔ applications/brain_annotation/figures/utils.py, lines 104–159 · score 0.69 · perceptual distance, perceptually uniform, LUV, MDS, matrix, colormaps
- [6] § Methods › TissueFormer architecture: hyperparameters of this study ↔ applications/brain_annotation/benchmarks.py, lines 740–824 · score 0.67 · learning rate warmup, PyTorch, decay, hidden, optimizer, Adam
- [7] § Results › Annotating cell groups with TissueFormer ↔ applications/brain_annotation/benchmarks.py, lines 827–889 · score 0.64 · Logistic regression, deep learning, random forest, bulk, CellCnn, ScRAT
- [8] § Methods › Deep learning benchmarks ↔ applications/brain_annotation/paper_figures/fig2/hyperparameters.ipynb, lines 139–170 · score 0.63 · CellCNN, ScAGG, ScRAT, Hyperparameters, benchmarked
- [9] § Methods › Data and tokenization › Data description: mouse brains ↔ applications/brain_annotation/paper_figures/fig3/cell_type_viz.ipynb, lines 518–657 · score 0.61 · brain slice, AnnData, ccf streamlines, axis
- [10] § Results › Single-cell profiles are not informative of cortical area ↔ applications/brain_annotation/paper_figures/fig2/single_cell_accuracy.py, lines 6–25 · score 0.61 · Murine Geneformer, logistic regression, random forest, neighbors, single cell, classifier
- [11] § Results › Predicting COVID-19 severity from blood transcriptomics ↔ applications/brain_annotation/paper_figures/fig2/hyperparameters.ipynb, lines 139–170 · score 0.56 · CellCNN, ScAGG, ScRAT, accuracy
- [12] § Results › Single-cell profiles are not informative of cortical area ↔ applications/brain_annotation/paper_figures/fig2/hyperparameters.ipynb, lines 172–212 · score 0.53 · logistic regression, brain accuracy, Validation accuracy, LR, bars, Figure 2
- [13] § Results › Predicted cortical maps ↔ applications/brain_annotation/paper_figures/fig3/gene_viz.ipynb, lines 489–498 · score 0.53 · Brinp3, Coro6, Nell1, Rcan2, Tshz2, Zfpm2
Paper
Loaded from Europe PMC by your browser, not stored by OSCR: doi.org · Europe PMC
The paper is loaded when this pane is shown.
The authors' code
Python · 892 lines · 32 KB · MIT · 4 matches
- """
- Benchmarking Configuration
- =========================
- This configuration file controls the benchmarking pipeline for brain cell type classification models.
- It defines parameters for Random Forest and Logistic Regression classifiers, along with data handling
- and experiment tracking settings.
- Dependencies
- -----------
- - hydra-core
- - wandb (Weights & Biases)
- - scikit-learn
- - numpy
- - anndata (when using AnnData features)
- Configuration Sections
- --------------------
- random_forest:
- Parameters for sklearn RandomForestClassifier
- - n_estimators: Number of trees (default: 200)
- - max_depth: Maximum tree depth (default: 15)
- - min_samples_split: Minimum samples for split (default: 5)
- - min_samples_leaf: Minimum samples per leaf (default: 2)
- - max_features: Features per tree (default: 0.33)
- - bootstrap: Bootstrap samples (default: true)
- - class_weight: Class weighting strategy (default: null)
- logistic_regression:
- Parameters for sklearn LogisticRegression
- - max_iter: Maximum iterations (default: 1000)
- - multi_class: Classification strategy (default: 'multinomial')
- - solver: Optimization algorithm (default: 'lbfgs')
- Debug Options
- ------------
- debug: false # Enable debug mode with reduced dataset
- debug_args:
- on_adata: false # Use AnnData features
- resample_adata: true # Resample AnnData to match dataset size
- Experiment Types
- --------------
- run_bulk_expression_rf: false # Run Random Forest on bulk expression
- run_bulk_expression_lr: true # Run Logistic Regression on bulk expression
- run_h3type_rf: false # Run Random Forest on H3 types
- run_h3type_lr: false # Run Logistic Regression on H3 types
- Usage Examples
- ------------
- # Run default benchmarks
- python benchmarks.py
- # Enable debug mode
- python benchmarks.py debug=true
- # Run specific experiment
- python benchmarks.py run_bulk_expression_rf=true
- # Override Random Forest parameters
- python benchmarks.py random_forest.n_estimators=300 random_forest.max_depth=20
- Output
- ------
- Results are saved to ${output_dir} (default: 'benchmarks/') including:
- - Model performance metrics
- - Feature importance plots
- - Confusion matrices
- - Weights & Biases logs (if enabled)
- Notes
- -----
- - Set debug=true for quick iteration with reduced dataset
- - Enable wandb tracking by configuring wandb/default.yaml
- - Use debug_args.on_adata=true when working with AnnData features
- """
- import os
- import hydra
- import wandb
- import numpy as np
- import anndata as ad
- from omegaconf import DictConfig, OmegaConf
- from sklearn.ensemble import RandomForestClassifier
- from sklearn.linear_model import LogisticRegression, LogisticRegressionCV
- from sklearn.preprocessing import StandardScaler
- from collections import Counter
- from typing import Dict, List, Tuple, Optional
- from datasets import load_from_disk, DatasetDict, Dataset
- import torch
- from torch.utils.data import DataLoader
- from sklearn.model_selection import train_test_split
- from transformers import PreTrainedModel, TrainingArguments
- import json
- import scipy
- from sklearn.utils import check_random_state
- from sklearn.preprocessing import OneHotEncoder
- # Import necessary components from train.py
- from tissueformer.samplers import (
- GroupedSpatialTrainer
- )
- from transformers import PretrainedConfig
- class DummyConfig(PretrainedConfig):
- def __init__(self, **kwargs):
- super().__init__(**kwargs)
- class DummyModel(PreTrainedModel):
- config_class = DummyConfig
- def __init__(self, config):
- super().__init__(config)
- def forward(self, *args, **kwargs):
- return None
- def get_dataloaders(datasets: DatasetDict, cfg: DictConfig) -> Dict[str, DataLoader]:
- """
- Create dataloaders using GroupedSpatialTrainer infrastructure.
- """
- # Create minimal training arguments
- training_args = TrainingArguments(
- output_dir=cfg.output_dir,
- per_device_train_batch_size=cfg.data.group_size*32,
- per_device_eval_batch_size=cfg.data.group_size*32,
- remove_unused_columns=False, # Important for GroupedSpatialTrainer
- )
- dummy_config = DummyConfig()
- # Initialize trainer with dummy model
- trainer = GroupedSpatialTrainer(
- model=DummyModel(dummy_config),
- args=training_args,
- train_dataset=datasets["train"],
- eval_dataset=datasets["validation"],
- spatial_group_size=cfg.data.group_size,
- spatial_label_key="labels",
- coordinate_key='CCF_streamlines',
- additional_feature_keys=['raw_counts'],
- sampling_strategy=cfg.data.sampling.strategy,
- hex_scaling=cfg.data.sampling.hex_scaling,
- reflect_points=cfg.data.sampling.reflect_points,
- use_train_hex_grid_on_eval=cfg.data.sampling.use_train_hex_grid_on_eval,
- max_radius_expansions=cfg.data.sampling.max_radius_expansions,
- group_within_keys=cfg.data.sampling.group_within_keys
- )
- # Get dataloaders
- dataloaders = {
- "train": trainer.get_train_dataloader(),
- "validation": trainer.get_eval_dataloader(datasets["validation"]),
- "test": trainer.get_test_dataloader(datasets["test"])
- }
- return dataloaders
- def load_and_align_anndata(
- train_filenames: List[str],
- test_filenames: List[str],
- data_dir: str,
- dataset: DatasetDict,
- coordinate_key = "CCF_streamlines",
- ) -> Tuple[ad.AnnData, ad.AnnData]:
- """
- Load and concatenate AnnData files, ensuring alignment with dataset indices.
- Handles train and test files separately to match the original tokenization.
- Args:
- train_filenames: List of h5ad filenames for training data
- test_filenames: List of h5ad filenames for test data
- data_dir: Directory containing h5ad files
- dataset: HuggingFace dataset to align with
- Returns:
- Tuple of (train_adata, test_adata)
- """
- print("Loading and processing AnnData files...")
- def process_files(filenames: List[str]) -> ad.AnnData:
- """Helper function to process a list of files"""
- adatas = []
- for filename in filenames:
- filepath = os.path.join(data_dir, filename)
- print(f"Loading {filepath}")
- adata = ad.read_h5ad(filepath)
- # Filter cells with invalid CCF coordinates
- valid_mask = ~np.isnan(adata.obsm[coordinate_key]).any(axis=1)
- adata = adata[valid_mask]
- adatas.append(adata)
- return ad.concat(adatas, join='outer', fill_value=0)
- # Process train and test files separately
- train_adata = process_files(train_filenames)
- test_adata = process_files(test_filenames)
- # Verify alignment with dataset
- print("Verifying alignment with dataset...")
- # Check first 10000 cells in train dataset
- test_dataset = dataset['test']
- dataset_h3types = np.array(test_dataset[:10000]['H3_type'])
- adata_h3types = test_adata.obs['H3_type'].values[:10000]
- if not np.array_equal(dataset_h3types, adata_h3types):
- mismatches = np.where(dataset_h3types != adata_h3types)[0]
- mismatch_info = [
- f"Index {i}: Dataset H3_type: {dataset_h3types[i]}, AnnData H3_type: {adata_h3types[i]}"
- for i in mismatches
- ]
- raise ValueError(
- "Mismatch found in H3 types:\n" + "\n".join(mismatch_info)
- )
- print("Alignment verification passed!")
- return train_adata, test_adata
- def prepare_datasets(dataset_dict: DatasetDict, cfg: DictConfig) -> DatasetDict:
- """
- Prepare train/validation split from dataset.
- """
- train_dataset = dataset_dict["train"]
- # Create train/validation split
- train_idx, val_idx = train_test_split(
- np.arange(len(train_dataset)),
- test_size=cfg.data.validation_split,
- random_state=cfg.seed
- )
- # Add unique ids if not present
- if 'uuid' not in train_dataset.features:
- dataset_dict["test"] = dataset_dict["test"].add_column(
- "uuid",
- np.arange(len(dataset_dict["test"]))
- )
- train_dataset = train_dataset.add_column(
- "uuid",
- np.arange(len(dataset_dict["train"]))
- )
- # Select train and validation datasets
- val_dataset = train_dataset.select(val_idx)
- train_dataset = train_dataset.select(train_idx)
- # Limit dataset size if in debug mode
- if hasattr(cfg.data, 'max_train_samples') and cfg.data.max_train_samples is not None:
- train_dataset = train_dataset.select(range(min(len(train_dataset), cfg.data.max_train_samples)))
- # Limit validation/test size if specified
- if hasattr(cfg.data, 'max_eval_samples') and cfg.data.max_eval_samples is not None:
- val_dataset = val_dataset.select(range(min(len(val_dataset), cfg.data.max_eval_samples)))
- dataset_dict["test"] = dataset_dict["test"].select(range(min(len(dataset_dict["test"]), cfg.data.max_eval_samples)))
- # rename the labels to match the model's expected input
- if hasattr(cfg.data, 'label_key'):
- train_dataset = train_dataset.rename_column(cfg.data.label_key, "labels")
- val_dataset = val_dataset.rename_column(cfg.data.label_key, "labels")
- dataset_dict["test"] = dataset_dict["test"].rename_column(cfg.data.label_key, "labels")
- return DatasetDict({
- "train": train_dataset,
- "validation": val_dataset,
- "test": dataset_dict["test"]
- })
- def setup_wandb(cfg: DictConfig):
- """Initialize W&B logging"""
- wandb.init(
- project=cfg.wandb.project,
- entity=cfg.wandb.entity,
- name=cfg.wandb.name,
- group=cfg.wandb.group,
- tags=cfg.wandb.tags,
- notes=cfg.wandb.notes,
- config=OmegaConf.to_container(cfg, resolve=True),
- )
- def evaluate_method(
- predictions: np.ndarray,
- labels: np.ndarray,
- indices: np.ndarray,
- prefix: str,
- label_names: Dict,
- output_dir: str
- ) -> Dict:
- """Evaluate predictions and save results."""
- # Extract metrics
- from sklearn.metrics import balanced_accuracy_score
- metrics = {
- "accuracy": (predictions == labels).mean(),
- "balanced_accuracy": balanced_accuracy_score(labels, predictions),
- }
- # Save predictions
- output_dict = {
- 'predictions': predictions,
- 'labels': labels,
- 'indices': indices,
- 'label_names': label_names
- }
- np.save(os.path.join(output_dir, f"{prefix}_predictions.npy"), output_dict)
- # Log to wandb
- wandb.log({f"{prefix}_{k}": v for k, v in metrics.items()})
- return metrics
- def prepare_h3type_data(dataset: DatasetDict) -> Tuple[Dict[str, np.ndarray], Dict[str, int], Dict[str, Dict[int, int]]]:
- """
- Prepare H3 type data for fast access during training.
- Creates both type mapping and index mapping.
- """
- # Create type mapping
- all_h3_types = set()
- for split in dataset.values():
- all_h3_types.update(split['H3_type'])
- type_to_idx = {h3_type: idx for idx, h3_type in enumerate(sorted(list(all_h3_types)))}
- # Create numpy arrays and index mappings for fast access
- h3_arrays = {}
- index_maps = {}
- for split, dset in dataset.items():
- # Create mapping from dataset index to array index
- index_maps[split] = {
- idx: i for i, idx in enumerate(dset['uuid'])
- }
- # Create array with H3 type indices
- h3_arrays[split] = np.array([
- type_to_idx[t] for t in dset['H3_type']
- ])
- return h3_arrays, type_to_idx, index_maps
- def get_h3type_histogram(
- indices: np.ndarray,
- h3_array: np.ndarray,
- index_map: Dict[int, int],
- n_types: int
- ) -> np.ndarray:
- """
- Create histogram of H3 types for a group of cells using vectorized operations.
- Args:
- indices: Original dataset indices
- h3_array: Pre-computed array of H3 type indices
- index_map: Mapping from dataset indices to array indices
- n_types: Total number of H3 types
- """
- # Ensure indices is 2D: (n_groups, group_size)
- indices = np.asarray(indices)
- if indices.ndim == 1:
- indices = indices.reshape(1, -1) # One group with multiple cells
- batch_size = indices.shape[0]
- histogram = np.zeros((batch_size, n_types))
- for i, batch_indices in enumerate(indices):
- # Map dataset indices to array indices
- array_indices = [index_map[idx] for idx in batch_indices]
- # Get H3 types using mapped indices
- type_indices = h3_array[array_indices]
- histogram[i] = np.bincount(type_indices, minlength=n_types)
- total = histogram[i].sum()
- if total > 0:
- histogram[i] /= total
- return histogram
- def create_dataset_from_anndata(adata: ad.AnnData, cfg: DictConfig) -> Dataset:
- """
- Create a HuggingFace Dataset directly from AnnData object.
- """
- # Extract features (gene expression)
- features = np.array(adata.X.todense() if scipy.sparse.issparse(adata.X) else adata.X)
- print(f"Features shape: {features.shape}")
- # Extract coordinates
- coordinates = adata.obsm["CCF_streamlines"]
- # Extract area labels using the same logic as in tokenize_cells.py
- with open('data/files/area_ancestor_id_map.json', 'r') as f:
- area_ancestor_id_map = json.load(f)
- with open('data/files/area_name_map.json', 'r') as f:
- area_name_map = json.load(f)
- area_name_map['0'] = 'outside_brain'
- annotation2area_int = {0.0:0}
- for a in area_ancestor_id_map.keys():
- higher_area_id = area_ancestor_id_map[str(int(a))][1] if len(area_ancestor_id_map[str(int(a))])>1 else a
- annotation2area_int[float(a)] = int(higher_area_id)
- unique_areas = np.unique(list(annotation2area_int.values()))
- area_classes = np.arange(len(unique_areas))
- id2id = {float(k):v for (k,v) in zip(unique_areas, area_classes)}
- annotation2area_class = {k: id2id[int(v)] for k,v in annotation2area_int.items()}
- # Convert CCF annotations to area labels
- labels = np.array([annotation2area_class[x] for x in adata.obs['CCFano']])
- # Extract H3 types
- h3types = adata.obs['H3_type'].values
- # # Filter dataset to only include cells for which the CCF_streamlines is not nans
- # same as tokenized_dataset = tokenized_dataset.filter(lambda x: not np.isnan(np.sum(x['CCF_streamlines'])))
- # Filter out indices where CCF_streamlines contains NaN values
- valid_mask = ~np.isnan(coordinates).any(axis=1)
- features = features[valid_mask]
- coordinates = coordinates[valid_mask]
- labels = labels[valid_mask]
- h3types = h3types[valid_mask]
- indices = np.arange(len(adata))[valid_mask]
- # Create dataset
- return Dataset.from_dict({
- 'expression': features,
- 'CCF_streamlines': coordinates,
- 'labels': labels,
- 'H3_type': h3types,
- 'uuid': indices
- })
- def prepare_features_from_anndata(
- train_adata: ad.AnnData,
- test_adata: ad.AnnData,
- cfg: DictConfig,
- scaler: StandardScaler,
- feature_type: str
- ) -> Dict[str, Tuple[np.ndarray, np.ndarray, np.ndarray]]:
- """
- Prepare features directly from AnnData objects.
- Returns dict with keys 'train', 'validation', 'test', each containing (features, labels, indices)
- """
- # Create datasets from AnnData
- train_valid_dataset = create_dataset_from_anndata(train_adata, cfg)
- test_dataset = create_dataset_from_anndata(test_adata, cfg)
- rng = check_random_state(cfg.seed)
- # Split train/validation
- train_idx, val_idx = train_test_split(
- np.arange(len(train_valid_dataset)),
- test_size=cfg.data.validation_split,
- random_state=rng
- )
- # Select splits using Dataset.select()
- train_dataset = train_valid_dataset.select(train_idx)
- val_dataset = train_valid_dataset.select(val_idx)
- if feature_type == "h3type":
- # One-hot encode H3 types
- # Collect all unique H3 types from all splits
- all_h3_types = np.concatenate([
- train_dataset['H3_type'],
- val_dataset['H3_type'],
- test_dataset['H3_type']
- ])
- encoder = OneHotEncoder(sparse_output=False)
- # Fit on all types
- encoder.fit(all_h3_types.reshape(-1, 1))
- # Transform each split
- train_features = encoder.transform(np.array(train_dataset['H3_type']).reshape(-1, 1))
- val_features = encoder.transform(np.array(val_dataset['H3_type']).reshape(-1, 1))
- test_features = encoder.transform(np.array(test_dataset['H3_type']).reshape(-1, 1))
- else:
- # Extract and scale continuous features
- train_features = np.array(train_dataset['expression'])
- val_features = np.array(val_dataset['expression'])
- test_features = np.array(test_dataset['expression'])
- # Only apply scaling for non-categorical features
- train_features = scaler.fit_transform(train_features)
- val_features = scaler.transform(val_features)
- test_features = scaler.transform(test_features)
- # Extract features and labels as numpy arrays
- train_features = scaler.fit_transform(train_features)
- val_features = scaler.transform(val_features)
- test_features = scaler.transform(test_features)
- # Verify no NaN or infinite values
- if np.any(np.isnan(train_features)) or np.any(np.isinf(train_features)):
- raise ValueError("Training features contain NaN or infinite values")
- print(f"Train features final shape: {train_features.shape}")
- print(f"Test features final shape: {test_features.shape}")
- output = {
- 'train': (
- train_features,
- np.array(train_dataset['labels']),
- train_idx
- ),
- 'validation': (
- val_features,
- np.array(val_dataset['labels']),
- val_idx
- ),
- 'test': (
- test_features,
- np.array(test_dataset['labels']),
- np.arange(len(test_dataset))
- )
- }
- if cfg.debug_args.resample_adata:
- # Resample the data
- print("Resampling the adata")
- for name, (features, labels, indices) in output.items():
- resampled_indices = np.random.choice(range(len(indices)), len(indices), replace=True)
- output[name] = (features[resampled_indices], labels[resampled_indices], indices[resampled_indices])
- return output
- def run_classifier(
- datasets: DatasetDict,
- adata: Optional[Tuple[ad.AnnData, ad.AnnData]],
- cfg: DictConfig,
- classifier_type: str,
- feature_type: str
- ) -> None:
- """Generic function to run different classifiers on different feature types."""
- # Initialize classifier and scaler
- if classifier_type == "random_forest":
- clf = RandomForestClassifier(**cfg.random_forest)
- elif classifier_type == "logistic_regression":
- clf = LogisticRegression(**cfg.logistic_regression)
- else:
- raise ValueError(f"Unknown classifier type: {classifier_type}")
- scaler = StandardScaler()
- # If using direct AnnData path
- if cfg.data.group_size == 1 and cfg.debug_args.on_adata:
- train_adata, test_adata = adata
- splits = prepare_features_from_anndata(train_adata, test_adata, cfg, scaler, feature_type)
- # Train classifier
- train_features, train_labels, train_indices = splits['train']
- clf.fit(train_features, train_labels)
- # Evaluate
- for name, (features, labels, indices) in splits.items():
- predictions = clf.predict(features)
- evaluate_method(
- predictions,
- labels,
- indices,
- f"{classifier_type}_{feature_type}_{name}",
- cfg.data.label_names,
- cfg.output_dir
- )
- return
- # Original dataloader-based path
- dataloaders = get_dataloaders(datasets, cfg)
- # Prepare H3 type data if needed
- if feature_type == "h3type":
- h3_arrays, type_to_idx, index_maps = prepare_h3type_data(datasets)
- n_types = len(type_to_idx)
- # Training
- print(f"Collecting training features for {classifier_type} on {feature_type}...")
- train_features = []
- train_labels = []
- train_indices = []
- for batch in dataloaders["train"]:
- indices = batch['indices'].cpu().numpy() # these are uuid = indices in the original dataset['train'] before splitting
- train_indices.extend(indices)
- if feature_type == "bulk_expression":
- features = batch['raw_counts'].mean(1).cpu().numpy()
- else: # h3type
- features = get_h3type_histogram(indices, h3_arrays['train'], index_maps['train'], n_types)
- # get
- train_features.append(features)
- train_labels.append(batch['labels'].cpu().numpy())
- train_features = np.vstack(train_features)
- train_features = scaler.fit_transform(train_features)
- train_labels = np.concatenate(train_labels)
- print("Train features:", train_features.shape)
- print("Train labels:", train_labels.shape)
- print(f"Training {classifier_type}...")
- # Check for NaN/Inf values
- assert not np.any(np.isnan(train_features)), "Features contain NaN values"
- assert not np.any(np.isinf(train_features)), "Features contain infinite values"
- # Verify feature array is contiguous
- if not train_features.flags['C_CONTIGUOUS']:
- train_features = np.ascontiguousarray(train_features)
- # Fit with verbose logging
- clf.fit(train_features, train_labels)
- # Evaluate on all sets
- for name in ['train', 'validation', 'test']:
- loader = dataloaders[name]
- print(f"Evaluating on {name} set...")
- predictions = []
- labels = []
- indices = []
- for batch in loader:
- batch_indices = batch['indices'].cpu().numpy()
- indices.extend(batch_indices)
- if feature_type == "bulk_expression":
- features = batch['raw_counts'].mean(1).cpu().numpy()
- else: # h3type
- features = get_h3type_histogram(
- batch_indices,
- h3_arrays[name],
- index_maps[name],
- n_types
- )
- features = scaler.transform(features)
- pred = clf.predict(features)
- predictions.extend(pred)
- labels.extend(batch['labels'].cpu().numpy())
- evaluate_method(
- np.array(predictions),
- np.array(labels),
- np.array(indices),
- f"{classifier_type}_{feature_type}_{name}",
- cfg.data.label_names,
- cfg.output_dir
- )
- def run_dl_benchmark_brain(
- datasets: DatasetDict,
- cfg: DictConfig,
- model_name: str,
- ) -> None:
- """Run a DL benchmark on brain annotation data.
- Brain annotation differs from COVID: data is spatial (not donor-level),
- each cell has a label, and we group spatially adjacent cells into bags.
- """
- from torch.utils.data import DataLoader
- from tissueformer.benchmark_models.cellcnn import CellCnn
- from tissueformer.benchmark_models.scagg import ScAGG
- from tissueformer.benchmark_models.scrat import ScRAT
- from tissueformer.benchmark_models.data import MILDataset, CroppedMILDataset, mil_collate_fn
- from tissueformer.benchmark_models.trainer import BenchmarkTrainer
- import transformers
- print(f"\n--- DL benchmark: {model_name} ---")
- device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
- torch.manual_seed(cfg.seed)
- np.random.seed(cfg.seed)
- # Load model config
- model_cfg = OmegaConf.load(
- os.path.join(
- os.path.dirname(__file__),
- "config", "benchmark_models", f"{model_name}.yaml",
- )
- )
- if hasattr(cfg, "benchmark_models") and hasattr(cfg.benchmark_models, model_name):
- model_cfg = OmegaConf.merge(model_cfg, cfg.benchmark_models[model_name])
- # Use the spatial group dataloaders to get grouped bags
- dataloaders = get_dataloaders(datasets, cfg)
- # Collect all data through dataloaders into bags
- def collect_bags(loader, split_name):
- all_cells = []
- all_labels = []
- for batch in loader:
- raw = batch['raw_counts'].cpu().numpy() # (batch, group_size, n_genes)
- labs = batch['labels'].cpu().numpy() # (batch,)
- for i in range(raw.shape[0]):
- all_cells.append(raw[i])
- all_labels.append(labs[i])
- return all_cells, np.array(all_labels)
- train_cells, train_labels = collect_bags(dataloaders["train"], "train")
- val_cells, val_labels = collect_bags(dataloaders["validation"], "validation")
- test_cells, test_labels = collect_bags(dataloaders["test"], "test")
- # Each item in train_cells is already a (group_size, n_genes) array = one bag
- # Stack into flat expression matrix with sample indices
- def bags_to_mil_format(bags_list, labels_arr):
- expression_parts = []
- sample_indices = []
- cursor = 0
- for bag in bags_list:
- n = bag.shape[0]
- expression_parts.append(bag)
- sample_indices.append(np.arange(cursor, cursor + n))
- cursor += n
- expression = np.vstack(expression_parts).astype(np.float32)
- return expression, sample_indices, labels_arr
- train_X, train_si, train_y = bags_to_mil_format(train_cells, train_labels)
- val_X, val_si, val_y = bags_to_mil_format(val_cells, val_labels)
- test_X, test_si, test_y = bags_to_mil_format(test_cells, test_labels)
- # Apply preprocessing
- from tissueformer.benchmark_models.data import preprocess_zscore, preprocess_cp10k_log1p
- preprocess = model_cfg.get("preprocess", "raw")
- if preprocess == "zscore":
- scaler = StandardScaler()
- scaler.fit(train_X)
- train_X, _ = preprocess_zscore(train_X, scaler)
- val_X, _ = preprocess_zscore(val_X, scaler)
- test_X, _ = preprocess_zscore(test_X, scaler)
- elif preprocess == "cp10k_log1p":
- train_X = preprocess_cp10k_log1p(train_X)
- val_X = preprocess_cp10k_log1p(val_X)
- test_X = preprocess_cp10k_log1p(test_X)
- n_genes = train_X.shape[1]
- n_classes = int(max(train_y.max(), val_y.max(), test_y.max())) + 1
- # Compute class weights for imbalanced labels
- from tissueformer.class_weights import calculate_class_weights
- class_weights = calculate_class_weights(train_y, method="balanced", n_classes=n_classes)
- class_weights_tensor = torch.tensor(class_weights, dtype=torch.float32).to(device)
- # Build datasets (bags already formed, just wrap)
- train_ds = MILDataset(train_X, train_si, train_y)
- val_ds = MILDataset(val_X, val_si, val_y)
- test_ds = MILDataset(test_X, test_si, test_y)
- batch_size = model_cfg.batch_size
- train_loader = DataLoader(train_ds, batch_size=batch_size, shuffle=True, collate_fn=mil_collate_fn)
- val_loader = DataLoader(val_ds, batch_size=batch_size, shuffle=False, collate_fn=mil_collate_fn)
- test_loader = DataLoader(test_ds, batch_size=batch_size, shuffle=False, collate_fn=mil_collate_fn)
- # Build model
- if model_name == "cellcnn":
- cells_per_input = model_cfg.get("cells_per_input", 200)
- model = CellCnn(
- n_genes=n_genes, n_classes=n_classes,
- n_filters=model_cfg.n_filters,
- maxpool_percentage=model_cfg.maxpool_percentage,
- dropout=model_cfg.dropout,
- cells_per_input=cells_per_input,
- )
- optimizer = torch.optim.Adam(model.parameters(), lr=model_cfg.lr, weight_decay=model_cfg.l2_reg)
- l1_fn = lambda: model.l1_loss(model_cfg.l1_reg)
- loss_fn = torch.nn.CrossEntropyLoss(weight=class_weights_tensor)
- trainer = BenchmarkTrainer(
- model=model, optimizer=optimizer,
- train_loader=train_loader, val_loader=val_loader,
- device=device, n_epochs=model_cfg.n_epochs,
- early_stopping_patience=model_cfg.early_stopping_patience,
- l1_loss_fn=l1_fn, loss_fn=loss_fn,
- model_name=model_name, n_classes=n_classes,
- )
- elif model_name == "scagg":
- model = ScAGG(
- n_genes=n_genes, n_classes=n_classes,
- hidden_dim=model_cfg.hidden_dim,
- n_heads=model_cfg.n_heads, n_heads2=model_cfg.n_heads2,
- dropout=model_cfg.dropout,
- )
- optimizer = torch.optim.Adam(model.parameters(), lr=model_cfg.lr, weight_decay=model_cfg.weight_decay)
- loss_fn = torch.nn.CrossEntropyLoss(weight=class_weights_tensor, reduction=model_cfg.loss_reduction)
- trainer = BenchmarkTrainer(
- model=model, optimizer=optimizer,
- train_loader=train_loader, val_loader=val_loader,
- device=device, n_epochs=model_cfg.n_epochs,
- loss_fn=loss_fn, early_stopping_patience=None,
- model_name=model_name, n_classes=n_classes,
- )
- elif model_name == "scrat":
- model = ScRAT(
- n_genes=n_genes, n_classes=n_classes,
- hidden_dim=model_cfg.hidden_dim,
- n_heads=model_cfg.n_heads, n_layers=model_cfg.n_layers,
- dropout=model_cfg.dropout,
- )
- optimizer = torch.optim.Adam(model.parameters(), lr=model_cfg.lr, weight_decay=model_cfg.weight_decay)
- scheduler = transformers.get_cosine_schedule_with_warmup(
- optimizer, num_warmup_steps=model_cfg.lr_warmup_epochs,
- num_training_steps=model_cfg.n_epochs,
- )
- loss_fn = torch.nn.CrossEntropyLoss(weight=class_weights_tensor)
- trainer = BenchmarkTrainer(
- model=model, optimizer=optimizer,
- train_loader=train_loader, val_loader=val_loader,
- device=device, n_epochs=model_cfg.n_epochs,
- early_stopping_patience=model_cfg.early_stopping_patience,
- early_stopping_start_epoch=model_cfg.early_stopping_start_epoch,
- scheduler=scheduler, loss_fn=loss_fn,
- model_name=model_name, n_classes=n_classes,
- )
- else:
- raise ValueError(f"Unknown DL benchmark: {model_name}")
- # Train
- trainer.train()
- # Evaluate on test
- test_loss, test_preds, test_labels_out, test_probs = trainer.eval_epoch(test_loader)
- evaluate_method(
- test_preds, test_labels_out,
- np.arange(len(test_preds)),
- f"{model_name}_test",
- cfg.data.label_names if hasattr(cfg.data, 'label_names') else {},
- cfg.output_dir,
- )
- @hydra.main(version_base=None, config_path="config", config_name="benchmarks")
- def main(cfg: DictConfig) -> None:
- print("Running benchmarks...")
- # Print config
- print(OmegaConf.to_yaml(cfg))
- if cfg.debug:
- # Limit dataset size for faster iteration
- cfg.data.max_train_samples = 10000
- cfg.data.max_eval_samples = 10000
- # Setup wandb
- setup_wandb(cfg)
- # Load dataset
- dataset_dict = load_from_disk(cfg.data.dataset_path)
- datasets = prepare_datasets(dataset_dict, cfg)
- print(f"Loaded datasets: {datasets}")
- # Load AnnData if needed
- if cfg.debug_args.on_adata:
- adata = load_and_align_anndata(
- cfg.data.train_h5ad_files,
- cfg.data.test_h5ad_files,
- cfg.data.h5ad_directory,
- datasets
- )
- else:
- adata = None
- # ensure that the lengths are the same
- if cfg.debug_args.on_adata and cfg.debug:
- subsampled_idx_train = np.random.choice(len(adata[0]), len(datasets['train']), replace=False)
- subsampled_idx_test = np.random.choice(len(adata[1]), len(datasets['test']), replace=False)
- adata = (adata[0][subsampled_idx_train], adata[1][subsampled_idx_test])
- # Run benchmarks
- if cfg.run_bulk_expression_rf:
- run_classifier(datasets, adata, cfg, "random_forest", "bulk_expression")
- if cfg.run_bulk_expression_lr:
- run_classifier(datasets, adata, cfg, "logistic_regression", "bulk_expression")
- if cfg.run_h3type_rf:
- run_classifier(datasets, adata, cfg, "random_forest", "h3type")
- if cfg.run_h3type_lr:
- run_classifier(datasets, adata, cfg, "logistic_regression", "h3type")
- # Deep learning benchmarks
- dl_models = {
- "cellcnn": cfg.get("run_cellcnn", False),
- "scagg": cfg.get("run_scagg", False),
- "scrat": cfg.get("run_scrat", False),
- }
- for model_name, should_run in dl_models.items():
- if not should_run:
- continue
- run_dl_benchmark_brain(datasets, cfg, model_name)
- # Close wandb run
- wandb.finish()
- if __name__ == "__main__":
- main()
benchmarks.py at commit ff2b881, under MIT · at the source
Overview
Abstract
Background: Single-cell RNA sequencing technologies have enabled unprecedented insights into gene expression and opened new pathways for diagnostics and tissue annotation. At present, most computational approaches for interpreting single-cell data predict labels or properties based on isolated single-cell transcriptomic profiles. This approach overlooks the cellular composition within a sample, which is often critical for inferring tissue identity or other sample-level phenotypes.
Results: To address this limitation, we introduce TissueFormer, a Transformer-based neural network that infers population-level labels from groups of single-cell RNA profiles while retaining single-cell resolution. We applied TissueFormer to two tasks: predicting COVID-19 severity from single-cell RNA sequencing of blood samples, and predicting cortical area identity from spatial transcriptomic data in mouse brains. TissueFormer outperformed single-cell foundation models and machine learning methods applied to pseudobulk and cell type composition.
Conclusions: TissueFormer’s higher performance promises more accurate diagnostics and enables the automated construction of high-resolution brain region maps in individual mice directly from spatial transcriptomic data. Applied to mice with developmental perturbations to visual input, these maps revealed a significant reduction in predicted visual cortex area, illustrating how individual differences in neuroanatomy can be quantified. More broadly, TissueFormer provides a framework for predicting any population-level phenotypes which are influenced by cellular diversity and tissue-level organization.
Reproduced under the paper's license (CC BY), from the paper cited above.
Repositories
Its files are read in the Code ↔ Paper reader above, with 13 matches between paragraphs and lines of code.
ZadorLaboratory/TissueFormer
ff2b881ab40e09757facc4838eea7fce536327dd, 3 April 2026Availability: 1 check, the latest on 27 September 2026: the link answers
- 27 September 2026: the link answers
21 files
- applications/
__init__.py , Python, 1 line - applications/
brain_annotation/ , Python, 234 linesanalyze_attention.py - applications/
brain_annotation/ , Python, 892 lines, 4 matchesbenchmarks.py - applications/
brain_annotation/ , Jupyter, 666 linesdata/ area_annotation_plots.ip ynb - applications/
brain_annotation/ , Python, 149 linesdata/ calculate_class_weights. py - applications/
brain_annotation/ , Python, 168 linesdata/ mat_to_h5.py - applications/
brain_annotation/ , Python, 118 linesdata/ sort_by_CCF.py - applications/
brain_annotation/ , Python, 157 linesdata/ tokenize_cells.py - applications/
brain_annotation/ , Jupyter, 288 linesfigures/ cell_type_perf.ipynb - applications/
brain_annotation/ , Jupyter, 485 linesfigures/ flatmap_predictions.ipyn b - applications/
brain_annotation/ , Jupyter, 783 linesfigures/ hyperparameters.ipynb - applications/
brain_annotation/ , Python, 159 lines, 1 matchfigures/ utils.py - applications/
brain_annotation/ , Jupyter, 101 linespaper_figures/ fig2/ cell_type_accuracy_conve rgence.ipynb - applications/
brain_annotation/ , Jupyter, 243 linespaper_figures/ fig2/ fig2a.ipynb - applications/
brain_annotation/ , Jupyter, 440 lines, 3 matchespaper_figures/ fig2/ hyperparameters.ipynb - applications/
brain_annotation/ , Jupyter, 306 linespaper_figures/ fig2/ interactive_slice_viewer .ipynb - applications/
brain_annotation/ , Python, 25 lines, 1 matchpaper_figures/ fig2/ single_cell_accuracy.py - applications/
brain_annotation/ , Jupyter, 1,302 lines, 1 matchpaper_figures/ fig3/ cell_type_viz.ipynb - applications/
brain_annotation/ , Python, 535 lines, 1 matchpaper_figures/ fig3/ compute_visp_area.py - applications/
brain_annotation/ , Jupyter, 502 lines, 1 matchpaper_figures/ fig3/ gene_viz.ipynb - repository limit reached (2,000 files or 30 MB): the rest is at the source (44 files)
- README.md, Text, 170 lines
Zenodo 15595324
Availability: 1 check, the latest on 27 September 2026: the link answers (HTTP 200)
- 27 September 2026: the link answers (HTTP 200)
5 files
- colormycells/
__init__.py , Python, 7 lines - colormycells/
colormap.py , Python, 647 lines - setup.py, Python, 34 lines
- LICENSE, License, 674 lines
- README.md, Text, 123 lines
zadorlaboratory/colormycells
3edd4a073e9c5a052b536ea729cec5ee2cce5305, 17 July 2025Availability: 1 check, the latest on 27 September 2026: the link answers
- 27 September 2026: the link answers
4 files
- colormycells/
__init__.py , Python, 7 lines - colormycells/
colormap.py , Python, 647 lines, 1 match - LICENSE, License, 674 lines
- README.md, Text, 155 lines
The paper's code and data availability statement is in the Data section.
Tracing map
Proposed by the machine: these links were found in the paper and verified at the source, without human review. The map will receive a Zenodo DOI once one of the paper's authors has validated it with their ORCID.
What the map holds:
- 3 repositories of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
- 25 scripts, each with its path and the digest of its content;
- 13 matches between paragraphs of the paper and lines of the code (method lexical-v1);
- neither the text of the paper nor the code itself.
Its JSON (tracing-map.json) is deposited on Zenodo with its DOI once the map is validated.
Data
Datasets cited
- doi:10.17632/
5xfzcb4kn8.1 , at the source; found in “Data availability” - doi:10.17632/
8bhhk7c5n9.1 , at the source; found in “Data availability”
Data availability
All code is available at https://
Reproduced under the paper's license (CC BY), from the paper cited above.
Versions
The history of this record: each version stored by the harvester or made by a correction of its authors or of the maintainers of its code, and what changed in its facts. The texts of the paper (its abstract, its availability statements) are not part of it; versions that changed only those are not listed.
Version 1, 27 September 2026: the first record
Recorded: type, language, journal, volume, issue, pages, dates, 2 authors, 6 keywords, 11 MeSH terms, 1 funder, 58 references.
Cite
This paper
Benjamin, A. S., & Zador, A. (2026). Tissueformer: extending single-cell foundation models to predict population-level phenotypes. BMC bioinformatics, 27(1), 172. https://
BibTeX
@article{benjamin2026tis
author = {Benjamin, Ari S. and Zador, Anthony},
title = {{Tissueformer: extending single-cell foundation models to predict population-level phenotypes}},
journal = {BMC bioinformatics},
year = {2026},
month = jun,
volume = {27},
number = {1},
pages = {172},
publisher = {BMC},
issn = {1471-2105},
doi = {10.1186/
url = {https://
pmid = {42243667},
pmcid = {PMC13466306}
}
RIS
TY - JOUR
AU - Benjamin, Ari S.
AU - Zador, Anthony
TI - Tissueformer: extending single-cell foundation models to predict population-level phenotypes
T2 - BMC bioinformatics
J2 - BMC Bioinformatics
PY - 2026
DA - 2026/
VL - 27
IS - 1
SP - 172
SN - 1471-2105
PB - BMC
DO - 10.1186/
UR - https://
LA - en
ER -
CSL-JSON
{
"id": "10.1186/
"type": "article-journal",
"title": "Tissueformer: extending single-cell foundation models to predict population-level phenotypes",
"container-title": "BMC bioinformatics",
"author": [
{
"family": "Benjamin",
"given": "Ari S."
},
{
"family": "Zador",
"given": "Anthony"
}
],
"container-title-short":
"volume": "27",
"issue": "1",
"page": "172",
"DOI": "10.1186/
"PMID": "42243667",
"PMCID": "PMC13466306",
"ISSN": "1471-2105",
"publisher": "BMC",
"URL": "https://
"language": "en",
"issued": {
"date-parts": [
[
2026,
6,
4
]
]
}
}
The tracing map gets a citation of its own once an author has validated it and it has a DOI.
Similar papers
The papers with a page that share the most with this one: the tools found in their code, their categories, datasets, cited references and authors, the rarest counting most.
- [1] doi:10.1093/nar/gkag706 [code]
- scDifformer: diffusion-based post-training for virtual cell modeling across large-scale single-cell data.Journal: Nucleic acids researchIn common: anndata, PyTorch, seaborn, 5 other tools, 8 references
- [2] doi:10.1038/s41592-026-03194-8 [code]
- Beyond benchmarking: an expert-guided consensus approach to spatially aware clustering.Journal: Nature methodsIn common: anndata, PyTorch, seaborn, 5 other tools, genetics / omics, 8 references
- [3] doi:10.1038/s41467-026-71759-4 [code]
- CellNiche represents cellular microenvironments in atlas-scale spatial omics data with contrastive learning.Journal: Nature communicationsIn common: anndata, PyTorch, seaborn, 5 other tools, other condition, mouse, cellular / molecular, 7 references
- [4] doi:10.21203/rs.3.rs-9676637/v1 [code]
- A Comprehensive Benchmarking of Spatial Deconvolution and Domain Detection Methods across Diverse Tissues and Spatial Transcriptomic TechnologiesJournal: Research Square (preprint)In common: anndata, PyTorch, seaborn, 5 other tools, genetics / omics, 5 references
- [5] doi:10.1002/advs.77003 [code]
- SemanticST: A Scalable Multi-Contextual Graph Learning Framework for Uncovering Spatial Niches and Robust Multi-Sample Integration in Spatial Transcriptomics.Journal: Advanced science (Weinheim, Baden-Wurttemberg, Germany)In common: anndata, PyTorch, seaborn, 5 other tools, genetics / omics, 5 references
- [6] doi:10.1093/bib/bbag298 [code]
- Empowering multifaceted analysis of spatial transcriptomics data with RGAST.Journal: Briefings in bioinformaticsIn common: anndata, PyTorch, seaborn, 5 other tools, genetics / omics, mouse, 4 references
- [7] doi:10.1186/s13073-026-01704-z [code]
- Gene expression profiling enables refined parcellation of cortical layers in the heterogeneous human cerebral cortex.Journal: Genome medicineIn common: anndata, PyTorch, seaborn, 5 other tools, genetics / omics, mouse, cellular / molecular, 4 references
- [8] doi:10.1016/j.isci.2026.117206 [code]
- ReliST: A model-agnostic risk layer for spatial transcriptomics deconvolution.Journal: iScienceIn common: anndata, PyTorch, seaborn, 5 other tools, genetics / omics, mouse, 4 references
- [9] doi:10.64898/2026.03.30.714220 [code]
- An integrated single cell and spatial omics atlas of human prenatal developmentJournal: bioRxiv (preprint)In common: CuPy, Hugging Face Transformers, anndata, 7 other tools, 1 reference
- [10] doi:10.1038/s41467-026-74171-0 [code]
- Cluster replicability in single-cell and single-nucleus atlases of the mouse brain.Journal: Nature communicationsIn common: anndata, pandas, SciPy, 1 other tool, genetics / omics, mouse, cellular / molecular, 6 references
Contribute
The authors of this paper can claim it, correct its record and validate its tracing map, and the maintainers of its code (its owner, or a public member of its organization) correct what it says of their repository; anyone signed in can ask for its removal. Every request goes to OSCR's own machine, which answers it; your account page follows them.
Sign in with ORCID to claim this paper as one of its authors, correct its record or validate its tracing map: when the paper's metadata lists your ORCID iD, you are recognized at once. Maintainers of its code: sign in with GitHub, then claim the repository on your account page.
Claim this paper
Correct its record
Say what each link of this record is, remove the ones that are not the paper's, add the ones that are missing. The correction becomes a new version of the record, in its Versions section.
Validate its tracing map
You validate the map as this page shows it: 3 repositories of the authors' code, each at its verified commit and with its license, 25 scripts, and 13 matches between paragraphs and code (see the Code and Map sections). It then receives a DOI on Zenodo, with you (your ORCID iD) and OSCR as its creators; the code itself is not deposited.
The map's fingerprint: sha256:aa1dd1ef553adc1a…
Add the badge to its README
The badge links the code to this page. Copy one of these into the README of the paper's code: only you decide where it goes, and nothing is changed for you.
Markdown
[, paste the snippet at the top, then “Commit changes…” and, to review it first, “Create a new branch and start a pull request”. You open the pull request; OSCR asks for no permission.
Request its removal
To ask OSCR to remove this record, the copies of its authors' scripts or its tracing map, use the removal request page: signed in, you say who you are, what to remove and why, then review and confirm the request. Published rules decide every request (how).
Discussion, reproductions, activity
Discussion: questions and error reports about this paper and its code, from signed-in readers and its authors. It opens with sign-in.
Reproductions: reports from readers who ran the authors' code: what they reproduced, with which environment, commit and data. It opens with sign-in.
Activity: what happens around this paper: new versions of its record, its map's validation, discussions and reproductions. It opens with sign-in.
