Eagle: GPT-OSS Eagle v2 support (#14920)
Signed-off-by: Izzy Putterman <iputterman@nvidia.com>
This commit is contained in:
@@ -314,6 +314,32 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
self.remote_instance_transfer_engine = None
|
||||
self.remote_instance_transfer_engine_session_id = ""
|
||||
self.remote_instance_transfer_engine_weight_info = None
|
||||
# auxiliary hidden capture mode. TODO: expose this to server args?
|
||||
self.eagle_use_aux_hidden_state = False
|
||||
if self.spec_algorithm.is_eagle3() and not self.is_draft_worker:
|
||||
# load draft config
|
||||
draft_model_config = ModelConfig.from_server_args(
|
||||
server_args,
|
||||
model_path=(server_args.speculative_draft_model_path),
|
||||
model_revision=server_args.speculative_draft_model_revision,
|
||||
is_draft_model=True,
|
||||
)
|
||||
self.eagle_use_aux_hidden_state = True
|
||||
|
||||
try:
|
||||
# get the aux layer from draft model config
|
||||
eagle_config = getattr(
|
||||
draft_model_config.hf_config, "eagle_config", None
|
||||
)
|
||||
self.eagle_use_aux_hidden_state = eagle_config.get(
|
||||
"use_aux_hidden_state", True
|
||||
)
|
||||
self.eagle_aux_hidden_state_layer_ids = eagle_config[
|
||||
"eagle_aux_hidden_state_layer_ids"
|
||||
]
|
||||
except:
|
||||
# if there is no aux layer, set to None
|
||||
self.eagle_aux_hidden_state_layer_ids = None
|
||||
|
||||
# Apply the rank zero filter to logger
|
||||
if server_args.show_time_cost:
|
||||
@@ -550,30 +576,11 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
if server_args.forward_hooks:
|
||||
register_forward_hooks(self.model, server_args.forward_hooks)
|
||||
|
||||
# auxiliary hidden capture mode. TODO: expose this to server args?
|
||||
if self.spec_algorithm.is_eagle3() and not self.is_draft_worker:
|
||||
# load draft config
|
||||
draft_model_config = ModelConfig.from_server_args(
|
||||
server_args,
|
||||
model_path=(server_args.speculative_draft_model_path),
|
||||
model_revision=server_args.speculative_draft_model_revision,
|
||||
is_draft_model=True,
|
||||
if self.eagle_use_aux_hidden_state:
|
||||
self.model.set_eagle3_layers_to_capture(
|
||||
self.eagle_aux_hidden_state_layer_ids
|
||||
)
|
||||
|
||||
try:
|
||||
# get the aux layer from draft model config
|
||||
eagle_config = getattr(
|
||||
draft_model_config.hf_config, "eagle_config", None
|
||||
)
|
||||
eagle_aux_hidden_state_layer_ids = eagle_config[
|
||||
"eagle_aux_hidden_state_layer_ids"
|
||||
]
|
||||
except:
|
||||
# if there is no aux layer, set to None
|
||||
eagle_aux_hidden_state_layer_ids = None
|
||||
|
||||
self.model.set_eagle3_layers_to_capture(eagle_aux_hidden_state_layer_ids)
|
||||
|
||||
# Initialize piecewise CUDA graph
|
||||
self.init_piecewise_cuda_graphs()
|
||||
|
||||
@@ -1745,7 +1752,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
if self.server_args.enable_torch_compile:
|
||||
set_torch_compile_config()
|
||||
|
||||
if self.spec_algorithm.is_eagle3():
|
||||
if self.eagle_use_aux_hidden_state:
|
||||
self.model.set_eagle3_layers_to_capture()
|
||||
|
||||
require_mlp_tp_gather_ = require_mlp_tp_gather(self.server_args)
|
||||
|
||||
Reference in New Issue
Block a user