In-Silico Perturbation
import CancerStFormer as cstformer
from cstformer.perturbation.perturb import Perturber
from cstformer.perturbation.perturb_stats import PerturberStats
import pandas as pd
import pickle
import os
import re
0.1 Parameter Descriptions
1.1 Option A: Delete Single Gene and View both Gene and Cell Embedding Shifts
Here we look at the following option:
Delete Gene: PDCD1
Load in pretrained model on masked learning objective
want to look at the cosine similarity to null distribution when perturbing PDCD1 in both per gene and per cell embedding
batch our perturbations at each model forward pass
isp = Perturber(
mode = 'extended', #perturbing either spot or extended model (based on tokenization)
perturb_type="delete", # in-silico deletion
genes_to_perturb=['PDCD1'], #immune checkpiont
perturb_rank_shift = None,
model_type="GeneClassifier", # GeneClassifier Model (can use CellClassifier, GeneClassifier, Pretrained (ML))
num_classes=2, # classes used in classifier
emb_mode="cell_and_gene", # cell and gene shifts
cell_emb_style='mean_pool',
max_ncells=1000, #num cells use in perturbation
emb_layer=-1,
forward_batch_size=80, #minibatch size
nproc=1, # num threads
token_dictionary_file='new_token_dictionary.pickle', # gene -> token mapping
)
isp.perturb_dataset(
model_directory='run-8eb93bdf/checkpoint-1000', #which final model to choose
input_data_file='STFormer_TNBC_neighbor.dataset', #tokenized dataset
output_directory='output/perturb/cell_shift',
output_prefix='perturb_extended')
We want to create stats to evaluate the effect of PDCD1 gene expression on:
Cell embeddings similarity to unperturbed cells
Gene embedding similarity to unperturbed genes, where we aggregate the gene shifts for our cells for each other token in our dataset
ex: how does deletion of PDCD1 effect the expression of another gene like CCR5AS
These results are ranked by mean cosine similarity (lowest -> highest) where low similarity signifies a greater impact on gene expression and a high dependency on the gene expression of our perturbed gene
provides standeard deviation of cosine similarity between gene embeddings and the number of detections
from cstformer.perturbation.perturb_stats import PerturberStats
ispstats = PerturberStats(
mode='aggregate_gene_shifts', # view gene shifts due to perturbation
gene_id_name_dict='gene_id_dictionary.pickle', # maps gene -> ENS (prints in gene_name)
token_dictionary_file = 'new_token_dictionary.pickle' # gene -> token dict
)
ispstats.compute_stats(
input_data_dir = 'output/perturb', # where embeddings are located
null_data_dir=None, # compare to null distribution
output_directory = 'output/perturb/perturb_stats',
output_prefix= 'perturb_extended'
)
import pandas as pd
pd.read_csv('perturb_extended_emb_aggregate_gene_shifts.csv')
| Perturbed | Gene_name | Ensembl_ID | Affected | Affected_gene_name | Affected_Ensembl_ID | Cosine_sim_mean | Cosine_sim_stdev | N_Detections | |
|---|---|---|---|---|---|---|---|---|---|
| 0 | PDCD1 | PDCD1 | PDCD1 | NEK11 | NEK11 | NEK11 | 0.567151 | 0.0 | 2 |
| 1 | PDCD1 | PDCD1 | PDCD1 | SYT17 | SYT17 | SYT17 | 0.632933 | 0.0 | 2 |
| 2 | PDCD1 | PDCD1 | PDCD1 | CASS4 | CASS4 | CASS4 | 0.645120 | 0.0 | 2 |
| 3 | PDCD1 | PDCD1 | PDCD1 | CAPN10-DT | CAPN10-DT | CAPN10-DT | 0.679628 | 0.0 | 2 |
| 4 | PDCD1 | PDCD1 | PDCD1 | PDE1B | PDE1B | PDE1B | 0.706005 | 0.0 | 2 |
| ... | ... | ... | ... | ... | ... | ... | ... | ... | ... |
| 15370 | PDCD1 | PDCD1 | PDCD1 | AL133453.1 | AL133453.1 | AL133453.1 | 0.999678 | 0.0 | 2 |
| 15371 | PDCD1 | PDCD1 | PDCD1 | ACP7 | ACP7 | ACP7 | 0.999701 | 0.0 | 2 |
| 15372 | PDCD1 | PDCD1 | PDCD1 | AL589182.1 | AL589182.1 | AL589182.1 | 0.999710 | 0.0 | 2 |
| 15373 | PDCD1 | PDCD1 | PDCD1 | PIH1D2 | PIH1D2 | PIH1D2 | 0.999718 | 0.0 | 2 |
| 15374 | PDCD1 | PDCD1 | PDCD1 | FAM218A | FAM218A | FAM218A | 0.999774 | 0.0 | 2 |
15375 rows × 9 columns
1.2 Option B: Delete all Genes and compare Cell State Shifts (Group A -> Group B)
Data Filtering & Batching
perturb_dataset()loads and filters your input dataset.In
isp_perturb_set()orisp_perturb_all(), cells are grouped into batches of sizeforward_batch_size.
Define Transition
get_state_embs
Embedding Extraction
Original (
full_orig) and perturbed (full_pert) token sequences are passed through the model (viaget_embs) to collect hidden‐state tensors.Because overexpressed tokens were prepended, the first N positions of
full_pertare your newly-inserted genes; the rest align with the original token order.
Cosine‐Similarity Quantification
Gene-level:
pu.quant_cos_sims(pert_emb, original_emb, …, emb_mode="gene")computes how each remaining gene embedding shifts when your target gene is overexpressed.Cell-level (if
cell_states_to_modelis set): average the non-padding embeddings before & after overexpression and quantify that shift against your state-embedding targets.
Aggregation & Output
For each cell or state, cosine-similarities are averaged (or bucketed) and stored in a dictionary keyed by the perturbed gene(s).
Final results are written out as “cell_embs_dict…” and—if
emb_mode="cell_and_gene"—also “gene_embs_dict…” files you can analyze within_silico_perturber_stats.
cell_states_to_model = {'state_key':'classification',
'start_state': 'Invasive cancer',
'goal_state': 'Invasive cancer + lymphocytes',
'alt_states': ['Invasive cancer + stroma','Invasive cancer + stroma + lymphocytes']}
from cstformer.tokenization.embedding_extractor import EmbExtractor
embex = EmbExtractor(
model_type = 'Pretrained',
num_classes=2,
filter_data = None,
max_ncells = 1000,
emb_layer=0,
summary_stat='exact_mean',
forward_batch_size=100,
token_dictionary_file='output/token_dictionary.pickle',
nproc = 16
)
state_embs_dict = embex.get_state_embs(
cell_states_to_model,
model_directory='output/spot/models/250422_102707_cstformer_L6_E3/final',
input_data_file='output/spot/visium_spot.dataset',
output_dir='output/perturb/cell_shift',
output_prefix='spot_state_shift'
)
visualize state embeddings
for key,value in state_embs_dict.items():
print(key)
Invasive cancer
Invasive cancer + lymphocytes
Invasive cancer + stroma
Invasive cancer + stroma + lymphocytes
Perform Perturbtion of all genes and compare cosine similarity from start state to goal state and rank according to shift to goal state
from stFormer.perturbation.stFormer_perturb import Perturber
from stFormer.perturbation.perturb_stats import PerturberStats
isp = Perturber(
perturb_type="delete",
perturb_rank_shift=None,
combos=0,
anchor_gene=None,
genes_to_perturb='all',
cell_states_to_model=cell_states_to_model,
state_embs_dict = state_embs_dict,
model_type="Pretrained",
num_classes=0,
emb_mode="cell",
cell_emb_style='mean_pool',
max_ncells=None,
emb_layer=0,
forward_batch_size=100,
nproc=12,
token_dictionary_file='output/token_dictionary.pickle',
)
isp.perturb_dataset(
model_directory='output/spot/models/250422_102707_cstformer_L6_E3/final',
input_data_file='output/spot/visium_spot.dataset',
output_directory='output/perturb/cell_shift',
output_prefix='perturb_spot')
ispstats = PerturberStats(mode="goal_state_shift",
genes_perturbed="all",
combos=0,
anchor_gene=None,
cell_states_to_model=cell_states_to_model,
token_dictionary_file='output/token_dictionary.pickle',
gene_id_name_dict='output/ensembl_mapping_dict.pickle')
ispstats.compute_stats(
input_data_dir="output/perturb/cell_shift", # this should be the directory
null_data_dir=None,
output_dir="output/perturb/cell_shift/perturb_stats",
output_prefix="visium_spot")
Structure of this dataset follows:
Gene: Token IDGene Name: hgnc name for geneEnsembl_ID: matching ensembl gene idShift to goal end: cosine shift from start state towards goal end state in response to given perturbationGoal end vs random pvalue: pvalue of cosine shift from start state towards goal end state by Wilcoxon to random distribution of max 10,000 cellsN Detections: Number of cells where the perturbed gene was detected/perturbedShift to alternate end state: Cosine shift from start state towards alternate end state in response to perturbationGoal end FDR: Benjamini Hochberg correction of Goal State vs Null PvalueAlt End FDR: Multiple Hypothesis Test Correction of Alternate State End vs Random PvalueSig: Binarized False Discovery Rate below significant (FDR < 0.05)
import pandas as pd
cell_shift_stats = pd.read_csv('output/perturb/cell_shift/perturb_stats/visium_spot.csv',index_col=0)
cell_shift_stats
| Gene | Gene_name | Ensembl_ID | Shift_to_goal_end | Goal_end_vs_random_pval | N_Detections | Shift_to_alt_end_Invasive cancer + stroma | Alt_end_vs_random_pval_Invasive cancer + stroma | Shift_to_alt_end_Invasive cancer + stroma + lymphocytes | Alt_end_vs_random_pval_Invasive cancer + stroma + lymphocytes | Goal_end_FDR | Alt_end_FDR_Invasive cancer + stroma | Alt_end_FDR_Invasive cancer + stroma + lymphocytes | Sig | |
|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|
| 2979 | 3907 | ANKRD37 | ENSG00000186352 | 0.001014 | 0.027741 | 5 | 0.000712 | 0.174332 | 0.000750 | 0.090140 | 0.924671 | 0.913586 | 0.873315 | 0 |
| 11646 | 17551 | SPINK4 | ENSG00000122711 | 0.000910 | 0.175442 | 6 | 0.001164 | 0.132797 | 0.001139 | 0.133500 | 0.953533 | 0.886458 | 0.888611 | 0 |
| 7884 | 10285 | ARID4A | ENSG00000032219 | 0.000824 | 0.023335 | 10 | 0.000348 | 0.601956 | 0.000403 | 0.424766 | 0.924671 | 0.982160 | 0.959837 | 0 |
| 9497 | 12278 | HOXB-AS1 | ENSG00000230148 | 0.000818 | 0.024909 | 6 | 0.000750 | 0.044214 | 0.000770 | 0.026993 | 0.924671 | 0.823552 | 0.790271 | 0 |
| 10644 | 13690 | PBX4 | ENSG00000105717 | 0.000778 | 0.003643 | 7 | 0.001221 | 0.005439 | 0.001154 | 0.005040 | 0.924671 | 0.734657 | 0.715931 | 0 |
| ... | ... | ... | ... | ... | ... | ... | ... | ... | ... | ... | ... | ... | ... | ... |
| 53 | 64 | TTC34 | ENSG00000215912 | -0.000763 | 0.013673 | 9 | 0.000024 | 0.509069 | -0.000080 | 0.630417 | 0.924671 | 0.974761 | 0.974686 | 0 |
| 5066 | 6672 | IDO1 | ENSG00000131203 | -0.000782 | 0.055477 | 7 | -0.000479 | 0.125197 | -0.000522 | 0.074521 | 0.924671 | 0.886458 | 0.853843 | 0 |
| 5856 | 7701 | CAVIN3 | ENSG00000170955 | -0.000879 | 0.013125 | 6 | -0.000491 | 0.008692 | -0.000545 | 0.005454 | 0.924671 | 0.734657 | 0.715931 | 0 |
| 479 | 604 | FYB2 | ENSG00000187889 | -0.000925 | 0.005857 | 11 | -0.000171 | 0.891572 | -0.000269 | 0.694267 | 0.924671 | 0.996187 | 0.985234 | 0 |
| 4609 | 6032 | SHROOM2 | ENSG00000146950 | -0.001132 | 0.015731 | 6 | -0.000631 | 0.063054 | -0.000729 | 0.040795 | 0.924671 | 0.848050 | 0.828248 | 0 |
11664 rows × 14 columns
1.3 Option C: Overexpress Gene
Data Filtering & Batching
perturb_dataset()loads and filters your input dataset.In
isp_perturb_set()orisp_perturb_all(), cells are grouped into batches of sizeforward_batch_size.
Applying Overexpression
For each example,
pu.overexpress_tokens(...)takes the token indices of your gene(s) and inserts them at the front of the token list.
Embedding Extraction
Original (
full_orig) and perturbed (full_pert) token sequences are passed through the model (viaget_embs) to collect hidden‐state tensors.Because overexpressed tokens were prepended, the first N positions of
full_pertare your newly-inserted genes; the rest align with the original token order.
Cosine‐Similarity Quantification
Gene-level:
pu.quant_cos_sims(pert_emb, original_emb, …, emb_mode="gene")computes how each remaining gene embedding shifts when your target gene is overexpressed.Cell-level (if
cell_states_to_modelis set): average the non-padding embeddings before & after overexpression and quantify that shift against your state-embedding targets.
Aggregation & Output
For each cell or state, cosine-similarities are averaged (or bucketed) and stored in a dictionary keyed by the perturbed gene(s).
Final results are written out as “cell_embs_dict…” and—if
emb_mode="cell_and_gene"—also “gene_embs_dict…” files you can analyze within_silico_perturber_stats.
from cstformer.perturbation.cstformer_perturb import Perturber
from cstformer.perturbation.perturb_stats import PerturberStats
isp = Perturber(
perturb_type="overexpress", # simulate gene overexpression
mode = 'spot',
genes_to_perturb=['ENSG00000091831'], #ESR1,
#perturb_rank_shift = None,
model_type="Pretrained", # ML Trained Model
filter_data={'subtype':'TNBC'}, #include filtering to TNBC to see overexpression of ESR1 on these cells
num_classes=0,
emb_mode="cell_and_gene", # cell and gene embeddings
cell_emb_style='mean_pool',
max_ncells=1000,
emb_layer=0,
forward_batch_size=50,
nproc=12,
token_dictionary_file='output/token_dictionary.pickle',
)
isp.perturb_dataset(
model_directory='output/spot/models/250422_102707_cstformer_L6_E3/final',
input_data_file='output/spot/visium_spot.dataset',
output_directory='output/perturb/overexpress',
output_prefix='perturb_spot')
ispstats = PerturberStats(mode='aggregate_gene_shifts',
token_dictionary_file='output/token_dictionary.pickle',
gene_name_id_dictionary_file='output/ensembl_mapping_dict.pickle'
)
ispstats.compute_stats(
input_data_dir='output/perturb/overexpress',
null_data_dir=None,
output_dir='output/perturb/overexpress/stats',
output_prefix='visium_spot'
)