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,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}")