1- import base64
2- from io import BytesIO
3- from typing import Any , cast
4-
5- import numpy as np
6-
71from verifiers .types import (
82 AssistantMessage ,
93 Messages ,
104 Response ,
115 TrajectoryStepTokens ,
126)
137
8+ ROUTED_EXPERTS_DATA_PREFIX = b'"routed_experts":{"data":"'
149
15- def parse_routed_experts (raw : Any ) -> str | None :
16- if raw is None :
17- return None
18- return cast (str , raw )
19-
20-
21- def truncate_routed_experts (routed_experts : str | None , seq_len : int ) -> str | None :
22- if routed_experts is None :
23- return None
2410
25- array = np .load (BytesIO (base64 .b64decode (routed_experts )), allow_pickle = False )
26- assert array .ndim == 3
27- assert 0 <= seq_len <= array .shape [0 ]
11+ def strip_routed_experts_data (raw : bytes ) -> tuple [bytes , memoryview | None ]:
12+ data_start = raw .find (ROUTED_EXPERTS_DATA_PREFIX )
13+ if data_start < 0 :
14+ return raw , None
2815
29- buffer = BytesIO ()
30- np .save (buffer , np .ascontiguousarray (array [:seq_len ]), allow_pickle = False )
31- return base64 .b64encode (buffer .getvalue ()).decode ("ascii" )
16+ data_start += len (ROUTED_EXPERTS_DATA_PREFIX )
17+ data_end = raw .index (b'"' , data_start )
18+ routed_data = memoryview (raw )[data_start :data_end ]
19+ stripped = raw [:data_start ] + raw [data_end :]
20+ return stripped , routed_data
3221
3322
3423async def parse_response_message (response : Response ) -> Messages :
@@ -73,15 +62,11 @@ async def parse_response_tokens(
7362 completion_ids = []
7463 completion_mask = []
7564 completion_logprobs = []
76- routed_experts = truncate_routed_experts (routed_experts , len (prompt_ids ))
7765 elif prompt_len + completion_len > max_seq_len :
7866 is_truncated = True
7967 completion_ids = tokens .completion_ids [: max_seq_len - prompt_len ]
8068 completion_mask = tokens .completion_mask [: max_seq_len - prompt_len ]
8169 completion_logprobs = tokens .completion_logprobs [: max_seq_len - prompt_len ]
82- routed_experts = truncate_routed_experts (
83- routed_experts , prompt_len + len (completion_ids )
84- )
8570 else :
8671 is_truncated = False
8772 else :
@@ -104,4 +89,6 @@ async def parse_response_tokens(
10489 # step. Leaving it on ``response.message.tokens`` too means every
10590 # downstream pass (msgpack, save) has to dedupe the duplicate.
10691 tokens .multi_modal_data = None
92+ if routed_experts is not None :
93+ tokens .routed_experts = None
10794 return out
0 commit comments