Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 8 additions & 1 deletion tensorrt_llm/_torch/pyexecutor/kv_cache_transceiver.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,8 @@

from .llm_request import LlmRequest
from .mamba_cache_manager import (BaseMambaCacheManager,
CppMambaHybridCacheManager)
CppMambaHybridCacheManager,
MixedMambaHybridCacheManager)
from .resource_manager import KVCacheManager

CacheTransceiverCpp = tensorrt_llm.bindings.internal.batch_manager.CacheTransceiver
Expand Down Expand Up @@ -129,6 +130,12 @@ def create_kv_cache_transceiver(
if cache_transceiver_config.transceiver_runtime == "auto":
cache_transceiver_config.transceiver_runtime = None

if (cache_transceiver_config.transceiver_runtime != "PYTHON"
and isinstance(mamba_cache_manager, MixedMambaHybridCacheManager)):
raise ValueError(
"MixedMambaHybridCacheManager requires the Python transceiver "
"runtime with the NIXL backend in disaggregated serving.")

_validate_disagg_inflight_cancel_config(cache_transceiver_config)

if cache_transceiver_config.backend == "DEFAULT":
Expand Down
110 changes: 83 additions & 27 deletions tests/unittest/others/test_kv_cache_transceiver.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,6 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0

import gc
import multiprocessing
import sys
Expand Down Expand Up @@ -33,6 +36,23 @@
KV_TRANSFER_COMPLETION_MARGIN_S = 10.0


@pytest.mark.parametrize("transceiver_runtime", ["CPP", "auto"])
def test_cpp_transceiver_rejects_mixed_mamba_manager(transceiver_runtime):
config = CacheTransceiverConfig(backend="NIXL",
transceiver_runtime=transceiver_runtime)
mixed_manager = object.__new__(MixedMambaHybridCacheManager)

with pytest.raises(
ValueError,
match="MixedMambaHybridCacheManager requires the Python"):
create_kv_cache_transceiver(mapping=None,
dist=None,
kv_cache_manager=None,
attention_type=AttentionTypeCpp.DEFAULT,
cache_transceiver_config=config,
mamba_cache_manager=mixed_manager)


def create_kv_cache_manager(mapping,
dtype,
max_tokens=256,
Expand Down Expand Up @@ -716,7 +736,7 @@ def create_hybrid_cache_manager(mapping,
dtype,
mamba_conv_dtype=torch.float16,
mamba_ssm_dtype=torch.float16):
"""Create a MixedMambaHybridCacheManager for testing hybrid models.
"""Create a mixed hybrid manager for Python transceiver tests.

This manager handles both KV cache (attention layers) and Mamba cache (RNN layers).

Expand Down Expand Up @@ -790,7 +810,7 @@ def hybrid_dtypes(request):
Returns (kv_dtype, mamba_conv_dtype, mamba_ssm_dtype) based on the parametrized string.

KV dtype: fp8, bf16
Conv dtype: fp8, bf16, fp32
Conv dtype: bf16, fp32
SSM dtype: bf16, fp32
"""
kv_dtype_str, conv_dtype_str, ssm_dtype_str = request.param
Expand All @@ -813,7 +833,6 @@ def hybrid_dtypes(request):


@pytest.mark.timeout(120)
@pytest.mark.parametrize("backend", ["NIXL", "UCX"], ids=["NIXL", "UCX"])
@pytest.mark.parametrize(
"hybrid_dtypes",
[
Expand All @@ -839,18 +858,20 @@ def hybrid_dtypes(request):
],
indirect=["hybrid_dtypes"],
)
def test_hybrid_cache_transceiver_single_process(backend, hybrid_dtypes):
def test_hybrid_cache_transceiver_single_process(hybrid_dtypes, request):
mapping = Mapping(world_size=1, rank=0)
kv_dtype, mamba_conv_dtype, mamba_ssm_dtype = hybrid_dtypes

# Create hybrid cache managers (combines KV + Mamba) for context and generation
hybrid_cache_manager_ctx = create_hybrid_cache_manager(
mapping, kv_dtype, mamba_conv_dtype, mamba_ssm_dtype)
request.addfinalizer(hybrid_cache_manager_ctx.shutdown)
hybrid_cache_manager_gen = create_hybrid_cache_manager(
mapping, kv_dtype, mamba_conv_dtype, mamba_ssm_dtype)
request.addfinalizer(hybrid_cache_manager_gen.shutdown)

cache_transceiver_config = CacheTransceiverConfig(backend=backend,
max_tokens_in_buffer=512)
cache_transceiver_config = CacheTransceiverConfig(
backend="NIXL", transceiver_runtime="PYTHON", max_tokens_in_buffer=512)
dist = Distributed.get(mapping)

# Create transceivers - the hybrid manager serves as both kv_cache_manager and mamba_cache_manager
Expand All @@ -861,6 +882,7 @@ def test_hybrid_cache_transceiver_single_process(backend, hybrid_dtypes):
AttentionTypeCpp.DEFAULT,
cache_transceiver_config,
mamba_cache_manager=hybrid_cache_manager_ctx)
request.addfinalizer(cache_transceiver_ctx.shutdown)

cache_transceiver_gen = create_kv_cache_transceiver(
mapping,
Expand All @@ -869,6 +891,7 @@ def test_hybrid_cache_transceiver_single_process(backend, hybrid_dtypes):
AttentionTypeCpp.DEFAULT,
cache_transceiver_config,
mamba_cache_manager=hybrid_cache_manager_gen)
request.addfinalizer(cache_transceiver_gen.shutdown)

# Fill both KV and Mamba cache buffers with random data
fill_hybrid_cache_buffers(hybrid_cache_manager_ctx)
Expand All @@ -883,6 +906,9 @@ def test_hybrid_cache_transceiver_single_process(backend, hybrid_dtypes):
sampling_params._get_sampling_config()),
is_streaming=False,
llm_request_type=LlmRequestType.LLMREQUEST_TYPE_CONTEXT_ONLY)
ctx_request.py_disaggregated_params = tensorrt_llm.DisaggregatedParams(
request_type="context_only",
disagg_request_id=uuid.uuid4().int & 0x7FFFFFFFFFFFFFFF)

# Prepare resources for hybrid manager (handles both KV and Mamba)
scheduled_ctx = ScheduledRequests()
Expand All @@ -902,6 +928,14 @@ def test_hybrid_cache_transceiver_single_process(backend, hybrid_dtypes):
is_streaming=False,
llm_request_type=LlmRequestType.LLMREQUEST_TYPE_GENERATION_ONLY,
context_phase_params=ctx_request.context_phase_params)
gen_request.py_disaggregated_params = tensorrt_llm.DisaggregatedParams(
request_type="generation_only",
disagg_request_id=ctx_request.py_disaggregated_params.disagg_request_id,
ctx_request_id=ctx_request.request_id,
ctx_dp_rank=ctx_request.context_phase_params.ctx_dp_rank,
ctx_info_endpoint=ctx_request.context_phase_params.disagg_info_endpoint,
first_gen_tokens=ctx_request.context_phase_params.first_gen_tokens,
draft_tokens=ctx_request.context_phase_params.draft_tokens)

# Prepare resources for hybrid manager on gen side
scheduled_gen = ScheduledRequests()
Expand All @@ -911,6 +945,7 @@ def test_hybrid_cache_transceiver_single_process(backend, hybrid_dtypes):
cache_transceiver_gen.request_and_receive_async(gen_request)

completed_ctx_ids = set()
expected_ctx_id = get_context_completed_request_id(ctx_request, "PYTHON")

def poll_transfers():
completed, failed = cache_transceiver_ctx.check_context_transfer_status(
Expand All @@ -920,8 +955,7 @@ def poll_transfers():
cache_transceiver_gen.check_gen_transfer_status(1)

def transfers_done():
return (ctx_request.py_request_id in completed_ctx_ids
and gen_request.state
return (expected_ctx_id in completed_ctx_ids and gen_request.state
== LlmRequestState.DISAGG_GENERATION_TRANS_COMPLETE)

wait_for_transfer_completion(poll_transfers, transfers_done)
Expand All @@ -934,10 +968,10 @@ def transfers_done():
# independently-allocated slots on each side, so we check the
# request's own slot instead of the full state buffer (which has
# extra padding-dummy slots that only the ctx side touched).
slot_ctx = hybrid_cache_manager_ctx._impl.mamba_impl.get_cache_index(
ctx_request.py_request_id)
slot_gen = hybrid_cache_manager_gen._impl.mamba_impl.get_cache_index(
gen_request.py_request_id)
slot_ctx = hybrid_cache_manager_ctx.mamba_cache_index[
ctx_request.py_request_id]
slot_gen = hybrid_cache_manager_gen.mamba_cache_index[
gen_request.py_request_id]
assert torch.equal(
hybrid_cache_manager_gen.get_conv_states(1)[slot_gen],
hybrid_cache_manager_ctx.get_conv_states(1)[slot_ctx]), (
Expand All @@ -947,20 +981,22 @@ def transfers_done():
hybrid_cache_manager_gen.get_ssm_states(1)[slot_gen],
hybrid_cache_manager_ctx.get_ssm_states(1)[slot_ctx]), (
"different mamba ssm states")
shutdown_transceivers(cache_transceiver_gen, cache_transceiver_ctx)


@pytest.mark.timeout(120)
@pytest.mark.parametrize("backend", ["NIXL", "UCX"], ids=["NIXL", "UCX"])
def test_hybrid_cache_transceiver_cancel_request(backend):
def test_hybrid_cache_transceiver_cancel_request(request):

mapping = Mapping(world_size=1, rank=0)
dtype = DataType.HALF

hybrid_cache_manager_ctx = create_hybrid_cache_manager(mapping, dtype)
request.addfinalizer(hybrid_cache_manager_ctx.shutdown)
hybrid_cache_manager_gen = create_hybrid_cache_manager(mapping, dtype)
request.addfinalizer(hybrid_cache_manager_gen.shutdown)

cache_transceiver_config = CacheTransceiverConfig(backend=backend,
max_tokens_in_buffer=512)
cache_transceiver_config = CacheTransceiverConfig(
backend="NIXL", transceiver_runtime="PYTHON", max_tokens_in_buffer=512)
dist = Distributed.get(mapping)

cache_transceiver_ctx = create_kv_cache_transceiver(
Expand All @@ -970,6 +1006,7 @@ def test_hybrid_cache_transceiver_cancel_request(backend):
AttentionTypeCpp.DEFAULT,
cache_transceiver_config,
mamba_cache_manager=hybrid_cache_manager_ctx)
request.addfinalizer(cache_transceiver_ctx.shutdown)

cache_transceiver_gen = create_kv_cache_transceiver(
mapping,
Expand All @@ -978,6 +1015,7 @@ def test_hybrid_cache_transceiver_cancel_request(backend):
AttentionTypeCpp.DEFAULT,
cache_transceiver_config,
mamba_cache_manager=hybrid_cache_manager_gen)
request.addfinalizer(cache_transceiver_gen.shutdown)

fill_hybrid_cache_buffers(hybrid_cache_manager_ctx)

Expand All @@ -991,6 +1029,9 @@ def test_hybrid_cache_transceiver_cancel_request(backend):
sampling_params._get_sampling_config()),
is_streaming=False,
llm_request_type=LlmRequestType.LLMREQUEST_TYPE_CONTEXT_ONLY)
ctx_request.py_disaggregated_params = tensorrt_llm.DisaggregatedParams(
request_type="context_only",
disagg_request_id=uuid.uuid4().int & 0x7FFFFFFFFFFFFFFF)

scheduled_ctx = ScheduledRequests()
scheduled_ctx.context_requests_last_chunk = [ctx_request]
Expand All @@ -999,13 +1040,6 @@ def test_hybrid_cache_transceiver_cancel_request(backend):
# Send ctx request
cache_transceiver_ctx.respond_and_send_async(ctx_request)

# Wait for ctx request to be sent
time.sleep(2)

# Cancel ctx request
is_cancelled = cache_transceiver_ctx.cancel_request(ctx_request)
assert is_cancelled

# Init gen request
gen_request = LlmRequest(
request_id=0,
Expand All @@ -1016,6 +1050,14 @@ def test_hybrid_cache_transceiver_cancel_request(backend):
is_streaming=False,
llm_request_type=LlmRequestType.LLMREQUEST_TYPE_GENERATION_ONLY,
context_phase_params=ctx_request.context_phase_params)
gen_request.py_disaggregated_params = tensorrt_llm.DisaggregatedParams(
request_type="generation_only",
disagg_request_id=ctx_request.py_disaggregated_params.disagg_request_id,
ctx_request_id=ctx_request.request_id,
ctx_dp_rank=ctx_request.context_phase_params.ctx_dp_rank,
ctx_info_endpoint=ctx_request.context_phase_params.disagg_info_endpoint,
first_gen_tokens=ctx_request.context_phase_params.first_gen_tokens,
draft_tokens=ctx_request.context_phase_params.draft_tokens)

scheduled_gen = ScheduledRequests()
scheduled_gen.context_requests_last_chunk = [gen_request]
Expand All @@ -1024,7 +1066,21 @@ def test_hybrid_cache_transceiver_cancel_request(backend):
# Try to receive gen request
cache_transceiver_gen.request_and_receive_async(gen_request)

# Block the main thread due to the async operation
time.sleep(2)
cache_transceiver_gen.check_gen_transfer_status(0)
assert gen_request.state == LlmRequestState.DISAGG_TRANS_ERROR
generation_cancelled = [False]

def cancel_generation_transfer():
generation_cancelled[0] = cache_transceiver_gen.cancel_request(
gen_request)

wait_for_transfer_completion(cancel_generation_transfer,
lambda: generation_cancelled[0])
assert cache_transceiver_gen.check_gen_transfer_complete()

context_cancelled = [False]

def cancel_context_transfer():
context_cancelled[0] = cache_transceiver_ctx.cancel_request(ctx_request)

wait_for_transfer_completion(cancel_context_transfer,
lambda: context_cancelled[0])
shutdown_transceivers(cache_transceiver_gen, cache_transceiver_ctx)
Loading