mistral.rs 批量嵌入(Batch Embeddings)实战:用 EmbeddingRequestBuilder 并行编码海量文本向量
【免费下载链接】mistral.rsFast, flexible LLM inference项目地址: https://gitcode.com/GitHub_Trending/mi/mistral.rs
本篇技术指南围绕 mistral.rs 的batching_embeddings示例展开,讲解如何在 Rust SDK 中通过EmbeddingModelBuilder加载嵌入模型,并利用EmbeddingRequest::builder().add_prompts()将多条嵌入请求合并为一次批量调用,实现高效的并行编码;读完后你将掌握批量嵌入请求的完整写法、EmbeddingRequestBuilder的完整 API(含预分词输入与截断控制)、结果顺序保证机制,以及max_num_seqs并发上限的默认行为与源码位置。
示例场景与运行方式
batching_embeddings是 mistral.rs 官方 Rust 示例之一(examples/advanced分组),用于演示如何将多条嵌入请求打包为一次批量调用,让引擎并行编码而非逐条串行执行。官方文档页见 batching-embeddings.md,源码位于 mistralrs/examples/advanced/batching_embeddings/main.rs。
运行命令:
cargo run --release --example batching_embeddings -p mistralrs示例使用google/embeddinggemma-300m这一小型嵌入模型,分三步完成:先对两条查询各发起一次单条请求作为基准,再一次性批量编码 100 条交替重复的查询,最后逐条断言批量结果与单条请求结果逐位相等——这既验证了批量编码的正确性,也体现了嵌入推理的确定性。
完整示例代码解析
//! Batch multiple embedding requests for efficient parallel encoding. //! //! Run with: `cargo run --release --example batching_embeddings -p mistralrs` use anyhow::Result; use mistralrs::{EmbeddingModelBuilder, EmbeddingRequest}; #[tokio::main] async fn main() -> Result<()> { let model = EmbeddingModelBuilder::new("google/embeddinggemma-300m") .with_logging() .build() .await?; let a = model .generate_embeddings( EmbeddingRequest::builder() .add_prompt("task: search result | query: What is graphene?"), ) .await?; let b = model .generate_embeddings(EmbeddingRequest::builder().add_prompt( "task: search result | query: What is an apple's significance to gravity?", )) .await?; let batched = model .generate_embeddings(EmbeddingRequest::builder().add_prompts((0..100).map(|i| { if i % 2 == 0 { "task: search result | query: What is graphene?" } else { "task: search result | query: What is an apple's significance to gravity?" } }))) .await?; for (i, embedding) in batched.into_iter().enumerate() { if i % 2 == 0 { assert_eq!(embedding, a[0]); } else { assert_eq!(embedding, b[0]); } } Ok(()) }逐段说明:
- 模型加载:
EmbeddingModelBuilder::new("google/embeddinggemma-300m")指定 Hugging Face 模型 ID;.with_logging()开启日志;.build()异步加载并返回可直接推理的Model。 - 基准请求
a/b:每次generate_embeddings接收一个EmbeddingRequestBuilder,通过.add_prompt(...)加入单条文本,返回Vec<Vec<f32>>——外层按输入顺序对应,因此单条请求的结果取a[0]/b[0]。 - 批量请求:
.add_prompts(...)接受任意IntoIterator<Item = Into<String>>,这里用(0..100).map(...)生成 100 条交替文本,一次调用完成全部编码。 - 正确性校验:批量输出的第
i条与单条结果做浮点向量的逐位相等断言,证明批量与单条编码结果完全一致。
EmbeddingModelBuilder:加载与运行参数
构建器定义于 mistralrs/src/embedding_model.rs,new()的文档注释与字段初始化明确了以下默认行为:
| 配置项 | 默认值 | 说明 |
|---|---|---|
max_num_seqs | 32 | 同一时刻允许并行执行的最大序列数,决定了批量编码的真实并行度;可用with_max_num_seqs(n)调整 |
| Token 来源 | TokenSource::CacheToken | 从~/.cache/huggingface/token读取 Hugging Face 访问令牌;可用with_token_source覆盖 |
| 设备映射 | 自动 | 依据AutoDeviceMapParams自动做设备分配;可用with_device_mapping/with_device手动控制 |
| dtype | ModelDType::Auto | 可用with_dtype指定 |
其他常用配置方法(均为构建器方法):
with_topology/with_topology_from_path:加载时指定模型拓扑;与 ISQ 类型冲突时拓扑优先;with_isq/with_auto_isq/with_imatrix/with_calibration_file:在线量化(ISQ)相关,with_imatrix与with_calibration_file互斥;with_force_cpu:强制 CPU 推理(文档注明不要与 PagedAttention 同时使用);with_hf_revision、from_hf_cache_path、with_tokenizer_json:HF 下载与分词器相关;write_uqff/from_uqff(后者已废弃):UQFF 打包写入/读取,读取建议改用UqffEmbeddingModelBuilder;with_throughput_logging:开启吞吐日志。
build()的内部实现(embedding_model.rs 第 225–228 行)是两步式:先build_embedding_pipeline(self)构造嵌入流水线与调度配置,再由build_model_from_pipeline装载为Model,这与文本/多模态构建器共用同一套 pipeline 构建体系(model_builder_trait)。
EmbeddingRequestBuilder:批量请求的完整 API
批量能力的核心是 mistralrs/src/messages.rs(第 1407–1475 行)中的EmbeddingRequestBuilder,示例只用了其中两个方法,完整 API 如下:
| 方法 | 作用 |
|---|---|
add_prompt(impl Into<String>) | 追加单条文本输入(示例中用于基准请求) |
add_prompts(I: IntoIterator<Item = S: Into<String>>) | 一次性追加多条文本输入(示例中用于 100 条批量) |
add_tokens(impl Into<Vec<u32>>) | 追加单条预分词输入,跳过重复分词 |
add_tokens_batch(I: IntoIterator<Item = Vec<u32>>) | 一次性追加多条预分词输入 |
with_truncate_sequence(bool) | 控制超长输入是否按模型最大上下文截断,默认false |
build() -> anyhow::Result<EmbeddingRequest> | 校验并生成请求;输入为空时报错Embedding request must contain at least one input. |
请求最终表示为EmbeddingRequest { inputs: Vec<EmbeddingRequestInput>, truncate_sequence: bool },每条输入经into_request_message()转为引擎侧的RequestMessage::Embedding(文本)或RequestMessage::EmbeddingTokens(预分词)消息(messages.rs 第 1381–1388 行)。对已经持有 token 序列的下游任务,add_tokens_batch可以直接省去 100 次重复分词开销。
底层并行机制:顺序保证与确定性采样
generate_embeddings的实现在 mistralrs/src/model.rs(第 761–766 行),委托给generate_embeddings_with_model(第 772 行起)。从源码结构看,其关键行为有三点:
- 逐输入并行:
inputs.into_iter().map(...)为每条输入构造一个独立的异步任务,各自经 channel 与推理引擎交互后收集结果——批量调用因此可以在引擎侧并发执行,而不是串行等待; - 顺序保证:每条输入生成前会记录其在请求中的位置,返回值
Vec<Vec<f32>>严格按照输入添加顺序排列,这正是示例中i % 2断言能够成立的前提,API 文档也明确写明 "Returns one embedding vector per input in the same order they were added"; - 确定性参数:批量嵌入请求统一使用
SamplingParams::deterministic()、无 seed、无流式、无工具/约束等附加字段,保证同一输入在批量与单条场景下得到逐位相同的向量(示例中的assert_eq!即验证这一点)。
此外,generate_embeddings_with_model(request, model_id)支持在多模型实例中指定目标模型;model_id为None时发给默认模型。
并发上限与调优
批量编码的实际并行度受max_num_seqs约束:默认 32(见 embedding_model.rs 第 64 行),即示例中 100 条输入会按调度能力分批推进,而非同时占用 100 路。如果你的部署内存/显存更充裕且希望提高单批吞吐,可以:
let model = EmbeddingModelBuilder::new("google/embeddinggemma-300m") .with_logging() .with_max_num_seqs(64) // 提升并行序列上限 .build() .await?;若需观察批量编码的吞吐表现,可追加.with_throughput_logging()。对于 GPU 资源受限场景,保持默认 32 即可让调度器自动排队,接口调用方式无需任何改动——这正是"一次generate_embeddings传一个批量 builder"这一用法的价值所在:调用方代码与逐条调用完全同构,仅凭add_prompts就把串行请求变成了并行编码。
小结
- 用
EmbeddingModelBuilder::new(model_id)加载嵌入模型,默认并行序列上限 32、自动设备映射、HF 缓存令牌; - 用
EmbeddingRequest::builder()的add_prompt/add_prompts(或预分词的add_tokens/add_tokens_batch)组织输入,with_truncate_sequence控制截断,build()完成非空校验; generate_embeddings按输入顺序返回Vec<Vec<f32>>,内部以确定性参数逐输入并发执行,批量结果与单条结果逐位一致;- 完整可运行示例见 mistralrs/examples/advanced/batching_embeddings/main.rs,用
cargo run --release --example batching_embeddings -p mistralrs即可复现并验证。
【免费下载链接】mistral.rsFast, flexible LLM inference项目地址: https://gitcode.com/GitHub_Trending/mi/mistral.rs
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考