Pretain Hugging Face Dataset using Masked Learning Objective
Look over dataset tokenization tutorial prior to running this code, which will give you some of the prerequisites you need:
token_dictionary
Hugging Face Tokenized Dataset
Example Lengths
If using one of our pretrained models, this is an uneccessary step and you will just need to tokenize your own h5ad or looom file
from stFormer.pretrain.stFormer_pretrainer import PretrainML
1.1 Create Example Lengths File
This file specifies the number of tokens (genes) in each spot in your tokenized data. The maximum value should be specified by max_length in tokenization process (truncated tokens)
from datasets import load_from_disk
import pickle
ds = load_from_disk('output/spot/visium_spot.dataset')
def add_lengths(example):
example['length'] = len(example['input_ids'])
return example
length_ds = ds.map(add_lengths,num_proc=16)
lengths = length_ds['length']
with open('output/example_lengths.pickle', 'wb') as f:
pickle.dump(lengths, f)
1.2 Run Pretraining using BERT Framework
Masking Objective
Randomly mask out a fraction of tokens (genes) in each sequence.
Task: Predict the original gene ID at each masked position.
Loss: Cross-entropy between predicted token distribution and true gene ID.
Learns rich, unsupervised representations of spatial gene expression patterns.
Captures co-expression and neighborhood relationships without labels in neighbor mode.
Provides strong initialization for downstream tasks (e.g., cell-type classification).
Configuration
A standard BERT-style architecture (hidden size, layers, heads, etc.).
Dropout, layer-norm, and positional embeddings adapted for gene sequences.
Grouped batching by sequence length for efficient GPU utilization.
Training Loop
Iterates over masked sequences, computing MLM loss.
Periodic checkpointing of model weights.
Final model and tokenizer saved for later fine-tuning or inference.
trainer = PretrainML(
dataset_path="output/spot/visium_spot.dataset",
token_dict_path="output/token_dictionary.pickle",
example_lengths_path="output/example_lengths.pickle",
mode='spot',
output_dir="output/pretrained_model"
)
trainer.run_pretraining()
Run Pretraining with hyperparameter serach
trainer = PretrainML(
dataset_path="output/spot/visium_spot.dataset",
token_dict_path="output/token_dictionary.pickle",
example_lengths_path="output/example_lengths.pickle",
mode='spot',
output_dir="output/pretrained_model"
)
trainer.run_hyperparameter_train(
search_space={
"learning_rate": {"type": "loguniform", "low": 1e-5, "high": 1e-3},
"per_device_train_batch_size": {"type": "categorical", "values": [4, 8, 16]},
"weight_decay": {"type": "loguniform", "low": 1e-6, "high": 1e-2},
},
resources_per_trial={'cpu': 12,'gpu':1}
n_trials=10
)
1.3 Extract Embeddings
This module provides helper functions and a high-level class for turning tokenized gene sequences into fixed-size embedding vectors using a pretrained transformer model.
Encapsulates the end-to-end process of:
Loading
A pretrained model from
model_directory(withoutput_hidden_states=True)A HuggingFace disk‐based dataset
A token dictionary (gene ↔ token ID mapping)
Batching
Iterating in chunks of
forward_batch_sizeExtracting
input_idsand their lengths
Preprocessing
Applying
pad_sequencesto each batchGenerating the
attention_mask
Model Forward Pass
Running the model in
eval()mode without gradientsGathering all hidden states
Saving
Concatenate all batch embeddings into one tensor of shape
(N, hidden_dim)Write to disk as
output_prefix + ".pt"
from stFormer.tokenization.embedding_extractor import EmbeddingExtractor
from pathlib import Path
extractor = EmbeddingExtractor(
token_dict_path=Path('output/token_dictionary.pickle'),
emb_mode='cls',
emb_layer = -1,
forward_batch_size=64
)
embeddins = extractor.extract_embs(
model_directory='output/spot/models/250422_102707_stFormer_L6_E3/final',
dataset_path='output/spot/visium_spot.dataset',
output_directory='output/spot/embeddings',
output_prefix='visium_spot'
)