[Speculative] Fix Eagle3/DFLASH aux hidden state capture during CUDA graph init (#22836)
This commit is contained in:
@@ -645,25 +645,6 @@ class CudaGraphRunner:
|
|||||||
|
|
||||||
self.tbo_plugin = TboCudaGraphRunnerPlugin()
|
self.tbo_plugin = TboCudaGraphRunnerPlugin()
|
||||||
|
|
||||||
# Speculative_inference
|
|
||||||
if (
|
|
||||||
model_runner.spec_algorithm.is_eagle3()
|
|
||||||
and model_runner.eagle_use_aux_hidden_state
|
|
||||||
):
|
|
||||||
self.model_runner.model.set_eagle3_layers_to_capture()
|
|
||||||
if (
|
|
||||||
model_runner.spec_algorithm.is_dflash()
|
|
||||||
and model_runner.dflash_use_aux_hidden_state
|
|
||||||
):
|
|
||||||
if not hasattr(self.model_runner.model, "set_dflash_layers_to_capture"):
|
|
||||||
raise ValueError(
|
|
||||||
f"Model {self.model_runner.model.__class__.__name__} does not implement set_dflash_layers_to_capture, "
|
|
||||||
"which is required for DFLASH aux hidden capture."
|
|
||||||
)
|
|
||||||
self.model_runner.model.set_dflash_layers_to_capture(
|
|
||||||
self.model_runner.dflash_target_layer_ids
|
|
||||||
)
|
|
||||||
|
|
||||||
# Capture
|
# Capture
|
||||||
try:
|
try:
|
||||||
with model_capture_mode():
|
with model_capture_mode():
|
||||||
|
|||||||
@@ -711,6 +711,10 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
# Init routed experts capturer
|
# Init routed experts capturer
|
||||||
self.init_routed_experts_capturer()
|
self.init_routed_experts_capturer()
|
||||||
|
|
||||||
|
# Must be called BEFORE init_device_graphs() so CUDA graph capture
|
||||||
|
# runs with aux hidden state capture enabled.
|
||||||
|
self.init_aux_hidden_state_capture()
|
||||||
|
|
||||||
if self.device == "cuda" or self.device == "musa":
|
if self.device == "cuda" or self.device == "musa":
|
||||||
self.init_cublas()
|
self.init_cublas()
|
||||||
self.init_attention_backend()
|
self.init_attention_backend()
|
||||||
@@ -727,19 +731,6 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
if server_args.forward_hooks:
|
if server_args.forward_hooks:
|
||||||
register_forward_hooks(self.model, server_args.forward_hooks)
|
register_forward_hooks(self.model, server_args.forward_hooks)
|
||||||
|
|
||||||
if self.eagle_use_aux_hidden_state:
|
|
||||||
self.model.set_eagle3_layers_to_capture(
|
|
||||||
self.eagle_aux_hidden_state_layer_ids
|
|
||||||
)
|
|
||||||
|
|
||||||
if self.dflash_use_aux_hidden_state:
|
|
||||||
if not hasattr(self.model, "set_dflash_layers_to_capture"):
|
|
||||||
raise ValueError(
|
|
||||||
f"Model {self.model.__class__.__name__} does not implement set_dflash_layers_to_capture, "
|
|
||||||
"which is required for DFLASH."
|
|
||||||
)
|
|
||||||
self.model.set_dflash_layers_to_capture(self.dflash_target_layer_ids)
|
|
||||||
|
|
||||||
# Initialize piecewise CUDA graph
|
# Initialize piecewise CUDA graph
|
||||||
self.init_piecewise_cuda_graphs()
|
self.init_piecewise_cuda_graphs()
|
||||||
|
|
||||||
@@ -764,6 +755,24 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def init_aux_hidden_state_capture(self):
|
||||||
|
"""Configure auxiliary hidden state capture for speculative decoding.
|
||||||
|
|
||||||
|
Must be called before CUDA graph capture so the captured graphs
|
||||||
|
include aux hidden state output paths.
|
||||||
|
"""
|
||||||
|
if self.eagle_use_aux_hidden_state:
|
||||||
|
self.model.set_eagle3_layers_to_capture(
|
||||||
|
self.eagle_aux_hidden_state_layer_ids
|
||||||
|
)
|
||||||
|
if self.dflash_use_aux_hidden_state:
|
||||||
|
if not hasattr(self.model, "set_dflash_layers_to_capture"):
|
||||||
|
raise ValueError(
|
||||||
|
f"Model {self.model.__class__.__name__} does not implement "
|
||||||
|
"set_dflash_layers_to_capture, which is required for DFLASH."
|
||||||
|
)
|
||||||
|
self.model.set_dflash_layers_to_capture(self.dflash_target_layer_ids)
|
||||||
|
|
||||||
def remote_instance_init_transfer_engine(self):
|
def remote_instance_init_transfer_engine(self):
|
||||||
try:
|
try:
|
||||||
from mooncake.engine import TransferEngine
|
from mooncake.engine import TransferEngine
|
||||||
@@ -2259,10 +2268,8 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
)
|
)
|
||||||
self.server_args.enable_torch_compile = False
|
self.server_args.enable_torch_compile = False
|
||||||
|
|
||||||
if self.eagle_use_aux_hidden_state:
|
# NOTE: aux hidden state capture (eagle3/dflash) is already
|
||||||
self.model.set_eagle3_layers_to_capture()
|
# configured by init_aux_hidden_state_capture() in initialize().
|
||||||
if self.dflash_use_aux_hidden_state:
|
|
||||||
self.model.set_dflash_layers_to_capture(self.dflash_target_layer_ids)
|
|
||||||
|
|
||||||
require_mlp_tp_gather_ = require_mlp_tp_gather(self.server_args)
|
require_mlp_tp_gather_ = require_mlp_tp_gather(self.server_args)
|
||||||
if require_gathered_buffer(self.server_args):
|
if require_gathered_buffer(self.server_args):
|
||||||
|
|||||||
Reference in New Issue
Block a user