@@ -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"\x93 NUMPY\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"\x93 NUMPY\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"\x93 NUMPY\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