diff --git a/rust/src/server/src/grpc/convert.rs b/rust/src/server/src/grpc/convert.rs index 93669bd1b2ab..2a3f41ffa438 100644 --- a/rust/src/server/src/grpc/convert.rs +++ b/rust/src/server/src/grpc/convert.rs @@ -516,6 +516,11 @@ fn positions_to_proto( // KV transfer params conversion (serde_json::Value ↔ prost_types::Struct) // ======================================================================================== +/// Largest integer exactly representable as `f64` (2^53 - 1). Integral +/// numbers within this bound are emitted as JSON integers so consumers +/// expecting ints (e.g. `image_grid_thw`) don't see `16.0`. +const MAX_SAFE_INTEGER_F64: f64 = ((1u64 << 53) - 1) as f64; + fn proto_struct_to_json(s: &prost_types::Struct) -> serde_json::Value { serde_json::Value::Object( s.fields.iter().map(|(k, v)| (k.clone(), proto_value_to_json(v))).collect(), @@ -527,6 +532,9 @@ fn proto_value_to_json(v: &prost_types::Value) -> serde_json::Value { match v.kind.as_ref() { None | Some(Kind::NullValue(_)) => serde_json::Value::Null, Some(Kind::BoolValue(b)) => serde_json::Value::Bool(*b), + Some(Kind::NumberValue(n)) if n.fract() == 0.0 && n.abs() <= MAX_SAFE_INTEGER_F64 => { + serde_json::Value::Number(serde_json::Number::from(*n as i64)) + } Some(Kind::NumberValue(n)) => serde_json::json!(*n), Some(Kind::StringValue(s)) => serde_json::Value::String(s.clone()), Some(Kind::ListValue(list)) => { @@ -536,7 +544,7 @@ fn proto_value_to_json(v: &prost_types::Value) -> serde_json::Value { } } -fn json_to_proto_struct(value: &serde_json::Value) -> Option { +pub(super) fn json_to_proto_struct(value: &serde_json::Value) -> Option { match value { serde_json::Value::Object(map) => Some(prost_types::Struct { fields: map.iter().map(|(k, v)| (k.clone(), json_to_proto_value(v))).collect(), diff --git a/rust/src/server/src/grpc/inference.rs b/rust/src/server/src/grpc/inference.rs index c3f590af99f6..f864cbde5959 100644 --- a/rust/src/server/src/grpc/inference.rs +++ b/rust/src/server/src/grpc/inference.rs @@ -35,6 +35,65 @@ struct PreparedGrpcRequest { started_at: Instant, } +/// Keep producer metadata for matching EC items; leave unmatched inputs intact. +fn apply_encoder_cache_placeholders(text_request: &mut TextRequest) { + let is_decode_kv_consumer = text_request + .sampling_params + .vllm_xargs + .as_ref() + .and_then(|args| args.get("kv_transfer_params")) + .and_then(|params| params.get("do_remote_prefill")) + .and_then(serde_json::Value::as_bool) + .unwrap_or(false); + let ec_items = text_request + .sampling_params + .vllm_xargs + .as_ref() + .and_then(|args| args.get("ec_transfer_params")) + .and_then(|params| params.get("ec_items")) + .and_then(serde_json::Value::as_array) + .cloned(); + + // Decode uses EC metadata only to prepare the prompt; EngineCore consumes KV. + if is_decode_kv_consumer && let Some(args) = text_request.sampling_params.vllm_xargs.as_mut() { + args.remove("ec_transfer_params"); + } + + let Some(ec_items) = ec_items else { + return; + }; + let Some(features) = text_request.mm_features.as_mut() else { + return; + }; + let ec_items_by_hash: std::collections::HashMap<_, _> = ec_items + .iter() + .filter_map(serde_json::Value::as_object) + .filter_map(|item| { + item.get("mm_hash") + .and_then(serde_json::Value::as_str) + .map(|mm_hash| (mm_hash, item)) + }) + .collect(); + + for feature in features.iter_mut() { + let Some(item) = ec_items_by_hash.get(feature.identifier.as_str()) else { + continue; + }; + let Some(data) = feature.data.as_mut() else { + continue; + }; + let metadata_keys: Vec<_> = data + .keys() + .filter(|key| key.as_str() != "mm_hash" && item.contains_key(key.as_str())) + .cloned() + .collect(); + if metadata_keys.is_empty() { + continue; + } + data.retain(|key, _| metadata_keys.contains(key)); + } +} + impl InferenceServiceImpl { pub fn new(state: Arc) -> Self { Self { state } @@ -106,6 +165,7 @@ impl InferenceServiceImpl { text_request.prompt = Prompt::TokenIds(token_ids); text_request.mm_features = mm_features; } + apply_encoder_cache_placeholders(&mut text_request); Ok(text_request) } diff --git a/rust/src/server/src/grpc/tests.rs b/rust/src/server/src/grpc/tests.rs index 34667c31ea35..ed465e5c7c6c 100644 --- a/rust/src/server/src/grpc/tests.rs +++ b/rust/src/server/src/grpc/tests.rs @@ -49,6 +49,7 @@ use zeromq::prelude::{SocketRecv, SocketSend}; use zeromq::{DealerSocket, PushSocket, ZmqMessage}; use super::control::kv_event_source; +use super::convert::json_to_proto_struct; use super::pb::control_client::ControlClient; use super::pb::inference_client::InferenceClient; use super::{ControlServer, ControlServiceImpl, InferenceServer, InferenceServiceImpl, pb}; @@ -143,6 +144,29 @@ fn default_stream_output_specs() -> Vec<(Vec, Option prost_types::Struct { + let ec_items: Vec<_> = mm_hashes + .iter() + .map(|mm_hash| { + serde_json::json!({ + "image_grid_thw": [[1, 16, 16]], + "mm_hash": mm_hash, + }) + }) + .collect(); + json_to_proto_struct(&serde_json::json!({ "ec_items": ec_items })) + .expect("valid EC proto struct") +} + +fn decode_kv_proto_struct() -> prost_types::Struct { + json_to_proto_struct(&serde_json::json!({ + "do_remote_prefill": true, + "pp_size": 1, + "remote_block_ids": [[7]], + })) + .expect("valid KV proto struct") +} + async fn send_outputs(push: &mut PushSocket, outputs: EngineCoreOutputs) { push.send(ZmqMessage::from( rmp_serde::to_vec_named(&outputs).expect("encode outputs"), @@ -701,22 +725,51 @@ async fn unary_generate_prepares_multimodal_input_for_engine_core() { |request| { let token_ids = request.prompt_token_ids.as_ref().expect("prompt token ids"); let features = request.mm_features.as_ref().expect("multimodal features"); - assert_eq!(features.len(), 1); - - let feature = &features[0]; - assert_eq!(feature.modality, "image"); - assert_eq!(feature.identifier, "image-1"); - assert_eq!(feature.mm_position.offset, 1); - assert!(feature.mm_position.length > 1); - assert_eq!(token_ids.len(), feature.mm_position.length + 2); - assert_eq!(token_ids[0], 11); - assert_eq!(token_ids.last(), Some(&12)); - assert!( - token_ids[feature.mm_position.offset - ..feature.mm_position.offset + feature.mm_position.length] - .iter() - .all(|token_id| *token_id == QWEN_IMAGE_TOKEN_ID) + assert_eq!(features.len(), 2); + + for (feature, identifier) in features.iter().zip(["image-1", "image-2"]) { + assert_eq!(feature.modality, "image"); + assert_eq!(feature.identifier, identifier); + assert!(feature.mm_position.length > 1); + assert_eq!( + feature + .data + .as_ref() + .expect("multimodal feature data") + .keys() + .map(String::as_str) + .collect::>(), + vec!["image_grid_thw"] + ); + } + assert_eq!(features[0].mm_position.offset, 1); + let xargs = request + .sampling_params + .as_ref() + .and_then(|params| params.extra_args.as_ref()) + .expect("KV transfer args"); + let kv_transfer_params = + xargs.get("kv_transfer_params").expect("KV transfer params"); + assert_eq!(kv_transfer_params["pp_size"].as_i64(), Some(1)); + assert_eq!( + kv_transfer_params["remote_block_ids"][0][0].as_i64(), + Some(7) ); + assert!(!xargs.contains_key("ec_transfer_params")); + assert_eq!( + token_ids.len(), + features.iter().map(|feature| feature.mm_position.length).sum::() + 3 + ); + assert_eq!(token_ids[0], 11); + assert_eq!(token_ids.last(), Some(&13)); + for feature in features { + assert!( + token_ids[feature.mm_position.offset + ..feature.mm_position.offset + feature.mm_position.length] + .iter() + .all(|token_id| *token_id == QWEN_IMAGE_TOKEN_ID) + ); + } }, ) .await; @@ -734,16 +787,24 @@ async fn unary_generate_prepares_multimodal_input_for_engine_core() { request_id: "test-multimodal".to_string(), model: "test-model".to_string(), prompt: Some(pb::generate_request::Prompt::TokenIds(pb::TokenIds { - ids: vec![11, QWEN_IMAGE_TOKEN_ID, 12], + ids: vec![11, QWEN_IMAGE_TOKEN_ID, 12, QWEN_IMAGE_TOKEN_ID, 13], })), - media: vec![pb::MediaItem { - modality: pb::Modality::Image as i32, - source: Some(pb::media_item::Source::DataUri( - TINY_PNG_DATA_URI.to_string(), - )), - mime_type: String::new(), - uuid: "image-1".to_string(), - }], + media: ["image-1", "image-2"] + .into_iter() + .map(|uuid| pb::MediaItem { + modality: pb::Modality::Image as i32, + source: Some(pb::media_item::Source::DataUri( + TINY_PNG_DATA_URI.to_string(), + )), + mime_type: String::new(), + uuid: uuid.to_string(), + }) + .collect(), + kv: Some(pb::KvCacheParameters { + kv_transfer_params: Some(decode_kv_proto_struct()), + ec_transfer_params: Some(ec_proto_struct(&["image-2", "image-1"])), + ..Default::default() + }), stopping: Some(pb::StoppingCriteria { max_new_tokens: 10, ..Default::default()