Skip to content

Commit b2e12f9

Browse files
authored
Feat: merge routed experts for P/D (#29)
* feat: recover routed experts in vllm pd router * fix: emit full routed experts from pd merge * feat: expose routed experts cache clear route * fix: keep routed experts native shaped * fix: stitch native routed expert payloads * fix: satisfy clippy in routed experts merge * fix: keep routed experts merge stitch-only * fix: keep pd response shaping upstream * fix: merge routed experts in pd metadata pass * fix: replace decode prompt routed experts
1 parent 49780a9 commit b2e12f9

5 files changed

Lines changed: 381 additions & 40 deletions

File tree

Cargo.lock

Lines changed: 1 addition & 0 deletions
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

Cargo.toml

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -73,6 +73,7 @@ dotenvy = "0.15"
7373
tokenizers = { version = "0.22.2" }
7474
tiktoken-rs = { version = "0.7.0" }
7575
minijinja = { version = "2.0" }
76+
base64 = "0.22.1"
7677
rustls = { version = "0.23", default-features = false, features = [
7778
"ring",
7879
"std",

src/routers/http/mod.rs

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,7 @@ pub mod logprobs_merge;
55
pub mod openai_router;
66
pub mod pd_router;
77
pub mod pd_types;
8+
pub mod routed_experts_merge;
89
pub mod router;
910
pub(crate) mod usage_metrics;
1011
pub mod vllm_pd_router;
Lines changed: 297 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,297 @@
1+
//! Stitch compact routed-experts payloads for vLLM P/D disaggregation.
2+
3+
use base64::{engine::general_purpose::STANDARD, Engine as _};
4+
use serde_json::Value;
5+
6+
const NPY_MAGIC: &[u8; 6] = b"\x93NUMPY";
7+
const NPY_V1_HEADER_PREFIX_LEN: usize = 10;
8+
9+
#[derive(Clone, Debug, PartialEq, Eq)]
10+
struct RoutedExpertsPayload {
11+
seq_len: usize,
12+
layers: usize,
13+
topk: usize,
14+
descr: String,
15+
item_size: usize,
16+
data: Vec<u8>,
17+
}
18+
19+
impl RoutedExpertsPayload {
20+
fn suffix_rows(&self, row_start: usize) -> Result<Self, String> {
21+
let row_size = self.layers * self.topk * self.item_size;
22+
let byte_start = row_start * row_size;
23+
let data = self
24+
.data
25+
.get(byte_start..)
26+
.ok_or_else(|| {
27+
format!(
28+
"decode routed_experts has {} rows, expected at least {row_start}",
29+
self.seq_len
30+
)
31+
})?
32+
.to_vec();
33+
34+
Ok(Self {
35+
seq_len: self.seq_len - row_start,
36+
layers: self.layers,
37+
topk: self.topk,
38+
descr: self.descr.clone(),
39+
item_size: self.item_size,
40+
data,
41+
})
42+
}
43+
44+
fn concat_rows(&self, other: &Self) -> Result<Self, String> {
45+
if self.descr != other.descr || self.layers != other.layers || self.topk != other.topk {
46+
return Err(format!(
47+
"cannot concatenate routed_experts with shapes/dtypes ({}, {}, {}, {}) and ({}, {}, {}, {})",
48+
self.seq_len,
49+
self.layers,
50+
self.topk,
51+
self.descr,
52+
other.seq_len,
53+
other.layers,
54+
other.topk,
55+
other.descr,
56+
));
57+
}
58+
59+
let mut data = Vec::with_capacity(self.data.len() + other.data.len());
60+
data.extend_from_slice(&self.data);
61+
data.extend_from_slice(&other.data);
62+
63+
Ok(Self {
64+
seq_len: self.seq_len + other.seq_len,
65+
layers: self.layers,
66+
topk: self.topk,
67+
descr: self.descr.clone(),
68+
item_size: self.item_size,
69+
data,
70+
})
71+
}
72+
}
73+
74+
pub fn prefill_has_routed_experts(prefill_json: &Value) -> bool {
75+
prefill_choice_routed_experts(prefill_json).is_some()
76+
}
77+
78+
pub fn merge_routed_experts_in_json(
79+
prefill_json: &Value,
80+
decode_json: &mut Value,
81+
) -> Result<bool, String> {
82+
let prefill_routed = prefill_choice_routed_experts(prefill_json);
83+
if prefill_routed.is_none() && !decode_has_routed_experts(decode_json) {
84+
return Ok(false);
85+
}
86+
87+
let prompt = decode_routed_experts_value(
88+
prefill_routed.ok_or_else(|| {
89+
"decode response contained routed_experts, but prefill response did not".to_string()
90+
})?,
91+
"prefill routed_experts",
92+
)?;
93+
94+
let choices = decode_json["choices"]
95+
.as_array_mut()
96+
.ok_or_else(|| "decode response choices must be an array".to_string())?;
97+
98+
for choice in choices {
99+
let routed_experts = choice
100+
.get("routed_experts")
101+
.filter(|value| !value.is_null())
102+
.ok_or_else(|| "decode choice routed_experts is missing".to_string())?;
103+
let decode = decode_routed_experts_value(routed_experts, "decode routed_experts")?;
104+
let completion = decode.suffix_rows(prompt.seq_len)?;
105+
let merged = prompt.concat_rows(&completion)?;
106+
choice["routed_experts"] = Value::String(encode_routed_experts_payload(&merged)?);
107+
}
108+
109+
Ok(true)
110+
}
111+
112+
fn prefill_choice_routed_experts(prefill_json: &Value) -> Option<&Value> {
113+
prefill_json["choices"]
114+
.as_array()
115+
.and_then(|choices| choices.first())
116+
.and_then(|choice| choice.get("routed_experts"))
117+
.filter(|value| !value.is_null())
118+
}
119+
120+
fn decode_has_routed_experts(decode_json: &Value) -> bool {
121+
decode_json["choices"]
122+
.as_array()
123+
.map(|choices| {
124+
choices.iter().any(|choice| {
125+
choice
126+
.get("routed_experts")
127+
.filter(|value| !value.is_null())
128+
.is_some()
129+
})
130+
})
131+
.unwrap_or(false)
132+
}
133+
134+
fn decode_routed_experts_value(value: &Value, name: &str) -> Result<RoutedExpertsPayload, String> {
135+
let payload = value
136+
.as_str()
137+
.ok_or_else(|| format!("{name} must be a base64 .npy string"))?;
138+
let bytes = STANDARD
139+
.decode(payload)
140+
.map_err(|error| format!("{name} base64 decode failed: {error}"))?;
141+
parse_npy_payload(&bytes, name)
142+
}
143+
144+
fn parse_npy_payload(bytes: &[u8], name: &str) -> Result<RoutedExpertsPayload, String> {
145+
if bytes.len() < NPY_V1_HEADER_PREFIX_LEN || &bytes[..6] != NPY_MAGIC {
146+
return Err(format!("{name} is not a NumPy .npy payload"));
147+
}
148+
149+
let (header_len, data_start) = match bytes[6] {
150+
1 => (
151+
u16::from_le_bytes([bytes[8], bytes[9]]) as usize,
152+
NPY_V1_HEADER_PREFIX_LEN,
153+
),
154+
2 | 3 => (
155+
u32::from_le_bytes([bytes[8], bytes[9], bytes[10], bytes[11]]) as usize,
156+
12,
157+
),
158+
version => return Err(format!("{name} has unsupported .npy version {version}")),
159+
};
160+
161+
let header_end = data_start + header_len;
162+
let header = std::str::from_utf8(&bytes[data_start..header_end])
163+
.map_err(|error| format!("{name} header decode failed: {error}"))?;
164+
let descr = parse_descr(header, name)?;
165+
let item_size = dtype_item_size(&descr, name)?;
166+
let (seq_len, layers, topk) = parse_shape(header, name)?;
167+
let data = bytes[header_end..].to_vec();
168+
let expected_data_len = seq_len * layers * topk * item_size;
169+
if data.len() != expected_data_len {
170+
return Err(format!(
171+
"{name} has {} data bytes, expected {expected_data_len}",
172+
data.len()
173+
));
174+
}
175+
176+
Ok(RoutedExpertsPayload {
177+
seq_len,
178+
layers,
179+
topk,
180+
descr,
181+
item_size,
182+
data,
183+
})
184+
}
185+
186+
fn parse_descr(header: &str, name: &str) -> Result<String, String> {
187+
let after_key = header
188+
.split("'descr':")
189+
.nth(1)
190+
.or_else(|| header.split("\"descr\":").nth(1))
191+
.ok_or_else(|| format!("{name} header is missing descr"))?;
192+
let quote = after_key
193+
.find(['\'', '"'])
194+
.ok_or_else(|| format!("{name} descr is missing opening quote"))?;
195+
let after_quote = &after_key[quote + 1..];
196+
let end_quote = after_quote
197+
.find(['\'', '"'])
198+
.ok_or_else(|| format!("{name} descr is missing closing quote"))?;
199+
Ok(after_quote[..end_quote].to_string())
200+
}
201+
202+
fn dtype_item_size(descr: &str, name: &str) -> Result<usize, String> {
203+
match descr {
204+
"|u1" => Ok(1),
205+
"<i2" => Ok(2),
206+
"<i4" => Ok(4),
207+
_ => Err(format!("{name} has unsupported dtype {descr}")),
208+
}
209+
}
210+
211+
fn parse_shape(header: &str, name: &str) -> Result<(usize, usize, usize), String> {
212+
let shape_header = header
213+
.split("shape")
214+
.nth(1)
215+
.ok_or_else(|| format!("{name} header is missing shape"))?;
216+
let shape_start = shape_header
217+
.find('(')
218+
.ok_or_else(|| format!("{name} shape is missing '('"))?;
219+
let shape_end = shape_header[shape_start + 1..]
220+
.find(')')
221+
.ok_or_else(|| format!("{name} shape is missing ')'"))?
222+
+ shape_start
223+
+ 1;
224+
let dims = shape_header[shape_start + 1..shape_end]
225+
.split(',')
226+
.map(str::trim)
227+
.filter(|value| !value.is_empty())
228+
.map(|value| value.parse::<usize>().map_err(|error| error.to_string()))
229+
.collect::<Result<Vec<_>, _>>()
230+
.map_err(|error| format!("{name} shape parse failed: {error}"))?;
231+
232+
match dims.as_slice() {
233+
[seq_len, layers, topk] => Ok((*seq_len, *layers, *topk)),
234+
_ => Err(format!("{name} must have shape (seq, layers, topk)")),
235+
}
236+
}
237+
238+
fn encode_routed_experts_payload(payload: &RoutedExpertsPayload) -> Result<String, String> {
239+
let header_body = format!(
240+
"{{'descr': '{}', 'fortran_order': False, 'shape': ({}, {}, {}), }}",
241+
payload.descr, payload.seq_len, payload.layers, payload.topk
242+
);
243+
let mut header = header_body.into_bytes();
244+
let padding = (16 - ((NPY_V1_HEADER_PREFIX_LEN + header.len() + 1) % 16)) % 16;
245+
header.extend(std::iter::repeat_n(b' ', padding));
246+
header.push(b'\n');
247+
248+
let header_len = u16::try_from(header.len())
249+
.map_err(|_| "routed_experts NumPy header is too large for v1 .npy".to_string())?;
250+
let mut bytes =
251+
Vec::with_capacity(NPY_V1_HEADER_PREFIX_LEN + header.len() + payload.data.len());
252+
bytes.extend_from_slice(NPY_MAGIC);
253+
bytes.push(1);
254+
bytes.push(0);
255+
bytes.extend_from_slice(&header_len.to_le_bytes());
256+
bytes.extend_from_slice(&header);
257+
bytes.extend_from_slice(&payload.data);
258+
Ok(STANDARD.encode(bytes))
259+
}
260+
261+
#[cfg(test)]
262+
mod tests {
263+
use super::*;
264+
use serde_json::json;
265+
266+
fn uint8_payload(seq_len: usize, layers: usize, topk: usize, data: &[u8]) -> String {
267+
let payload = RoutedExpertsPayload {
268+
seq_len,
269+
layers,
270+
topk,
271+
descr: "|u1".to_string(),
272+
item_size: 1,
273+
data: data.to_vec(),
274+
};
275+
encode_routed_experts_payload(&payload).unwrap()
276+
}
277+
278+
#[test]
279+
fn merge_replaces_decode_prompt_routing_with_prefill_routing() {
280+
let prompt_payload = uint8_payload(2, 1, 2, &[10, 11, 20, 21]);
281+
let decode_payload = uint8_payload(3, 1, 2, &[0, 0, 1, 1, 30, 31]);
282+
let prefill = json!({
283+
"choices": [{"routed_experts": prompt_payload}],
284+
});
285+
let mut decode = json!({"choices": [{"routed_experts": decode_payload}]});
286+
287+
assert!(merge_routed_experts_in_json(&prefill, &mut decode).unwrap());
288+
let merged = decode_routed_experts_value(
289+
decode["choices"][0].get("routed_experts").unwrap(),
290+
"merged routed_experts",
291+
)
292+
.unwrap();
293+
294+
assert_eq!(merged.seq_len, 3);
295+
assert_eq!(merged.data, vec![10, 11, 20, 21, 30, 31]);
296+
}
297+
}

0 commit comments

Comments
 (0)