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:

Summary

  • ESMC exposes hidden states through the ESMCOutput.hidden_states field, automatically stacking per-layer representations into a tensor of shape [n_layers, batch, seq_len, d_model].
  • The TransformerStack collects raw hidden states during the forward pass, which ESMC.forward then formats and returns.
  • For memory-efficient extraction of single layers, use LogitsConfig with the ith_hidden_layer parameter rather than fetching all layers.
  • Set output_attentions=True to 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:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →