import importlib.util import sys import unittest from pathlib import Path _REPO_ROOT = Path(__file__).resolve().parents[4] sys.path.insert(0, str(_REPO_ROOT / "python")) _BENCH_PATH = ( _REPO_ROOT / "benchmark" / "hicache" / "bench_prefill_scheduler_admission.py" ) def _load_bench_module(): spec = importlib.util.spec_from_file_location( "bench_prefill_scheduler_admission", _BENCH_PATH ) module = importlib.util.module_from_spec(spec) assert spec.loader is not None sys.modules[spec.name] = module spec.loader.exec_module(module) return module class TestPrefillSchedulerAdmissionBench(unittest.TestCase): def test_l1_l2_and_extend_tokens_are_reported_with_real_prefill_semantics(self): bench = _load_bench_module() cfg = bench.SchedulerBenchConfig( page_size=64, available_tokens=10000, evictable_tokens=0, max_prefill_tokens=1024, enable_cp_shared_kv_prefill_bs_gt1=True, cp_shared_kv_prefill_max_batch_requests=8, cp_shared_kv_prefill_max_total_extend_tokens=4096, ) trace = bench.run_scheduler_admission_trace( [ bench.RequestSpec( rid="hit-l1-l2", l1_cached_tokens=128, l2_cached_tokens=128, extend_tokens=64, max_new_tokens=1, ) ], cfg, ) self.assertEqual(len(trace.ticks), 1) tick = trace.ticks[0] self.assertEqual([req.rid for req in tick.accepted], ["hit-l1-l2"]) accepted = tick.accepted[0] self.assertEqual(accepted.initial_extend_tokens, 192) self.assertEqual(accepted.effective_extend_tokens, 64) self.assertEqual(accepted.l1_cached_tokens, 128) self.assertEqual(accepted.l2_cached_tokens, 128) self.assertEqual(accepted.loaded_l2_tokens, 128) self.assertEqual(tick.log_hit_tokens, 256) self.assertEqual(tick.log_input_tokens, 64) def test_cp_total_extend_limit_controls_batching_not_generic_max_prefill_tokens(self): bench = _load_bench_module() cfg = bench.SchedulerBenchConfig( page_size=64, available_tokens=10000, evictable_tokens=0, max_prefill_tokens=192, enable_cp_shared_kv_prefill_bs_gt1=True, cp_shared_kv_prefill_max_batch_requests=8, cp_shared_kv_prefill_max_total_extend_tokens=256, ) trace = bench.run_scheduler_admission_trace( [ bench.RequestSpec("a", 0, 0, 128, max_new_tokens=1), bench.RequestSpec("b", 0, 0, 128, max_new_tokens=1), ], cfg, ) self.assertEqual(len(trace.ticks), 1) self.assertEqual([req.rid for req in trace.ticks[0].accepted], ["a", "b"]) self.assertEqual(trace.ticks[0].cp_total_extend_tokens, 256) def test_l2_load_back_consumes_l1_capacity_and_can_stop_later_requests(self): bench = _load_bench_module() cfg = bench.SchedulerBenchConfig( page_size=64, available_tokens=320, evictable_tokens=0, max_prefill_tokens=4096, enable_cp_shared_kv_prefill_bs_gt1=True, cp_shared_kv_prefill_max_batch_requests=8, cp_shared_kv_prefill_max_total_extend_tokens=4096, ) trace = bench.run_scheduler_admission_trace( [ bench.RequestSpec("l2-heavy", 0, 128, 64, max_new_tokens=1), bench.RequestSpec("next", 0, 0, 64, max_new_tokens=1), ], cfg, ) self.assertEqual([req.rid for req in trace.ticks[0].accepted], ["l2-heavy"]) self.assertEqual(trace.ticks[0].stopped_on_rid, "next") self.assertEqual(trace.ticks[0].stopped_result, "NO_TOKEN") self.assertEqual(trace.ticks[0].allocator_available_after_tick, 192) self.assertEqual(trace.ticks[0].load_back_events[0].loaded_tokens, 128) def test_total_cached_limit_is_observable_in_scheduler_trace(self): bench = _load_bench_module() cfg = bench.SchedulerBenchConfig( page_size=64, available_tokens=10000, evictable_tokens=0, max_prefill_tokens=4096, enable_cp_shared_kv_prefill_bs_gt1=True, cp_shared_kv_prefill_max_batch_requests=8, cp_shared_kv_prefill_max_total_extend_tokens=4096, cp_shared_kv_prefill_max_total_cached_tokens=4096, ) trace = bench.run_scheduler_admission_trace( [ bench.RequestSpec("a", 4096, 0, 64, max_new_tokens=1), bench.RequestSpec("b", 4096, 0, 64, max_new_tokens=1), ], cfg, ) self.assertEqual([req.rid for req in trace.ticks[0].accepted], ["a"]) self.assertEqual(trace.ticks[0].stopped_on_rid, "b") self.assertEqual(trace.ticks[0].stopped_result, "OTHER") self.assertEqual(trace.ticks[0].cp_total_cached_tokens, 4096) def test_max_buffer_size_is_observable_in_scheduler_trace(self): bench = _load_bench_module() cfg = bench.SchedulerBenchConfig( page_size=64, available_tokens=10000, evictable_tokens=0, max_prefill_tokens=4096, enable_cp_shared_kv_prefill_bs_gt1=True, cp_shared_kv_prefill_max_batch_requests=8, cp_shared_kv_prefill_max_total_extend_tokens=4096, cp_shared_kv_prefill_max_total_cached_tokens=4096, cp_shared_kv_prefill_max_buffer_size=1, kv_cache_dim=1, kv_dtype_bytes=2, index_head_dim=8, vocab_size=16, tp_size=1, ) trace = bench.run_scheduler_admission_trace( [ bench.RequestSpec("a", 0, 0, 64, max_new_tokens=1), bench.RequestSpec("b", 0, 0, 64, max_new_tokens=1), ], cfg, ) tick = trace.ticks[0] self.assertEqual([req.rid for req in tick.accepted], ["a"]) self.assertEqual(tick.stopped_on_rid, "b") self.assertEqual(tick.stopped_result, "OTHER") self.assertGreater(tick.cp_estimated_peak_buffer_bytes, 1) self.assertIn("layer_forward_peak_bytes", tick.cp_buffer_breakdown) if __name__ == "__main__": unittest.main()