interp-engine 0.0.24__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- interp_engine-0.0.24.dist-info/METADATA +21 -0
- interp_engine-0.0.24.dist-info/RECORD +28 -0
- interp_engine-0.0.24.dist-info/WHEEL +4 -0
- interp_engine-0.0.24.dist-info/licenses/LICENSE +19 -0
- neuron_explainer/__init__.py +0 -0
- neuron_explainer/activations/__init__.py +0 -0
- neuron_explainer/activations/activation_records.py +130 -0
- neuron_explainer/activations/activations.py +311 -0
- neuron_explainer/activations/attention_utils.py +121 -0
- neuron_explainer/activations/token_connections.py +59 -0
- neuron_explainer/api_client.py +190 -0
- neuron_explainer/azure.py +5 -0
- neuron_explainer/explanations/__init__.py +0 -0
- neuron_explainer/explanations/calibrated_simulator.py +194 -0
- neuron_explainer/explanations/explainer.py +2585 -0
- neuron_explainer/explanations/explanations.py +230 -0
- neuron_explainer/explanations/few_shot_examples.py +3125 -0
- neuron_explainer/explanations/prompt_builder.py +118 -0
- neuron_explainer/explanations/puzzles.json +399 -0
- neuron_explainer/explanations/puzzles.py +50 -0
- neuron_explainer/explanations/scoring.py +155 -0
- neuron_explainer/explanations/simulator.py +1121 -0
- neuron_explainer/explanations/test_explainer.py +227 -0
- neuron_explainer/explanations/test_simulator.py +269 -0
- neuron_explainer/explanations/token_space_few_shot_examples.py +212 -0
- neuron_explainer/fast_dataclasses/__init__.py +3 -0
- neuron_explainer/fast_dataclasses/fast_dataclasses.py +85 -0
- neuron_explainer/fast_dataclasses/test_fast_dataclasses.py +83 -0
|
@@ -0,0 +1,118 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from enum import Enum
|
|
4
|
+
from typing import TypedDict, Union
|
|
5
|
+
|
|
6
|
+
import tiktoken
|
|
7
|
+
|
|
8
|
+
HarmonyMessage = TypedDict(
|
|
9
|
+
"HarmonyMessage",
|
|
10
|
+
{
|
|
11
|
+
"role": str,
|
|
12
|
+
"content": str,
|
|
13
|
+
},
|
|
14
|
+
)
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
class PromptFormat(str, Enum):
|
|
18
|
+
"""
|
|
19
|
+
Different ways of formatting the components of a prompt into the format accepted by the relevant
|
|
20
|
+
API server endpoint.
|
|
21
|
+
"""
|
|
22
|
+
|
|
23
|
+
NONE = "none"
|
|
24
|
+
"""Suitable for use with models that don't use special tokens for instructions."""
|
|
25
|
+
INSTRUCTION_FOLLOWING = "instruction_following"
|
|
26
|
+
"""Suitable for IF models that use <|endofprompt|>."""
|
|
27
|
+
HARMONY_V4 = "harmony_v4"
|
|
28
|
+
"""
|
|
29
|
+
Suitable for Harmony models that use a structured turn-taking role+content format. Generates a
|
|
30
|
+
list of HarmonyMessage dicts that can be sent to the /chat/completions endpoint.
|
|
31
|
+
"""
|
|
32
|
+
|
|
33
|
+
@classmethod
|
|
34
|
+
def from_string(cls, s: str) -> PromptFormat:
|
|
35
|
+
for prompt_format in cls:
|
|
36
|
+
if prompt_format.value == s:
|
|
37
|
+
return prompt_format
|
|
38
|
+
raise ValueError(f"{s} is not a valid PromptFormat")
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
class Role(str, Enum):
|
|
42
|
+
"""See https://platform.openai.com/docs/guides/chat"""
|
|
43
|
+
|
|
44
|
+
SYSTEM = "system"
|
|
45
|
+
USER = "user"
|
|
46
|
+
ASSISTANT = "assistant"
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
class PromptBuilder:
|
|
50
|
+
"""Class for accumulating components of a prompt and then formatting them into an output."""
|
|
51
|
+
|
|
52
|
+
def __init__(self) -> None:
|
|
53
|
+
self._messages: list[HarmonyMessage] = []
|
|
54
|
+
|
|
55
|
+
def add_message(self, role: Role, message: str) -> None:
|
|
56
|
+
self._messages.append(HarmonyMessage(role=role, content=message))
|
|
57
|
+
|
|
58
|
+
def prompt_length_in_tokens(self, prompt_format: PromptFormat) -> int:
|
|
59
|
+
# TODO(sbills): Make the model/encoding configurable. This implementation assumes GPT-4.
|
|
60
|
+
encoding = tiktoken.get_encoding("cl100k_base")
|
|
61
|
+
if prompt_format == PromptFormat.HARMONY_V4:
|
|
62
|
+
# Approximately-correct implementation adapted from this documentation:
|
|
63
|
+
# https://platform.openai.com/docs/guides/chat/introduction
|
|
64
|
+
num_tokens = 0
|
|
65
|
+
for message in self._messages:
|
|
66
|
+
num_tokens += (
|
|
67
|
+
4 # every message follows <|im_start|>{role/name}\n{content}<|im_end|>\n
|
|
68
|
+
)
|
|
69
|
+
num_tokens += len(encoding.encode(message["content"], allowed_special="all"))
|
|
70
|
+
num_tokens += 2 # every reply is primed with <|im_start|>assistant
|
|
71
|
+
return num_tokens
|
|
72
|
+
else:
|
|
73
|
+
prompt_str = self.build(prompt_format)
|
|
74
|
+
assert isinstance(prompt_str, str)
|
|
75
|
+
return len(encoding.encode(prompt_str, allowed_special="all"))
|
|
76
|
+
|
|
77
|
+
def build(
|
|
78
|
+
self, prompt_format: PromptFormat, *, allow_extra_system_messages: bool = False
|
|
79
|
+
) -> Union[str, list[HarmonyMessage]]:
|
|
80
|
+
"""
|
|
81
|
+
Validates the messages added so far (reasonable alternation of assistant vs. user, etc.)
|
|
82
|
+
and returns either a regular string (maybe with <|endofprompt|> tokens) or a list of
|
|
83
|
+
HarmonyMessages suitable for use with the /chat/completions endpoint.
|
|
84
|
+
|
|
85
|
+
The `allow_extra_system_messages` parameter allows the caller to specify that the prompt
|
|
86
|
+
should be allowed to contain system messages after the very first one.
|
|
87
|
+
"""
|
|
88
|
+
# Create a deep copy of the messages so we can modify it and so that the caller can't
|
|
89
|
+
# modify the internal state of this object.
|
|
90
|
+
messages = [message.copy() for message in self._messages]
|
|
91
|
+
|
|
92
|
+
expected_next_role = Role.SYSTEM
|
|
93
|
+
for message in messages:
|
|
94
|
+
role = message["role"]
|
|
95
|
+
assert role == expected_next_role or (
|
|
96
|
+
allow_extra_system_messages and role == Role.SYSTEM
|
|
97
|
+
), f"Expected message from {expected_next_role} but got message from {role}"
|
|
98
|
+
if role == Role.SYSTEM:
|
|
99
|
+
expected_next_role = Role.USER
|
|
100
|
+
elif role == Role.USER:
|
|
101
|
+
expected_next_role = Role.ASSISTANT
|
|
102
|
+
elif role == Role.ASSISTANT:
|
|
103
|
+
expected_next_role = Role.USER
|
|
104
|
+
|
|
105
|
+
if prompt_format == PromptFormat.INSTRUCTION_FOLLOWING:
|
|
106
|
+
last_user_message = None
|
|
107
|
+
for message in messages:
|
|
108
|
+
if message["role"] == Role.USER:
|
|
109
|
+
last_user_message = message
|
|
110
|
+
assert last_user_message is not None
|
|
111
|
+
last_user_message["content"] += "<|endofprompt|>"
|
|
112
|
+
|
|
113
|
+
if prompt_format == PromptFormat.HARMONY_V4:
|
|
114
|
+
return messages
|
|
115
|
+
elif prompt_format in [PromptFormat.NONE, PromptFormat.INSTRUCTION_FOLLOWING]:
|
|
116
|
+
return "".join(message["content"] for message in messages)
|
|
117
|
+
else:
|
|
118
|
+
raise ValueError(f"Unknown prompt format: {prompt_format}")
|