krishnateja95's picture
Convert Inkling DFlash from sglang to speculators format
b8f2de3 verified
Raw
History Blame Contribute Delete
3.06 kB
from typing import Any, Literal
from pydantic import Field, field_serializer, field_validator
from transformers import AutoConfig, PretrainedConfig
from transformers.models.qwen3.modeling_qwen3 import (
Qwen3Config,
)
from speculators import SpeculatorModelConfig
__all__ = [
"DFlashSpeculatorConfig",
]
@SpeculatorModelConfig.register("dflash")
class DFlashSpeculatorConfig(SpeculatorModelConfig):
"""
Configuration for DFlash speculator with Inkling support.
Extends standard DFlash config with embed_norm and mup scaling.
"""
speculators_model_type: Literal["dflash"] = "dflash"
architectures: list[str] = Field(
default_factory=lambda: ["DFlashDraftModel"],
description="Model architectures that can load these weights",
)
transformer_layer_config: PretrainedConfig = Field(
default_factory=Qwen3Config,
description="Configuration for the transformer decoder layer",
)
draft_vocab_size: int = Field(
default=201024,
description="Size of draft model vocabulary for speculation",
)
block_size: int = Field(
default=16,
description="Default size of the draft block predicted with a forward pass",
)
max_anchors: int = Field(
default=256,
description="Maximum number of anchor positions to sample during training",
)
target_hidden_size: int | None = Field(
default=None,
description="Hidden size of the target model (if different from draft model)",
)
aux_hidden_state_layer_ids: list[int] | None = Field(
default=None,
description="Layer IDs of the DFlash auxiliary hidden state layers",
)
mask_token_id: int | None = Field(
default=None,
description="Token ID used for masking",
)
use_embed_norm: bool = Field(
default=False,
description="Apply RMSNorm after token embedding (Inkling-specific)",
)
logits_mup_width_multiplier: float | None = Field(
default=None,
description="muP width multiplier for logit scaling (Inkling-specific)",
)
@field_serializer("transformer_layer_config")
def serialize_transformer_config(self, value: PretrainedConfig) -> dict:
"""Serialize transformer config to dict."""
return value.to_diff_dict()
@field_validator("transformer_layer_config", mode="before")
@classmethod
def validate_transformer_config(cls, value: Any) -> PretrainedConfig:
"""Validate and convert transformer config."""
if isinstance(value, dict):
config_class: type[PretrainedConfig] = Qwen3Config
if "model_type" in value:
config_class = AutoConfig.for_model(
model_type=value["model_type"]
).__class__
return config_class(**value)
return value
@property
def target_vocab_size(self) -> int:
"""Get target vocabulary size from transformer config."""
return self.transformer_layer_config.vocab_size