From 4028a73c10835ad40b9bd47b7dd744629da23aee Mon Sep 17 00:00:00 2001 From: Vladislav Nosivskoy Date: Mon, 20 Apr 2026 12:34:38 +0300 Subject: [PATCH] [KV-Events] Fix kv events events publishing for CP (#22983) Signed-off-by: Vladislav Nosivskoy --- .../observability/scheduler_metrics_mixin.py | 4 +- test/manual/test_kv_events.py | 132 ++++++++++++++++++ 2 files changed, 135 insertions(+), 1 deletion(-) diff --git a/python/sglang/srt/observability/scheduler_metrics_mixin.py b/python/sglang/srt/observability/scheduler_metrics_mixin.py index 30a8a6802..0c9da952e 100644 --- a/python/sglang/srt/observability/scheduler_metrics_mixin.py +++ b/python/sglang/srt/observability/scheduler_metrics_mixin.py @@ -171,7 +171,9 @@ class SchedulerMetricsMixin: ) def init_kv_events(self: Scheduler, kv_events_config: Optional[str]): - self.enable_kv_cache_events = bool(kv_events_config and self.attn_tp_rank == 0) + self.enable_kv_cache_events = bool( + kv_events_config and self.attn_tp_rank == 0 and self.attn_cp_rank == 0 + ) if self.enable_kv_cache_events: self.kv_event_publisher = EventPublisherFactory.create( diff --git a/test/manual/test_kv_events.py b/test/manual/test_kv_events.py index 0f657333c..95367cef0 100644 --- a/test/manual/test_kv_events.py +++ b/test/manual/test_kv_events.py @@ -21,6 +21,8 @@ from sglang.test.test_utils import ( popen_launch_server, ) +QWEN3_30B_MODEL_PATH = "Qwen/Qwen3-30B-A3B-FP8" + class TestKvEvents(CustomTestCase): def test_kv_events_enabled(self): @@ -287,6 +289,136 @@ class TestKvEvents(CustomTestCase): context.term() kill_process_tree(process.pid) + def test_kv_events_attn_cp_single_stream_per_dp_rank(self): + """Test that CP replicas do not publish duplicate KV events for one DP rank.""" + + decoder = Decoder(type=KVEventBatch) + context = zmq.Context() + + sub_dp0 = context.socket(zmq.SUB) + sub_dp0.connect("tcp://localhost:5557") + topic = "kv-events" + sub_dp0.setsockopt_string(zmq.SUBSCRIBE, topic) + + # There is only one DP rank in this test, so CP must not create another stream. + sub_unexpected = context.socket(zmq.SUB) + sub_unexpected.connect("tcp://localhost:5558") + sub_unexpected.setsockopt_string(zmq.SUBSCRIBE, topic) + + process = popen_launch_server( + QWEN3_30B_MODEL_PATH, + DEFAULT_URL_FOR_TEST, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + other_args=[ + "--kv-events-config", + '{"publisher": "zmq", "topic": "kv-events"}', + "--tp-size", + 2, + "--attn-cp-size", + 2, + "--moe-dp-size", + 2, + "--enable-prefill-context-parallel", + "--trust-remote-code", + "--max-total-tokens", + 4096, + "--max-running-requests", + 4, + "--disable-cuda-graph", + "--cuda-graph-max-bs", + 4, + "--model-loader-extra-config", + '{"enable_multithread_load": true, "num_threads": 64}', + ], + ) + + try: + response = requests.get(f"{DEFAULT_URL_FOR_TEST}/health_generate") + self.assertEqual(response.status_code, 200) + + for i in range(4): + response = requests.post( + f"{DEFAULT_URL_FOR_TEST}/generate", + json={ + "text": ( + f"KV event context parallelism request {i}: " + "write a concise fact about distributed inference." + ), + "sampling_params": { + "temperature": 0, + "max_new_tokens": 16, + }, + }, + ) + self.assertEqual(response.status_code, 200) + + batches = [] + stored_hashes = set() + duplicate_hashes = set() + unexpected_batches = [] + start = time.time() + max_wait_s = 15 + min_stored_blocks = 3 + + while (time.time() - start) < max_wait_s and ( + len(stored_hashes) < min_stored_blocks + ): + if sub_dp0.poll(timeout=100): + _, seq_bytes, payload = sub_dp0.recv_multipart() + event_batch = decoder.decode(payload) + print( + f"DP Rank 0 - EventBatch: ts={event_batch.ts}, " + f"attn_dp_rank={event_batch.attn_dp_rank}" + ) + self.assertEqual( + event_batch.attn_dp_rank, + 0, + "CP mode with one DP rank should publish events as attn_dp_rank=0", + ) + batches.append(event_batch) + + for event in event_batch.events: + print(f" DP0 - {event}") + self.assertIsInstance( + event, + (BlockStored, BlockRemoved, AllBlocksCleared), + f"Event should be a KV cache event, got {type(event)}", + ) + if isinstance(event, BlockStored): + for block_hash in event.block_hashes: + if block_hash in stored_hashes: + duplicate_hashes.add(block_hash) + stored_hashes.add(block_hash) + + if sub_unexpected.poll(timeout=0): + _, seq_bytes, payload = sub_unexpected.recv_multipart() + unexpected_batches.append(decoder.decode(payload)) + + self.assertGreater( + len(batches), 0, "Should have received KV cache event batches" + ) + self.assertGreaterEqual( + len(stored_hashes), + min_stored_blocks, + f"Expected at least {min_stored_blocks} stored KV blocks", + ) + self.assertEqual( + unexpected_batches, + [], + "CP ranks within one DP rank should not create a second KV event stream", + ) + self.assertEqual( + duplicate_hashes, + set(), + "CP ranks should not publish duplicate BlockStored events for replicated KV blocks", + ) + + finally: + sub_dp0.close() + sub_unexpected.close() + context.term() + kill_process_tree(process.pid) + if __name__ == "__main__": unittest.main()