[EC Connector] Add EC Transfer Params (#42433)
Signed-off-by: omerpaz95 <omerpaz95@gmail.com> Co-authored-by: Or Ozeri <oro@il.ibm.com>
This commit is contained in:
@@ -89,6 +89,7 @@ steps:
|
||||
- tests/v1/simple_kv_offload
|
||||
- tests/v1/worker
|
||||
- tests/v1/kv_connector/unit
|
||||
- tests/v1/ec_connector/unit
|
||||
- tests/v1/metrics
|
||||
- tests/entrypoints/openai/correctness/test_lmeval.py
|
||||
commands:
|
||||
@@ -101,6 +102,7 @@ steps:
|
||||
- pytest -v -s v1/simple_kv_offload
|
||||
- pytest -v -s v1/worker
|
||||
- pytest -v -s -m 'not cpu_test' v1/kv_connector/unit
|
||||
- pytest -v -s -m 'not cpu_test' v1/ec_connector/unit
|
||||
- pytest -v -s -m 'not cpu_test' v1/metrics
|
||||
# Integration test for streaming correctness (requires special branch).
|
||||
- pip install -U git+https://github.com/vllm-project/lm-evaluation-harness.git@streaming-api
|
||||
|
||||
@@ -107,6 +107,9 @@ message KVCacheParameters {
|
||||
|
||||
// KV Connector transfer parameters
|
||||
google.protobuf.Struct kv_transfer_params = 3;
|
||||
|
||||
// Encoder cache connector transfer parameters
|
||||
google.protobuf.Struct ec_transfer_params = 4;
|
||||
}
|
||||
|
||||
// Controls which extra candidate tokens at each position should be returned
|
||||
@@ -173,6 +176,7 @@ message FinishInfo {
|
||||
|
||||
google.protobuf.Struct kv_transfer_params = 6;
|
||||
//uint64 seed = 7;
|
||||
google.protobuf.Struct ec_transfer_params = 8;
|
||||
}
|
||||
|
||||
// Info for candidate tokens other than the input/sampled
|
||||
|
||||
@@ -202,5 +202,8 @@ pub enum ChatEvent {
|
||||
finish_reason: FinishReason,
|
||||
/// Connector-specific KV transfer parameters for disaggregated serving.
|
||||
kv_transfer_params: Option<serde_json::Value>,
|
||||
/// Connector-specific encoder cache transfer parameters for
|
||||
/// disaggregated serving.
|
||||
ec_transfer_params: Option<serde_json::Value>,
|
||||
},
|
||||
}
|
||||
|
||||
@@ -257,6 +257,7 @@ pub(crate) async fn unified_event_stream(
|
||||
usage: finished.usage,
|
||||
finish_reason: finished.finish_reason,
|
||||
kv_transfer_params: finished.kv_transfer_params,
|
||||
ec_transfer_params: finished.ec_transfer_params,
|
||||
})
|
||||
.await;
|
||||
}
|
||||
@@ -387,6 +388,7 @@ mod tests {
|
||||
usage: vllm_llm::TokenUsage::default(),
|
||||
finish_reason: crate::FinishReason::Stop(None),
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
}),
|
||||
}
|
||||
}
|
||||
@@ -628,6 +630,7 @@ mod tests {
|
||||
usage: vllm_llm::TokenUsage::default(),
|
||||
finish_reason: crate::FinishReason::Stop(None),
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
},
|
||||
]
|
||||
);
|
||||
@@ -671,6 +674,7 @@ mod tests {
|
||||
usage: vllm_llm::TokenUsage::default(),
|
||||
finish_reason: crate::FinishReason::Stop(None),
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
},
|
||||
]
|
||||
);
|
||||
|
||||
@@ -370,6 +370,7 @@ async fn harmony_assistant_event_stream(
|
||||
usage: finished.usage,
|
||||
finish_reason: finished.finish_reason,
|
||||
kv_transfer_params: finished.kv_transfer_params,
|
||||
ec_transfer_params: finished.ec_transfer_params,
|
||||
})
|
||||
.await;
|
||||
}
|
||||
|
||||
@@ -52,6 +52,7 @@ fn finished() -> Finished {
|
||||
},
|
||||
finish_reason: FinishReason::stop_eos(),
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -115,6 +116,7 @@ fn interrupted_final_message_is_preserved() {
|
||||
},
|
||||
finish_reason: FinishReason::stop_eos(),
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
})
|
||||
);
|
||||
}
|
||||
@@ -175,6 +177,7 @@ fn interrupted_analysis_message_is_preserved() {
|
||||
},
|
||||
finish_reason: FinishReason::stop_eos(),
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
})
|
||||
);
|
||||
}
|
||||
|
||||
@@ -48,6 +48,9 @@ pub(crate) enum AssistantEvent {
|
||||
finish_reason: FinishReason,
|
||||
/// Connector-specific KV transfer parameters for disaggregated serving.
|
||||
kv_transfer_params: Option<serde_json::Value>,
|
||||
/// Connector-specific encoder cache transfer parameters for
|
||||
/// disaggregated serving.
|
||||
ec_transfer_params: Option<serde_json::Value>,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@@ -146,6 +146,7 @@ impl StructuredEventState {
|
||||
usage: vllm_llm::TokenUsage,
|
||||
finish_reason: FinishReason,
|
||||
kv_transfer_params: Option<serde_json::Value>,
|
||||
ec_transfer_params: Option<serde_json::Value>,
|
||||
) -> Result<Vec<ChatEvent>> {
|
||||
let mut events = Vec::new();
|
||||
self.close_open_text_block(&mut events);
|
||||
@@ -155,6 +156,7 @@ impl StructuredEventState {
|
||||
usage,
|
||||
finish_reason,
|
||||
kv_transfer_params,
|
||||
ec_transfer_params,
|
||||
});
|
||||
Ok(events)
|
||||
}
|
||||
@@ -296,8 +298,11 @@ pub(crate) async fn structured_chat_event_stream(
|
||||
usage,
|
||||
finish_reason,
|
||||
kv_transfer_params,
|
||||
ec_transfer_params,
|
||||
} => {
|
||||
for next in state.finish(usage, finish_reason, kv_transfer_params)? {
|
||||
for next in
|
||||
state.finish(usage, finish_reason, kv_transfer_params, ec_transfer_params)?
|
||||
{
|
||||
y.yield_ok(next).await;
|
||||
}
|
||||
}
|
||||
@@ -334,6 +339,7 @@ mod tests {
|
||||
},
|
||||
finish_reason: FinishReason::stop_eos(),
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
}),
|
||||
]);
|
||||
|
||||
@@ -388,6 +394,7 @@ mod tests {
|
||||
},
|
||||
finish_reason: FinishReason::stop_eos(),
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
}),
|
||||
]);
|
||||
|
||||
@@ -439,6 +446,7 @@ mod tests {
|
||||
},
|
||||
finish_reason: FinishReason::stop_eos(),
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
}),
|
||||
]);
|
||||
|
||||
@@ -490,6 +498,7 @@ mod tests {
|
||||
},
|
||||
finish_reason: FinishReason::stop_eos(),
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
}),
|
||||
]);
|
||||
|
||||
@@ -557,6 +566,7 @@ mod tests {
|
||||
},
|
||||
finish_reason: FinishReason::stop_eos(),
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
}),
|
||||
]);
|
||||
|
||||
|
||||
@@ -22,6 +22,9 @@ pub struct CollectedAssistantMessage {
|
||||
pub finish_reason: FinishReason,
|
||||
/// Connector-specific KV transfer parameters for disaggregated serving.
|
||||
pub kv_transfer_params: Option<serde_json::Value>,
|
||||
/// Connector-specific encoder cache transfer parameters for disaggregated
|
||||
/// serving.
|
||||
pub ec_transfer_params: Option<serde_json::Value>,
|
||||
}
|
||||
|
||||
/// Per-request stream of chat events.
|
||||
@@ -77,6 +80,7 @@ impl ChatEventStream {
|
||||
usage,
|
||||
finish_reason,
|
||||
kv_transfer_params,
|
||||
ec_transfer_params,
|
||||
} => {
|
||||
return Ok(CollectedAssistantMessage {
|
||||
message: done,
|
||||
@@ -89,6 +93,7 @@ impl ChatEventStream {
|
||||
usage,
|
||||
finish_reason,
|
||||
kv_transfer_params,
|
||||
ec_transfer_params,
|
||||
});
|
||||
}
|
||||
ChatEvent::ToolCallEnd { call, .. } => {
|
||||
@@ -194,6 +199,7 @@ mod tests {
|
||||
},
|
||||
finish_reason: FinishReason::stop_eos(),
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
}),
|
||||
]),
|
||||
);
|
||||
@@ -234,6 +240,7 @@ mod tests {
|
||||
},
|
||||
finish_reason: FinishReason::stop_eos(),
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
}
|
||||
);
|
||||
}
|
||||
|
||||
@@ -50,6 +50,7 @@ fn request_output(
|
||||
stop_reason,
|
||||
events: None,
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
trace_headers: None,
|
||||
prefill_stats: None,
|
||||
routed_experts: None,
|
||||
@@ -75,6 +76,7 @@ fn request_output_with_logprobs(
|
||||
stop_reason,
|
||||
events: None,
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
trace_headers: None,
|
||||
prefill_stats: None,
|
||||
routed_experts: None,
|
||||
|
||||
@@ -676,6 +676,7 @@ fn decoded_completion_stream(
|
||||
usage: Default::default(),
|
||||
finish_reason: FinishReason::stop_eos(),
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
}),
|
||||
}
|
||||
});
|
||||
@@ -686,6 +687,7 @@ fn decoded_completion_stream(
|
||||
usage: Default::default(),
|
||||
finish_reason: FinishReason::stop_eos(),
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
});
|
||||
events.push(DecodedTextEvent::TextDelta {
|
||||
delta: chunk.delta,
|
||||
|
||||
@@ -97,6 +97,8 @@ pub struct EngineCoreOutput {
|
||||
#[serde(default)]
|
||||
pub kv_transfer_params: Option<serde_json::Value>,
|
||||
#[serde(default)]
|
||||
pub ec_transfer_params: Option<serde_json::Value>,
|
||||
#[serde(default)]
|
||||
pub trace_headers: Option<OpaqueValue>,
|
||||
/// Breakdown of the scheduled prefill computation, set on the first output
|
||||
/// of a newly scheduled prefill and elided for subsequent decode outputs.
|
||||
@@ -374,6 +376,7 @@ mod tests {
|
||||
stop_reason: Some(StopReason::Text("stop".to_string())),
|
||||
events: None,
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
trace_headers: None,
|
||||
prefill_stats: None,
|
||||
routed_experts: None,
|
||||
@@ -426,6 +429,7 @@ mod tests {
|
||||
stop_reason: None,
|
||||
events: None,
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
trace_headers: None,
|
||||
prefill_stats: None,
|
||||
routed_experts: None,
|
||||
|
||||
@@ -226,6 +226,7 @@ fn request_output(
|
||||
stop_reason: None,
|
||||
events: None,
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
trace_headers: None,
|
||||
prefill_stats: None,
|
||||
routed_experts: None,
|
||||
@@ -2517,6 +2518,7 @@ fn python_msgpack_fixtures_match_rust_encoding() {
|
||||
stop_reason: None,
|
||||
events: None,
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
trace_headers: None,
|
||||
prefill_stats: None,
|
||||
routed_experts: None,
|
||||
|
||||
@@ -91,6 +91,7 @@ class EngineCoreOutput(
|
||||
stop_reason: int | str | None = None
|
||||
events: object | None = None
|
||||
kv_transfer_params: object | None = None
|
||||
ec_transfer_params: object | None = None
|
||||
trace_headers: object | None = None
|
||||
prefill_stats: object | None = None
|
||||
routed_experts: object | None = None
|
||||
|
||||
@@ -38,6 +38,9 @@ pub struct CollectedGenerateOutput {
|
||||
pub usage: TokenUsage,
|
||||
/// Connector-specific KV transfer parameters for disaggregated serving.
|
||||
pub kv_transfer_params: Option<serde_json::Value>,
|
||||
/// Connector-specific encoder cache transfer parameters for disaggregated
|
||||
/// serving.
|
||||
pub ec_transfer_params: Option<serde_json::Value>,
|
||||
}
|
||||
|
||||
/// Prompt-scoped metadata emitted only once on the first [`GenerateOutput`] for
|
||||
@@ -146,6 +149,9 @@ pub struct GenerateOutput {
|
||||
pub cached_token_count: usize,
|
||||
/// Connector-specific KV transfer parameters for disaggregated serving.
|
||||
pub kv_transfer_params: Option<serde_json::Value>,
|
||||
/// Connector-specific encoder cache transfer parameters for disaggregated
|
||||
/// serving.
|
||||
pub ec_transfer_params: Option<serde_json::Value>,
|
||||
}
|
||||
|
||||
impl GenerateOutput {
|
||||
@@ -192,6 +198,7 @@ impl GenerateOutput {
|
||||
finish_reason,
|
||||
cached_token_count: 0,
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -280,6 +287,7 @@ impl Stream for GenerateOutputStream {
|
||||
finish_reason,
|
||||
cached_token_count,
|
||||
kv_transfer_params: raw.kv_transfer_params,
|
||||
ec_transfer_params: raw.ec_transfer_params,
|
||||
};
|
||||
|
||||
Poll::Ready(Some(Ok(output)))
|
||||
@@ -362,6 +370,7 @@ impl<T: Stream<Item = Result<GenerateOutput>> + Send> T {
|
||||
cached_token_count,
|
||||
},
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
});
|
||||
}
|
||||
|
||||
@@ -374,6 +383,7 @@ impl<T: Stream<Item = Result<GenerateOutput>> + Send> T {
|
||||
cached_token_count,
|
||||
};
|
||||
collected.kv_transfer_params = output.kv_transfer_params;
|
||||
collected.ec_transfer_params = output.ec_transfer_params;
|
||||
return Ok(collected);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -51,6 +51,7 @@ fn request_output_with_events(
|
||||
stop_reason: None,
|
||||
events,
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
trace_headers: None,
|
||||
prefill_stats: None,
|
||||
routed_experts: None,
|
||||
@@ -75,6 +76,7 @@ fn request_output_with_logprobs(
|
||||
stop_reason: None,
|
||||
events: None,
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
trace_headers: None,
|
||||
prefill_stats: None,
|
||||
routed_experts: None,
|
||||
@@ -89,6 +91,7 @@ fn request_output_with_logprobs_and_kv(
|
||||
new_logprobs: Option<Logprobs>,
|
||||
prompt_logprobs: Option<Logprobs>,
|
||||
kv_transfer_params: Option<serde_json::Value>,
|
||||
ec_transfer_params: Option<serde_json::Value>,
|
||||
) -> EngineCoreOutput {
|
||||
EngineCoreOutput {
|
||||
request_id: request_id.to_string(),
|
||||
@@ -100,6 +103,7 @@ fn request_output_with_logprobs_and_kv(
|
||||
stop_reason: None,
|
||||
events: None,
|
||||
kv_transfer_params,
|
||||
ec_transfer_params,
|
||||
trace_headers: None,
|
||||
prefill_stats: None,
|
||||
routed_experts: None,
|
||||
@@ -356,6 +360,7 @@ async fn collect_output_aggregates_raw_tokens_logprobs_and_terminal_metadata() {
|
||||
Some(logprobs_for_position(44, -0.3, 1, 88, -0.4)),
|
||||
None,
|
||||
Some(serde_json::json!({"connector": "x"})),
|
||||
None,
|
||||
),
|
||||
],
|
||||
..Default::default()
|
||||
|
||||
@@ -69,6 +69,11 @@ pub fn to_text_request(
|
||||
let map = sampling_params.vllm_xargs.get_or_insert_with(Default::default);
|
||||
map.insert("kv_transfer_params".to_string(), kv_json);
|
||||
}
|
||||
if let Some(ec_struct) = kv.ec_transfer_params.as_ref() {
|
||||
let ec_json = proto_struct_to_json(ec_struct);
|
||||
let map = sampling_params.vllm_xargs.get_or_insert_with(Default::default);
|
||||
map.insert("ec_transfer_params".to_string(), ec_json);
|
||||
}
|
||||
if kv.bypass_prefix_cache {
|
||||
sampling_params.skip_reading_prefix_cache = Some(true);
|
||||
}
|
||||
@@ -343,6 +348,7 @@ fn to_finish_info(finished: &Finished, token_ids: &[u32]) -> pb::FinishInfo {
|
||||
finish_reason,
|
||||
stop_reason,
|
||||
kv_transfer_params: finished.kv_transfer_params.as_ref().and_then(json_to_proto_struct),
|
||||
ec_transfer_params: finished.ec_transfer_params.as_ref().and_then(json_to_proto_struct),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -586,6 +592,7 @@ mod tests {
|
||||
},
|
||||
finish_reason: reason,
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -71,6 +71,7 @@ impl pb::generate_server::Generate for GenerateServiceImpl {
|
||||
usage: collected.usage,
|
||||
finish_reason: collected.finish_reason,
|
||||
kv_transfer_params: collected.kv_transfer_params,
|
||||
ec_transfer_params: collected.ec_transfer_params,
|
||||
};
|
||||
|
||||
let outputs = convert::to_sequence_output(
|
||||
|
||||
@@ -104,6 +104,7 @@ fn request_output(
|
||||
stop_reason: None,
|
||||
events: None,
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
trace_headers: None,
|
||||
prefill_stats: None,
|
||||
routed_experts: None,
|
||||
|
||||
@@ -100,6 +100,7 @@ fn request_output(
|
||||
stop_reason: None,
|
||||
events: None,
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
trace_headers: None,
|
||||
prefill_stats: None,
|
||||
routed_experts: None,
|
||||
|
||||
@@ -266,6 +266,7 @@ fn collect_generate(
|
||||
}],
|
||||
prompt_logprobs,
|
||||
kv_transfer_params: collected.kv_transfer_params,
|
||||
ec_transfer_params: collected.ec_transfer_params,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -404,6 +405,7 @@ mod tests {
|
||||
finish_reason: None,
|
||||
cached_token_count: 0,
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
}),
|
||||
Ok(GenerateOutput {
|
||||
request_id: String::new(),
|
||||
@@ -416,6 +418,7 @@ mod tests {
|
||||
finish_reason: Some(FinishReason::stop_eos()),
|
||||
cached_token_count: 2,
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
}),
|
||||
]);
|
||||
|
||||
|
||||
@@ -4,7 +4,7 @@ use super::types::GenerateRequest;
|
||||
use super::validate;
|
||||
use crate::error::ApiError;
|
||||
use crate::lora::LoraModelResolution;
|
||||
use crate::utils::{ResolvedRequestContext, merge_kv_transfer_params};
|
||||
use crate::utils::{ResolvedRequestContext, merge_ec_transfer_params, merge_kv_transfer_params};
|
||||
|
||||
/// Lowered generate request plus the response request ID.
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
@@ -56,6 +56,10 @@ pub(super) fn prepare_generate_request(
|
||||
sampling_params.vllm_xargs,
|
||||
request.kv_transfer_params.as_ref(),
|
||||
);
|
||||
sampling_params.vllm_xargs = merge_ec_transfer_params(
|
||||
sampling_params.vllm_xargs,
|
||||
request.ec_transfer_params.as_ref(),
|
||||
);
|
||||
|
||||
let text_request = TextRequest {
|
||||
request_id: ctx.request_id.clone(),
|
||||
|
||||
@@ -22,6 +22,7 @@ pub struct GenerateRequest {
|
||||
#[serde(default)]
|
||||
pub priority: i32,
|
||||
pub kv_transfer_params: Option<HashMap<String, Value>>,
|
||||
pub ec_transfer_params: Option<HashMap<String, Value>>,
|
||||
#[serde(flatten)]
|
||||
pub other: Map<String, Value>,
|
||||
}
|
||||
@@ -66,6 +67,7 @@ pub(super) struct GenerateResponse {
|
||||
pub choices: Vec<GenerateResponseChoice>,
|
||||
pub prompt_logprobs: Option<Vec<Option<HashMap<u32, GenerateLogprob>>>>,
|
||||
pub kv_transfer_params: Option<Value>,
|
||||
pub ec_transfer_params: Option<Value>,
|
||||
}
|
||||
|
||||
/// Mirrors the Python vLLM `Logprob` class used in prompt-logprobs payloads.
|
||||
|
||||
@@ -144,6 +144,7 @@ async fn collect_chat_completion(
|
||||
usage,
|
||||
finish_reason,
|
||||
kv_transfer_params,
|
||||
ec_transfer_params,
|
||||
} = collected;
|
||||
let stop_reason = finish_reason.as_stop_reason().map(stop_reason_to_json);
|
||||
let saw_tool_calls = message.tool_calls().next().is_some();
|
||||
@@ -224,6 +225,7 @@ async fn collect_chat_completion(
|
||||
prompt_logprobs,
|
||||
prompt_token_ids: return_token_ids.then(|| prompt_token_ids.to_vec()),
|
||||
kv_transfer_params,
|
||||
ec_transfer_params,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -951,6 +953,7 @@ mod tests {
|
||||
},
|
||||
finish_reason: FinishReason::stop_eos(),
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
}),
|
||||
]);
|
||||
|
||||
@@ -1031,6 +1034,7 @@ mod tests {
|
||||
},
|
||||
finish_reason: FinishReason::stop_eos(),
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
}),
|
||||
]);
|
||||
|
||||
@@ -1086,6 +1090,7 @@ mod tests {
|
||||
},
|
||||
finish_reason: FinishReason::stop_eos(),
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
}),
|
||||
]);
|
||||
|
||||
@@ -1167,6 +1172,7 @@ mod tests {
|
||||
},
|
||||
finish_reason: FinishReason::stop_eos(),
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
}),
|
||||
]);
|
||||
|
||||
@@ -1300,6 +1306,7 @@ mod tests {
|
||||
},
|
||||
finish_reason: FinishReason::stop_eos(),
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
}),
|
||||
]);
|
||||
|
||||
@@ -1381,6 +1388,7 @@ mod tests {
|
||||
},
|
||||
finish_reason: FinishReason::stop_eos(),
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
}),
|
||||
]);
|
||||
|
||||
|
||||
@@ -13,7 +13,9 @@ use crate::routes::openai::utils::structured_outputs::convert_from_response_form
|
||||
use crate::routes::openai::utils::types::{
|
||||
ChatMessage, ContentPart, MessageContent, Tool, ToolChoice, ToolChoiceValue,
|
||||
};
|
||||
use crate::utils::{ResolvedRequestContext, convert_logit_bias, merge_kv_transfer_params};
|
||||
use crate::utils::{
|
||||
ResolvedRequestContext, convert_logit_bias, merge_ec_transfer_params, merge_kv_transfer_params,
|
||||
};
|
||||
|
||||
/// Lowered chat request plus the public response metadata carried by every SSE
|
||||
/// chunk.
|
||||
@@ -132,7 +134,7 @@ pub(super) fn prepare_chat_request(
|
||||
structured_outputs,
|
||||
skip_reading_prefix_cache: None,
|
||||
vllm_xargs: merge_kv_transfer_params(
|
||||
request.vllm_xargs,
|
||||
merge_ec_transfer_params(request.vllm_xargs, request.ec_transfer_params.as_ref()),
|
||||
request.kv_transfer_params.as_ref(),
|
||||
),
|
||||
},
|
||||
|
||||
@@ -232,6 +232,9 @@ pub struct ChatCompletionRequest {
|
||||
/// KV transfer parameters for disaggregated serving
|
||||
pub kv_transfer_params: Option<HashMap<String, Value>>,
|
||||
|
||||
/// Encoder cache transfer parameters for disaggregated serving
|
||||
pub ec_transfer_params: Option<HashMap<String, Value>>,
|
||||
|
||||
/// Additional request parameters with string or numeric values for custom
|
||||
/// extensions
|
||||
pub vllm_xargs: Option<HashMap<String, Value>>,
|
||||
@@ -299,6 +302,7 @@ impl Default for ChatCompletionRequest {
|
||||
return_token_ids: None,
|
||||
cache_salt: None,
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
vllm_xargs: None,
|
||||
repetition_detection: None,
|
||||
}
|
||||
@@ -346,6 +350,7 @@ pub(super) struct ChatCompletionResponse {
|
||||
pub prompt_logprobs: Option<Vec<Option<HashMap<String, f32>>>>,
|
||||
pub prompt_token_ids: Option<Vec<u32>>,
|
||||
pub kv_transfer_params: Option<Value>,
|
||||
pub ec_transfer_params: Option<Value>,
|
||||
}
|
||||
|
||||
/// Mirrors the Python vLLM `ChatCompletionResponseChoice` class.
|
||||
|
||||
@@ -212,6 +212,7 @@ async fn collect_completion(
|
||||
usage: Some(usage),
|
||||
system_fingerprint: None,
|
||||
kv_transfer_params: collected.kv_transfer_params,
|
||||
ec_transfer_params: collected.ec_transfer_params,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -682,6 +683,7 @@ mod tests {
|
||||
"repetition_detected".to_string(),
|
||||
))),
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
}),
|
||||
}),
|
||||
]);
|
||||
@@ -786,6 +788,7 @@ mod tests {
|
||||
},
|
||||
finish_reason: FinishReason::Length,
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
}),
|
||||
}),
|
||||
]);
|
||||
@@ -837,6 +840,7 @@ mod tests {
|
||||
},
|
||||
finish_reason: FinishReason::Length,
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
}),
|
||||
}),
|
||||
]);
|
||||
@@ -891,6 +895,7 @@ mod tests {
|
||||
},
|
||||
finish_reason: FinishReason::Length,
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
}),
|
||||
}),
|
||||
]);
|
||||
@@ -962,6 +967,7 @@ mod tests {
|
||||
},
|
||||
finish_reason: FinishReason::Length,
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
}),
|
||||
}),
|
||||
]);
|
||||
@@ -1035,6 +1041,7 @@ mod tests {
|
||||
},
|
||||
finish_reason: FinishReason::Length,
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
}),
|
||||
}),
|
||||
]);
|
||||
|
||||
@@ -7,7 +7,9 @@ use crate::error::ApiError;
|
||||
use crate::lora::LoraModelResolution;
|
||||
use crate::routes::openai::completions::validate;
|
||||
use crate::routes::openai::utils::structured_outputs::convert_from_response_format_value;
|
||||
use crate::utils::{ResolvedRequestContext, convert_logit_bias, merge_kv_transfer_params};
|
||||
use crate::utils::{
|
||||
ResolvedRequestContext, convert_logit_bias, merge_ec_transfer_params, merge_kv_transfer_params,
|
||||
};
|
||||
|
||||
/// Lowered completion request plus the public response metadata carried by
|
||||
/// every SSE chunk.
|
||||
@@ -127,7 +129,7 @@ pub(super) fn prepare_completion_request(
|
||||
structured_outputs,
|
||||
skip_reading_prefix_cache: None,
|
||||
vllm_xargs: merge_kv_transfer_params(
|
||||
request.vllm_xargs,
|
||||
merge_ec_transfer_params(request.vllm_xargs, request.ec_transfer_params.as_ref()),
|
||||
request.kv_transfer_params.as_ref(),
|
||||
),
|
||||
},
|
||||
|
||||
@@ -174,6 +174,9 @@ pub struct CompletionRequest {
|
||||
/// KV transfer parameters for disaggregated serving
|
||||
pub kv_transfer_params: Option<HashMap<String, Value>>,
|
||||
|
||||
/// Encoder cache transfer parameters for disaggregated serving
|
||||
pub ec_transfer_params: Option<HashMap<String, Value>>,
|
||||
|
||||
/// Additional request parameters with string or numeric values for custom
|
||||
/// extensions
|
||||
pub vllm_xargs: Option<HashMap<String, Value>>,
|
||||
@@ -209,6 +212,7 @@ pub(super) struct CompletionResponse {
|
||||
pub usage: Option<Usage>,
|
||||
pub system_fingerprint: Option<String>,
|
||||
pub kv_transfer_params: Option<Value>,
|
||||
pub ec_transfer_params: Option<Value>,
|
||||
}
|
||||
|
||||
/// Mirrors the Python vLLM `CompletionResponseChoice` class.
|
||||
|
||||
@@ -76,6 +76,7 @@ fn request_output_with_stop_reason(
|
||||
stop_reason,
|
||||
events: None,
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
trace_headers: None,
|
||||
prefill_stats: None,
|
||||
routed_experts: None,
|
||||
@@ -101,6 +102,7 @@ fn request_output_with_logprobs(
|
||||
stop_reason,
|
||||
events: None,
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
trace_headers: None,
|
||||
prefill_stats: None,
|
||||
routed_experts: None,
|
||||
@@ -116,6 +118,7 @@ fn request_output_with_logprobs_and_kv(
|
||||
new_logprobs: Option<Logprobs>,
|
||||
new_prompt_logprobs_tensors: Option<Logprobs>,
|
||||
kv_transfer_params: Option<serde_json::Value>,
|
||||
ec_transfer_params: Option<serde_json::Value>,
|
||||
) -> EngineCoreOutput {
|
||||
EngineCoreOutput {
|
||||
request_id: request_id.to_string(),
|
||||
@@ -127,6 +130,7 @@ fn request_output_with_logprobs_and_kv(
|
||||
stop_reason,
|
||||
events: None,
|
||||
kv_transfer_params,
|
||||
ec_transfer_params,
|
||||
trace_headers: None,
|
||||
prefill_stats: None,
|
||||
routed_experts: None,
|
||||
@@ -3634,6 +3638,7 @@ async fn non_stream_raw_generate_returns_token_output_envelope() {
|
||||
Some(sample_logprobs_for_token(44, 45)),
|
||||
None,
|
||||
Some(json!({"connector": "x"})),
|
||||
None,
|
||||
),
|
||||
],
|
||||
..Default::default()
|
||||
|
||||
@@ -45,6 +45,24 @@ pub fn merge_kv_transfer_params(
|
||||
xargs
|
||||
}
|
||||
|
||||
/// Merge `ec_transfer_params` into the `vllm_xargs` map, mirroring the Python
|
||||
/// vLLM behavior where `ec_transfer_params` is injected into `extra_args` for
|
||||
/// engine-core consumption.
|
||||
pub fn merge_ec_transfer_params(
|
||||
mut xargs: Option<HashMap<String, Value>>,
|
||||
ec_transfer_params: Option<&HashMap<String, Value>>,
|
||||
) -> Option<HashMap<String, Value>> {
|
||||
if let Some(ec_params) = ec_transfer_params {
|
||||
let map = xargs.get_or_insert_with(HashMap::new);
|
||||
map.insert(
|
||||
"ec_transfer_params".to_string(),
|
||||
// This is safe because we know that `ec_params` is already valid JSON.
|
||||
serde_json::to_value(ec_params).unwrap(),
|
||||
);
|
||||
}
|
||||
xargs
|
||||
}
|
||||
|
||||
/// Convert OpenAI-style `logit_bias` with string token-ID keys into the
|
||||
/// internal `HashMap<u32, f32>` representation, validating that every key
|
||||
/// parses as a `u32`.
|
||||
|
||||
@@ -44,6 +44,9 @@ pub struct Finished {
|
||||
pub finish_reason: FinishReason,
|
||||
/// Connector-specific KV transfer parameters for disaggregated serving.
|
||||
pub kv_transfer_params: Option<serde_json::Value>,
|
||||
/// Connector-specific encoder cache transfer parameters for disaggregated
|
||||
/// serving.
|
||||
pub ec_transfer_params: Option<serde_json::Value>,
|
||||
}
|
||||
|
||||
/// Internal decoded-text event emitted before higher-level assistant
|
||||
@@ -151,6 +154,7 @@ pub async fn decoded_text_event_stream(
|
||||
let decoder = decoder.as_mut().unwrap();
|
||||
|
||||
let kv_transfer_params = output.kv_transfer_params;
|
||||
let ec_transfer_params = output.ec_transfer_params;
|
||||
let mut finish_reason = output.finish_reason;
|
||||
let mut stop_str_matched = false;
|
||||
let suppress_terminal_stop_token = finish_reason.as_ref().is_some_and(|r| r.is_stop())
|
||||
@@ -275,6 +279,7 @@ pub async fn decoded_text_event_stream(
|
||||
},
|
||||
finish_reason: reason,
|
||||
kv_transfer_params,
|
||||
ec_transfer_params,
|
||||
}),
|
||||
})
|
||||
.await;
|
||||
|
||||
@@ -26,6 +26,9 @@ pub struct CollectedTextOutput {
|
||||
pub usage: vllm_llm::TokenUsage,
|
||||
/// Connector-specific KV transfer parameters for disaggregated serving.
|
||||
pub kv_transfer_params: Option<serde_json::Value>,
|
||||
/// Connector-specific encoder cache transfer parameters for disaggregated
|
||||
/// serving.
|
||||
pub ec_transfer_params: Option<serde_json::Value>,
|
||||
}
|
||||
|
||||
#[allow(clippy::manual_async_fn, reason = "specify `Send` bound")]
|
||||
@@ -77,6 +80,7 @@ impl<T: TextOutputStream> T {
|
||||
finish_reason: FinishReason::Error,
|
||||
usage: vllm_llm::TokenUsage::default(),
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
})
|
||||
};
|
||||
|
||||
@@ -85,6 +89,7 @@ impl<T: TextOutputStream> T {
|
||||
collected.finish_reason = finished.finish_reason;
|
||||
collected.usage = finished.usage;
|
||||
collected.kv_transfer_params = finished.kv_transfer_params;
|
||||
collected.ec_transfer_params = finished.ec_transfer_params;
|
||||
return Ok(collected);
|
||||
}
|
||||
}
|
||||
@@ -156,6 +161,7 @@ mod tests {
|
||||
},
|
||||
finish_reason: FinishReason::stop_eos(),
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
}),
|
||||
}),
|
||||
]);
|
||||
@@ -273,6 +279,7 @@ mod tests {
|
||||
},
|
||||
finish_reason: FinishReason::stop_eos(),
|
||||
kv_transfer_params: None,
|
||||
ec_transfer_params: None,
|
||||
}),
|
||||
}),
|
||||
]);
|
||||
|
||||
@@ -294,7 +294,7 @@ def test_abort_request_when_structured_output_fsm_cannot_advance():
|
||||
def free_request(req, delay_free_blocks=False):
|
||||
scheduler.finished_req_ids.add(req.request_id)
|
||||
scheduler.requests.pop(req.request_id, None)
|
||||
return None
|
||||
return None, None
|
||||
|
||||
scheduler._free_request = Mock(side_effect=free_request)
|
||||
|
||||
|
||||
@@ -2968,7 +2968,7 @@ def test_abort_request_when_structured_output_fsm_cannot_advance():
|
||||
def free_request(req: Request, delay_free_blocks: bool = False):
|
||||
scheduler.finished_req_ids.add(req.request_id)
|
||||
scheduler.requests.pop(req.request_id, None)
|
||||
return None
|
||||
return None, None
|
||||
|
||||
scheduler._free_request = Mock(side_effect=free_request)
|
||||
|
||||
|
||||
@@ -0,0 +1,118 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""Sanity tests for ec_transfer_params protocol plumbing.
|
||||
|
||||
No running engine required.
|
||||
"""
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from tests.v1.core.utils import create_scheduler
|
||||
from vllm.entrypoints.openai.chat_completion.protocol import (
|
||||
ChatCompletionRequest,
|
||||
)
|
||||
from vllm.outputs import CompletionOutput, RequestOutput
|
||||
from vllm.sampling_params import SamplingParams
|
||||
from vllm.v1.request import Request, RequestStatus
|
||||
|
||||
EC_PARAMS: dict = {"mm_hash_abc": {"peer_host": "10.0.0.1", "peer_port": 5501}}
|
||||
|
||||
|
||||
def test_ec_transfer_params_routed_to_sampling_params_extra_args():
|
||||
"""ec_transfer_params on the request must land in SamplingParams.extra_args."""
|
||||
req = ChatCompletionRequest(
|
||||
model="test-model",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
max_tokens=5,
|
||||
ec_transfer_params=EC_PARAMS,
|
||||
)
|
||||
sp = req.to_sampling_params(max_tokens=5, default_sampling_params={})
|
||||
assert sp.extra_args is not None
|
||||
assert sp.extra_args.get("ec_transfer_params") == EC_PARAMS
|
||||
|
||||
|
||||
def test_request_output_add_propagates_ec_transfer_params():
|
||||
"""RequestOutput.add() must carry ec_transfer_params forward to the caller."""
|
||||
|
||||
def _out(ec_params):
|
||||
return RequestOutput(
|
||||
request_id="r1",
|
||||
prompt="p",
|
||||
prompt_token_ids=[1],
|
||||
prompt_logprobs=None,
|
||||
outputs=[
|
||||
CompletionOutput(
|
||||
index=0,
|
||||
text="",
|
||||
token_ids=[],
|
||||
cumulative_logprob=0.0,
|
||||
logprobs=None,
|
||||
finish_reason=None,
|
||||
)
|
||||
],
|
||||
finished=False,
|
||||
ec_transfer_params=ec_params,
|
||||
)
|
||||
|
||||
accumulated = _out(None)
|
||||
accumulated.add(_out(EC_PARAMS), aggregate=True)
|
||||
assert accumulated.ec_transfer_params == EC_PARAMS
|
||||
|
||||
|
||||
def test_request_reads_ec_transfer_params_from_extra_args():
|
||||
"""v1 Request must pull ec_transfer_params out of SamplingParams.extra_args."""
|
||||
sp = SamplingParams(extra_args={"ec_transfer_params": EC_PARAMS})
|
||||
req = Request(
|
||||
request_id="r1",
|
||||
prompt_token_ids=[1, 2, 3],
|
||||
sampling_params=sp,
|
||||
pooling_params=None,
|
||||
)
|
||||
assert req.ec_transfer_params == EC_PARAMS
|
||||
|
||||
|
||||
def test_free_request_calls_ec_connector_and_surfaces_params():
|
||||
"""_free_request must call ec_connector.request_finished() and return its params."""
|
||||
sp = SamplingParams(max_tokens=1)
|
||||
sp.update_from_generation_config({}, 50256)
|
||||
request = Request(
|
||||
request_id="test-req",
|
||||
prompt_token_ids=[1, 2, 3],
|
||||
sampling_params=sp,
|
||||
pooling_params=None,
|
||||
client_index=0,
|
||||
)
|
||||
scheduler = create_scheduler(use_ec_connector=True, ec_role="ec_producer")
|
||||
scheduler.add_request(request)
|
||||
request.status = RequestStatus.FINISHED_STOPPED
|
||||
|
||||
mock_ec = MagicMock()
|
||||
mock_ec.request_finished.return_value = (False, EC_PARAMS)
|
||||
scheduler.ec_connector = mock_ec
|
||||
|
||||
kv_params, ec_params = scheduler._free_request(request)
|
||||
|
||||
mock_ec.request_finished.assert_called_once_with(request)
|
||||
assert ec_params == EC_PARAMS
|
||||
assert kv_params is None
|
||||
|
||||
|
||||
def test_free_request_without_ec_connector_returns_none():
|
||||
"""When no EC connector is configured, ec_transfer_params must be None."""
|
||||
sp = SamplingParams(max_tokens=1)
|
||||
sp.update_from_generation_config({}, 50256)
|
||||
request = Request(
|
||||
request_id="test-req",
|
||||
prompt_token_ids=[1, 2, 3],
|
||||
sampling_params=sp,
|
||||
pooling_params=None,
|
||||
client_index=0,
|
||||
)
|
||||
scheduler = create_scheduler(use_ec_connector=True, ec_role="ec_producer")
|
||||
scheduler.add_request(request)
|
||||
request.status = RequestStatus.FINISHED_STOPPED
|
||||
|
||||
kv_params, ec_params = scheduler._free_request(request)
|
||||
|
||||
assert ec_params is None
|
||||
assert kv_params is None
|
||||
@@ -137,6 +137,12 @@ class AnthropicMessagesRequest(BaseModel):
|
||||
default=None,
|
||||
description="KVTransfer parameters used for disaggregated serving.",
|
||||
)
|
||||
ec_transfer_params: dict[str, Any] | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"ECTransfer parameters used for encoder-cache disaggregated serving."
|
||||
),
|
||||
)
|
||||
chat_template_kwargs: dict[str, Any] | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
@@ -218,6 +224,9 @@ class AnthropicMessagesResponse(BaseModel):
|
||||
kv_transfer_params: dict[str, Any] | None = Field(
|
||||
default=None, description="KVTransfer parameters."
|
||||
)
|
||||
ec_transfer_params: dict[str, Any] | None = Field(
|
||||
default=None, description="ECTransfer parameters."
|
||||
)
|
||||
|
||||
def model_post_init(self, __context):
|
||||
if not self.id:
|
||||
|
||||
@@ -491,6 +491,7 @@ class AnthropicServingMessages(OpenAIServingChat):
|
||||
top_p=anthropic_request.top_p,
|
||||
top_k=anthropic_request.top_k,
|
||||
kv_transfer_params=anthropic_request.kv_transfer_params,
|
||||
ec_transfer_params=anthropic_request.ec_transfer_params,
|
||||
chat_template_kwargs=anthropic_request.chat_template_kwargs,
|
||||
)
|
||||
|
||||
@@ -630,6 +631,7 @@ class AnthropicServingMessages(OpenAIServingChat):
|
||||
generator.usage,
|
||||
),
|
||||
kv_transfer_params=generator.kv_transfer_params,
|
||||
ec_transfer_params=generator.ec_transfer_params,
|
||||
)
|
||||
choice = generator.choices[0]
|
||||
if choice.finish_reason == "stop":
|
||||
|
||||
@@ -134,6 +134,9 @@ class ChatCompletionResponse(OpenAIBaseModel):
|
||||
kv_transfer_params: dict[str, Any] | None = Field(
|
||||
default=None, description="KVTransfer parameters."
|
||||
)
|
||||
ec_transfer_params: dict[str, Any] | None = Field(
|
||||
default=None, description="ECTransfer parameters."
|
||||
)
|
||||
metrics: PerRequestTimingMetrics | None = None
|
||||
|
||||
|
||||
@@ -439,6 +442,13 @@ class ChatCompletionRequest(OpenAIBaseModel):
|
||||
description="KVTransfer parameters used for disaggregated serving.",
|
||||
)
|
||||
|
||||
ec_transfer_params: dict[str, Any] | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"ECTransfer parameters used for encoder-cache disaggregated serving."
|
||||
),
|
||||
)
|
||||
|
||||
vllm_xargs: dict[str, str | int | float | list[str | int | float]] | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
@@ -665,6 +675,9 @@ class ChatCompletionRequest(OpenAIBaseModel):
|
||||
if self.kv_transfer_params:
|
||||
# Pass in kv_transfer_params via extra_args
|
||||
extra_args["kv_transfer_params"] = self.kv_transfer_params
|
||||
if self.ec_transfer_params:
|
||||
# Pass in ec_transfer_params via extra_args
|
||||
extra_args["ec_transfer_params"] = self.ec_transfer_params
|
||||
return SamplingParams.from_optional(
|
||||
n=self.n,
|
||||
presence_penalty=self.presence_penalty,
|
||||
|
||||
@@ -1057,6 +1057,7 @@ class OpenAIServingChat(GenerateBaseServing):
|
||||
),
|
||||
prompt_text=prompt_text,
|
||||
kv_transfer_params=final_res.kv_transfer_params,
|
||||
ec_transfer_params=final_res.ec_transfer_params,
|
||||
metrics=per_request_metrics,
|
||||
)
|
||||
|
||||
|
||||
@@ -189,6 +189,13 @@ class CompletionRequest(OpenAIBaseModel):
|
||||
description="KVTransfer parameters used for disaggregated serving.",
|
||||
)
|
||||
|
||||
ec_transfer_params: dict[str, Any] | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"ECTransfer parameters used for encoder-cache disaggregated serving."
|
||||
),
|
||||
)
|
||||
|
||||
vllm_xargs: dict[str, str | int | float] | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
@@ -346,6 +353,9 @@ class CompletionRequest(OpenAIBaseModel):
|
||||
if self.kv_transfer_params:
|
||||
# Pass in kv_transfer_params via extra_args
|
||||
extra_args["kv_transfer_params"] = self.kv_transfer_params
|
||||
if self.ec_transfer_params:
|
||||
# Pass in ec_transfer_params via extra_args
|
||||
extra_args["ec_transfer_params"] = self.ec_transfer_params
|
||||
return SamplingParams.from_optional(
|
||||
n=self.n,
|
||||
presence_penalty=self.presence_penalty,
|
||||
@@ -595,6 +605,9 @@ class CompletionResponse(OpenAIBaseModel):
|
||||
kv_transfer_params: dict[str, Any] | None = Field(
|
||||
default=None, description="KVTransfer parameters."
|
||||
)
|
||||
ec_transfer_params: dict[str, Any] | None = Field(
|
||||
default=None, description="ECTransfer parameters."
|
||||
)
|
||||
metrics: PerRequestTimingMetrics | None = None
|
||||
|
||||
|
||||
|
||||
@@ -510,6 +510,7 @@ class OpenAIServingCompletion(GenerateBaseServing):
|
||||
num_prompt_tokens = 0
|
||||
num_generated_tokens = 0
|
||||
kv_transfer_params = None
|
||||
ec_transfer_params = None
|
||||
last_final_res = None
|
||||
for final_res in final_res_batch:
|
||||
last_final_res = final_res
|
||||
@@ -632,6 +633,8 @@ class OpenAIServingCompletion(GenerateBaseServing):
|
||||
|
||||
if final_res_batch:
|
||||
kv_transfer_params = final_res_batch[0].kv_transfer_params
|
||||
ec_transfer_params = final_res_batch[0].ec_transfer_params
|
||||
|
||||
return CompletionResponse(
|
||||
id=request_id,
|
||||
created=created_time,
|
||||
@@ -640,6 +643,7 @@ class OpenAIServingCompletion(GenerateBaseServing):
|
||||
usage=usage,
|
||||
system_fingerprint=self.system_fingerprint,
|
||||
kv_transfer_params=kv_transfer_params,
|
||||
ec_transfer_params=ec_transfer_params,
|
||||
metrics=per_request_metrics,
|
||||
)
|
||||
|
||||
|
||||
@@ -200,6 +200,7 @@ class SimpleContext(ConversationContext):
|
||||
|
||||
self.input_messages: list[ResponseRawMessageAndToken] = []
|
||||
self.kv_transfer_params: dict[str, Any] | None = None
|
||||
self.ec_transfer_params: dict[str, Any] | None = None
|
||||
|
||||
def append_output(self, output) -> None:
|
||||
self.last_output = output
|
||||
@@ -210,6 +211,8 @@ class SimpleContext(ConversationContext):
|
||||
self.num_output_tokens += len(output.outputs[0].token_ids or [])
|
||||
if output.kv_transfer_params is not None:
|
||||
self.kv_transfer_params = output.kv_transfer_params
|
||||
if output.ec_transfer_params is not None:
|
||||
self.ec_transfer_params = output.ec_transfer_params
|
||||
|
||||
# Accumulate text, token_ids, and logprobs for streaming mode
|
||||
delta_output = output.outputs[0]
|
||||
@@ -328,6 +331,7 @@ class ParsableContext(ConversationContext):
|
||||
self.output_messages: list[ResponseRawMessageAndToken] = []
|
||||
self._accumulated_token_ids: list[int] = []
|
||||
self.kv_transfer_params: dict[str, Any] | None = None
|
||||
self.ec_transfer_params: dict[str, Any] | None = None
|
||||
|
||||
def append_output(self, output: RequestOutput) -> None:
|
||||
self.num_prompt_tokens = len(output.prompt_token_ids or [])
|
||||
@@ -336,6 +340,9 @@ class ParsableContext(ConversationContext):
|
||||
if output.kv_transfer_params is not None:
|
||||
self.kv_transfer_params = output.kv_transfer_params
|
||||
|
||||
if output.ec_transfer_params is not None:
|
||||
self.ec_transfer_params = output.ec_transfer_params
|
||||
|
||||
completion = output.outputs[0]
|
||||
self.finish_reason = completion.finish_reason
|
||||
|
||||
@@ -630,6 +637,7 @@ class HarmonyContext(ConversationContext):
|
||||
self.is_first_turn = True
|
||||
self.first_tok_of_message = True
|
||||
self.kv_transfer_params: dict[str, Any] | None = None
|
||||
self.ec_transfer_params: dict[str, Any] | None = None
|
||||
|
||||
def append_output(self, output: RequestOutput) -> None:
|
||||
if self.first_tok_of_message:
|
||||
@@ -645,6 +653,8 @@ class HarmonyContext(ConversationContext):
|
||||
self._update_decode_token_usage(output)
|
||||
if output.kv_transfer_params is not None:
|
||||
self.kv_transfer_params = output.kv_transfer_params
|
||||
if output.ec_transfer_params is not None:
|
||||
self.ec_transfer_params = output.ec_transfer_params
|
||||
|
||||
if output.finished:
|
||||
self.finish_reason = output.outputs[0].finish_reason
|
||||
|
||||
@@ -285,6 +285,12 @@ class ResponsesRequest(OpenAIBaseModel):
|
||||
default=None,
|
||||
description="KVTransfer parameters used for disaggregated serving.",
|
||||
)
|
||||
ec_transfer_params: dict[str, Any] | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"ECTransfer parameters used for encoder-cache disaggregated serving."
|
||||
),
|
||||
)
|
||||
chat_template_kwargs: dict[str, Any] | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
@@ -409,6 +415,8 @@ class ResponsesRequest(OpenAIBaseModel):
|
||||
extra_args: dict[str, Any] = self.vllm_xargs if self.vllm_xargs else {}
|
||||
if self.kv_transfer_params:
|
||||
extra_args["kv_transfer_params"] = self.kv_transfer_params
|
||||
if self.ec_transfer_params:
|
||||
extra_args["ec_transfer_params"] = self.ec_transfer_params
|
||||
|
||||
return SamplingParams.from_optional(
|
||||
temperature=temperature,
|
||||
@@ -675,6 +683,9 @@ class ResponsesResponse(OpenAIBaseModel):
|
||||
kv_transfer_params: dict[str, Any] | None = Field(
|
||||
default=None, description="KVTransfer parameters."
|
||||
)
|
||||
ec_transfer_params: dict[str, Any] | None = Field(
|
||||
default=None, description="ECTransfer parameters."
|
||||
)
|
||||
|
||||
# --8<-- [start:responses-response-extra-params]
|
||||
# These are populated when enable_response_messages is set to True
|
||||
@@ -720,6 +731,7 @@ class ResponsesResponse(OpenAIBaseModel):
|
||||
input_messages: ResponseInputOutputMessage | None = None,
|
||||
output_messages: ResponseInputOutputMessage | None = None,
|
||||
kv_transfer_params: dict[str, Any] | None = None,
|
||||
ec_transfer_params: dict[str, Any] | None = None,
|
||||
) -> "ResponsesResponse":
|
||||
incomplete_details: IncompleteDetails | None = None
|
||||
if status == "incomplete":
|
||||
@@ -758,6 +770,7 @@ class ResponsesResponse(OpenAIBaseModel):
|
||||
user=request.user,
|
||||
usage=usage,
|
||||
kv_transfer_params=kv_transfer_params,
|
||||
ec_transfer_params=ec_transfer_params,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -932,6 +932,7 @@ class OpenAIServingResponses(GenerateBaseServing):
|
||||
status=status,
|
||||
usage=usage,
|
||||
kv_transfer_params=context.kv_transfer_params,
|
||||
ec_transfer_params=context.ec_transfer_params,
|
||||
)
|
||||
|
||||
if request.store:
|
||||
|
||||
@@ -132,6 +132,12 @@ class GenerateRequest(BaseModel):
|
||||
default=None,
|
||||
description="KVTransfer parameters used for disaggregated serving.",
|
||||
)
|
||||
ec_transfer_params: dict[str, Any] | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"ECTransfer parameters used for encoder-cache disaggregated serving."
|
||||
),
|
||||
)
|
||||
|
||||
# Tracks which keys the caller explicitly set inside ``sampling_params``
|
||||
# when the request was parsed from a JSON body. Lets the server tell
|
||||
@@ -238,6 +244,12 @@ class GenerateResponse(BaseModel):
|
||||
default=None,
|
||||
description="KVTransfer parameters used for disaggregated serving.",
|
||||
)
|
||||
ec_transfer_params: dict[str, Any] | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"ECTransfer parameters used for encoder-cache disaggregated serving."
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
####### Derender (postprocessing) #######
|
||||
|
||||
@@ -330,6 +330,7 @@ class ServingTokens(GenerateBaseServing):
|
||||
usage=usage,
|
||||
prompt_logprobs=clamp_prompt_logprobs(final_res.prompt_logprobs),
|
||||
kv_transfer_params=final_res.kv_transfer_params,
|
||||
ec_transfer_params=final_res.ec_transfer_params,
|
||||
)
|
||||
|
||||
# Log complete response if output logging is enabled
|
||||
|
||||
@@ -207,6 +207,8 @@ if TYPE_CHECKING:
|
||||
VLLM_DISABLE_REQUEST_ID_RANDOMIZATION: bool = False
|
||||
VLLM_NIXL_SIDE_CHANNEL_HOST: str = "localhost"
|
||||
VLLM_NIXL_SIDE_CHANNEL_PORT: int = 5600
|
||||
VLLM_EC_SIDE_CHANNEL_HOST: str = "localhost"
|
||||
VLLM_EC_SIDE_CHANNEL_PORT: int = 5601
|
||||
VLLM_MOONCAKE_BOOTSTRAP_PORT: int = 8998
|
||||
VLLM_MOONCAKE_STORE_TIER_LOG: bool = False
|
||||
VLLM_MOONCAKE_LOAD_RECV_THREADS: int = 1
|
||||
@@ -1565,6 +1567,16 @@ environment_variables: dict[str, Callable[[], Any]] = {
|
||||
"VLLM_NIXL_SIDE_CHANNEL_PORT": lambda: int(
|
||||
os.getenv("VLLM_NIXL_SIDE_CHANNEL_PORT", "5600")
|
||||
),
|
||||
# IP address used for the EC connector's ZMQ side channel
|
||||
# (producer ROUTER bind, consumer DEALER dial).
|
||||
"VLLM_EC_SIDE_CHANNEL_HOST": lambda: os.getenv(
|
||||
"VLLM_EC_SIDE_CHANNEL_HOST", "localhost"
|
||||
),
|
||||
# Port for the EC connector's ZMQ side channel; advertised to peers
|
||||
# via `ec_transfer_params.peer_port` on the producer's response.
|
||||
"VLLM_EC_SIDE_CHANNEL_PORT": lambda: int(
|
||||
os.getenv("VLLM_EC_SIDE_CHANNEL_PORT", "5601")
|
||||
),
|
||||
# Port used for Mooncake handshake between remote agents.
|
||||
"VLLM_MOONCAKE_BOOTSTRAP_PORT": lambda: int(
|
||||
os.getenv("VLLM_MOONCAKE_BOOTSTRAP_PORT", "8998")
|
||||
|
||||
@@ -104,6 +104,7 @@ class RequestOutput:
|
||||
None if decoder-only.
|
||||
num_cached_tokens: The number of tokens with prefix cache hit.
|
||||
kv_transfer_params: The params for remote K/V transfer.
|
||||
ec_transfer_params: The params for remote encoder-cache transfer.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
@@ -121,6 +122,7 @@ class RequestOutput:
|
||||
num_cached_tokens: int | None = None,
|
||||
*,
|
||||
kv_transfer_params: dict[str, Any] | None = None,
|
||||
ec_transfer_params: dict[str, Any] | None = None,
|
||||
# Forward compatibility, code that uses args added in new release can
|
||||
# still run with older versions of vLLM without breaking.
|
||||
**kwargs: Any,
|
||||
@@ -141,12 +143,14 @@ class RequestOutput:
|
||||
self.encoder_prompt_token_ids = encoder_prompt_token_ids
|
||||
self.num_cached_tokens = num_cached_tokens
|
||||
self.kv_transfer_params = kv_transfer_params
|
||||
self.ec_transfer_params = ec_transfer_params
|
||||
|
||||
def add(self, next_output: "RequestOutput", aggregate: bool) -> None:
|
||||
"""Merge subsequent RequestOutput into this one"""
|
||||
|
||||
self.finished |= next_output.finished
|
||||
self.kv_transfer_params = next_output.kv_transfer_params
|
||||
self.ec_transfer_params = next_output.ec_transfer_params
|
||||
|
||||
for next_completion in next_output.outputs:
|
||||
for i, completion in enumerate(self.outputs):
|
||||
|
||||
@@ -9,6 +9,7 @@ from vllm.multimodal import MULTIMODAL_REGISTRY, MultiModalRegistry
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from vllm.config import VllmConfig
|
||||
from vllm.distributed.ec_transfer.ec_connector.base import ECConnectorBase
|
||||
from vllm.distributed.kv_transfer.kv_connector.v1 import KVConnectorBase_V1
|
||||
from vllm.v1.core.sched.output import GrammarOutput, SchedulerOutput
|
||||
from vllm.v1.engine import EngineCoreOutputs
|
||||
@@ -248,3 +249,6 @@ class SchedulerInterface(ABC):
|
||||
|
||||
def get_kv_connector(self) -> "KVConnectorBase_V1 | None":
|
||||
return None
|
||||
|
||||
def get_ec_connector(self) -> "ECConnectorBase | None":
|
||||
return None
|
||||
|
||||
@@ -10,6 +10,7 @@ from typing import Any
|
||||
from vllm.compilation.cuda_graph import CUDAGraphStat
|
||||
from vllm.config import VllmConfig
|
||||
from vllm.distributed.ec_transfer.ec_connector.base import (
|
||||
ECConnectorBase,
|
||||
ECConnectorMetadata,
|
||||
ECConnectorRole,
|
||||
)
|
||||
@@ -1675,6 +1676,7 @@ class Scheduler(SchedulerInterface):
|
||||
new_token_ids = generated_token_ids
|
||||
pooler_output = pooler_outputs[req_index] if pooler_outputs else None
|
||||
kv_transfer_params = None
|
||||
ec_transfer_params = None
|
||||
status_before_stop = request.status
|
||||
num_output_tokens_before = len(request._output_token_ids)
|
||||
|
||||
@@ -1761,7 +1763,7 @@ class Scheduler(SchedulerInterface):
|
||||
finish_reason = request.get_finished_reason()
|
||||
finished = self._handle_stopped_request(request)
|
||||
if finished:
|
||||
kv_transfer_params = self._free_request(request)
|
||||
kv_transfer_params, ec_transfer_params = self._free_request(request)
|
||||
|
||||
if status_before_stop == RequestStatus.RUNNING:
|
||||
stopped_running_reqs.add(request)
|
||||
@@ -1785,6 +1787,7 @@ class Scheduler(SchedulerInterface):
|
||||
new_token_ids
|
||||
or pooler_output is not None
|
||||
or kv_transfer_params
|
||||
or ec_transfer_params
|
||||
or stopped
|
||||
):
|
||||
# Add EngineCoreOutput for this Request.
|
||||
@@ -1800,6 +1803,7 @@ class Scheduler(SchedulerInterface):
|
||||
events=request.take_events(),
|
||||
prefill_stats=request.take_prefill_stats(),
|
||||
kv_transfer_params=kv_transfer_params,
|
||||
ec_transfer_params=ec_transfer_params,
|
||||
trace_headers=request.trace_headers,
|
||||
routed_experts=routed_experts,
|
||||
num_nans_in_logits=request.num_nans_in_logits,
|
||||
@@ -2151,11 +2155,21 @@ class Scheduler(SchedulerInterface):
|
||||
|
||||
def _free_request(
|
||||
self, request: Request, delay_free_blocks: bool = False
|
||||
) -> dict[str, Any] | None:
|
||||
) -> tuple[dict[str, Any] | None, dict[str, Any] | None]:
|
||||
assert request.is_finished()
|
||||
|
||||
self._inflight_prefills.discard(request)
|
||||
connector_delay_free_blocks, kv_xfer_params = self._connector_finished(request)
|
||||
|
||||
# EC Connector: mirror the KV hook. The contract requires firing
|
||||
# before the encoder cache is freed so the connector can inspect
|
||||
# per-request state (e.g. which mm_hashes it recorded during
|
||||
# save_caches()) and emit ec_transfer_params for the response body.
|
||||
ec_xfer_params: dict[str, Any] | None = None
|
||||
if self.ec_connector is not None:
|
||||
ec_delay_free, ec_xfer_params = self.ec_connector.request_finished(request)
|
||||
connector_delay_free_blocks |= ec_delay_free
|
||||
|
||||
self.encoder_cache_manager.free(request)
|
||||
request_id = request.request_id
|
||||
self.finished_req_ids.add(request_id)
|
||||
@@ -2166,7 +2180,7 @@ class Scheduler(SchedulerInterface):
|
||||
if not delay_free_blocks:
|
||||
self._free_blocks(request)
|
||||
|
||||
return kv_xfer_params
|
||||
return kv_xfer_params, ec_xfer_params
|
||||
|
||||
def _free_blocks(self, request: Request):
|
||||
assert request.is_finished()
|
||||
@@ -2414,6 +2428,9 @@ class Scheduler(SchedulerInterface):
|
||||
def get_kv_connector(self) -> KVConnectorBase_V1 | None:
|
||||
return self.connector
|
||||
|
||||
def get_ec_connector(self) -> ECConnectorBase | None:
|
||||
return self.ec_connector
|
||||
|
||||
def _connector_finished(
|
||||
self, request: Request
|
||||
) -> tuple[bool, dict[str, Any] | None]:
|
||||
|
||||
@@ -190,6 +190,7 @@ class EngineCoreOutput(
|
||||
stop_reason: int | str | None = None
|
||||
events: list[EngineCoreEvent] | None = None
|
||||
kv_transfer_params: dict[str, Any] | None = None
|
||||
ec_transfer_params: dict[str, Any] | None = None
|
||||
|
||||
trace_headers: Mapping[str, str] | None = None
|
||||
|
||||
|
||||
@@ -400,6 +400,15 @@ class EngineCore:
|
||||
"Disabling KVTransfer for this request."
|
||||
)
|
||||
|
||||
if (
|
||||
request.ec_transfer_params is not None
|
||||
and self.scheduler.get_ec_connector() is None
|
||||
):
|
||||
logger.warning(
|
||||
"Got ec_transfer_params, but no ECConnector found. "
|
||||
"Disabling ECTransfer for this request."
|
||||
)
|
||||
|
||||
self.scheduler.add_request(request)
|
||||
if request.abort_immediately:
|
||||
# Immediately abort so the connector's request_finished hook runs
|
||||
|
||||
@@ -276,6 +276,7 @@ class RequestState:
|
||||
finish_reason: FinishReason | None,
|
||||
stop_reason: int | str | None,
|
||||
kv_transfer_params: dict[str, Any] | None = None,
|
||||
ec_transfer_params: dict[str, Any] | None = None,
|
||||
) -> RequestOutput | PoolingRequestOutput | None:
|
||||
finished = finish_reason is not None
|
||||
final_only = self.output_kind == RequestOutputKind.FINAL_ONLY
|
||||
@@ -327,7 +328,11 @@ class RequestState:
|
||||
external_req_id = self.parent_req.external_req_id
|
||||
|
||||
return self._new_request_output(
|
||||
external_req_id, outputs, finished, kv_transfer_params
|
||||
external_req_id,
|
||||
outputs,
|
||||
finished,
|
||||
kv_transfer_params,
|
||||
ec_transfer_params,
|
||||
)
|
||||
|
||||
def _new_request_output(
|
||||
@@ -336,6 +341,7 @@ class RequestState:
|
||||
outputs: list[CompletionOutput] | list[PoolingOutput],
|
||||
finished: bool,
|
||||
kv_transfer_params: dict[str, Any] | None = None,
|
||||
ec_transfer_params: dict[str, Any] | None = None,
|
||||
) -> RequestOutput | PoolingRequestOutput:
|
||||
# If prompt embeds were used, put placeholder prompt token ids
|
||||
prompt_token_ids = self.prompt_token_ids
|
||||
@@ -369,6 +375,7 @@ class RequestState:
|
||||
outputs=cast(list[CompletionOutput], outputs),
|
||||
finished=finished,
|
||||
kv_transfer_params=kv_transfer_params,
|
||||
ec_transfer_params=ec_transfer_params,
|
||||
num_cached_tokens=self.num_cached_tokens,
|
||||
metrics=self.stats,
|
||||
)
|
||||
@@ -497,6 +504,7 @@ class OutputProcessor:
|
||||
finish_reason=FinishReason.ABORT,
|
||||
stop_reason=None,
|
||||
kv_transfer_params=None,
|
||||
ec_transfer_params=None,
|
||||
)
|
||||
):
|
||||
req_state.queue.put(request_output)
|
||||
@@ -620,6 +628,7 @@ class OutputProcessor:
|
||||
finish_reason = engine_core_output.finish_reason
|
||||
stop_reason = engine_core_output.stop_reason
|
||||
kv_transfer_params = engine_core_output.kv_transfer_params
|
||||
ec_transfer_params = engine_core_output.ec_transfer_params
|
||||
if engine_core_output.routed_experts is not None:
|
||||
req_state.routed_experts_chunks.append(
|
||||
engine_core_output.routed_experts
|
||||
@@ -654,6 +663,7 @@ class OutputProcessor:
|
||||
finish_reason,
|
||||
stop_reason,
|
||||
kv_transfer_params,
|
||||
ec_transfer_params,
|
||||
):
|
||||
if req_state.streaming_input:
|
||||
request_output.finished = False
|
||||
|
||||
@@ -100,6 +100,8 @@ class Request:
|
||||
|
||||
# P/D: Connector-specific KV transfer parameters.
|
||||
self.kv_transfer_params: dict[str, Any] | None = None
|
||||
# E/P/D: Connector-specific encoder-cache transfer parameters.
|
||||
self.ec_transfer_params: dict[str, Any] | None = None
|
||||
|
||||
if pooling_params is not None:
|
||||
# Pooling models.
|
||||
@@ -115,6 +117,9 @@ class Request:
|
||||
self.kv_transfer_params = sampling_params.extra_args.get(
|
||||
"kv_transfer_params"
|
||||
)
|
||||
self.ec_transfer_params = sampling_params.extra_args.get(
|
||||
"ec_transfer_params"
|
||||
)
|
||||
self.kv_cache_report_mode = sampling_params.extra_args.get(
|
||||
"kv_cache_report_mode", "incremental"
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user