Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
23 commits
Select commit Hold shift + click to select a range
0273e82
[None][feat] MLA-backboned standalone DSpark drafter (Inferact/Kimi-K…
dc3671 Sep 9, 2026
9918cf7
[None][doc] Correct the draft-mirror saturation comments
dc3671 Sep 11, 2026
81dff1f
[None][fix] Charge the external drafter's KV budget at its allocated …
dc3671 Sep 11, 2026
a248176
[None][fix] Bound the standalone drafter's draft-KV writes per request
dc3671 Sep 11, 2026
c1c19a0
[None][chore] Address review on the MLA drafter port
dc3671 Sep 11, 2026
cdca348
[None][doc] State which V2 backend makes the draft page count per-req…
dc3671 Sep 11, 2026
01c4e32
[None][chore] Build the MLA drafter's YaRN table with RopeEmbeddingUtils
dc3671 Sep 11, 2026
dd00d57
[None][chore] Trim the MLA drafter's comments to what cannot be re-de…
dc3671 Sep 11, 2026
50039e2
[None][feat] Make the standalone drafter's block-decode backend selec…
dc3671 Sep 16, 2026
029e0a6
[None][fix] Keep degenerate rows out of the rejection sampling kernel
dc3671 Sep 16, 2026
9a76ec0
[None][fix] DSV4 DSpark CUDA graph RoPE bounds
reasonsolo Sep 16, 2026
d4ac584
[None][fix] Let each drafter family judge its own attention backend
dc3671 Sep 16, 2026
77a3230
[None][fix] Never reset a live DSpark slot from prepare()
dc3671 Sep 16, 2026
4c0d7ae
[None][fix] Fix drafter test fixtures against the rebased main
dc3671 Sep 17, 2026
373624f
[None][chore] Regenerate the LLM-args golden manifest
dc3671 Sep 17, 2026
dfd0308
[None][fix] Raise on a declared backend that loads no op set
dc3671 Sep 17, 2026
d576296
[None][chore] Cut the long inline comments to the facts
dc3671 Sep 17, 2026
bb87859
[None][chore] Cap docstrings and the backend description by point
dc3671 Sep 17, 2026
e54aea2
[None][fix] Derive KDA CuTe argument alignment
dc3671 Sep 17, 2026
494b3cf
[None][feat] Join DFlash/DSpark to the paired draft KV reuse protocol
dc3671 Sep 17, 2026
b614879
[None][chore] Report whether the DFlash ctx cache can reach a reused …
dc3671 Sep 17, 2026
1ead943
[None][fix] Give the EPLB config fixture a max_seq_len
dc3671 Sep 17, 2026
dbfadd0
[None][fix] Keep Eagle/MTP in draft_prompt_lookahead
dc3671 Sep 17, 2026
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
5 changes: 5 additions & 0 deletions tensorrt_llm/_torch/configs/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,7 @@
Gemma4UnifiedTextConfig,
Gemma4UnifiedVisionConfig,
)
from tensorrt_llm._torch.configs.k3_dspark import K3DsparkConfig
from tensorrt_llm._torch.configs.kimi_k3 import KimiK3Config, KimiK3VisionConfig
from tensorrt_llm._torch.configs.kimi_linear import KimiLinearConfig
from tensorrt_llm._torch.configs.laguna import LagunaConfig
Expand Down Expand Up @@ -67,6 +68,9 @@ def _register_custom_configs_with_transformers() -> None:
# sub-configs and multimodal is not disabled, and otherwise flattens to
# the text config. Registering both here lets AutoConfig / AutoTokenizer
# resolve them without trust_remote_code.
# The MLA DSpark drafter checkpoint ships no auto_map, so AutoConfig
# cannot resolve its model_type on its own.
"k3_dspark": K3DsparkConfig,
"kimi_k3": KimiK3Config,
"kimi_linear": KimiLinearConfig,
"laguna": LagunaConfig,
Expand Down Expand Up @@ -104,6 +108,7 @@ def _register_custom_configs_with_transformers() -> None:
"Gemma4UnifiedConfig",
"Gemma4UnifiedTextConfig",
"Gemma4UnifiedVisionConfig",
"K3DsparkConfig",
"KimiK3Config",
"KimiK3VisionConfig",
"KimiLinearConfig",
Expand Down
24 changes: 24 additions & 0 deletions tensorrt_llm/_torch/configs/k3_dspark.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,24 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.

from transformers.configuration_utils import PretrainedConfig


# The MLA-backboned DSpark drafter (Inferact/Kimi-K3-DSpark) ships a config.json
# with model_type "k3_dspark", no auto_map and no modeling code, so
# AutoConfig.from_pretrained cannot resolve it. Same workaround as LagunaConfig:
# the fields are plain attributes, and MLADSparkForCausalLM reads them directly.
class K3DsparkConfig(PretrainedConfig):
model_type = "k3_dspark"
10 changes: 9 additions & 1 deletion tensorrt_llm/_torch/custom_ops/cute_dsl_custom_ops.py
Original file line number Diff line number Diff line change
Expand Up @@ -11559,12 +11559,20 @@ def forward(

compiled_mla = CuteDSLNVMlaDecodeBlackwellRunner.kernel_cache[
cache_key]
page_table_arg = page_table
if page_table.shape[0] == 1 and page_table.shape[1] == 1:
# leading_dim=0 does not survive TensorAdapter's call-time re-adapt
# (cute/runtime.py:915), and a (1, 1) table has no extent > 1, so
# deduction raises "Can't deduce the leading dimension from layout".
page_table_arg = cute.runtime.from_dlpack(
page_table,
assumed_align=16).mark_layout_dynamic(leading_dim=0)
runtime_args = [
q_latent,
q_rope,
c_latent,
c_rope,
page_table,
page_table_arg,
o,
lse,
]
Expand Down
157 changes: 88 additions & 69 deletions tensorrt_llm/_torch/custom_ops/cute_dsl_kimi_k3_kda_mtp_ops.py
Original file line number Diff line number Diff line change
Expand Up @@ -53,6 +53,7 @@
(2, 12, 32), num_spec == 2``; other shapes compile the general variant.
"""

from math import gcd
from typing import Optional, Tuple

import torch
Expand Down Expand Up @@ -302,30 +303,50 @@ def _fits_32bit_stride(tensor: torch.Tensor) -> bool:
return True


def _from_dlpack_arg(tensor: torch.Tensor):
def _from_dlpack_arg(tensor: torch.Tensor, *, assumed_align: int = 16):
return from_dlpack(
tensor,
assumed_align=16,
assumed_align=assumed_align,
use_32bit_stride=_fits_32bit_stride(tensor),
)


def _dlpack_arg(tensor: torch.Tensor):
def _beta_cache_assumed_align(beta_cache: torch.Tensor) -> int:
"""Return the alignment shared by KDA per-layer beta-cache views.

The producer allocates ``[layers, slots, num_spec, local_heads]``. After
selecting a layer, ``slots * stride(0)`` is the physical layer span even
when the head rows are padded. Combine that span with the current pointer
and dtype, capped at the CuTe bridge's useful 16-byte guarantee.
"""
layer_span_bytes = beta_cache.shape[0] * beta_cache.stride(0) * beta_cache.element_size()
return gcd(16, beta_cache.data_ptr(), layer_span_bytes)


def _dlpack_arg(tensor: torch.Tensor, *, assumed_align: int):
# Alignment is deliberately mandatory: layout dynamism does not imply
# arbitrary pointer alignment. Each call site must state the guarantee
# provided by that tensor's producer and view pattern.
arg = _from_dlpack_arg(tensor, assumed_align=assumed_align)
for dim, stride in enumerate(tensor.stride()):
if stride == 1:
return _from_dlpack_arg(tensor).mark_layout_dynamic(dim)
return _from_dlpack_arg(tensor).mark_layout_dynamic()
return arg.mark_layout_dynamic(dim)
return arg.mark_layout_dynamic()


def _layout_key(tensor: torch.Tensor, dynamic_layout: bool = False):
arg = _dlpack_arg(tensor) if dynamic_layout else _from_dlpack_arg(tensor)
def _layout_key(tensor: torch.Tensor, dynamic_layout: bool = False, *, assumed_align: int = 16):
arg = (
_dlpack_arg(tensor, assumed_align=assumed_align)
if dynamic_layout
else _from_dlpack_arg(tensor, assumed_align=assumed_align)
)
shape_mask = arg.dynamic_shapes_mask
stride_mask = arg.dynamic_strides_mask
shape = tuple(None if dynamic else size for size, dynamic in zip(tensor.shape, shape_mask))
stride = tuple(
None if dynamic else value for value, dynamic in zip(tensor.stride(), stride_mask)
)
return (tensor.dtype, shape, stride, _fits_32bit_stride(tensor))
return (tensor.dtype, shape, stride, _fits_32bit_stride(tensor), assumed_align)


# (device_index, enabled) -> persistent int32 [1] control tensor. Keys are
Expand Down Expand Up @@ -460,9 +481,6 @@ def kda_mtp_decode_impl(
out = torch.zeros(1, T_total, HV, V_dim, dtype=x_q.dtype, device=x_q.device)
if num_accepted_tokens.dtype != torch.int32:
num_accepted_tokens = num_accepted_tokens.to(torch.int32)
if ssm_state_indices.data_ptr() % 16 != 0:
raise ValueError("ssm_state_indices must be 16-byte aligned before CuTe DLPack conversion")

_require_stride_layout(
x_q=x_q,
x_k=x_k,
Expand Down Expand Up @@ -531,6 +549,7 @@ def kda_mtp_decode_impl(
"buffer before enabling PROFILE_STAGES"
)
stage_timing_arg = out
beta_cache_assumed_align = _beta_cache_assumed_align(beta_cache)

key = (
x_q.dtype,
Expand All @@ -544,26 +563,26 @@ def kda_mtp_decode_impl(
lower_bound,
use_flat_layout,
_layout_key(h0_arg),
_layout_key(x_q_arg, dynamic_layout=True),
_layout_key(x_k_arg, dynamic_layout=True),
_layout_key(x_v_arg, dynamic_layout=True),
_layout_key(x_q_arg, dynamic_layout=True, assumed_align=16),
_layout_key(x_k_arg, dynamic_layout=True, assumed_align=16),
_layout_key(x_v_arg, dynamic_layout=True, assumed_align=16),
_layout_key(w_q),
_layout_key(w_k),
_layout_key(w_v),
_layout_key(cs_q),
_layout_key(cs_k),
_layout_key(cs_v),
_layout_key(A_log),
_layout_key(g, dynamic_layout=True),
_layout_key(g, dynamic_layout=True, assumed_align=16),
_layout_key(dt_bias),
_layout_key(beta, dynamic_layout=True),
_layout_key(out, dynamic_layout=True),
_layout_key(beta, dynamic_layout=True, assumed_align=16),
_layout_key(out, dynamic_layout=True, assumed_align=16),
_layout_key(qkg_cache),
_layout_key(v_cache),
_layout_key(beta_cache),
_layout_key(ssm_state_indices, dynamic_layout=True),
_layout_key(cu_seqlens, dynamic_layout=True),
_layout_key(num_accepted_tokens, dynamic_layout=True),
_layout_key(beta_cache, assumed_align=beta_cache_assumed_align),
_layout_key(ssm_state_indices, dynamic_layout=True, assumed_align=4),
_layout_key(cu_seqlens, dynamic_layout=True, assumed_align=4),
_layout_key(num_accepted_tokens, dynamic_layout=True, assumed_align=4),
use_setmaxreg,
use_regular_metadata,
use_reg_q_weights,
Expand All @@ -579,30 +598,30 @@ def kda_mtp_decode_impl(
)
_compiled_cache[key] = cute.compile(
_run_kda_decode_mtp,
_from_dlpack_arg(h0_arg),
_dlpack_arg(x_q_arg),
_dlpack_arg(x_k_arg),
_dlpack_arg(x_v_arg),
_from_dlpack_arg(w_q),
_from_dlpack_arg(w_k),
_from_dlpack_arg(w_v),
_from_dlpack_arg(cs_q),
_from_dlpack_arg(cs_k),
_from_dlpack_arg(cs_v),
_from_dlpack_arg(A_log),
_dlpack_arg(g),
_from_dlpack_arg(dt_bias),
_dlpack_arg(beta),
_dlpack_arg(out),
_from_dlpack_arg(h0_arg),
_from_dlpack_arg(qkg_cache),
_from_dlpack_arg(v_cache),
_from_dlpack_arg(beta_cache),
_dlpack_arg(stage_timing_arg),
_dlpack_arg(ssm_state_indices),
_dlpack_arg(cu_seqlens),
_dlpack_arg(num_accepted_tokens),
_from_dlpack_arg(precompute_control),
_from_dlpack_arg(h0_arg, assumed_align=16),
_dlpack_arg(x_q_arg, assumed_align=16),
_dlpack_arg(x_k_arg, assumed_align=16),
_dlpack_arg(x_v_arg, assumed_align=16),
_from_dlpack_arg(w_q, assumed_align=16),
_from_dlpack_arg(w_k, assumed_align=16),
_from_dlpack_arg(w_v, assumed_align=16),
_from_dlpack_arg(cs_q, assumed_align=16),
_from_dlpack_arg(cs_k, assumed_align=16),
_from_dlpack_arg(cs_v, assumed_align=16),
_from_dlpack_arg(A_log, assumed_align=16),
_dlpack_arg(g, assumed_align=16),
_from_dlpack_arg(dt_bias, assumed_align=16),
_dlpack_arg(beta, assumed_align=16),
_dlpack_arg(out, assumed_align=16),
_from_dlpack_arg(h0_arg, assumed_align=16),
_from_dlpack_arg(qkg_cache, assumed_align=16),
_from_dlpack_arg(v_cache, assumed_align=16),
_from_dlpack_arg(beta_cache, assumed_align=beta_cache_assumed_align),
_dlpack_arg(stage_timing_arg, assumed_align=16),
_dlpack_arg(ssm_state_indices, assumed_align=4),
_dlpack_arg(cu_seqlens, assumed_align=4),
_dlpack_arg(num_accepted_tokens, assumed_align=4),
_from_dlpack_arg(precompute_control, assumed_align=16),
scale=scale,
HV=HV,
K=K,
Expand All @@ -624,30 +643,30 @@ def kda_mtp_decode_impl(
)

_compiled_cache[key](
_dlpack_arg(h0_arg),
_dlpack_arg(x_q_arg),
_dlpack_arg(x_k_arg),
_dlpack_arg(x_v_arg),
_dlpack_arg(w_q),
_dlpack_arg(w_k),
_dlpack_arg(w_v),
_dlpack_arg(cs_q),
_dlpack_arg(cs_k),
_dlpack_arg(cs_v),
_dlpack_arg(A_log),
_dlpack_arg(g),
_dlpack_arg(dt_bias),
_dlpack_arg(beta),
_dlpack_arg(out),
_dlpack_arg(h0_arg),
_dlpack_arg(qkg_cache),
_dlpack_arg(v_cache),
_dlpack_arg(beta_cache),
_dlpack_arg(stage_timing_arg),
_dlpack_arg(ssm_state_indices),
_dlpack_arg(cu_seqlens),
_dlpack_arg(num_accepted_tokens),
_dlpack_arg(precompute_control),
_dlpack_arg(h0_arg, assumed_align=16),
_dlpack_arg(x_q_arg, assumed_align=16),
_dlpack_arg(x_k_arg, assumed_align=16),
_dlpack_arg(x_v_arg, assumed_align=16),
_dlpack_arg(w_q, assumed_align=16),
_dlpack_arg(w_k, assumed_align=16),
_dlpack_arg(w_v, assumed_align=16),
_dlpack_arg(cs_q, assumed_align=16),
_dlpack_arg(cs_k, assumed_align=16),
_dlpack_arg(cs_v, assumed_align=16),
_dlpack_arg(A_log, assumed_align=16),
_dlpack_arg(g, assumed_align=16),
_dlpack_arg(dt_bias, assumed_align=16),
_dlpack_arg(beta, assumed_align=16),
_dlpack_arg(out, assumed_align=16),
_dlpack_arg(h0_arg, assumed_align=16),
_dlpack_arg(qkg_cache, assumed_align=16),
_dlpack_arg(v_cache, assumed_align=16),
_dlpack_arg(beta_cache, assumed_align=beta_cache_assumed_align),
_dlpack_arg(stage_timing_arg, assumed_align=16),
_dlpack_arg(ssm_state_indices, assumed_align=4),
_dlpack_arg(cu_seqlens, assumed_align=4),
_dlpack_arg(num_accepted_tokens, assumed_align=4),
_dlpack_arg(precompute_control, assumed_align=16),
N,
stream,
)
Expand Down
Loading
Loading