|
| 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