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
¶
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 |
|---|---|---|---|
|
str | AllLayersSplitter
|
Hugging Face repository ID or configured model wrapper used to collect and project all layer states. |
required |
|
int
|
Maximum number of token or class scores returned per prediction. |
5
|
Raises:
| Type | Description |
|---|---|
ValueError
|
If |
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
explain
¶
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 |
|---|---|---|---|
|
str
|
One text passed to the wrapped model. |
required |
|
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
forward
¶
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 |
|---|---|---|---|
|
str
|
Prompt passed to the wrapped causal language model. |
required |
|
int
|
Maximum number of tokens to generate. |
10
|
|
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 |
Source code in interpreto/lens/logit_lens.py
TunedLens¶
interpreto.TunedLens
¶
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 |
|---|---|---|---|
|
str | AllLayersSplitter
|
Hugging Face repository ID or configured model wrapper used to collect and project all layer states. |
required |
|
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
fit
¶
fit(inputs, epochs=1, learning_rate=0.001, weight_decay=0.0)
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 |
|---|---|---|---|
|
str | Iterable[str]
|
Texts used to train the translators. |
required |
|
int
|
Number of passes over the texts. |
1
|
|
float
|
AdamW learning rate. |
0.001
|
|
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 |
Source code in interpreto/lens/tuned_lens.py
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 |
|---|---|---|---|
|
LensResults
|
Output returned by |
required |
|
str
|
Complete text used to produce |
required |
|
PreTrainedTokenizerBase
|
Tokenizer used by the lens splitter. |
required |
|
LabelNames | None
|
Optional display names for classification labels. |
None
|
|
bool
|
Whether causal predictions target the displayed column tokens. |
True
|
|
str
|
Additional CSS appended to the visualization styles. |
''
|
|
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 |
Examples: