You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

自定义ChatModel中bind_tools函数的正确实现方案

自定义LangChain ChatModel工具调用失效的解决方法

你的核心问题是:bind_tools仅完成了工具参数的绑定,但自定义的_generate方法既没有读取这些参数,也没有按模型要求注入工具调用prompt,更没有解析模型输出的工具调用指令。以下是具体修复步骤:

1. 在_generate中读取并处理绑定的工具参数

修改CustomChatModel的_generate方法,取出绑定的tools和tool_choice参数,将工具信息转换成Llama-3能识别的格式插入对话prompt:

import json
from typing import List, Optional, Any
from langchain.schema import BaseMessage, SystemMessage, AIMessage, ChatResult, ChatGeneration

def _generate(
    self,
    messages: List[BaseMessage],
    stop: Optional[List[str]] = None,
    run_manager: Optional[CallbackManagerForLLMRun] = None,
    **kwargs: Any,
) -> ChatResult:
    # 取出绑定的工具配置
    tools = kwargs.get("tools", [])
    tool_choice = kwargs.get("tool_choice", "auto")
    
    # 构造工具调用说明的prompt片段
    tool_prompt = ""
    if tools:
        # 将OpenAI格式工具转为Llama-3可读的自然语言描述
        tool_desc = []
        for tool in tools:
            func = tool["function"]
            param_str = json.dumps(func["parameters"], ensure_ascii=False, indent=2)
            tool_desc.append(f"- **{func['name']}**: {func['description']}\n参数格式:\n{param_str}")
        
        # 定义模型调用工具的严格格式
        tool_prompt = f"""你可以使用以下工具获取实时信息:
{'\n'.join(tool_desc)}
### 工具调用规则
当需要调用工具时,必须严格按照以下格式输出:
<|tool_call_begin|>
{{"name": "工具名称", "parameters": {{"参数名": "参数值"}}}}
<|tool_call_end|>
如果不需要调用工具,直接回答用户问题即可。"""
        
        # 将工具prompt注入对话历史(优先合并到已有SystemMessage)
        updated_messages = []
        system_found = False
        for msg in messages:
            if isinstance(msg, SystemMessage):
                updated_messages.append(SystemMessage(content=f"{msg.content}\n{tool_prompt}"))
                system_found = True
            else:
                updated_messages.append(msg)
        if not system_found:
            updated_messages.insert(0, SystemMessage(content=tool_prompt))
        messages = updated_messages
    
    # 调用模型生成响应
    res = llama3_instruct(messages)
    
    # 解析模型输出中的工具调用指令
    tool_calls = []
    if "<|tool_call_begin|>" in res and "<|tool_call_end|>" in res:
        # 提取工具调用的JSON片段
        start_pos = res.index("<|tool_call_begin|>") + len("<|tool_call_begin|>")
        end_pos = res.index("<|tool_call_end|>")
        call_json = res[start_pos:end_pos].strip()
        try:
            call_data = json.loads(call_json)
            tool_calls.append({
                "name": call_data["name"],
                "args": call_data["parameters"]
            })
            # 工具调用场景下清空普通回答内容
            res = ""
        except json.JSONDecodeError:
            # 解析失败,视为普通回答
            pass
    
    # 构造包含工具调用的AIMessage
    message = AIMessage(
        content=res,
        tool_calls=tool_calls,  # 直接赋值tool_calls字段,而非additional_kwargs
        response_metadata={"time_in_seconds": 3}
    )

    generation = ChatGeneration(message=message)
    return ChatResult(generations=[generation])

2. 确保模型生成时保留工具调用标记

调整llama3_instruct函数的生成参数,避免工具调用标记被误判为终止符:

def llama3_instruct(messages):
    msgs = convert_messages(messages)
    input_ids = tokenizer.apply_chat_template(msgs, add_generation_prompt=True, return_tensors="pt").to(model.device)

    # 仅使用官方终止符,不要添加自定义工具标记
    terminators = [
        tokenizer.eos_token_id,
        tokenizer.convert_tokens_to_ids("<|eot_id|>")
    ]

    outputs = model.generate(
        input_ids,
        max_new_tokens=512,  # 适当增大生成长度,容纳工具调用JSON
        eos_token_id=terminators,
        do_sample=True,
        temperature=0.3,  # 降低温度,提升格式遵循度
        top_p=0.9,
    )
    response = outputs[0][input_ids.shape[-1]:]
    return tokenizer.decode(response, skip_special_tokens=False)  # 不要跳过特殊标记,方便解析

3. 简化bind_tools实现(可选)

无需完全照搬ChatOpenAI的复杂逻辑,BaseChatModel的bind方法已支持参数传递,可简化为:

from langchain.tools.convert_to_openai import convert_to_openai_tool
from typing import Sequence, Union, Dict, Type, Callable, BaseModel, Literal

def bind_tools(
    self,
    tools: Sequence[Union[Dict[str, Any], Type[BaseModel], Callable, BaseTool]],
    *,
    tool_choice: Optional[Union[dict, str, Literal["auto", "any", "none"], bool]] = None,
    **kwargs: Any,
) -> Runnable[LanguageModelInput, BaseMessage]:
    formatted_tools = [convert_to_openai_tool(tool) for tool in tools]
    # 处理强制工具调用逻辑
    if tool_choice is not None and tool_choice != "none":
        if isinstance(tool_choice, bool) and tool_choice:
            tool_choice = {"type": "function", "function": {"name": formatted_tools[0]["function"]["name"]}}
        elif isinstance(tool_choice, str) and tool_choice not in ("auto", "any"):
            tool_choice = {"type": "function", "function": {"name": tool_choice}}
        kwargs["tool_choice"] = tool_choice
    return super().bind(tools=formatted_tools, **kwargs)

4. 验证效果

调整后调用模型,当需要实时信息时,模型会输出工具调用指令,response.tool_calls将包含正确的调用参数:

ContentString: 
ToolCalls: [{'name': 'tavily_search', 'args': {'query': '旧金山今日天气'}}]

内容的提问来源于stack exchange,提问作者AVJ

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.20 02:24:52