interp-engine 0.0.24__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- interp_engine-0.0.24.dist-info/METADATA +21 -0
- interp_engine-0.0.24.dist-info/RECORD +28 -0
- interp_engine-0.0.24.dist-info/WHEEL +4 -0
- interp_engine-0.0.24.dist-info/licenses/LICENSE +19 -0
- neuron_explainer/__init__.py +0 -0
- neuron_explainer/activations/__init__.py +0 -0
- neuron_explainer/activations/activation_records.py +130 -0
- neuron_explainer/activations/activations.py +311 -0
- neuron_explainer/activations/attention_utils.py +121 -0
- neuron_explainer/activations/token_connections.py +59 -0
- neuron_explainer/api_client.py +190 -0
- neuron_explainer/azure.py +5 -0
- neuron_explainer/explanations/__init__.py +0 -0
- neuron_explainer/explanations/calibrated_simulator.py +194 -0
- neuron_explainer/explanations/explainer.py +2585 -0
- neuron_explainer/explanations/explanations.py +230 -0
- neuron_explainer/explanations/few_shot_examples.py +3125 -0
- neuron_explainer/explanations/prompt_builder.py +118 -0
- neuron_explainer/explanations/puzzles.json +399 -0
- neuron_explainer/explanations/puzzles.py +50 -0
- neuron_explainer/explanations/scoring.py +155 -0
- neuron_explainer/explanations/simulator.py +1121 -0
- neuron_explainer/explanations/test_explainer.py +227 -0
- neuron_explainer/explanations/test_simulator.py +269 -0
- neuron_explainer/explanations/token_space_few_shot_examples.py +212 -0
- neuron_explainer/fast_dataclasses/__init__.py +3 -0
- neuron_explainer/fast_dataclasses/fast_dataclasses.py +85 -0
- neuron_explainer/fast_dataclasses/test_fast_dataclasses.py +83 -0
|
@@ -0,0 +1,2585 @@
|
|
|
1
|
+
"""Uses API calls to generate explanations of neuron behavior."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import logging
|
|
6
|
+
import re
|
|
7
|
+
from abc import ABC, abstractmethod
|
|
8
|
+
from enum import Enum
|
|
9
|
+
from typing import Any, List, Optional, Sequence, Union
|
|
10
|
+
|
|
11
|
+
import numpy as np
|
|
12
|
+
|
|
13
|
+
from neuron_explainer.activations.activation_records import (
|
|
14
|
+
calculate_max_activation,
|
|
15
|
+
format_activation_records,
|
|
16
|
+
non_zero_activation_proportion,
|
|
17
|
+
)
|
|
18
|
+
from neuron_explainer.activations.activations import ActivationRecord
|
|
19
|
+
from neuron_explainer.activations.attention_utils import (
|
|
20
|
+
convert_flattened_index_to_unflattened_index,
|
|
21
|
+
)
|
|
22
|
+
from neuron_explainer.api_client import ApiClient
|
|
23
|
+
from neuron_explainer.explanations.few_shot_examples import (
|
|
24
|
+
ATTENTION_HEAD_FEW_SHOT_EXAMPLES,
|
|
25
|
+
AttentionTokenPairExample,
|
|
26
|
+
FewShotExampleSet,
|
|
27
|
+
)
|
|
28
|
+
from neuron_explainer.explanations.prompt_builder import (
|
|
29
|
+
HarmonyMessage,
|
|
30
|
+
PromptBuilder,
|
|
31
|
+
PromptFormat,
|
|
32
|
+
Role,
|
|
33
|
+
)
|
|
34
|
+
from neuron_explainer.explanations.token_space_few_shot_examples import (
|
|
35
|
+
TokenSpaceFewShotExampleSet,
|
|
36
|
+
)
|
|
37
|
+
|
|
38
|
+
logger = logging.getLogger(__name__)
|
|
39
|
+
ATTENTION_EXPLANATION_PREFIX = "this attention head"
|
|
40
|
+
ATTENTION_SEQUENCE_SEPARATOR = "<|sequence_separator|>"
|
|
41
|
+
|
|
42
|
+
|
|
43
|
+
# TODO(williamrs): This prefix may not work well for some things, like predicting the next token.
|
|
44
|
+
# Try other options like "this neuron activates for".
|
|
45
|
+
EXPLANATION_PREFIX = "the main thing this neuron does is find"
|
|
46
|
+
|
|
47
|
+
# we keep it blank to so the model just fills out: Explanation of neuron 4 behavior:
|
|
48
|
+
EXPLANATION_PREFIX_LOGITS = "my explanation for this neuron is "
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
def _split_numbered_list(text: str) -> list[str]:
|
|
52
|
+
"""Split a numbered list into a list of strings."""
|
|
53
|
+
lines = re.split(r"\n\d+\.", text)
|
|
54
|
+
# Strip the leading whitespace from each line.
|
|
55
|
+
return [line.lstrip() for line in lines]
|
|
56
|
+
|
|
57
|
+
|
|
58
|
+
def _remove_final_period(text: str) -> str:
|
|
59
|
+
"""Strip a final period or period-space from a string."""
|
|
60
|
+
if text.endswith("."):
|
|
61
|
+
return text[:-1]
|
|
62
|
+
elif text.endswith(". "):
|
|
63
|
+
return text[:-2]
|
|
64
|
+
return text
|
|
65
|
+
|
|
66
|
+
|
|
67
|
+
def _remove_method_from_explanation(explanation: str) -> str:
|
|
68
|
+
# attempt to remove "method" from the explanation if model outputted this (gpt has a lot of variations of this)
|
|
69
|
+
|
|
70
|
+
# case: "beginning. (Method 3)"
|
|
71
|
+
explanation = re.sub(
|
|
72
|
+
r"\s*[–—-]?\s*\(\s*[Mm]ethod\s+\d+\s*\)\s*$", "", explanation
|
|
73
|
+
).strip()
|
|
74
|
+
|
|
75
|
+
# case: "blah \nMethod 1"
|
|
76
|
+
explanation = re.sub(r"\s*\n\s*[Mm]ethod\s+\d+.*$", "", explanation).strip()
|
|
77
|
+
|
|
78
|
+
# case: "time words. Method 3 – temporal tokens – ts"
|
|
79
|
+
explanation = re.sub(r"\.\s*[Mm]ethod\s+\d+.*$", "", explanation).strip()
|
|
80
|
+
|
|
81
|
+
# case: "available — Method 1" or "say thought — Method 2, tokens after max are thought‑related"
|
|
82
|
+
explanation = re.sub(r"\s*[–—-]\s*[Mm]ethod\s+\d+.*$", "", explanation).strip()
|
|
83
|
+
|
|
84
|
+
# case: "Crystal (used Method 1)"
|
|
85
|
+
explanation = re.sub(
|
|
86
|
+
r"\s*\(\s*used\s+[Mm]ethod\s+\d+\s*\)\s*$", "", explanation
|
|
87
|
+
).strip()
|
|
88
|
+
|
|
89
|
+
# case: "www (used 1)"
|
|
90
|
+
explanation = re.sub(r"\s*\(\s*used\s+\d+\s*\)\s*$", "", explanation).strip()
|
|
91
|
+
|
|
92
|
+
# case: "homework (3)" - only at the end
|
|
93
|
+
explanation = re.sub(r"\s*\(\s*\d+\s*\)\s*$", "", explanation).strip()
|
|
94
|
+
|
|
95
|
+
# case: "code instructionsMethod 4"
|
|
96
|
+
explanation = re.sub(r"[Mm]ethod\s+\d+\s*$", "", explanation).strip()
|
|
97
|
+
|
|
98
|
+
# case: "non-Latin textMethod: 3"
|
|
99
|
+
explanation = re.sub(r"[Mm]ethod:\s*\d+\s*$", "", explanation).strip()
|
|
100
|
+
|
|
101
|
+
# strip ending "."
|
|
102
|
+
explanation = explanation.strip(".")
|
|
103
|
+
return explanation
|
|
104
|
+
|
|
105
|
+
|
|
106
|
+
# TODO: should pull from API and/or combine with the HARMONY_V4_MODELS
|
|
107
|
+
class ContextSize(int, Enum):
|
|
108
|
+
TWO_K = 2049
|
|
109
|
+
FOUR_K = 4097
|
|
110
|
+
SIXTEEN_K = 16384
|
|
111
|
+
ONETWENTYEIGHT_K = 128000
|
|
112
|
+
|
|
113
|
+
@classmethod
|
|
114
|
+
def from_int(cls, i: int) -> ContextSize:
|
|
115
|
+
for context_size in cls:
|
|
116
|
+
if context_size.value == i:
|
|
117
|
+
return context_size
|
|
118
|
+
raise ValueError(f"{i} is not a valid ContextSize")
|
|
119
|
+
|
|
120
|
+
|
|
121
|
+
# TODO: should pull these from API
|
|
122
|
+
HARMONY_V4_MODELS = [
|
|
123
|
+
"gpt-3.5-turbo",
|
|
124
|
+
"gpt-4",
|
|
125
|
+
"gpt-4o",
|
|
126
|
+
"gpt-4-turbo",
|
|
127
|
+
"gpt-4o-2024-05-13",
|
|
128
|
+
"gpt-4-1106-preview",
|
|
129
|
+
"gpt-4-turbo-2024-04-09",
|
|
130
|
+
]
|
|
131
|
+
|
|
132
|
+
|
|
133
|
+
class NeuronExplainer(ABC):
|
|
134
|
+
"""
|
|
135
|
+
Abstract base class for Explainer classes that generate explanations from subclass-specific
|
|
136
|
+
input data.
|
|
137
|
+
"""
|
|
138
|
+
|
|
139
|
+
def __init__(
|
|
140
|
+
self,
|
|
141
|
+
model_name: str,
|
|
142
|
+
prompt_format: PromptFormat = PromptFormat.HARMONY_V4,
|
|
143
|
+
# This parameter lets us adjust the length of the prompt when we're generating explanations
|
|
144
|
+
# using older models with shorter context windows. In the future we can use it to experiment
|
|
145
|
+
# with longer context windows.
|
|
146
|
+
context_size: ContextSize = ContextSize.ONETWENTYEIGHT_K,
|
|
147
|
+
max_concurrent: Optional[int] = 10,
|
|
148
|
+
cache: bool = False,
|
|
149
|
+
base_api_url: str = ApiClient.BASE_API_URL,
|
|
150
|
+
override_api_key: str | None = None,
|
|
151
|
+
):
|
|
152
|
+
# if prompt_format == PromptFormat.HARMONY_V4:
|
|
153
|
+
# assert model_name in HARMONY_V4_MODELS
|
|
154
|
+
if prompt_format in [PromptFormat.NONE, PromptFormat.INSTRUCTION_FOLLOWING]:
|
|
155
|
+
assert model_name not in HARMONY_V4_MODELS
|
|
156
|
+
# else:
|
|
157
|
+
# raise ValueError(f"Unhandled prompt format {prompt_format}")
|
|
158
|
+
|
|
159
|
+
self.model_name = model_name
|
|
160
|
+
self.prompt_format = prompt_format
|
|
161
|
+
self.context_size = context_size
|
|
162
|
+
self.client = ApiClient(
|
|
163
|
+
model_name=model_name,
|
|
164
|
+
max_concurrent=max_concurrent,
|
|
165
|
+
cache=cache,
|
|
166
|
+
base_api_url=base_api_url,
|
|
167
|
+
override_api_key=override_api_key,
|
|
168
|
+
)
|
|
169
|
+
|
|
170
|
+
async def generate_explanations(
|
|
171
|
+
self,
|
|
172
|
+
*,
|
|
173
|
+
num_samples: int = 5,
|
|
174
|
+
max_tokens: int = 60,
|
|
175
|
+
temperature: float = 1.0,
|
|
176
|
+
top_p: float = 1.0,
|
|
177
|
+
reasoning_effort: str | None = None,
|
|
178
|
+
**prompt_kwargs: Any,
|
|
179
|
+
) -> list[Any]:
|
|
180
|
+
"""Generate explanations based on subclass-specific input data."""
|
|
181
|
+
prompt = self.make_explanation_prompt(
|
|
182
|
+
max_tokens_for_completion=max_tokens, **prompt_kwargs
|
|
183
|
+
)
|
|
184
|
+
|
|
185
|
+
logger.info(prompt)
|
|
186
|
+
|
|
187
|
+
generate_kwargs: dict[str, Any] = {
|
|
188
|
+
"n": num_samples,
|
|
189
|
+
"max_tokens": max_tokens,
|
|
190
|
+
"temperature": temperature,
|
|
191
|
+
"top_p": top_p,
|
|
192
|
+
"reasoning_effort": reasoning_effort,
|
|
193
|
+
}
|
|
194
|
+
|
|
195
|
+
if self.prompt_format == PromptFormat.HARMONY_V4:
|
|
196
|
+
assert isinstance(prompt, list)
|
|
197
|
+
assert isinstance(prompt[0], dict) # Really a HarmonyMessage
|
|
198
|
+
generate_kwargs["messages"] = prompt
|
|
199
|
+
else:
|
|
200
|
+
assert isinstance(prompt, str)
|
|
201
|
+
generate_kwargs["prompt"] = prompt
|
|
202
|
+
|
|
203
|
+
response = await self.client.make_request(**generate_kwargs)
|
|
204
|
+
# logger.error("response in generate_explanations is %s", response)
|
|
205
|
+
|
|
206
|
+
if self.prompt_format == PromptFormat.HARMONY_V4:
|
|
207
|
+
# usually a content filter case
|
|
208
|
+
if "choices" not in response or "message" not in response["choices"][0]:
|
|
209
|
+
# print(f"error response: {response}")
|
|
210
|
+
explanations = []
|
|
211
|
+
else:
|
|
212
|
+
explanations = [x["message"]["content"] for x in response["choices"]]
|
|
213
|
+
elif self.prompt_format in [
|
|
214
|
+
PromptFormat.NONE,
|
|
215
|
+
PromptFormat.INSTRUCTION_FOLLOWING,
|
|
216
|
+
]:
|
|
217
|
+
explanations = [x["text"] for x in response["choices"]]
|
|
218
|
+
else:
|
|
219
|
+
raise ValueError(f"Unhandled prompt format {self.prompt_format}")
|
|
220
|
+
|
|
221
|
+
return self.postprocess_explanations(explanations, prompt_kwargs)
|
|
222
|
+
|
|
223
|
+
@abstractmethod
|
|
224
|
+
def make_explanation_prompt(
|
|
225
|
+
self, **kwargs: Any
|
|
226
|
+
) -> Union[str, list[HarmonyMessage]]:
|
|
227
|
+
"""
|
|
228
|
+
Create a prompt to send to the API to generate one or more explanations.
|
|
229
|
+
|
|
230
|
+
A prompt can be a simple string, or a list of HarmonyMessages, depending on the PromptFormat
|
|
231
|
+
used by this instance.
|
|
232
|
+
"""
|
|
233
|
+
...
|
|
234
|
+
|
|
235
|
+
@staticmethod
|
|
236
|
+
def _simple_clean_explanation(explanation: str) -> str:
|
|
237
|
+
# gpt oss outputs \u202f and \u2003 as space sometimes, so we need to replace it with a normal space
|
|
238
|
+
explanation = (
|
|
239
|
+
explanation.replace("\u202f", " ")
|
|
240
|
+
.replace("\u2003", "")
|
|
241
|
+
.replace("\u2002", "")
|
|
242
|
+
.strip()
|
|
243
|
+
)
|
|
244
|
+
if explanation.endswith("."):
|
|
245
|
+
explanation = explanation[:-1]
|
|
246
|
+
return explanation
|
|
247
|
+
|
|
248
|
+
def strip_explanation(self, explanation: str) -> str:
|
|
249
|
+
replaced = explanation
|
|
250
|
+
# Remove common prefixes
|
|
251
|
+
prefixes_to_remove = [
|
|
252
|
+
"References to ",
|
|
253
|
+
"Associated with ",
|
|
254
|
+
"Relates to ",
|
|
255
|
+
"Relating to ",
|
|
256
|
+
"Occurrences of ",
|
|
257
|
+
"Mentions of ",
|
|
258
|
+
"Related to ",
|
|
259
|
+
"Words related to ",
|
|
260
|
+
"Concepts related to ",
|
|
261
|
+
"Variations of the word ",
|
|
262
|
+
"Words indicating ",
|
|
263
|
+
"Words ",
|
|
264
|
+
"The word ",
|
|
265
|
+
"The phrase ",
|
|
266
|
+
"The tokens ",
|
|
267
|
+
"This neuron detects",
|
|
268
|
+
"This neuron predicts",
|
|
269
|
+
"This neuron activates for",
|
|
270
|
+
]
|
|
271
|
+
for prefix in prefixes_to_remove:
|
|
272
|
+
if replaced.startswith(prefix):
|
|
273
|
+
replaced = replaced[len(prefix) :]
|
|
274
|
+
break
|
|
275
|
+
|
|
276
|
+
# Remove common suffixes
|
|
277
|
+
suffixes_to_remove = [
|
|
278
|
+
" or related terms",
|
|
279
|
+
" and related terms",
|
|
280
|
+
" or its variations",
|
|
281
|
+
" and its variations",
|
|
282
|
+
" or related forms",
|
|
283
|
+
" and related forms",
|
|
284
|
+
]
|
|
285
|
+
for suffix in suffixes_to_remove:
|
|
286
|
+
if replaced.endswith(suffix):
|
|
287
|
+
replaced = replaced[: -len(suffix)]
|
|
288
|
+
break
|
|
289
|
+
replaced = self._simple_clean_explanation(replaced)
|
|
290
|
+
return replaced.strip()
|
|
291
|
+
|
|
292
|
+
def postprocess_explanations(
|
|
293
|
+
self, completions: list[str], prompt_kwargs: dict[str, Any]
|
|
294
|
+
) -> list[Any]:
|
|
295
|
+
"""Postprocess the completions returned by the API into a list of explanations."""
|
|
296
|
+
return completions # no-op by default
|
|
297
|
+
|
|
298
|
+
def _prompt_is_too_long(
|
|
299
|
+
self, prompt_builder: PromptBuilder, max_tokens_for_completion: int
|
|
300
|
+
) -> bool:
|
|
301
|
+
# We'll get a context size error if the prompt itself plus the maximum number of tokens for
|
|
302
|
+
# the completion is longer than the context size.
|
|
303
|
+
prompt_length = prompt_builder.prompt_length_in_tokens(self.prompt_format)
|
|
304
|
+
if prompt_length + max_tokens_for_completion > self.context_size.value:
|
|
305
|
+
print(
|
|
306
|
+
f"Prompt is too long: {prompt_length} + {max_tokens_for_completion} > "
|
|
307
|
+
f"{self.context_size.value}"
|
|
308
|
+
)
|
|
309
|
+
return True
|
|
310
|
+
return False
|
|
311
|
+
|
|
312
|
+
|
|
313
|
+
class TokenActivationPairExplainer(NeuronExplainer):
|
|
314
|
+
"""
|
|
315
|
+
Generate explanations of neuron behavior using a prompt with lists of token/activation pairs.
|
|
316
|
+
"""
|
|
317
|
+
|
|
318
|
+
def __init__(
|
|
319
|
+
self,
|
|
320
|
+
model_name: str,
|
|
321
|
+
prompt_format: PromptFormat = PromptFormat.HARMONY_V4,
|
|
322
|
+
# This parameter lets us adjust the length of the prompt when we're generating explanations
|
|
323
|
+
# using older models with shorter context windows. In the future we can use it to experiment
|
|
324
|
+
# with 8k+ context windows.
|
|
325
|
+
context_size: ContextSize = ContextSize.ONETWENTYEIGHT_K,
|
|
326
|
+
few_shot_example_set: FewShotExampleSet = FewShotExampleSet.ORIGINAL,
|
|
327
|
+
repeat_non_zero_activations: bool = True,
|
|
328
|
+
max_concurrent: Optional[int] = 10,
|
|
329
|
+
cache: bool = False,
|
|
330
|
+
base_api_url: str = ApiClient.BASE_API_URL,
|
|
331
|
+
override_api_key: str | None = None,
|
|
332
|
+
):
|
|
333
|
+
super().__init__(
|
|
334
|
+
model_name=model_name,
|
|
335
|
+
prompt_format=prompt_format,
|
|
336
|
+
max_concurrent=max_concurrent,
|
|
337
|
+
cache=cache,
|
|
338
|
+
base_api_url=base_api_url,
|
|
339
|
+
override_api_key=override_api_key,
|
|
340
|
+
)
|
|
341
|
+
self.context_size = context_size
|
|
342
|
+
self.few_shot_example_set = few_shot_example_set
|
|
343
|
+
self.repeat_non_zero_activations = repeat_non_zero_activations
|
|
344
|
+
|
|
345
|
+
def make_explanation_prompt(
|
|
346
|
+
self, **kwargs: Any
|
|
347
|
+
) -> Union[str, list[HarmonyMessage]]:
|
|
348
|
+
original_kwargs = kwargs.copy()
|
|
349
|
+
all_activation_records: Sequence[ActivationRecord] = kwargs.pop(
|
|
350
|
+
"all_activation_records"
|
|
351
|
+
)
|
|
352
|
+
max_activation: float = kwargs.pop("max_activation")
|
|
353
|
+
kwargs.setdefault("numbered_list_of_n_explanations", None)
|
|
354
|
+
numbered_list_of_n_explanations: Optional[int] = kwargs.pop(
|
|
355
|
+
"numbered_list_of_n_explanations"
|
|
356
|
+
)
|
|
357
|
+
if numbered_list_of_n_explanations is not None:
|
|
358
|
+
assert numbered_list_of_n_explanations > 0, numbered_list_of_n_explanations
|
|
359
|
+
# This parameter lets us dynamically shrink the prompt if our initial attempt to create it
|
|
360
|
+
# results in something that's too long. It's only implemented for the 4k context size.
|
|
361
|
+
kwargs.setdefault("omit_n_activation_records", 0)
|
|
362
|
+
omit_n_activation_records: int = kwargs.pop("omit_n_activation_records")
|
|
363
|
+
max_tokens_for_completion: int = kwargs.pop("max_tokens_for_completion")
|
|
364
|
+
assert not kwargs, f"Unexpected kwargs: {kwargs}"
|
|
365
|
+
|
|
366
|
+
prompt_builder = PromptBuilder()
|
|
367
|
+
prompt_builder.add_message(
|
|
368
|
+
Role.SYSTEM,
|
|
369
|
+
"We're studying neurons in a neural network. Each neuron looks for some particular "
|
|
370
|
+
"thing in a short document. Look at the parts of the document the neuron activates for "
|
|
371
|
+
"and summarize in a single sentence what the neuron is looking for. Don't list "
|
|
372
|
+
"examples of words.\n\nThe activation format is token<tab>activation. Activation "
|
|
373
|
+
"values range from 0 to 10. A neuron finding what it's looking for is represented by a "
|
|
374
|
+
"non-zero activation value. The higher the activation value, the stronger the match.",
|
|
375
|
+
)
|
|
376
|
+
few_shot_examples = self.few_shot_example_set.get_examples()
|
|
377
|
+
num_omitted_activation_records = 0
|
|
378
|
+
for i, few_shot_example in enumerate(few_shot_examples):
|
|
379
|
+
few_shot_activation_records = few_shot_example.activation_records
|
|
380
|
+
if self.context_size == ContextSize.TWO_K:
|
|
381
|
+
# If we're using a 2k context window, we only have room for one activation record
|
|
382
|
+
# per few-shot example. (Two few-shot examples with one activation record each seems
|
|
383
|
+
# to work better than one few-shot example with two activation records, in local
|
|
384
|
+
# testing.)
|
|
385
|
+
few_shot_activation_records = few_shot_activation_records[:1]
|
|
386
|
+
elif (
|
|
387
|
+
self.context_size == ContextSize.FOUR_K
|
|
388
|
+
and num_omitted_activation_records < omit_n_activation_records
|
|
389
|
+
):
|
|
390
|
+
# Drop the last activation record for this few-shot example to save tokens, assuming
|
|
391
|
+
# there are at least two activation records.
|
|
392
|
+
if len(few_shot_activation_records) > 1:
|
|
393
|
+
print(
|
|
394
|
+
f"Warning: omitting activation record from few-shot example {i}"
|
|
395
|
+
)
|
|
396
|
+
few_shot_activation_records = few_shot_activation_records[:-1]
|
|
397
|
+
num_omitted_activation_records += 1
|
|
398
|
+
self._add_per_neuron_explanation_prompt(
|
|
399
|
+
prompt_builder,
|
|
400
|
+
few_shot_activation_records,
|
|
401
|
+
i,
|
|
402
|
+
calculate_max_activation(few_shot_example.activation_records),
|
|
403
|
+
numbered_list_of_n_explanations=numbered_list_of_n_explanations,
|
|
404
|
+
explanation=few_shot_example.explanation,
|
|
405
|
+
)
|
|
406
|
+
self._add_per_neuron_explanation_prompt(
|
|
407
|
+
prompt_builder,
|
|
408
|
+
# If we're using a 2k context window, we only have room for two of the activation
|
|
409
|
+
# records.
|
|
410
|
+
(
|
|
411
|
+
all_activation_records[:2]
|
|
412
|
+
if self.context_size == ContextSize.TWO_K
|
|
413
|
+
else all_activation_records
|
|
414
|
+
),
|
|
415
|
+
len(few_shot_examples),
|
|
416
|
+
max_activation,
|
|
417
|
+
numbered_list_of_n_explanations=numbered_list_of_n_explanations,
|
|
418
|
+
explanation=None,
|
|
419
|
+
)
|
|
420
|
+
# If the prompt is too long *and* we omitted the specified number of activation records, try
|
|
421
|
+
# again, omitting one more. (If we didn't make the specified number of omissions, we're out
|
|
422
|
+
# of opportunities to omit records, so we just return the prompt as-is.)
|
|
423
|
+
if (
|
|
424
|
+
self._prompt_is_too_long(prompt_builder, max_tokens_for_completion)
|
|
425
|
+
and num_omitted_activation_records == omit_n_activation_records
|
|
426
|
+
):
|
|
427
|
+
original_kwargs["omit_n_activation_records"] = omit_n_activation_records + 1
|
|
428
|
+
return self.make_explanation_prompt(**original_kwargs)
|
|
429
|
+
return prompt_builder.build(self.prompt_format)
|
|
430
|
+
|
|
431
|
+
def _add_per_neuron_explanation_prompt(
|
|
432
|
+
self,
|
|
433
|
+
prompt_builder: PromptBuilder,
|
|
434
|
+
activation_records: Sequence[ActivationRecord],
|
|
435
|
+
index: int,
|
|
436
|
+
max_activation: float,
|
|
437
|
+
# When set, this indicates that the prompt should solicit a numbered list of the given
|
|
438
|
+
# number of explanations, rather than a single explanation.
|
|
439
|
+
numbered_list_of_n_explanations: Optional[int],
|
|
440
|
+
explanation: Optional[str], # None means this is the end of the full prompt.
|
|
441
|
+
) -> None:
|
|
442
|
+
max_activation = calculate_max_activation(activation_records)
|
|
443
|
+
user_message = f"""
|
|
444
|
+
|
|
445
|
+
Neuron {index + 1}
|
|
446
|
+
Activations:{format_activation_records(activation_records, max_activation, omit_zeros=False)}"""
|
|
447
|
+
# We repeat the non-zero activations only if it was requested and if the proportion of
|
|
448
|
+
# non-zero activations isn't too high.
|
|
449
|
+
if (
|
|
450
|
+
self.repeat_non_zero_activations
|
|
451
|
+
and non_zero_activation_proportion(activation_records, max_activation) < 0.2
|
|
452
|
+
):
|
|
453
|
+
user_message += (
|
|
454
|
+
f"\nSame activations, but with all zeros filtered out:"
|
|
455
|
+
f"{format_activation_records(activation_records, max_activation, omit_zeros=True)}"
|
|
456
|
+
)
|
|
457
|
+
|
|
458
|
+
if numbered_list_of_n_explanations is None:
|
|
459
|
+
user_message += f"\nExplanation of neuron {index + 1} behavior:"
|
|
460
|
+
assistant_message = ""
|
|
461
|
+
# For the IF format, we want <|endofprompt|> to come before the explanation prefix.
|
|
462
|
+
if self.prompt_format == PromptFormat.INSTRUCTION_FOLLOWING:
|
|
463
|
+
assistant_message += f" {EXPLANATION_PREFIX}"
|
|
464
|
+
else:
|
|
465
|
+
user_message += f" {EXPLANATION_PREFIX}"
|
|
466
|
+
prompt_builder.add_message(Role.USER, user_message)
|
|
467
|
+
|
|
468
|
+
if explanation is not None:
|
|
469
|
+
assistant_message += f" {explanation}."
|
|
470
|
+
if assistant_message:
|
|
471
|
+
prompt_builder.add_message(Role.ASSISTANT, assistant_message)
|
|
472
|
+
else:
|
|
473
|
+
if explanation is None:
|
|
474
|
+
# For the final neuron, we solicit a numbered list of explanations.
|
|
475
|
+
prompt_builder.add_message(
|
|
476
|
+
Role.USER,
|
|
477
|
+
f"""\nHere are {numbered_list_of_n_explanations} possible explanations for neuron {index + 1} behavior, each beginning with "{EXPLANATION_PREFIX}":\n1. {EXPLANATION_PREFIX}""",
|
|
478
|
+
)
|
|
479
|
+
else:
|
|
480
|
+
# For the few-shot examples, we only present one explanation, but we present it as a
|
|
481
|
+
# numbered list.
|
|
482
|
+
prompt_builder.add_message(
|
|
483
|
+
Role.USER,
|
|
484
|
+
f"""\nHere is 1 possible explanation for neuron {index + 1} behavior, beginning with "{EXPLANATION_PREFIX}":\n1. {EXPLANATION_PREFIX}""",
|
|
485
|
+
)
|
|
486
|
+
prompt_builder.add_message(Role.ASSISTANT, f" {explanation}.")
|
|
487
|
+
|
|
488
|
+
def postprocess_explanations(
|
|
489
|
+
self, completions: list[str], prompt_kwargs: dict[str, Any]
|
|
490
|
+
) -> list[Any]:
|
|
491
|
+
"""Postprocess the explanations returned by the API"""
|
|
492
|
+
numbered_list_of_n_explanations = prompt_kwargs.get(
|
|
493
|
+
"numbered_list_of_n_explanations"
|
|
494
|
+
)
|
|
495
|
+
if numbered_list_of_n_explanations is None:
|
|
496
|
+
return completions
|
|
497
|
+
else:
|
|
498
|
+
all_explanations = []
|
|
499
|
+
for completion in completions:
|
|
500
|
+
for explanation in _split_numbered_list(completion):
|
|
501
|
+
explanation = self.strip_explanation(explanation)
|
|
502
|
+
if explanation.startswith(EXPLANATION_PREFIX):
|
|
503
|
+
explanation = explanation[len(EXPLANATION_PREFIX) :]
|
|
504
|
+
all_explanations.append(explanation.strip())
|
|
505
|
+
return all_explanations
|
|
506
|
+
|
|
507
|
+
|
|
508
|
+
class TokenActivationPairLogitsExplainer(NeuronExplainer):
|
|
509
|
+
"""
|
|
510
|
+
Generate explanations of neuron behavior using a prompt with lists of token/activation pairs, with these changes:
|
|
511
|
+
- Don't tell the model to not specify specific words.
|
|
512
|
+
- Adding the top positive logits to the prompt.
|
|
513
|
+
- Telling the model to keep the explanation concise.
|
|
514
|
+
- Telling the model sometimes the neuron activates right before a specific word, token, or phrase, and to explain this with the format: 'say [the specific word, token or phrase]'
|
|
515
|
+
- Postprocessing will strip the ending period from the explanation.
|
|
516
|
+
- The additional explanation prefix is now "this neuron activates for".
|
|
517
|
+
"""
|
|
518
|
+
|
|
519
|
+
def __init__(
|
|
520
|
+
self,
|
|
521
|
+
model_name: str,
|
|
522
|
+
prompt_format: PromptFormat = PromptFormat.HARMONY_V4,
|
|
523
|
+
# This parameter lets us adjust the length of the prompt when we're generating explanations
|
|
524
|
+
# using older models with shorter context windows. In the future we can use it to experiment
|
|
525
|
+
# with 8k+ context windows.
|
|
526
|
+
context_size: ContextSize = ContextSize.ONETWENTYEIGHT_K,
|
|
527
|
+
few_shot_example_set: FewShotExampleSet = FewShotExampleSet.ORIGINAL,
|
|
528
|
+
repeat_non_zero_activations: bool = True,
|
|
529
|
+
max_concurrent: Optional[int] = 10,
|
|
530
|
+
cache: bool = False,
|
|
531
|
+
base_api_url: str = ApiClient.BASE_API_URL,
|
|
532
|
+
override_api_key: str | None = None,
|
|
533
|
+
):
|
|
534
|
+
super().__init__(
|
|
535
|
+
model_name=model_name,
|
|
536
|
+
prompt_format=prompt_format,
|
|
537
|
+
max_concurrent=max_concurrent,
|
|
538
|
+
cache=cache,
|
|
539
|
+
base_api_url=base_api_url,
|
|
540
|
+
override_api_key=override_api_key,
|
|
541
|
+
)
|
|
542
|
+
self.context_size = context_size
|
|
543
|
+
self.few_shot_example_set = few_shot_example_set
|
|
544
|
+
self.repeat_non_zero_activations = repeat_non_zero_activations
|
|
545
|
+
|
|
546
|
+
def make_explanation_prompt(
|
|
547
|
+
self, **kwargs: Any
|
|
548
|
+
) -> Union[str, list[HarmonyMessage]]:
|
|
549
|
+
original_kwargs = kwargs.copy()
|
|
550
|
+
all_activation_records: Sequence[ActivationRecord] = kwargs.pop(
|
|
551
|
+
"all_activation_records"
|
|
552
|
+
)
|
|
553
|
+
max_activation: float = kwargs.pop("max_activation")
|
|
554
|
+
top_positive_logits: List[str] = kwargs.pop("top_positive_logits")
|
|
555
|
+
kwargs.setdefault("numbered_list_of_n_explanations", None)
|
|
556
|
+
numbered_list_of_n_explanations: Optional[int] = kwargs.pop(
|
|
557
|
+
"numbered_list_of_n_explanations"
|
|
558
|
+
)
|
|
559
|
+
if numbered_list_of_n_explanations is not None:
|
|
560
|
+
assert numbered_list_of_n_explanations > 0, numbered_list_of_n_explanations
|
|
561
|
+
# This parameter lets us dynamically shrink the prompt if our initial attempt to create it
|
|
562
|
+
# results in something that's too long. It's only implemented for the 4k context size.
|
|
563
|
+
kwargs.setdefault("omit_n_activation_records", 0)
|
|
564
|
+
omit_n_activation_records: int = kwargs.pop("omit_n_activation_records")
|
|
565
|
+
max_tokens_for_completion: int = kwargs.pop("max_tokens_for_completion")
|
|
566
|
+
assert not kwargs, f"Unexpected kwargs: {kwargs}"
|
|
567
|
+
|
|
568
|
+
prompt_builder = PromptBuilder()
|
|
569
|
+
prompt_builder.add_message(
|
|
570
|
+
Role.SYSTEM,
|
|
571
|
+
"We're studying neurons in a neural network. Each neuron looks for some particular "
|
|
572
|
+
"thing in a short document or predicts the next word or token in a sentence. Your task is to "
|
|
573
|
+
"summarize in a single short phrase or word what the neuron is either looking for or predicting.\n\n"
|
|
574
|
+
"You will see the documents first, which are split up into tokens. The document activation format is token<tab>activation. Activation "
|
|
575
|
+
"values range from 0 to 10. A neuron finding what it's looking for is represented by a "
|
|
576
|
+
"non-zero activation value. The higher the activation value, the stronger the match.\n\n"
|
|
577
|
+
"After the documents, you will see 'Top Positive Logits', which predict the most likely next word or token after the activated tokens. "
|
|
578
|
+
"The top positive logits may have a specific pattern, like 'starts with a certain letter'. If you use the top positive logits, then format your response exactly like this: 'this neuron activates for say [the predicted text or pattern]'.\n\n"
|
|
579
|
+
"You should take both the documents and top positive logits into account when generating your explanation. Pay attention to the token immediately after the highest activating token. If there is no clear pattern in documents, then just explain what the top positive logits are predicting.\n\n"
|
|
580
|
+
"Finally, your explanation should not be a full sentence - it should be very concise, and should not include unnecessary "
|
|
581
|
+
"phrases like 'the neuron is looking for' or 'the neuron predicts the word' or 'words related to' "
|
|
582
|
+
"or 'concepts related to', or 'the word' etc. Simply say what it is the neuron is looking for or predicting, which can be "
|
|
583
|
+
"as short as a single word. If the neuron or pattern or prediction is a single word, then just say that word only.",
|
|
584
|
+
)
|
|
585
|
+
few_shot_examples = self.few_shot_example_set.get_examples()
|
|
586
|
+
num_omitted_activation_records = 0
|
|
587
|
+
for i, few_shot_example in enumerate(few_shot_examples):
|
|
588
|
+
few_shot_activation_records = few_shot_example.activation_records
|
|
589
|
+
if self.context_size == ContextSize.TWO_K:
|
|
590
|
+
# If we're using a 2k context window, we only have room for one activation record
|
|
591
|
+
# per few-shot example. (Two few-shot examples with one activation record each seems
|
|
592
|
+
# to work better than one few-shot example with two activation records, in local
|
|
593
|
+
# testing.)
|
|
594
|
+
few_shot_activation_records = few_shot_activation_records[:1]
|
|
595
|
+
elif (
|
|
596
|
+
self.context_size == ContextSize.FOUR_K
|
|
597
|
+
and num_omitted_activation_records < omit_n_activation_records
|
|
598
|
+
):
|
|
599
|
+
# Drop the last activation record for this few-shot example to save tokens, assuming
|
|
600
|
+
# there are at least two activation records.
|
|
601
|
+
if len(few_shot_activation_records) > 1:
|
|
602
|
+
print(
|
|
603
|
+
f"Warning: omitting activation record from few-shot example {i}"
|
|
604
|
+
)
|
|
605
|
+
few_shot_activation_records = few_shot_activation_records[:-1]
|
|
606
|
+
num_omitted_activation_records += 1
|
|
607
|
+
self._add_per_neuron_explanation_prompt(
|
|
608
|
+
prompt_builder,
|
|
609
|
+
few_shot_activation_records,
|
|
610
|
+
i,
|
|
611
|
+
calculate_max_activation(few_shot_example.activation_records),
|
|
612
|
+
numbered_list_of_n_explanations=numbered_list_of_n_explanations,
|
|
613
|
+
top_positive_logits=few_shot_example.top_positive_logits,
|
|
614
|
+
explanation=few_shot_example.explanation,
|
|
615
|
+
)
|
|
616
|
+
self._add_per_neuron_explanation_prompt(
|
|
617
|
+
prompt_builder,
|
|
618
|
+
# If we're using a 2k context window, we only have room for two of the activation
|
|
619
|
+
# records.
|
|
620
|
+
(
|
|
621
|
+
all_activation_records[:2]
|
|
622
|
+
if self.context_size == ContextSize.TWO_K
|
|
623
|
+
else all_activation_records
|
|
624
|
+
),
|
|
625
|
+
len(few_shot_examples),
|
|
626
|
+
max_activation,
|
|
627
|
+
numbered_list_of_n_explanations=numbered_list_of_n_explanations,
|
|
628
|
+
top_positive_logits=top_positive_logits,
|
|
629
|
+
explanation=None,
|
|
630
|
+
)
|
|
631
|
+
# If the prompt is too long *and* we omitted the specified number of activation records, try
|
|
632
|
+
# again, omitting one more. (If we didn't make the specified number of omissions, we're out
|
|
633
|
+
# of opportunities to omit records, so we just return the prompt as-is.)
|
|
634
|
+
if (
|
|
635
|
+
self._prompt_is_too_long(prompt_builder, max_tokens_for_completion)
|
|
636
|
+
and num_omitted_activation_records == omit_n_activation_records
|
|
637
|
+
):
|
|
638
|
+
original_kwargs["omit_n_activation_records"] = omit_n_activation_records + 1
|
|
639
|
+
return self.make_explanation_prompt(**original_kwargs)
|
|
640
|
+
return prompt_builder.build(self.prompt_format)
|
|
641
|
+
|
|
642
|
+
def _add_per_neuron_explanation_prompt(
|
|
643
|
+
self,
|
|
644
|
+
prompt_builder: PromptBuilder,
|
|
645
|
+
activation_records: Sequence[ActivationRecord],
|
|
646
|
+
index: int,
|
|
647
|
+
max_activation: float,
|
|
648
|
+
top_positive_logits: Optional[List[str]],
|
|
649
|
+
# When set, this indicates that the prompt should solicit a numbered list of the given
|
|
650
|
+
# number of explanations, rather than a single explanation.
|
|
651
|
+
numbered_list_of_n_explanations: Optional[int],
|
|
652
|
+
explanation: Optional[str], # None means this is the end of the full prompt.
|
|
653
|
+
) -> None:
|
|
654
|
+
max_activation = calculate_max_activation(activation_records)
|
|
655
|
+
user_message = f"""
|
|
656
|
+
|
|
657
|
+
Neuron {index + 1}
|
|
658
|
+
Activations:{format_activation_records(activation_records, max_activation, omit_zeros=False)}"""
|
|
659
|
+
# We repeat the non-zero activations only if it was requested and if the proportion of
|
|
660
|
+
# non-zero activations isn't too high.
|
|
661
|
+
if (
|
|
662
|
+
self.repeat_non_zero_activations
|
|
663
|
+
and non_zero_activation_proportion(activation_records, max_activation) < 0.2
|
|
664
|
+
):
|
|
665
|
+
user_message += (
|
|
666
|
+
f"\nSame activations, but with all zeros filtered out:"
|
|
667
|
+
f"{format_activation_records(activation_records, max_activation, omit_zeros=True)}"
|
|
668
|
+
)
|
|
669
|
+
|
|
670
|
+
user_message += f"\nTop Positive Logits: {top_positive_logits}\n"
|
|
671
|
+
|
|
672
|
+
if numbered_list_of_n_explanations is None:
|
|
673
|
+
user_message += f"\nExplanation of neuron {index + 1} behavior:"
|
|
674
|
+
assistant_message = ""
|
|
675
|
+
# For the IF format, we want <|endofprompt|> to come before the explanation prefix.
|
|
676
|
+
if self.prompt_format == PromptFormat.INSTRUCTION_FOLLOWING:
|
|
677
|
+
assistant_message += f" {EXPLANATION_PREFIX_LOGITS}"
|
|
678
|
+
else:
|
|
679
|
+
user_message += f" {EXPLANATION_PREFIX_LOGITS}"
|
|
680
|
+
prompt_builder.add_message(Role.USER, user_message)
|
|
681
|
+
|
|
682
|
+
if explanation is not None:
|
|
683
|
+
assistant_message += f" {explanation}."
|
|
684
|
+
if assistant_message:
|
|
685
|
+
prompt_builder.add_message(Role.ASSISTANT, assistant_message)
|
|
686
|
+
else:
|
|
687
|
+
if explanation is None:
|
|
688
|
+
# For the final neuron, we solicit a numbered list of explanations.
|
|
689
|
+
prompt_builder.add_message(
|
|
690
|
+
Role.USER,
|
|
691
|
+
f"""\nHere are {numbered_list_of_n_explanations} possible explanations for neuron {index + 1} behavior:\n1. {EXPLANATION_PREFIX_LOGITS}""",
|
|
692
|
+
)
|
|
693
|
+
else:
|
|
694
|
+
# For the few-shot examples, we only present one explanation, but we present it as a
|
|
695
|
+
# numbered list.
|
|
696
|
+
prompt_builder.add_message(
|
|
697
|
+
Role.USER,
|
|
698
|
+
f"""\nHere is 1 possible explanation for neuron {index + 1} behavior:\n1. {EXPLANATION_PREFIX_LOGITS}""",
|
|
699
|
+
)
|
|
700
|
+
prompt_builder.add_message(Role.ASSISTANT, f" {explanation}.")
|
|
701
|
+
|
|
702
|
+
def postprocess_explanations(
|
|
703
|
+
self, completions: list[str], prompt_kwargs: dict[str, Any]
|
|
704
|
+
) -> list[Any]:
|
|
705
|
+
"""Postprocess the explanations returned by the API"""
|
|
706
|
+
numbered_list_of_n_explanations = prompt_kwargs.get(
|
|
707
|
+
"numbered_list_of_n_explanations"
|
|
708
|
+
)
|
|
709
|
+
if numbered_list_of_n_explanations is None:
|
|
710
|
+
return completions
|
|
711
|
+
else:
|
|
712
|
+
all_explanations = []
|
|
713
|
+
for completion in completions:
|
|
714
|
+
for explanation in _split_numbered_list(completion):
|
|
715
|
+
explanation = self.strip_explanation(explanation)
|
|
716
|
+
if explanation.startswith(EXPLANATION_PREFIX_LOGITS):
|
|
717
|
+
explanation = explanation[len(EXPLANATION_PREFIX_LOGITS) :]
|
|
718
|
+
all_explanations.append(explanation)
|
|
719
|
+
return all_explanations
|
|
720
|
+
|
|
721
|
+
|
|
722
|
+
class TokenActivationPairLogitsNewExplainer(NeuronExplainer):
|
|
723
|
+
"""
|
|
724
|
+
Generate explanations of neuron behavior using a prompt with lists of token/activation pairs, with these changes:
|
|
725
|
+
- Don't tell the model to not specify specific words.
|
|
726
|
+
- Adding the top positive logits to the prompt.
|
|
727
|
+
- Telling the model to keep the explanation concise.
|
|
728
|
+
- Telling the model sometimes the neuron activates right before a specific word, token, or phrase, and to explain this with the format: 'say [the specific word, token or phrase]'
|
|
729
|
+
- Postprocessing will strip the ending period from the explanation.
|
|
730
|
+
- The additional explanation prefix is now "this neuron activates for".
|
|
731
|
+
"""
|
|
732
|
+
|
|
733
|
+
def __init__(
|
|
734
|
+
self,
|
|
735
|
+
model_name: str,
|
|
736
|
+
prompt_format: PromptFormat = PromptFormat.HARMONY_V4,
|
|
737
|
+
# This parameter lets us adjust the length of the prompt when we're generating explanations
|
|
738
|
+
# using older models with shorter context windows. In the future we can use it to experiment
|
|
739
|
+
# with 8k+ context windows.
|
|
740
|
+
context_size: ContextSize = ContextSize.ONETWENTYEIGHT_K,
|
|
741
|
+
few_shot_example_set: FewShotExampleSet = FewShotExampleSet.LOGITS,
|
|
742
|
+
repeat_non_zero_activations: bool = False,
|
|
743
|
+
max_concurrent: Optional[int] = 10,
|
|
744
|
+
cache: bool = False,
|
|
745
|
+
base_api_url: str = ApiClient.BASE_API_URL,
|
|
746
|
+
override_api_key: str | None = None,
|
|
747
|
+
):
|
|
748
|
+
super().__init__(
|
|
749
|
+
model_name=model_name,
|
|
750
|
+
prompt_format=prompt_format,
|
|
751
|
+
max_concurrent=max_concurrent,
|
|
752
|
+
cache=cache,
|
|
753
|
+
base_api_url=base_api_url,
|
|
754
|
+
override_api_key=override_api_key,
|
|
755
|
+
)
|
|
756
|
+
self.context_size = context_size
|
|
757
|
+
self.few_shot_example_set = few_shot_example_set
|
|
758
|
+
self.repeat_non_zero_activations = repeat_non_zero_activations
|
|
759
|
+
|
|
760
|
+
def make_explanation_prompt(
|
|
761
|
+
self, **kwargs: Any
|
|
762
|
+
) -> Union[str, list[HarmonyMessage]]:
|
|
763
|
+
original_kwargs = kwargs.copy()
|
|
764
|
+
all_activation_records: Sequence[ActivationRecord] = kwargs.pop(
|
|
765
|
+
"all_activation_records"
|
|
766
|
+
)
|
|
767
|
+
max_activation: float = kwargs.pop("max_activation")
|
|
768
|
+
top_positive_logits: List[str] = kwargs.pop("top_positive_logits")
|
|
769
|
+
kwargs.setdefault("numbered_list_of_n_explanations", None)
|
|
770
|
+
numbered_list_of_n_explanations: Optional[int] = kwargs.pop(
|
|
771
|
+
"numbered_list_of_n_explanations"
|
|
772
|
+
)
|
|
773
|
+
if numbered_list_of_n_explanations is not None:
|
|
774
|
+
assert numbered_list_of_n_explanations > 0, numbered_list_of_n_explanations
|
|
775
|
+
# This parameter lets us dynamically shrink the prompt if our initial attempt to create it
|
|
776
|
+
# results in something that's too long. It's only implemented for the 4k context size.
|
|
777
|
+
kwargs.setdefault("omit_n_activation_records", 0)
|
|
778
|
+
omit_n_activation_records: int = kwargs.pop("omit_n_activation_records")
|
|
779
|
+
max_tokens_for_completion: int = kwargs.pop("max_tokens_for_completion")
|
|
780
|
+
assert not kwargs, f"Unexpected kwargs: {kwargs}"
|
|
781
|
+
|
|
782
|
+
prompt_builder = PromptBuilder()
|
|
783
|
+
prompt_builder.add_message(
|
|
784
|
+
Role.SYSTEM,
|
|
785
|
+
"You are explaining the behavior of a neuron in a neural network. Your response should be a single short phrase (< 10 words) or word that summarizes what the neuron is looking for or predicting.\n\n"
|
|
786
|
+
"There are two types of responses you can give. Think carefully about which one to respond with.\n\n"
|
|
787
|
+
"Your first option is to respond with 'say [a specific word, phrase, or pattern, like 'a word that starts with a certain letter']'.\n\n"
|
|
788
|
+
"Your second option is to respond with '[a specific word, phrase, or pattern, like 'a word that starts with a certain letter']'.\n\n"
|
|
789
|
+
"To determine your response, you are given two types of information:\n\n"
|
|
790
|
+
"1. DOCUMENTS, which are split up into tokens. The document activation format is token<tab>activation. Activation "
|
|
791
|
+
"values range from 0 to 10. A neuron finding what it's looking for is represented by a "
|
|
792
|
+
"non-zero activation value. The higher the activation value, the stronger the match.\n\n"
|
|
793
|
+
"2. TOP POSITIVE LOGITS, which are the most likely word or token associated with this neuron.\n\n"
|
|
794
|
+
"How you should think:\n"
|
|
795
|
+
"1. Look at the tokens immediately AFTER the highest activating tokens in the DOCUMENTS. If these tokens seem to have a pattern or similarity, like 'starts with a certain letter', then respond with 'say [the predicted text or pattern]' and end there.\n"
|
|
796
|
+
"2. Look at the highest activating tokens and their context in the DOCUMENTS. If these tokens seem to have a pattern, VERY BRIEFLY describe what the neuron is looking for or predicting in the context of the DOCUMENTS.\n"
|
|
797
|
+
"3. Look at both the TOP POSITIVE LOGITS and the DOCUMENTS together. Try to find some similarity in them, and VERY BRIEFLY respond with the most likely option.\n\n"
|
|
798
|
+
"Your explanation should not be a full sentence - it should be very concise, and should not include unnecessary "
|
|
799
|
+
"phrases like 'the neuron is looking for' or 'the neuron predicts the word' or 'words related to' or 'variations of the word' "
|
|
800
|
+
"or 'concepts related to', or 'the word' etc. Simply say what it is the neuron is looking for or predicting, which can be "
|
|
801
|
+
"as short as a single word. If the neuron or pattern or prediction is a single word, then just say that word only.",
|
|
802
|
+
)
|
|
803
|
+
few_shot_examples = self.few_shot_example_set.get_examples()
|
|
804
|
+
num_omitted_activation_records = 0
|
|
805
|
+
for i, few_shot_example in enumerate(few_shot_examples):
|
|
806
|
+
few_shot_activation_records = few_shot_example.activation_records
|
|
807
|
+
self._add_per_neuron_explanation_prompt(
|
|
808
|
+
prompt_builder,
|
|
809
|
+
few_shot_activation_records,
|
|
810
|
+
i,
|
|
811
|
+
calculate_max_activation(few_shot_example.activation_records),
|
|
812
|
+
numbered_list_of_n_explanations=numbered_list_of_n_explanations,
|
|
813
|
+
top_positive_logits=few_shot_example.top_positive_logits,
|
|
814
|
+
explanation=few_shot_example.explanation,
|
|
815
|
+
)
|
|
816
|
+
self._add_per_neuron_explanation_prompt(
|
|
817
|
+
prompt_builder,
|
|
818
|
+
all_activation_records,
|
|
819
|
+
len(few_shot_examples),
|
|
820
|
+
max_activation,
|
|
821
|
+
numbered_list_of_n_explanations=numbered_list_of_n_explanations,
|
|
822
|
+
top_positive_logits=top_positive_logits,
|
|
823
|
+
explanation=None,
|
|
824
|
+
)
|
|
825
|
+
# If the prompt is too long *and* we omitted the specified number of activation records, try
|
|
826
|
+
# again, omitting one more. (If we didn't make the specified number of omissions, we're out
|
|
827
|
+
# of opportunities to omit records, so we just return the prompt as-is.)
|
|
828
|
+
if (
|
|
829
|
+
self._prompt_is_too_long(prompt_builder, max_tokens_for_completion)
|
|
830
|
+
and num_omitted_activation_records == omit_n_activation_records
|
|
831
|
+
):
|
|
832
|
+
original_kwargs["omit_n_activation_records"] = omit_n_activation_records + 1
|
|
833
|
+
return self.make_explanation_prompt(**original_kwargs)
|
|
834
|
+
return prompt_builder.build(self.prompt_format)
|
|
835
|
+
|
|
836
|
+
def _add_per_neuron_explanation_prompt(
|
|
837
|
+
self,
|
|
838
|
+
prompt_builder: PromptBuilder,
|
|
839
|
+
activation_records: Sequence[ActivationRecord],
|
|
840
|
+
index: int,
|
|
841
|
+
max_activation: float,
|
|
842
|
+
top_positive_logits: Optional[List[str]],
|
|
843
|
+
# When set, this indicates that the prompt should solicit a numbered list of the given
|
|
844
|
+
# number of explanations, rather than a single explanation.
|
|
845
|
+
numbered_list_of_n_explanations: Optional[int],
|
|
846
|
+
explanation: Optional[str], # None means this is the end of the full prompt.
|
|
847
|
+
) -> None:
|
|
848
|
+
max_activation = calculate_max_activation(activation_records)
|
|
849
|
+
user_message = f"""
|
|
850
|
+
|
|
851
|
+
Neuron {index + 1}
|
|
852
|
+
|
|
853
|
+
[START DOCUMENTS]
|
|
854
|
+
|
|
855
|
+
Activations:{format_activation_records(activation_records, max_activation, omit_zeros=False)}
|
|
856
|
+
|
|
857
|
+
[END DOCUMENTS]"""
|
|
858
|
+
# We repeat the non-zero activations only if it was requested and if the proportion of
|
|
859
|
+
# non-zero activations isn't too high.
|
|
860
|
+
if (
|
|
861
|
+
self.repeat_non_zero_activations
|
|
862
|
+
and non_zero_activation_proportion(activation_records, max_activation) < 0.2
|
|
863
|
+
):
|
|
864
|
+
user_message += (
|
|
865
|
+
f"\nSame activations, but with all zeros filtered out:"
|
|
866
|
+
f"{format_activation_records(activation_records, max_activation, omit_zeros=True)}"
|
|
867
|
+
)
|
|
868
|
+
|
|
869
|
+
user_message += f"\n\n[START TOP POSITIVE LOGITS]\n{top_positive_logits}\n[END TOP POSITIVE LOGITS]\n\n"
|
|
870
|
+
|
|
871
|
+
if numbered_list_of_n_explanations is None:
|
|
872
|
+
user_message += f"\nExplanation of neuron {index + 1} behavior:"
|
|
873
|
+
assistant_message = ""
|
|
874
|
+
# For the IF format, we want <|endofprompt|> to come before the explanation prefix.
|
|
875
|
+
if self.prompt_format == PromptFormat.INSTRUCTION_FOLLOWING:
|
|
876
|
+
assistant_message += f" {EXPLANATION_PREFIX_LOGITS}"
|
|
877
|
+
else:
|
|
878
|
+
user_message += f" {EXPLANATION_PREFIX_LOGITS}"
|
|
879
|
+
prompt_builder.add_message(Role.USER, user_message)
|
|
880
|
+
|
|
881
|
+
if explanation is not None:
|
|
882
|
+
assistant_message += f" {explanation}."
|
|
883
|
+
if assistant_message:
|
|
884
|
+
prompt_builder.add_message(Role.ASSISTANT, assistant_message)
|
|
885
|
+
else:
|
|
886
|
+
prompt_builder.add_message(
|
|
887
|
+
Role.USER,
|
|
888
|
+
f"""\nExplanation for neuron {index + 1} behavior: {EXPLANATION_PREFIX_LOGITS}""",
|
|
889
|
+
)
|
|
890
|
+
if explanation is not None:
|
|
891
|
+
prompt_builder.add_message(Role.ASSISTANT, f"{explanation}")
|
|
892
|
+
|
|
893
|
+
def postprocess_explanations(
|
|
894
|
+
self, completions: list[str], prompt_kwargs: dict[str, Any]
|
|
895
|
+
) -> list[Any]:
|
|
896
|
+
"""Postprocess the explanations returned by the API"""
|
|
897
|
+
numbered_list_of_n_explanations = prompt_kwargs.get(
|
|
898
|
+
"numbered_list_of_n_explanations"
|
|
899
|
+
)
|
|
900
|
+
if numbered_list_of_n_explanations is None:
|
|
901
|
+
return completions
|
|
902
|
+
else:
|
|
903
|
+
all_explanations = []
|
|
904
|
+
for completion in completions:
|
|
905
|
+
for explanation in _split_numbered_list(completion):
|
|
906
|
+
explanation = self.strip_explanation(explanation)
|
|
907
|
+
if explanation.startswith(EXPLANATION_PREFIX_LOGITS):
|
|
908
|
+
explanation = explanation[len(EXPLANATION_PREFIX_LOGITS) :]
|
|
909
|
+
all_explanations.append(explanation)
|
|
910
|
+
return all_explanations
|
|
911
|
+
|
|
912
|
+
|
|
913
|
+
class MaxActivationAndLogitsExplainer(NeuronExplainer):
|
|
914
|
+
"""
|
|
915
|
+
This is a very concise explainer (1 to 6 words) that attempts to replicate Anthropic's attribution graphs explainer.
|
|
916
|
+
It shows the model both activations and top positive logits.
|
|
917
|
+
This explainer is expected to be used for the last 1/3 of layers in a model, since it has heavy focus on predicting the next token.
|
|
918
|
+
This explainer's tries to explain using one of these options:
|
|
919
|
+
- "say [the next predicted token after the max activating token]"
|
|
920
|
+
- a brief description of the max activating token (which can simply be the max activating token itself)
|
|
921
|
+
- a brief description of the top positive logits
|
|
922
|
+
- a brief description of the top activating texts
|
|
923
|
+
|
|
924
|
+
We force the explainer to try and explain using each method, then return when an explanation is found. Forcing it to do this made the explanations much more accurate.
|
|
925
|
+
|
|
926
|
+
See make_explanation_prompt below for the full prompt.
|
|
927
|
+
|
|
928
|
+
A weakness of this explainer is that it is less good at explaining the whole context - more for immediate words/characters on or after the top activating token.
|
|
929
|
+
You can increase the "tokens_around_max_activating_token" to try to improve this behavior.
|
|
930
|
+
|
|
931
|
+
We mostly tested using this explainer with Gemini-2.0-Flash.
|
|
932
|
+
"""
|
|
933
|
+
|
|
934
|
+
def __init__(
|
|
935
|
+
self,
|
|
936
|
+
model_name: str,
|
|
937
|
+
prompt_format: PromptFormat = PromptFormat.HARMONY_V4,
|
|
938
|
+
context_size: ContextSize = ContextSize.ONETWENTYEIGHT_K,
|
|
939
|
+
few_shot_example_set: FewShotExampleSet = FewShotExampleSet.LOGITS,
|
|
940
|
+
tokens_around_max_activating_token: int = 24,
|
|
941
|
+
repeat_non_zero_activations: bool = False,
|
|
942
|
+
max_concurrent: Optional[int] = 10,
|
|
943
|
+
cache: bool = False,
|
|
944
|
+
base_api_url: str = ApiClient.BASE_API_URL,
|
|
945
|
+
override_api_key: str | None = None,
|
|
946
|
+
):
|
|
947
|
+
super().__init__(
|
|
948
|
+
model_name=model_name,
|
|
949
|
+
prompt_format=prompt_format,
|
|
950
|
+
max_concurrent=max_concurrent,
|
|
951
|
+
cache=cache,
|
|
952
|
+
base_api_url=base_api_url,
|
|
953
|
+
override_api_key=override_api_key,
|
|
954
|
+
)
|
|
955
|
+
self.context_size = context_size
|
|
956
|
+
self.few_shot_example_set = few_shot_example_set
|
|
957
|
+
self.repeat_non_zero_activations = repeat_non_zero_activations
|
|
958
|
+
self.tokens_around_max_activating_token = tokens_around_max_activating_token
|
|
959
|
+
|
|
960
|
+
def format_tokens_after_max_activating_token(
|
|
961
|
+
self, activation_records: Sequence[ActivationRecord]
|
|
962
|
+
) -> str:
|
|
963
|
+
"""
|
|
964
|
+
Format the tokens immediately after the max activating token.
|
|
965
|
+
"""
|
|
966
|
+
formatted_texts = []
|
|
967
|
+
for record in activation_records:
|
|
968
|
+
tokens = record.tokens
|
|
969
|
+
activations = record.activations
|
|
970
|
+
max_activation_index = activations.index(max(activations))
|
|
971
|
+
# Only get the first token after the max activating token
|
|
972
|
+
if max_activation_index + 1 < len(tokens):
|
|
973
|
+
token_after_max_activating_token = (
|
|
974
|
+
tokens[max_activation_index + 1].replace("\n", "").strip()
|
|
975
|
+
)
|
|
976
|
+
formatted_texts.append(f"{token_after_max_activating_token}")
|
|
977
|
+
else:
|
|
978
|
+
# Handle case where max activation is the last token
|
|
979
|
+
formatted_texts.append("")
|
|
980
|
+
return "\n".join(formatted_texts)
|
|
981
|
+
|
|
982
|
+
def format_max_activating_tokens(
|
|
983
|
+
self, activation_records: Sequence[ActivationRecord]
|
|
984
|
+
) -> str:
|
|
985
|
+
"""
|
|
986
|
+
Format the max activating tokens.
|
|
987
|
+
"""
|
|
988
|
+
formatted_tokens = []
|
|
989
|
+
for record in activation_records:
|
|
990
|
+
tokens = record.tokens
|
|
991
|
+
activations = record.activations
|
|
992
|
+
max_activation_index = activations.index(max(activations))
|
|
993
|
+
max_activating_token = (
|
|
994
|
+
tokens[max_activation_index].replace("\n", "").strip()
|
|
995
|
+
)
|
|
996
|
+
formatted_tokens.append(f"{max_activating_token}")
|
|
997
|
+
return "\n".join(formatted_tokens)
|
|
998
|
+
|
|
999
|
+
def format_top_activating_texts(
|
|
1000
|
+
self, activation_records: Sequence[ActivationRecord]
|
|
1001
|
+
) -> str:
|
|
1002
|
+
"""
|
|
1003
|
+
Format activation records into a bullet point list of texts, with each text trimmed to
|
|
1004
|
+
8 tokens to the left and right of the maximum activating token. Replace line breaks with two spaces.
|
|
1005
|
+
"""
|
|
1006
|
+
formatted_texts = []
|
|
1007
|
+
|
|
1008
|
+
for record in activation_records:
|
|
1009
|
+
tokens = record.tokens
|
|
1010
|
+
activations = record.activations
|
|
1011
|
+
|
|
1012
|
+
# Find the index of the maximum activation
|
|
1013
|
+
max_activation_index = activations.index(max(activations))
|
|
1014
|
+
|
|
1015
|
+
# Calculate the start and end indices for the window
|
|
1016
|
+
start_index = max(
|
|
1017
|
+
0, max_activation_index - self.tokens_around_max_activating_token
|
|
1018
|
+
)
|
|
1019
|
+
end_index = min(
|
|
1020
|
+
len(tokens),
|
|
1021
|
+
max_activation_index + self.tokens_around_max_activating_token + 1,
|
|
1022
|
+
) # +1 to include the token at end_index-1
|
|
1023
|
+
|
|
1024
|
+
# Create the trimmed text with the max activating token surrounded by ^^
|
|
1025
|
+
trimmed_tokens = (
|
|
1026
|
+
tokens[start_index:max_activation_index]
|
|
1027
|
+
+ [f"{tokens[max_activation_index]}"]
|
|
1028
|
+
+ tokens[max_activation_index + 1 : end_index]
|
|
1029
|
+
)
|
|
1030
|
+
|
|
1031
|
+
trimmed_text = "".join(trimmed_tokens).replace("\n", " ")
|
|
1032
|
+
formatted_texts.append(f"{trimmed_text}")
|
|
1033
|
+
|
|
1034
|
+
return "\n".join(formatted_texts)
|
|
1035
|
+
|
|
1036
|
+
def make_explanation_prompt(
|
|
1037
|
+
self, **kwargs: Any
|
|
1038
|
+
) -> Union[str, list[HarmonyMessage]]:
|
|
1039
|
+
original_kwargs = kwargs.copy()
|
|
1040
|
+
all_activation_records: Sequence[ActivationRecord] = kwargs.pop(
|
|
1041
|
+
"all_activation_records"
|
|
1042
|
+
)
|
|
1043
|
+
# Replace all ▁ characters in tokens with spaces
|
|
1044
|
+
processed_activation_records = []
|
|
1045
|
+
for record in all_activation_records:
|
|
1046
|
+
# Create a new ActivationRecord with processed tokens
|
|
1047
|
+
processed_tokens = [token.replace("▁", " ") for token in record.tokens]
|
|
1048
|
+
processed_activation_records.append(
|
|
1049
|
+
ActivationRecord(
|
|
1050
|
+
tokens=processed_tokens, activations=record.activations
|
|
1051
|
+
)
|
|
1052
|
+
)
|
|
1053
|
+
|
|
1054
|
+
# Use the processed records for the rest of the function
|
|
1055
|
+
all_activation_records = processed_activation_records
|
|
1056
|
+
max_activation: float = kwargs.pop("max_activation")
|
|
1057
|
+
top_positive_logits: List[str] = kwargs.pop("top_positive_logits")
|
|
1058
|
+
|
|
1059
|
+
# Replace all ▁ characters in top_positive_logits with spaces
|
|
1060
|
+
processed_top_positive_logits = [
|
|
1061
|
+
logit.replace("▁", " ").replace("\n", "").strip()
|
|
1062
|
+
for logit in top_positive_logits
|
|
1063
|
+
]
|
|
1064
|
+
top_positive_logits = processed_top_positive_logits
|
|
1065
|
+
kwargs.setdefault("numbered_list_of_n_explanations", None)
|
|
1066
|
+
numbered_list_of_n_explanations: Optional[int] = kwargs.pop(
|
|
1067
|
+
"numbered_list_of_n_explanations"
|
|
1068
|
+
)
|
|
1069
|
+
if numbered_list_of_n_explanations is not None:
|
|
1070
|
+
assert numbered_list_of_n_explanations > 0, numbered_list_of_n_explanations
|
|
1071
|
+
# This parameter lets us dynamically shrink the prompt if our initial attempt to create it
|
|
1072
|
+
# results in something that's too long. It's only implemented for the 4k context size.
|
|
1073
|
+
kwargs.setdefault("omit_n_activation_records", 0)
|
|
1074
|
+
omit_n_activation_records: int = kwargs.pop("omit_n_activation_records")
|
|
1075
|
+
max_tokens_for_completion: int = kwargs.pop("max_tokens_for_completion")
|
|
1076
|
+
assert not kwargs, f"Unexpected kwargs: {kwargs}"
|
|
1077
|
+
|
|
1078
|
+
prompt_builder = PromptBuilder()
|
|
1079
|
+
# TODO: this is pretty verbose and can probably be shortened
|
|
1080
|
+
prompt_builder.add_message(
|
|
1081
|
+
Role.SYSTEM,
|
|
1082
|
+
"You are explaining the behavior of a neuron in a neural network. Your response should be a very concise explanation (1-6 words) that captures what the neuron detects or predicts by finding patterns in lists.\n\n"
|
|
1083
|
+
"To determine the explanation, you are given four lists:\n\n"
|
|
1084
|
+
"- MAX_ACTIVATING_TOKENS, which are the top activating tokens in the top activating texts.\n"
|
|
1085
|
+
"- TOKENS_AFTER_MAX_ACTIVATING_TOKEN, which are the tokens immediately after the max activating token.\n"
|
|
1086
|
+
"- TOP_POSITIVE_LOGITS, which are the most likely words or tokens associated with this neuron.\n"
|
|
1087
|
+
"- TOP_ACTIVATING_TEXTS, which are top activating texts.\n\n"
|
|
1088
|
+
"You should look for a pattern by trying the following methods in order. Once you find a pattern, stop and return that pattern. Do not proceed to the later methods.\n"
|
|
1089
|
+
"Method 1: Look at MAX_ACTIVATING_TOKENS. If they share something specific in common, or are all the same token or a variation of the same token (like different cases or conjugations), respond with that token.\n"
|
|
1090
|
+
"Method 2: Look at TOKENS_AFTER_MAX_ACTIVATING_TOKEN. Try to find a specific pattern or similarity in all the tokens. A common pattern is that they all start with the same letter. If you find a pattern (like 's word', 'the ending -ing', 'number 8'), respond with 'say [the pattern]'. You can ignore uppercase/lowercase differences for this.\n"
|
|
1091
|
+
"Method 3: Look at TOP_POSITIVE_LOGITS for similarities and describe it very briefly (1-3 words).\n"
|
|
1092
|
+
"Method 4: Look at TOP_ACTIVATING_TEXTS and make a best guess by describing the broad theme or context, ignoring the max activating tokens.\n\n"
|
|
1093
|
+
"Rules:\n"
|
|
1094
|
+
"- Keep your explanation extremely concise (1-6 words, mostly 1-3 words).\n"
|
|
1095
|
+
'- Do not add unnecessary phrases like "words related to", "concepts related to", or "variations of the word".\n'
|
|
1096
|
+
'- Do not mention "tokens" or "patterns" in your explanation.\n'
|
|
1097
|
+
'- The explanation should be specific. For example, "unique words" is not a specific enough pattern, nor is "foreign words".\n'
|
|
1098
|
+
"- Remember to use the 'say [the pattern]' when using Method 2 above (pattern found in TOKENS_AFTER_MAX_ACTIVATING_TOKEN).\n"
|
|
1099
|
+
"- If you absolutely cannot make any guesses, return the first token in MAX_ACTIVATING_TOKENS.\n\n"
|
|
1100
|
+
"Respond by going through each method number until you find one that helps you find an explanation for what this neuron is detecting or predicting. If a method does not help you find an explanation, briefly explain why it does not, then go on to the next method. "
|
|
1101
|
+
"Finally, end your response with the method number you used, the reason for your explanation, and then the explanation.",
|
|
1102
|
+
)
|
|
1103
|
+
few_shot_examples = self.few_shot_example_set.get_examples()
|
|
1104
|
+
num_omitted_activation_records = 0
|
|
1105
|
+
for i, few_shot_example in enumerate(few_shot_examples):
|
|
1106
|
+
few_shot_activation_records = few_shot_example.activation_records
|
|
1107
|
+
self._add_per_neuron_explanation_prompt(
|
|
1108
|
+
prompt_builder,
|
|
1109
|
+
few_shot_activation_records,
|
|
1110
|
+
i,
|
|
1111
|
+
calculate_max_activation(few_shot_example.activation_records),
|
|
1112
|
+
numbered_list_of_n_explanations=numbered_list_of_n_explanations,
|
|
1113
|
+
top_positive_logits=few_shot_example.top_positive_logits,
|
|
1114
|
+
explanation=few_shot_example.explanation,
|
|
1115
|
+
)
|
|
1116
|
+
self._add_per_neuron_explanation_prompt(
|
|
1117
|
+
prompt_builder,
|
|
1118
|
+
all_activation_records,
|
|
1119
|
+
len(few_shot_examples),
|
|
1120
|
+
max_activation,
|
|
1121
|
+
numbered_list_of_n_explanations=numbered_list_of_n_explanations,
|
|
1122
|
+
top_positive_logits=top_positive_logits,
|
|
1123
|
+
explanation=None,
|
|
1124
|
+
)
|
|
1125
|
+
# If the prompt is too long *and* we omitted the specified number of activation records, try
|
|
1126
|
+
# again, omitting one more. (If we didn't make the specified number of omissions, we're out
|
|
1127
|
+
# of opportunities to omit records, so we just return the prompt as-is.)
|
|
1128
|
+
if (
|
|
1129
|
+
self._prompt_is_too_long(prompt_builder, max_tokens_for_completion)
|
|
1130
|
+
and num_omitted_activation_records == omit_n_activation_records
|
|
1131
|
+
):
|
|
1132
|
+
original_kwargs["omit_n_activation_records"] = omit_n_activation_records + 1
|
|
1133
|
+
return self.make_explanation_prompt(**original_kwargs)
|
|
1134
|
+
built_prompt = prompt_builder.build(self.prompt_format)
|
|
1135
|
+
|
|
1136
|
+
# ## debug only
|
|
1137
|
+
# import json
|
|
1138
|
+
|
|
1139
|
+
# if isinstance(built_prompt, list):
|
|
1140
|
+
# logger.error(json.dumps({"built_prompt": built_prompt}))
|
|
1141
|
+
# else:
|
|
1142
|
+
# logger.error(json.dumps({"built_prompt": built_prompt}))
|
|
1143
|
+
# import sys
|
|
1144
|
+
|
|
1145
|
+
# sys.exit(1)
|
|
1146
|
+
|
|
1147
|
+
return built_prompt
|
|
1148
|
+
|
|
1149
|
+
def format_top_logits(self, top_positive_logits: List[str]) -> str:
|
|
1150
|
+
return "\n".join([f"{logit.strip()}" for logit in top_positive_logits])
|
|
1151
|
+
|
|
1152
|
+
def _add_per_neuron_explanation_prompt(
|
|
1153
|
+
self,
|
|
1154
|
+
prompt_builder: PromptBuilder,
|
|
1155
|
+
activation_records: Sequence[ActivationRecord],
|
|
1156
|
+
index: int,
|
|
1157
|
+
max_activation: float,
|
|
1158
|
+
top_positive_logits: Optional[List[str]],
|
|
1159
|
+
# When set, this indicates that the prompt should solicit a numbered list of the given
|
|
1160
|
+
# number of explanations, rather than a single explanation.
|
|
1161
|
+
numbered_list_of_n_explanations: Optional[int],
|
|
1162
|
+
explanation: Optional[str], # None means this is the end of the full prompt.
|
|
1163
|
+
) -> None:
|
|
1164
|
+
user_message = f"""
|
|
1165
|
+
|
|
1166
|
+
Neuron {index + 1}
|
|
1167
|
+
|
|
1168
|
+
<TOKENS_AFTER_MAX_ACTIVATING_TOKEN>
|
|
1169
|
+
|
|
1170
|
+
{self.format_tokens_after_max_activating_token(activation_records)}
|
|
1171
|
+
|
|
1172
|
+
</TOKENS_AFTER_MAX_ACTIVATING_TOKEN>
|
|
1173
|
+
|
|
1174
|
+
|
|
1175
|
+
<MAX_ACTIVATING_TOKENS>
|
|
1176
|
+
|
|
1177
|
+
{self.format_max_activating_tokens(activation_records)}
|
|
1178
|
+
|
|
1179
|
+
</MAX_ACTIVATING_TOKENS>
|
|
1180
|
+
|
|
1181
|
+
|
|
1182
|
+
<TOP_POSITIVE_LOGITS>
|
|
1183
|
+
|
|
1184
|
+
{self.format_top_logits(top_positive_logits) if top_positive_logits else ""}
|
|
1185
|
+
|
|
1186
|
+
</TOP_POSITIVE_LOGITS>
|
|
1187
|
+
|
|
1188
|
+
|
|
1189
|
+
<TOP_ACTIVATING_TEXTS>
|
|
1190
|
+
|
|
1191
|
+
{self.format_top_activating_texts(activation_records)}
|
|
1192
|
+
|
|
1193
|
+
</TOP_ACTIVATING_TEXTS>
|
|
1194
|
+
|
|
1195
|
+
"""
|
|
1196
|
+
# logger.error(f"user_message: {user_message}")
|
|
1197
|
+
|
|
1198
|
+
if numbered_list_of_n_explanations is None:
|
|
1199
|
+
user_message += f"\nExplanation of neuron {index + 1} behavior: "
|
|
1200
|
+
assistant_message = ""
|
|
1201
|
+
# # For the IF format, we want <|endofprompt|> to come before the explanation prefix.
|
|
1202
|
+
# if self.prompt_format == PromptFormat.INSTRUCTION_FOLLOWING:
|
|
1203
|
+
# assistant_message += f"{EXPLANATION_PREFIX_LOGITS}"
|
|
1204
|
+
# else:
|
|
1205
|
+
# user_message += f"{EXPLANATION_PREFIX_LOGITS}"
|
|
1206
|
+
prompt_builder.add_message(Role.USER, user_message)
|
|
1207
|
+
|
|
1208
|
+
if explanation is not None:
|
|
1209
|
+
assistant_message += f"{explanation}"
|
|
1210
|
+
if assistant_message:
|
|
1211
|
+
prompt_builder.add_message(Role.ASSISTANT, assistant_message)
|
|
1212
|
+
else:
|
|
1213
|
+
prompt_builder.add_message(
|
|
1214
|
+
Role.USER,
|
|
1215
|
+
f"""\nExplanation for neuron {index + 1} behavior: """,
|
|
1216
|
+
)
|
|
1217
|
+
if explanation is not None:
|
|
1218
|
+
prompt_builder.add_message(Role.ASSISTANT, f"{explanation}")
|
|
1219
|
+
|
|
1220
|
+
def postprocess_explanations(
|
|
1221
|
+
self, completions: list[str], prompt_kwargs: dict[str, Any]
|
|
1222
|
+
) -> list[Any]:
|
|
1223
|
+
"""Postprocess the explanations returned by the API"""
|
|
1224
|
+
numbered_list_of_n_explanations = prompt_kwargs.get(
|
|
1225
|
+
"numbered_list_of_n_explanations"
|
|
1226
|
+
)
|
|
1227
|
+
if numbered_list_of_n_explanations is None:
|
|
1228
|
+
all_explanations = []
|
|
1229
|
+
for explanation in completions:
|
|
1230
|
+
explanation = self.strip_explanation(explanation)
|
|
1231
|
+
# Split by "Explanation: " and take the last segment if it exists
|
|
1232
|
+
if "Explanation: " in explanation:
|
|
1233
|
+
explanation = explanation.split("Explanation: ")[-1]
|
|
1234
|
+
elif "explanation: " in explanation:
|
|
1235
|
+
explanation = explanation.split("explanation: ")[-1]
|
|
1236
|
+
else:
|
|
1237
|
+
logger.error(
|
|
1238
|
+
f"Error parsing response explanation, no explanation string found: {explanation}"
|
|
1239
|
+
)
|
|
1240
|
+
all_explanations.append("")
|
|
1241
|
+
continue
|
|
1242
|
+
|
|
1243
|
+
all_explanations.append(_remove_method_from_explanation(explanation))
|
|
1244
|
+
return all_explanations
|
|
1245
|
+
else:
|
|
1246
|
+
all_explanations = []
|
|
1247
|
+
for completion in completions:
|
|
1248
|
+
for explanation in _split_numbered_list(completion):
|
|
1249
|
+
explanation = self.strip_explanation(explanation)
|
|
1250
|
+
if explanation.endswith("."):
|
|
1251
|
+
explanation = explanation[:-1]
|
|
1252
|
+
# Split by "Explanation: " and take the last segment if it exists
|
|
1253
|
+
if "Explanation: " in explanation:
|
|
1254
|
+
explanation = explanation.split("Explanation: ")[-1]
|
|
1255
|
+
elif "explanation: " in explanation:
|
|
1256
|
+
explanation = explanation.split("explanation: ")[-1]
|
|
1257
|
+
else:
|
|
1258
|
+
logger.error(
|
|
1259
|
+
f"Error parsing response explanation, no explanation string found: {explanation}"
|
|
1260
|
+
)
|
|
1261
|
+
all_explanations.append("")
|
|
1262
|
+
continue
|
|
1263
|
+
|
|
1264
|
+
all_explanations.append(
|
|
1265
|
+
_remove_method_from_explanation(explanation)
|
|
1266
|
+
)
|
|
1267
|
+
return all_explanations
|
|
1268
|
+
|
|
1269
|
+
|
|
1270
|
+
class MaxActivationAndLogitsGeneralExplainer(NeuronExplainer):
|
|
1271
|
+
"""
|
|
1272
|
+
This is the MaxActivationAndLogitsExplainer, but with a more general explanation (5 to 20 words), not targeted for extreme conciseness. It also doesn't show the model examples.
|
|
1273
|
+
It shows the model both activations and top positive logits.
|
|
1274
|
+
This explainer is expected to be used for the last 1/3 of layers in a model, since it has heavy focus on predicting the next token.
|
|
1275
|
+
This explainer assumes you are using a more intelligent model (eg Gemini-2.0-Flash or above), since it gives more general instructions.
|
|
1276
|
+
|
|
1277
|
+
Method:
|
|
1278
|
+
- We show the top activating token in the context of each snippet
|
|
1279
|
+
- We show the top positive logits (and explain what this means)
|
|
1280
|
+
- We ask the model, in a single short phrase, to explain the behavior of this neuron.
|
|
1281
|
+
|
|
1282
|
+
See make_explanation_prompt below for the full prompt.
|
|
1283
|
+
|
|
1284
|
+
We mostly tested using this explainer with Gemini-2.0-Flash.
|
|
1285
|
+
"""
|
|
1286
|
+
|
|
1287
|
+
def __init__(
|
|
1288
|
+
self,
|
|
1289
|
+
model_name: str,
|
|
1290
|
+
prompt_format: PromptFormat = PromptFormat.HARMONY_V4,
|
|
1291
|
+
context_size: ContextSize = ContextSize.ONETWENTYEIGHT_K,
|
|
1292
|
+
few_shot_example_set: FewShotExampleSet = FewShotExampleSet.LOGITS,
|
|
1293
|
+
tokens_around_max_activating_token: int = 24,
|
|
1294
|
+
repeat_non_zero_activations: bool = False,
|
|
1295
|
+
max_concurrent: Optional[int] = 10,
|
|
1296
|
+
cache: bool = False,
|
|
1297
|
+
base_api_url: str = ApiClient.BASE_API_URL,
|
|
1298
|
+
override_api_key: str | None = None,
|
|
1299
|
+
):
|
|
1300
|
+
super().__init__(
|
|
1301
|
+
model_name=model_name,
|
|
1302
|
+
prompt_format=prompt_format,
|
|
1303
|
+
max_concurrent=max_concurrent,
|
|
1304
|
+
cache=cache,
|
|
1305
|
+
base_api_url=base_api_url,
|
|
1306
|
+
override_api_key=override_api_key,
|
|
1307
|
+
)
|
|
1308
|
+
self.context_size = context_size
|
|
1309
|
+
self.few_shot_example_set = few_shot_example_set
|
|
1310
|
+
self.repeat_non_zero_activations = repeat_non_zero_activations
|
|
1311
|
+
self.tokens_around_max_activating_token = tokens_around_max_activating_token
|
|
1312
|
+
|
|
1313
|
+
def format_tokens_after_max_activating_token(
|
|
1314
|
+
self, activation_records: Sequence[ActivationRecord]
|
|
1315
|
+
) -> str:
|
|
1316
|
+
"""
|
|
1317
|
+
Format the tokens immediately after the max activating token.
|
|
1318
|
+
"""
|
|
1319
|
+
formatted_texts = []
|
|
1320
|
+
for record in activation_records:
|
|
1321
|
+
tokens = record.tokens
|
|
1322
|
+
activations = record.activations
|
|
1323
|
+
max_activation_index = activations.index(max(activations))
|
|
1324
|
+
# Only get the first token after the max activating token
|
|
1325
|
+
if max_activation_index + 1 < len(tokens):
|
|
1326
|
+
token_after_max_activating_token = (
|
|
1327
|
+
tokens[max_activation_index + 1].replace("\n", "").strip()
|
|
1328
|
+
)
|
|
1329
|
+
formatted_texts.append(f"{token_after_max_activating_token}")
|
|
1330
|
+
else:
|
|
1331
|
+
# Handle case where max activation is the last token
|
|
1332
|
+
formatted_texts.append("")
|
|
1333
|
+
return "\n".join(formatted_texts)
|
|
1334
|
+
|
|
1335
|
+
def format_max_activating_tokens(
|
|
1336
|
+
self, activation_records: Sequence[ActivationRecord]
|
|
1337
|
+
) -> str:
|
|
1338
|
+
"""
|
|
1339
|
+
Format the max activating tokens.
|
|
1340
|
+
"""
|
|
1341
|
+
formatted_tokens = []
|
|
1342
|
+
for record in activation_records:
|
|
1343
|
+
tokens = record.tokens
|
|
1344
|
+
activations = record.activations
|
|
1345
|
+
max_activation_index = activations.index(max(activations))
|
|
1346
|
+
max_activating_token = (
|
|
1347
|
+
tokens[max_activation_index].replace("\n", "").strip()
|
|
1348
|
+
)
|
|
1349
|
+
formatted_tokens.append(f"{max_activating_token}")
|
|
1350
|
+
return "\n".join(formatted_tokens)
|
|
1351
|
+
|
|
1352
|
+
def format_top_activating_texts(
|
|
1353
|
+
self, activation_records: Sequence[ActivationRecord]
|
|
1354
|
+
) -> str:
|
|
1355
|
+
"""
|
|
1356
|
+
Format activation records into a bullet point list of texts, with each text trimmed to
|
|
1357
|
+
8 tokens to the left and right of the maximum activating token. Replace line breaks with two spaces.
|
|
1358
|
+
"""
|
|
1359
|
+
formatted_texts = []
|
|
1360
|
+
|
|
1361
|
+
for record in activation_records:
|
|
1362
|
+
tokens = record.tokens
|
|
1363
|
+
activations = record.activations
|
|
1364
|
+
|
|
1365
|
+
# Find the index of the maximum activation
|
|
1366
|
+
max_activation_index = activations.index(max(activations))
|
|
1367
|
+
|
|
1368
|
+
# Calculate the start and end indices for the window
|
|
1369
|
+
start_index = max(
|
|
1370
|
+
0, max_activation_index - self.tokens_around_max_activating_token
|
|
1371
|
+
)
|
|
1372
|
+
end_index = min(
|
|
1373
|
+
len(tokens),
|
|
1374
|
+
max_activation_index + self.tokens_around_max_activating_token + 1,
|
|
1375
|
+
) # +1 to include the token at end_index-1
|
|
1376
|
+
|
|
1377
|
+
# Create the trimmed text with the max activating token surrounded by ^^
|
|
1378
|
+
trimmed_tokens = (
|
|
1379
|
+
tokens[start_index:max_activation_index]
|
|
1380
|
+
+ [f"{tokens[max_activation_index]}"]
|
|
1381
|
+
+ tokens[max_activation_index + 1 : end_index]
|
|
1382
|
+
)
|
|
1383
|
+
|
|
1384
|
+
trimmed_text = "".join(trimmed_tokens).replace("\n", " ")
|
|
1385
|
+
formatted_texts.append(f"{trimmed_text}")
|
|
1386
|
+
|
|
1387
|
+
return "\n".join(formatted_texts)
|
|
1388
|
+
|
|
1389
|
+
def make_explanation_prompt(
|
|
1390
|
+
self, **kwargs: Any
|
|
1391
|
+
) -> Union[str, list[HarmonyMessage]]:
|
|
1392
|
+
original_kwargs = kwargs.copy()
|
|
1393
|
+
all_activation_records: Sequence[ActivationRecord] = kwargs.pop(
|
|
1394
|
+
"all_activation_records"
|
|
1395
|
+
)
|
|
1396
|
+
# Replace all ▁ characters in tokens with spaces
|
|
1397
|
+
processed_activation_records = []
|
|
1398
|
+
for record in all_activation_records:
|
|
1399
|
+
# Create a new ActivationRecord with processed tokens
|
|
1400
|
+
processed_tokens = [token.replace("▁", " ") for token in record.tokens]
|
|
1401
|
+
processed_activation_records.append(
|
|
1402
|
+
ActivationRecord(
|
|
1403
|
+
tokens=processed_tokens, activations=record.activations
|
|
1404
|
+
)
|
|
1405
|
+
)
|
|
1406
|
+
|
|
1407
|
+
# Use the processed records for the rest of the function
|
|
1408
|
+
all_activation_records = processed_activation_records
|
|
1409
|
+
max_activation: float = kwargs.pop("max_activation")
|
|
1410
|
+
top_positive_logits: List[str] = kwargs.pop("top_positive_logits")
|
|
1411
|
+
|
|
1412
|
+
# Replace all ▁ characters in top_positive_logits with spaces
|
|
1413
|
+
processed_top_positive_logits = [
|
|
1414
|
+
logit.replace("▁", " ").replace("\n", "").strip()
|
|
1415
|
+
for logit in top_positive_logits
|
|
1416
|
+
]
|
|
1417
|
+
top_positive_logits = processed_top_positive_logits
|
|
1418
|
+
kwargs.setdefault("numbered_list_of_n_explanations", None)
|
|
1419
|
+
numbered_list_of_n_explanations: Optional[int] = kwargs.pop(
|
|
1420
|
+
"numbered_list_of_n_explanations"
|
|
1421
|
+
)
|
|
1422
|
+
if numbered_list_of_n_explanations is not None:
|
|
1423
|
+
assert numbered_list_of_n_explanations > 0, numbered_list_of_n_explanations
|
|
1424
|
+
# This parameter lets us dynamically shrink the prompt if our initial attempt to create it
|
|
1425
|
+
# results in something that's too long. It's only implemented for the 4k context size.
|
|
1426
|
+
kwargs.setdefault("omit_n_activation_records", 0)
|
|
1427
|
+
omit_n_activation_records: int = kwargs.pop("omit_n_activation_records")
|
|
1428
|
+
max_tokens_for_completion: int = kwargs.pop("max_tokens_for_completion")
|
|
1429
|
+
assert not kwargs, f"Unexpected kwargs: {kwargs}"
|
|
1430
|
+
|
|
1431
|
+
prompt_builder = PromptBuilder()
|
|
1432
|
+
# TODO: this is pretty verbose and can probably be shortened
|
|
1433
|
+
prompt_builder.add_message(
|
|
1434
|
+
Role.SYSTEM,
|
|
1435
|
+
"You are explaining the behavior of a neuron in a neural network. Your response should be a concise explanation (3 to 20 words) that captures what the neuron detects or predicts by finding patterns in lists.\n\n"
|
|
1436
|
+
"To determine the explanation, you are given four lists:\n\n"
|
|
1437
|
+
"- TOP_POSITIVE_LOGITS, which are the most likely words or tokens associated with this neuron.\n"
|
|
1438
|
+
"- TOP_ACTIVATING_TEXTS, which are top activating texts.\n\n"
|
|
1439
|
+
"- MAX_ACTIVATING_TOKENS, which are the top activating tokens in the top activating texts.\n"
|
|
1440
|
+
"- TOKENS_AFTER_MAX_ACTIVATING_TOKEN, which are the tokens immediately after the max activating token.\n"
|
|
1441
|
+
"Your job is to explain the behavior of the neuron in a single short phrase. You should look at the lists and find a pattern that helps you explain the behavior of the neuron.\n\n"
|
|
1442
|
+
"Rules:\n"
|
|
1443
|
+
"- Keep your explanation concise (3 to 20 words).\n"
|
|
1444
|
+
"- The explanation could be a single word, or phrase, or pattern.\n"
|
|
1445
|
+
"- The explanation could be about tokens following or preceding certain tokens.\n"
|
|
1446
|
+
"- The explanation could be about words starting with a sequence.\n"
|
|
1447
|
+
"- Avoid simply listing all the tokens. Instead, try to find patterns.\n"
|
|
1448
|
+
'- Just say the pattern itself, and do not start with phrases like "words related to", "concepts related to", or "variations of the word".\n'
|
|
1449
|
+
'- Do not start your explanation with "This neuron detects/predicts".\n'
|
|
1450
|
+
'- Do not mention "tokens" or "patterns" in your explanation.\n'
|
|
1451
|
+
"- Do not capitalize the first letter unless it is a proper noun.\n"
|
|
1452
|
+
'- The explanation should be specific. For example, "unique words" is not a specific enough pattern, nor is "foreign words".\n'
|
|
1453
|
+
"- Not ALL top activating texts/tokens have to match the exact same pattern, but a majority should.\n"
|
|
1454
|
+
"- If you absolutely cannot make any guesses, return the first token in MAX_ACTIVATING_TOKENS.\n\n"
|
|
1455
|
+
"Your response should be exactly a short phrase that explains the behavior of the neuron, not a full sentence.",
|
|
1456
|
+
)
|
|
1457
|
+
num_omitted_activation_records = 0
|
|
1458
|
+
self._add_per_neuron_explanation_prompt(
|
|
1459
|
+
prompt_builder,
|
|
1460
|
+
all_activation_records,
|
|
1461
|
+
0,
|
|
1462
|
+
max_activation,
|
|
1463
|
+
numbered_list_of_n_explanations=numbered_list_of_n_explanations,
|
|
1464
|
+
top_positive_logits=top_positive_logits,
|
|
1465
|
+
explanation=None,
|
|
1466
|
+
)
|
|
1467
|
+
# If the prompt is too long *and* we omitted the specified number of activation records, try
|
|
1468
|
+
# again, omitting one more. (If we didn't make the specified number of omissions, we're out
|
|
1469
|
+
# of opportunities to omit records, so we just return the prompt as-is.)
|
|
1470
|
+
if (
|
|
1471
|
+
self._prompt_is_too_long(prompt_builder, max_tokens_for_completion)
|
|
1472
|
+
and num_omitted_activation_records == omit_n_activation_records
|
|
1473
|
+
):
|
|
1474
|
+
original_kwargs["omit_n_activation_records"] = omit_n_activation_records + 1
|
|
1475
|
+
return self.make_explanation_prompt(**original_kwargs)
|
|
1476
|
+
built_prompt = prompt_builder.build(self.prompt_format)
|
|
1477
|
+
|
|
1478
|
+
## debug only
|
|
1479
|
+
# import json
|
|
1480
|
+
|
|
1481
|
+
# if isinstance(built_prompt, list):
|
|
1482
|
+
# logger.error(json.dumps({"built_prompt": built_prompt}))
|
|
1483
|
+
# else:
|
|
1484
|
+
# logger.error(json.dumps({"built_prompt": built_prompt}))
|
|
1485
|
+
|
|
1486
|
+
# import sys
|
|
1487
|
+
# sys.exit(1)
|
|
1488
|
+
|
|
1489
|
+
return built_prompt
|
|
1490
|
+
|
|
1491
|
+
def format_top_logits(self, top_positive_logits: List[str]) -> str:
|
|
1492
|
+
return "\n".join([f"{logit.strip()}" for logit in top_positive_logits])
|
|
1493
|
+
|
|
1494
|
+
def _add_per_neuron_explanation_prompt(
|
|
1495
|
+
self,
|
|
1496
|
+
prompt_builder: PromptBuilder,
|
|
1497
|
+
activation_records: Sequence[ActivationRecord],
|
|
1498
|
+
index: int,
|
|
1499
|
+
max_activation: float,
|
|
1500
|
+
top_positive_logits: Optional[List[str]],
|
|
1501
|
+
# When set, this indicates that the prompt should solicit a numbered list of the given
|
|
1502
|
+
# number of explanations, rather than a single explanation.
|
|
1503
|
+
numbered_list_of_n_explanations: Optional[int],
|
|
1504
|
+
explanation: Optional[str], # None means this is the end of the full prompt.
|
|
1505
|
+
) -> None:
|
|
1506
|
+
user_message = f"""
|
|
1507
|
+
|
|
1508
|
+
<MAX_ACTIVATING_TOKENS>
|
|
1509
|
+
|
|
1510
|
+
{self.format_max_activating_tokens(activation_records)}
|
|
1511
|
+
|
|
1512
|
+
</MAX_ACTIVATING_TOKENS>
|
|
1513
|
+
|
|
1514
|
+
|
|
1515
|
+
<TOKENS_AFTER_MAX_ACTIVATING_TOKEN>
|
|
1516
|
+
|
|
1517
|
+
{self.format_tokens_after_max_activating_token(activation_records)}
|
|
1518
|
+
|
|
1519
|
+
</TOKENS_AFTER_MAX_ACTIVATING_TOKEN>
|
|
1520
|
+
|
|
1521
|
+
|
|
1522
|
+
<TOP_POSITIVE_LOGITS>
|
|
1523
|
+
|
|
1524
|
+
{self.format_top_logits(top_positive_logits) if top_positive_logits else ""}
|
|
1525
|
+
|
|
1526
|
+
</TOP_POSITIVE_LOGITS>
|
|
1527
|
+
|
|
1528
|
+
|
|
1529
|
+
<TOP_ACTIVATING_TEXTS>
|
|
1530
|
+
|
|
1531
|
+
{self.format_top_activating_texts(activation_records)}
|
|
1532
|
+
|
|
1533
|
+
</TOP_ACTIVATING_TEXTS>
|
|
1534
|
+
|
|
1535
|
+
"""
|
|
1536
|
+
# logger.error(f"user_message: {user_message}")
|
|
1537
|
+
|
|
1538
|
+
message_requesting_explanation = (
|
|
1539
|
+
"\nExplain the neuron above with a word or phrase, not a complete sentence."
|
|
1540
|
+
)
|
|
1541
|
+
if numbered_list_of_n_explanations is None:
|
|
1542
|
+
user_message += message_requesting_explanation
|
|
1543
|
+
assistant_message = ""
|
|
1544
|
+
prompt_builder.add_message(Role.USER, user_message)
|
|
1545
|
+
|
|
1546
|
+
if explanation is not None:
|
|
1547
|
+
assistant_message += f"{explanation}"
|
|
1548
|
+
if assistant_message:
|
|
1549
|
+
prompt_builder.add_message(Role.ASSISTANT, assistant_message)
|
|
1550
|
+
else:
|
|
1551
|
+
prompt_builder.add_message(
|
|
1552
|
+
Role.USER,
|
|
1553
|
+
message_requesting_explanation,
|
|
1554
|
+
)
|
|
1555
|
+
if explanation is not None:
|
|
1556
|
+
prompt_builder.add_message(Role.ASSISTANT, f"{explanation}")
|
|
1557
|
+
|
|
1558
|
+
def postprocess_explanations(
|
|
1559
|
+
self, completions: list[str], prompt_kwargs: dict[str, Any]
|
|
1560
|
+
) -> list[Any]:
|
|
1561
|
+
"""Postprocess the explanations returned by the API"""
|
|
1562
|
+
numbered_list_of_n_explanations = prompt_kwargs.get(
|
|
1563
|
+
"numbered_list_of_n_explanations"
|
|
1564
|
+
)
|
|
1565
|
+
if numbered_list_of_n_explanations is None:
|
|
1566
|
+
all_explanations = []
|
|
1567
|
+
for explanation in completions:
|
|
1568
|
+
# print(f"explanation: {explanation}")
|
|
1569
|
+
explanation = self.strip_explanation(explanation)
|
|
1570
|
+
if explanation.endswith("."):
|
|
1571
|
+
explanation = explanation[:-1]
|
|
1572
|
+
|
|
1573
|
+
all_explanations.append(explanation)
|
|
1574
|
+
continue
|
|
1575
|
+
return all_explanations
|
|
1576
|
+
else:
|
|
1577
|
+
all_explanations = []
|
|
1578
|
+
for completion in completions:
|
|
1579
|
+
for explanation in _split_numbered_list(completion):
|
|
1580
|
+
# print(f"explanation: {explanation}")
|
|
1581
|
+
explanation = self.strip_explanation(explanation)
|
|
1582
|
+
if explanation.endswith("."):
|
|
1583
|
+
explanation = explanation[:-1]
|
|
1584
|
+
all_explanations.append(explanation)
|
|
1585
|
+
continue
|
|
1586
|
+
return all_explanations
|
|
1587
|
+
|
|
1588
|
+
|
|
1589
|
+
class PythonCodeExplainer(NeuronExplainer):
|
|
1590
|
+
"""
|
|
1591
|
+
This explainer is used to explain OpenAI's CircuitGPT-python neurons, a sparse, interpretable model that is trained on Python code. (https://huggingface.co/openai/circuit-sparsity/tree/main)
|
|
1592
|
+
We need this because other explainers just end up saying generic variations of "code" or "symbols", which is not helpful.
|
|
1593
|
+
|
|
1594
|
+
Method:
|
|
1595
|
+
- We tell the model the context that it's Python code and that it should try to explain what specifically this neuron does in the Python code.
|
|
1596
|
+
- We ask the model specifically to avoid generic "code" or "symbols" explanations.
|
|
1597
|
+
- We show the top activating token in the context of each snippet
|
|
1598
|
+
- We show the model the snippets of code, with the top activating token highlighted between the symbol ★.
|
|
1599
|
+
- By default we show more tokens around the max activating token because code is dense.
|
|
1600
|
+
|
|
1601
|
+
This explainer assumes you are using a more intelligent model (eg Gemini-2.0-Flash or above), since it gives more general instructions.
|
|
1602
|
+
|
|
1603
|
+
See make_explanation_prompt below for the full prompt.
|
|
1604
|
+
|
|
1605
|
+
We mostly tested using this explainer with Gemini-2.5-Flash-Lite on Low thinking.
|
|
1606
|
+
"""
|
|
1607
|
+
|
|
1608
|
+
def __init__(
|
|
1609
|
+
self,
|
|
1610
|
+
model_name: str,
|
|
1611
|
+
prompt_format: PromptFormat = PromptFormat.HARMONY_V4,
|
|
1612
|
+
context_size: ContextSize = ContextSize.ONETWENTYEIGHT_K,
|
|
1613
|
+
few_shot_example_set: FewShotExampleSet = FewShotExampleSet.LOGITS,
|
|
1614
|
+
tokens_around_max_activating_token: int = 96,
|
|
1615
|
+
repeat_non_zero_activations: bool = False,
|
|
1616
|
+
max_concurrent: Optional[int] = 10,
|
|
1617
|
+
cache: bool = False,
|
|
1618
|
+
base_api_url: str = ApiClient.BASE_API_URL,
|
|
1619
|
+
override_api_key: str | None = None,
|
|
1620
|
+
):
|
|
1621
|
+
super().__init__(
|
|
1622
|
+
model_name=model_name,
|
|
1623
|
+
prompt_format=prompt_format,
|
|
1624
|
+
max_concurrent=max_concurrent,
|
|
1625
|
+
cache=cache,
|
|
1626
|
+
base_api_url=base_api_url,
|
|
1627
|
+
override_api_key=override_api_key,
|
|
1628
|
+
)
|
|
1629
|
+
self.context_size = context_size
|
|
1630
|
+
self.few_shot_example_set = few_shot_example_set
|
|
1631
|
+
self.repeat_non_zero_activations = repeat_non_zero_activations
|
|
1632
|
+
self.tokens_around_max_activating_token = tokens_around_max_activating_token
|
|
1633
|
+
|
|
1634
|
+
def format_tokens_after_max_activating_token(
|
|
1635
|
+
self, activation_records: Sequence[ActivationRecord]
|
|
1636
|
+
) -> str:
|
|
1637
|
+
"""
|
|
1638
|
+
Format the tokens immediately after the max activating token.
|
|
1639
|
+
"""
|
|
1640
|
+
formatted_texts = []
|
|
1641
|
+
for record in activation_records:
|
|
1642
|
+
tokens = record.tokens
|
|
1643
|
+
activations = record.activations
|
|
1644
|
+
max_activation_index = activations.index(max(activations))
|
|
1645
|
+
# Only get the first token after the max activating token
|
|
1646
|
+
if max_activation_index + 1 < len(tokens):
|
|
1647
|
+
token_after_max_activating_token = (
|
|
1648
|
+
tokens[max_activation_index + 1].replace("\n", "").strip()
|
|
1649
|
+
)
|
|
1650
|
+
formatted_texts.append(f"{token_after_max_activating_token}")
|
|
1651
|
+
else:
|
|
1652
|
+
# Handle case where max activation is the last token
|
|
1653
|
+
formatted_texts.append("")
|
|
1654
|
+
return "\n".join(formatted_texts)
|
|
1655
|
+
|
|
1656
|
+
def format_max_activating_tokens(
|
|
1657
|
+
self, activation_records: Sequence[ActivationRecord]
|
|
1658
|
+
) -> str:
|
|
1659
|
+
"""
|
|
1660
|
+
Format the max activating tokens.
|
|
1661
|
+
"""
|
|
1662
|
+
formatted_tokens = []
|
|
1663
|
+
for record in activation_records:
|
|
1664
|
+
tokens = record.tokens
|
|
1665
|
+
activations = record.activations
|
|
1666
|
+
max_activation_index = activations.index(max(activations))
|
|
1667
|
+
max_activating_token = (
|
|
1668
|
+
tokens[max_activation_index].replace("\n", "").strip()
|
|
1669
|
+
)
|
|
1670
|
+
formatted_tokens.append(f"{max_activating_token}")
|
|
1671
|
+
return "\n".join(formatted_tokens)
|
|
1672
|
+
|
|
1673
|
+
def format_top_activating_texts(
|
|
1674
|
+
self, activation_records: Sequence[ActivationRecord]
|
|
1675
|
+
) -> str:
|
|
1676
|
+
"""
|
|
1677
|
+
Format activation records into a bullet point list of texts, with each text trimmed to
|
|
1678
|
+
8 tokens to the left and right of the maximum activating token. Replace line breaks with two spaces.
|
|
1679
|
+
"""
|
|
1680
|
+
formatted_texts = []
|
|
1681
|
+
|
|
1682
|
+
for record in activation_records:
|
|
1683
|
+
tokens = record.tokens
|
|
1684
|
+
activations = record.activations
|
|
1685
|
+
|
|
1686
|
+
# Find the index of the maximum activation
|
|
1687
|
+
max_activation_index = activations.index(max(activations))
|
|
1688
|
+
|
|
1689
|
+
# Calculate the start and end indices for the window
|
|
1690
|
+
start_index = max(
|
|
1691
|
+
0, max_activation_index - self.tokens_around_max_activating_token
|
|
1692
|
+
)
|
|
1693
|
+
end_index = min(
|
|
1694
|
+
len(tokens),
|
|
1695
|
+
max_activation_index + self.tokens_around_max_activating_token + 1,
|
|
1696
|
+
) # +1 to include the token at end_index-1
|
|
1697
|
+
|
|
1698
|
+
# Create the trimmed text with the max activating token surrounded by ★
|
|
1699
|
+
trimmed_tokens = (
|
|
1700
|
+
tokens[start_index:max_activation_index]
|
|
1701
|
+
+ [f"★{tokens[max_activation_index]}★"]
|
|
1702
|
+
+ tokens[max_activation_index + 1 : end_index]
|
|
1703
|
+
)
|
|
1704
|
+
|
|
1705
|
+
trimmed_text = "".join(trimmed_tokens).replace("\n", " ")
|
|
1706
|
+
formatted_texts.append(f"{trimmed_text}")
|
|
1707
|
+
|
|
1708
|
+
return "\n".join(formatted_texts)
|
|
1709
|
+
|
|
1710
|
+
def make_explanation_prompt(
|
|
1711
|
+
self, **kwargs: Any
|
|
1712
|
+
) -> Union[str, list[HarmonyMessage]]:
|
|
1713
|
+
original_kwargs = kwargs.copy()
|
|
1714
|
+
all_activation_records: Sequence[ActivationRecord] = kwargs.pop(
|
|
1715
|
+
"all_activation_records"
|
|
1716
|
+
)
|
|
1717
|
+
# Replace all ▁ characters in tokens with spaces
|
|
1718
|
+
processed_activation_records = []
|
|
1719
|
+
for record in all_activation_records:
|
|
1720
|
+
# Create a new ActivationRecord with processed tokens
|
|
1721
|
+
processed_tokens = [token.replace("▁", " ") for token in record.tokens]
|
|
1722
|
+
processed_activation_records.append(
|
|
1723
|
+
ActivationRecord(
|
|
1724
|
+
tokens=processed_tokens, activations=record.activations
|
|
1725
|
+
)
|
|
1726
|
+
)
|
|
1727
|
+
|
|
1728
|
+
# Use the processed records for the rest of the function
|
|
1729
|
+
all_activation_records = processed_activation_records
|
|
1730
|
+
max_activation: float = kwargs.pop("max_activation")
|
|
1731
|
+
top_positive_logits: List[str] = kwargs.pop("top_positive_logits")
|
|
1732
|
+
|
|
1733
|
+
# Replace all ▁ characters in top_positive_logits with spaces
|
|
1734
|
+
processed_top_positive_logits = [
|
|
1735
|
+
logit.replace("▁", " ").replace("\n", "").strip()
|
|
1736
|
+
for logit in top_positive_logits
|
|
1737
|
+
]
|
|
1738
|
+
top_positive_logits = processed_top_positive_logits
|
|
1739
|
+
kwargs.setdefault("numbered_list_of_n_explanations", None)
|
|
1740
|
+
numbered_list_of_n_explanations: Optional[int] = kwargs.pop(
|
|
1741
|
+
"numbered_list_of_n_explanations"
|
|
1742
|
+
)
|
|
1743
|
+
if numbered_list_of_n_explanations is not None:
|
|
1744
|
+
assert numbered_list_of_n_explanations > 0, numbered_list_of_n_explanations
|
|
1745
|
+
# This parameter lets us dynamically shrink the prompt if our initial attempt to create it
|
|
1746
|
+
# results in something that's too long. It's only implemented for the 4k context size.
|
|
1747
|
+
kwargs.setdefault("omit_n_activation_records", 0)
|
|
1748
|
+
omit_n_activation_records: int = kwargs.pop("omit_n_activation_records")
|
|
1749
|
+
max_tokens_for_completion: int = kwargs.pop("max_tokens_for_completion")
|
|
1750
|
+
assert not kwargs, f"Unexpected kwargs: {kwargs}"
|
|
1751
|
+
|
|
1752
|
+
prompt_builder = PromptBuilder()
|
|
1753
|
+
# TODO: this is pretty verbose and can probably be shortened
|
|
1754
|
+
prompt_builder.add_message(
|
|
1755
|
+
Role.SYSTEM,
|
|
1756
|
+
"You are explaining the behavior of a neuron in a neural network. This specific neuron is trained on Python code. Your response should be a concise explanation (3 to 20 words) that captures what the neuron detects or predicts by finding patterns in lists.\n\n"
|
|
1757
|
+
"To determine the explanation, you are given three lists:\n\n"
|
|
1758
|
+
"- MAX_ACTIVATING_TOKENS, which are the top activating tokens in the top activating texts.\n"
|
|
1759
|
+
"- TOKENS_AFTER_MAX_ACTIVATING_TOKEN, which are the tokens immediately after the max activating token.\n"
|
|
1760
|
+
"- TOP_ACTIVATING_TEXTS, which are top activating texts (snippets of Python code). The top activating token is highlighted between the symbol ★. For example, if the top activating token is the beginning parenthesis character '(', the snippet would be 'print★(★\"Hello world\")'. \n\n"
|
|
1761
|
+
"Your job is to explain the behavior of the neuron in a single short phrase. You should look at the lists and find a pattern that helps you explain the behavior of the neuron.\n\n"
|
|
1762
|
+
"Rules:\n"
|
|
1763
|
+
"- Keep your explanation concise (3 to 20 words).\n"
|
|
1764
|
+
'- The explanation should be related a specific Python concept, like a specific code-related pattern or symbol. For example, it could be a loop, or keyword, a symbol, or a part of code like "second argument of a method definition".\n'
|
|
1765
|
+
"- The explanation could be about tokens following or preceding certain tokens.\n"
|
|
1766
|
+
"- Avoid simply listing all the tokens. Instead, try to find patterns.\n"
|
|
1767
|
+
'- Just say the pattern itself, and do not start with phrases like "words related to", "concepts related to", or "variations of the word".\n'
|
|
1768
|
+
'- Do not start your explanation with "This neuron detects/predicts".\n'
|
|
1769
|
+
'- Do not mention "tokens" or "patterns" in your explanation.\n'
|
|
1770
|
+
"- Do not capitalize the first letter unless it is a proper noun.\n"
|
|
1771
|
+
'- The explanation should be specific. For example, "code snippets" is not a specific enough pattern, nor is "various symbols".\n'
|
|
1772
|
+
"- Not ALL top activating texts/tokens have to match the exact same pattern, but a majority should.\n"
|
|
1773
|
+
"- If you absolutely cannot make any guesses, return the first token in MAX_ACTIVATING_TOKENS.\n\n"
|
|
1774
|
+
"Your response should be exactly a short phrase that explains the behavior of the neuron, not a full sentence.",
|
|
1775
|
+
)
|
|
1776
|
+
num_omitted_activation_records = 0
|
|
1777
|
+
self._add_per_neuron_explanation_prompt(
|
|
1778
|
+
prompt_builder,
|
|
1779
|
+
all_activation_records,
|
|
1780
|
+
0,
|
|
1781
|
+
max_activation,
|
|
1782
|
+
numbered_list_of_n_explanations=numbered_list_of_n_explanations,
|
|
1783
|
+
top_positive_logits=top_positive_logits,
|
|
1784
|
+
explanation=None,
|
|
1785
|
+
)
|
|
1786
|
+
# If the prompt is too long *and* we omitted the specified number of activation records, try
|
|
1787
|
+
# again, omitting one more. (If we didn't make the specified number of omissions, we're out
|
|
1788
|
+
# of opportunities to omit records, so we just return the prompt as-is.)
|
|
1789
|
+
if (
|
|
1790
|
+
self._prompt_is_too_long(prompt_builder, max_tokens_for_completion)
|
|
1791
|
+
and num_omitted_activation_records == omit_n_activation_records
|
|
1792
|
+
):
|
|
1793
|
+
original_kwargs["omit_n_activation_records"] = omit_n_activation_records + 1
|
|
1794
|
+
return self.make_explanation_prompt(**original_kwargs)
|
|
1795
|
+
built_prompt = prompt_builder.build(self.prompt_format)
|
|
1796
|
+
|
|
1797
|
+
# # debug only
|
|
1798
|
+
# import json
|
|
1799
|
+
|
|
1800
|
+
# if isinstance(built_prompt, list):
|
|
1801
|
+
# print(json.dumps({"built_prompt": built_prompt}))
|
|
1802
|
+
# else:
|
|
1803
|
+
# print(json.dumps({"built_prompt": built_prompt}))
|
|
1804
|
+
|
|
1805
|
+
# import sys
|
|
1806
|
+
# import time
|
|
1807
|
+
|
|
1808
|
+
# time.sleep(5)
|
|
1809
|
+
|
|
1810
|
+
# sys.exit(1)
|
|
1811
|
+
|
|
1812
|
+
return built_prompt
|
|
1813
|
+
|
|
1814
|
+
def format_top_logits(self, top_positive_logits: List[str]) -> str:
|
|
1815
|
+
return "\n".join([f"{logit.strip()}" for logit in top_positive_logits])
|
|
1816
|
+
|
|
1817
|
+
def _add_per_neuron_explanation_prompt(
|
|
1818
|
+
self,
|
|
1819
|
+
prompt_builder: PromptBuilder,
|
|
1820
|
+
activation_records: Sequence[ActivationRecord],
|
|
1821
|
+
index: int,
|
|
1822
|
+
max_activation: float,
|
|
1823
|
+
top_positive_logits: Optional[List[str]],
|
|
1824
|
+
# When set, this indicates that the prompt should solicit a numbered list of the given
|
|
1825
|
+
# number of explanations, rather than a single explanation.
|
|
1826
|
+
numbered_list_of_n_explanations: Optional[int],
|
|
1827
|
+
explanation: Optional[str], # None means this is the end of the full prompt.
|
|
1828
|
+
) -> None:
|
|
1829
|
+
user_message = f"""
|
|
1830
|
+
|
|
1831
|
+
<MAX_ACTIVATING_TOKENS>
|
|
1832
|
+
|
|
1833
|
+
{self.format_max_activating_tokens(activation_records)}
|
|
1834
|
+
|
|
1835
|
+
</MAX_ACTIVATING_TOKENS>
|
|
1836
|
+
|
|
1837
|
+
<TOKENS_AFTER_MAX_ACTIVATING_TOKEN>
|
|
1838
|
+
|
|
1839
|
+
{self.format_tokens_after_max_activating_token(activation_records)}
|
|
1840
|
+
|
|
1841
|
+
</TOKENS_AFTER_MAX_ACTIVATING_TOKEN>
|
|
1842
|
+
|
|
1843
|
+
<TOP_ACTIVATING_TEXTS>
|
|
1844
|
+
|
|
1845
|
+
{self.format_top_activating_texts(activation_records)}
|
|
1846
|
+
|
|
1847
|
+
</TOP_ACTIVATING_TEXTS>
|
|
1848
|
+
|
|
1849
|
+
"""
|
|
1850
|
+
# logger.error(f"user_message: {user_message}")
|
|
1851
|
+
|
|
1852
|
+
message_requesting_explanation = (
|
|
1853
|
+
"\nExplain the neuron above with a word or phrase, not a complete sentence."
|
|
1854
|
+
)
|
|
1855
|
+
if numbered_list_of_n_explanations is None:
|
|
1856
|
+
user_message += message_requesting_explanation
|
|
1857
|
+
assistant_message = ""
|
|
1858
|
+
prompt_builder.add_message(Role.USER, user_message)
|
|
1859
|
+
|
|
1860
|
+
if explanation is not None:
|
|
1861
|
+
assistant_message += f"{explanation}"
|
|
1862
|
+
if assistant_message:
|
|
1863
|
+
prompt_builder.add_message(Role.ASSISTANT, assistant_message)
|
|
1864
|
+
else:
|
|
1865
|
+
prompt_builder.add_message(
|
|
1866
|
+
Role.USER,
|
|
1867
|
+
message_requesting_explanation,
|
|
1868
|
+
)
|
|
1869
|
+
if explanation is not None:
|
|
1870
|
+
prompt_builder.add_message(Role.ASSISTANT, f"{explanation}")
|
|
1871
|
+
|
|
1872
|
+
def postprocess_explanations(
|
|
1873
|
+
self, completions: list[str], prompt_kwargs: dict[str, Any]
|
|
1874
|
+
) -> list[Any]:
|
|
1875
|
+
"""Postprocess the explanations returned by the API"""
|
|
1876
|
+
numbered_list_of_n_explanations = prompt_kwargs.get(
|
|
1877
|
+
"numbered_list_of_n_explanations"
|
|
1878
|
+
)
|
|
1879
|
+
if numbered_list_of_n_explanations is None:
|
|
1880
|
+
all_explanations = []
|
|
1881
|
+
for explanation in completions:
|
|
1882
|
+
# print(f"explanation: {explanation}")
|
|
1883
|
+
explanation = self.strip_explanation(explanation)
|
|
1884
|
+
if explanation.endswith("."):
|
|
1885
|
+
explanation = explanation[:-1]
|
|
1886
|
+
|
|
1887
|
+
all_explanations.append(explanation)
|
|
1888
|
+
continue
|
|
1889
|
+
return all_explanations
|
|
1890
|
+
else:
|
|
1891
|
+
all_explanations = []
|
|
1892
|
+
for completion in completions:
|
|
1893
|
+
for explanation in _split_numbered_list(completion):
|
|
1894
|
+
# print(f"explanation: {explanation}")
|
|
1895
|
+
explanation = self.strip_explanation(explanation)
|
|
1896
|
+
if explanation.endswith("."):
|
|
1897
|
+
explanation = explanation[:-1]
|
|
1898
|
+
all_explanations.append(explanation)
|
|
1899
|
+
continue
|
|
1900
|
+
return all_explanations
|
|
1901
|
+
|
|
1902
|
+
|
|
1903
|
+
class MaxActivationExplainer(NeuronExplainer):
|
|
1904
|
+
"""
|
|
1905
|
+
This is a trimmed down version of the MaxActivationAndLogitsExplainer. It's the same except it doesn't show the model top positive logits or the immediate tokens after the max activating token.
|
|
1906
|
+
|
|
1907
|
+
This is a very concise explainer (1 to 6 words) that attempts to replicate Anthropic's attribution graphs explainer.
|
|
1908
|
+
It shows the model activations.
|
|
1909
|
+
This explainer is expected to be used for the first 2/3 of layers in a model.
|
|
1910
|
+
This explainer's tries to explain using one of these options:
|
|
1911
|
+
- a brief description of the max activating token (which can simply be the max activating token itself)
|
|
1912
|
+
- a brief description of the top activating texts
|
|
1913
|
+
|
|
1914
|
+
We force the explainer to try and explain using each method, then return when an explanation is found. Forcing it to do this made the explanations much more accurate.
|
|
1915
|
+
|
|
1916
|
+
See make_explanation_prompt below for the full prompt.
|
|
1917
|
+
|
|
1918
|
+
A weakness of this explainer is that it is less good at explaining the whole context - more for immediate words/characters on the top activating token.
|
|
1919
|
+
You can increase the "tokens_around_max_activating_token" to try to improve this behavior.
|
|
1920
|
+
|
|
1921
|
+
We mostly tested using this explainer with Gemini-2.0-Flash.
|
|
1922
|
+
"""
|
|
1923
|
+
|
|
1924
|
+
def __init__(
|
|
1925
|
+
self,
|
|
1926
|
+
model_name: str,
|
|
1927
|
+
prompt_format: PromptFormat = PromptFormat.HARMONY_V4,
|
|
1928
|
+
context_size: ContextSize = ContextSize.ONETWENTYEIGHT_K,
|
|
1929
|
+
few_shot_example_set: FewShotExampleSet = FewShotExampleSet.ACTIVATIONS,
|
|
1930
|
+
tokens_around_max_activating_token: int = 24,
|
|
1931
|
+
repeat_non_zero_activations: bool = False,
|
|
1932
|
+
max_concurrent: Optional[int] = 10,
|
|
1933
|
+
cache: bool = False,
|
|
1934
|
+
base_api_url: str = ApiClient.BASE_API_URL,
|
|
1935
|
+
override_api_key: str | None = None,
|
|
1936
|
+
):
|
|
1937
|
+
super().__init__(
|
|
1938
|
+
model_name=model_name,
|
|
1939
|
+
prompt_format=prompt_format,
|
|
1940
|
+
max_concurrent=max_concurrent,
|
|
1941
|
+
cache=cache,
|
|
1942
|
+
base_api_url=base_api_url,
|
|
1943
|
+
override_api_key=override_api_key,
|
|
1944
|
+
)
|
|
1945
|
+
self.context_size = context_size
|
|
1946
|
+
self.few_shot_example_set = few_shot_example_set
|
|
1947
|
+
self.repeat_non_zero_activations = repeat_non_zero_activations
|
|
1948
|
+
self.tokens_around_max_activating_token = tokens_around_max_activating_token
|
|
1949
|
+
|
|
1950
|
+
def format_max_activating_tokens(
|
|
1951
|
+
self, activation_records: Sequence[ActivationRecord]
|
|
1952
|
+
) -> str:
|
|
1953
|
+
"""
|
|
1954
|
+
Format the max activating tokens.
|
|
1955
|
+
"""
|
|
1956
|
+
formatted_tokens = []
|
|
1957
|
+
for record in activation_records:
|
|
1958
|
+
tokens = record.tokens
|
|
1959
|
+
activations = record.activations
|
|
1960
|
+
max_activation_index = activations.index(max(activations))
|
|
1961
|
+
max_activating_token = (
|
|
1962
|
+
tokens[max_activation_index].replace("\n", "").strip()
|
|
1963
|
+
)
|
|
1964
|
+
formatted_tokens.append(f"{max_activating_token}")
|
|
1965
|
+
return "\n".join(formatted_tokens)
|
|
1966
|
+
|
|
1967
|
+
def format_top_activating_texts(
|
|
1968
|
+
self, activation_records: Sequence[ActivationRecord]
|
|
1969
|
+
) -> str:
|
|
1970
|
+
"""
|
|
1971
|
+
Format activation records into a bullet point list of texts, with each text trimmed to
|
|
1972
|
+
8 tokens to the left and right of the maximum activating token. Replace line breaks with two spaces.
|
|
1973
|
+
"""
|
|
1974
|
+
formatted_texts = []
|
|
1975
|
+
|
|
1976
|
+
for record in activation_records:
|
|
1977
|
+
tokens = record.tokens
|
|
1978
|
+
activations = record.activations
|
|
1979
|
+
|
|
1980
|
+
# Find the index of the maximum activation
|
|
1981
|
+
max_activation_index = activations.index(max(activations))
|
|
1982
|
+
|
|
1983
|
+
# Calculate the start and end indices for the window
|
|
1984
|
+
start_index = max(
|
|
1985
|
+
0, max_activation_index - self.tokens_around_max_activating_token
|
|
1986
|
+
)
|
|
1987
|
+
end_index = min(
|
|
1988
|
+
len(tokens),
|
|
1989
|
+
max_activation_index + self.tokens_around_max_activating_token + 1,
|
|
1990
|
+
) # +1 to include the token at end_index-1
|
|
1991
|
+
|
|
1992
|
+
# Create the trimmed text with the max activating token surrounded by ^^
|
|
1993
|
+
trimmed_tokens = (
|
|
1994
|
+
tokens[start_index:max_activation_index]
|
|
1995
|
+
+ [f"{tokens[max_activation_index]}"]
|
|
1996
|
+
+ tokens[max_activation_index + 1 : end_index]
|
|
1997
|
+
)
|
|
1998
|
+
|
|
1999
|
+
trimmed_text = "".join(trimmed_tokens).replace("\n", " ")
|
|
2000
|
+
formatted_texts.append(f"{trimmed_text}")
|
|
2001
|
+
|
|
2002
|
+
return "\n".join(formatted_texts)
|
|
2003
|
+
|
|
2004
|
+
def make_explanation_prompt(
|
|
2005
|
+
self, **kwargs: Any
|
|
2006
|
+
) -> Union[str, list[HarmonyMessage]]:
|
|
2007
|
+
original_kwargs = kwargs.copy()
|
|
2008
|
+
all_activation_records: Sequence[ActivationRecord] = kwargs.pop(
|
|
2009
|
+
"all_activation_records"
|
|
2010
|
+
)
|
|
2011
|
+
# Replace all ▁ characters in tokens with spaces
|
|
2012
|
+
processed_activation_records = []
|
|
2013
|
+
for record in all_activation_records:
|
|
2014
|
+
# Create a new ActivationRecord with processed tokens
|
|
2015
|
+
processed_tokens = [token.replace("▁", " ") for token in record.tokens]
|
|
2016
|
+
processed_activation_records.append(
|
|
2017
|
+
ActivationRecord(
|
|
2018
|
+
tokens=processed_tokens, activations=record.activations
|
|
2019
|
+
)
|
|
2020
|
+
)
|
|
2021
|
+
|
|
2022
|
+
# Use the processed records for the rest of the function
|
|
2023
|
+
all_activation_records = processed_activation_records
|
|
2024
|
+
max_activation: float = kwargs.pop("max_activation")
|
|
2025
|
+
|
|
2026
|
+
kwargs.setdefault("numbered_list_of_n_explanations", None)
|
|
2027
|
+
numbered_list_of_n_explanations: Optional[int] = kwargs.pop(
|
|
2028
|
+
"numbered_list_of_n_explanations"
|
|
2029
|
+
)
|
|
2030
|
+
if numbered_list_of_n_explanations is not None:
|
|
2031
|
+
assert numbered_list_of_n_explanations > 0, numbered_list_of_n_explanations
|
|
2032
|
+
# This parameter lets us dynamically shrink the prompt if our initial attempt to create it
|
|
2033
|
+
# results in something that's too long. It's only implemented for the 4k context size.
|
|
2034
|
+
kwargs.setdefault("omit_n_activation_records", 0)
|
|
2035
|
+
omit_n_activation_records: int = kwargs.pop("omit_n_activation_records")
|
|
2036
|
+
max_tokens_for_completion: int = kwargs.pop("max_tokens_for_completion")
|
|
2037
|
+
assert not kwargs, f"Unexpected kwargs: {kwargs}"
|
|
2038
|
+
|
|
2039
|
+
prompt_builder = PromptBuilder()
|
|
2040
|
+
# TODO: this is pretty verbose and can probably be shortened
|
|
2041
|
+
prompt_builder.add_message(
|
|
2042
|
+
Role.SYSTEM,
|
|
2043
|
+
"You are explaining the behavior of a neuron in a neural network. Your response should be a very concise explanation (1-6 words) that captures what the neuron detects or predicts by finding patterns in lists.\n\n"
|
|
2044
|
+
"To determine the explanation, you are given two lists:\n\n"
|
|
2045
|
+
"- MAX_ACTIVATING_TOKENS, which are the top activating tokens in the top activating texts.\n"
|
|
2046
|
+
"- TOP_ACTIVATING_TEXTS, which are top activating texts.\n\n"
|
|
2047
|
+
"You should look for a pattern by trying the following methods in order. Once you find a pattern, stop and return that pattern. Do not proceed to the later methods.\n"
|
|
2048
|
+
"Method 1: Look at MAX_ACTIVATING_TOKENS. If they share something specific in common, or are all the same token or a variation of the same token (like different cases or conjugations), respond with that token.\n"
|
|
2049
|
+
"Method 2: Look at TOP_ACTIVATING_TEXTS and make a best guess by describing the broad theme or context, ignoring the max activating tokens.\n\n"
|
|
2050
|
+
"Rules:\n"
|
|
2051
|
+
"- Keep your explanation extremely concise (1-6 words, mostly 1-3 words).\n"
|
|
2052
|
+
'- Do not add unnecessary phrases like "words related to", "concepts related to", or "variations of the word".\n'
|
|
2053
|
+
'- Do not mention "tokens" or "patterns" in your explanation.\n'
|
|
2054
|
+
'- The explanation should be specific. For example, "unique words" is not a specific enough pattern, nor is "foreign words".\n'
|
|
2055
|
+
"- If you absolutely cannot make any guesses, return the first token in MAX_ACTIVATING_TOKENS.\n\n"
|
|
2056
|
+
"Respond by going through each method number until you find one that helps you find an explanation for what this neuron is detecting or predicting. If a method does not help you find an explanation, briefly explain why it does not, then go on to the next method. "
|
|
2057
|
+
"Finally, end your response with the method number you used, the reason for your explanation, and then the explanation.",
|
|
2058
|
+
)
|
|
2059
|
+
few_shot_examples = self.few_shot_example_set.get_examples()
|
|
2060
|
+
num_omitted_activation_records = 0
|
|
2061
|
+
for i, few_shot_example in enumerate(few_shot_examples):
|
|
2062
|
+
few_shot_activation_records = few_shot_example.activation_records
|
|
2063
|
+
self._add_per_neuron_explanation_prompt(
|
|
2064
|
+
prompt_builder,
|
|
2065
|
+
few_shot_activation_records,
|
|
2066
|
+
i,
|
|
2067
|
+
calculate_max_activation(few_shot_example.activation_records),
|
|
2068
|
+
numbered_list_of_n_explanations=numbered_list_of_n_explanations,
|
|
2069
|
+
explanation=few_shot_example.explanation,
|
|
2070
|
+
)
|
|
2071
|
+
self._add_per_neuron_explanation_prompt(
|
|
2072
|
+
prompt_builder,
|
|
2073
|
+
all_activation_records,
|
|
2074
|
+
len(few_shot_examples),
|
|
2075
|
+
max_activation,
|
|
2076
|
+
numbered_list_of_n_explanations=numbered_list_of_n_explanations,
|
|
2077
|
+
explanation=None,
|
|
2078
|
+
)
|
|
2079
|
+
# If the prompt is too long *and* we omitted the specified number of activation records, try
|
|
2080
|
+
# again, omitting one more. (If we didn't make the specified number of omissions, we're out
|
|
2081
|
+
# of opportunities to omit records, so we just return the prompt as-is.)
|
|
2082
|
+
if (
|
|
2083
|
+
self._prompt_is_too_long(prompt_builder, max_tokens_for_completion)
|
|
2084
|
+
and num_omitted_activation_records == omit_n_activation_records
|
|
2085
|
+
):
|
|
2086
|
+
original_kwargs["omit_n_activation_records"] = omit_n_activation_records + 1
|
|
2087
|
+
return self.make_explanation_prompt(**original_kwargs)
|
|
2088
|
+
built_prompt = prompt_builder.build(self.prompt_format)
|
|
2089
|
+
|
|
2090
|
+
# ## debug only
|
|
2091
|
+
# import json
|
|
2092
|
+
|
|
2093
|
+
# if isinstance(built_prompt, list):
|
|
2094
|
+
# logger.error(json.dumps({"built_prompt": built_prompt}))
|
|
2095
|
+
# else:
|
|
2096
|
+
# logger.error(json.dumps({"built_prompt": built_prompt}))
|
|
2097
|
+
# import sys
|
|
2098
|
+
|
|
2099
|
+
# sys.exit(1)
|
|
2100
|
+
|
|
2101
|
+
return built_prompt
|
|
2102
|
+
|
|
2103
|
+
def _add_per_neuron_explanation_prompt(
|
|
2104
|
+
self,
|
|
2105
|
+
prompt_builder: PromptBuilder,
|
|
2106
|
+
activation_records: Sequence[ActivationRecord],
|
|
2107
|
+
index: int,
|
|
2108
|
+
max_activation: float,
|
|
2109
|
+
# When set, this indicates that the prompt should solicit a numbered list of the given
|
|
2110
|
+
# number of explanations, rather than a single explanation.
|
|
2111
|
+
numbered_list_of_n_explanations: Optional[int],
|
|
2112
|
+
explanation: Optional[str], # None means this is the end of the full prompt.
|
|
2113
|
+
) -> None:
|
|
2114
|
+
user_message = f"""
|
|
2115
|
+
|
|
2116
|
+
Neuron {index + 1}
|
|
2117
|
+
|
|
2118
|
+
<MAX_ACTIVATING_TOKENS>
|
|
2119
|
+
|
|
2120
|
+
{self.format_max_activating_tokens(activation_records)}
|
|
2121
|
+
|
|
2122
|
+
</MAX_ACTIVATING_TOKENS>
|
|
2123
|
+
|
|
2124
|
+
|
|
2125
|
+
<TOP_ACTIVATING_TEXTS>
|
|
2126
|
+
|
|
2127
|
+
{self.format_top_activating_texts(activation_records)}
|
|
2128
|
+
|
|
2129
|
+
</TOP_ACTIVATING_TEXTS>
|
|
2130
|
+
|
|
2131
|
+
"""
|
|
2132
|
+
# logger.error(f"user_message: {user_message}")
|
|
2133
|
+
|
|
2134
|
+
if numbered_list_of_n_explanations is None:
|
|
2135
|
+
user_message += f"\nExplanation of neuron {index + 1} behavior: "
|
|
2136
|
+
assistant_message = ""
|
|
2137
|
+
# # For the IF format, we want <|endofprompt|> to come before the explanation prefix.
|
|
2138
|
+
# if self.prompt_format == PromptFormat.INSTRUCTION_FOLLOWING:
|
|
2139
|
+
# assistant_message += f"{EXPLANATION_PREFIX_LOGITS}"
|
|
2140
|
+
# else:
|
|
2141
|
+
# user_message += f"{EXPLANATION_PREFIX_LOGITS}"
|
|
2142
|
+
prompt_builder.add_message(Role.USER, user_message)
|
|
2143
|
+
|
|
2144
|
+
if explanation is not None:
|
|
2145
|
+
assistant_message += f"{explanation}"
|
|
2146
|
+
if assistant_message:
|
|
2147
|
+
prompt_builder.add_message(Role.ASSISTANT, assistant_message)
|
|
2148
|
+
else:
|
|
2149
|
+
prompt_builder.add_message(
|
|
2150
|
+
Role.USER,
|
|
2151
|
+
f"""\nExplanation for neuron {index + 1} behavior: """,
|
|
2152
|
+
)
|
|
2153
|
+
if explanation is not None:
|
|
2154
|
+
prompt_builder.add_message(Role.ASSISTANT, f"{explanation}")
|
|
2155
|
+
|
|
2156
|
+
def postprocess_explanations(
|
|
2157
|
+
self, completions: list[str], prompt_kwargs: dict[str, Any]
|
|
2158
|
+
) -> list[Any]:
|
|
2159
|
+
"""Postprocess the explanations returned by the API"""
|
|
2160
|
+
numbered_list_of_n_explanations = prompt_kwargs.get(
|
|
2161
|
+
"numbered_list_of_n_explanations"
|
|
2162
|
+
)
|
|
2163
|
+
if numbered_list_of_n_explanations is None:
|
|
2164
|
+
all_explanations = []
|
|
2165
|
+
for explanation in completions:
|
|
2166
|
+
explanation = self.strip_explanation(explanation)
|
|
2167
|
+
# logger.error(f"explanation: {explanation}")
|
|
2168
|
+
# Split by "Explanation: " and take the last segment if it exists
|
|
2169
|
+
if "Explanation: " in explanation:
|
|
2170
|
+
explanation = explanation.split("Explanation: ")[-1]
|
|
2171
|
+
elif "explanation: " in explanation:
|
|
2172
|
+
explanation = explanation.split("explanation: ")[-1]
|
|
2173
|
+
else:
|
|
2174
|
+
logger.error(
|
|
2175
|
+
f"Error parsing response explanation, no explanation string found: {explanation}"
|
|
2176
|
+
)
|
|
2177
|
+
all_explanations.append("")
|
|
2178
|
+
continue
|
|
2179
|
+
|
|
2180
|
+
all_explanations.append(_remove_method_from_explanation(explanation))
|
|
2181
|
+
return all_explanations
|
|
2182
|
+
else:
|
|
2183
|
+
all_explanations = []
|
|
2184
|
+
for completion in completions:
|
|
2185
|
+
for explanation in _split_numbered_list(completion):
|
|
2186
|
+
explanation = self.strip_explanation(explanation)
|
|
2187
|
+
# Split by "Explanation: " and take the last segment if it exists
|
|
2188
|
+
if "Explanation: " in explanation:
|
|
2189
|
+
explanation = explanation.split("Explanation: ")[-1]
|
|
2190
|
+
elif "explanation: " in explanation:
|
|
2191
|
+
explanation = explanation.split("explanation: ")[-1]
|
|
2192
|
+
else:
|
|
2193
|
+
logger.error(
|
|
2194
|
+
f"Error parsing response explanation, no explanation string found: {explanation}"
|
|
2195
|
+
)
|
|
2196
|
+
all_explanations.append("")
|
|
2197
|
+
continue
|
|
2198
|
+
|
|
2199
|
+
all_explanations.append(
|
|
2200
|
+
_remove_method_from_explanation(explanation)
|
|
2201
|
+
)
|
|
2202
|
+
return all_explanations
|
|
2203
|
+
|
|
2204
|
+
|
|
2205
|
+
class TokenSpaceRepresentationExplainer(NeuronExplainer):
|
|
2206
|
+
"""
|
|
2207
|
+
Generate explanations of arbitrary lists of tokens which disproportionately activate a
|
|
2208
|
+
particular neuron. These lists of tokens can be generated in various ways. As an example, in one
|
|
2209
|
+
set of experiments, we compute the average activation for each neuron conditional on each token
|
|
2210
|
+
that appears in an internet text corpus. We then sort the tokens by their average activation,
|
|
2211
|
+
and show 50 of the top 100 tokens. Other techniques that could be used include taking the top
|
|
2212
|
+
tokens in the logit lens or tuned lens representations of a neuron.
|
|
2213
|
+
"""
|
|
2214
|
+
|
|
2215
|
+
def __init__(
|
|
2216
|
+
self,
|
|
2217
|
+
model_name: str,
|
|
2218
|
+
prompt_format: PromptFormat = PromptFormat.HARMONY_V4,
|
|
2219
|
+
context_size: ContextSize = ContextSize.FOUR_K,
|
|
2220
|
+
few_shot_example_set: TokenSpaceFewShotExampleSet = TokenSpaceFewShotExampleSet.ORIGINAL,
|
|
2221
|
+
use_few_shot: bool = False,
|
|
2222
|
+
output_numbered_list: bool = False,
|
|
2223
|
+
max_concurrent: Optional[int] = 10,
|
|
2224
|
+
cache: bool = False,
|
|
2225
|
+
):
|
|
2226
|
+
super().__init__(
|
|
2227
|
+
model_name=model_name,
|
|
2228
|
+
prompt_format=prompt_format,
|
|
2229
|
+
context_size=context_size,
|
|
2230
|
+
max_concurrent=max_concurrent,
|
|
2231
|
+
cache=cache,
|
|
2232
|
+
)
|
|
2233
|
+
self.use_few_shot = use_few_shot
|
|
2234
|
+
self.output_numbered_list = output_numbered_list
|
|
2235
|
+
if self.use_few_shot:
|
|
2236
|
+
assert few_shot_example_set is not None
|
|
2237
|
+
self.few_shot_examples: Optional[TokenSpaceFewShotExampleSet] = (
|
|
2238
|
+
few_shot_example_set
|
|
2239
|
+
)
|
|
2240
|
+
else:
|
|
2241
|
+
self.few_shot_examples = None
|
|
2242
|
+
self.prompt_prefix = (
|
|
2243
|
+
"We're studying neurons in a neural network. Each neuron looks for some particular "
|
|
2244
|
+
"kind of token (which can be a word, or part of a word). Look at the tokens the neuron "
|
|
2245
|
+
"activates for (listed below) and summarize in a single sentence what the neuron is "
|
|
2246
|
+
"looking for. Don't list examples of words."
|
|
2247
|
+
)
|
|
2248
|
+
|
|
2249
|
+
def make_explanation_prompt(
|
|
2250
|
+
self, **kwargs: Any
|
|
2251
|
+
) -> Union[str, list[HarmonyMessage]]:
|
|
2252
|
+
tokens: list[str] = kwargs.pop("tokens")
|
|
2253
|
+
max_tokens_for_completion = kwargs.pop("max_tokens_for_completion")
|
|
2254
|
+
assert not kwargs, f"Unexpected kwargs: {kwargs}"
|
|
2255
|
+
# Note that this does not preserve the precise tokens, as e.g.
|
|
2256
|
+
# f" {token_with_no_leading_space}" may be tokenized as "f{token_with_leading_space}".
|
|
2257
|
+
# TODO(dan): Try out other variants, including "\n".join(...) and ",".join(...)
|
|
2258
|
+
stringified_tokens = ", ".join([f"'{t}'" for t in tokens])
|
|
2259
|
+
|
|
2260
|
+
prompt_builder = PromptBuilder()
|
|
2261
|
+
prompt_builder.add_message(Role.SYSTEM, self.prompt_prefix)
|
|
2262
|
+
if self.use_few_shot:
|
|
2263
|
+
self._add_few_shot_examples(prompt_builder)
|
|
2264
|
+
self._add_neuron_specific_prompt(
|
|
2265
|
+
prompt_builder, stringified_tokens, explanation=None
|
|
2266
|
+
)
|
|
2267
|
+
|
|
2268
|
+
if self._prompt_is_too_long(prompt_builder, max_tokens_for_completion):
|
|
2269
|
+
raise ValueError(
|
|
2270
|
+
f"Prompt too long: {prompt_builder.build(self.prompt_format)}"
|
|
2271
|
+
)
|
|
2272
|
+
else:
|
|
2273
|
+
return prompt_builder.build(self.prompt_format)
|
|
2274
|
+
|
|
2275
|
+
def _add_few_shot_examples(self, prompt_builder: PromptBuilder) -> None:
|
|
2276
|
+
"""
|
|
2277
|
+
Append few-shot examples to the prompt. Each one consists of a comma-delimited list of
|
|
2278
|
+
tokens and corresponding explanations, as saved in
|
|
2279
|
+
alignment/neuron_explainer/weight_explainer/token_space_few_shot_examples.py.
|
|
2280
|
+
"""
|
|
2281
|
+
assert self.few_shot_examples is not None
|
|
2282
|
+
few_shot_example_list = self.few_shot_examples.get_examples()
|
|
2283
|
+
if self.output_numbered_list:
|
|
2284
|
+
raise NotImplementedError(
|
|
2285
|
+
"Numbered list output not supported for few-shot examples"
|
|
2286
|
+
)
|
|
2287
|
+
else:
|
|
2288
|
+
for few_shot_example in few_shot_example_list:
|
|
2289
|
+
self._add_neuron_specific_prompt(
|
|
2290
|
+
prompt_builder,
|
|
2291
|
+
", ".join([f"'{t}'" for t in few_shot_example.tokens]),
|
|
2292
|
+
explanation=few_shot_example.explanation,
|
|
2293
|
+
)
|
|
2294
|
+
|
|
2295
|
+
def _add_neuron_specific_prompt(
|
|
2296
|
+
self,
|
|
2297
|
+
prompt_builder: PromptBuilder,
|
|
2298
|
+
stringified_tokens: str,
|
|
2299
|
+
explanation: Optional[str],
|
|
2300
|
+
) -> None:
|
|
2301
|
+
"""
|
|
2302
|
+
Append a neuron-specific prompt to the prompt builder. The prompt consists of a list of
|
|
2303
|
+
tokens followed by either an explanation (if one is passed, for few shot examples) or by
|
|
2304
|
+
the beginning of a completion, to be completed by the model with an explanation.
|
|
2305
|
+
"""
|
|
2306
|
+
user_message = f"\n\n\n\nTokens:\n{stringified_tokens}\n\nExplanation:\n"
|
|
2307
|
+
assistant_message = ""
|
|
2308
|
+
looking_for = "This neuron is looking for"
|
|
2309
|
+
if self.prompt_format == PromptFormat.INSTRUCTION_FOLLOWING:
|
|
2310
|
+
# We want <|endofprompt|> to come before "This neuron is looking for" in the IF format.
|
|
2311
|
+
assistant_message += looking_for
|
|
2312
|
+
else:
|
|
2313
|
+
user_message += looking_for
|
|
2314
|
+
if self.output_numbered_list:
|
|
2315
|
+
start_of_list = "\n1."
|
|
2316
|
+
if self.prompt_format == PromptFormat.INSTRUCTION_FOLLOWING:
|
|
2317
|
+
assistant_message += start_of_list
|
|
2318
|
+
else:
|
|
2319
|
+
user_message += start_of_list
|
|
2320
|
+
if explanation is not None:
|
|
2321
|
+
assistant_message += f"{explanation}."
|
|
2322
|
+
prompt_builder.add_message(Role.USER, user_message)
|
|
2323
|
+
if assistant_message:
|
|
2324
|
+
prompt_builder.add_message(Role.ASSISTANT, assistant_message)
|
|
2325
|
+
|
|
2326
|
+
def postprocess_explanations(
|
|
2327
|
+
self, completions: list[str], prompt_kwargs: dict[str, Any]
|
|
2328
|
+
) -> list[str]:
|
|
2329
|
+
if self.output_numbered_list:
|
|
2330
|
+
# Each list in the top-level list will have multiple explanations (multiple strings).
|
|
2331
|
+
all_explanations = []
|
|
2332
|
+
for completion in completions:
|
|
2333
|
+
for explanation in _split_numbered_list(completion):
|
|
2334
|
+
explanation = self.strip_explanation(explanation)
|
|
2335
|
+
if explanation.startswith(EXPLANATION_PREFIX):
|
|
2336
|
+
explanation = explanation[len(EXPLANATION_PREFIX) :]
|
|
2337
|
+
all_explanations.append(explanation.strip())
|
|
2338
|
+
return all_explanations
|
|
2339
|
+
else:
|
|
2340
|
+
# Each element in the top-level list will be an explanation as a string.
|
|
2341
|
+
return [_remove_final_period(explanation) for explanation in completions]
|
|
2342
|
+
|
|
2343
|
+
|
|
2344
|
+
def format_attention_head_token_pairs(
|
|
2345
|
+
token_pair_examples: list[AttentionTokenPairExample], omit_zeros: bool = False
|
|
2346
|
+
) -> str:
|
|
2347
|
+
if omit_zeros:
|
|
2348
|
+
return ", ".join(
|
|
2349
|
+
[
|
|
2350
|
+
", ".join(
|
|
2351
|
+
[
|
|
2352
|
+
f"({example.tokens[coords[1]]}, {example.tokens[coords[0]]})"
|
|
2353
|
+
for coords in example.token_pair_coordinates
|
|
2354
|
+
]
|
|
2355
|
+
)
|
|
2356
|
+
for example in token_pair_examples
|
|
2357
|
+
]
|
|
2358
|
+
)
|
|
2359
|
+
else:
|
|
2360
|
+
return f"\n{ATTENTION_SEQUENCE_SEPARATOR}\n".join(
|
|
2361
|
+
[
|
|
2362
|
+
f"\n{ATTENTION_SEQUENCE_SEPARATOR}\n".join(
|
|
2363
|
+
[
|
|
2364
|
+
f"{format_attention_head_token_pair_string(example.tokens, coords)}"
|
|
2365
|
+
for coords in example.token_pair_coordinates
|
|
2366
|
+
]
|
|
2367
|
+
)
|
|
2368
|
+
for example in token_pair_examples
|
|
2369
|
+
]
|
|
2370
|
+
)
|
|
2371
|
+
|
|
2372
|
+
|
|
2373
|
+
def format_attention_head_token_pair_string(
|
|
2374
|
+
token_list: list[str], pair_coordinates: tuple[int, int]
|
|
2375
|
+
) -> str:
|
|
2376
|
+
def format_activated_token(i: int, token: str) -> str:
|
|
2377
|
+
if i == pair_coordinates[0] and i == pair_coordinates[1]:
|
|
2378
|
+
return f"[[**{token}**]]" # from and to
|
|
2379
|
+
if i == pair_coordinates[0]:
|
|
2380
|
+
return f"[[{token}]]" # from
|
|
2381
|
+
if i == pair_coordinates[1]:
|
|
2382
|
+
return f"**{token}**" # to
|
|
2383
|
+
return token
|
|
2384
|
+
|
|
2385
|
+
return "".join(
|
|
2386
|
+
[format_activated_token(i, token) for i, token in enumerate(token_list)]
|
|
2387
|
+
)
|
|
2388
|
+
|
|
2389
|
+
|
|
2390
|
+
def get_top_attention_coordinates(
|
|
2391
|
+
activation_records: list[ActivationRecord], top_k: int = 5
|
|
2392
|
+
) -> list[tuple[int, float, tuple[int, int]]]:
|
|
2393
|
+
candidates = []
|
|
2394
|
+
for i, record in enumerate(activation_records):
|
|
2395
|
+
top_activation_flat_indices = np.argsort(record.activations)[::-1][:top_k]
|
|
2396
|
+
top_vals: list[float] = [
|
|
2397
|
+
record.activations[idx] for idx in top_activation_flat_indices
|
|
2398
|
+
]
|
|
2399
|
+
top_coordinates = [
|
|
2400
|
+
convert_flattened_index_to_unflattened_index(flat_index)
|
|
2401
|
+
for flat_index in top_activation_flat_indices
|
|
2402
|
+
]
|
|
2403
|
+
candidates.extend(
|
|
2404
|
+
[(i, top_val, coords) for top_val, coords in zip(top_vals, top_coordinates)]
|
|
2405
|
+
)
|
|
2406
|
+
return sorted(candidates, key=lambda x: x[1], reverse=True)[:top_k]
|
|
2407
|
+
|
|
2408
|
+
|
|
2409
|
+
class AttentionHeadExplainer(NeuronExplainer):
|
|
2410
|
+
"""
|
|
2411
|
+
Generate explanations of attention head behavior using a prompt with lists of
|
|
2412
|
+
strongly attending to/from token pairs.
|
|
2413
|
+
Takes in NeuronRecord's corresponding to a single attention head. Extracts strongly
|
|
2414
|
+
activating to/from token pairs.
|
|
2415
|
+
"""
|
|
2416
|
+
|
|
2417
|
+
def __init__(
|
|
2418
|
+
self,
|
|
2419
|
+
model_name: str,
|
|
2420
|
+
prompt_format: PromptFormat = PromptFormat.HARMONY_V4,
|
|
2421
|
+
# This parameter lets us adjust the length of the prompt when we're generating explanations
|
|
2422
|
+
# using older models with shorter context windows. In the future we can use it to experiment
|
|
2423
|
+
# with 8k+ context windows.
|
|
2424
|
+
context_size: ContextSize = ContextSize.ONETWENTYEIGHT_K,
|
|
2425
|
+
repeat_strongly_attending_pairs: bool = False,
|
|
2426
|
+
max_concurrent: int | None = 10,
|
|
2427
|
+
cache: bool = False,
|
|
2428
|
+
base_api_url: str = ApiClient.BASE_API_URL,
|
|
2429
|
+
override_api_key: str | None = None,
|
|
2430
|
+
):
|
|
2431
|
+
super().__init__(
|
|
2432
|
+
model_name=model_name,
|
|
2433
|
+
prompt_format=prompt_format,
|
|
2434
|
+
max_concurrent=max_concurrent,
|
|
2435
|
+
cache=cache,
|
|
2436
|
+
base_api_url=base_api_url,
|
|
2437
|
+
override_api_key=override_api_key,
|
|
2438
|
+
)
|
|
2439
|
+
assert (
|
|
2440
|
+
context_size != ContextSize.TWO_K
|
|
2441
|
+
), "2k context size not supported for attention explanation"
|
|
2442
|
+
self.context_size = context_size
|
|
2443
|
+
self.repeat_strongly_attending_pairs = repeat_strongly_attending_pairs
|
|
2444
|
+
|
|
2445
|
+
def make_explanation_prompt(self, **kwargs: Any) -> str | list[HarmonyMessage]:
|
|
2446
|
+
original_kwargs = kwargs.copy()
|
|
2447
|
+
all_activation_records: list[ActivationRecord] = kwargs.pop(
|
|
2448
|
+
"all_activation_records"
|
|
2449
|
+
)
|
|
2450
|
+
# This parameter lets us dynamically shrink the prompt if our initial attempt to create it
|
|
2451
|
+
# results in something that's too long.
|
|
2452
|
+
kwargs.setdefault("omit_n_token_pair_examples", 0)
|
|
2453
|
+
omit_n_token_pair_examples: int = kwargs.pop("omit_n_token_pair_examples")
|
|
2454
|
+
|
|
2455
|
+
max_tokens_for_completion: int = kwargs.pop("max_tokens_for_completion")
|
|
2456
|
+
|
|
2457
|
+
kwargs.setdefault("num_top_pairs_to_display", 0)
|
|
2458
|
+
num_top_pairs_to_display: int = kwargs.pop("num_top_pairs_to_display")
|
|
2459
|
+
|
|
2460
|
+
assert not kwargs, f"Unexpected kwargs: {kwargs}"
|
|
2461
|
+
|
|
2462
|
+
prompt_builder = PromptBuilder()
|
|
2463
|
+
prompt_builder.add_message(
|
|
2464
|
+
Role.SYSTEM,
|
|
2465
|
+
"We're studying attention heads in a neural network. Each head looks at every pair of tokens "
|
|
2466
|
+
"in a short token sequence and activates for pairs of tokens that fit what it is looking for. "
|
|
2467
|
+
"Attention heads always attend from a token to a token earlier in the sequence (or from a "
|
|
2468
|
+
'token to itself). We will display multiple instances of sequences with the "to" token '
|
|
2469
|
+
'surrounded by double asterisks (e.g., **token**) and the "from" token surrounded by double '
|
|
2470
|
+
"square brackets (e.g., [[token]]). If a token attends from itself to itself, it will be "
|
|
2471
|
+
"surrounded by both (e.g., [[**token**]]). Look at the pairs of tokens the head activates for "
|
|
2472
|
+
"and summarize in a single sentence what pattern the head is looking for. We do not display "
|
|
2473
|
+
"every activating pair of tokens in a sequence; you must generalize from limited examples. "
|
|
2474
|
+
"Remember, the head always attends to tokens earlier in the sentence (marked with ** **) from "
|
|
2475
|
+
"tokens later in the sentence (marked with [[ ]]), except when the head attends from a token to "
|
|
2476
|
+
'itself (marked with [[** **]]). The explanation takes the form: "This attention head attends '
|
|
2477
|
+
"to {pattern of tokens marked with ** **, which appear earlier} from {pattern of tokens marked with "
|
|
2478
|
+
'[[ ]], which appear later}." The explanation does not include any of the markers (** **, [[ ]]), '
|
|
2479
|
+
f"as these are just for your reference. Sequences are separated by `{ATTENTION_SEQUENCE_SEPARATOR}`.",
|
|
2480
|
+
)
|
|
2481
|
+
num_omitted_token_pair_examples = 0
|
|
2482
|
+
for i, few_shot_example in enumerate(ATTENTION_HEAD_FEW_SHOT_EXAMPLES):
|
|
2483
|
+
few_shot_token_pair_examples = few_shot_example.token_pair_examples
|
|
2484
|
+
if num_omitted_token_pair_examples < omit_n_token_pair_examples:
|
|
2485
|
+
# Drop the last activation record for this few-shot example to save tokens, assuming
|
|
2486
|
+
# there are at least two activation records.
|
|
2487
|
+
if len(few_shot_token_pair_examples) > 1:
|
|
2488
|
+
print(
|
|
2489
|
+
f"Warning: omitting activation record from few-shot example {i}"
|
|
2490
|
+
)
|
|
2491
|
+
few_shot_token_pair_examples = few_shot_token_pair_examples[:-1]
|
|
2492
|
+
num_omitted_token_pair_examples += 1
|
|
2493
|
+
few_shot_explanation: str = few_shot_example.explanation
|
|
2494
|
+
self._add_per_head_explanation_prompt(
|
|
2495
|
+
prompt_builder,
|
|
2496
|
+
few_shot_token_pair_examples,
|
|
2497
|
+
i,
|
|
2498
|
+
explanation=few_shot_explanation,
|
|
2499
|
+
)
|
|
2500
|
+
|
|
2501
|
+
# ================================
|
|
2502
|
+
# Comment (Johnny): the original code does not seem to work (or I am not using it correctly?). Re-written below
|
|
2503
|
+
# ================================
|
|
2504
|
+
# # each element is (record_index, attention value, (from_token_index, to_token_index))
|
|
2505
|
+
# coords = get_top_attention_coordinates(
|
|
2506
|
+
# all_activation_records, top_k=num_top_pairs_to_display
|
|
2507
|
+
# )
|
|
2508
|
+
# prompt_examples = {}
|
|
2509
|
+
# for record_index, _, (from_token_index, to_token_index) in coords:
|
|
2510
|
+
# if record_index not in prompt_examples:
|
|
2511
|
+
# prompt_examples[record_index] = AttentionTokenPairExample(
|
|
2512
|
+
# tokens=all_activation_records[record_index].tokens,
|
|
2513
|
+
# token_pair_coordinates=[(from_token_index, to_token_index)],
|
|
2514
|
+
# )
|
|
2515
|
+
# else:
|
|
2516
|
+
# prompt_examples[record_index].token_pair_coordinates.append(
|
|
2517
|
+
# (from_token_index, to_token_index)
|
|
2518
|
+
# )
|
|
2519
|
+
# current_head_token_pair_examples = list(prompt_examples.values())
|
|
2520
|
+
|
|
2521
|
+
# make list of attention token pair examples
|
|
2522
|
+
attention_token_pair_examples = []
|
|
2523
|
+
for i, activation_record in enumerate(all_activation_records):
|
|
2524
|
+
# from (first value) is the dfaTargetIndex
|
|
2525
|
+
from_index = activation_record.dfa_target_index
|
|
2526
|
+
# to (second value) is the index of the max dfa
|
|
2527
|
+
to_index = np.argmax(activation_record.dfa_values)
|
|
2528
|
+
attention_token_pair_examples.append(
|
|
2529
|
+
AttentionTokenPairExample(
|
|
2530
|
+
tokens=activation_record.tokens,
|
|
2531
|
+
token_pair_coordinates=[(from_index, to_index)],
|
|
2532
|
+
)
|
|
2533
|
+
)
|
|
2534
|
+
|
|
2535
|
+
self._add_per_head_explanation_prompt(
|
|
2536
|
+
prompt_builder,
|
|
2537
|
+
attention_token_pair_examples,
|
|
2538
|
+
len(ATTENTION_HEAD_FEW_SHOT_EXAMPLES),
|
|
2539
|
+
explanation=None,
|
|
2540
|
+
)
|
|
2541
|
+
# If the prompt is too long *and* we omitted the specified number of activation records, try
|
|
2542
|
+
# again, omitting one more. (If we didn't make the specified number of omissions, we're out
|
|
2543
|
+
# of opportunities to omit records, so we just return the prompt as-is.)
|
|
2544
|
+
# if (
|
|
2545
|
+
# self._prompt_is_too_long(prompt_builder, max_tokens_for_completion)
|
|
2546
|
+
# and num_omitted_token_pair_examples == omit_n_token_pair_examples
|
|
2547
|
+
# ):
|
|
2548
|
+
# original_kwargs["omit_n_token_pair_examples"] = (
|
|
2549
|
+
# omit_n_token_pair_examples + 1
|
|
2550
|
+
# )
|
|
2551
|
+
# return self.make_explanation_prompt(**original_kwargs)
|
|
2552
|
+
return prompt_builder.build(self.prompt_format)
|
|
2553
|
+
|
|
2554
|
+
def _add_per_head_explanation_prompt(
|
|
2555
|
+
self,
|
|
2556
|
+
prompt_builder: PromptBuilder,
|
|
2557
|
+
token_pair_examples: list[
|
|
2558
|
+
AttentionTokenPairExample
|
|
2559
|
+
], # each dict has keys "tokens" and "token_pair_coordinates"
|
|
2560
|
+
index: int,
|
|
2561
|
+
explanation: str | None, # None means this is the end of the full prompt.
|
|
2562
|
+
) -> None:
|
|
2563
|
+
user_message = f"""
|
|
2564
|
+
|
|
2565
|
+
Attention head {index + 1}
|
|
2566
|
+
Activations:\n{format_attention_head_token_pairs(token_pair_examples, omit_zeros=False)}"""
|
|
2567
|
+
if self.repeat_strongly_attending_pairs:
|
|
2568
|
+
user_message += (
|
|
2569
|
+
f"\nThe same list of strongly activating token pairs, presented as (to_token, from_token):"
|
|
2570
|
+
f"{format_attention_head_token_pairs(token_pair_examples, omit_zeros=True)}"
|
|
2571
|
+
)
|
|
2572
|
+
|
|
2573
|
+
user_message += f"\nExplanation of attention head {index + 1} behavior:"
|
|
2574
|
+
assistant_message = ""
|
|
2575
|
+
# For the IF format, we want <|endofprompt|> to come before the explanation prefix.
|
|
2576
|
+
if self.prompt_format == PromptFormat.INSTRUCTION_FOLLOWING:
|
|
2577
|
+
assistant_message += f" {ATTENTION_EXPLANATION_PREFIX}"
|
|
2578
|
+
else:
|
|
2579
|
+
user_message += f" {ATTENTION_EXPLANATION_PREFIX}"
|
|
2580
|
+
prompt_builder.add_message(Role.USER, user_message)
|
|
2581
|
+
|
|
2582
|
+
if explanation is not None:
|
|
2583
|
+
assistant_message += f" {explanation}."
|
|
2584
|
+
if assistant_message:
|
|
2585
|
+
prompt_builder.add_message(Role.ASSISTANT, assistant_message)
|