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 toTrueto 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 toTrueto 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 usingreturn_hidden_states(use-1for 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 likeESM3InferenceClientorESMCInferenceClient(defined inesm/sdk/forge.pyandesm/sdk/base_forge_client.py) builds aLogitsOutputdataclass during the forward pass. When anyreturn_*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
LogitsConfiginesm/sdk/api.pycontrols which embeddings and hidden states are returned during inference.- Set
sequence=Trueandreturn_embeddings=Trueto extract per-residue embeddings for sequence-level downstream tasks. - Enable
return_mean_embeddingto obtain a single pooled vector per protein suitable for classification and clustering. - Use
return_hidden_stateswithith_hidden_layerto access intermediate transformer representations for layer-specific analysis. - The
logits()method populatesLogitsOutputfields 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:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →