Split TRTLLM MHA decode batches by KV sequence length (#34888)

This commit is contained in:
YAMY
2026-08-20 00:44:26 -07:00
committed by GitHub
parent b8996a5ab2
commit ae23423b46
3 changed files with 96 additions and 23 deletions
@@ -2,6 +2,7 @@ import unittest
import torch
from sglang.srt.environ import envs
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.srt.utils import is_flashinfer_available
from sglang.srt.utils.common import (
@@ -176,8 +177,11 @@ class TestTRTLLMMHADenseAttentionBackendCorrectness(CustomTestCase):
)
def test_projected_dense_decode_cases(self):
for case in self.DECODE_CASES:
with self.subTest(case=case.name, backend=case.backend):
for case_index, case in enumerate(self.DECODE_CASES):
splits = 2 if case_index == 0 else 1
with self.subTest(
case=case.name, backend=case.backend
), envs.SGLANG_TRTLLM_MHA_DECODE_SEQ_LEN_SPLITS.override(splits):
run_dense_attention_case(
self,
case,
@@ -226,7 +230,9 @@ class TestTRTLLMMHADenseAttentionBackendCorrectness(CustomTestCase):
def test_runner_mode_frozen_kv_mtp_cuda_graph_runner_cases(self):
for case in self.FROZEN_KV_MTP_RUNNER_CASES:
with self.subTest(case=case.name, backend=case.backend):
with self.subTest(
case=case.name, backend=case.backend
), envs.SGLANG_TRTLLM_MHA_DECODE_SEQ_LEN_SPLITS.override(2):
run_dense_frozen_kv_mtp_cuda_graph_runner_case(
self,
case,