From a0a8f1473c07921fca0469d922a8c745fb79419a Mon Sep 17 00:00:00 2001 From: Alison Shao <54658187+alisonshao@users.noreply.github.com> Date: Wed, 25 Feb 2026 18:45:54 -0800 Subject: [PATCH] [Benchmark] Fix generated_shared_prefix attribute naming and remove args dependency (#19363) Co-authored-by: Alison Shao Co-authored-by: sglang-bot --- benchmark/hicache/data_processing.py | 10 ++- .../datasets/generated_shared_prefix.py | 67 ++++++++++++------- .../bench_fn/test_benchmark_datasets_api.py | 13 +++- 3 files changed, 61 insertions(+), 29 deletions(-) diff --git a/benchmark/hicache/data_processing.py b/benchmark/hicache/data_processing.py index 6695dab36..8c4b8cd1b 100644 --- a/benchmark/hicache/data_processing.py +++ b/benchmark/hicache/data_processing.py @@ -442,7 +442,15 @@ def sample_generated_shared_prefix_requests( disable_shuffle: bool = False, ) -> SampleOutput: """Generate benchmark requests with shared system prompts using random tokens and caching.""" - cache_path = get_gen_prefix_cache_path(args, tokenizer) + cache_path = get_gen_prefix_cache_path( + args.seed, + num_groups, + prompts_per_group, + system_prompt_len, + question_len, + output_len, + tokenizer, + ) # Try to load from cache first if cache_path.exists(): diff --git a/python/sglang/benchmark/datasets/generated_shared_prefix.py b/python/sglang/benchmark/datasets/generated_shared_prefix.py index ab10524ad..51c4e18ae 100644 --- a/python/sglang/benchmark/datasets/generated_shared_prefix.py +++ b/python/sglang/benchmark/datasets/generated_shared_prefix.py @@ -1,4 +1,3 @@ -import argparse import pickle import random import uuid @@ -29,10 +28,10 @@ class GeneratedSharedPrefixDataset(BaseDataset): output_len: int range_ratio: float seed: int - gsp_fast_prepare: bool - gsp_send_routing_key: bool - gsp_num_turns: int - gsp_ordered: bool + fast_prepare: bool + send_routing_key: bool + num_turns: int + ordered: bool @classmethod def from_args(cls, args: Namespace) -> "GeneratedSharedPrefixDataset": @@ -45,10 +44,10 @@ class GeneratedSharedPrefixDataset(BaseDataset): output_len=args.gsp_output_len, range_ratio=getattr(args, "gsp_range_ratio", 1.0), seed=args.seed, - gsp_fast_prepare=getattr(args, "gsp_fast_prepare", False), - gsp_send_routing_key=getattr(args, "gsp_send_routing_key", False), - gsp_num_turns=getattr(args, "gsp_num_turns", 1), - gsp_ordered=getattr(args, "gsp_ordered", False), + fast_prepare=getattr(args, "gsp_fast_prepare", False), + send_routing_key=getattr(args, "gsp_send_routing_key", False), + num_turns=getattr(args, "gsp_num_turns", 1), + ordered=getattr(args, "gsp_ordered", False), ) def load( @@ -62,18 +61,29 @@ class GeneratedSharedPrefixDataset(BaseDataset): output_len=self.output_len, range_ratio=self.range_ratio, tokenizer=tokenizer, - args=self, + seed=self.seed, + send_routing_key=self.send_routing_key, + num_turns=self.num_turns, + fast_prepare=self.fast_prepare, + ordered=self.ordered, ) -def get_gen_prefix_cache_path(args, tokenizer): +def get_gen_prefix_cache_path( + seed: int, + num_groups: int, + prompts_per_group: int, + system_prompt_len: int, + question_len: int, + output_len: int, + tokenizer, +): """Create cache directory under ~/.cache/sglang/benchmark""" cache_dir = Path.home() / ".cache" / "sglang" / "benchmark" - # Create a unique cache filename based on the generation parameters cache_key = ( - f"gen_shared_prefix_{args.seed}_{args.gsp_num_groups}_{args.gsp_prompts_per_group}_" - f"{args.gsp_system_prompt_len}_{args.gsp_question_len}_{args.gsp_output_len}_" + f"gen_shared_prefix_{seed}_{num_groups}_{prompts_per_group}_" + f"{system_prompt_len}_{question_len}_{output_len}_" f"{tokenizer.__class__.__name__}.pkl" ) return cache_dir / cache_key @@ -87,13 +97,22 @@ def sample_generated_shared_prefix_requests( output_len: int, range_ratio: float, tokenizer: PreTrainedTokenizerBase, - args: argparse.Namespace, + seed: int, + send_routing_key: bool = False, + num_turns: int = 1, + fast_prepare: bool = False, + ordered: bool = False, ) -> List[DatasetRow]: """Generate benchmark requests with shared system prompts using random tokens and caching.""" - send_routing_key = getattr(args, "gsp_send_routing_key", False) - num_turns = getattr(args, "gsp_num_turns", 1) - - cache_path = get_gen_prefix_cache_path(args, tokenizer) + cache_path = get_gen_prefix_cache_path( + seed, + num_groups, + prompts_per_group, + system_prompt_len, + question_len, + output_len, + tokenizer, + ) should_cache = (range_ratio == 1) and not send_routing_key and num_turns == 1 # Try to load from cache first @@ -168,11 +187,7 @@ def sample_generated_shared_prefix_requests( 1: ] full_prompt = turn_prompts[0] if num_turns == 1 else turn_prompts - prompt_len = ( - 1 - if getattr(args, "gsp_fast_prepare", False) - else len(tokenizer.encode(turn_prompts[0])) - ) + prompt_len = 1 if fast_prepare else len(tokenizer.encode(turn_prompts[0])) output_len_val = int(output_lens[group_idx, prompt_idx]) input_requests.append( @@ -186,7 +201,7 @@ def sample_generated_shared_prefix_requests( total_input_tokens += prompt_len total_output_tokens += output_len_val - if not getattr(args, "gsp_ordered", False): + if not ordered: random.shuffle(input_requests) # Print statistics @@ -195,7 +210,7 @@ def sample_generated_shared_prefix_requests( print(f"Prompts per group: {prompts_per_group}") print(f"Number of turns: {num_turns}") print(f"Total prompts: {len(input_requests)}") - if not getattr(args, "gsp_fast_prepare", False): + if not fast_prepare: print(f"Total input tokens: {total_input_tokens}") print(f"Total output tokens: {total_output_tokens}") print( diff --git a/test/registered/bench_fn/test_benchmark_datasets_api.py b/test/registered/bench_fn/test_benchmark_datasets_api.py index a5aac8a37..96f781a22 100644 --- a/test/registered/bench_fn/test_benchmark_datasets_api.py +++ b/test/registered/bench_fn/test_benchmark_datasets_api.py @@ -294,7 +294,7 @@ class TestBenchmarkDatasetsAPI(unittest.TestCase): self.assertIn("tools", rows[1].extra_request_body) def test_generated_shared_prefix_sampler(self): - args = make_args(gsp_range_ratio=0.0, gsp_num_groups=2, gsp_prompts_per_group=2) + args = make_args(gsp_num_groups=2, gsp_prompts_per_group=2) rows = sample_generated_shared_prefix_requests( num_groups=args.gsp_num_groups, prompts_per_group=args.gsp_prompts_per_group, @@ -303,7 +303,7 @@ class TestBenchmarkDatasetsAPI(unittest.TestCase): output_len=args.gsp_output_len, range_ratio=args.gsp_range_ratio, tokenizer=self.tokenizer, - args=args, + seed=args.seed, ) self.assertEqual(len(rows), 4) self.assertTrue(all(isinstance(row, DatasetRow) for row in rows)) @@ -416,6 +416,15 @@ class TestBenchmarkDatasetsAPI(unittest.TestCase): mmmu_rows = get_dataset(mmmu_args, self.tokenizer, model_id="dummy-model") self.assertEqual(len(mmmu_rows), 1) + gsp_args = make_args( + dataset_name="generated-shared-prefix", + gsp_num_groups=2, + gsp_prompts_per_group=2, + ) + gsp_rows = get_dataset(gsp_args, self.tokenizer, model_id="dummy-model") + self.assertEqual(len(gsp_rows), 4) + self.assertTrue(all(isinstance(row, DatasetRow) for row in gsp_rows)) + def test_get_dataset_unknown_dataset(self): args = make_args(dataset_name="not-a-dataset") with self.assertRaises(ValueError):