sdougbrown/North-Mini-Code-1.0-dflash
024
1from typing import Any, Literal2 3from pydantic import Field, field_serializer, field_validator4from transformers import AutoConfig, PretrainedConfig5from transformers.models.qwen3.modeling_qwen3 import (6 Qwen3Config,7)8 9from speculators import SpeculatorModelConfig10 11__all__ = [12 "DFlashSpeculatorConfig",13]14 15 16@SpeculatorModelConfig.register("dflash")17class DFlashSpeculatorConfig(SpeculatorModelConfig):18 """19 Configuration for DFlash speculator with vocabulary mapping.20 21 DFlash features vocabulary mapping between the 32K draft vocabulary and22 the 262,144-token target vocabulary, enabling reduced-vocabulary speculation.23 24 :param transformer_layer_config: Configuration for the transformer decoder layer25 :param draft_vocab_size: Size of draft model vocabulary for speculation26 """27 28 speculators_model_type: Literal["dflash"] = "dflash"29 architectures: list[str] = Field(30 default_factory=lambda: ["DFlashSpeculator"],31 description="Model architectures that can load these weights",32 )33 34 transformer_layer_config: PretrainedConfig = Field(35 default_factory=Qwen3Config,36 description="Configuration for the transformer decoder layer",37 )38 39 draft_vocab_size: int = Field(40 default=32000,41 description="Size of draft model vocabulary for speculation",42 )43 44 block_size: int = Field(45 default=8,46 description=(47 "Default size of the draft block predicted with a forward pass of the model"48 ),49 )50 51 target_hidden_size: int | None = Field(52 default=None,53 description="Hidden size of the target model (if different from draft model)",54 )55 56 aux_hidden_state_layer_ids: list[int] | None = Field(57 default=None,58 description="Layer IDs of the DFlash auxiliary hidden state layers",59 )60 61 mask_token_id: int | None = Field(62 default=None,63 description="Token ID used for masking",64 )65 66 sliding_window_non_causal: bool = Field(67 default=False,68 description="Use non-causal (bidirectional) masking within draft blocks for "69 "sliding window attention layers. Full attention layers are always "70 "bidirectional.",71 )72 73 sample_from_anchor: bool = Field(74 default=False,75 description=(76 "Whether to sample from the anchor position. "77 "False: anchor is the bonus token, only mask tokens predict "78 "(block_size-1 speculative tokens). "79 "True: sample from anchor and all mask positions "80 "(block_size speculative tokens). "81 ),82 )83 84 @field_serializer("transformer_layer_config")85 def serialize_transformer_config(self, value: PretrainedConfig) -> dict:86 """Serialize transformer config to dict."""87 return value.to_diff_dict()88 89 @field_validator("transformer_layer_config", mode="before")90 @classmethod91 def validate_transformer_config(cls, value: Any) -> PretrainedConfig:92 """Validate and convert transformer config."""93 if isinstance(value, dict):94 config_class: type[PretrainedConfig] = Qwen3Config95 if "model_type" in value:96 config_class = AutoConfig.for_model(97 model_type=value["model_type"]98 ).__class__99 return config_class(**value)100 return value101 102 @property103 def target_vocab_size(self) -> int:104 """Get target vocabulary size from transformer config."""105 return self.transformer_layer_config.vocab_size106 