| name | residual-stream-state |
| description | Use this skill when working with transformer model interpretability, analyzing layer-by-layer predictions, training tuned lenses to understand intermediate representations, or peeking into iterative computations of transformers |
Demo Scripts
scripts/analyze_predictions.py
"""
Analyze Layer-wise Predictions using Tuned Lens
This script demonstrates how to use a tuned lens to analyze how predictions
evolve through the layers of a transformer model, providing insights into
the model's iterative computation process.
Requires: pip install tuned-lens torch transformers matplotlib
"""
import torch
import numpy as np
import matplotlib.pyplot as plt
from transformers import AutoModelForCausalLM, AutoTokenizer
from typing import List, Dict, Tuple, Optional
from dataclasses import dataclass
import json
from pathlib import Path
@dataclass
class PredictionTrajectory:
"""
Represents the evolution of predictions through transformer layers.
"""
text: str
tokens: List[str]
layer_predictions: Dict[int, np.ndarray]
final_prediction: np.ndarray
top_k_tokens: Dict[int, List[Tuple[str, float]]]
class TunedLensAnalyzer:
"""
Analyzer for understanding transformer predictions using tuned lenses.
"""
def __init__(
self,
model_name: str = "gpt2",
device: str = "cuda" if torch.cuda.is_available() else "cpu"
):
"""
Initialize the analyzer with a model.
Args:
model_name: HuggingFace model identifier
device: Device to run analysis on
"""
self.device = device
self.model_name = model_name
self.model = AutoModelForCausalLM.from_pretrained(model_name).to(device)
self.tokenizer = AutoTokenizer.from_pretrained(model_name)
if self.tokenizer.pad_token is None:
self.tokenizer.pad_token = self.tokenizer.eos_token
self.n_layers = (
self.model.config.n_layer
if hasattr(self.model.config, 'n_layer')
else self.model.config.num_hidden_layers
)
self.hidden_size = self.model.config.hidden_size
self.vocab_size = self.model.config.vocab_size
self.use_simple_lens = True
def get_unembedding_matrix(self) -> torch.Tensor:
"""
Get the unembedding matrix from the model.
Returns:
Unembedding weight matrix
"""
if hasattr(self.model, 'lm_head'):
return self.model.lm_head.weight
elif hasattr(self.model, 'embed_out'):
return self.model.embed_out.weight
else:
return self.model.transformer.wte.weight.T
def apply_lens(
self,
hidden_state: torch.Tensor,
layer_idx: int
) -> torch.Tensor:
"""
Apply lens (tuned or simple) to hidden state.
Args:
hidden_state: Hidden state from layer
layer_idx: Index of the layer
Returns:
Logits after applying lens
"""
if self.use_simple_lens:
unembedding = self.get_unembedding_matrix()
if hidden_state.shape[-1] != unembedding.shape[1]:
logits = torch.matmul(hidden_state, unembedding[:hidden_state.shape[-1]].T)
else:
logits = torch.matmul(hidden_state, unembedding.T)
else:
raise NotImplementedError("Load trained translators for tuned lens")
return logits
def analyze_text(
self,
text: str,
top_k: int = 5
) -> PredictionTrajectory:
"""
Analyze how predictions evolve through layers for given text.
Args:
text: Input text to analyze
top_k: Number of top predictions to track
Returns:
PredictionTrajectory object with analysis results
"""
inputs = self.tokenizer(text, return_tensors='pt').to(self.device)
input_ids = inputs['input_ids']
with torch.no_grad():
outputs = self.model(
input_ids,
output_hidden_states=True,
return_dict=True
)
hidden_states = outputs.hidden_states[1:]
final_logits = outputs.logits
tokens = [self.tokenizer.decode([tid]) for tid in input_ids[0]]
layer_predictions = {}
top_k_tokens = {}
for layer_idx, hidden_state in enumerate(hidden_states):
logits = self.apply_lens(hidden_state, layer_idx)
probs = torch.nn.functional.softmax(logits, dim=-1)
layer_predictions[layer_idx] = probs[0, -1].cpu().numpy()
top_probs, top_indices = torch.topk(probs[0, -1], top_k)
top_k_tokens[layer_idx] = [
(self.tokenizer.decode([idx.item()]), prob.item())
for idx, prob in zip(top_indices, top_probs)
]
final_probs = torch.nn.functional.softmax(final_logits[0, -1], dim=-1)
return PredictionTrajectory(
text=text,
tokens=tokens,
layer_predictions=layer_predictions,
final_prediction=final_probs.cpu().numpy(),
top_k_tokens=top_k_tokens
)
def plot_prediction_trajectory(
self,
trajectory: PredictionTrajectory,
target_tokens: Optional[List[str]] = None,
save_path: Optional[Path] = None
):
"""
Plot how predictions for specific tokens evolve through layers.
Args:
trajectory: Prediction trajectory to plot
target_tokens: Specific tokens to track (if None, use top final predictions)
save_path: Path to save the plot
"""
if target_tokens is None:
final_probs = trajectory.final_prediction
top_indices = np.argsort(final_probs)[-5:][::-1]
target_tokens = [self.tokenizer.decode([idx]) for idx in top_indices]
token_ids = [
self.tokenizer.encode(token, add_special_tokens=False)[0]
for token in target_tokens
]
fig, ax = plt.subplots(figsize=(12, 6))
layers = list(trajectory.layer_predictions.keys())
for token, token_id in zip(target_tokens, token_ids):
probs = [
trajectory.layer_predictions[layer][token_id]
for layer in layers
]
ax.plot(layers, probs, marker='o', label=f'"{token}"')
ax.set_xlabel('Layer')
ax.set_ylabel('Probability')
ax.set_title(f'Prediction Trajectory for: "{trajectory.text}"')
ax.legend()
ax.grid(True, alpha=0.3)
if save_path:
plt.savefig(save_path, dpi=150, bbox_inches='tight')
plt.show()
def compare_layer_predictions(
self,
trajectory: PredictionTrajectory,
layers_to_compare: Optional[List[int]] = None
) -> Dict[int, Dict[str, float]]:
"""
Compare top predictions across different layers.
Args:
trajectory: Prediction trajectory
layers_to_compare: Specific layers to compare (if None, use first, middle, last)
Returns:
Dictionary mapping layer indices to top predictions
"""
if layers_to_compare is None:
layers_to_compare = [0, self.n_layers // 2, self.n_layers - 1]
comparison = {}
for layer_idx in layers_to_compare:
if layer_idx in trajectory.top_k_tokens:
comparison[layer_idx] = {
token: prob
for token, prob in trajectory.top_k_tokens[layer_idx]
}
return comparison
def entropy_analysis(
self,
trajectory: PredictionTrajectory
) -> Dict[int, float]:
"""
Compute entropy of predictions at each layer.
Args:
trajectory: Prediction trajectory
Returns:
Dictionary mapping layer indices to entropy values
"""
entropies = {}
for layer_idx, probs in trajectory.layer_predictions.items():
probs = probs + 1e-10
entropy = -np.sum(probs * np.log(probs))
entropies[layer_idx] = entropy
return entropies
def main():
"""
Main function demonstrating tuned lens analysis capabilities.
"""
analyzer = TunedLensAnalyzer(model_name="gpt2")
test_texts = [
"The capital of France is",
"Two plus two equals",
"The sky is usually",
"Water freezes at zero degrees"
]
for text in test_texts:
print(f"\n{'='*60}")
print(f"Analyzing: '{text}'")
print('='*60)
trajectory = analyzer.analyze_text(text, top_k=5)
comparison = analyzer.compare_layer_predictions(trajectory)
print("\nTop predictions at different layers:")
for layer_idx, predictions in comparison.items():
print(f"\nLayer {layer_idx}:")
for token, prob in list(predictions.items())[:3]:
print(f" {token:15s}: {prob:.4f}")
entropies = analyzer.entropy_analysis(trajectory)
print("\nEntropy evolution:")
for layer_idx in [0, len(entropies)//2, len(entropies)-1]:
print(f" Layer {layer_idx:2d}: {entropies[layer_idx]:.4f}")
if text == test_texts[0]:
analyzer.plot_prediction_trajectory(
trajectory,
save_path=Path("prediction_trajectory.png")
)
results = {
"model": analyzer.model_name,
"n_layers": analyzer.n_layers,
"analyses": []
}
for text in test_texts[:2]:
trajectory = analyzer.analyze_text(text)
results["analyses"].append({
"text": text,
"final_top_tokens": trajectory.top_k_tokens[analyzer.n_layers - 1]
})
with open("analysis_results.json", "w") as f:
json.dump(results, f, indent=2)
print("\nAnalysis complete! Results saved to analysis_results.json")
if __name__ == "__main__":
main()
scripts/train_tuned_lens.py
"""
Train a Tuned Lens for Transformer Model Interpretability
This script demonstrates how to train a tuned lens for a transformer model,
allowing you to peek at intermediate layer predictions and understand how
the model builds its predictions layer-by-layer.
Requires: pip install tuned-lens torch transformers datasets
"""
import torch
from torch.utils.data import DataLoader
from transformers import AutoModelForCausalLM, AutoTokenizer
from datasets import load_dataset
import numpy as np
from pathlib import Path
from typing import Optional, List, Tuple
import logging
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
class TunedLensTrainer:
"""
A trainer for tuned lenses that learn to predict final outputs
from intermediate transformer representations.
"""
def __init__(
self,
model_name: str = "gpt2",
device: str = "cuda" if torch.cuda.is_available() else "cpu",
learning_rate: float = 1e-3,
batch_size: int = 8
):
"""
Initialize the tuned lens trainer.
Args:
model_name: HuggingFace model identifier
device: Device to run training on
learning_rate: Learning rate for optimizer
batch_size: Batch size for training
"""
self.device = device
.learning_rate = learning_rate
.batch_size = batch_size
logger.info()
.model = AutoModelForCausalLM.from_pretrained(model_name).to(device)
.tokenizer = AutoTokenizer.from_pretrained(model_name)
.tokenizer.pad_token :
.tokenizer.pad_token = .tokenizer.eos_token
.n_layers = .model.config.n_layer (.model.config, ) .model.config.num_hidden_layers
.hidden_size = .model.config.hidden_size
.vocab_size = .model.config.vocab_size
.translators = ._init_translators()
() -> torch.nn.ModuleList:
translators = torch.nn.ModuleList([
torch.nn.Linear(.hidden_size, .vocab_size, bias=)
_ (.n_layers)
]).to(.device)
translator translators:
torch.nn.init.normal_(translator.weight, std=)
torch.nn.init.zeros_(translator.bias)
translators
() -> [[torch.Tensor], torch.Tensor]:
torch.no_grad():
outputs = .model(
input_ids,
output_hidden_states=,
return_dict=
)
hidden_states = outputs.hidden_states[:]
logits = outputs.logits
hidden_states, logits
() -> torch.Tensor:
pred_log_probs = torch.nn.functional.log_softmax(pred_logits / temperature, dim=-)
target_probs = torch.nn.functional.softmax(target_logits / temperature, dim=-)
kl_div = torch.nn.functional.kl_div(
pred_log_probs,
target_probs,
reduction=,
log_target=
)
kl_div * (temperature ** )
() -> :
input_ids = batch[].to(.device)
hidden_states, target_logits = .extract_hidden_states(input_ids)
losses = {}
layer_idx, (hidden_state, translator, optimizer) (
(hidden_states, .translators, optimizers)
):
pred_logits = translator(hidden_state)
loss = .compute_kl_loss(pred_logits, target_logits)
optimizer.zero_grad()
loss.backward()
optimizer.step()
losses[] = loss.item()
losses
():
logger.info()
dataset = load_dataset(dataset_name, dataset_config, split=)
():
.tokenizer(
examples[],
padding=,
truncation=,
max_length=,
return_tensors=
)
tokenized_dataset = dataset.(tokenize_function, batched=)
tokenized_dataset.set_format(, columns=[])
dataloader = DataLoader(
tokenized_dataset,
batch_size=.batch_size,
shuffle=
)
optimizers = [
torch.optim.Adam(translator.parameters(), lr=.learning_rate)
translator .translators
]
logger.info()
step =
batch dataloader:
step >= n_steps:
losses = .train_step(batch, optimizers)
step % eval_interval == :
avg_loss = np.mean((losses.values()))
logger.info()
step +=
logger.info()
():
output_dir = Path(output_dir)
output_dir.mkdir(parents=, exist_ok=)
layer_idx, translator (.translators):
torch.save(
translator.state_dict(),
output_dir /
)
logger.info()
() -> torch.Tensor:
inputs = .tokenizer(text, return_tensors=).to(.device)
hidden_states, _ = .extract_hidden_states(inputs[])
torch.no_grad():
translator_logits = .translators[layer_idx](hidden_states[layer_idx])
probs = torch.nn.functional.softmax(translator_logits, dim=-)
probs
():
trainer = TunedLensTrainer(
model_name=,
device= torch.cuda.is_available() ,
learning_rate=,
batch_size=
)
trainer.train(
dataset_name=,
dataset_config=,
n_steps=,
eval_interval=
)
trainer.save_translators(Path())
sample_text =
layer_idx [, trainer.n_layers // , trainer.n_layers - ]:
probs = trainer.evaluate_lens(sample_text, layer_idx)
top_k =
top_probs, top_indices = torch.topk(probs[, -], top_k)
()
prob, idx (top_probs, top_indices):
token = trainer.tokenizer.decode([idx.item()])
()
__name__ == :
main()