[scheduler] fix: correcting extend_logprob_start_len calculation (#15922)

This commit is contained in:
Cheng Wan
2025-12-28 14:57:04 -08:00
committed by GitHub
parent d7a3336ebe
commit 6f9d0a89a0
8 changed files with 84 additions and 53 deletions

View File

@@ -93,7 +93,7 @@ class TestForwardSplitPrefill(CustomTestCase):
)
req.fill_ids = req.origin_input_ids
req.extend_input_len = len(req.fill_ids) - len(req.prefix_indices)
req.logprob_start_len = len(req.origin_input_ids) - 1
req.logprob_start_len = -1
reqs.append(req)
# Create dummy tree_cache for tests (no prefix caching, just allocation)

View File

@@ -3,8 +3,10 @@ from types import SimpleNamespace
import requests
from sglang.srt.environ import envs
from sglang.srt.utils import kill_process_tree
from sglang.test.few_shot_gsm8k import run_eval as run_eval_few_shot_gsm8k
from sglang.test.kits.radix_cache_server_kit import run_radix_attention_test
from sglang.test.run_eval import run_eval
from sglang.test.test_utils import (
DEFAULT_MLA_MODEL_NAME_FOR_TEST,
@@ -58,6 +60,41 @@ class TestDPAttentionDP2TP2(CustomTestCase):
self.assertGreater(metrics["score"], 0.8)
class TestDPRetract(CustomTestCase):
@classmethod
def setUpClass(cls):
cls.model = DEFAULT_MLA_MODEL_NAME_FOR_TEST
cls.base_url = DEFAULT_URL_FOR_TEST
cls.process = popen_launch_server(
cls.model,
cls.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=[
"--trust-remote-code",
"--tp",
"2",
"--enable-dp-attention",
"--dp",
"2",
"--max-total-tokens",
"4500",
"--max-running-requests",
"128",
"--chunked-prefill-size",
"256",
],
)
@classmethod
def tearDownClass(cls):
kill_process_tree(cls.process.pid)
def test_radix_attention(self):
with envs.SGLANG_TEST_RETRACT.override(True):
run_radix_attention_test(self.base_url)
self.assertIsNone(self.process.poll())
class TestDPAttentionDP2TP2DeepseekV3MTP(CustomTestCase):
@classmethod
def setUpClass(cls):