Skip to content

Lens

Lens methods decode the residual stream at every transformer block boundary. Pass a model repository ID directly for the common case, or provide a configured AllLayersSplitter when you need to select a model class, tokenizer, device, or transformer block path.

Inference accepts one text at a time. To explain several texts, iterate over them. TunedLens.fit() accepts either one text or an iterable and trains on each text sequentially.

Classification with Logit Lens

from transformers import AutoModelForSequenceClassification

from interpreto import AllLayersSplitter, LogitLens, plot_lens

text = "Interpreto makes model decisions easier to inspect."
splitter = AllLayersSplitter(
    "distilbert-base-uncased-finetuned-sst-2-english",
    automodel=AutoModelForSequenceClassification,
)
lens = LogitLens(splitter, top_k=2)
results = lens(text)

plot_lens(results, text, tokenizer=splitter.tokenizer, label_names=["negative", "positive"])

Generation with Tuned Lens

from interpreto import TunedLens, plot_lens

text = "Paris is the capital of"
lens = TunedLens("distilgpt2", top_k=3)
lens.fit(["Paris is the capital of France.", "Rome is the capital of Italy."], epochs=1)
results = lens(text)

plot_lens(results, text, tokenizer=lens.splitter.tokenizer)

For causal language models, predictions are aligned with their observed target tokens by default. Each column therefore shows a target token, while the bottom Input row shows the preceding token that produced its prediction. Correct top predictions are outlined in green.

generate() uses the same convention. Include the prompt when plotting so the visualization can display the input to the first generated prediction:

prompt = "Although the committee initially rejected the proposal,"
generated_text, generated_results = lens.generate(prompt, max_new_tokens=10)

plot_lens(generated_results, prompt + generated_text, tokenizer=lens.splitter.tokenizer)

Set align=False on both calls for the conventional next-token view, where each column contains the prediction made after its displayed input token:

generated_text, generated_results = lens.generate(prompt, max_new_tokens=10, align=False)
plot_lens(
    generated_results,
    prompt + generated_text,
    tokenizer=lens.splitter.tokenizer,
    align=False,
)

LogitLens

interpreto.LogitLens

LogitLens(splitter, top_k=5)

Bases: Module

Project every residual-stream state through the model prediction head.

Logit Lens was introduced by nostalgebraist. It has no learned parameters: residual states are collected in one model trace and projected together through the model's native prediction path.

The prediction head was trained on final states, so early-layer scores are useful for rankings and within-model comparisons rather than as calibrated probabilities. Inference processes one text at a time; callers can iterate over several texts when needed. :meth:generate greedily completes a prompt and explains the generated continuation.

Parameters:

Name Type Description Default

splitter

str | AllLayersSplitter

Hugging Face repository ID or configured model wrapper used to collect and project all layer states.

required

top_k

int

Maximum number of token or class scores returned per prediction.

5

Raises:

Type Description
ValueError

If top_k is not positive.

Examples:

>>> from interpreto import LogitLens
>>> lens = LogitLens("hf-internal-testing/tiny-random-gpt2", top_k=3)
>>> results = lens("Interpreto is useful.")
>>> list(results) == lens.splitter.activation_names
True
Source code in interpreto/lens/logit_lens.py
def __init__(self, splitter: str | AllLayersSplitter, top_k: int = 5) -> None:
    super().__init__()
    if top_k < 1:
        raise ValueError("`top_k` must be positive.")

    self.splitter = splitter if isinstance(splitter, AllLayersSplitter) else AllLayersSplitter(splitter)
    self.top_k = top_k
    self.splitter._model.eval()

explain

explain(inputs, align=True)

Return top predictions at every transformer block boundary.

Causal language-model predictions are aligned with their observed next tokens by default. Classification outputs are unchanged.

Parameters:

Name Type Description Default

inputs

str

One text passed to the wrapped model.

required

align

bool

Whether causal predictions should target the token in the corresponding column.

True

Returns:

Name Type Description
LensResults LensResults

Top indices and normalized scores for each residual-stream state.

Source code in interpreto/lens/logit_lens.py
@torch.inference_mode()
def explain(self, inputs: str, align: bool = True) -> LensResults:
    """Return top predictions at every transformer block boundary.

    Causal language-model predictions are aligned with their observed next
    tokens by default. Classification outputs are unchanged.

    Args:
        inputs (str): One text passed to the wrapped model.
        align (bool): Whether causal predictions should target the token in
            the corresponding column.

    Returns:
        LensResults: Top indices and normalized scores for each residual-stream state.
    """
    logits = self._get_logits(inputs)
    if align and logits.ndim == 3:
        if logits.shape[1] < 2:
            raise ValueError("Aligned explanations require at least two tokens.")
        logits = logits[:, :-1]
    return self._format_outputs(logits)

forward

forward(inputs, align=True)

Alias for :meth:explain.

Source code in interpreto/lens/logit_lens.py
def forward(self, inputs: str, align: bool = True) -> LensResults:
    """Alias for :meth:`explain`."""
    return self.explain(inputs, align)

generate

generate(inputs, max_new_tokens=10, align=True)

Generate a continuation and explain its tokens.

Predictions target the generated token displayed in the same column by default. Set align=False to show the prediction made after each generated token instead.

Parameters:

Name Type Description Default

inputs

str

Prompt passed to the wrapped causal language model.

required

max_new_tokens

int

Maximum number of tokens to generate.

10

align

bool

Whether predictions should target the generated tokens displayed in the same columns.

True

Returns:

Type Description
tuple[str, LensResults]

tuple[str, LensResults]: Generated continuation and its layer predictions.

Raises:

Type Description
ValueError

If max_new_tokens is not positive or the wrapped model does not support generation.

Source code in interpreto/lens/logit_lens.py
@torch.inference_mode()
def generate(
    self,
    inputs: str,
    max_new_tokens: int = 10,
    align: bool = True,
) -> tuple[str, LensResults]:
    """Generate a continuation and explain its tokens.

    Predictions target the generated token displayed in the same column by
    default. Set ``align=False`` to show the prediction made after each
    generated token instead.

    Args:
        inputs (str): Prompt passed to the wrapped causal language model.
        max_new_tokens (int): Maximum number of tokens to generate.
        align (bool): Whether predictions should target the generated tokens
            displayed in the same columns.

    Returns:
        tuple[str, LensResults]: Generated continuation and its layer predictions.

    Raises:
        ValueError: If ``max_new_tokens`` is not positive or the wrapped
            model does not support generation.
    """
    if max_new_tokens < 1:
        raise ValueError("`max_new_tokens` must be positive.")

    model = self.splitter._model
    if not model.can_generate():
        raise ValueError("The wrapped model does not support generation.")

    tokenizer = self.splitter.tokenizer
    prompt_length = tokenizer(inputs, return_tensors="pt")["input_ids"].shape[1]
    with self.splitter.generate(inputs, max_new_tokens=max_new_tokens, do_sample=False) as tracer:
        sequence = tracer.result.save()
    sequence = sequence[0]
    generated_ids = sequence[prompt_length:]
    generated_text = tokenizer.decode(
        generated_ids,
        skip_special_tokens=False,
        clean_up_tokenization_spaces=False,
    )
    start = prompt_length - 1 if align else prompt_length
    stop = -1 if align else None
    logits = self._get_logits(sequence.unsqueeze(0))[:, start:stop]
    return generated_text, self._format_outputs(logits)

TunedLens

interpreto.TunedLens

TunedLens(splitter, top_k=5)

Bases: LogitLens

Learn one affine residual translator for each non-final model state.

Tuned Lens follows Belrose et al. (2023). Its translators are initialized to zero, making a new Tuned Lens identical to a Logit Lens. During fitting, all translators are trained together to match the model's final prediction distribution. Texts are processed one at a time while all model depths share one prediction-head call.

Use separate training and evaluation texts when assessing a fitted lens. Because this class is a regular :class:torch.nn.Module, translators can be persisted with state_dict() and load_state_dict().

Parameters:

Name Type Description Default

splitter

str | AllLayersSplitter

Hugging Face repository ID or configured model wrapper used to collect and project all layer states.

required

top_k

int

Maximum number of token or class scores returned per prediction.

5

Examples:

>>> from interpreto import TunedLens
>>> lens = TunedLens("hf-internal-testing/tiny-random-gpt2", top_k=3)
>>> losses = lens.fit(["Interpreto is useful."], epochs=1)
>>> results = lens("Interpreto is useful.")
Source code in interpreto/lens/tuned_lens.py
def __init__(self, splitter: str | AllLayersSplitter, top_k: int = 5) -> None:
    super().__init__(splitter, top_k)
    hidden_size = self.splitter._model.config.hidden_size
    reference_parameter = next(
        parameter for parameter in self.splitter._model.parameters() if parameter.is_floating_point()
    )
    device = None if reference_parameter.is_meta else reference_parameter.device
    self.translators = nn.ModuleList(
        [
            nn.Linear(
                hidden_size,
                hidden_size,
                device=device,
                dtype=reference_parameter.dtype,
            )
            for _ in self.splitter.split_points
        ]
    )
    for translator in self.translators:
        nn.init.zeros_(translator.weight)
        nn.init.zeros_(translator.bias)

fit

Fit every translator on a sequence of texts.

Each text is traced independently, while all model depths are optimized together in one prediction-head call.

Parameters:

Name Type Description Default

inputs

str | Iterable[str]

Texts used to train the translators.

required

epochs

int

Number of passes over the texts.

1

learning_rate

float

AdamW learning rate.

0.001

weight_decay

float

AdamW weight decay.

0.0

Returns:

Type Description
list[float]

list[float]: Mean loss for each epoch.

Raises:

Type Description
ValueError

If no training text is provided or epochs is not positive.

Source code in interpreto/lens/tuned_lens.py
def fit(
    self,
    inputs: str | Iterable[str],
    epochs: int = 1,
    learning_rate: float = 1e-3,
    weight_decay: float = 0.0,
) -> list[float]:
    """Fit every translator on a sequence of texts.

    Each text is traced independently, while all model depths are optimized
    together in one prediction-head call.

    Args:
        inputs (str | Iterable[str]): Texts used to train the translators.
        epochs (int): Number of passes over the texts.
        learning_rate (float): AdamW learning rate.
        weight_decay (float): AdamW weight decay.

    Returns:
        list[float]: Mean loss for each epoch.

    Raises:
        ValueError: If no training text is provided or `epochs` is not positive.
    """
    texts = [inputs] if isinstance(inputs, str) else list(inputs)
    if not texts:
        raise ValueError("Tuned Lens fitting requires at least one text.")
    if epochs < 1:
        raise ValueError("`epochs` must be positive.")

    model_parameters = list(self.splitter._model.parameters())
    requires_grad = [parameter.requires_grad for parameter in model_parameters]
    self.splitter._model.requires_grad_(False)
    losses = []
    optimizer = None

    try:
        self.train()
        for _ in range(epochs):
            epoch_loss = 0.0
            for text in texts:
                activations = torch.cat(self.splitter.get_activations(text), dim=0)
                transformed = self._transform(activations)
                if optimizer is None:
                    optimizer = torch.optim.AdamW(
                        self.translators.parameters(),
                        lr=learning_rate,
                        weight_decay=weight_decay,
                    )
                optimizer.zero_grad(set_to_none=True)
                loss = self._loss(self.splitter.apply_head(transformed))
                loss.backward()
                optimizer.step()
                epoch_loss += loss.item()
            losses.append(epoch_loss / len(texts))
    finally:
        self.eval()
        for parameter, original_requires_grad in zip(model_parameters, requires_grad, strict=True):
            parameter.requires_grad_(original_requires_grad)

    return losses

plot_lens

plot_lens renders model outputs from the final layer down to the embeddings, followed by the actual input tokens. The Embeddings row is the prediction head applied before the first transformer block; it is not the raw model input. The visible cells contain only the top prediction and use color intensity to show relative confidence. A green outline marks a correct prediction. Hover over a cell to inspect its numerical score and remaining top-k results.

Display model-depth predictions from the final layer to the input.

Color intensity shows relative confidence. Hover over a cell to see its numerical score and the remaining top-k predictions. Language-model plots outline correct top predictions and end with the actual input tokens below the embedding-state predictions.

Parameters:

Name Type Description Default

results

LensResults

Output returned by LogitLens.explain() or TunedLens.explain().

required

inputs

str

Complete text used to produce results, including the prompt when plotting a generated continuation.

required

tokenizer

PreTrainedTokenizerBase

Tokenizer used by the lens splitter.

required

label_names

LabelNames | None

Optional display names for classification labels.

None

align

bool

Whether causal predictions target the displayed column tokens.

True

custom_css

str

Additional CSS appended to the visualization styles.

''

save_path

str | PathLike[str] | None

Optional path for the rendered HTML.

None

Returns:

Name Type Description
None None

This function displays HTML and saves it when requested.

Raises:

Type Description
ValueError

If results is empty, has an unsupported output shape, or lacks the preceding text required for alignment.

Examples:

>>> results = lens.explain("Interpreto is useful.")
>>> plot_lens(results, "Interpreto is useful.", tokenizer=splitter.tokenizer)