feat(grpc): support disaggregated generation requests (#30440)
Signed-off-by: Connor Carpenter <connorc@nvidia.com> Co-authored-by: Ishan Dhanani <ishandhanani@gmail.com>
This commit is contained in:
co-authored by
Ishan Dhanani
parent
ca8f15cd70
commit
cc7d6659fd
@@ -34,6 +34,20 @@ service SglangService {
|
|||||||
rpc UpdateWeightsFromDisk(UpdateWeightsRequest) returns (UpdateWeightsResponse);
|
rpc UpdateWeightsFromDisk(UpdateWeightsRequest) returns (UpdateWeightsResponse);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Disaggregated prefill/decode rendezvous parameters.
|
||||||
|
//
|
||||||
|
// Carried on Generate/TextGenerate requests when the worker is part of a
|
||||||
|
// disaggregated deployment. The prefill leg writes its KV cache to the
|
||||||
|
// bootstrap server identified by (bootstrap_host, bootstrap_port) under
|
||||||
|
// the room id; the decode leg fetches it from the same room. Both legs
|
||||||
|
// receive the same `bootstrap_room` from the router so they rendezvous
|
||||||
|
// on the KV transfer.
|
||||||
|
message DisaggregatedParams {
|
||||||
|
string bootstrap_host = 1;
|
||||||
|
int32 bootstrap_port = 2;
|
||||||
|
int64 bootstrap_room = 3;
|
||||||
|
}
|
||||||
|
|
||||||
// Sampling parameters shared across text and tokenized RPCs.
|
// Sampling parameters shared across text and tokenized RPCs.
|
||||||
message SamplingParams {
|
message SamplingParams {
|
||||||
optional float temperature = 1;
|
optional float temperature = 1;
|
||||||
@@ -69,6 +83,7 @@ message TextGenerateRequest {
|
|||||||
optional int32 routed_dp_rank = 11;
|
optional int32 routed_dp_rank = 11;
|
||||||
map<string, string> trace_headers = 12;
|
map<string, string> trace_headers = 12;
|
||||||
optional string session_id = 13;
|
optional string session_id = 13;
|
||||||
|
optional DisaggregatedParams disaggregated_params = 14;
|
||||||
}
|
}
|
||||||
|
|
||||||
message TextGenerateResponse {
|
message TextGenerateResponse {
|
||||||
@@ -92,6 +107,7 @@ message GenerateRequest {
|
|||||||
optional int32 routed_dp_rank = 10;
|
optional int32 routed_dp_rank = 10;
|
||||||
map<string, string> trace_headers = 11;
|
map<string, string> trace_headers = 11;
|
||||||
optional string session_id = 12;
|
optional string session_id = 12;
|
||||||
|
optional DisaggregatedParams disaggregated_params = 13;
|
||||||
}
|
}
|
||||||
|
|
||||||
message GenerateResponse {
|
message GenerateResponse {
|
||||||
|
|||||||
@@ -66,6 +66,26 @@ fn trace_headers_to_json(headers: &HashMap<String, String>) -> Option<serde_json
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn insert_disaggregated_params(
|
||||||
|
request: &mut HashMap<String, serde_json::Value>,
|
||||||
|
params: &Option<proto::DisaggregatedParams>,
|
||||||
|
) {
|
||||||
|
if let Some(params) = params {
|
||||||
|
request.insert(
|
||||||
|
"bootstrap_host".into(),
|
||||||
|
serde_json::json!(params.bootstrap_host),
|
||||||
|
);
|
||||||
|
request.insert(
|
||||||
|
"bootstrap_port".into(),
|
||||||
|
serde_json::json!(params.bootstrap_port),
|
||||||
|
);
|
||||||
|
request.insert(
|
||||||
|
"bootstrap_room".into(),
|
||||||
|
serde_json::json!(params.bootstrap_room),
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
fn now_timestamp() -> f64 {
|
fn now_timestamp() -> f64 {
|
||||||
std::time::SystemTime::now()
|
std::time::SystemTime::now()
|
||||||
.duration_since(std::time::UNIX_EPOCH)
|
.duration_since(std::time::UNIX_EPOCH)
|
||||||
@@ -131,6 +151,7 @@ pub(crate) fn build_text_generate_dict(
|
|||||||
if let Some(ref session_id) = req.session_id {
|
if let Some(ref session_id) = req.session_id {
|
||||||
d.insert("session_id".into(), serde_json::json!(session_id));
|
d.insert("session_id".into(), serde_json::json!(session_id));
|
||||||
}
|
}
|
||||||
|
insert_disaggregated_params(&mut d, &req.disaggregated_params);
|
||||||
if let Some(trace) = trace_headers_to_json(&req.trace_headers) {
|
if let Some(trace) = trace_headers_to_json(&req.trace_headers) {
|
||||||
d.insert("external_trace_header".into(), trace);
|
d.insert("external_trace_header".into(), trace);
|
||||||
}
|
}
|
||||||
@@ -178,6 +199,7 @@ pub(crate) fn build_generate_dict(
|
|||||||
if let Some(ref session_id) = req.session_id {
|
if let Some(ref session_id) = req.session_id {
|
||||||
d.insert("session_id".into(), serde_json::json!(session_id));
|
d.insert("session_id".into(), serde_json::json!(session_id));
|
||||||
}
|
}
|
||||||
|
insert_disaggregated_params(&mut d, &req.disaggregated_params);
|
||||||
if let Some(trace) = trace_headers_to_json(&req.trace_headers) {
|
if let Some(trace) = trace_headers_to_json(&req.trace_headers) {
|
||||||
d.insert("external_trace_header".into(), trace);
|
d.insert("external_trace_header".into(), trace);
|
||||||
}
|
}
|
||||||
@@ -269,4 +291,52 @@ mod tests {
|
|||||||
Some(&serde_json::json!("session-1"))
|
Some(&serde_json::json!("session-1"))
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn generate_dicts_include_disaggregated_params() {
|
||||||
|
let disaggregated_params = Some(proto::DisaggregatedParams {
|
||||||
|
bootstrap_host: "10.0.0.1".to_string(),
|
||||||
|
bootstrap_port: 8998,
|
||||||
|
bootstrap_room: i64::MAX,
|
||||||
|
});
|
||||||
|
let text_req = proto::TextGenerateRequest {
|
||||||
|
disaggregated_params: disaggregated_params.clone(),
|
||||||
|
..Default::default()
|
||||||
|
};
|
||||||
|
let token_req = proto::GenerateRequest {
|
||||||
|
disaggregated_params,
|
||||||
|
..Default::default()
|
||||||
|
};
|
||||||
|
|
||||||
|
for request in [
|
||||||
|
build_text_generate_dict("request-1", &text_req),
|
||||||
|
build_generate_dict("request-2", &token_req),
|
||||||
|
] {
|
||||||
|
assert_eq!(
|
||||||
|
request.get("bootstrap_host"),
|
||||||
|
Some(&serde_json::json!("10.0.0.1"))
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
request.get("bootstrap_port"),
|
||||||
|
Some(&serde_json::json!(8998))
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
request.get("bootstrap_room"),
|
||||||
|
Some(&serde_json::json!(i64::MAX))
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn generate_dicts_omit_disaggregated_params_when_absent() {
|
||||||
|
let text_request =
|
||||||
|
build_text_generate_dict("request-1", &proto::TextGenerateRequest::default());
|
||||||
|
let token_request = build_generate_dict("request-2", &proto::GenerateRequest::default());
|
||||||
|
|
||||||
|
for request in [text_request, token_request] {
|
||||||
|
assert!(!request.contains_key("bootstrap_host"));
|
||||||
|
assert!(!request.contains_key("bootstrap_port"));
|
||||||
|
assert!(!request.contains_key("bootstrap_room"));
|
||||||
|
}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user