使用databricks_langchain时无法通过MLflow记录LangChain模型的问题
解决MLflow记录LangChain链时的
langchain_databricks模块错误 问题原因
MLflow的LangChain集成目前默认仅识别旧的langchain_databricks包中的组件,而你使用的是替代它的databricks_langchain包,导致序列化时MLflow无法找到对应模块路径,抛出No module named 'langchain_databricks'错误。
解决方案
方案1:注册自定义组件序列化映射
手动告诉MLflow如何序列化databricks_langchain中的ChatDatabricks组件,修改代码如下:
from databricks_langchain import DatabricksVectorSearch, ChatDatabricks from langchain.prompts import PromptTemplate from langchain.schema.runnable import RunnableMap, RunnableLambda from langchain.schema.output_parser import StrOutputParser from operator import itemgetter import mlflow from mlflow.langchain import _runnables # 关键:注册ChatDatabricks的序列化信息 _runnables._SERIALIZABLE_RUNNABLES[ChatDatabricks] = { "module": "databricks_langchain", "type": "ChatDatabricks", } vs_endpoint = "your_vector_search_endpoint" my_index_name = "your_index_name" def retriever_loader(): my_index = DatabricksVectorSearch( endpoint=vs_endpoint, index_name=my_index_name, columns=["ID", "TEXT"] ) return my_index.as_retriever(search_kwargs={"k": 3, "query_type": "HYBRID"}) my_retriever = retriever_loader() prompt = PromptTemplate.from_template( template="""Some template: {query} and {context} """ ) def format_context(text): return modified(text) # 确保modified函数已定义 llm_endpoint = ChatDatabricks(endpoint="databricks-meta-llama-3-3-70b-instruct") chain = ( RunnableMap({ "query": RunnableLambda(itemgetter("messages")), "context": RunnableLambda(itemgetter("messages")) | my_retriever | RunnableLambda(format_context), }) | prompt | llm_endpoint | StrOutputParser() ) model_name = "some_model_name" input_example = {"messages": "Your example query here"} resp = chain.invoke(input_example) with mlflow.start_run(run_name="run_name") as run: model_info = mlflow.langchain.log_model( chain, loader_fn=retriever_loader, artifact_path="path_to_artifact", registered_model_name=model_name, input_example=input_example )
方案2:使用MLflow PyFunc包装链
如果方案1无效,改用更灵活的PyFunc方式绕过LangChain专属序列化限制:
from databricks_langchain import DatabricksVectorSearch, ChatDatabricks from langchain.prompts import PromptTemplate from langchain.schema.runnable import RunnableMap, RunnableLambda from langchain.schema.output_parser import StrOutputParser from operator import itemgetter import mlflow import pandas as pd class LangChainPyFunc(mlflow.pyfunc.PythonModel): def __init__(self, chain): self.chain = chain def predict(self, context, model_input): # 适配PyFunc的输入格式,转换为链需要的结构 input_data = model_input.to_dict(orient="records")[0] return self.chain.invoke(input_data) vs_endpoint = "your_vector_search_endpoint" my_index_name = "your_index_name" def retriever_loader(): my_index = DatabricksVectorSearch( endpoint=vs_endpoint, index_name=my_index_name, columns=["ID", "TEXT"] ) return my_index.as_retriever(search_kwargs={"k": 3, "query_type": "HYBRID"}) my_retriever = retriever_loader() prompt = PromptTemplate.from_template( template="""Some template: {query} and {context} """ ) def format_context(text): return modified(text) # 确保modified函数已定义 llm_endpoint = ChatDatabricks(endpoint="databricks-meta-llama-3-3-70b-instruct") chain = ( RunnableMap({ "query": RunnableLambda(itemgetter("messages")), "context": RunnableLambda(itemgetter("messages")) | my_retriever | RunnableLambda(format_context), }) | prompt | llm_endpoint | StrOutputParser() ) model_name = "some_model_name" input_example = {"messages": "Your example query here"} resp = chain.invoke(input_example) with mlflow.start_run(run_name="run_name") as run: mlflow.pyfunc.log_model( artifact_path="path_to_artifact", python_model=LangChainPyFunc(chain), registered_model_name=model_name, input_example=input_example, # 显式声明依赖环境,避免加载时缺失包 conda_env={ "channels": ["conda-forge"], "dependencies": [ "python=3.10", {"pip": [ "mlflow==2.20.2", "langchain_core==0.3.35", "databricks_langchain==0.3.0", "pandas" ]} ] } )
方案3:升级MLflow到最新版本
检查MLflow官方更新,部分新版本已经适配了databricks_langchain包。尝试升级MLflow:
pip install --upgrade mlflow
升级后直接运行原代码,若新版本已支持则问题自动解决。
内容的提问来源于stack exchange,提问作者skad00sh
相关产品推荐
相关产品推荐

