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:
sychen52
2026-10-01 00:37:08 +00:00
committed by GitHub
parent bc5d2c3610
commit 333ace1bc9
5 changed files with 617 additions and 47 deletions
+45 -2
View File
@@ -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)