Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 9 additions & 1 deletion rust/src/server/src/grpc/convert.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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(),
Expand All @@ -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)) => {
Expand All @@ -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<prost_types::Struct> {
pub(super) fn json_to_proto_struct(value: &serde_json::Value) -> Option<prost_types::Struct> {
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(),
Expand Down
60 changes: 60 additions & 0 deletions rust/src/server/src/grpc/inference.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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<AppState>) -> Self {
Self { state }
Expand Down Expand Up @@ -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)
}
Expand Down
109 changes: 85 additions & 24 deletions rust/src/server/src/grpc/tests.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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};
Expand Down Expand Up @@ -143,6 +144,29 @@ fn default_stream_output_specs() -> Vec<(Vec<u32>, Option<EngineCoreFinishReason
]
}

fn ec_proto_struct(mm_hashes: &[&str]) -> 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"),
Expand Down Expand Up @@ -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<_>>(),
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::<usize>() + 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;
Expand All @@ -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()
Expand Down
Loading