How to Extract Hidden States from All ESMC Transformer Layers for Downstream Analysis
To extract hidden states from all ESMC transformer layers, run a forward pass through the ESMC model and access the hidden_states field of the returned ESMCOutput object, which contains a tensor of shape [n_layers, batch, seq_len, d_model].
The ESMC (Evolutionary Scale Modeling - Compact) architecture in the Biohub/esm repository exposes intermediate representations through its TransformerStack implementation. When you need per-layer embeddings for protein language model probing, transfer learning, or visualization, the model naturally aggregates these states during the forward pass without requiring gradient computation.
How ESMC Exposes Hidden States
The ESM-C model constructs a TransformerStack composed of unified transformer blocks. During inference, each block returns its output representation along with optional attention weights. The collection mechanism works in two stages: first the stack accumulates states, then the main model formats them into a usable tensor.
TransformerStack Collection
In esm/layers/transformer_stack.py, the forward method (lines 99-115) iterates through each transformer block and appends the hidden state to an all_hidden_states list. After processing all blocks, it returns these as a tuple. This raw collection captures the representation from every layer in the stack before any final formatting.
Tensor Stacking in ESMC.forward
The ESMC.forward method in esm/models/esmc.py (lines 81-84) receives the list of hidden states from the stack. If Flash Attention caused padding alterations, it restores the original sequence shape, then stacks the list into a single 4-D tensor using torch.stack. This produces the final output shape: [n_layers, batch_size, sequence_length, d_model].
The ESMCOutput Dataclass
The output structure is defined in esm/models/esmc.py (lines 37-42) as the ESMCOutput dataclass. This container holds hidden_states as a primary field alongside embeddings and attentions, making it straightforward to access representations for downstream analysis without parsing complex nested dictionaries.
Extracting Hidden States: Three Methods
Depending on whether you need all layers, a specific layer, or attention weights alongside hidden states, the Biohub/esm SDK provides three distinct patterns.
Method 1: Direct Forward Pass (All Layers)
For comprehensive analysis requiring every layer's representation, use the standard forward pass with torch.no_grad() for inference efficiency.
import torch
from esm.models.esmc import ESMC
from esm.tokenization.sequence_tokenizer import EsmSequenceTokenizer
# Load tokenizer and pretrained model (600M parameter variant)
tokenizer = EsmSequenceTokenizer()
model = ESMC.from_pretrained("esmc_600m")
model.eval()
# Tokenize protein sequence
sequence = "MKTAYIAKQRQISFVKSHFSRQDILDLWIYHTQGYFPDWQNY"
tokens = tokenizer.encode(sequence)
seq_tokens = torch.tensor([tokens])
# Forward pass capturing all hidden states
with torch.no_grad():
output = model(sequence_tokens=seq_tokens, output_attentions=False)
# Access hidden states: shape (n_layers, batch, seq_len, d_model)
hidden_states = output.hidden_states
print(hidden_states.shape) # Example: torch.Size([33, 1, 50, 1280])
This returns a tensor containing the representation from every transformer block (33 layers for the 600M model), suitable for layer-wise probing or concatenation.
Method 2: Single Layer via LogitsConfig
When you only need a specific layer's hidden state to conserve memory, use the logits helper with LogitsConfig. This avoids materializing the full 4-D tensor.
from esm.sdk.api import LogitsConfig
# Configure to extract only layer 5 (0-indexed)
config = LogitsConfig(
ith_hidden_layer=5,
return_hidden_states=True,
)
# Pass embeddings through the logits pathway
logits_output = model.logits(input=output.embeddings, config=config)
# Result is shape (batch, seq_len, d_model)
layer_5_hidden = logits_output.hidden_states
print(layer_5_hidden.shape) # torch.Size([1, 50, 1280])
The ith_hidden_layer parameter indexes into the stack, allowing targeted extraction without computing the full forward pass overhead.
Method 3: Per-Layer Attention Weights
To simultaneously extract hidden states and attention matrices for visualization or attention-based analysis, enable the output_attentions flag.
with torch.no_grad():
output = model(sequence_tokens=seq_tokens, output_attentions=True)
# Hidden states (same as Method 1)
hidden_states = output.hidden_states # Shape: (n_layers, batch, seq_len, d_model)
# Attention weights: tuple of length n_layers, each shape (batch, n_heads, seq_len, seq_len)
attentions = output.attentions
print(f"Number of layers: {len(attentions)}")
print(f"Attention shape: {attentions[0].shape}") # Example: torch.Size([1, 20, 50, 50])
This approach provides the complete set of attention distributions alongside the hidden representations, essential for interpretability studies.
Key Source Files for Hidden State Analysis
Understanding the implementation requires examining these specific files in the Biohub/esm repository:
esm/models/esmc.py– Contains theESMCclass,forwardmethod (lines 81-84), andESMCOutputdataclass definition (lines 37-42)esm/layers/transformer_stack.py– Implements theTransformerStackthat collects hidden states during the forward pass (lines 99-115)esm/tokenization/sequence_tokenizer.py– ProvidesEsmSequenceTokenizerfor converting protein strings to model tokensesm/utils/constants/models.py– Defines model identifier constants likeESMC_600Mfor thefrom_pretrainedloader
Summary
- ESMC exposes hidden states through the
ESMCOutput.hidden_statesfield, automatically stacking per-layer representations into a tensor of shape[n_layers, batch, seq_len, d_model]. - The
TransformerStackcollects raw hidden states during the forward pass, whichESMC.forwardthen formats and returns. - For memory-efficient extraction of single layers, use
LogitsConfigwith theith_hidden_layerparameter rather than fetching all layers. - Set
output_attentions=Trueto simultaneously retrieve attention weight tensors alongside hidden states for interpretability research.
Frequently Asked Questions
What is the shape of the hidden_states tensor?
The hidden_states tensor returned in ESMCOutput has shape [n_layers, batch_size, sequence_length, d_model]. For the 600M parameter ESMC model with 33 transformer layers and embedding dimension 1280, processing a single sequence of length 50 would yield a tensor of shape [33, 1, 50, 1280].
Can I extract hidden states from specific layers only?
Yes. Instead of running the full forward pass, use the logits method with a LogitsConfig object specifying ith_hidden_layer to the 0-based index of your target layer. This approach avoids allocating memory for all 33 layers when you only need one specific representation.
Does extracting hidden states require additional memory?
Yes, storing the full hidden_states tensor requires significant GPU memory proportional to the number of layers, batch size, and sequence length. For the 600M model with 33 layers, this typically requires 33 times the memory of a single layer's activations. Use gradient checkpointing or the LogitsConfig single-layer method if memory constraints are critical.
How do I access attention weights alongside hidden states?
Pass output_attentions=True to the model.forward() call. The returned ESMCOutput will contain both hidden_states (the 4-D tensor of per-layer representations) and attentions (a tuple of attention weight tensors, one per layer). Each attention tensor has shape [batch_size, num_heads, sequence_length, sequence_length].
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 →