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.
@@ -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)