Skip to content

Commit 75c824c

Browse files
author
P90-RushB
committed
style: format exact token I/O tests
1 parent 1b4c00f commit 75c824c

2 files changed

Lines changed: 66 additions & 87 deletions

File tree

tests/rollout/test_exact_token_io.py

Lines changed: 36 additions & 56 deletions
Original file line numberDiff line numberDiff line change
@@ -1,8 +1,8 @@
11
import unittest
22
from 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)
66
from swift.rollout.agent_loop import run_multi_turn
77
from swift.rollout.multi_turn import MultiTurnScheduler
88

@@ -50,15 +50,16 @@ def step(self, infer_request, response_choice, current_turn):
5050

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

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

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

161155
def 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):
320305
def 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\nAction Input: 1 + 2\n', [-0.2, -0.4], finish_reason=None)
308+
first_output = make_output([11, 12], 'Action: calculator\nAction 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]]

tests/rollout/test_server_exact_token_io.py

Lines changed: 30 additions & 31 deletions
Original file line numberDiff line numberDiff line change
@@ -1,21 +1,14 @@
11
import asyncio
22
import sys
3+
import torch
34
import unittest
45
from contextlib import ExitStack, nullcontext
56
from copy import deepcopy
67
from types import ModuleType, SimpleNamespace
78
from unittest.mock import patch
89

9-
import torch
10-
11-
from swift.infer_engine.protocol import (
12-
ChatCompletionResponse,
13-
ChatCompletionResponseChoice,
14-
ChatMessage,
15-
RequestConfig,
16-
RolloutInferRequest,
17-
UsageInfo,
18-
)
10+
from swift.infer_engine.protocol import (ChatCompletionResponse, ChatCompletionResponseChoice, ChatMessage,
11+
RequestConfig, RolloutInferRequest, UsageInfo)
1912
from swift.rl_core.data import GRPOBatch, GRPOSample
2013
from swift.rlhf_trainers.grpo_trainer import GRPOTrainer
2114
from swift.rollout.multi_turn import MultiTurnScheduler, RolloutScheduler
@@ -62,16 +55,12 @@ def encode(self, data, **kwargs):
6255
response_mask = [1]
6356
return {
6457
'input_ids': [100, *response_ids],
65-
'labels': [-100, *[token_id if mask else -100
66-
for token_id, mask in zip(response_ids, response_mask)]],
58+
'labels': [-100, *[token_id if mask else -100 for token_id, mask in zip(response_ids, response_mask)]],
6759
}
6860

6961
def data_collator(self, encoded_data, padding_to=None):
7062
self.encoded_data = encoded_data
71-
return {
72-
key: torch.tensor([item[key] for item in encoded_data])
73-
for key in ('input_ids', 'labels')
74-
}
63+
return {key: torch.tensor([item[key] for item in encoded_data]) for key in ('input_ids', 'labels')}
7564

7665

7766
class AsyncTwoTurnEngine:
@@ -108,21 +97,24 @@ def make_response(token_ids, text, logprobs, finish_reason=None):
10897
0,
10998
ChatMessage('assistant', text),
11099
finish_reason,
111-
logprobs={'content': [{'logprob': value} for value in logprobs]},
100+
logprobs={'content': [{
101+
'logprob': value
102+
} for value in logprobs]},
112103
token_ids=token_ids)
113104
return ChatCompletionResponse('fake-model', [choice], UsageInfo(0, len(token_ids), len(token_ids)))
114105

115106

116107
def test_server_scheduler_preserves_exact_token_history():
117108
request = RolloutInferRequest(
118-
messages=[{'role': 'user', 'content': 'question'}],
119-
chat_template_kwargs={'response_prefix': '<prefix>'})
109+
messages=[{
110+
'role': 'user',
111+
'content': 'question'
112+
}], chat_template_kwargs={'response_prefix': '<prefix>'})
120113
engine = AsyncTwoTurnEngine([
121114
make_response([11, 12], 'first action', [-0.2, -0.4]),
122115
make_response([13], 'second action', [-0.7], finish_reason='stop'),
123116
])
124-
scheduler = ServerBoundaryScheduler(
125-
infer_engine=engine, tokenizer=PrefixTokenizer(), template=PrefixTemplate())
117+
scheduler = ServerBoundaryScheduler(infer_engine=engine, tokenizer=PrefixTokenizer(), template=PrefixTemplate())
126118

127119
result = asyncio.run(scheduler.run(request, RequestConfig(n=1)))
128120

@@ -141,16 +133,19 @@ def test_server_scheduler_preserves_exact_token_history():
141133
def test_multimodal_chunk_rebuild_preserves_exact_response_tokens():
142134
sample = GRPOSample(
143135
messages=[
144-
{'role': 'user', 'content': 'question'},
145-
{'role': 'assistant', 'content': 'decoded response'},
136+
{
137+
'role': 'user',
138+
'content': 'question'
139+
},
140+
{
141+
'role': 'assistant',
142+
'content': 'decoded response'
143+
},
146144
],
147145
images=[object()],
148146
response_token_ids=[[9, 8, 11, 12]],
149147
response_loss_mask=[[0, 0, 1, 1]])
150-
batch = GRPOBatch(
151-
completion_mask=torch.ones((1, 5)),
152-
truncated_mask=torch.zeros(1),
153-
seq_lengths=torch.tensor([5]))
148+
batch = GRPOBatch(completion_mask=torch.ones((1, 5)), truncated_mask=torch.zeros(1), seq_lengths=torch.tensor([5]))
154149
template = ExactTokenTemplate()
155150
trainer = SimpleNamespace(
156151
is_multimodal=True,
@@ -169,10 +164,14 @@ def test_token_backed_response_deduplicates_template_separator_overlap():
169164
template = object.__new__(SwiftTemplate)
170165
template.processor = SeparatorTokenizer()
171166

172-
assert template._remove_response_separator_overlap(
173-
{'token_ids': [1, 9], 'loss_scale': [1, 1]}, ['<end>\n']) == [[10]]
174-
assert template._remove_response_separator_overlap(
175-
{'token_ids': [1, 7, 8], 'loss_scale': [1, 1, 1]}, ['<pair>\n']) == [[10]]
167+
assert template._remove_response_separator_overlap({
168+
'token_ids': [1, 9],
169+
'loss_scale': [1, 1]
170+
}, ['<end>\n']) == [[10]]
171+
assert template._remove_response_separator_overlap({
172+
'token_ids': [1, 7, 8],
173+
'loss_scale': [1, 1, 1]
174+
}, ['<pair>\n']) == [[10]]
176175
assert template._remove_response_separator_overlap([1, 9, 10], ['<end>\n']) == []
177176
no_overlap = ['<end>\n']
178177
assert template._remove_response_separator_overlap({'token_ids': [1, 2]}, no_overlap) is no_overlap

0 commit comments

Comments
 (0)