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
29 changes: 28 additions & 1 deletion tensorrt_llm/_torch/models/modeling_deepseekv3.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,7 +28,7 @@
import copy
import math
import os
from typing import Dict, List, Optional, Tuple
from typing import Any, Dict, List, Literal, Optional, Tuple

import torch
import triton
Expand Down Expand Up @@ -1893,6 +1893,33 @@ def forward(
class DeepseekV3ForCausalLM(SpecDecOneEngineForCausalLM[DeepseekV3Model,
PretrainedConfig]):

@classmethod
def get_preferred_transceiver_runtime(cls,
pretrained_config: Any = None
) -> Optional[Literal["PYTHON"]]:
"""GLM-5 family checkpoints default to the Python (v2) KV-cache
transceiver.

This implementation class is shared by DeepSeek-V3/V3.2 and the GLM-5
family — both GLM-5 and GLM-5.2 declare ``GlmMoeDsaForCausalLM`` /
``glm_moe_dsa`` — so the preference is differentiated per checkpoint:
only GLM checkpoints opt into the Python transceiver.
The MLA backbone transfers a large latent KV, which the Python
transceiver handles better in disaggregated serving. This is only
adopted when the user leaves
``cache_transceiver_config.transceiver_runtime`` at 'auto' and the
effective backend is NIXL; otherwise the C++ transceiver is used.
"""
if pretrained_config is None:
return None
architectures = getattr(pretrained_config, 'architectures', None) or []
# model_type is checked as a fallback: it is 'glm_moe_dsa' on GLM
# checkpoints until __init__ rewrites it to 'deepseek_v32'.
if ("GlmMoeDsaForCausalLM" in architectures or getattr(
pretrained_config, 'model_type', None) == 'glm_moe_dsa'):
return "PYTHON"
return None

def __init__(self, model_config: ModelConfig[PretrainedConfig]):
self.mapping_with_cp = None
# Note: Currently the usage of mapping is all over the place making its usage brittle
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,8 @@ context_servers:
cuda_graph_config: null
print_iter_log: true
cache_transceiver_config:
backend: DEFAULT
backend: NIXL
transceiver_runtime: PYTHON
max_tokens_in_buffer: 16384
generation_servers:
num_instances: 1
Expand Down Expand Up @@ -55,5 +56,6 @@ generation_servers:
- 1024
print_iter_log: true
cache_transceiver_config:
backend: DEFAULT
backend: NIXL
transceiver_runtime: PYTHON
max_tokens_in_buffer: 16384
Original file line number Diff line number Diff line change
Expand Up @@ -61,6 +61,7 @@ worker_config:
cache_transceiver_config:
max_tokens_in_buffer: 16384
backend: NIXL
transceiver_runtime: PYTHON
disable_overlap_scheduler: true
speculative_config: &id001
decoding_type: MTP
Expand Down Expand Up @@ -89,5 +90,6 @@ worker_config:
cache_transceiver_config:
max_tokens_in_buffer: 16384
backend: NIXL
transceiver_runtime: PYTHON
disable_overlap_scheduler: true
speculative_config: *id001
Original file line number Diff line number Diff line change
Expand Up @@ -66,6 +66,7 @@ worker_config:
cache_transceiver_config:
max_tokens_in_buffer: 16384
backend: NIXL
transceiver_runtime: PYTHON
disable_overlap_scheduler: true
speculative_config: &id001
decoding_type: MTP
Expand Down Expand Up @@ -93,5 +94,6 @@ worker_config:
cache_transceiver_config:
max_tokens_in_buffer: 16384
backend: NIXL
transceiver_runtime: PYTHON
disable_overlap_scheduler: true
speculative_config: *id001
Original file line number Diff line number Diff line change
Expand Up @@ -63,6 +63,7 @@ worker_config:
cache_transceiver_config:
max_tokens_in_buffer: 16384
backend: NIXL
transceiver_runtime: PYTHON
disable_overlap_scheduler: true
speculative_config: &id001
decoding_type: MTP
Expand Down Expand Up @@ -90,5 +91,6 @@ worker_config:
cache_transceiver_config:
max_tokens_in_buffer: 16384
backend: NIXL
transceiver_runtime: PYTHON
disable_overlap_scheduler: true
speculative_config: *id001
Original file line number Diff line number Diff line change
Expand Up @@ -66,6 +66,7 @@ worker_config:
cache_transceiver_config:
max_tokens_in_buffer: 16384
backend: NIXL
transceiver_runtime: PYTHON
disable_overlap_scheduler: true
speculative_config: &id001
decoding_type: MTP
Expand Down Expand Up @@ -93,5 +94,6 @@ worker_config:
cache_transceiver_config:
max_tokens_in_buffer: 16384
backend: NIXL
transceiver_runtime: PYTHON
disable_overlap_scheduler: true
speculative_config: *id001
Original file line number Diff line number Diff line change
Expand Up @@ -61,6 +61,7 @@ worker_config:
cache_transceiver_config:
max_tokens_in_buffer: 16384
backend: NIXL
transceiver_runtime: PYTHON
disable_overlap_scheduler: true
speculative_config: &id001
decoding_type: MTP
Expand Down Expand Up @@ -89,5 +90,6 @@ worker_config:
cache_transceiver_config:
max_tokens_in_buffer: 16384
backend: NIXL
transceiver_runtime: PYTHON
disable_overlap_scheduler: true
speculative_config: *id001
Original file line number Diff line number Diff line change
Expand Up @@ -63,6 +63,7 @@ worker_config:
cache_transceiver_config:
max_tokens_in_buffer: 16384
backend: NIXL
transceiver_runtime: PYTHON
disable_overlap_scheduler: true
speculative_config: &id001
decoding_type: MTP
Expand Down Expand Up @@ -90,5 +91,6 @@ worker_config:
cache_transceiver_config:
max_tokens_in_buffer: 16384
backend: NIXL
transceiver_runtime: PYTHON
disable_overlap_scheduler: true
speculative_config: *id001
Original file line number Diff line number Diff line change
Expand Up @@ -61,6 +61,7 @@ worker_config:
cache_transceiver_config:
max_tokens_in_buffer: 16384
backend: NIXL
transceiver_runtime: PYTHON
disable_overlap_scheduler: true
speculative_config: &id001
decoding_type: MTP
Expand Down Expand Up @@ -89,5 +90,6 @@ worker_config:
cache_transceiver_config:
max_tokens_in_buffer: 16384
backend: NIXL
transceiver_runtime: PYTHON
disable_overlap_scheduler: true
speculative_config: *id001
Original file line number Diff line number Diff line change
Expand Up @@ -66,6 +66,7 @@ worker_config:
cache_transceiver_config:
max_tokens_in_buffer: 16384
backend: NIXL
transceiver_runtime: PYTHON
disable_overlap_scheduler: true
speculative_config: &id001
decoding_type: MTP
Expand Down Expand Up @@ -93,5 +94,6 @@ worker_config:
cache_transceiver_config:
max_tokens_in_buffer: 16384
backend: NIXL
transceiver_runtime: PYTHON
disable_overlap_scheduler: true
speculative_config: *id001
Original file line number Diff line number Diff line change
Expand Up @@ -63,6 +63,7 @@ worker_config:
cache_transceiver_config:
max_tokens_in_buffer: 16384
backend: NIXL
transceiver_runtime: PYTHON
disable_overlap_scheduler: true
speculative_config: &id001
decoding_type: MTP
Expand Down Expand Up @@ -90,5 +91,6 @@ worker_config:
cache_transceiver_config:
max_tokens_in_buffer: 16384
backend: NIXL
transceiver_runtime: PYTHON
disable_overlap_scheduler: true
speculative_config: *id001
Original file line number Diff line number Diff line change
Expand Up @@ -66,6 +66,7 @@ worker_config:
cache_transceiver_config:
max_tokens_in_buffer: 16384
backend: NIXL
transceiver_runtime: PYTHON
disable_overlap_scheduler: true
speculative_config: &id001
decoding_type: MTP
Expand Down Expand Up @@ -93,5 +94,6 @@ worker_config:
cache_transceiver_config:
max_tokens_in_buffer: 16384
backend: NIXL
transceiver_runtime: PYTHON
disable_overlap_scheduler: true
speculative_config: *id001
Original file line number Diff line number Diff line change
Expand Up @@ -61,6 +61,7 @@ worker_config:
cache_transceiver_config:
max_tokens_in_buffer: 16384
backend: NIXL
transceiver_runtime: PYTHON
disable_overlap_scheduler: true
speculative_config: &id001
decoding_type: MTP
Expand Down Expand Up @@ -89,5 +90,6 @@ worker_config:
cache_transceiver_config:
max_tokens_in_buffer: 16384
backend: NIXL
transceiver_runtime: PYTHON
disable_overlap_scheduler: true
speculative_config: *id001
Original file line number Diff line number Diff line change
Expand Up @@ -63,6 +63,7 @@ worker_config:
cache_transceiver_config:
max_tokens_in_buffer: 16384
backend: NIXL
transceiver_runtime: PYTHON
disable_overlap_scheduler: true
speculative_config: &id001
decoding_type: MTP
Expand Down Expand Up @@ -90,5 +91,6 @@ worker_config:
cache_transceiver_config:
max_tokens_in_buffer: 16384
backend: NIXL
transceiver_runtime: PYTHON
disable_overlap_scheduler: true
speculative_config: *id001
58 changes: 58 additions & 0 deletions tests/unittest/llmapi/test_llm_args.py
Original file line number Diff line number Diff line change
Expand Up @@ -3338,3 +3338,61 @@ def test_resolve_default_backend_env_priority(self, monkeypatch):
# An explicit backend bypasses the env vars entirely.
assert CacheTransceiverConfig(
backend="UCX")._resolve_default_backend() == ("UCX", None)


class TestGlm5TransceiverPreference:
"""GLM-5 defaults to the Python KV-cache transceiver in disagg.

DeepseekV3ForCausalLM is shared by DeepSeek-V3/V3.2 and GLM-5
(GlmMoeDsaForCausalLM); the preference must apply to GLM checkpoints
only.
"""

@staticmethod
def _pretrained_config(architectures, model_type):
from unittest.mock import MagicMock
cfg = MagicMock()
cfg.architectures = architectures
cfg.model_type = model_type
return cfg

@pytest.mark.parametrize(
"architectures,model_type,expected",
[
(["GlmMoeDsaForCausalLM"], "glm_moe_dsa", "PYTHON"),
(["DeepseekV3ForCausalLM"], "deepseek_v3", None),
(["DeepseekV32ForCausalLM"], "deepseek_v32", None),
# Each predicate in isolation: the architecture match and the
# model_type fallback must each suffice on their own.
(["GlmMoeDsaForCausalLM"], "deepseek_v32", "PYTHON"),
(["DeepseekV32ForCausalLM"], "glm_moe_dsa", "PYTHON"),
])
def test_preference_per_architecture(self, architectures, model_type,
expected):
from tensorrt_llm._torch.models.modeling_deepseekv3 import \
DeepseekV3ForCausalLM
cfg = self._pretrained_config(architectures, model_type)
assert DeepseekV3ForCausalLM.get_preferred_transceiver_runtime(
cfg) == expected

def test_no_config_defers_to_cpp(self):
"""Without a pretrained config the class defers to the C++ default."""
from tensorrt_llm._torch.models.modeling_deepseekv3 import \
DeepseekV3ForCausalLM
assert DeepseekV3ForCausalLM.get_preferred_transceiver_runtime() is None

def test_glm5_resolves_auto_to_python_on_nixl(self):
"""GLM-5 on NIXL adopts the Python transceiver from 'auto'.

End-to-end through _resolve_transceiver_runtime_auto with the real
model class and a GLM pretrained config.
"""
from tensorrt_llm._torch.models.modeling_deepseekv3 import \
DeepseekV3ForCausalLM
args = TorchLlmArgs(
model="/tmp/dummy_model",
cache_transceiver_config=CacheTransceiverConfig(backend="NIXL"),
)
cfg = self._pretrained_config(["GlmMoeDsaForCausalLM"], "glm_moe_dsa")
_resolve_transceiver_runtime_auto(args, DeepseekV3ForCausalLM, cfg)
assert args.cache_transceiver_config.transceiver_runtime == "PYTHON"
Loading