Skip to content

Latest commit

 

History

2 Commits

Folders and files

NameName
Last commit message
Last commit date
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

SpatioVelo

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 .

Requirement

  • Julia >= 1.10

  • Python >= 3.7

  • scvelo

The minimum example of using SpatioVelo

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"
)

Evaluation (python code)

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}")

About

No description, website, or topics provided.

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages