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);
|
||||
}
|
||||
|
||||
// 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.
|
||||
message SamplingParams {
|
||||
optional float temperature = 1;
|
||||
@@ -69,6 +83,7 @@ message TextGenerateRequest {
|
||||
optional int32 routed_dp_rank = 11;
|
||||
map<string, string> trace_headers = 12;
|
||||
optional string session_id = 13;
|
||||
optional DisaggregatedParams disaggregated_params = 14;
|
||||
}
|
||||
|
||||
message TextGenerateResponse {
|
||||
@@ -92,6 +107,7 @@ message GenerateRequest {
|
||||
optional int32 routed_dp_rank = 10;
|
||||
map<string, string> trace_headers = 11;
|
||||
optional string session_id = 12;
|
||||
optional DisaggregatedParams disaggregated_params = 13;
|
||||
}
|
||||
|
||||
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 {
|
||||
std::time::SystemTime::now()
|
||||
.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 {
|
||||
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) {
|
||||
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 {
|
||||
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) {
|
||||
d.insert("external_trace_header".into(), trace);
|
||||
}
|
||||
@@ -269,4 +291,52 @@ mod tests {
|
||||
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