Skip to content
Merged
Show file tree
Hide file tree
Changes from 2 commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
28 changes: 14 additions & 14 deletions transformer_lens/ActivationCache.py
Original file line number Diff line number Diff line change
Expand Up @@ -71,9 +71,9 @@ class ActivationCache:
the model predicting "road". This kind of analysis commonly falls under the category of "logit
attribution" or "direct logit attribution" (DLA).

>>> from transformer_lens import HookedTransformer
>>> model = HookedTransformer.from_pretrained("tiny-stories-1M")
Loaded pretrained model tiny-stories-1M into HookedTransformer
>>> from transformer_lens.model_bridge import TransformerBridge
>>> model = TransformerBridge.boot_transformers("roneneldan/TinyStories-1M")
>>> model.enable_compatibility_mode()

>>> _logits, cache = model.run_with_cache("Why did the chicken cross the")
>>> residual_stream, labels = cache.decompose_resid(return_labels=True, mode="attn")
Expand Down Expand Up @@ -270,12 +270,12 @@ def keys(self):

Examples:

>>> from transformer_lens import HookedTransformer
>>> model = HookedTransformer.from_pretrained("tiny-stories-1M")
Loaded pretrained model tiny-stories-1M into HookedTransformer
>>> from transformer_lens.model_bridge import TransformerBridge
>>> model = TransformerBridge.boot_transformers("roneneldan/TinyStories-1M")
>>> model.enable_compatibility_mode()
>>> _logits, cache = model.run_with_cache("Some prompt")
>>> list(cache.keys())[0:3]
['hook_embed', 'hook_pos_embed', 'blocks.0.hook_resid_pre']
['embed.hook_in', 'hook_embed', 'embed.hook_out']
Comment thread
jlarson4 marked this conversation as resolved.
Outdated

Returns:
List of all keys.
Expand Down Expand Up @@ -306,16 +306,16 @@ def __iter__(self) -> Iterator[str]:

Examples:

>>> from transformer_lens import HookedTransformer
>>> model = HookedTransformer.from_pretrained("tiny-stories-1M")
Loaded pretrained model tiny-stories-1M into HookedTransformer
>>> from transformer_lens.model_bridge import TransformerBridge
>>> model = TransformerBridge.boot_transformers("roneneldan/TinyStories-1M")
>>> model.enable_compatibility_mode()
>>> _logits, cache = model.run_with_cache("Some prompt")
>>> cache_interesting_names = []
>>> for key in cache:
... if not key.startswith("blocks.") or key.startswith("blocks.0"):
... cache_interesting_names.append(key)
>>> print(cache_interesting_names[0:3])
['hook_embed', 'hook_pos_embed', 'blocks.0.hook_resid_pre']
['embed.hook_in', 'hook_embed', 'embed.hook_out']

Returns:
Iterator over the cache.
Expand Down Expand Up @@ -392,12 +392,12 @@ def accumulated_resid(

Logit Lens analysis can be done as follows:

>>> from transformer_lens import HookedTransformer
>>> from transformer_lens.model_bridge import TransformerBridge
>>> import torch
>>> import pandas as pd

>>> model = HookedTransformer.from_pretrained("tiny-stories-1M", device="cpu", fold_ln=True)
Loaded pretrained model tiny-stories-1M into HookedTransformer
>>> model = TransformerBridge.boot_transformers("roneneldan/TinyStories-1M", device="cpu")
>>> model.enable_compatibility_mode()

>>> prompt = "Why did the chicken cross the"
>>> answer = " road"
Expand Down
11 changes: 6 additions & 5 deletions transformer_lens/evals.py
Original file line number Diff line number Diff line change
Expand Up @@ -323,10 +323,10 @@ class IOIDataset(Dataset):
.. code-block:: python

>>> from transformer_lens.evals import ioi_eval, IOIDataset
>>> from transformer_lens.HookedTransformer import HookedTransformer
>>> from transformer_lens.model_bridge import TransformerBridge

>>> model = HookedTransformer.from_pretrained('gpt2-small')
Loaded pretrained model gpt2-small into HookedTransformer
>>> model = TransformerBridge.boot_transformers("gpt2")
>>> model.enable_compatibility_mode()

>>> # Evaluate on a deterministic dataset (seed makes results reproducible)
>>> ds = IOIDataset(tokenizer=model.tokenizer, num_samples=100, seed=42)
Expand Down Expand Up @@ -551,10 +551,11 @@ def mmlu_eval(

.. code-block:: python

>>> from transformer_lens import HookedTransformer
>>> from transformer_lens.model_bridge import TransformerBridge
>>> from transformer_lens.evals import mmlu_eval

>>> model = HookedTransformer.from_pretrained("gpt2-small") # doctest: +SKIP
>>> model = TransformerBridge.boot_transformers("gpt2") # doctest: +SKIP
>>> model.enable_compatibility_mode() # doctest: +SKIP
>>> results = mmlu_eval(model, subjects="abstract_algebra", num_samples=10) # doctest: +SKIP
>>> print(f"Accuracy: {results['accuracy']:.2%}") # doctest: +SKIP
"""
Expand Down
7 changes: 4 additions & 3 deletions transformer_lens/utilities/exploratory_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,9 +31,10 @@ def test_prompt(

Examples:

>>> from transformer_lens import HookedTransformer, utilities
>>> model = HookedTransformer.from_pretrained("tiny-stories-1M")
Loaded pretrained model tiny-stories-1M into HookedTransformer
>>> from transformer_lens import utilities
>>> from transformer_lens.model_bridge import TransformerBridge
>>> model = TransformerBridge.boot_transformers("roneneldan/TinyStories-1M")
>>> model.enable_compatibility_mode()

>>> prompt = "Why did the elephant cross the"
>>> answer = "road"
Expand Down
Loading