From 1e0dceb74e5e18d92294c6a62965c7043d87d3dd Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Thu, 1 Oct 2026 21:28:59 -0700 Subject: [PATCH] DFlash: start the dummy slot at context length 0 every step CUDA-graph padding rows and warmup dummies share the drafter's dummy slot. Their accepted tokens grow its context length like any request's, but nothing reset it, while their page-table rows are the padding request's pages with every other entry mapped to page 0, another request's page. Once the dummy length passed the padding request's own pages, the padding rows' context K / V (k3_ctx_kv, or the torch path) landed on that page: a live request's drafter context overwritten on every padded step, costing acceptance. prepare() now writes 0 to the dummy slot every step, with the evicted slots. test_dflash_dummy_slot.py: padded steps keep the dummy slot at 0 and the real slots untouched; an evicted request and the dummy reset in one step. It fails without the reset. Signed-off-by: Vasanth Sabavat --- tensorrt_llm/_torch/speculative/dflash.py | 4 + .../hw_agnostic/test_dflash_dummy_slot.py | 83 +++++++++++++++++++ 2 files changed, 87 insertions(+) create mode 100644 tests/unittest/_torch/speculative/hw_agnostic/test_dflash_dummy_slot.py diff --git a/tensorrt_llm/_torch/speculative/dflash.py b/tensorrt_llm/_torch/speculative/dflash.py index d8097dffbd58..d760e15fbc78 100644 --- a/tensorrt_llm/_torch/speculative/dflash.py +++ b/tensorrt_llm/_torch/speculative/dflash.py @@ -551,6 +551,10 @@ def prepare(self): evicted[slot] = 0 worker._req_ctx_pos.pop(rid, None) worker._free_slots.append(slot) + # The dummy slot starts every step empty. Its rows (CUDA-graph padding, warmup dummies) add their accepted + # tokens to its length like any request, and their table rows are the padding request's pages, the rest + # mapped to page 0 (another request's). Left growing, a padding row's context K / V would land there. + evicted[worker._dummy_slot] = 0 worker._write_ctx_len(evicted) # A disagg generation worker receives prompt KV instead of diff --git a/tests/unittest/_torch/speculative/hw_agnostic/test_dflash_dummy_slot.py b/tests/unittest/_torch/speculative/hw_agnostic/test_dflash_dummy_slot.py new file mode 100644 index 000000000000..37d837aa7c28 --- /dev/null +++ b/tests/unittest/_torch/speculative/hw_agnostic/test_dflash_dummy_slot.py @@ -0,0 +1,83 @@ +# 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. +"""DFlash's dummy slot (CUDA-graph padding and warmup dummies) starts every step at context length 0. + +Padding rows add their accepted tokens to the dummy slot's length like any request, while their page-table rows are +the padding request's pages with the rest mapped to page 0, which another request owns. A dummy length that kept +growing would put the padding rows' context K / V on that page. ``DFlashSpecMetadata.prepare`` runs before every step +(eager or replayed graph); after it the dummy slot's length is 0 and every real slot's is untouched. +""" + +from functools import partial +from types import SimpleNamespace + +import pytest +import torch + +from tensorrt_llm._torch.speculative.dflash import DFlashSpecMetadata, DFlashWorker + +pytestmark = pytest.mark.skipif( + not torch.cuda.is_available(), reason="the slot lengths live on the GPU" +) + +FLOOR = 1 << 40 # request ids from here on are CUDA-graph padding dummies + + +def _worker(lengths): + w = SimpleNamespace( + _ctx_buf_inited=True, + _req_to_slot={11: 0, 12: 1, 13: 2}, + _req_ctx_pos={}, + _free_slots=[3], + _graph_dummy_id_floor=FLOOR, + _dummy_slot=len(lengths) - 1, + _ctx_len=torch.tensor(lengths, dtype=torch.long, device="cuda"), + _ctx_len_host=list(lengths), + _batch_to_slot=torch.zeros(8, dtype=torch.long, device="cuda"), + ) + w._write_ctx_len = partial(DFlashWorker._write_ctx_len, w) + w._assign_slot = lambda rid, *a, **k: None + return w + + +def _prepare(worker, request_ids, num_generations): + meta = SimpleNamespace( + request_ids=request_ids, + num_generations=num_generations, + batch_indices_cuda=torch.zeros(8, dtype=torch.int, device="cuda"), + _dflash_worker=worker, + ) + DFlashSpecMetadata.prepare(meta) + torch.cuda.synchronize() + + +def test_padding_steps_keep_the_dummy_slot_empty(): + worker = _worker([40, 300, 7, 0, 0]) + padded = [11, 12, 13, FLOOR + 1, FLOOR + 1] # three requests in a graph of five + for step in range(4): + # What the step's acceptance does to the slots of its rows (k3_ctx_kv / the torch path): + accepted. + _prepare(worker, padded, num_generations=5) + assert worker._ctx_len.tolist() == [40 + 8 * step, 300 + 8 * step, 7 + 8 * step, 0, 0], step + assert worker._batch_to_slot[:5].tolist() == [0, 1, 2, 4, 4] + assert worker._ctx_len_host[4] == 0 + worker._ctx_len[[0, 1, 2]] += 8 + worker._ctx_len[4] += 2 * 8 # both padding rows land on the dummy slot + + +def test_evicted_request_and_dummy_reset_together(): + worker = _worker([40, 300, 7, 0, 96]) + _prepare(worker, [11, 13, FLOOR + 1], num_generations=3) # request 12 left; one padding row + assert worker._ctx_len.tolist() == [40, 0, 7, 0, 0] + assert 1 in worker._free_slots and 12 not in worker._req_to_slot