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:
@@ -271,18 +271,61 @@ First, prepare input conversation skeletons using `--mode generate` (default) fr
|
||||
|
||||
```bash
|
||||
pip install vllm
|
||||
vllm serve meta-llama/Llama-3.2-1B-Instruct --api-key token-abc123 --port 8000 --tensor-parallel-size 1
|
||||
vllm serve meta-llama/Llama-3.2-1B-Instruct --api-key token-abc123 --port 8000 --generation-config vllm --tensor-parallel-size 1
|
||||
```
|
||||
|
||||
Note: Add `--quantization=modelopt` flag for quantized models.
|
||||
|
||||
The `--generation-config vllm` server option avoids a checkpoint-specific output cap.
|
||||
|
||||
Then, we generate conversations with the base model using the prepared prompts:
|
||||
|
||||
```bash
|
||||
python scripts/server_generate.py --data_path input_conversations/train.jsonl --output_path synthetic/train.jsonl
|
||||
```
|
||||
|
||||
To add a system prompt, use the `--system_prompt <system_prompt_text>` argument.
|
||||
Inputs can use `conversations` or `messages`, with either full conversations or user-only
|
||||
skeletons. Every user turn gets a fresh response using the previously generated responses
|
||||
as history. Input system messages are preserved; `--system_prompt <system_prompt_text>`
|
||||
overrides them and emits a warning when an input system message is replaced.
|
||||
Output remains one full conversation per JSONL record, with no train/eval/test
|
||||
split or expansion into separate assistant-turn examples. Existing output IDs are skipped on resume.
|
||||
|
||||
Use `--extra_body` to pass model-specific chat parameters, including thinking controls and
|
||||
sampling settings. For example, with a compatible Qwen server and reasoning parser:
|
||||
|
||||
```bash
|
||||
python scripts/server_generate.py --data_path input_conversations/train.jsonl \
|
||||
--output_path synthetic/train.jsonl --model qwen3.8-27b --temperature 1.0 --max_tokens 8192 \
|
||||
--extra_body '{"reasoning_effort":"medium","top_p":0.95,"top_k":20,"chat_template_kwargs":{"enable_thinking":true,"preserve_thinking":true}}'
|
||||
```
|
||||
|
||||
The client uses the server's default thinking mode and effort unless overridden through
|
||||
`--extra_body`. Returned reasoning is saved in the
|
||||
assistant's `reasoning_content` field and included in subsequent requests. Token-capped
|
||||
conversations retain the existing `truncated: true` flag and need filtering before training.
|
||||
Failed conversations and partial answers are excluded from training output. Generation continues
|
||||
after per-conversation failures and records each one in `<output_path>.failures`, a JSONL journal
|
||||
whose filename deliberately does not end in `.jsonl` so the shard combiner ignores it.
|
||||
|
||||
Rerun the same command with the same input ordering and output path to resume. Completed
|
||||
conversation IDs are skipped. Connection errors, timeouts, rate limits, and temporary server
|
||||
errors remain retryable on resume. Empty final answers also remain retryable when `--temperature`
|
||||
is greater than zero. At temperature zero, empty final answers are recorded as rejected to
|
||||
avoid repeating greedy-generation failures. Other rejected inputs and responses, including
|
||||
unsupported tool roles or calls and HTTP 400/422 responses, are skipped on ordinary resume.
|
||||
Inspect the journal and use `--retry_failed` to retry these rejections after correcting the
|
||||
input or request settings. This script has no tool-execution loop.
|
||||
|
||||
Use `--fail_on_error` to return a nonzero exit status if failures remain after the batch finishes.
|
||||
Authentication failures, missing endpoints or models, and unexpected internal or output-write
|
||||
errors still exit nonzero by default. When `--log_empty_conversations` is enabled, the `finished` marker means
|
||||
every conversation is either saved or recorded as rejected; no marker is appended while
|
||||
retryable failures remain. Always inspect the failure journal before using the generated data.
|
||||
|
||||
For chat generation, `--max_tokens 0` sends no fixed response cap; the server determines
|
||||
the budget from its configured context window and generation defaults. Increase
|
||||
`--request_timeout` (seconds, default 600) when long responses need more time.
|
||||
|
||||
For large scale data generation, please see [SLURM prepare data](SLURM_prepare_data.md) for SLURM support.
|
||||
|
||||
|
||||
@@ -30,8 +30,44 @@ To process the next 40 shards
|
||||
bash distributed_generate/launch.sh $SLURM_JOB_ID vllm TinyLlama/TinyLlama-1.1B-Chat-v1.0 /data/train/ /data/output /scripts/ 40 10 n1,n2,n3,n4
|
||||
```
|
||||
|
||||
## Failures and resuming
|
||||
|
||||
Workers continue to later shards after per-conversation failures by default. Successful
|
||||
conversations stay in the output shard, while failures are logged to stderr and recorded in
|
||||
`<output_shard>.failures`. Failed conversations and partial answers are never written as
|
||||
training rows. The journal records the conversation ID, error, and whether ordinary resume
|
||||
will retry it.
|
||||
|
||||
Rerun the launch command for the same shard range and output directory to resume. Keep shard
|
||||
names and input ordering unchanged, since conversation IDs are positions within each shard.
|
||||
Completed IDs are skipped. Temporary connection, timeout, rate-limit, and server errors remain
|
||||
retryable. Empty final answers are retryable at positive temperatures, but recorded as rejected
|
||||
at temperature zero. Unsupported tool roles or calls, malformed conversations, and HTTP 400/422
|
||||
responses are also recorded as rejected and skipped on ordinary resume. Inspect these
|
||||
rejections before training; they can indicate bad input or incompatible request settings.
|
||||
|
||||
Set `RETRY_FAILED=1` when launching to retry rejected conversations after fixing their input
|
||||
or request settings. Set `FAIL_ON_ERROR=1` to make any unresolved failures return a nonzero
|
||||
status after the current shard finishes, stopping that worker before later shards. These
|
||||
variables enable the generator's `--retry_failed` and `--fail_on_error` flags:
|
||||
|
||||
```sh
|
||||
RETRY_FAILED=1 FAIL_ON_ERROR=1 bash distributed_generate/launch.sh $SLURM_JOB_ID vllm TinyLlama/TinyLlama-1.1B-Chat-v1.0 /data/train/ /data/output /scripts/ 0 10 n1,n2,n3,n4
|
||||
```
|
||||
|
||||
Authentication failures, missing endpoints or models, and unexpected internal or output-write
|
||||
errors stop workers even without strict mode. The worker enables `--log_empty_conversations`; its
|
||||
`finished` marker means all inputs are saved or recorded as rejected, not that every input
|
||||
succeeded. No new marker is written while retryable failures remain. A retry can append rows
|
||||
after an older marker, so the generator reads the entire output when resuming.
|
||||
|
||||
## Combining shards
|
||||
|
||||
To combine the shards back
|
||||
|
||||
```sh
|
||||
python3 distributed_generate/sharding_utils.py --input_dir /data/output/ --output_path /data/output.jsonl --combine
|
||||
```
|
||||
|
||||
The combiner ignores the `.failures` journals and completion markers. Review the journals and
|
||||
filter conversations marked `truncated: true` before using the combined data for training.
|
||||
|
||||
@@ -190,7 +190,15 @@ if [ "$mpi_rank" -eq 0 ]; then
|
||||
if [ -n "$SYSTEM_PROMPT" ]; then
|
||||
cmd+=(--system_prompt "$SYSTEM_PROMPT")
|
||||
fi
|
||||
if [ "${FAIL_ON_ERROR:-0}" = "1" ]; then
|
||||
cmd+=(--fail_on_error)
|
||||
fi
|
||||
if [ "${RETRY_FAILED:-0}" = "1" ]; then
|
||||
cmd+=(--retry_failed)
|
||||
fi
|
||||
echo "Running: ${cmd[*]}"
|
||||
# Sample failures are journaled separately by default, so later shards still run.
|
||||
# Fatal errors and opt-in strict failures retain the worker's nonzero exit behavior.
|
||||
"${cmd[@]}"
|
||||
done
|
||||
}
|
||||
|
||||
@@ -34,7 +34,7 @@ import os
|
||||
import sys
|
||||
|
||||
import tqdm
|
||||
from openai import OpenAI
|
||||
from openai import APIConnectionError, APIStatusError, OpenAI
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--data_path", type=str, help="Path to the data file")
|
||||
@@ -44,8 +44,12 @@ parser.add_argument(
|
||||
)
|
||||
parser.add_argument("--temperature", type=float, default=0.0, help="Temperature for the model")
|
||||
parser.add_argument(
|
||||
"--max_tokens", type=int, default=2048, help="Maximum number of tokens to generate"
|
||||
"--max_tokens",
|
||||
type=int,
|
||||
default=2048,
|
||||
help="Maximum generated tokens; 0 lets the server determine the remaining context budget",
|
||||
)
|
||||
parser.add_argument("--request_timeout", type=float, default=600, help="API timeout in seconds")
|
||||
parser.add_argument("--chat", default=True, type=bool, help="Use chat mode")
|
||||
parser.add_argument("--model", type=str, default="model", help="Model name")
|
||||
parser.add_argument("--url", type=str, default="http://localhost:8000/v1", help="URL of the API")
|
||||
@@ -54,7 +58,22 @@ parser.add_argument(
|
||||
"--log_empty_conversations", action="store_true", help="Log empty conversations"
|
||||
)
|
||||
parser.add_argument("--system_prompt", nargs="+", type=str, default="", help="System prompt")
|
||||
parser.add_argument(
|
||||
"--extra_body", type=json.loads, help="JSON object of additional chat request parameters"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--fail_on_error",
|
||||
action="store_true",
|
||||
help="Exit nonzero after processing a batch with failures",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--retry_failed", action="store_true", help="Retry previously rejected conversations on resume"
|
||||
)
|
||||
args = parser.parse_args()
|
||||
if args.extra_body is not None and not isinstance(args.extra_body, dict):
|
||||
parser.error("--extra_body must be a JSON object")
|
||||
if args.max_tokens < 0:
|
||||
parser.error("--max_tokens must be nonnegative")
|
||||
|
||||
|
||||
if args.data_path.endswith("jsonl"):
|
||||
@@ -66,11 +85,26 @@ else:
|
||||
client = OpenAI(
|
||||
base_url=args.url,
|
||||
api_key=args.api_key,
|
||||
timeout=args.request_timeout,
|
||||
)
|
||||
|
||||
|
||||
def generate_data(messages, idx, system_prompt):
|
||||
class RejectedConversationError(ValueError):
|
||||
"""A conversation cannot be generated with the current input and request settings."""
|
||||
|
||||
|
||||
class RetryableConversationError(RuntimeError):
|
||||
"""A generation failure may succeed when the conversation is retried on resume."""
|
||||
|
||||
|
||||
def generate_data(sample, idx, system_prompt):
|
||||
"""Generate a complete conversation, retaining reasoning and marking truncated responses."""
|
||||
try:
|
||||
if not isinstance(sample, dict):
|
||||
raise RejectedConversationError("Expected a conversation object.")
|
||||
messages = sample.get("conversations", sample.get("messages"))
|
||||
if not isinstance(messages, list):
|
||||
raise RejectedConversationError("Expected a conversations or messages list.")
|
||||
model_name = args.model
|
||||
|
||||
if args.chat:
|
||||
@@ -81,47 +115,80 @@ def generate_data(messages, idx, system_prompt):
|
||||
system_message = {"role": "system", "content": system_prompt}
|
||||
output_messages.append(system_message)
|
||||
|
||||
for message in messages[::2]:
|
||||
for message in messages:
|
||||
# Detect message format
|
||||
if not isinstance(message, dict):
|
||||
raise RejectedConversationError("Expected a message object.")
|
||||
if "from" in message and "value" in message:
|
||||
role = message["from"].lower()
|
||||
role = message["from"]
|
||||
content = message["value"]
|
||||
elif "role" in message and "content" in message:
|
||||
role = message["role"].lower()
|
||||
role = message["role"]
|
||||
content = message["content"]
|
||||
else:
|
||||
raise ValueError(f"Message format not recognized: {message}")
|
||||
raise RejectedConversationError("Message format not recognized.")
|
||||
if not isinstance(role, str):
|
||||
raise RejectedConversationError("Expected a string message role.")
|
||||
role = role.lower()
|
||||
|
||||
if message.get("tool_calls") or message.get("function_call"):
|
||||
raise RejectedConversationError("Tool calls require a tool-execution loop.")
|
||||
if role == "system":
|
||||
if not system_prompt:
|
||||
output_messages.append({"role": "system", "content": content})
|
||||
else:
|
||||
print(
|
||||
f"Warning: conversation {idx}: --system_prompt overrides the input system message.",
|
||||
file=sys.stderr,
|
||||
)
|
||||
continue
|
||||
if role in ["assistant", "gpt"]:
|
||||
continue
|
||||
if role not in ["user", "human"]:
|
||||
return
|
||||
raise RejectedConversationError(f"Unsupported message role: {role}")
|
||||
output_messages.append(
|
||||
{
|
||||
"role": "user",
|
||||
"content": content,
|
||||
}
|
||||
)
|
||||
try:
|
||||
response = client.chat.completions.create(
|
||||
model=model_name,
|
||||
messages=output_messages,
|
||||
max_tokens=args.max_tokens,
|
||||
temperature=args.temperature,
|
||||
response = client.chat.completions.create(
|
||||
model=model_name,
|
||||
messages=output_messages,
|
||||
max_tokens=args.max_tokens or None,
|
||||
temperature=args.temperature,
|
||||
extra_body=args.extra_body,
|
||||
)
|
||||
choice = response.choices[0]
|
||||
if choice.message.tool_calls or choice.message.function_call:
|
||||
raise RejectedConversationError("Tool calls require a tool-execution loop.")
|
||||
generated_message = {
|
||||
"role": "assistant",
|
||||
"content": (choice.message.content or "").strip(),
|
||||
}
|
||||
reasoning = getattr(choice.message, "reasoning_content", None) or getattr(
|
||||
choice.message, "reasoning", None
|
||||
)
|
||||
if reasoning:
|
||||
generated_message["reasoning_content"] = reasoning
|
||||
truncated = (
|
||||
choice.finish_reason in ("length", "repetition")
|
||||
or getattr(choice, "stop_reason", None) == "repetition_detected"
|
||||
)
|
||||
if not generated_message["content"] and not truncated:
|
||||
error_type = (
|
||||
RetryableConversationError
|
||||
if args.temperature > 0
|
||||
else RejectedConversationError
|
||||
)
|
||||
choice = response.choices[0]
|
||||
response = (choice.message.content or "").strip()
|
||||
output_messages.append(
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": response,
|
||||
}
|
||||
raise error_type(
|
||||
f"Model returned an empty final answer (finish_reason={choice.finish_reason}, "
|
||||
f"reasoning_characters={len(reasoning or '')})."
|
||||
)
|
||||
if choice.finish_reason == "length":
|
||||
truncated = True
|
||||
break
|
||||
except Exception as e:
|
||||
print(e)
|
||||
output_messages.append(generated_message)
|
||||
if truncated:
|
||||
break
|
||||
if len(output_messages) == 1 or (system_prompt and len(output_messages) == 2):
|
||||
if not any(message["role"] == "assistant" for message in output_messages):
|
||||
if not args.log_empty_conversations:
|
||||
return
|
||||
to_write = {"conversation_id": idx}
|
||||
@@ -129,6 +196,8 @@ def generate_data(messages, idx, system_prompt):
|
||||
to_write = {"conversation_id": idx, "conversations": output_messages}
|
||||
if truncated:
|
||||
to_write["truncated"] = True
|
||||
to_write["finish_reason"] = choice.finish_reason
|
||||
to_write["stop_reason"] = getattr(choice, "stop_reason", None)
|
||||
with open(args.output_path, "a") as f:
|
||||
# write in share gpt format
|
||||
f.write(json.dumps(to_write) + "\n")
|
||||
@@ -158,46 +227,92 @@ def generate_data(messages, idx, system_prompt):
|
||||
to_write = {"text": prompt + response}
|
||||
f.write(json.dumps(to_write) + "\n")
|
||||
except Exception as e:
|
||||
print(e)
|
||||
print(prompt)
|
||||
print("Failed to generate data")
|
||||
raise RuntimeError(f"Failed to generate conversation {idx}") from e
|
||||
|
||||
|
||||
# if output_path exists identify the conversation_ids that have already been generated
|
||||
finished_ids = []
|
||||
finished_ids = set()
|
||||
done = False
|
||||
if os.path.exists(args.output_path):
|
||||
with open(args.output_path) as f:
|
||||
for line in f:
|
||||
outdata = json.loads(line)
|
||||
finished_ids.append(outdata.get("conversation_id", -1))
|
||||
if outdata.get("finished", False):
|
||||
done = True
|
||||
break
|
||||
finished_ids = set(finished_ids)
|
||||
if "conversation_id" in outdata:
|
||||
finished_ids.add(outdata["conversation_id"])
|
||||
done = outdata.get("finished", False)
|
||||
|
||||
# Keep the JSONL failure journal outside the shard combiner's *.jsonl input set.
|
||||
failures_path = args.output_path + ".failures"
|
||||
failures = {}
|
||||
if os.path.exists(failures_path):
|
||||
with open(failures_path) as f:
|
||||
for line in f:
|
||||
failure = json.loads(line)
|
||||
if failure["conversation_id"] not in finished_ids:
|
||||
failures[failure["conversation_id"]] = failure
|
||||
rejected_ids = {idx for idx, failure in failures.items() if not failure["retryable"]}
|
||||
if failures:
|
||||
print(f"Found {len(failures)} unresolved failures in {failures_path}.", file=sys.stderr)
|
||||
|
||||
# Ensure the output directory exists before writing to the output file
|
||||
output_dir = os.path.dirname(args.output_path)
|
||||
if output_dir and not os.path.exists(output_dir):
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
|
||||
if done:
|
||||
print("All conversations already generated")
|
||||
sys.exit()
|
||||
if (
|
||||
done
|
||||
and not (args.retry_failed and rejected_ids)
|
||||
and not any(failure["retryable"] for failure in failures.values())
|
||||
):
|
||||
print("All conversations already processed")
|
||||
sys.exit(1 if args.fail_on_error and failures else 0)
|
||||
|
||||
fatal_error = False
|
||||
with concurrent.futures.ThreadPoolExecutor(max_workers=args.num_threads) as executor:
|
||||
futures = []
|
||||
futures = {}
|
||||
system_prompt = " ".join(args.system_prompt)
|
||||
|
||||
for idx, sample in enumerate(data):
|
||||
if idx in finished_ids:
|
||||
if idx in finished_ids or (idx in rejected_ids and not args.retry_failed):
|
||||
continue
|
||||
future = executor.submit(generate_data, sample["conversations"], idx, system_prompt)
|
||||
futures.append(future)
|
||||
future = executor.submit(generate_data, sample, idx, system_prompt)
|
||||
futures[future] = idx
|
||||
|
||||
for future in tqdm.tqdm(concurrent.futures.as_completed(futures), total=len(futures)):
|
||||
future.result()
|
||||
idx = futures[future]
|
||||
try:
|
||||
future.result()
|
||||
except Exception as exc:
|
||||
cause = exc.__cause__ or exc
|
||||
status = cause.status_code if isinstance(cause, APIStatusError) else None
|
||||
rejected = isinstance(cause, RejectedConversationError) or status in (400, 422)
|
||||
transient = isinstance(cause, (APIConnectionError, RetryableConversationError)) or (
|
||||
status is not None and (status in (408, 409, 429) or status >= 500)
|
||||
)
|
||||
fatal_error |= not (rejected or transient)
|
||||
failure = {
|
||||
"conversation_id": idx,
|
||||
"retryable": not rejected,
|
||||
"error_type": type(cause).__name__,
|
||||
"error": str(cause),
|
||||
}
|
||||
failures[idx] = failure
|
||||
with open(failures_path, "a") as f:
|
||||
f.write(json.dumps(failure) + "\n")
|
||||
print(f"Failed conversation {idx}: {cause}", file=sys.stderr)
|
||||
else:
|
||||
failures.pop(idx, None)
|
||||
|
||||
if args.log_empty_conversations:
|
||||
if failures:
|
||||
retryable = sum(failure["retryable"] for failure in failures.values())
|
||||
print(
|
||||
f"{len(failures)} conversations failed ({retryable} retryable); see {failures_path}.",
|
||||
file=sys.stderr,
|
||||
)
|
||||
|
||||
if args.log_empty_conversations and not any(failure["retryable"] for failure in failures.values()):
|
||||
with open(args.output_path, "a") as f:
|
||||
f.write(json.dumps({"finished": True}) + "\n")
|
||||
|
||||
if fatal_error or (args.fail_on_error and failures):
|
||||
sys.exit(1)
|
||||
|
||||
Reference in New Issue
Block a user