code/05-agents/tool_design.py
80 lignes · 3.8 KoLe code et les sorties des programmes sont reproduits tels qu’ils ont tourné : commentaires et sorties sont donc en chinois.
"""同样三个工具,一组名字和说明写得含糊,一组写得清楚。看模型第一步选的工具对不对。
在 AI-Course/code/05-agents 目录下运行:python tool_design.py
只看模型第一步想调用哪个工具,不真的执行。10 个问题 × 2 组 × 3 次,共 60 次调用。
"""
import os
from collections import Counter
from concurrent.futures import ThreadPoolExecutor
from openai import OpenAI
client = OpenAI(api_key=os.environ["LLM_API_KEY"], base_url=os.environ.get("LLM_BASE_URL", "https://api.deepseek.com"))
MODEL = os.environ.get("LLM_MODEL", "deepseek-flash")
def fn(tool_name, description, **params):
return {"type": "function", "function": {"name": tool_name, "description": description, "parameters": {
"type": "object", "properties": {k: {"type": "string", "description": v} for k, v in params.items()},
"required": list(params)}}}
# 含糊的一组:名字是通用的动词,说明只有两三个字
VAGUE = [
fn("search", "搜索", q="内容"),
fn("read", "读取", x="要读的东西"),
fn("lookup", "查找信息", name="名字"),
]
# 清楚的一组:名字说明对象,说明写清楚能做什么、什么时候用、参数怎么填
CLEAR = [
fn("search_docs", "在 httpx 官方文档里按英文关键词全文搜索,返回匹配的文件名和行号。"
"用户问 httpx 某个功能怎么用、某个参数是什么意思时,先用它。",
keyword="英文关键词,例如 timeout、proxy、follow_redirects"),
fn("read_doc", "读取 httpx 文档里某个文件的内容。通常在 search_docs 找到文件名之后使用。",
path="文档文件路径,例如 advanced/timeouts.md"),
fn("get_pypi_info", "查询某个 Python 包在 PyPI 上的最新版本号和发布信息。只在用户问版本号、是否已发布新版本时使用。",
package="PyPI 上的包名,例如 httpx"),
]
# 两组工具一一对应:含糊组的第 i 个和清楚组的第 i 个是同一个功能
SAME = {"search": "search_docs", "read": "read_doc", "lookup": "get_pypi_info"}
QUESTIONS = [
("httpx 怎么设置代理?", "search_docs"),
("httpx 最新版本是多少?", "get_pypi_info"),
("帮我看看 advanced/ssl.md 里写了什么", "read_doc"),
("follow_redirects 参数是干什么的?", "search_docs"),
("requests 现在出到哪个版本了?", "get_pypi_info"),
("httpx 怎么上传文件?", "search_docs"),
("你好", None),
("把 quickstart.md 的内容给我看一下", "read_doc"),
("httpx 有没有发布 1.0 正式版?", "get_pypi_info"),
("httpx 的 event hooks 怎么用?", "search_docs"),
]
SYSTEM = "你是 httpx 的答疑助手。需要查资料时使用工具。"
def first_tool(tools, question):
response = client.chat.completions.create(
model=MODEL,
messages=[{"role": "system", "content": SYSTEM}, {"role": "user", "content": question}],
tools=tools,
extra_body={"thinking": {"type": "disabled"}},
)
calls = response.choices[0].message.tool_calls
if not calls:
return None
name = calls[0].function.name
return SAME.get(name, name) # 含糊组的工具名换成对应的清楚组名字,方便比较
for label, tools in [("含糊的工具", VAGUE), ("清楚的工具", CLEAR)]:
jobs = [(q, want) for q, want in QUESTIONS for _ in range(3)]
with ThreadPoolExecutor(10) as pool:
got = list(pool.map(lambda j: first_tool(tools, j[0]), jobs))
correct = sum(g == want for g, (_, want) in zip(got, jobs))
print(f"{label}:{correct}/{len(jobs)} 次选对")
for i, (question, want) in enumerate(QUESTIONS):
answers = got[i * 3:(i + 1) * 3]
if any(a != want for a in answers):
print(f" {question} 应该用 {want},实际 {dict(Counter(map(str, answers)))}")