From 278012caa0bbee0b090e8e6da670a8f4480e8cba Mon Sep 17 00:00:00 2001 From: jiapingW <56055330+jiapingW@users.noreply.github.com> Date: Tue, 30 Dec 2025 02:06:11 +0800 Subject: [PATCH] [Feature] support bench jsonl files with sharegpt format (#15057) Co-authored-by: Xinyuan Tong <115166877+JustinTong0323@users.noreply.github.com> --- python/sglang/bench_serving.py | 108 +++++++++++++++++++++++++++++++++ 1 file changed, 108 insertions(+) diff --git a/python/sglang/bench_serving.py b/python/sglang/bench_serving.py index e6b3b1149..cf62da0a7 100644 --- a/python/sglang/bench_serving.py +++ b/python/sglang/bench_serving.py @@ -870,6 +870,17 @@ def get_dataset(args, tokenizer, model_id=None): # Limit the number of requests based on --num-prompts input_requests = all_requests_data[: args.num_prompts] + elif args.dataset_name == "custom": + assert not tokenize_prompt + input_requests = sample_custom_requests( + dataset_path=args.dataset_path, + num_requests=args.num_prompts, + tokenizer=tokenizer, + fixed_output_len=args.sharegpt_output_len, + context_len=args.sharegpt_context_len, + prompt_suffix=args.prompt_suffix, + apply_chat_template=args.apply_chat_template, + ) else: raise ValueError(f"Unknown dataset: {args.dataset_name}") return input_requests @@ -1274,6 +1285,102 @@ def sample_sharegpt_requests( return filtered_dataset +def sample_custom_requests( + dataset_path: str, + num_requests: int, + tokenizer: PreTrainedTokenizerBase, + fixed_output_len: Optional[int] = None, + context_len: Optional[int] = None, + prompt_suffix: Optional[str] = "", + apply_chat_template=False, +) -> List[DatasetRow]: + """ + Sample requests from a custom JSONL dataset: supports 'content'/'value' as conversation keys. + """ + if fixed_output_len is not None and fixed_output_len < 4: + raise ValueError("output_len too small") + + # Load the dataset + dataset = [] + if not os.path.isfile(dataset_path): + raise FileNotFoundError(f"Dataset not found at {dataset_path}") + + with open(dataset_path, "r", encoding="utf-8") as f: + for line in f: + line = line.strip() + if line: # skip empty lines + try: + dataset.append(json.loads(line)) + except json.JSONDecodeError: + continue # skip lines with JSON errors + + # Filter out the conversations with less than 2 turns. + processed_dataset = [] + for data in dataset: + convs = data.get("conversations", data.get("conversation", [])) + if len(convs) >= 2: + user_turn = convs[0].get("content", convs[0].get("value", "")) + assist_turn = convs[1].get("content", convs[1].get("value", "")) + processed_dataset.append((user_turn, assist_turn)) + dataset = processed_dataset + random.shuffle(dataset) + + # Filter out sequences that are too long or too short + filtered_dataset: List[DatasetRow] = [] + + for i in range(len(dataset)): + if len(filtered_dataset) == num_requests: + break + + # Tokenize the prompts and completions. + prompt = dataset[i][0] + + if prompt_suffix: + prompt = ( + remove_suffix(prompt, ASSISTANT_SUFFIX) + + prompt_suffix + + ASSISTANT_SUFFIX + ) + + if apply_chat_template: + prompt = tokenizer.apply_chat_template( + [{"role": "user", "content": prompt}], + add_generation_prompt=True, + tokenize=False, + return_dict=False, + ) + if tokenizer.bos_token: + prompt = prompt.replace(tokenizer.bos_token, "") + + prompt_token_ids = tokenizer.encode(prompt) + completion = dataset[i][1] + completion_token_ids = tokenizer.encode(completion) + prompt_len = len(prompt_token_ids) + output_len = ( + len(completion_token_ids) if fixed_output_len is None else fixed_output_len + ) + + if prompt_len < 2 or output_len < 2: + # Prune too short sequences. + continue + + if context_len and prompt_len + output_len > context_len: + # Prune too long sequences. + continue + + filtered_dataset.append( + DatasetRow( + prompt=prompt, + prompt_len=prompt_len, + output_len=output_len, + ) + ) + + print(f"#Input tokens: {np.sum([x.prompt_len for x in filtered_dataset])}") + print(f"#Output tokens: {np.sum([x.output_len for x in filtered_dataset])}") + return filtered_dataset + + def compute_random_lens(full_len: int, range_ratio: float, num: int): return np.random.randint( max(int(full_len * range_ratio), 1), @@ -2687,6 +2794,7 @@ if __name__ == "__main__": default="sharegpt", choices=[ "sharegpt", + "custom", "random", "random-ids", "generated-shared-prefix",