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_hidden_mode = CaptureHiddenMode.NULL
|
||||
|
||||
# If returning hidden states is enabled, set initial capture hidden mode to full to avoid double-capture on startup
|
||||
if model_runner.server_args.enable_return_hidden_states:
|
||||
# If returning hidden states is enabled, or if speculative prefill needs
|
||||
# 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.max_num_tokens = (
|
||||
@@ -416,7 +420,7 @@ class PiecewiseCudaGraphRunner:
|
||||
mrope_positions=mrope_positions,
|
||||
spec_algorithm=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_cpu=num_tokens,
|
||||
global_forward_mode=ForwardMode.EXTEND,
|
||||
@@ -587,7 +591,7 @@ class PiecewiseCudaGraphRunner:
|
||||
mrope_positions=mrope_positions,
|
||||
spec_algorithm=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_cpu=num_tokens,
|
||||
global_forward_mode=ForwardMode.EXTEND,
|
||||
|
||||
Reference in New Issue
Block a user