LangChain中如何通过回调从自定义LLM提取响应数据?
LangChain自定义LLM回调收集响应信息方案(含Watsonx示例)
问题核心原因
你遇到的on_llm_end传入response为None的问题,本质是第三方实现的LLM类(如WatsonxLLM)未在生成响应时正确构建LLMResult对象,也未将完整响应数据传递给回调管理器。LangChain的回调机制依赖LLM类主动触发on_llm_end并传入有效LLMResult,否则回调无法获取任何响应信息。
回调方案修复思路(无需修改原LLM类)
由于无法直接修改WatsonxLLM的源码,可通过**包装类(Wrapper)**继承原LLM类,重写生成逻辑,手动构建包含stop_reason的LLMResult,再传递给回调。
步骤1:编写包装类,补充LLMResult构建逻辑
from typing import List, Optional, Any from langchain.llms.base import LLMResult, CallbackManagerForLLMRun from ibm_watson_machine_learning.foundation_models.extensions.langchain import WatsonxLLM class WrappedWatsonxLLM(WatsonxLLM): def _generate( self, prompts: List[str], stop: Optional[List[str]] = None, run_manager: Optional[CallbackManagerForLLMRun] = None, **kwargs: Any, ) -> LLMResult: # 调用父类方法获取原始生成结果 raw_result = super()._generate(prompts, stop, run_manager, **kwargs) # 从Watsonx底层响应提取stop_reason(需根据实际API返回结构调整) # 示例:假设通过客户端调用的原始响应中包含stop_reason字段 stop_reasons = [] for prompt in prompts: # 这里替换为实际提取逻辑,比如通过self.client直接调用API获取完整响应 raw_response = self.client.generate( prompt=prompt, parameters=self._get_parameters(stop=stop, **kwargs) ) stop_reason = raw_response.get("results", [{}])[0].get("stop_reason") stop_reasons.append(stop_reason) # 构建包含stop_reason的LLMResult updated_llm_output = raw_result.llm_output or {} updated_llm_output["stop_reasons"] = stop_reasons return LLMResult( generations=raw_result.generations, llm_output=updated_llm_output, run=raw_result.run )
步骤2:调整回调类,提取stop_reason
from langchain.callbacks.base import BaseCallbackHandler, LLMResult from typing import Any, Optional from contextvars import ContextVar, Generator class _WatsonXCallbackHandler(BaseCallbackHandler): def on_llm_end(self, response: LLMResult, **kwargs: Any) -> None: """Run when LLM ends running.""" if not response.llm_output: return # 提取并处理stop_reason stop_reasons = response.llm_output.get("stop_reasons", []) for idx, reason in enumerate(stop_reasons): print(f"Prompt {idx+1} 的stop_reason: {reason}") def __copy__(self) -> "_WatsonXCallbackHandler": return self def __deepcopy__(self, memo: Any) -> "_WatsonXCallbackHandler": return self watsonx_callback_var: ContextVar[Optional[_WatsonXCallbackHandler]] = ContextVar( "watsonx_callback", default=None ) class MyServiceIBMWatsonx(MyService): @contextmanager def get_langchain_callback(self) -> Generator[_WatsonXCallbackHandler, None, None]: cb = _WatsonXCallbackHandler() watsonx_callback_var.set(cb) yield cb watsonx_callback_var.set(None)
步骤3:使用包装后的LLM类
# 初始化包装后的LLM llm = WrappedWatsonxLLM( model_id="granite-13b-instruct-v2", credentials={"apikey": "your_api_key", "url": "your_url"}, params={"max_new_tokens": 200} ) # 结合上下文管理器使用回调 with MyServiceIBMWatsonx().get_langchain_callback() as cb: llm.generate(["你的测试Prompt"])
无需回调的替代方案
如果不想用回调机制,可直接从响应中提取stop_reason,有以下几种方式:
1. 直接解析LLMResult的llm_output
若原LLM类已部分填充llm_output,可直接读取:
llm = WatsonxLLM(...) result = llm.generate(["你的Prompt"]) if result.llm_output: stop_reason = result.llm_output.get("stop_reason") print(f"stop_reason: {stop_reason}")
2. 直接调用Watsonx底层API
绕开LangChain封装,直接使用IBM官方客户端调用模型,获取完整原始响应:
from ibm_watson_machine_learning.foundation_models import Model model = Model( model_id="granite-13b-instruct-v2", credentials={"apikey": "your_api_key", "url": "your_url"}, params={"max_new_tokens": 200} ) raw_response = model.generate_text(prompt="你的Prompt", return_full_response=True) stop_reason = raw_response.get("results", [{}])[0].get("stop_reason") print(f"stop_reason: {stop_reason}")
3. Monkey Patch临时修改原LLM类
通过动态修改原类的_generate方法,补充stop_reason到LLMResult中(适合快速验证,不推荐生产环境):
from ibm_watson_machine_learning.foundation_models.extensions.langchain import WatsonxLLM original_generate = WatsonxLLM._generate def patched_generate(self, prompts, stop=None, run_manager=None, **kwargs): result = original_generate(self, prompts, stop, run_manager, **kwargs) # 提取stop_reason并补充到llm_output stop_reasons = [] for prompt in prompts: raw_response = self.client.generate(prompt=prompt, parameters=self._get_parameters(stop=stop, **kwargs)) stop_reasons.append(raw_response.get("results", [{}])[0].get("stop_reason")) result.llm_output = result.llm_output or {} result.llm_output["stop_reasons"] = stop_reasons return result WatsonxLLM._generate = patched_generate
内容的提问来源于stack exchange,提问作者Sjoerd222888
相关产品推荐
相关产品推荐

