Skip to content

Commit 76fe721

Browse files
committed
merge: update routed experts core
Signed-off-by: Biswa Panda <biswa.panda@gmail.com>
2 parents 6b70c1e + 5a3165c commit 76fe721

1 file changed

Lines changed: 86 additions & 0 deletions

File tree

rust/src/engine-core-client/src/protocol/routed_experts.rs

Lines changed: 86 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -40,6 +40,64 @@ impl RoutedExperts {
4040
self.data.append(&mut other.data);
4141
Ok(())
4242
}
43+
44+
/// Serialize the tensor using NumPy's version 1.0 `.npy` format.
45+
pub fn to_npy_bytes(&self) -> Result<Vec<u8>> {
46+
if self.shape.len() != 3 {
47+
bail_ext_value_decode!(
48+
"routed_experts: expected a rank-3 ndarray, got shape {:?}",
49+
self.shape
50+
);
51+
}
52+
let (descriptor, item_size) = match self.dtype.as_str() {
53+
"uint8" => ("|u1", 1usize),
54+
"uint16" => ("<u2", 2usize),
55+
other => bail_ext_value_decode!(
56+
"routed_experts: expected normalized uint8 or uint16 dtype, got {other:?}"
57+
),
58+
};
59+
let expected = self
60+
.shape
61+
.checked_numel()
62+
.and_then(|numel| numel.checked_mul(item_size))
63+
.ok_or_else(|| {
64+
ext_value_decode!(
65+
"routed_experts: shape byte length overflowed usize: {:?}",
66+
self.shape
67+
)
68+
})?;
69+
if self.data.len() != expected {
70+
bail_ext_value_decode!(
71+
"routed_experts: byte length mismatch: expected {expected}, got {}",
72+
self.data.len()
73+
);
74+
}
75+
76+
let dictionary = format!(
77+
"{{'descr': '{descriptor}', 'fortran_order': False, 'shape': ({}, {}, {}), }}",
78+
self.shape[0], self.shape[1], self.shape[2]
79+
);
80+
const PREAMBLE_LEN: usize = 10;
81+
const ARRAY_ALIGNMENT: usize = 64;
82+
let padding = ARRAY_ALIGNMENT - ((PREAMBLE_LEN + dictionary.len() + 1) % ARRAY_ALIGNMENT);
83+
let header_len = dictionary
84+
.len()
85+
.checked_add(padding)
86+
.and_then(|length| length.checked_add(1))
87+
.ok_or_else(|| ext_value_decode!("routed_experts: NumPy header length overflow"))?;
88+
let header_len = u16::try_from(header_len).map_err(|_| {
89+
ext_value_decode!("routed_experts: NumPy v1 header does not fit in uint16")
90+
})?;
91+
92+
let mut encoded = Vec::with_capacity(PREAMBLE_LEN + usize::from(header_len) + expected);
93+
encoded.extend_from_slice(b"\x93NUMPY\x01\x00");
94+
encoded.extend_from_slice(&header_len.to_le_bytes());
95+
encoded.extend_from_slice(dictionary.as_bytes());
96+
encoded.resize(encoded.len() + padding, b' ');
97+
encoded.push(b'\n');
98+
encoded.extend_from_slice(&self.data);
99+
Ok(encoded)
100+
}
43101
}
44102

45103
/// Routed-experts output is initially decoded from Python's ndarray wire
@@ -161,4 +219,32 @@ mod tests {
161219

162220
assert!(error.to_string().contains("expected a rank-3 ndarray"));
163221
}
222+
223+
#[test]
224+
fn serializes_uint8_as_numpy_v1() {
225+
let routed = RoutedExperts {
226+
dtype: "uint8".to_string(),
227+
shape: vec![2, 1, 2],
228+
data: vec![1, 2, 3, 4],
229+
};
230+
231+
assert_eq!(
232+
routed.to_npy_bytes().expect("serialize routed experts"),
233+
b"\x93NUMPY\x01\x00\x76\x00{'descr': '|u1', 'fortran_order': False, 'shape': (2, 1, 2), } \n\x01\x02\x03\x04"
234+
);
235+
}
236+
237+
#[test]
238+
fn serializes_uint16_as_numpy_v1() {
239+
let routed = RoutedExperts {
240+
dtype: "uint16".to_string(),
241+
shape: vec![2, 1, 2],
242+
data: vec![1, 0, 2, 0, 3, 0, 4, 0],
243+
};
244+
245+
assert_eq!(
246+
routed.to_npy_bytes().expect("serialize routed experts"),
247+
b"\x93NUMPY\x01\x00\x76\x00{'descr': '<u2', 'fortran_order': False, 'shape': (2, 1, 2), } \n\x01\x00\x02\x00\x03\x00\x04\x00"
248+
);
249+
}
164250
}

0 commit comments

Comments
 (0)