Prhub

#47581 [Rust Frontend] Avoid extra copies for multimodal tensors

原始 PR 作者 reidliu41 合并时间 2026-07-07 11:09 文件变更 2 提交数 2 评论 4 代码增减 +67 / -4

执行摘要

消除多模态张量的两次额外拷贝

多模态请求携带的 tensor payload(如 image pixel_values)可能很大。PR body 指出此前存在两处不必要的克隆:lower_text_request 中克隆了 request.mm_featuresGenerateRequest,以及 RawView 序列化时通过 Value::Ext(…, bytes.clone()) 克隆整个字节缓冲区。这导致 per-request CPU 负载和临时内存压力上升。

推荐阅读此 PR 以了解 Rust 后端中零拷贝序列化的典型模式:借字节 newtype 与 Serde 的集成。设计上值得关注的是如何不引入外部依赖而实现 serialize_bytes。对于后续优化 AuxIndex 路径也有启发。

讨论亮点

Review 中 BugenZhao 询问 ByteSlice newtype 是否必要。作者 reidliu41 解释:如果不使用 newtype,那只能通过 serde_bytes crate 实现相同的 serialize_bytes 行为,而 ByteSlice 无需新增依赖即可实现借字节的序列化,保持线格式不变。BugenZhao 随后批准。

实现拆解

  1. rust/src/text/src/lower.rs: 使用 take() 移动 mm_features
    - 将 lower_text_requestrequest 参数改为 mut,用 request.mm_features.take() 替代 clone(),将 mm_features 的所有权转移给 GenerateRequest
    - 对应的 PreparedTextRequest 中的 text_requestmm_features 字段变为 None,与 Python 侧响应解码路径保持一致。

  2. rust/src/engine-core-client/src/protocol/tensor.rs: 新增 MsgpackExtRef 零拷贝序列化
    - 引入 MsgpackExtRef<'a> 结构体(通过 _ExtStruct 标记),它持有一个 (i8, ByteSlice<'a>) 元组。ByteSlice<'a> 是一个 &[u8] 的新类型包装器,其 Serialize 实现调用 serialize_bytes,确保 rmp-serde 将其编码为 MessagePack 扩展类型值,而非字节序列。
    - 修改 WireArrayData::RawView 的分支,从 Value::Ext(CUSTOM_TYPE_RAW_VIEW, bytes.clone()).serialize(serializer) 改为 MsgpackExtRef((CUSTOM_TYPE_RAW_VIEW, ByteSlice(bytes))).serialize(serializer),完全避免了字节克隆。

  3. 测试补充
    - 在 tensor.rs 新增 raw_view_serializes_as_msgpack_ext 测试,验证 WireArrayData::RawView 序列化结果与 Value::Ext 编码完全一致。
    - 在 lower.rs 新增 lower_text_request_moves_multimodal_features_to_generate_request 测试,验证 mm_features 成功移动到 generate_request,而 text_request 中变为 None

文件 模块 状态 重要度
rust/src/engine-core-client/src/protocol/tensor.rs 客户端协议 modified 7.54
rust/src/text/src/lower.rs 文本处理 modified 7.04

关键符号

lower_text_request serialize

关键源码片段

rust/src/engine-core-client/src/protocol/tensor.rs core-logic

核心序列化路径:新增 `MsgpackExtRef` 和 `ByteSlice` 实现零拷贝 MessagePack 扩展编码,修改 `WireArrayData::RawView` 序列化避免克隆。

/// MessagePack extension struct that serializes without copying bytes.
///
/// `rmp-serde`'s `_ExtStruct` path expects the second tuple element to use
/// `serialize_bytes()` to produce a bin/str value inside the ext.  By wrapping
/// `&[u8]` in `ByteSlice` (which delegates to `serialize_bytes`), we
/// avoid both a `serde_bytes` dependency and cloning the entire byte buffer.
#[derive(Serialize)]
#[serde(rename = "_ExtStruct")]
struct MsgpackExtRef<'a>((i8, ByteSlice<'a>));/// Newtype over `&[u8]` to force `serde`'s bytes serialization path.
struct ByteSlice<'a>(&'a [u8]);impl Serialize for ByteSlice<'_> {
    fn serialize<S>(&self, serializer: S) -> std::result::Result<S::Ok, S::Error>
    where
        S: Serializer,
    {
        // `serialize_bytes` is what `rmp-serde`'s `ExtStruct` expects for the
        // binary payload; a plain `&[u8]` would serialize as a sequence.
        serializer.serialize_bytes(self.0)
    }
}// … inside the `Serialize for WireArrayData` impl:
impl Serialize for WireArrayData {
    fn serialize<S>(&self, serializer: S) -> std::result::Result<S::Ok, S::Error>
    where
        S: Serializer,
    {
        match self {
            Self::AuxIndex(index) => serializer.serialize_u64(*index as u64),
            Self::RawView(bytes) => {
                // Before: `Value::Ext(CUSTOM_TYPE_RAW_VIEW, bytes.clone()).serialize(serializer)`
                // Now: borrow `bytes` without cloning.
                MsgpackExtRef((CUSTOM_TYPE_RAW_VIEW, ByteSlice(bytes))).serialize(serializer)
            }
        }
    }
}
rust/src/text/src/lower.rs core-logic

多模态特征所有权转移:将 `mm_features` 的克隆改为 `take()`,避免大缓冲区拷贝,并调整函数签名为 `mut request`。

/// Convert a high-level [`TextRequest`] into one lower-level
/// [`GenerateRequest`] ready for the `llm` crate.
pub fn lower_text_request(
    mut request: TextRequest, // was `request: TextRequest`
    prompt_token_ids: Vec<u32>,
    sampling_hints: SamplingHints,
    sampling_limits: SamplingLimits,
    tokenizer: &dyn Tokenizer,
) -> Result<PreparedTextRequest> {
    let prompt_len = prompt_token_ids.len() as u32;
    validate_prompt_token_ids(&prompt_token_ids, &sampling_limits)?;    let generate_request = GenerateRequest {
        request_id: request.request_id.clone(),
        prompt_token_ids,
        // Before: `mm_features: request.mm_features.clone()` caused an extra
        // allocation of the potentially-large multimodal tensor bytes.
        // After: move the feature data into the engine request so the retained
        // `text_request` no longer holds it — matching Python's response path.
        mm_features: request.mm_features.take(),
        sampling_params: lower_sampling_params(/* … */)?,
        // … other fields unchanged …
        arrival_time: None,
        trace_headers: None,
    };    Ok(PreparedTextRequest {
        text_request: request, // `mm_features` is now `None` here
        generate_request,
    })
}

评论区精华

ByteSlice newtype 必要性讨论 设计

BugenZhao 问 `ByteSlice` 是否必要,因为单纯 `&[u8]` 也可行。作者解释:没有 `ByteSlice` 时,`&[u8]` 默认序列化为序列(sequence);而 `_ExtStruct` 需要 `serialize_bytes`,所以必须通过 newtype 强制使用字节序列化,避免引入 `serde_bytes` 依赖。

结论:保留 `ByteSlice` 设计,BugenZhao 同意后 LGTM。 · 已解决

风险与影响

风险极低:变更集中在两个 Rust 源文件的序列化路径上,且通过单元测试验证了线格式兼容性。唯一需关注的场景是是否有外部消费者依赖 WireArrayData::RawView 的旧序列化实现(例如通过 Python 端反序列化),但 PR 作者已确保输出完全一致。对 engine 侧行为无影响。

直接影响 Rust 前端处理多模态请求时的内存分配和 CPU 开销,对包含大 tensor 的请求(如图像)可减少多次大缓冲区拷贝。不影响 Python 前端或 engine 侧。团队协作上,由于仅修改了 Rust 前端核心路径,且合入时需确保 CI 中 Rust 测试通过。

借字节序列化依赖运行时生命周期

关联 Issue

未识别关联 Issue

当前没有检测到明确关联的 Issue 链接,后续同步到相关引用后会出现在这里。

完整报告

参与讨论