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

使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.14 09:10:59