# How to Extract Hidden States from All ESMC Transformer Layers for Downstream Analysis

> Easily extract hidden states from all ESMC transformer layers for downstream analysis. Run a forward pass and access the hidden_states tensor for in-depth protein model insights.

- Repository: [Biohub/esm](https://github.com/Biohub/esm)
- Tags: how-to-guide
- Published: 2026-05-30

---

**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`](https://github.com/Biohub/esm/blob/main/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`](https://github.com/Biohub/esm/blob/main/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`](https://github.com/Biohub/esm/blob/main/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.

```python
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.

```python
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.

```python
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`](https://github.com/Biohub/esm/blob/main/esm/models/esmc.py)** – Contains the `ESMC` class, `forward` method (lines 81-84), and `ESMCOutput` dataclass definition (lines 37-42)
- **[`esm/layers/transformer_stack.py`](https://github.com/Biohub/esm/blob/main/esm/layers/transformer_stack.py)** – Implements the `TransformerStack` that collects hidden states during the forward pass (lines 99-115)
- **[`esm/tokenization/sequence_tokenizer.py`](https://github.com/Biohub/esm/blob/main/esm/tokenization/sequence_tokenizer.py)** – Provides `EsmSequenceTokenizer` for converting protein strings to model tokens
- **[`esm/utils/constants/models.py`](https://github.com/Biohub/esm/blob/main/esm/utils/constants/models.py)** – Defines model identifier constants like `ESMC_600M` for the `from_pretrained` loader

## 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]`.