自定义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
相关产品推荐
相关产品推荐

