Skip to content

Commit fdbdf2b

Browse files
authored
Update models and address auto import of agents (#35)
* Experimental changes * updating models * Dealing with auto import of agents/. Add opus 4.6 models
1 parent 513579e commit fdbdf2b

5 files changed

Lines changed: 93 additions & 49 deletions

File tree

agents/llm.py

Lines changed: 15 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -25,6 +25,20 @@
2525
" When stuck, try using the `help` command to see what commands are available."
2626
)
2727

28+
CLAUDE_MODELS = [
29+
"claude-3.5-haiku",
30+
"claude-3.5-sonnet",
31+
"claude-3.5-sonnet-latest",
32+
"claude-3.7-sonnet",
33+
"claude-4-sonnet",
34+
"claude-4-opus",
35+
"claude-opus-4.5",
36+
"claude-opus-4.6",
37+
"claude-sonnet-4.5",
38+
"claude-sonnet-4.6",
39+
"claude-haiku-4.5",
40+
]
41+
2842

2943
class LLMAgent(tales.Agent):
3044

@@ -98,12 +112,7 @@ def act(self, obs, reward, done, infos):
98112
"seed": self.seed,
99113
"stream": False,
100114
}
101-
if self.llm in [
102-
"claude-3.5-haiku",
103-
"claude-3.5-sonnet",
104-
"claude-3.5-sonnet-latest",
105-
"claude-3.7-sonnet",
106-
]:
115+
if self.llm in CLAUDE_MODELS:
107116
# For these models, we cannot set the seed.
108117
llm_kwargs.pop("seed")
109118

agents/reasoning.py

Lines changed: 50 additions & 38 deletions
Original file line numberDiff line numberDiff line change
@@ -28,6 +28,36 @@
2828

2929
DEEPSEEK_CHAT_TEMPLATE_NO_THINK = "{% if not add_generation_prompt is defined %}{% set add_generation_prompt = false %}{% endif %}{% set ns = namespace(is_first=false, is_tool=false, is_output_first=true, system_prompt='') %}{%- for message in messages %}{%- if message['role'] == 'system' %}{% set ns.system_prompt = message['content'] %}{%- endif %}{%- endfor %}{{bos_token}}{{ns.system_prompt}}{%- for message in messages %}{%- if message['role'] == 'user' %}{%- set ns.is_tool = false -%}{{'<|User|>' + message['content']}}{%- endif %}{%- if message['role'] == 'assistant' and message['content'] is none %}{%- set ns.is_tool = false -%}{%- for tool in message['tool_calls']%}{%- if not ns.is_first %}{{'<|Assistant|><|tool▁calls▁begin|><|tool▁call▁begin|>' + tool['type'] + '<|tool▁sep|>' + tool['function']['name'] + '\\n' + '```json' + '\\n' + tool['function']['arguments'] + '\\n' + '```' + '<|tool▁call▁end|>'}}{%- set ns.is_first = true -%}{%- else %}{{'\\n' + '<|tool▁call▁begin|>' + tool['type'] + '<|tool▁sep|>' + tool['function']['name'] + '\\n' + '```json' + '\\n' + tool['function']['arguments'] + '\\n' + '```' + '<|tool▁call▁end|>'}}{{'<|tool▁calls▁end|><|end▁of▁sentence|>'}}{%- endif %}{%- endfor %}{%- endif %}{%- if message['role'] == 'assistant' and message['content'] is not none %}{%- if ns.is_tool %}{{'<|tool▁outputs▁end|>' + message['content'] + '<|end▁of▁sentence|>'}}{%- set ns.is_tool = false -%}{%- else %}{% set content = message['content'] %}{% if '</think>' in content %}{% set content = content.split('</think>')[-1] %}{% endif %}{{'<|Assistant|>' + content + '<|end▁of▁sentence|>'}}{%- endif %}{%- endif %}{%- if message['role'] == 'tool' %}{%- set ns.is_tool = true -%}{%- if ns.is_output_first %}{{'<|tool▁outputs▁begin|><|tool▁output▁begin|>' + message['content'] + '<|tool▁output▁end|>'}}{%- set ns.is_output_first = false %}{%- else %}{{'\\n<|tool▁output▁begin|>' + message['content'] + '<|tool▁output▁end|>'}}{%- endif %}{%- endif %}{%- endfor -%}{% if ns.is_tool %}{{'<|tool▁outputs▁end|>'}}{% endif %}{% if add_generation_prompt and not ns.is_tool %}{{'<|Assistant|><think>\\n</think>\\n'}}{% endif %}"
3030

31+
CLAUDE_MODELS = [
32+
"claude-3.7-sonnet",
33+
"claude-4-sonnet",
34+
"claude-4-opus",
35+
"claude-sonnet-4.5",
36+
"claude-haiku-4.5",
37+
"claude-opus-4.5",
38+
"claude-opus-4.6",
39+
"claude-sonnet-4.6",
40+
]
41+
42+
OPENAI_MODELS = [
43+
"o1",
44+
"o1-mini",
45+
"o1-preview",
46+
"o3-mini",
47+
"o4-mini",
48+
"o3",
49+
"gpt-5.1",
50+
"gpt-5.2",
51+
"gpt-5",
52+
"gpt-5-mini",
53+
"gpt-5-nano",
54+
]
55+
56+
GEMINI_MODELS = [
57+
"gemini-2.5-pro",
58+
"gemini-3-pro-preview",
59+
]
60+
3161

3262
class ReasoningAgent(tales.Agent):
3363

@@ -117,42 +147,25 @@ def act(self, obs, reward, done, infos):
117147
"stream": True, # Should prevent openai.APITimeoutError
118148
}
119149
if isinstance(self.reasoning_effort, int):
120-
if self.llm in ["claude-3.7-sonnet"]:
150+
if self.llm in CLAUDE_MODELS:
121151
llm_kwargs["thinking_budget"] = self.reasoning_effort
122152
else:
123153
llm_kwargs["max_tokens"] = self.reasoning_effort
124154

125-
elif self.llm in [
126-
"o1",
127-
"o1-preview",
128-
"o3-mini",
129-
"o4-mini",
130-
"o3",
131-
"gpt-5",
132-
"gpt-5-mini",
133-
"gpt-5-nano",
134-
]:
155+
elif self.llm in OPENAI_MODELS:
135156
llm_kwargs["reasoning_effort"] = self.reasoning_effort
136157

137-
if self.llm in [
138-
"o1",
139-
"o1-mini",
140-
"o1-preview",
141-
"o3-mini",
142-
"o4-mini",
143-
"o3",
144-
"claude-3.7-sonnet",
145-
"gpt-5",
146-
"gpt-5-mini",
147-
"gpt-5-nano",
148-
]:
158+
elif self.llm in CLAUDE_MODELS:
159+
llm_kwargs["thinking_effort"] = self.reasoning_effort
160+
161+
if self.llm in OPENAI_MODELS + CLAUDE_MODELS:
149162
# For these models, we cannot set the temperature.
150163
llm_kwargs.pop("temperature")
151164

152165
if self.llm in ["o3-mini"]:
153166
llm_kwargs.pop("stream")
154167

155-
if self.llm in ["claude-3.7-sonnet"]:
168+
if self.llm in CLAUDE_MODELS:
156169
llm_kwargs["thinking"] = 1
157170
llm_kwargs.pop("seed")
158171

@@ -238,7 +251,7 @@ def act(self, obs, reward, done, infos):
238251
# Extract the action part from the response.
239252
action = action[reasoning_end:].strip()
240253

241-
elif self.llm in ["claude-3.7-sonnet"]:
254+
elif self.llm in CLAUDE_MODELS:
242255
# Extract the thinking part from the response JSON.
243256
thinking = "".join(
244257
[item.get("thinking", "") for item in response.json()["content"]]
@@ -253,24 +266,14 @@ def act(self, obs, reward, done, infos):
253266
"response": response_text,
254267
}
255268

256-
if self.llm in ["gemini-2.5-pro-preview-03-25", "gemini-2.5-pro-preview-05-06"]:
269+
if self.llm in GEMINI_MODELS:
257270
stats["nb_tokens_prompt"] = response.usage().input
258271
stats["nb_tokens_thinking"] = response.usage().details.get(
259272
"thoughtsTokenCount", 0
260273
)
261274
stats["nb_tokens_response"] = response.usage().output
262275

263-
elif self.llm in [
264-
"o1",
265-
"o1-mini",
266-
"o1-preview",
267-
"o3-mini",
268-
"o4-mini",
269-
"o3",
270-
"gpt-5",
271-
"gpt-5-mini",
272-
"gpt-5-nano",
273-
]:
276+
elif self.llm in OPENAI_MODELS:
274277
# stats["nb_tokens_prompt"] = self.token_counter(messages=messages),
275278
# stats["nb_tokens_response"] = self.token_counter(text=response_text)
276279
stats["nb_tokens_prompt"] = response.usage().input
@@ -281,8 +284,17 @@ def act(self, obs, reward, done, infos):
281284
"completion_tokens_details"
282285
]["reasoning_tokens"]
283286

287+
elif self.llm in CLAUDE_MODELS:
288+
stats["nb_tokens_prompt"] = self.token_counter(messages=messages)
289+
stats["nb_tokens_response"] = self.token_counter(text=response_text)
290+
stats["nb_tokens_thinking"] = 0
291+
if thinking:
292+
stats["nb_tokens_thinking"] = (
293+
response.usage().output - self.token_counter(text=response_text)
294+
)
295+
284296
else:
285-
stats["nb_tokens_prompt"] = (self.token_counter(messages=messages),)
297+
stats["nb_tokens_prompt"] = self.token_counter(messages=messages)
286298
stats["nb_tokens_thinking"] = (
287299
self.token_counter(text=thinking) if thinking else 0
288300
)

benchmark.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -405,6 +405,7 @@ def _maybe_load_agent_module():
405405
if args.agent:
406406
print(f"Importing agent(s) from {args.agent}.")
407407
for agent_file in glob.glob(args.agent):
408+
print(f"Importing {agent_file}...")
408409
agent_dirname = os.path.dirname(agent_file)
409410
agent_filename, _ = os.path.splitext(os.path.basename(agent_file))
410411
if f"{agent_dirname}.{agent_filename}" in sys.modules:

tales/agent.py

Lines changed: 5 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -33,7 +33,7 @@ def params(self):
3333
AGENTS = {}
3434

3535

36-
def register(name: str, desc: str, klass: callable, add_arguments: callable) -> None:
36+
def register(name: str, desc: str, klass: type, add_arguments: callable) -> None:
3737
""" Register a new type of Agent.
3838
3939
Arguments:
@@ -60,7 +60,9 @@ def register(name: str, desc: str, klass: callable, add_arguments: callable) ->
6060
>>> klass=RandomAgent,
6161
>>> add_arguments=_add_arguments)
6262
"""
63-
if name in AGENTS:
64-
raise ValueError(f"Agent '{name}' already registered.")
63+
if name in AGENTS and str(klass) != str(AGENTS[name][1]):
64+
raise ValueError(
65+
f"Agent '{name}' already registered from {AGENTS.get(name)[1]}."
66+
)
6567

6668
AGENTS[name] = (desc, klass, add_arguments)

tales/token.py

Lines changed: 22 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -46,7 +46,15 @@ def __init__(self, model: str):
4646
self.model = model
4747
if self.model in tiktoken.model.MODEL_TO_ENCODING:
4848
self.tokenize = tiktoken.encoding_for_model(self.model).encode
49-
elif self.model in ("o4-mini", "o3", "gpt-5", "gpt-5-mini", "gpt-5-nano"):
49+
elif self.model in (
50+
"o4-mini",
51+
"o3",
52+
"gpt-5",
53+
"gpt-5-mini",
54+
"gpt-5-nano",
55+
"gpt-5.1",
56+
"gpt-5.2",
57+
):
5058
self.tokenize = tiktoken.encoding_for_model("o3-mini").encode
5159
elif self.model in ("gpt-4.1", "gpt-4.1-nano", "gpt-4.1-mini"):
5260
self.tokenize = tiktoken.encoding_for_model("gpt-4o").encode
@@ -99,19 +107,31 @@ def __call__(self, *, messages=None, text=None):
99107
system = messages[0]["content"]
100108
messages.pop(0)
101109

102-
return self.client.beta.messages.count_tokens(
110+
nb_tokens = self.client.messages.count_tokens(
103111
model=self.model,
104112
messages=messages,
105113
system=system,
106114
).input_tokens
107115

116+
if text is not None:
117+
# Remove the boilerplate tokens needed for the count_tokens(...).
118+
# i.e. [{"role": "assistant", "content": ""}]
119+
nb_tokens -= (
120+
16 - 3
121+
) # 16 tokens for the boilerplate, minus 3 for <|im_*|> tokens.
122+
123+
return nb_tokens
124+
108125

109126
class GeminiTokenCounter(TokenCounter):
110127

111128
def __init__(self, model: Model):
112129
from google import genai
113130

114131
self.model = model.model_id
132+
# Strip 'gemini/' prefix if present
133+
if self.model.startswith("gemini/"):
134+
self.model = self.model[len("gemini/") :]
115135
self.client = genai.Client(api_key=model.get_key())
116136

117137
def __call__(self, *, messages=None, text=None):

0 commit comments

Comments
 (0)