diff --git a/scripts/train.py b/scripts/train.py index 0697baaef..8b5d4bfc8 100644 --- a/scripts/train.py +++ b/scripts/train.py @@ -620,10 +620,9 @@ def parse_args(): parser.add_argument( "--draft-arch", type=str, - default="llama", + default="qwen3", choices=list(DRAFT_ARCH_CONFIGS.keys()), - help="Architecture for draft decoder layers. Defaults to 'llama'. " - "Note: only 'llama' is currently supported in vLLM for inference.", + help="Architecture for draft decoder layers. Defaults to 'qwen3'.", ) parser.add_argument( "--draft-hidden-act", diff --git a/src/speculators/models/eagle3/config.py b/src/speculators/models/eagle3/config.py index a3201c452..dbcdb7bdb 100644 --- a/src/speculators/models/eagle3/config.py +++ b/src/speculators/models/eagle3/config.py @@ -2,7 +2,7 @@ from pydantic import Field, field_serializer, field_validator from transformers import AutoConfig, PretrainedConfig -from transformers.models.llama.configuration_llama import LlamaConfig +from transformers.models.qwen3.configuration_qwen3 import Qwen3Config from speculators import SpeculatorModelConfig @@ -31,7 +31,7 @@ class Eagle3SpeculatorConfig(SpeculatorModelConfig): ) transformer_layer_config: PretrainedConfig = Field( - default_factory=LlamaConfig, + default_factory=Qwen3Config, description="Configuration for the transformer decoder layer", ) @@ -79,7 +79,7 @@ def serialize_transformer_config(self, value: PretrainedConfig) -> dict: def validate_transformer_config(cls, value: Any) -> PretrainedConfig: """Validate and convert transformer config.""" if isinstance(value, dict): - config_class: type[PretrainedConfig] = LlamaConfig + config_class: type[PretrainedConfig] = Qwen3Config if "model_type" in value: config_class = AutoConfig.for_model( model_type=value["model_type"]