How to Configure LogitsConfig to Return Embeddings for Downstream ML Tasks

Set sequence=True and return_embeddings=True in the LogitsConfig dataclass to extract dense per-residue embeddings from ESM models via the client.logits() method.

The Biohub/esm repository provides the LogitsConfig dataclass in esm/sdk/api.py to control which representations are computed during inference. When you configure LogitsConfig to return embeddings for downstream ML tasks, you enable the extraction of high-dimensional protein representations suitable for classification, clustering, or transfer learning without requiring separate model forward passes.

Core LogitsConfig Fields for Embedding Extraction

Located in [esm/sdk/api.py](https://github.com/Biohub/esm/blob/main/esm/sdk/api.py#L14-L38), the LogitsConfig dataclass contains several boolean flags that determine which tensors are populated in the returned LogitsOutput:

  • sequence – Must be set to True to request per-residue logits for the amino-acid sequence. This is required when extracting per-residue embeddings.
  • return_embeddings – Returns the dense per-residue embeddings tensor (shape [L, D] where L is sequence length and D is embedding dimension). Set this to True to obtain embeddings for downstream classifiers or regressors.
  • return_mean_embedding – Returns a single mean-pooled embedding vector (shape [D]) averaged over all residues. This is ideal for protein-level prediction tasks.
  • return_hidden_states – Returns hidden states from each transformer layer. Useful when you need intermediate representations for analysis or layer-specific fine-tuning.
  • ith_hidden_layer – Specific layer index to expose when using return_hidden_states (use -1 for all layers, or zero-based index for specific layers).
  • return_mean_hidden_states – Returns mean-pooled hidden states from the selected layer(s).

Architectural Note: The logits() method in clients like ESM3InferenceClient or ESMCInferenceClient (defined in esm/sdk/forge.py and esm/sdk/base_forge_client.py) builds a LogitsOutput dataclass during the forward pass. When any return_* flag is enabled, the implementation populates the corresponding tensor fields without requiring additional model runs, making the same forward pass serve both traditional language modeling and representation learning.

Extracting Per-Residue Embeddings

To obtain embeddings for each amino acid position in a protein sequence, enable both the sequence track and the embeddings flag:

from esm.sdk.api import LogitsConfig, LogitsOutput
from esm.sdk.forge import client

# Encode a protein to tensor format

protein_tensor = client.encode(protein)

# Configure to return per-residue embeddings

config = LogitsConfig(sequence=True, return_embeddings=True)
output: LogitsOutput = client.logits(protein_tensor, config=config)

# embeddings tensor has shape [L, D]

per_residue_embeddings = output.embeddings

This pattern is validated in the test suite at [tests/oss_pytests/test_oss_client.py](https://github.com/Biohub/esm/blob/main/tests/oss_pytests/test_oss_client.py#L38-L40) and demonstrated in the cookbook example [cookbook/snippets/esm3.py](https://github.com/Biohub/esm/blob/main/cookbook/snippets/esm3.py#L6-L8).

Generating Mean-Pooled Protein Embeddings

For downstream tasks requiring a single fixed-size vector per protein (such as whole-protein classification or clustering), configure the client to return mean-pooled embeddings:

config = LogitsConfig(
    sequence=True,
    return_embeddings=True,
    return_mean_embedding=True,
)
output = client.logits(protein_tensor, config=config)

# Mean embedding is a [D] dimensional vector

protein_vector = output.mean_embedding

This approach aggregates residue-level information into a compact representation suitable for scikit-learn classifiers or deep learning architectures that expect fixed-length inputs.

Extracting Intermediate Hidden States

When you need representations from specific transformer layers rather than the final output embeddings, use the hidden states configuration:

config = LogitsConfig(
    sequence=True,
    return_hidden_states=True,
    ith_hidden_layer=5,  # Zero-based index for layer 5

    return_mean_hidden_states=True,
)
output = client.logits(protein_tensor, config=config)

# Layer-specific hidden states [L, D]

layer_activations = output.hidden_states

# Mean-pooled version [D]

layer_mean = output.mean_hidden_state

This technique is particularly valuable for probing tasks or when implementing layer-wise feature extraction for transfer learning pipelines.

Summary

  • LogitsConfig in esm/sdk/api.py controls which embeddings and hidden states are returned during inference.
  • Set sequence=True and return_embeddings=True to extract per-residue embeddings for sequence-level downstream tasks.
  • Enable return_mean_embedding to obtain a single pooled vector per protein suitable for classification and clustering.
  • Use return_hidden_states with ith_hidden_layer to access intermediate transformer representations for layer-specific analysis.
  • The logits() method populates LogitsOutput fields based on these flags without requiring multiple forward passes.

Frequently Asked Questions

What is the difference between return_embeddings and return_hidden_states?

return_embeddings returns the final output representations from the model's last layer, typically used as the default protein embedding. return_hidden_states returns intermediate activations from specified transformer layers, which are useful for analyzing how information propagates through the network or for layer-specific fine-tuning strategies.

Do I need to set sequence=True to get embeddings?

Yes. The sequence field must be set to True because the embedding extraction logic is tied to the sequence track computation. Without enabling the sequence track, the model does not compute the representations necessary for populating the embeddings or hidden_states fields in the output.

Can I use LogitsConfig with both ESM3 and ESMC models?

Yes. The LogitsConfig dataclass is designed to work across different ESM model implementations including ESM3InferenceClient and ESMCInferenceClient. However, note that the sae_config field is specific to ESM-C models for Sparse Auto-Encoder configurations and is not required for standard embedding extraction.

What shape should I expect for the returned embedding tensors?

Per-residue embeddings return a tensor of shape [L, D] where $L$ is the protein sequence length and $D$ is the model's hidden dimension (e.g., 512, 1024, or 1536 depending on the ESM variant). Mean embeddings return a vector of shape [D]. Hidden states return [L, D] for specific layers or a stacked tensor when requesting all layers.

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 →