CoolFace
Modelpublic

sdougbrown/North-Mini-Code-1.0-dflash

sourceHugging Faceapache-2.0updated 1mo agoView on Hugging Face
0likes24downloads
config.py106 linesDownload Raw Back to root
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