首页 > 教程攻略 > ai资讯 >Langgraph实战--自定义embeding

Langgraph实战--自定义embeding

来源:互联网 时间:2026-07-22 15:02:16

在LangGraph中整合第三方平台的Embedding接口,其实是一个很常见的需求。市面上开箱即用的LangChain类,主要就是AzureOpenAIEmbeddings和OpenAIEmbeddings这两个,模型选择范围比较有限。如果你想把硅基流之类的平台接进来,就得自己动手封装一个类,让它继承LangChain的Embedding接口。下面就来拆解具体怎么做。

Langgraph实战--自定义embeding

概述

先交代一下背景。当我们想在LangGraph里使用第三方的Embedding接口时,会发现LangChain默认给的两个类——OpenAIEmbeddings和AzureOpenAIEmbeddings——都是针对对应平台的。想通过ChatOpenAI去调用硅基流的接口?行不通。唯一靠谱的方式是自己封装一个类,让它继承LangChain的Embeddings抽象类,然后在这个类里实现对接第三方Embedding API的逻辑。这才是让LangGraph支持任意Embedding平台的正确姿势。

实现思路

核心思路其实特别直白:继承langchain_core.embeddings里的Embeddings类,然后按需实现两个关键方法——embed_documentsembed_query。前者负责批量处理多个文档,后者负责对单条查询进行向量化。把第三方平台的API调用逻辑塞进这两个方法里,事情就成了。

来看具体的代码实现。这里以硅基流平台为例,写一个自定义的Embedding类:

import requests
import os
from typing import List
from langchain_core.embeddings import Embeddings
from dotenv import load_dotenv

class CustomSiliconFlowEmbeddings(Embeddings):
    def __init__(
        self,
        api_key: str,
        base_url: str = "https://api.siliconflow.cn/v1/embeddings",
        model: str = "BAAI/bge-large-zh-v1.5"
    ):
        self.api_key = api_key
        self.base_url = base_url
        self.model = model

    def embed_documents(self, texts: List[str]) -> List[List[float]]:
        """Embed a list of documents."""
        embeddings = []
        for text in texts:
            embedding = self.embed_query(text)
            embeddings.append(embedding)
        return embeddings

    def embed_query(self, text: str) -> List[float]:
        """Embed a query."""
        headers = {
            "Authorization": f"Bearer {self.api_key}",
            "Content-Type": "application/json"
        }
        
        payload = {
            "model": self.model,
            "input": text,
            "encoding_format": "float"
        }
        
        response = requests.post(
            self.base_url,
            json=payload,
            headers=headers
        )
        
        if response.status_code == 200:
            return response.json()["data"][0]["embedding"]
        else:
            raise Exception(f"Error in embedding: {response.text}")

注意看,embed_documents本质上是在循环调用embed_query来实现的,代码简洁实用。而embed_query负责构造HTTP请求,解析返回的JSON数据,拿到向量结果。容错逻辑也加了——如果请求失败,直接抛出异常,方便排查问题。

使用CustomSiliconFlowEmbeddings嵌入类

类写好了,怎么用?大致流程是这样:加载环境变量,初始化自定义的Embedding模型,然后就能像用LangChain原生类一样去调用了。参数包括API密钥、模型名称、接口地址,一个都不能少。

# Load environment variables
load_dotenv()
SL_API_KEY = os.getenv("SL_API_KEY")

# Initialize embedding model
embedding_model = CustomSiliconFlowEmbeddings(
    base_url="https://api.siliconflow.cn/v1/embeddings",
    api_key=SL_API_KEY,
    model="BAAI/bge-large-zh-v1.5"
)

# Test the embedding
if __name__ == "__main__":
    test_text = "您好世界!"
    result = embedding_model.embed_query(test_text)
    print(f"Embedding dimension: {len(result)}")
    print(f"First few values: {result[:10]}")

    # 获取网页中的数据,并进行分割,然后存储到FAISS中
    urls = [
        "https://lilianweng.github.io/posts/2023-06-23-Agent/",
        "https://lilianweng.github.io/posts/2023-03-15-prompt-engineering/",
        "https://lilianweng.github.io/posts/2023-10-25-adv-attack-llm/"
    ]

    docs = [WebBaseLoader(url).load() for url in urls]
    docs_list = [item for sublist in docs for item in sublist]
    text_splitter = RecursiveCharacterTextSplitter.from_tiktoken_encoder(chunk_size=250, chunk_overlap=0)
    doc_splits = text_splitter.split_documents(docs_list)

    vectorstore = FAISS.from_documents(documents=doc_splits, embedding=embedding_model)
    retriever = vectorstore.as_retriever()

    # 测试检索功能,查询与问题最相关的分块文档
    resp = retriever.invoke("什么是prompt engineering?")
    # 返回的是一个个Document对象
    for doc in resp:
        print(doc.id + ": " + doc.page_content)

这段代码演示了一个完整链路:先测试单个文本的向量化效果,然后抓取几个网页内容,做文本分块,再通过FAISS构建向量检索库。最后,用检索器去查询“什么是prompt engineering?”,能看到实际返回的相关文档块。

输出效果也很直观:

Embedding dimension: 1024
First few values: [0.021915348, 0.0048826355, -0.09566349, -0.010307786, -0.0025656442, 0.043084737, -0.045955546, 0.011641469, 0.02809776, -0.012489148]

从结果看,模型的向量维度是1024,前10个浮点数也显示正常。这说明自定义的Embedding类已经成功接入了硅基流平台,并且可以无缝嵌入LangChain的向量检索体系。整个过程下来,核心逻辑其实就是继承、实现、调用这三步,并不复杂。但有了这个能力,LangGraph在模型选择和部署灵活性上,就打开了更大的空间。