Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
25 commits
Select commit Hold shift + click to select a range
facbc44
[None][feat] MLA-backboned standalone DSpark drafter (Inferact/Kimi-K…
dc3671 Sep 9, 2026
56d48dc
[None][doc] Correct the draft-mirror saturation comments
dc3671 Sep 11, 2026
98c701b
[None][fix] Charge the external drafter's KV budget at its allocated …
dc3671 Sep 11, 2026
b8e6ef9
[None][fix] Bound the standalone drafter's draft-KV writes per request
dc3671 Sep 11, 2026
fa47925
[None][chore] Address review on the MLA drafter port
dc3671 Sep 11, 2026
39432a2
[None][doc] State which V2 backend makes the draft page count per-req…
dc3671 Sep 11, 2026
c3533e9
[None][chore] Build the MLA drafter's YaRN table with RopeEmbeddingUtils
dc3671 Sep 11, 2026
9e9011f
[None][chore] Trim the MLA drafter's comments to what cannot be re-de…
dc3671 Sep 11, 2026
363414e
[None][feat] Make the standalone drafter's block-decode backend selec…
dc3671 Sep 16, 2026
443cb4a
[None][fix] Keep degenerate rows out of the rejection sampling kernel
dc3671 Sep 16, 2026
61e8c74
[None][fix] DSV4 DSpark CUDA graph RoPE bounds
reasonsolo Sep 16, 2026
9be64dc
[None][fix] Let each drafter family judge its own attention backend
dc3671 Sep 16, 2026
2a642b1
[None][fix] Never reset a live DSpark slot from prepare()
dc3671 Sep 16, 2026
12ffde9
[None][fix] Fix drafter test fixtures against the rebased main
dc3671 Sep 17, 2026
47df6a3
[None][chore] Regenerate the LLM-args golden manifest
dc3671 Sep 17, 2026
8ec2608
[None][fix] Raise on a declared backend that loads no op set
dc3671 Sep 17, 2026
b7ab453
[None][chore] Cut the long inline comments to the facts
dc3671 Sep 17, 2026
3793117
[None][chore] Cap docstrings and the backend description by point
dc3671 Sep 17, 2026
fb5ecf4
[None][fix] Derive KDA CuTe argument alignment
dc3671 Sep 17, 2026
717bfb2
[None][feat] Join DFlash/DSpark to the paired draft KV reuse protocol
dc3671 Sep 17, 2026
c48121b
[None][chore] Report whether the DFlash ctx cache can reach a reused …
dc3671 Sep 17, 2026
f17891f
[None][fix] Give the EPLB config fixture a max_seq_len
dc3671 Sep 17, 2026
fe16449
[None][fix] Keep Eagle/MTP in draft_prompt_lookahead
dc3671 Sep 17, 2026
b44802a
[None][chore] Name both dummy_slot_row publishers in the padding-row …
dc3671 Sep 18, 2026
8b47f68
[None][fix] Bound the DSv4 DSpark drafter's RoPE positions
dc3671 Sep 21, 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 @@ -11684,12 +11684,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 @@ -320,30 +321,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 @@ -478,9 +499,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 @@ -549,6 +567,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 @@ -562,26 +581,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 @@ -597,30 +616,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 @@ -642,30 +661,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