How to Decompose ESMC Representations with Sparse Autoencoders (SAEs): A Complete Guide
You can decompose ESMC protein embeddings into interpretable sparse features by passing an SAEConfig to LogitsConfig and calling client.logits(), which returns sparse COO tensors via LogitsOutput.sae_outputs that are processed using the helper utilities in cookbook/snippets/sae.py.
The Biohub/esm repository provides built-in support for sparse autoencoders that factorize dense per-residue embeddings from ESMC (Evolutionary-Scale Modeling for Proteins) models. By leveraging the Forge inference API and the sae_config parameter, researchers can extract highly sparse activation patterns that reveal distinct biochemical features across protein sequences without modifying the underlying transformer architecture in esm/models/esmc.py.
Architecture of SAE Decomposition in Biohub/esm
The SAE pathway is fully integrated into the Forge inference stack. When you request SAE features, the server runs the selected autoencoders on the hidden states produced by the ESMC transformer and returns the activations as sparse tensors to minimize payload size.
Core API Contracts
LogitsConfig – Located in esm/sdk/api.py#L14-L38, this configuration class exposes the optional sae_config field. Setting this field tells the server to run SAE inference on the ESMC hidden states immediately after the forward pass.
SAEConfig – Defined in esm/sdk/api.py#L40-L58, this class validates your SAE selections. You specify model identifiers (e.g., "esmc-sae-1k-6b") and control feature normalization. Note that for 300M-size SAEs, you must set normalize_features=False to ensure compatibility.
LogitsOutput.sae_outputs – As implemented in esm/sdk/api.py#L81-L83, the API response contains a dictionary mapping each SAE name to a sparse COO tensor storing activations for every residue in the sequence.
Client-Side Processing Utilities
cookbook/snippets/sae.py – The get_sae_features helper (lines 10-39) wraps the Forge client call, extracts the sparse tensors, and handles batching across multiple sequences.
cookbook/snippets/sparse_utils.py – Low-level utilities remove_indexes (lines 7-24) and max_pool (lines 86-94) operate directly on the sparse COO tensors, allowing you to remove special tokens (BOS/EOS) and pool across residues without materializing dense matrices.
Server-Side Model Integration
esm/models/esmc.py (lines 45-62) generates the hidden states that feed into the SAE branch. This integration happens server-side; no client-side modifications to the model weights or architecture are required to access SAE features.
Step-by-Step Implementation Guide
Step 1: Configure the Forge Client and SAE Settings
First, instantiate the ESMCForgeInferenceClient and define which SAE models you want to apply to your sequences.
from esm.sdk.forge import ESMCForgeInferenceClient
from esm.sdk.api import SAEConfig
# Initialize the Forge client for an ESMC model that supports SAEs
forge = ESMCForgeInferenceClient(
model="esmc-6b-2024-12",
url="https://biohub.ai",
token="YOUR_ESM_API_KEY", # Replace with your actual API key
)
# Configure SAE: request the 1k-feature SAE for the 6B model
sae_cfg = SAEConfig(
models=["esmc-sae-1k-6b"],
normalize_features=True
)
Important: If you are using a 300M parameter ESMC model, set normalize_features=False in the SAEConfig to match the expected input distribution according to the validation logic in esm/sdk/api.py#L40-L58.
Step 2: Extract Sparse Activations
Use the get_sae_features helper from cookbook/snippets/sae.py to retrieve sparse tensors. Set pool=False to keep the per-residue activation matrix, or pool=True to automatically max-pool across the sequence length.
from cookbook.snippets.sae import get_sae_features
sequences = ["MKTLLILAVL...", "GAVMVL..."]
# Extract sparse SAE features (keep per-residue matrix)
sparse_features = get_sae_features(
client=forge,
sae_config=sae_cfg,
sequences=sequences,
pool=False,
)
# Each element is a torch.sparse_coo_tensor with shape (L+2, F)
# where L is sequence length and F is the SAE feature dimension
print(sparse_features[0]) # torch.sparse_coo_tensor of size (L+2, 1000)
The +2 in the first dimension accounts for the BOS (beginning-of-sequence) and EOS (end-of-sequence) tokens added by the tokenizer.
Step 3: Post-Process Sparse Tensors
To remove special tokens and obtain a single feature vector per protein, use the utilities in cookbook/snippets/sparse_utils.py. These functions avoid densifying the full tensor until necessary, preserving memory efficiency.
from cookbook.snippets.sparse_utils import remove_indexes, max_pool
pooled_vectors = []
for sparse_tensor in sparse_features:
# Remove BOS (index 0) and EOS (index -1) tokens
cleaned = remove_indexes(sparse_tensor, {0, -1})
# Max-pool across the residue axis (axis=0) to get a single vector per protein
protein_vector = max_pool(cleaned, axis=0)
pooled_vectors.append(protein_vector)
print(pooled_vectors[0].shape) # torch.Size([1000])
The remove_indexes function manipulates the sparse indices directly, while max_pool computes the maximum activation per feature across all positions without converting to a dense matrix first.
Working with Multiple SAEs and Batch Processing
You can request multiple SAEs simultaneously by passing multiple model identifiers to SAEConfig. The get_sae_features helper returns a list of dictionaries when multiple SAEs are requested.
# Request two different SAEs
multi_sae_cfg = SAEConfig(
models=["esmc-sae-1k-6b", "esmc-sae-2k-6b"],
normalize_features=True
)
# Returns list of dicts: [{"esmc-sae-1k-6b": tensor, "esmc-sae-2k-6b": tensor}, ...]
sparse_outputs = get_sae_features(
client=forge,
sae_config=multi_sae_cfg,
sequences=["AAAAA", "KKKKK"],
pool=False,
)
# Access specific SAE activations
first_sae_activations = sparse_outputs[0]["esmc-sae-1k-6b"]
second_sae_activations = sparse_outputs[0]["esmc-sae-2k-6b"]
The helper automatically handles batching for long sequence lists via the Forge API's batch endpoint, parallelizing requests server-side while maintaining the sparse tensor format throughout the pipeline.
Why Sparse COO Tensors Matter
SAE activations are inherently sparse—typically over 90% of features are zero for any given residue. By returning results as COO (Coordinate Format) tensors in LogitsOutput.sae_outputs, the Biohub API reduces network payload size by more than 90% compared to dense arrays. The cookbook/snippets/sparse_utils.py utilities allow you to perform downstream operations (index removal, pooling, feature selection) while staying in the sparse domain, only converting to dense arrays when you need specific values for analysis.
Summary
- Configure your SAE request using
SAEConfiginesm/sdk/api.py#L40-L58, specifying model names and settingnormalize_features=Falsefor 300M models. - Extract sparse activations via
LogitsConfigandclient.logits(), receiving COO tensors throughLogitsOutput.sae_outputsas defined inesm/sdk/api.py#L81-L83. - Process sparse tensors efficiently using
get_sae_featuresfromcookbook/snippets/sae.py#L10-L39and utilities likeremove_indexesandmax_poolfromcookbook/snippets/sparse_utils.py. - Pool per-residue features across the sequence length to obtain fixed-length protein representations without materializing dense activation maps.
- Scale to large datasets by leveraging automatic batching in the Forge client while maintaining sparse representations throughout the pipeline.
Frequently Asked Questions
What is the difference between standard ESMC embeddings and SAE-decomposed features?
Standard ESMC embeddings are dense vectors produced by the final hidden layer of the transformer in esm/models/esmc.py. SAE-decomposed features are sparse linear combinations of these dense representations, learned to isolate interpretable biochemical concepts such as specific amino acid preferences or structural motifs. While dense embeddings mix information across thousands of dimensions, SAE features typically activate only on specific biological patterns, making them more interpretable for downstream analysis.
Why must I disable feature normalization for 300M parameter SAE models?
According to the validation logic in esm/sdk/api.py#L40-L58, the 300M-size SAEs were trained without input normalization, and their weights expect raw hidden state magnitudes. Setting normalize_features=True for these models would scale the activations incorrectly, leading to degraded feature extraction. The 6B parameter SAEs require normalization enabled, so always verify your model size when configuring SAEConfig.
How do I convert the sparse COO tensors to dense numpy arrays for downstream analysis?
While keeping tensors sparse preserves memory efficiency during processing, you can densify specific samples when needed. PyTorch provides the .to_dense() method on sparse_coo_tensor objects. For example: dense_array = sparse_features[0].to_dense().numpy(). Be cautious with large batches, as densifying long protein sequences with high-dimensional SAEs can consume significant memory.
Can I run SAE decomposition on local ESMC models, or is the Forge API required?
The current SAE inference pathway is integrated into the Forge server architecture and accessed via ESMCForgeInferenceClient. The SAE weights run on the server side after the ESMC forward pass (esm/models/esmc.py#L45-62) returns hidden states. Local execution would require loading the SAE weights manually and applying them to hidden states from a local ESMC forward pass, which is not exposed in the current cookbook/snippets utilities.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →