11import unittest
22from copy import deepcopy
33
4- from swift .infer_engine .protocol import (ChatCompletionResponse , ChatCompletionResponseChoice , ChatMessage , RequestConfig ,
5- RolloutInferRequest , RolloutOutput , UsageInfo )
4+ from swift .infer_engine .protocol import (ChatCompletionResponse , ChatCompletionResponseChoice , ChatMessage ,
5+ RequestConfig , RolloutInferRequest , RolloutOutput , UsageInfo )
66from swift .rollout .agent_loop import run_multi_turn
77from swift .rollout .multi_turn import MultiTurnScheduler
88
@@ -50,15 +50,16 @@ def step(self, infer_request, response_choice, current_turn):
5050
5151def make_choice (token_ids ):
5252 """Build an inference response choice containing exact sampled IDs."""
53- return ChatCompletionResponseChoice (
54- 0 , ChatMessage ('assistant' , 'sampled text' ), 'stop' , token_ids = token_ids )
53+ return ChatCompletionResponseChoice (0 , ChatMessage ('assistant' , 'sampled text' ), 'stop' , token_ids = token_ids )
5554
5655
5756def make_request ():
5857 """Build a request with an explicit deterministic assistant prefix."""
5958 return RolloutInferRequest (
60- messages = [{'role' : 'user' , 'content' : 'question' }],
61- chat_template_kwargs = {'response_prefix' : '<prefix>' })
59+ messages = [{
60+ 'role' : 'user' ,
61+ 'content' : 'question'
62+ }], chat_template_kwargs = {'response_prefix' : '<prefix>' })
6263
6364
6465class NonBijectiveTokenizer :
@@ -127,10 +128,7 @@ def test_explicit_masked_prefix_is_not_duplicated():
127128 scheduler = Scheduler (tokenizer = PrefixTokenizer (), template = PrefixTemplate ())
128129
129130 ids , mask = scheduler .get_response_token_data (
130- make_request (),
131- make_choice ([11 ]),
132- response_token_ids = [9 , 8 , 7 ],
133- response_loss_mask = [0 , 0 , 0 ])
131+ make_request (), make_choice ([11 ]), response_token_ids = [9 , 8 , 7 ], response_loss_mask = [0 , 0 , 0 ])
134132
135133 assert ids == [9 , 8 , 7 ]
136134 assert mask == [0 , 0 , 0 ]
@@ -140,8 +138,7 @@ def test_continuation_does_not_repeat_response_prefix():
140138 """Treat continuation IDs as part of the current assistant message."""
141139 scheduler = Scheduler (tokenizer = PrefixTokenizer (), template = PrefixTemplate ())
142140
143- ids , mask = scheduler .get_response_token_data (
144- make_request (), make_choice ([12 ]), is_continuation = True )
141+ ids , mask = scheduler .get_response_token_data (make_request (), make_choice ([12 ]), is_continuation = True )
145142
146143 assert ids == [12 ]
147144 assert mask == [1 ]
@@ -152,10 +149,7 @@ def assert_invalid_response_loss_mask_is_rejected(loss_mask):
152149 scheduler = Scheduler (tokenizer = PrefixTokenizer (), template = PrefixTemplate ())
153150
154151 scheduler .get_response_token_data (
155- make_request (),
156- make_choice ([11 , 12 ]),
157- response_token_ids = [11 , 12 ],
158- response_loss_mask = loss_mask )
152+ make_request (), make_choice ([11 , 12 ]), response_token_ids = [11 , 12 ], response_loss_mask = loss_mask )
159153
160154
161155def make_output (token_ids , text , logprobs , finish_reason = None ):
@@ -164,10 +158,11 @@ def make_output(token_ids, text, logprobs, finish_reason=None):
164158 0 ,
165159 ChatMessage ('assistant' , text ),
166160 finish_reason ,
167- logprobs = {'content' : [{'logprob' : value } for value in logprobs ]},
161+ logprobs = {'content' : [{
162+ 'logprob' : value
163+ } for value in logprobs ]},
168164 token_ids = token_ids )
169- response = ChatCompletionResponse (
170- 'fake-model' , [choice ], UsageInfo (0 , len (token_ids ), len (token_ids )))
165+ response = ChatCompletionResponse ('fake-model' , [choice ], UsageInfo (0 , len (token_ids ), len (token_ids )))
171166 return RolloutOutput (response = response )
172167
173168
@@ -220,13 +215,11 @@ def rollout_fn(requests, request_config):
220215 assert requests [0 ].messages [- 1 ] == {'role' : 'user' , 'content' : 'observation' }
221216 return next (outputs_by_turn )
222217
223- result = run_multi_turn (
224- [request ],
225- [first_output ],
226- TwoTurnScheduler (tokenizer = PrefixTokenizer (), template = PrefixTemplate ()),
227- rollout_fn ,
228- RequestConfig (n = 1 ),
229- max_turns = 2 )
218+ result = run_multi_turn ([request ], [first_output ],
219+ TwoTurnScheduler (tokenizer = PrefixTokenizer (), template = PrefixTemplate ()),
220+ rollout_fn ,
221+ RequestConfig (n = 1 ),
222+ max_turns = 2 )
230223
231224 assert result [0 ].response_token_ids == [[9 , 8 , 11 , 12 ], [9 , 8 , 13 ]]
232225 assert result [0 ].response_loss_mask == [[0 , 0 , 1 , 1 ], [0 , 0 , 1 ]]
@@ -244,13 +237,11 @@ def rollout_fn(requests, request_config):
244237 return []
245238 return [second_output ]
246239
247- result = run_multi_turn (
248- [request ],
249- [first_output ],
250- MutatingPrefixScheduler (tokenizer = PrefixTokenizer (), template = PrefixTemplate ()),
251- rollout_fn ,
252- RequestConfig (n = 1 ),
253- max_turns = 2 )
240+ result = run_multi_turn ([request ], [first_output ],
241+ MutatingPrefixScheduler (tokenizer = PrefixTokenizer (), template = PrefixTemplate ()),
242+ rollout_fn ,
243+ RequestConfig (n = 1 ),
244+ max_turns = 2 )
254245
255246 assert result [0 ].response_token_ids == [[9 , 8 , 11 , 12 ], [7 , 13 ]]
256247 assert result [0 ].response_loss_mask == [[0 , 0 , 1 , 1 ], [0 , 1 ]]
@@ -271,13 +262,7 @@ def rollout_fn(requests, request_config):
271262 inference_messages .append (deepcopy (requests [0 ].messages ))
272263 return [second_output ]
273264
274- result = run_multi_turn (
275- [request ],
276- [first_output ],
277- scheduler ,
278- rollout_fn ,
279- RequestConfig (n = 1 ),
280- max_turns = 2 )
265+ result = run_multi_turn ([request ], [first_output ], scheduler , rollout_fn , RequestConfig (n = 1 ), max_turns = 2 )
281266
282267 assert scheduler .hook_messages [0 ][- 1 ] == {'role' : 'assistant' , 'content' : 'first action' }
283268 assert inference_messages [0 ][1 ] == {'role' : 'assistant' , 'content' : [9 , 8 , 11 , 12 ]}
@@ -320,8 +305,7 @@ def step(self, infer_request, response_choice, current_turn):
320305def test_tool_call_scheduler_preserves_sampled_tokens_and_masks_tool_result ():
321306 """Keep tool observations out of the loss while preserving exact history."""
322307 request = make_request ()
323- first_output = make_output (
324- [11 , 12 ], 'Action: calculator\n Action Input: 1 + 2\n ' , [- 0.2 , - 0.4 ], finish_reason = None )
308+ first_output = make_output ([11 , 12 ], 'Action: calculator\n Action Input: 1 + 2\n ' , [- 0.2 , - 0.4 ], finish_reason = None )
325309 second_output = make_output ([31 , 32 ], 'The answer is 3' , [- 0.7 , - 0.8 ], finish_reason = 'stop' )
326310 inference_messages = []
327311
@@ -332,13 +316,11 @@ def rollout_fn(requests, request_config):
332316 inference_messages .append (deepcopy (requests [0 ].messages ))
333317 return [second_output ]
334318
335- result = run_multi_turn (
336- [request ],
337- [first_output ],
338- ToolCallScheduler (tokenizer = ToolResultTokenizer (), template = PrefixTemplate ()),
339- rollout_fn ,
340- RequestConfig (n = 1 ),
341- max_turns = 2 )
319+ result = run_multi_turn ([request ], [first_output ],
320+ ToolCallScheduler (tokenizer = ToolResultTokenizer (), template = PrefixTemplate ()),
321+ rollout_fn ,
322+ RequestConfig (n = 1 ),
323+ max_turns = 2 )
342324
343325 assert inference_messages [0 ][1 ] == {
344326 'role' : 'assistant' ,
@@ -363,13 +345,11 @@ def rollout_fn(requests, request_config):
363345 assert requests [0 ].messages [- 1 ] == {'role' : 'assistant' , 'content' : [9 , 8 , 11 ]}
364346 return next (outputs_by_turn )
365347
366- result = run_multi_turn (
367- [request ],
368- [first_output ],
369- ContinuationScheduler (tokenizer = PrefixTokenizer (), template = PrefixTemplate ()),
370- rollout_fn ,
371- RequestConfig (n = 1 ),
372- max_turns = 2 )
348+ result = run_multi_turn ([request ], [first_output ],
349+ ContinuationScheduler (tokenizer = PrefixTokenizer (), template = PrefixTemplate ()),
350+ rollout_fn ,
351+ RequestConfig (n = 1 ),
352+ max_turns = 2 )
373353
374354 assert result [0 ].response_token_ids == [[9 , 8 , 11 , 12 ]]
375355 assert result [0 ].response_loss_mask == [[0 , 0 , 1 , 1 ]]
0 commit comments