dflash piecewise cuda graphs support (#27468)
This commit is contained in:
@@ -217,8 +217,12 @@ class PiecewiseCudaGraphRunner:
|
|||||||
self.capture_forward_mode = ForwardMode.EXTEND
|
self.capture_forward_mode = ForwardMode.EXTEND
|
||||||
self.capture_hidden_mode = CaptureHiddenMode.NULL
|
self.capture_hidden_mode = CaptureHiddenMode.NULL
|
||||||
|
|
||||||
# If returning hidden states is enabled, set initial capture hidden mode to full to avoid double-capture on startup
|
# If returning hidden states is enabled, or if speculative prefill needs
|
||||||
if model_runner.server_args.enable_return_hidden_states:
|
# aux hidden states (DFLASH), capture the FULL variant up front.
|
||||||
|
if (
|
||||||
|
model_runner.server_args.enable_return_hidden_states
|
||||||
|
or model_runner.spec_algorithm.is_dflash()
|
||||||
|
):
|
||||||
self.capture_hidden_mode = CaptureHiddenMode.FULL
|
self.capture_hidden_mode = CaptureHiddenMode.FULL
|
||||||
|
|
||||||
self.max_num_tokens = (
|
self.max_num_tokens = (
|
||||||
@@ -416,7 +420,7 @@ class PiecewiseCudaGraphRunner:
|
|||||||
mrope_positions=mrope_positions,
|
mrope_positions=mrope_positions,
|
||||||
spec_algorithm=None,
|
spec_algorithm=None,
|
||||||
spec_info=None,
|
spec_info=None,
|
||||||
capture_hidden_mode=CaptureHiddenMode.NULL,
|
capture_hidden_mode=self.capture_hidden_mode,
|
||||||
num_token_non_padded=None,
|
num_token_non_padded=None,
|
||||||
num_token_non_padded_cpu=num_tokens,
|
num_token_non_padded_cpu=num_tokens,
|
||||||
global_forward_mode=ForwardMode.EXTEND,
|
global_forward_mode=ForwardMode.EXTEND,
|
||||||
@@ -587,7 +591,7 @@ class PiecewiseCudaGraphRunner:
|
|||||||
mrope_positions=mrope_positions,
|
mrope_positions=mrope_positions,
|
||||||
spec_algorithm=None,
|
spec_algorithm=None,
|
||||||
spec_info=None,
|
spec_info=None,
|
||||||
capture_hidden_mode=CaptureHiddenMode.NULL,
|
capture_hidden_mode=self.capture_hidden_mode,
|
||||||
num_token_non_padded=None,
|
num_token_non_padded=None,
|
||||||
num_token_non_padded_cpu=num_tokens,
|
num_token_non_padded_cpu=num_tokens,
|
||||||
global_forward_mode=ForwardMode.EXTEND,
|
global_forward_mode=ForwardMode.EXTEND,
|
||||||
|
|||||||
@@ -0,0 +1,47 @@
|
|||||||
|
"""Test piecewise CUDA graph coexisting with speculative decoding (DFLASH).
|
||||||
|
|
||||||
|
PCG handles prefill/extend path while DFlash needs target aux hidden states
|
||||||
|
from prefill to materialize draft KV cache. This verifies PCG captures that
|
||||||
|
path with the DFlash hidden-state variant enabled.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import unittest
|
||||||
|
|
||||||
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
|
from sglang.test.server_fixtures.pcg_spec_fixture import PCGSpecBase
|
||||||
|
from sglang.test.test_utils import (
|
||||||
|
DEFAULT_DRAFT_MODEL_DFLASH,
|
||||||
|
DEFAULT_TARGET_MODEL_DFLASH,
|
||||||
|
CustomTestCase,
|
||||||
|
)
|
||||||
|
|
||||||
|
register_cuda_ci(est_time=531, stage="base-b", runner_config="1-gpu-small")
|
||||||
|
|
||||||
|
|
||||||
|
class TestPCGWithDFlash(PCGSpecBase, CustomTestCase):
|
||||||
|
"""PCG + DFLASH on Llama-3.1-8B-Instruct."""
|
||||||
|
|
||||||
|
model = DEFAULT_TARGET_MODEL_DFLASH
|
||||||
|
server_args = [
|
||||||
|
"--trust-remote-code",
|
||||||
|
"--attention-backend",
|
||||||
|
"flashinfer",
|
||||||
|
"--enforce-piecewise-cuda-graph",
|
||||||
|
"--speculative-algorithm",
|
||||||
|
"DFLASH",
|
||||||
|
"--speculative-draft-model-path",
|
||||||
|
DEFAULT_DRAFT_MODEL_DFLASH,
|
||||||
|
"--page-size",
|
||||||
|
"1",
|
||||||
|
"--max-running-requests",
|
||||||
|
"64",
|
||||||
|
"--cuda-graph-bs",
|
||||||
|
*[str(i) for i in range(1, 65)],
|
||||||
|
]
|
||||||
|
server_env = {"SGLANG_ALLOW_OVERWRITE_LONGER_CONTEXT_LEN": "1"}
|
||||||
|
accuracy_threshold = 0.75
|
||||||
|
speedup_threshold = 2.8
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user