code/02-prompting/few_shot.py

76 lines · 3.4 KB

Code and program output are shown exactly as they ran, so comments and printed output are in Chinese.

"""把用户留言分成四类。比较不给例子(零样本)和给 4 个例子(少样本)的效果。"""
import os
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")
LABELS = ["缺陷", "功能建议", "使用问题", "其他"]

# 20 条测试留言和人工标注的正确类别
TESTS = [
    ("用 AsyncClient 并发 100 个请求,程序直接卡死,CPU 占满", "缺陷"),
    ("能不能加一个像 requests 那样的 Session 重试适配器?", "功能建议"),
    ("怎么给单个请求设置不同的超时时间?", "使用问题"),
    ("你们的文档网站打不开了", "其他"),
    ("升级到 0.28 之后,proxies 参数报 TypeError", "缺陷"),
    ("希望 Response 对象能直接支持 .json() 返回 dataclass", "功能建议"),
    ("上传文件的时候怎么同时带上表单字段?", "使用问题"),
    ("谢谢作者,这个库太好用了", "其他"),
    ("HTTP/2 开启后,偶尔会收到空的响应体,但状态码是 200", "缺陷"),
    ("cookie 在重定向之后丢了,是我用法不对吗", "使用问题"),
    ("建议命令行工具支持把响应保存到文件", "功能建议"),
    ("有没有中文版文档?", "其他"),
    ("在 Windows 上用 verify=False 仍然报 SSL 错误", "缺陷"),
    ("如何查看实际发出去的请求头?", "使用问题"),
    ("能否提供同步和异步通用的中间件接口", "功能建议"),
    ("请问你们招人吗", "其他"),
    ("stream 模式下 iter_lines 会把中文按字节截断成乱码", "缺陷"),
    ("代理需要用户名密码,应该怎么写?", "使用问题"),
    ("希望能内置请求耗时统计", "功能建议"),
    ("我的 PR 什么时候能合并?", "其他"),
]

ZERO_SHOT = """把用户留言分成以下四类之一:缺陷、功能建议、使用问题、其他。
只输出类别名称。"""

FEW_SHOT = ZERO_SHOT + """

例子:
留言:调用 client.close() 之后再发请求没有报错,而是静默返回了旧的响应
类别:缺陷

留言:想要一个参数,能在请求失败时自动打印完整的请求和响应
类别:功能建议

留言:base_url 和 url 拼接的规则是什么?结尾的斜杠有没有影响
类别:使用问题

留言:这个项目和 aiohttp 比哪个更好
类别:其他"""


def classify(system, text):
    response = client.chat.completions.create(
        model=MODEL,
        messages=[{"role": "system", "content": system}, {"role": "user", "content": f"留言:{text}\n类别:"}],
        temperature=0,
        extra_body={"thinking": {"type": "disabled"}},
    )
    return response.choices[0].message.content.strip()


for name, system in [("零样本", ZERO_SHOT), ("少样本", FEW_SHOT)]:
    with ThreadPoolExecutor(10) as pool:
        outputs = list(pool.map(lambda t: classify(system, t[0]), TESTS))
    correct = sum(out == label for out, (_, label) in zip(outputs, TESTS))
    off_format = [out for out in outputs if out not in LABELS]
    print(f"{name}:答对 {correct}/20,格式不对 {len(off_format)} 条 {off_format}")
    for out, (text, label) in zip(outputs, TESTS):
        if out != label:
            print(f"    {text}  标注={label}  模型={out}")