使用 LlamaIndex 接入 Vertex AI Endpoint 自定义 Embedding 模型:VertexEndpointEmbedding 全解析
【免费下载链接】llama_indexLlamaIndex is the leading document agent and OCR platform项目地址: https://gitcode.com/GitHub_Trending/ll/llama_index
导读
VertexEndpointEmbedding 是 LlamaIndex 为 Google Cloud Vertex AI Endpoint 提供的 Embedding 集成类,用于把部署在 Vertex AI 上的自定义/私有化文本嵌入模型无缝接入 LlamaIndex 的检索与索引体系。阅读本文你将掌握:如何安装与初始化该组件、每个构造参数的含义与默认值、同步与异步两种推理路径的实现原理、输入输出如何通过IOHandler进行序列化适配,以及它如何被纳入 LlamaIndex 文档级源码与测试验证的完整链路。
一、为什么需要VertexEndpointEmbedding?
Google Cloud 的 Vertex AI Endpoint 允许用户部署经过微调或自定义训练的模型,并通过统一的predict接口对外提供推理服务。对于不在 LlamaIndex 官方模型列表中的私有 embedding 模型,用户需要一种方式让 LlamaIndex 的VectorStoreIndex、检索器等上层能力直接消费该端点的输出向量。
VertexEndpointEmbedding的作用正是建立这条桥接通道:它在内部持有google.cloud.aiplatform.Endpoint客户端,把 LlamaIndex 的BaseEmbedding抽象请求翻译成对 Vertex AI Endpoint 的预测调用,再返回标准化的向量结果。从测试可见其设计完全遵循 LlamaIndex 的嵌入抽象——test_embeddings_vertex_endpoint.py 通过遍历VertexEndpointEmbedding.__mro__断言BaseEmbedding是其基类之一,从继承关系上确认了它属于标准的 LlamaIndex Embedding 体系。
二、安装与包结构
该集成作为独立包发布,安装命令如下:
pip install llama-index-embeddings-vertex-endpoint根据 pyproject.toml 中的依赖声明,它依赖两个关键运行库:
| 依赖 | 版本约束 | 用途 |
|---|---|---|
google-cloud-aiplatform | >=1.69.0,<2 | 创建aiplatform.Endpoint客户端并发送预测请求 |
llama-index-core | >=0.13.0,<0.15 | 提供BaseEmbedding、CallbackManager等核心抽象 |
包本身要求python >= 3.10,<4.0,采用 MIT 许可证。包的导入路径在 pyproject.toml 中注册为llama_index.embeddings.vertex_endpoint,顶层的init.py 只导出一个公开符号:
from llama_index.embeddings.vertex_endpoint.base import VertexEndpointEmbedding __all__ = ["VertexEndpointEmbedding"]对应的文档页 vertex_endpoint.md 也正是以VertexEndpointEmbedding为唯一公开成员生成 API 参考,因此使用时代码写作:
from llama_index.embeddings.vertex_endpoint import VertexEndpointEmbedding三、构造参数详解:字段与默认值
VertexEndpointEmbedding的核心入口是__init__构造函数(见 base.py)。所有关键参数同时以 PydanticField形式声明为类属性(base.py),既保证参数可校验、可序列化,也方便后续作为 LlamaIndex 组件被保存与恢复。
3.1 必填参数
| 参数 | 类型 | 说明 |
|---|---|---|
endpoint_id | str | Vertex AI Endpoint 的 ID,用于定位已部署的模型端点 |
project_id | str | 承载该 Endpoint 的 GCP 项目 ID |
location | str | Vertex AI 所在 GCP 区域(Region),如us-central1 |
这三个值在__init__中被直接用于构造客户端:
self._client = aiplatform.Endpoint( endpoint_name=endpoint_id, project=project_id, location=location, credentials=credentials, )3.2 请求级参数
| 参数 | 类型 | 默认值 | 说明 |
|---|---|---|---|
endpoint_kwargs | Dict[str, Any] | {} | 传给predict请求的附加关键字参数 |
model_kwargs | Dict[str, Any] | {} | 传给模型的参数(映射为 Vertex AI 的parameters) |
timeout | float | 60.0 | API 请求超时时间(秒),Pydantic 约束ge=0,即不可为负 |
embed_batch_size | int | DEFAULT_EMBED_BATCH_SIZE | 每次批量送入模型的文本条数,来自 llama_index.core.constants 体系的默认批大小常量 |
3.3 凭据与高级参数
| 参数 | 类型 | 默认值 | 说明 |
|---|---|---|---|
service_account_file | str \| None | None | 服务账号 JSON 文件路径 |
service_account_info | Dict[str, str] \| None | None | 直接以字典形式提供服务账号凭据内容 |
content_handler | BaseIOHandler | IOHandler() | 负责输入序列化与输出反序列化的处理器 |
callback_manager | CallbackManager \| None | None | LlamaIndex 回调管理器,接入可观测体系 |
verbose | bool | False | 是否开启调试输出,保存在私有属性_verbose中 |
3.4 凭据解析优先级
构造函数的凭据处理逻辑清晰(base.py):
- 若传入了
service_account_file,调用service_account.Credentials.from_service_account_file(file)从 JSON 文件加载凭据; - 否则若传入了
service_account_info,调用service_account.Credentials.from_service_account_info(info)直接从字典加载; - 两者都未提供时
credentials = None,回退使用 Google 默认应用凭据(Application Default Credentials, ADC),例如通过gcloud auth application-default login或环境变量注入的身份。
客户端创建若失败,异常会被包装为ValueError("Please verify the provided credentials.")抛出——这说明该项目的大多数连接问题根源都在凭据环节,排查时应优先检查以上三种凭据来源是否有效。
四、IOHandler:输入输出适配的关键
不同模型端点对请求体和返回体的格式约定并不一致。为此,源码在 utils.py 中定义了可插拔的 I/O 处理器抽象:
4.1 抽象基类BaseIOHandler
它通过abc.ABCMeta定义了协议,同时用__subclasshook__支持结构化鸭子类型——任何同时实现了serialize_input与deserialize_output两个可调用方法的类都会被视为其子类(不必显式继承):
| 抽象方法 | 签名 | 职责 |
|---|---|---|
serialize_input | (request: List[str]) -> ... | 把待嵌入的文本列表转换为端点期望的请求实例结构 |
deserialize_output | (response: Any) -> List[List[float]] | 从端点返回的预测结果中抽取向量列表 |
4.2 默认实现IOHandler
默认处理器实现的输入/输出约定为:
def serialize_input(self, request: List[str]) -> List[Dict[str, Any]]: return [{"inputs": text} for text in request] def deserialize_output(self, response: Any) -> List[List[float]]: return [prediction[0] for prediction in response.predictions]即:每个文本被包装成{"inputs": <text>}一条实例发送;返回时逐条读取response.predictions,并取出每条 prediction 的第 0 个元素作为该文本的嵌入向量。若你的自建模型使用不同的字段名(例如text代替inputs)或输出结构不同,就应当自定义一个BaseIOHandler子类并通过content_handler传入。模块级默认实例在 base.py 中创建:
DEFAULT_IO_HANDLER = IOHandler()五、核心推理链路:同步与异步
_get_embedding方法(base.py)是所有调用的汇聚点:
def _get_embedding(self, payload: List[str], **kwargs: Any) -> List[Embedding]: # 合并 endpoint 级参数:始终附加超时时间 endpoint_kwargs = {**self.endpoint_kwargs, **{"timeout": self.timeout}} # 合并模型参数:传入的方法 kwargs 拥有更高优先级 model_kwargs = {**self.model_kwargs, **kwargs} response = self._client.predict( instances=self.content_handler.serialize_input(payload), parameters=model_kwargs, **endpoint_kwargs, ) return self.content_handler.deserialize_output(response)可以拆解出三层参数合并规则:
- 端点参数:
endpoint_kwargs与自动注入的timeout合并,作为predict的关键字参数; - 模型参数:
model_kwargs与每次调用额外传入的**kwargs合并,作为 Vertex AI 的parameters,调用级参数可覆盖构造时的默认模型参数; - 实例数据:
payload(文本列表)先经content_handler.serialize_input转换,再作为instances传入。
对应的异步版本_aget_embedding(base.py)逻辑完全一致,仅把_client.predict替换为_client.predict_async并await其返回。这一对方法由BaseEmbedding基类的同步/异步公共 API 分派调用。
5.1 四种嵌入方法及文本预处理
类内实现了BaseEmbedding要求的六个方法。其中四个查询/文本单条方法以及两条批量路径(base.py)在调用底层推理前统一执行了一条预处理规则:
text = text.replace("\n", " ")将文本中的换行符替换为空格,降低换行对嵌入质量与批次格式的干扰:
| 方法 | 签名 | 用途 |
|---|---|---|
_get_query_embedding | (query: str) -> Embedding | 生成查询向量(同步) |
_get_text_embedding | (text: str) -> Embedding | 生成单条文档向量(同步) |
_get_text_embeddings | (texts: List[str]) -> List[Embedding] | 批量生成文档向量(同步) |
_aget_query_embedding | (query: str) -> Embedding | 查询向量(异步) |
_aget_text_embedding | (text: str) -> Embedding | 单条文档向量(异步) |
_aget_text_embeddings | (texts: List[str]) -> List[Embedding] | 批量文档向量(异步) |
此外,class_name()类方法(base.py)返回"VertexEndpointEmbedding",这是 LlamaIndex 系列组件在序列化、日志与类型识别上的约定接口。
六、完整使用示例
6.1 基础用法(走默认凭据)
from llama_index.embeddings.vertex_endpoint import VertexEndpointEmbedding from llama_index.core import VectorStoreIndex # 前提:已在对应 GCP 项目的 us-central1 部署好 embedding 端点, # 且当前环境具备默认应用凭据(ADC)。 embed_model = VertexEndpointEmbedding( endpoint_id="projects/123456789012/locations/us-central1/endpoints/987654321", project_id="my-gcp-project", location="us-central1", )6.2 指定服务账号文件
embed_model = VertexEndpointEmbedding( endpoint_id="my-endpoint-id", project_id="my-gcp-project", location="us-central1", service_account_file="/path/to/service-account.json", timeout=120.0, # 覆盖默认 60 秒超时 embed_batch_size=32, verbose=True, )初始化完成后,该对象即可作为 LlamaIndex 全体系的标准 embedding 使用:
from llama_index.core import Settings Settings.embed_model = embed_model # 全局生效 index = VectorStoreIndex.from_documents(documents) # 索引阶段自动调用批量嵌入 retriever = index.as_retriever(similarity_top_k=5)6.3 针对自定义端点格式定制 IOHandler
若模型端点的输入字段是"text"而非"inputs",自定义处理器并传入即可:
from llama_index.embeddings.vertex_endpoint import VertexEndpointEmbedding from llama_index.embeddings.vertex_endpoint.utils import BaseIOHandler class MyHandler(BaseIOHandler): def serialize_input(self, request): return [{"text": text} for text in request] def deserialize_output(self, response): return [prediction["embedding"] for prediction in response.predictions] embed_model = VertexEndpointEmbedding( endpoint_id=..., project_id=..., location=..., content_handler=MyHandler(), )得益于__subclasshook__,MyHandler无需显式继承BaseIOHandler也会被接受,只要它实现了上述两个方法。
七、源码质量与集成验证
该集成虽小,但保持了与 LlamaIndex 核心一致的工程质量:
- 类型标注完整:全部方法均带类型注解,
mypy配置disallow_untyped_defs = true(pyproject.toml); - 测试回归:tests/test_embeddings_vertex_endpoint.py 验证类层次符合 LlamaIndex
BaseEmbedding抽象契约,防止未来核心 API 变更导致集成失效; - 文档集成:API 参考页 vertex_endpoint.md 通过 mkdocstrings 指令自动渲染
VertexEndpointEmbedding的字段与方法签名,保证文档与源码同步演进; - 编排规范:在 LlamaIndex 的集成目录体系中位于
llama-index-integrations/embeddings/llama-index-embeddings-vertex-endpoint/,遵循统一的包命名、Makefile与uv.lock工作流,便于按官方方式构建与发布。
八、注意事项与限制
- 端点必须已部署且可访问:
VertexEndpointEmbedding只负责调用,不负责模型部署;使用前需在 Vertex AI 控制台或通过aiplatform完成模型的线上部署并拿到 endpoint。 - 凭据错误提示较宽泛:客户端创建失败统一报
Please verify the provided credentials.,无法区分网络、项目或端点名错误,需要借助verbose=True与 GCP 日志进一步排查。 - 换行符会被替换:文本中的
\n在进入模型前统一替换为空格,对依赖原始换行的场景(如代码嵌入)需自行评估影响。 - 返回结构依赖端点实现:默认
IOHandler假定predictions每个元素的第一项即向量;若端点返回多输出结构,务必自定义deserialize_output。 - 版本前提:以当前仓库为准,该包面向
llama-index-core 0.13~0.15与google-cloud-aiplatform 1.x设计,升级大版本前请核对兼容性约束。
【免费下载链接】llama_indexLlamaIndex is the leading document agent and OCR platform项目地址: https://gitcode.com/GitHub_Trending/ll/llama_index
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考