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,1121 @@
1
+ """Uses API calls to simulate neuron activations based on an explanation."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import asyncio
6
+ import logging
7
+ import json
8
+ from abc import ABC, abstractmethod
9
+ from collections import OrderedDict
10
+ from enum import Enum
11
+ from typing import Any, Optional, Sequence, Union
12
+
13
+ import numpy as np
14
+ from neuron_explainer.activations.activation_records import (
15
+ calculate_max_activation,
16
+ format_activation_records,
17
+ format_sequences_for_simulation,
18
+ normalize_activations,
19
+ )
20
+ from neuron_explainer.activations.activations import ActivationRecord
21
+ from neuron_explainer.api_client import ApiClient
22
+ from neuron_explainer.explanations.explainer import EXPLANATION_PREFIX
23
+ from neuron_explainer.explanations.explanations import (
24
+ ActivationScale,
25
+ SequenceSimulation,
26
+ )
27
+ from neuron_explainer.explanations.few_shot_examples import FewShotExampleSet
28
+ from neuron_explainer.explanations.prompt_builder import (
29
+ HarmonyMessage,
30
+ PromptBuilder,
31
+ PromptFormat,
32
+ Role,
33
+ )
34
+
35
+ logger = logging.getLogger(__name__)
36
+
37
+ # Our prompts use normalized activation values, which map any range of positive activations to the
38
+ # integers from 0 to 10.
39
+ MAX_NORMALIZED_ACTIVATION = 10
40
+ VALID_ACTIVATION_TOKENS_ORDERED = list(
41
+ str(i) for i in range(MAX_NORMALIZED_ACTIVATION + 1)
42
+ )
43
+ VALID_ACTIVATION_TOKENS = set(VALID_ACTIVATION_TOKENS_ORDERED)
44
+
45
+ # Edge Case #3: The chat-based simulator is confused by end token. Replace it with a "not end token"
46
+ END_OF_TEXT_TOKEN = "<|endoftext|>"
47
+ END_OF_TEXT_TOKEN_REPLACEMENT = "<|not_endoftext|>"
48
+
49
+
50
+ class SimulationType(str, Enum):
51
+ """How to simulate neuron activations. Values correspond to subclasses of NeuronSimulator."""
52
+
53
+ ALL_AT_ONCE = "all_at_once"
54
+ """
55
+ Use a single prompt with <unknown> tokens; calculate EVs using logprobs.
56
+
57
+ Implemented by ExplanationNeuronSimulator.
58
+ """
59
+
60
+ ONE_AT_A_TIME = "one_at_a_time"
61
+ """
62
+ Use a separate prompt for each token being simulated; calculate EVs using logprobs.
63
+
64
+ Implemented by ExplanationTokenByTokenSimulator.
65
+ """
66
+
67
+ @classmethod
68
+ def from_string(cls, s: str) -> SimulationType:
69
+ for simulation_type in SimulationType:
70
+ if simulation_type.value == s:
71
+ return simulation_type
72
+ raise ValueError(f"Invalid simulation type: {s}")
73
+
74
+
75
+ def compute_expected_value(
76
+ norm_probabilities_by_distribution_value: OrderedDict[int, float]
77
+ ) -> float:
78
+ """
79
+ Given a map from distribution values (integers on the range [0, 10]) to normalized
80
+ probabilities, return an expected value for the distribution.
81
+ """
82
+ return np.dot(
83
+ np.array(list(norm_probabilities_by_distribution_value.keys())),
84
+ np.array(list(norm_probabilities_by_distribution_value.values())),
85
+ )
86
+
87
+
88
+ def parse_top_logprobs(top_logprobs: dict[str, float]) -> OrderedDict[int, float]:
89
+ """
90
+ Given a map from tokens to logprobs, return a map from distribution values (integers on the
91
+ range [0, 10]) to unnormalized probabilities (in the sense that they may not sum to 1).
92
+ """
93
+ probabilities_by_distribution_value = OrderedDict()
94
+ for token, logprob in top_logprobs.items():
95
+ if token in VALID_ACTIVATION_TOKENS:
96
+ token_as_int = int(token)
97
+ probabilities_by_distribution_value[token_as_int] = np.exp(logprob)
98
+ return probabilities_by_distribution_value
99
+
100
+
101
+ def compute_predicted_activation_stats_for_token(
102
+ top_logprobs: dict[str, float],
103
+ ) -> tuple[OrderedDict[int, float], float]:
104
+ probabilities_by_distribution_value = parse_top_logprobs(top_logprobs)
105
+ total_p_of_distribution_values = sum(probabilities_by_distribution_value.values())
106
+ norm_probabilities_by_distribution_value = OrderedDict(
107
+ {
108
+ distribution_value: p / total_p_of_distribution_values
109
+ for distribution_value, p in probabilities_by_distribution_value.items()
110
+ }
111
+ )
112
+ expected_value = compute_expected_value(norm_probabilities_by_distribution_value)
113
+ return (
114
+ norm_probabilities_by_distribution_value,
115
+ expected_value,
116
+ )
117
+
118
+
119
+ # Adapted from tether/tether/core/encoder.py.
120
+ def convert_to_byte_array(s: str) -> bytearray:
121
+ byte_array = bytearray()
122
+ assert s.startswith("bytes:"), s
123
+ s = s[6:]
124
+ while len(s) > 0:
125
+ if s[0] == "\\":
126
+ # Hex encoding.
127
+ assert s[1] == "x"
128
+ assert len(s) >= 4
129
+ byte_array.append(int(s[2:4], 16))
130
+ s = s[4:]
131
+ else:
132
+ # Regular ascii encoding.
133
+ byte_array.append(ord(s[0]))
134
+ s = s[1:]
135
+ return byte_array
136
+
137
+
138
+ def handle_byte_encoding(
139
+ response_tokens: Sequence[str], merged_response_index: int
140
+ ) -> tuple[str, int]:
141
+ """
142
+ Handle the case where the current token is a sequence of bytes. This may involve merging
143
+ multiple response tokens into a single token.
144
+ """
145
+ response_token = response_tokens[merged_response_index]
146
+ if response_token.startswith("bytes:"):
147
+ byte_array = bytearray()
148
+ while True:
149
+ byte_array = convert_to_byte_array(response_token) + byte_array
150
+ try:
151
+ # If we can decode the byte array as utf-8, then we're done.
152
+ response_token = byte_array.decode("utf-8")
153
+ break
154
+ except UnicodeDecodeError:
155
+ # If not, then we need to merge the previous response token into the byte
156
+ # array.
157
+ merged_response_index -= 1
158
+ response_token = response_tokens[merged_response_index]
159
+ return response_token, merged_response_index
160
+
161
+
162
+ def was_token_split(
163
+ current_token: str, response_tokens: Sequence[str], start_index: int
164
+ ) -> bool:
165
+ """
166
+ Return whether current_token (a token from the subject model) was split into multiple tokens by
167
+ the simulator model (as represented by the tokens in response_tokens). start_index is the index
168
+ in response_tokens at which to begin looking backward to form a complete token. It is usually
169
+ the first token *before* the delimiter that separates the token from the normalized activation,
170
+ barring some unusual cases.
171
+
172
+ This mainly happens if the subject model uses a different tokenizer than the simulator model.
173
+ But it can also happen in cases where Unicode characters are split. This function handles both
174
+ cases.
175
+ """
176
+ merged_response_tokens = ""
177
+ merged_response_index = start_index
178
+ while len(merged_response_tokens) < len(current_token):
179
+ response_token = response_tokens[merged_response_index]
180
+ response_token, merged_response_index = handle_byte_encoding(
181
+ response_tokens, merged_response_index
182
+ )
183
+ merged_response_tokens = response_token + merged_response_tokens
184
+ merged_response_index -= 1
185
+ # It's possible that merged_response_tokens is longer than current_token at this point,
186
+ # since the between-lines delimiter may have been merged into the original token. But it
187
+ # should always be the case that merged_response_tokens ends with current_token.
188
+ assert merged_response_tokens.endswith(current_token)
189
+ num_merged_tokens = start_index - merged_response_index
190
+ token_was_split = num_merged_tokens > 1
191
+ if token_was_split:
192
+ logger.debug(
193
+ "Warning: token from the subject model was split into 2+ tokens by the simulator model."
194
+ )
195
+ return token_was_split
196
+
197
+
198
+ def parse_simulation_response(
199
+ response: dict[str, Any],
200
+ prompt_format: PromptFormat,
201
+ tokens: Sequence[str],
202
+ ) -> SequenceSimulation:
203
+ """
204
+ Parse an API response to a simulation prompt.
205
+
206
+ Args:
207
+ response: response from the API
208
+ prompt_format: how the prompt was formatted
209
+ tokens: list of tokens as strings in the sequence where the neuron is being simulated
210
+ """
211
+ choice = response["choices"][0]
212
+ if prompt_format == PromptFormat.HARMONY_V4:
213
+ text = choice["message"]["content"]
214
+ elif prompt_format in [
215
+ PromptFormat.NONE,
216
+ PromptFormat.INSTRUCTION_FOLLOWING,
217
+ ]:
218
+ text = choice["text"]
219
+ else:
220
+ raise ValueError(f"Unhandled prompt format {prompt_format}")
221
+ response_tokens = choice["logprobs"]["tokens"]
222
+ choice["logprobs"]["token_logprobs"]
223
+ top_logprobs = choice["logprobs"]["top_logprobs"]
224
+ token_text_offset = choice["logprobs"]["text_offset"]
225
+ # This only works because the sequence "<start>" tokenizes into multiple tokens if it appears in
226
+ # a text sequence in the prompt.
227
+ scoring_start = text.rfind("<start>")
228
+ expected_values = []
229
+ original_sequence_tokens: list[str] = []
230
+ distribution_values: list[list[float]] = []
231
+ distribution_probabilities: list[list[float]] = []
232
+ for i in range(2, len(response_tokens)):
233
+ if len(original_sequence_tokens) == len(tokens):
234
+ # Make sure we haven't hit some sort of off-by-one error.
235
+ # TODO(sbills): Generalize this to handle different tokenizers.
236
+ reached_end = (
237
+ response_tokens[i + 1] == "<" and response_tokens[i + 2] == "end"
238
+ )
239
+ assert reached_end, f"{response_tokens[i-3:i+3]}"
240
+ break
241
+ if token_text_offset[i] >= scoring_start:
242
+ # We're looking for the first token after a tab. This token should be the text
243
+ # "unknown" if hide_activations=True or a normalized activation (0-10) otherwise.
244
+ # If it isn't, that means that the tab is not appearing as a delimiter, but rather
245
+ # as a token, in which case we should move on to the next response token.
246
+ if response_tokens[i - 1] == "\t":
247
+ if response_tokens[i] != "unknown":
248
+ logger.debug(
249
+ "Ignoring tab token that is not followed by an 'unknown' token."
250
+ )
251
+ continue
252
+
253
+ # j represents the index of the token in a "token<tab>activation" line, barring
254
+ # one of the unusual cases handled below.
255
+ j = i - 2
256
+
257
+ current_token = tokens[len(original_sequence_tokens)]
258
+ if current_token == response_tokens[j] or was_token_split(
259
+ current_token, response_tokens, j
260
+ ):
261
+ # We're in the normal case where the tokenization didn't throw off the
262
+ # formatting or in the token-was-split case, which we handle the usual way.
263
+ current_top_logprobs = top_logprobs[i]
264
+
265
+ (
266
+ norm_probabilities_by_distribution_value,
267
+ expected_value,
268
+ ) = compute_predicted_activation_stats_for_token(
269
+ current_top_logprobs,
270
+ )
271
+ current_distribution_values = list(
272
+ norm_probabilities_by_distribution_value.keys()
273
+ )
274
+ current_distribution_probabilities = list(
275
+ norm_probabilities_by_distribution_value.values()
276
+ )
277
+ else:
278
+ # We're in a case where the tokenization resulted in a newline being folded into
279
+ # the token. We can't do our usual prediction of activation stats for the token,
280
+ # since the model did not observe the original token. Instead, we use dummy
281
+ # values. See the TODO elsewhere in this file about coming up with a better
282
+ # prompt format that avoids this situation.
283
+ newline_folded_into_token = "\n" in response_tokens[j]
284
+ assert (
285
+ newline_folded_into_token
286
+ ), f"`{current_token=}` {response_tokens[j-3:j+3]=}"
287
+ logger.debug(
288
+ "Warning: newline before a token<tab>activation line was folded into the token"
289
+ )
290
+ current_distribution_values = []
291
+ current_distribution_probabilities = []
292
+ expected_value = 0.0
293
+
294
+ original_sequence_tokens.append(current_token)
295
+ distribution_values.append(
296
+ [float(v) for v in current_distribution_values]
297
+ )
298
+ distribution_probabilities.append(current_distribution_probabilities)
299
+ expected_values.append(expected_value)
300
+
301
+ return SequenceSimulation(
302
+ tokens=original_sequence_tokens,
303
+ expected_activations=expected_values,
304
+ activation_scale=ActivationScale.SIMULATED_NORMALIZED_ACTIVATIONS,
305
+ distribution_values=distribution_values,
306
+ distribution_probabilities=distribution_probabilities,
307
+ )
308
+
309
+
310
+ class NeuronSimulator(ABC):
311
+ """Abstract base class for simulating neuron behavior."""
312
+
313
+ @abstractmethod
314
+ async def simulate(self, tokens: Sequence[str]) -> SequenceSimulation:
315
+ """Simulate the behavior of a neuron based on an explanation."""
316
+ ...
317
+
318
+
319
+ class ExplanationNeuronSimulator(NeuronSimulator):
320
+ """
321
+ Simulate neuron behavior based on an explanation.
322
+
323
+ This class uses a few-shot prompt with examples of other explanations and activations. This
324
+ prompt allows us to score all of the tokens at once using a nifty trick involving logprobs.
325
+ """
326
+
327
+ def __init__(
328
+ self,
329
+ model_name: str,
330
+ explanation: str,
331
+ max_concurrent: Optional[int] = 10,
332
+ few_shot_example_set: FewShotExampleSet = FewShotExampleSet.ORIGINAL,
333
+ prompt_format: PromptFormat = PromptFormat.INSTRUCTION_FOLLOWING,
334
+ cache: bool = False,
335
+ ):
336
+ self.api_client = ApiClient(
337
+ model_name=model_name, max_concurrent=max_concurrent, cache=cache
338
+ )
339
+ self.explanation = explanation
340
+ self.few_shot_example_set = few_shot_example_set
341
+ self.prompt_format = prompt_format
342
+
343
+ async def simulate(
344
+ self,
345
+ tokens: Sequence[str],
346
+ ) -> SequenceSimulation:
347
+ prompt = self.make_simulation_prompt(tokens)
348
+
349
+ generate_kwargs: dict[str, Any] = {
350
+ "max_tokens": 0,
351
+ "echo": True,
352
+ "logprobs": 15,
353
+ }
354
+ if self.prompt_format == PromptFormat.HARMONY_V4:
355
+ assert isinstance(prompt, list)
356
+ assert isinstance(prompt[0], dict) # Really a HarmonyMessage
357
+ generate_kwargs["messages"] = prompt
358
+ else:
359
+ assert isinstance(prompt, str)
360
+ generate_kwargs["prompt"] = prompt
361
+
362
+ response = await self.api_client.make_request(**generate_kwargs)
363
+ logger.debug("response in score_explanation_by_activations is %s", response)
364
+ result = parse_simulation_response(response, self.prompt_format, tokens)
365
+ logger.debug("result in score_explanation_by_activations is %s", result)
366
+ return result
367
+
368
+ # TODO(sbills): The current token<tab>activation format can result in improper tokenization.
369
+ # In particular, if the token is itself a tab, we may get a single "\t\t" token rather than two
370
+ # "\t" tokens. Consider using a separator that does not appear in any multi-character tokens.
371
+ def make_simulation_prompt(
372
+ self, tokens: Sequence[str]
373
+ ) -> Union[str, list[HarmonyMessage]]:
374
+ """Create a few-shot prompt for predicting neuron activations for the given tokens."""
375
+
376
+ # TODO(sbills): The prompts in this file are subtly different from the ones in explainer.py.
377
+ # Consider reconciling them.
378
+ prompt_builder = PromptBuilder()
379
+ prompt_builder.add_message(
380
+ Role.SYSTEM,
381
+ """We're studying neurons in a neural network.
382
+ Each neuron looks for some particular thing in a short document.
383
+ Look at summary of what the neuron does, and try to predict how it will fire on each token.
384
+
385
+ The activation format is token<tab>activation, activations go from 0 to 10, "unknown" indicates an unknown activation. Most activations will be 0.
386
+ """,
387
+ )
388
+
389
+ few_shot_examples = self.few_shot_example_set.get_examples()
390
+ for i, example in enumerate(few_shot_examples):
391
+ prompt_builder.add_message(
392
+ Role.USER,
393
+ f"\n\nNeuron {i + 1}\nExplanation of neuron {i + 1} behavior: {EXPLANATION_PREFIX} "
394
+ f"{example.explanation}",
395
+ )
396
+ formatted_activation_records = format_activation_records(
397
+ example.activation_records,
398
+ calculate_max_activation(example.activation_records),
399
+ start_indices=example.first_revealed_activation_indices,
400
+ )
401
+ prompt_builder.add_message(
402
+ Role.ASSISTANT, f"\nActivations: {formatted_activation_records}\n"
403
+ )
404
+
405
+ prompt_builder.add_message(
406
+ Role.USER,
407
+ f"\n\nNeuron {len(few_shot_examples) + 1}\nExplanation of neuron "
408
+ f"{len(few_shot_examples) + 1} behavior: {EXPLANATION_PREFIX} "
409
+ f"{self.explanation.strip()}",
410
+ )
411
+ prompt_builder.add_message(
412
+ Role.ASSISTANT,
413
+ f"\nActivations: {format_sequences_for_simulation([tokens])}",
414
+ )
415
+ return prompt_builder.build(self.prompt_format)
416
+
417
+
418
+ class ExplanationTokenByTokenSimulator(NeuronSimulator):
419
+ """
420
+ Simulate neuron behavior based on an explanation.
421
+
422
+ Unlike ExplanationNeuronSimulator, this class uses one few-shot prompt per token to calculate
423
+ expected activations. This is slower. This class gets a one-token completion and calculates an
424
+ expected value from that token's logprobs.
425
+ """
426
+
427
+ def __init__(
428
+ self,
429
+ model_name: str,
430
+ explanation: str,
431
+ max_concurrent: Optional[int] = 10,
432
+ few_shot_example_set: FewShotExampleSet = FewShotExampleSet.NEWER,
433
+ prompt_format: PromptFormat = PromptFormat.INSTRUCTION_FOLLOWING,
434
+ cache: bool = False,
435
+ ):
436
+ assert (
437
+ few_shot_example_set != FewShotExampleSet.ORIGINAL
438
+ ), "This simulator doesn't support the ORIGINAL few-shot example set."
439
+ self.api_client = ApiClient(
440
+ model_name=model_name, max_concurrent=max_concurrent, cache=cache
441
+ )
442
+ self.explanation = explanation
443
+ self.few_shot_example_set = few_shot_example_set
444
+ self.prompt_format = prompt_format
445
+
446
+ async def simulate(
447
+ self,
448
+ tokens: Sequence[str],
449
+ ) -> SequenceSimulation:
450
+ responses_by_token = await asyncio.gather(
451
+ *[
452
+ self._get_activation_stats_for_single_token(
453
+ tokens, self.explanation, token_index
454
+ )
455
+ for token_index in range(len(tokens))
456
+ ]
457
+ )
458
+ expected_values, distribution_values, distribution_probabilities = [], [], []
459
+ for response in responses_by_token:
460
+ activation_logprobs = response["choices"][0]["logprobs"]["top_logprobs"][0]
461
+ (
462
+ norm_probabilities_by_distribution_value,
463
+ expected_value,
464
+ ) = compute_predicted_activation_stats_for_token(
465
+ activation_logprobs,
466
+ )
467
+ distribution_values.append(
468
+ [float(v) for v in norm_probabilities_by_distribution_value.keys()]
469
+ )
470
+ distribution_probabilities.append(
471
+ list(norm_probabilities_by_distribution_value.values())
472
+ )
473
+ expected_values.append(expected_value)
474
+
475
+ result = SequenceSimulation(
476
+ tokens=list(tokens), # SequenceSimulation expects List type
477
+ expected_activations=expected_values,
478
+ activation_scale=ActivationScale.SIMULATED_NORMALIZED_ACTIVATIONS,
479
+ distribution_values=distribution_values,
480
+ distribution_probabilities=distribution_probabilities,
481
+ )
482
+ logger.debug("result in score_explanation_by_activations is %s", result)
483
+ return result
484
+
485
+ async def _get_activation_stats_for_single_token(
486
+ self,
487
+ tokens: Sequence[str],
488
+ explanation: str,
489
+ token_index_to_score: int,
490
+ ) -> dict:
491
+ prompt = self.make_single_token_simulation_prompt(
492
+ tokens,
493
+ explanation,
494
+ token_index_to_score=token_index_to_score,
495
+ )
496
+ return await self.api_client.make_request(
497
+ prompt=prompt, max_tokens=1, echo=False, logprobs=15
498
+ )
499
+
500
+ def _add_single_token_simulation_subprompt(
501
+ self,
502
+ prompt_builder: PromptBuilder,
503
+ activation_record: ActivationRecord,
504
+ neuron_index: int,
505
+ explanation: str,
506
+ token_index_to_score: int,
507
+ end_of_prompt: bool,
508
+ ) -> None:
509
+ trimmed_activation_record = ActivationRecord(
510
+ tokens=activation_record.tokens[: token_index_to_score + 1],
511
+ activations=activation_record.activations[: token_index_to_score + 1],
512
+ )
513
+ prompt_builder.add_message(
514
+ Role.USER,
515
+ f"""
516
+ Neuron {neuron_index}
517
+ Explanation of neuron {neuron_index} behavior: {EXPLANATION_PREFIX} {explanation.strip()}
518
+ Text:
519
+ {"".join(trimmed_activation_record.tokens)}
520
+
521
+ Last token in the text:
522
+ {trimmed_activation_record.tokens[-1]}
523
+
524
+ Last token activation, considering the token in the context in which it appeared in the text:
525
+ """,
526
+ )
527
+ if not end_of_prompt:
528
+ normalized_activations = normalize_activations(
529
+ trimmed_activation_record.activations,
530
+ calculate_max_activation([activation_record]),
531
+ )
532
+ prompt_builder.add_message(
533
+ Role.ASSISTANT,
534
+ str(normalized_activations[-1]) + ("" if end_of_prompt else "\n\n"),
535
+ )
536
+
537
+ def make_single_token_simulation_prompt(
538
+ self,
539
+ tokens: Sequence[str],
540
+ explanation: str,
541
+ token_index_to_score: int,
542
+ ) -> Union[str, list[HarmonyMessage]]:
543
+ """Make a few-shot prompt for predicting the neuron's activation on a single token."""
544
+ assert explanation != ""
545
+ prompt_builder = PromptBuilder()
546
+ prompt_builder.add_message(
547
+ Role.SYSTEM,
548
+ """We're studying neurons in a neural network. Each neuron looks for some particular thing in a short document. Look at an explanation of what the neuron does, and try to predict its activations on a particular token.
549
+
550
+ The activation format is token<tab>activation, and activations range from 0 to 10. Most activations will be 0.
551
+
552
+ """,
553
+ )
554
+
555
+ few_shot_examples = self.few_shot_example_set.get_examples()
556
+ for i, example in enumerate(few_shot_examples):
557
+ prompt_builder.add_message(
558
+ Role.USER,
559
+ f"Neuron {i + 1}\nExplanation of neuron {i + 1} behavior: {EXPLANATION_PREFIX} "
560
+ f"{example.explanation}\n",
561
+ )
562
+ formatted_activation_records = format_activation_records(
563
+ example.activation_records,
564
+ calculate_max_activation(example.activation_records),
565
+ start_indices=None,
566
+ )
567
+ prompt_builder.add_message(
568
+ Role.ASSISTANT,
569
+ f"Activations: {formatted_activation_records}\n\n",
570
+ )
571
+
572
+ prompt_builder.add_message(
573
+ Role.SYSTEM,
574
+ "Now, we're going predict the activation of a new neuron on a single token, "
575
+ "following the same rules as the examples above. Activations still range from 0 to 10.",
576
+ )
577
+ single_token_example = (
578
+ self.few_shot_example_set.get_single_token_prediction_example()
579
+ )
580
+ assert single_token_example.token_index_to_score is not None
581
+ self._add_single_token_simulation_subprompt(
582
+ prompt_builder,
583
+ single_token_example.activation_records[0],
584
+ len(few_shot_examples) + 1,
585
+ explanation,
586
+ token_index_to_score=single_token_example.token_index_to_score,
587
+ end_of_prompt=False,
588
+ )
589
+
590
+ activation_record = ActivationRecord(
591
+ tokens=list(
592
+ tokens[: token_index_to_score + 1]
593
+ ), # ActivationRecord expects List type.
594
+ activations=[0.0] * len(tokens),
595
+ )
596
+ self._add_single_token_simulation_subprompt(
597
+ prompt_builder,
598
+ activation_record,
599
+ len(few_shot_examples) + 2,
600
+ explanation,
601
+ token_index_to_score,
602
+ end_of_prompt=True,
603
+ )
604
+ return prompt_builder.build(
605
+ self.prompt_format, allow_extra_system_messages=True
606
+ )
607
+
608
+
609
+ def _format_record_for_logprob_free_simulation(
610
+ activation_record: ActivationRecord,
611
+ include_activations: bool = False,
612
+ max_activation: Optional[float] = None,
613
+ ) -> str:
614
+ response = ""
615
+ if include_activations:
616
+ assert max_activation is not None
617
+ assert len(activation_record.tokens) == len(
618
+ activation_record.activations
619
+ ), f"{len(activation_record.tokens)=}, {len(activation_record.activations)=}"
620
+ normalized_activations = normalize_activations(
621
+ activation_record.activations, max_activation=max_activation
622
+ )
623
+ for i, token in enumerate(activation_record.tokens):
624
+ # Edge Case #3: End tokens confuse the chat-based simulator. Replace end token with "not end token".
625
+ if token.strip() == END_OF_TEXT_TOKEN:
626
+ token = END_OF_TEXT_TOKEN_REPLACEMENT
627
+ # We use a weird unicode character here to make it easier to parse the response (can split on "༗\n").
628
+ if include_activations:
629
+ response += f"{token}\t{normalized_activations[i]}༗\n"
630
+ else:
631
+ response += f"{token}\t༗\n"
632
+ return response
633
+
634
+
635
+ def _format_record_for_logprob_free_simulation_json(
636
+ explanation: str,
637
+ activation_record: ActivationRecord,
638
+ include_activations: bool = False,
639
+ ) -> str:
640
+ if include_activations:
641
+ assert len(activation_record.tokens) == len(
642
+ activation_record.activations
643
+ ), f"{len(activation_record.tokens)=}, {len(activation_record.activations)=}"
644
+ return json.dumps(
645
+ {
646
+ "to_find": explanation,
647
+ "document": "".join(activation_record.tokens),
648
+ "activations": [
649
+ {
650
+ "token": token,
651
+ "activation": (
652
+ activation_record.activations[i]
653
+ if include_activations
654
+ else None
655
+ ),
656
+ }
657
+ for i, token in enumerate(activation_record.tokens)
658
+ ],
659
+ }
660
+ )
661
+
662
+
663
+ def _parse_no_logprobs_completion_json(
664
+ completion: str,
665
+ tokens: Sequence[str],
666
+ ) -> Sequence[float]:
667
+ """
668
+ Parse a completion into a list of simulated activations. If the model did not faithfully
669
+ reproduce the token sequence, return a list of 0s. If the model's activation for a token
670
+ is not a number between 0 and 10 (inclusive), substitute 0.
671
+
672
+ Args:
673
+ completion: completion from the API
674
+ tokens: list of tokens as strings in the sequence where the neuron is being simulated
675
+ """
676
+
677
+ logger.debug("for tokens:\n%s", tokens)
678
+ logger.debug("received completion:\n%s", completion)
679
+
680
+ zero_prediction = [0] * len(tokens)
681
+
682
+ try:
683
+ completion = json.loads(completion)
684
+ if "activations" not in completion:
685
+ logger.error(
686
+ "The key 'activations' is not in the completion:\n%s\nExpected Tokens:\n%s",
687
+ json.dumps(completion),
688
+ tokens,
689
+ )
690
+ return zero_prediction
691
+ activations = completion["activations"]
692
+ if len(activations) != len(tokens):
693
+ logger.error(
694
+ "Tokens and activations length did not match:\n%s\nExpected Tokens:\n%s",
695
+ json.dumps(completion),
696
+ tokens,
697
+ )
698
+ return zero_prediction
699
+ predicted_activations = []
700
+ # check that there is a token and activation value
701
+ # no need to double check the token matches exactly
702
+ for i, activation in enumerate(activations):
703
+ if "token" not in activation:
704
+ logger.error(
705
+ "The key 'token' is not in activation:\n%s\nCompletion:%s\nExpected Tokens:\n%s",
706
+ activation,
707
+ json.dumps(completion),
708
+ tokens,
709
+ )
710
+ predicted_activations.append(0)
711
+ continue
712
+ if "activation" not in activation:
713
+ logger.error(
714
+ "The key 'activation' is not in activation:\n%s\nCompletion:%s\nExpected Tokens:\n%s",
715
+ activation,
716
+ json.dumps(completion),
717
+ tokens,
718
+ )
719
+ predicted_activations.append(0)
720
+ continue
721
+ # Ensure activation value is between 0-10 inclusive
722
+ try:
723
+ predicted_activation_float = float(activation["activation"])
724
+ if (
725
+ predicted_activation_float < 0
726
+ or predicted_activation_float > MAX_NORMALIZED_ACTIVATION
727
+ ):
728
+ logger.error(
729
+ "activation value out of range: %s\nCompletion:%s\nExpected Tokens:\n%s",
730
+ predicted_activation_float,
731
+ json.dumps(completion),
732
+ tokens,
733
+ )
734
+ predicted_activations.append(0)
735
+ else:
736
+ predicted_activations.append(predicted_activation_float)
737
+ except ValueError:
738
+ logger.error(
739
+ "activation value invalid: %s\nCompletion:%s\nExpected Tokens:\n%s",
740
+ activation["activation"],
741
+ json.dumps(completion),
742
+ tokens,
743
+ )
744
+ predicted_activations.append(0)
745
+ except TypeError:
746
+ logger.error(
747
+ "activation value incorrect type: %s\nCompletion:%s\nExpected Tokens:\n%s",
748
+ activation["activation"],
749
+ json.dumps(completion),
750
+ tokens,
751
+ )
752
+ predicted_activations.append(0)
753
+ logger.debug("predicted activations: %s", predicted_activations)
754
+ return predicted_activations
755
+
756
+ except json.JSONDecodeError:
757
+ logger.warning(
758
+ "Failed to parse completion JSON:\n%s\nExpected Tokens:\n%s",
759
+ completion,
760
+ tokens,
761
+ )
762
+ return zero_prediction
763
+
764
+
765
+ def _parse_no_logprobs_completion(
766
+ completion: str,
767
+ tokens: Sequence[str],
768
+ ) -> Sequence[float]:
769
+ """
770
+ Parse a completion into a list of simulated activations. If the model did not faithfully
771
+ reproduce the token sequence, return a list of 0s. If the model's activation for a token
772
+ is not a number between 0 and 10 (inclusive), substitute 0.
773
+
774
+ Args:
775
+ completion: completion from the API
776
+ tokens: list of tokens as strings in the sequence where the neuron is being simulated
777
+ """
778
+
779
+ logger.debug("for tokens:\n%s", tokens)
780
+ logger.debug("received completion:\n%s", completion)
781
+
782
+ zero_prediction = [0] * len(tokens)
783
+ # FIX: Strip the last ༗\n, otherwise all last activations are invalid
784
+ token_lines = completion.strip("\n").strip("༗\n").split("༗\n")
785
+ # Edge Case #2: Sometimes GPT doesn't use the special character when it answers, it only uses the \n"
786
+ # The fix is to try splitting by \n if we detect that the response isn't the right format
787
+ # TODO: If there are also line breaks in the text, this will probably break
788
+ if (len(token_lines)) == 1:
789
+ token_lines = completion.strip("\n").strip("༗\n").split("\n")
790
+ logger.debug("parsed completion into token_lines as:\n%s", token_lines)
791
+
792
+ start_line_index = None
793
+ for i, token_line in enumerate(token_lines):
794
+ if (
795
+ token_line.startswith(f"{tokens[0]}\t")
796
+ # Edge Case #1: GPT often omits the space before the first token.
797
+ # Allow the returned token line to be either " token" or "token".
798
+ or f" {token_line}".startswith(f"{tokens[0]}\t")
799
+ # Edge Case #3: Allow our "not end token" replacement
800
+ or (
801
+ token_line.startswith(END_OF_TEXT_TOKEN_REPLACEMENT)
802
+ and tokens[0].strip() == END_OF_TEXT_TOKEN
803
+ )
804
+ ):
805
+ logger.debug("start_line_index is: %s", start_line_index)
806
+ logger.debug("matched token %s with token_line %s", tokens[0], token_line)
807
+ start_line_index = i
808
+ break
809
+
810
+ # If we didn't find the first token, or if the number of lines in the completion doesn't match
811
+ # the number of tokens, return a list of 0s.
812
+ if start_line_index is None or len(token_lines) - start_line_index != len(tokens):
813
+ logger.debug(
814
+ "didn't find first token or number of lines didn't match, returning all zeroes"
815
+ )
816
+ return zero_prediction
817
+
818
+ predicted_activations = []
819
+ for i, token_line in enumerate(token_lines[start_line_index:]):
820
+ if (
821
+ not token_line.startswith(f"{tokens[i]}\t")
822
+ # Edge Case #1: GPT often omits the space before the token.
823
+ # Allow the returned token line to be either " token" or "token".
824
+ and not f" {token_line}".startswith(f"{tokens[i]}\t")
825
+ # Edge Case #3: Allow our "not end token" replacement
826
+ and not token_line.startswith(END_OF_TEXT_TOKEN_REPLACEMENT)
827
+ ):
828
+ logger.debug(
829
+ "failed to match token %s with token_line %s, returning all zeroes",
830
+ tokens[i],
831
+ token_line,
832
+ )
833
+ return zero_prediction
834
+ predicted_activation_split = token_line.split("\t")
835
+ # Ensure token line has correct size after splitting. If not then assume it's a zero.
836
+ if len(predicted_activation_split) != 2:
837
+ logger.debug("tokenline split invalid size: %s", token_line)
838
+ predicted_activations.append(0)
839
+ continue
840
+ predicted_activation = predicted_activation_split[1]
841
+ # Sometimes GPT the activation value is not a float (GPT likes to append an extra ༗).
842
+ # In all cases if the activation is not numerically parseable, set it to 0
843
+ try:
844
+ predicted_activation_float = float(predicted_activation)
845
+ if (
846
+ predicted_activation_float < 0
847
+ or predicted_activation_float > MAX_NORMALIZED_ACTIVATION
848
+ ):
849
+ logger.debug(
850
+ "activation value out of range: %s", predicted_activation_float
851
+ )
852
+ predicted_activations.append(0)
853
+ else:
854
+ predicted_activations.append(predicted_activation_float)
855
+ except ValueError:
856
+ logger.debug("activation value not numeric: %s", predicted_activation)
857
+ predicted_activations.append(0)
858
+ logger.debug("predicted activations: %s", predicted_activations)
859
+ return predicted_activations
860
+
861
+
862
+ class LogprobFreeExplanationTokenSimulator(NeuronSimulator):
863
+ """
864
+ Simulate neuron behavior based on an explanation.
865
+
866
+ Unlike ExplanationNeuronSimulator and ExplanationTokenByTokenSimulator, this class does not rely on
867
+ logprobs to calculate expected activations. Instead, it uses a few-shot prompt that displays all of the
868
+ tokens at once, and request that the model repeat the tokens with the activations appended. Sampling
869
+ is with temperature = 0. Thus, the activations are deterministic. Also, each activation for a token
870
+ is a function of all the activations that came previously and all of the tokens in the sequence, not
871
+ just the current and previous tokens. In the case where the model does not faithfully reproduce the
872
+ token sequence, the simulator will return a response where every predicted activation is 0. Example prompt as follows:
873
+
874
+ Explanation: Explanation 1
875
+
876
+ Sequence 1 Tokens Without Activations:
877
+
878
+ A\t_
879
+ B\t_
880
+ C\t_
881
+
882
+ Sequence 1 Tokens With Activations:
883
+
884
+ A\t4_
885
+ B\t10_
886
+ C\t0_
887
+
888
+ Sequence 2 Tokens Without Activations:
889
+
890
+ D\t_
891
+ E\t_
892
+ F\t_
893
+
894
+ Sequence 2 Tokens With Activations:
895
+
896
+ D\t3_
897
+ E\t6_
898
+ F\t9_
899
+
900
+ Explanation: Explanation 2
901
+
902
+ Sequence 1 Tokens Without Activations:
903
+
904
+ G\t_
905
+ H\t_
906
+ I\t_
907
+
908
+ Sequence 1 Tokens With Activations:
909
+ <start sampling here>
910
+
911
+ G\t2_
912
+ H\t0_
913
+ I\t3_
914
+
915
+ """
916
+
917
+ def __init__(
918
+ self,
919
+ model_name: str,
920
+ explanation: str,
921
+ max_concurrent: Optional[int] = 10,
922
+ json_mode: Optional[bool] = True,
923
+ few_shot_example_set: FewShotExampleSet = FewShotExampleSet.NEWER,
924
+ prompt_format: PromptFormat = PromptFormat.HARMONY_V4,
925
+ cache: bool = False,
926
+ ):
927
+ assert (
928
+ few_shot_example_set != FewShotExampleSet.ORIGINAL
929
+ ), "This simulator doesn't support the ORIGINAL few-shot example set."
930
+ self.api_client = ApiClient(
931
+ model_name=model_name, max_concurrent=max_concurrent, cache=cache
932
+ )
933
+ self.json_mode = json_mode
934
+ self.explanation = explanation
935
+ self.few_shot_example_set = few_shot_example_set
936
+ self.prompt_format = prompt_format
937
+
938
+ async def simulate(
939
+ self,
940
+ tokens: Sequence[str],
941
+ ) -> SequenceSimulation:
942
+ if self.json_mode:
943
+ prompt = self._make_simulation_prompt_json(
944
+ tokens,
945
+ self.explanation,
946
+ )
947
+ response = await self.api_client.make_request(
948
+ messages=prompt, max_tokens=2000, temperature=0, json_mode=True
949
+ )
950
+ assert len(response["choices"]) == 1
951
+ choice = response["choices"][0]
952
+ completion = choice["message"]["content"]
953
+ predicted_activations = _parse_no_logprobs_completion_json(
954
+ completion, tokens
955
+ )
956
+ else:
957
+ prompt = self._make_simulation_prompt(
958
+ tokens,
959
+ self.explanation,
960
+ )
961
+ response = await self.api_client.make_request(
962
+ messages=prompt, max_tokens=1000, temperature=0
963
+ )
964
+ assert len(response["choices"]) == 1
965
+ choice = response["choices"][0]
966
+ completion = choice["message"]["content"]
967
+ predicted_activations = _parse_no_logprobs_completion(completion, tokens)
968
+
969
+ result = SequenceSimulation(
970
+ activation_scale=ActivationScale.SIMULATED_NORMALIZED_ACTIVATIONS,
971
+ expected_activations=predicted_activations,
972
+ # Since the predicted activation is just a sampled token, we don't have a distribution.
973
+ distribution_values=[],
974
+ distribution_probabilities=[],
975
+ tokens=list(tokens), # SequenceSimulation expects List type
976
+ )
977
+ logger.debug("result in score_explanation_by_activations is %s", result)
978
+ return result
979
+
980
+ def _make_simulation_prompt_json(
981
+ self,
982
+ tokens: Sequence[str],
983
+ explanation: str,
984
+ ) -> Union[str, list[HarmonyMessage]]:
985
+ """Make a few-shot prompt for predicting the neuron's activations on a sequence."""
986
+ """NOTE: The JSON version does not give GPT multiple sequence examples per neuron."""
987
+ assert explanation != ""
988
+ prompt_builder = PromptBuilder()
989
+ prompt_builder.add_message(
990
+ Role.SYSTEM,
991
+ """We're studying neurons in a neural network. Each neuron looks for certain things in a short document. Your task is to read the explanation of what the neuron does, and predict the neuron's activations for each token in the document.
992
+
993
+ For each document, you will see the full text of the document, then the tokens in the document with the activation left blank. You will print, in valid json, the exact same tokens verbatim, but with the activation values filled in according to the explanation. Pay special attention to the explanation's description of the context and order of tokens or words.
994
+
995
+ Fill out the activation values from 0 to 10. Please think carefully.";
996
+ """,
997
+ )
998
+
999
+ few_shot_examples = self.few_shot_example_set.get_examples()
1000
+ for example in few_shot_examples:
1001
+ """
1002
+ {
1003
+ "to_find": "hello",
1004
+ "document": "The",
1005
+ "activations": [
1006
+ {
1007
+ "token": "The",
1008
+ "activation": null
1009
+ }
1010
+ ]
1011
+ }
1012
+ """
1013
+ prompt_builder.add_message(
1014
+ Role.USER,
1015
+ _format_record_for_logprob_free_simulation_json(
1016
+ explanation=example.explanation,
1017
+ activation_record=example.activation_records[0],
1018
+ include_activations=False,
1019
+ ),
1020
+ )
1021
+ """
1022
+ {
1023
+ "to_find": "hello",
1024
+ "document": "The",
1025
+ "activations": [
1026
+ {
1027
+ "token": "The",
1028
+ "activation": 10
1029
+ }
1030
+ ]
1031
+ }
1032
+ """
1033
+ prompt_builder.add_message(
1034
+ Role.ASSISTANT,
1035
+ _format_record_for_logprob_free_simulation_json(
1036
+ explanation=example.explanation,
1037
+ activation_record=example.activation_records[0],
1038
+ include_activations=True,
1039
+ ),
1040
+ )
1041
+ """
1042
+ {
1043
+ "to_find": "hello",
1044
+ "document": "The",
1045
+ "activations": [
1046
+ {
1047
+ "token": "The",
1048
+ "activation": null
1049
+ }
1050
+ ]
1051
+ }
1052
+ """
1053
+ prompt_builder.add_message(
1054
+ Role.USER,
1055
+ _format_record_for_logprob_free_simulation_json(
1056
+ explanation=explanation,
1057
+ activation_record=ActivationRecord(tokens=tokens, activations=[]),
1058
+ include_activations=False,
1059
+ ),
1060
+ )
1061
+ return prompt_builder.build(
1062
+ self.prompt_format, allow_extra_system_messages=True
1063
+ )
1064
+
1065
+ def _make_simulation_prompt(
1066
+ self,
1067
+ tokens: Sequence[str],
1068
+ explanation: str,
1069
+ ) -> Union[str, list[HarmonyMessage]]:
1070
+ """Make a few-shot prompt for predicting the neuron's activations on a sequence."""
1071
+ assert explanation != ""
1072
+ prompt_builder = PromptBuilder()
1073
+ prompt_builder.add_message(
1074
+ Role.SYSTEM,
1075
+ """We're studying neurons in a neural network. Each neuron looks for some particular thing in a short document. Look at an explanation of what the neuron does, and try to predict its activations on a particular token.
1076
+
1077
+ The activation format is token<tab>activation, and activations range from 0 to 10. Most activations will be 0.
1078
+ For each sequence, you will see the tokens in the sequence where the activations are left blank. You will print the exact same tokens verbatim, but with the activations filled in according to the explanation.
1079
+ """,
1080
+ )
1081
+
1082
+ few_shot_examples = self.few_shot_example_set.get_examples()
1083
+ for i, example in enumerate(few_shot_examples):
1084
+ few_shot_example_max_activation = calculate_max_activation(
1085
+ example.activation_records
1086
+ )
1087
+
1088
+ prompt_builder.add_message(
1089
+ Role.USER,
1090
+ f"Neuron {i + 1}\nExplanation of neuron {i + 1} behavior: {EXPLANATION_PREFIX} "
1091
+ f"{example.explanation}\n\n"
1092
+ f"Sequence 1 Tokens without Activations:\n{_format_record_for_logprob_free_simulation(example.activation_records[0], include_activations=False)}\n\n"
1093
+ f"Sequence 1 Tokens with Activations:\n",
1094
+ )
1095
+ prompt_builder.add_message(
1096
+ Role.ASSISTANT,
1097
+ f"{_format_record_for_logprob_free_simulation(example.activation_records[0], include_activations=True, max_activation=few_shot_example_max_activation)}\n\n",
1098
+ )
1099
+
1100
+ for record_index, record in enumerate(example.activation_records[1:]):
1101
+ prompt_builder.add_message(
1102
+ Role.USER,
1103
+ f"Sequence {record_index + 2} Tokens without Activations:\n{_format_record_for_logprob_free_simulation(record, include_activations=False)}\n\n"
1104
+ f"Sequence {record_index + 2} Tokens with Activations:\n",
1105
+ )
1106
+ prompt_builder.add_message(
1107
+ Role.ASSISTANT,
1108
+ f"{_format_record_for_logprob_free_simulation(record, include_activations=True, max_activation=few_shot_example_max_activation)}\n\n",
1109
+ )
1110
+
1111
+ neuron_index = len(few_shot_examples) + 1
1112
+ prompt_builder.add_message(
1113
+ Role.USER,
1114
+ f"Neuron {neuron_index}\nExplanation of neuron {neuron_index} behavior: {EXPLANATION_PREFIX} "
1115
+ f"{explanation}\n\n"
1116
+ f"Sequence 1 Tokens without Activations:\n{_format_record_for_logprob_free_simulation(ActivationRecord(tokens=tokens, activations=[]), include_activations=False)}\n\n"
1117
+ f"Sequence 1 Tokens with Activations:\n",
1118
+ )
1119
+ return prompt_builder.build(
1120
+ self.prompt_format, allow_extra_system_messages=True
1121
+ )