A local RNA velocity model for cell differentiation inference on spatial transcriptomics based on Julia programming language. This project is currently WIP and highly inspired by cellDancer and JuloVelo .
-
Julia >= 1.10
-
Python >= 3.7
-
scvelo
Preprocessing (python code)
import scvelo as scv
import scanpy as sc
import numpy as np
adata = sc.read_h5ad("sample_data/Mouse_brain/HybISS_adata.h5ad")
sc.pp.pca(adata)
sc.pp.neighbors(adata, n_pcs = 50, n_neighbors = 30)
# Clustering
sc.tl.leiden(adata)
sc.tl.umap(adata)
sc.pl.umap(adata, color = "leiden")
adata.uns["iroot"] = np.flatnonzero(adata.obs["Subclass"] == 'Midbrain-hindbrain boundary')[0]
root_idx = np.flatnonzero(adata.obs['Subclass'] == 'Midbrain-hindbrain boundary')[0]
adata.uns["root"] = root_idx #Base on root leiden cluster
sc.tl.dpt(adata)
adata.var["highly_variable"] = "true"
adata.write_h5ad("sample_data/Mouse_brain/for_spatioVelo/HybISS_adata.h5ad")
Training model (julia code)
using SpatioVelo
using Plots, Statistics
using DataFrames
using CSV
using Flux
# ========================================
# configuration
# ========================================
epoch = 100
lr = 0.0001
α_spatial = 0.1
β_transcriptome = 0.9
neighbor = 30
suffix = "Mouse_brain"
basis = "xy_loc"
spatial_key = "X_$(basis)"
sample_number = 4628
# input file path
st_path = "sample_data/Mouse_brain/for_spatioVelo/HybISS_adata.h5ad"
# output file path
pre_path = "Result/Spatiovelo/$(suffix)"
if !isdir(pre_path)
mkpath(pre_path)
@info "Created output directory: $pre_path"
end
output_dir = joinpath(pre_path, "model_epochs$(epoch)_lr$(lr)_train_results")
println("RNA VELOCITY TRAINING PIPELINE")
println("Configuration:")
println(" Epochs: $epoch")
println(" Learning rate: $lr")
println(" Alpha (spatial): $α_spatial")
println(" Beta (transcriptome): $β_transcriptome")
println(" Output: $pre_path")
# ========================================
# PHASE 1: ALL GENES(No Filtering)
# ========================================
println("PHASE 1: FULL TRAINING (ALL GENES)")
@info "[1/7] Loading and preprocessing data..."
adata = read_adata(st_path)
n_cells_original = size(adata, 1)
n_genes_original = size(adata, 2)
@info " Loaded: $n_cells_original cells × $n_genes_original genes"
@info "[2/7] Normalizing..."
normalize(adata)
@info "[3/7] Filtering genes and determining kinetics..."
filter_and_gene_kinetics_predetermination(adata)
n_train_genes = sum(adata.var[!, "train_genes"])
@info " Train genes: $n_train_genes"
@info "[4/7] Reshaping data..."
reshape_data(adata)
@info "[5/7] Density sampling..."
density_sampling_with_indices(adata, sample_number = sample_number)
n_cells_sampled = size(adata.obs, 1)
@info " Sampled: $n_cells_sampled cells"
@info "[6/7] Creating train/validation split..."
split_train_validation(adata; val_ratio=0.2, random_seed=42)
n_train_cells = sum(adata.obs[!, "split"] .== "train")
n_val_cells = sum(adata.obs[!, "split"] .== "validation")
@info " Train: $n_train_cells cells"
@info " Validation: $n_val_cells cells"
@info "[7/7] Training model..."
results = train(
adata;
epochs = epoch,
neighbor_number = neighbor,
learning_rate = lr,
spatial_key = spatial_key,
dt = 0.5f0,
use_gpu = true,
α_spatial = α_spatial,
β_transcriptome = β_transcriptome,
output_dir = output_dir
)
@info "Training completed successfully!"
@info " Model saved to: $output_dir"
# ========================================
# Computing velocity and calculating psedotime
# ========================================
@info "Computing velocity for all genes..."
velocity_estimation(adata, results.model)
compute_cell_velocity(adata, basis = basis)
estimate_pseudotime(adata, basis = basis)
output_filename = joinpath(pre_path, "lr$(lr)_nofilter_spatioVelo_$(suffix).h5ad")
write_adata(adata, filename=output_filename)
println("\n Phase 1 Complete!")
println(" Saved: $output_filename")
println(" Genes: $n_train_genes (all)")
# ========================================
# PHASE 2: Select top n genes
# ========================================
println("PHASE 2: SELECTIVE FILTERING (TOP N GENES)")
if n_train_genes < 50
println("Genes counts: $(n_train_genes) < 50, do not process phase 2.")
return (
pre_path = pre_path,
output_dir = output_dir,
n_train_genes = n_train_genes
)
end
# depends on data
top_list = [50, 100, 150, 200, 300]
model_dir = output_dir
for (idx, top) in enumerate(top_list)
println("Processing: TOP $(top) GENES ($(idx)/$(length(top_list)))")
adata = read_adata(st_path)
normalize(adata)
filter_and_gene_kinetics_predetermination(adata)
reshape_data(adata)
density_sampling_with_indices(adata, sample_number = sample_number)
split_train_validation(adata; val_ratio=0.2, random_seed=42)
idx_result = prepare_selected_gene_indices(adata, model_dir, top)
if !idx_result.is_valid
@error "Critical Index Error at Top $top. Skipping this iteration."
continue
end
adata.uns["gene_indices"] = idx_result.data
converged_indices = idx_result.data["converged_indices"]
converged_in_original = idx_result.data["converged_in_original"]
adata.uns["X"] = adata.uns["X"][:, :, converged_indices]
adata.uns["train_X"] = adata.uns["train_X"][:, :, converged_indices]
adata.var[!, "train_genes"] .= false
adata.var[converged_in_original, "train_genes"] .= true
@info " Filtered data shape: $(size(adata.uns["X"]))"
# ----------------------------------------
# Load model
# ----------------------------------------
@info "[Step 7/9] Loading trained model..."
model = load_model(joinpath(model_dir, "model.bson"))
# ----------------------------------------
# Create filtered model
# ----------------------------------------
@info "[Step 8/9] Creating filtered model..."
Kinetics_converged = Chain(
l1 = SpatialDense(
model.layers.l1.weight[:, :, converged_indices],
model.layers.l1.bias[:, :, converged_indices],
leakyrelu
),
l2 = SpatialDense(
model.layers.l2.weight[:, :, converged_indices],
model.layers.l2.bias[:, :, converged_indices],
leakyrelu
),
l3 = SpatialDense(
model.layers.l3.weight[:, :, converged_indices],
model.layers.l3.bias[:, :, converged_indices],
sigmoid
)
)
# ----------------------------------------
# Computing velocity and calculating psedotime
# ----------------------------------------
@info "[Step 9/9] Computing velocity with filtered model..."
velocity_estimation(adata, Kinetics_converged)
compute_cell_velocity(adata, basis = basis)
estimate_pseudotime(adata, basis = basis)
output_filename = joinpath(
pre_path,
"lr$(lr)_ntop$(top)_filter_spatioVelo_$(suffix).h5ad"
)
write_adata(adata, filename=output_filename)
println("\n Top $top Complete!")
println(" Saved: $output_filename")
println(" Genes: $(length(converged_indices))")
GC.gc()
end
println("ALL PROCESSING COMPLETE!")
println("\n Data Output Directory: $pre_path")
println("\nModel Output Directory:")
println(" $output_dir")
println(" ├─ model.bson")
println(" ├─ training_stats.csv")
println(" ├─ train_loss_history.csv")
println(" └─ val_loss_history.csv")
Visualization (python code)
import scvelo as scv
import scanpy as sc
path ="/Result/Spatiovelo/Mouse_brain/lr0.0001_nofilter_spatioVelo_Mouse_brain.h5ad"
spatio_mouse_brain_velocity = sc.read_h5ad(path)
ax = scv.pl.velocity_embedding_stream(
spatio_mouse_brain_velocity,
basis="X_xy_loc",
color="Subclass",
legend_loc='none',
show=False,
figsize=(12,8),
s=200,
alpha=0.7
)
ax = scv.pl.velocity_embedding_stream(
spatio_mouse_brain_velocity,
basis="umap",
color="Subclass",
legend_loc='right'
)
ax = scv.pl.scatter(
spatio_mouse_brain_velocity, basis="X_xy_loc",
color="velocity_pseudotime",
figsize=(12,8),
legend_loc='none',
size=150,
color_map="gnuplot"
)
import scanpy as sc
from evaluation import spatial_time_consistency, spatial_velocity_consistency, gen_cross_boundary_correctness
path = "/Result/Spatiovelo/Mouse_brain/lr0.0001_nofilter_spatioVelo_Mouse_brain.h5ad"
vkey = "velocity"
tkey = "velocity_pseudotime"
spatial_key = "X_xy_loc"
spatial_graph_key = 'spatial_connectivities'
cluster_key = "Subclass"
cluster_edge = [
("Midbrain-hindbrain boundary", "Midbrain"),
("Midbrain-hindbrain boundary", "Ventral midbrain"),
("Midbrain-hindbrain boundary", "Dorsal hindbrain"),
("Dorsal hindbrain", "Hindbrain"),
("Forebrain", "Dorsal diencephalon"),
("Forebrain", "Cortical hem"),
("Cortical hem", "Cajal-Retzius")
]
adata = sc.read_h5ad(path)
print(f"spatial velocity consistency: {spatial_velocity_consistency(adata,vkey,spatial_graph_key)}")
print(f"spatial time consistency: {spatial_time_consistency(adata,tkey,spatial_graph_key)}")
def cbdir(adata, cluster_edges,vkey="velocity", tkey="velocity_pseudotime", cluster_key=cluster_key,
graph_key=spatial_graph_key, x_emb = spatial_key):
results = {}
# original cbdir
scores, avg_score = gen_cross_boundary_correctness(
adata=adata,
k_cluster=cluster_key,
k_velocity=vkey,
cluster_edges=cluster_edges,
tkey=tkey,
spatial_graph_key = graph_key,
x_emb=x_emb
)
results["original_cbdir"] = avg_score
# k-hop CBDir scores
scores, avg_scores = gen_cross_boundary_correctness(
adata=adata,
k_cluster=cluster_key,
k_velocity=vkey,
cluster_edges=cluster_edges,
tkey=tkey,
spatial_graph_key=graph_key,
k_hop=5,
dir_test=True,
x_emb=x_emb,
n_prune=30
)
for k, score in enumerate(avg_scores, 1):
results[f"{k}hop_cbdir"] = score
return results
cbdir_results = cbdir(adata,
vkey=vkey,
tkey=tkey,
graph_key = spatial_graph_key,
cluster_edges=cluster_edge,
x_emb=spatial_key,
cluster_key=cluster_key
)
print(f"CBDir Results: {cbdir_results}")