mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
Fix multi-turn synthetic generation and mark incomplete outputs (#2571)
### What does this PR do?
Type of change: Bug fix
Fixes synthetic conversation generation that assumed alternating user
and assistant messages. That assumption skipped user turns in prompt
skeletons and mishandled
leading system messages.
• Preserve conversation history. Regenerate every user turn while
retaining system messages and generated reasoning for subsequent
requests.
• Expose generation controls. Support model-specific request parameters,
configurable timeouts, and server-managed response budgets.
• Handle failures explicitly. Reject empty final answers and unsupported
tool calls. Failed conversations remain retryable without duplicating
saved output.
• Identify incomplete outputs. Mark length- and repetition-stopped
conversations as truncated, preserve stop metadata, and stop generating
follow-up turns.
### Usage
Run from the repository root against a compatible Qwen server with
reasoning parsing enabled:
python examples/speculative_decoding/scripts/server_generate.py \
--data_path input_conversations/train.jsonl \
--output_path synthetic/train.jsonl \
--url http://localhost:8000/v1 \
--model model \
--max_tokens 0 \
--request_timeout 3600 \
--extra_body
'{"chat_template_kwargs":{"enable_thinking":true,"preserve_thinking":true}}'
The model name must match the server’s configured name. Filter truncated
conversations before training.
### Testing
Focused regression tests: 15 passed.
The tests execute the command-line entry point using the real OpenAI
client library with mocked HTTP transport.
Coverage includes multi-turn generation, system prompts, reasoning
preservation, request parameters, failure recovery, resume
deduplication, truncation, and invalid
responses.
python -m pytest \
--confcutdir=tests/examples/speculative_decoding \
tests/examples/speculative_decoding/test_server_generate.py -q
The isolated test configuration avoids an unrelated parent configuration
import failure. All applicable pre-commit checks passed for the
generator, tests, and
documentation.
### Before your PR is "Ready for review"
• Is this change backward compatible?: ✅ Existing valid inputs,
defaults, conversation output structure, and resume behavior remain
supported. Invalid inputs and
failed requests now raise errors instead of being silently accepted.
• Copied code or new PIP dependencies?: N/A. No new third-party code or
dependencies were added.
• Did you write any new necessary tests?: ✅ Added focused command-line
regression tests.
• Did you update Changelog?: N/A. These are example-script correctness
fixes, not critical released library fixes.
• Did you get Claude approval on this PR?: ❌ Not yet obtained.
### Additional Information
<!-- This is an auto-generated comment: release notes by coderabbit.ai
-->
## Summary by CodeRabbit
* **New Features**
* Data generation supports `conversations` and `messages` inputs,
preserves reasoning content, and accepts additional chat settings and
configurable request timeouts.
* Failed conversations are recorded separately, with options to retry
failures or exit when errors occur. Resume behavior distinguishes
retryable failures from rejected inputs.
* Outputs identify conversations truncated by length or repetition
limits.
* **Documentation**
* Updated data preparation guides with generation setup, input formats,
failure handling, resuming, and training guidance.
<!-- end of auto-generated comment: release notes by coderabbit.ai -->
---------
Signed-off-by: Shiyang Chen <shiychen@nvidia.com>
This commit is contained in:
@@ -0,0 +1,368 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
"""Exercise generator control flow with a stub client, without requiring the OpenAI SDK."""
|
||||
|
||||
import copy
|
||||
import json
|
||||
import runpy
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from types import ModuleType, SimpleNamespace
|
||||
from unittest.mock import Mock
|
||||
|
||||
import pytest
|
||||
|
||||
_EXAMPLE = Path(__file__).resolve().parents[3] / "examples/speculative_decoding"
|
||||
|
||||
|
||||
class _ConnectionError(Exception):
|
||||
"""Stand in for a network failure without importing the client SDK."""
|
||||
|
||||
|
||||
class _StatusError(Exception):
|
||||
"""Expose the HTTP status used by the generator's failure classifier."""
|
||||
|
||||
def __init__(self, status_code):
|
||||
super().__init__(f"Request failed with status {status_code}")
|
||||
self.status_code = status_code
|
||||
|
||||
|
||||
def _response(content="answer", finish_reason="stop", stop_reason=None, **message_fields):
|
||||
"""Build only the response attributes consumed by the generator."""
|
||||
message = SimpleNamespace(
|
||||
content=content, tool_calls=None, function_call=None, **message_fields
|
||||
)
|
||||
choice = SimpleNamespace(message=message, finish_reason=finish_reason, stop_reason=stop_reason)
|
||||
return SimpleNamespace(choices=[choice])
|
||||
|
||||
|
||||
def _sample(prompt):
|
||||
"""Build a user-only conversation without loading a dataset."""
|
||||
return {"messages": [{"role": "user", "content": prompt}]}
|
||||
|
||||
|
||||
def _read_jsonl(path):
|
||||
"""Read the generated output or journal, treating absent files as empty."""
|
||||
return [json.loads(line) for line in path.read_text().splitlines()] if path.exists() else []
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def run_generator(monkeypatch, tmp_path):
|
||||
"""Run the real CLI and files, replacing only the optional network client boundary."""
|
||||
data_path = tmp_path / "input.json"
|
||||
output_path = tmp_path / "shards/output.jsonl"
|
||||
|
||||
def run(samples, responses, *options):
|
||||
data_path.write_text(json.dumps(samples))
|
||||
responses = iter(responses)
|
||||
requests = []
|
||||
|
||||
def create(**kwargs):
|
||||
# Later turns mutate the same history list passed to the client.
|
||||
requests.append(copy.deepcopy(kwargs))
|
||||
response = next(responses)
|
||||
if isinstance(response, Exception):
|
||||
raise response
|
||||
return response
|
||||
|
||||
client = SimpleNamespace(chat=SimpleNamespace(completions=SimpleNamespace(create=create)))
|
||||
openai_stub = ModuleType("openai")
|
||||
openai_stub.OpenAI = Mock(return_value=client)
|
||||
openai_stub.APIConnectionError = _ConnectionError
|
||||
openai_stub.APIStatusError = _StatusError
|
||||
monkeypatch.setitem(sys.modules, "openai", openai_stub)
|
||||
monkeypatch.setattr(
|
||||
sys,
|
||||
"argv",
|
||||
[
|
||||
"server_generate.py",
|
||||
"--data_path",
|
||||
str(data_path),
|
||||
"--output_path",
|
||||
str(output_path),
|
||||
"--num_threads",
|
||||
"1",
|
||||
"--log_empty_conversations",
|
||||
*options,
|
||||
],
|
||||
)
|
||||
exit_code = 0
|
||||
try:
|
||||
runpy.run_path(str(_EXAMPLE / "scripts/server_generate.py"), run_name="__main__")
|
||||
except SystemExit as exc:
|
||||
exit_code = exc.code
|
||||
return SimpleNamespace(
|
||||
exit_code=exit_code,
|
||||
rows=_read_jsonl(output_path),
|
||||
failures=_read_jsonl(Path(str(output_path) + ".failures")),
|
||||
requests=requests,
|
||||
client_factory=openai_stub.OpenAI,
|
||||
output_path=output_path,
|
||||
)
|
||||
|
||||
return run
|
||||
|
||||
|
||||
@pytest.mark.parametrize("sharegpt", [False, True])
|
||||
def test_multi_turn_history_and_request_options(run_generator, sharegpt):
|
||||
"""Regenerate each user turn using generated history and configured request options."""
|
||||
messages = [
|
||||
{"role": "system", "content": "Be helpful."},
|
||||
{"role": "user", "content": "first"},
|
||||
{"role": "assistant", "content": "discard this reference answer"},
|
||||
{"role": "user", "content": "second"},
|
||||
]
|
||||
sample = {"messages": messages}
|
||||
if sharegpt:
|
||||
roles = {"system": "system", "user": "human", "assistant": "gpt"}
|
||||
sample = {
|
||||
"conversations": [{"from": roles[m["role"]], "value": m["content"]} for m in messages]
|
||||
}
|
||||
result = run_generator(
|
||||
[sample],
|
||||
[_response(" first answer ", reasoning="thinking"), _response("second answer")],
|
||||
"--max_tokens",
|
||||
"0",
|
||||
"--request_timeout",
|
||||
"90",
|
||||
"--extra_body",
|
||||
'{"chat_template_kwargs": {"enable_thinking": true}}',
|
||||
)
|
||||
generated = {"role": "assistant", "content": "first answer", "reasoning_content": "thinking"}
|
||||
assert result.exit_code == 0
|
||||
assert result.requests[0]["messages"] == messages[:2]
|
||||
assert result.requests[1]["messages"] == [*messages[:2], generated, messages[3]]
|
||||
assert all(request["max_tokens"] is None for request in result.requests)
|
||||
assert all(
|
||||
request["extra_body"] == {"chat_template_kwargs": {"enable_thinking": True}}
|
||||
for request in result.requests
|
||||
)
|
||||
assert result.client_factory.call_args.kwargs["timeout"] == 90
|
||||
assert result.rows == [
|
||||
{
|
||||
"conversation_id": 0,
|
||||
"conversations": [
|
||||
*messages[:2],
|
||||
generated,
|
||||
messages[3],
|
||||
{"role": "assistant", "content": "second answer"},
|
||||
],
|
||||
},
|
||||
{"finished": True},
|
||||
]
|
||||
assert not result.failures
|
||||
|
||||
|
||||
@pytest.mark.parametrize("has_system_message", [False, True])
|
||||
def test_system_prompt_override_warns_only_when_replacing_input(
|
||||
run_generator, capsys, has_system_message
|
||||
):
|
||||
"""Warn when overriding a dataset system message without exposing its contents."""
|
||||
sample = _sample("question")
|
||||
if has_system_message:
|
||||
sample["messages"].insert(0, {"role": "system", "content": "dataset instructions"})
|
||||
result = run_generator([sample], [_response()], "--system_prompt", "override instructions")
|
||||
assert result.exit_code == 0
|
||||
assert result.requests[0]["messages"] == [
|
||||
{"role": "system", "content": "override instructions"},
|
||||
{"role": "user", "content": "question"},
|
||||
]
|
||||
stderr = capsys.readouterr().err
|
||||
assert ("--system_prompt overrides the input system message" in stderr) == has_system_message
|
||||
assert "dataset instructions" not in stderr
|
||||
|
||||
|
||||
@pytest.mark.parametrize("temperature", [0.0, 0.7])
|
||||
@pytest.mark.parametrize("strict", [False, True])
|
||||
def test_empty_answer_resume_depends_on_temperature(run_generator, temperature, strict):
|
||||
"""Retry sampled empty answers on ordinary resume without keeping partial conversations."""
|
||||
samples = [
|
||||
{"messages": [*_sample("first")["messages"], *_sample("second")["messages"]]},
|
||||
_sample("keep me"),
|
||||
]
|
||||
options = ["--temperature", str(temperature)]
|
||||
if strict:
|
||||
options.append("--fail_on_error")
|
||||
first = run_generator(
|
||||
samples,
|
||||
[_response(), _response(content=" ", reasoning_content="no final answer"), _response()],
|
||||
*options,
|
||||
)
|
||||
retryable = temperature > 0
|
||||
assert first.exit_code == int(strict)
|
||||
assert all(request["temperature"] == temperature for request in first.requests)
|
||||
assert [row["conversation_id"] for row in first.rows if "conversation_id" in row] == [1]
|
||||
assert len(first.failures) == 1
|
||||
assert first.failures[0]["conversation_id"] == 0
|
||||
assert first.failures[0]["retryable"] == retryable
|
||||
assert (first.rows[-1] == {"finished": True}) == (not retryable)
|
||||
|
||||
responses = [_response(), _response()] if retryable else []
|
||||
resumed = run_generator(samples, responses, *options)
|
||||
assert resumed.exit_code == int(strict and not retryable)
|
||||
assert len(resumed.requests) == len(responses)
|
||||
assert [row["conversation_id"] for row in resumed.rows if "conversation_id" in row] == (
|
||||
[1, 0] if retryable else [1]
|
||||
)
|
||||
assert resumed.rows[-1] == {"finished": True}
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"error", [_ConnectionError("connection lost"), _StatusError(429), _StatusError(503)]
|
||||
)
|
||||
def test_transient_failure_resumes_without_duplicates(run_generator, error):
|
||||
"""Retry transient failures while preserving previously written conversations."""
|
||||
samples = [_sample("retry me"), _sample("keep me")]
|
||||
first = run_generator(samples, [error, _response()])
|
||||
assert first.exit_code == 0
|
||||
assert [row["conversation_id"] for row in first.rows] == [1]
|
||||
assert len(first.failures) == 1
|
||||
assert first.failures[0]["conversation_id"] == 0
|
||||
assert first.failures[0]["retryable"] is True
|
||||
|
||||
resumed = run_generator(samples, [_response()])
|
||||
assert resumed.exit_code == 0
|
||||
assert len(resumed.requests) == 1
|
||||
assert resumed.requests[0]["messages"] == samples[0]["messages"]
|
||||
assert [row["conversation_id"] for row in resumed.rows[:-1]] == [1, 0]
|
||||
assert resumed.rows[-1] == {"finished": True}
|
||||
completed = run_generator(samples, [], "--fail_on_error")
|
||||
assert completed.exit_code == 0
|
||||
assert completed.requests == []
|
||||
assert completed.rows == resumed.rows
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"error",
|
||||
[
|
||||
_StatusError(400),
|
||||
_StatusError(422),
|
||||
_response(content="", reasoning_content="no final answer"),
|
||||
],
|
||||
)
|
||||
def test_rejected_sample_continues_and_requires_explicit_retry(
|
||||
run_generator, error, monkeypatch, tmp_path
|
||||
):
|
||||
"""Skip rejections on ordinary resume and keep their journal out of combined data."""
|
||||
samples = [_sample("reject me"), _sample("keep me")]
|
||||
first = run_generator(samples, [error, _response()])
|
||||
assert first.exit_code == 0
|
||||
assert [row.get("conversation_id") for row in first.rows] == [1, None]
|
||||
assert first.rows[-1] == {"finished": True}
|
||||
assert len(first.failures) == 1
|
||||
assert first.failures[0]["retryable"] is False
|
||||
skipped = run_generator(samples, [])
|
||||
assert skipped.exit_code == 0
|
||||
assert skipped.requests == []
|
||||
|
||||
combined = tmp_path / "combined.jsonl"
|
||||
monkeypatch.setattr(
|
||||
sys,
|
||||
"argv",
|
||||
[
|
||||
"sharding_utils.py",
|
||||
"--combine",
|
||||
"--input_dir",
|
||||
str(first.output_path.parent),
|
||||
"--output_path",
|
||||
str(combined),
|
||||
],
|
||||
)
|
||||
runpy.run_path(str(_EXAMPLE / "distributed_generate/sharding_utils.py"), run_name="__main__")
|
||||
assert _read_jsonl(combined) == [{"conversations": first.rows[0]["conversations"]}]
|
||||
|
||||
retried = run_generator(samples, [_response()], "--retry_failed", "--fail_on_error")
|
||||
assert retried.exit_code == 0
|
||||
assert len(retried.requests) == 1
|
||||
assert [row["conversation_id"] for row in retried.rows if "conversation_id" in row] == [1, 0]
|
||||
assert run_generator(samples, [], "--fail_on_error").exit_code == 0
|
||||
|
||||
|
||||
def test_retryable_failure_overrides_old_completion_marker(run_generator):
|
||||
"""Resume a transient failure even when an earlier run wrote a completion marker."""
|
||||
samples = [_sample("retry me")]
|
||||
rejected = run_generator(samples, [_StatusError(400)])
|
||||
assert rejected.rows == [{"finished": True}]
|
||||
retry = run_generator(samples, [_StatusError(503)], "--retry_failed")
|
||||
assert retry.exit_code == 0
|
||||
assert retry.rows == rejected.rows
|
||||
assert retry.failures[-1]["retryable"] is True
|
||||
resumed = run_generator(samples, [_response()])
|
||||
assert resumed.exit_code == 0
|
||||
assert len(resumed.requests) == 1
|
||||
assert [row["conversation_id"] for row in resumed.rows if "conversation_id" in row] == [0]
|
||||
|
||||
|
||||
def test_strict_mode_reports_all_failures_and_keeps_successes(run_generator, capsys):
|
||||
"""Report the whole batch and return nonzero until all strict-mode failures resolve."""
|
||||
samples = [_sample("reject"), _sample("transient"), _sample("success")]
|
||||
first = run_generator(
|
||||
samples, [_StatusError(400), _StatusError(503), _response()], "--fail_on_error"
|
||||
)
|
||||
assert first.exit_code == 1
|
||||
assert {failure["conversation_id"] for failure in first.failures} == {0, 1}
|
||||
assert [row["conversation_id"] for row in first.rows] == [2]
|
||||
assert "2 conversations failed (1 retryable)" in capsys.readouterr().err
|
||||
resumed = run_generator(samples, [_response()], "--fail_on_error")
|
||||
assert resumed.exit_code == 1
|
||||
assert len(resumed.requests) == 1
|
||||
assert resumed.rows[-1] == {"finished": True}
|
||||
completed = run_generator(samples, [], "--fail_on_error")
|
||||
assert completed.exit_code == 1
|
||||
assert completed.requests == []
|
||||
|
||||
|
||||
@pytest.mark.parametrize("status", [401, 404])
|
||||
def test_fatal_request_error_exits_nonzero(run_generator, status):
|
||||
"""Keep server configuration errors fatal even without strict mode."""
|
||||
result = run_generator([_sample("fatal")], [_StatusError(status)])
|
||||
assert result.exit_code == 1
|
||||
assert result.rows == []
|
||||
assert result.failures[0]["retryable"] is True
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("finish_reason", "stop_reason"),
|
||||
[("length", None), ("repetition", None), ("stop", "repetition_detected")],
|
||||
)
|
||||
def test_truncation_stops_remaining_turns(run_generator, finish_reason, stop_reason):
|
||||
"""Persist the incomplete response and its stop metadata without generating later turns."""
|
||||
sample = {"messages": [*_sample("first")["messages"], *_sample("second")["messages"]]}
|
||||
result = run_generator([sample], [_response("partial", finish_reason, stop_reason)])
|
||||
assert result.exit_code == 0
|
||||
assert len(result.requests) == 1
|
||||
assert result.rows[0] == {
|
||||
"conversation_id": 0,
|
||||
"conversations": [sample["messages"][0], {"role": "assistant", "content": "partial"}],
|
||||
"truncated": True,
|
||||
"finish_reason": finish_reason,
|
||||
"stop_reason": stop_reason,
|
||||
}
|
||||
assert not result.failures
|
||||
|
||||
|
||||
def test_unsupported_role_does_not_write_partial_conversation(run_generator):
|
||||
"""Reject unsupported roles without persisting an earlier successful turn."""
|
||||
sample = {
|
||||
"messages": [*_sample("first")["messages"], {"role": "tool", "content": "unsupported"}]
|
||||
}
|
||||
result = run_generator(
|
||||
[sample, _sample("success")], [_response(), _response()], "--temperature", "0.7"
|
||||
)
|
||||
assert result.exit_code == 0
|
||||
assert [row.get("conversation_id") for row in result.rows] == [1, None]
|
||||
assert result.failures[0]["conversation_id"] == 0
|
||||
assert result.failures[0]["retryable"] is False
|
||||
assert "Unsupported message role" in result.failures[0]["error"]
|
||||
Reference in New Issue
Block a user