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

Rust中如何借助Serde反序列化时直接生成含CLIP嵌入的Embedding结构体

直接在Serde反序列化流程中生成Embedding的方案

核心思路

利用Serde的DeserializeSeed trait携带预加载的ClipModel引用,在反序列化每个字符串时直接调用embed_text生成嵌入向量,无需额外定义中间结构体。DeserializeSeed的作用是允许反序列化过程携带外部状态(这里就是模型),完美适配需要依赖外部资源的转换场景。

具体实现代码

1. 基础定义(假设已存在)

use serde::de::{self, Deserialize, DeserializeSeed, Deserializer, Visitor};
use serde_yaml;
use std::fmt;

// 替换为实际的ClipModel定义
struct ClipModel;
type Result<T> = std::result::Result<T, Box<dyn std::error::Error>>;

/// 用户提供的嵌入生成函数
fn embed_text(clip_model: &ClipModel, text: &str) -> Result<Vec<f64>> {
    // 实际调用CLIP模型的逻辑,此处为示例返回值
    Ok(vec![0.1, 0.2, 0.3])
}

/// 目标结构体
#[derive(Debug)]
struct Embedding {
    values: Vec<f64>,
    original: String,
}

2. 实现带模型的反序列化种子

/// 携带ClipModel引用的反序列化种子
struct EmbeddingSeed<'a> {
    model: &'a ClipModel,
}

impl<'a> DeserializeSeed<'a> for EmbeddingSeed<'a> {
    type Value = Embedding;

    fn deserialize<D>(self, deserializer: D) -> Result<Self::Value, D::Error>
    where
        D: Deserializer<'a>,
    {
        /// 自定义Visitor处理字符串到Embedding的转换
        struct EmbeddingVisitor<'b> {
            model: &'b ClipModel,
        }

        impl<'b> Visitor<'b> for EmbeddingVisitor<'b> {
            type Value = Embedding;

            // 反序列化时的预期提示
            fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
                formatter.write_str("a string to generate embedding for")
            }

            // 处理字符串类型的YAML元素
            fn visit_str<E>(self, s: &str) -> Result<Self::Value, E>
            where
                E: de::Error,
            {
                // 生成嵌入向量并转换为Embedding结构体
                let values = embed_text(self.model, s)
                    .map_err(|e| de::Error::custom(format!("embedding failed: {}", e)))?;
                Ok(Embedding {
                    values,
                    original: s.to_string(),
                })
            }
        }

        // 告诉反序列化器用自定义Visitor处理字符串
        deserializer.deserialize_str(EmbeddingVisitor { model: self.model })
    }
}

3. 读取YAML并生成Embedding列表

fn load_embeddings(model: &ClipModel, yaml_path: &str) -> Result<Vec<Embedding>> {
    let yaml_content = std::fs::read_to_string(yaml_path)?;
    let mut deserializer = serde_yaml::Deserializer::from_str(&yaml_content);

    let mut embeddings = Vec::new();
    // 遍历YAML中的每个字符串元素,用种子反序列化为Embedding
    while let Some(embedding) = serde::de::Deserializer::next_element_seed(&mut deserializer, EmbeddingSeed { model })? {
        embeddings.push(embedding);
    }

    Ok(embeddings)
}

方案优势

  • 无需定义额外的中间结构体(如之前的RawEmbedding),代码更简洁。
  • 预加载的ClipModel仅通过引用传递,避免重复加载的高额成本。
  • 嵌入生成的错误被统一封装到Serde的反序列化错误体系中,便于集中处理。

内容的提问来源于stack exchange,提问作者beyarkay

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.30 01:10:35