File size: 1,473 Bytes
d0888d9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
from typing import Literal

from pydantic import Field, model_validator

from speculators import SpeculatorModelConfig
from speculators.models.dflash.config import DFlashSpeculatorConfig


@SpeculatorModelConfig.register("dflash2")
class DFlash2SpeculatorConfig(DFlashSpeculatorConfig):
    """DFlash2 draft-model configuration."""

    speculators_model_type: Literal["dflash2"] = "dflash2"  # type: ignore[assignment]

    architectures: list[str] = Field(
        default_factory=lambda: ["DFlash2DraftModel"],
    )

    sliding_window_non_causal: bool = True

    conv_kernel_size: int = Field(default=2, ge=1)
    conv_group_size: int = Field(default=16, ge=1)
    selector_rank: int = Field(default=256, ge=1)
    selector_top_k: int = Field(default=16, ge=1)

    draft_ffn_type: Literal["dense", "moe"] = "dense"
    num_experts: int = Field(default=256, ge=1)
    num_experts_per_tok: int = Field(default=8, ge=1)
    moe_intermediate_size: int = Field(default=512, ge=1)
    shared_expert_intermediate_size: int = Field(default=512, ge=1)

    @model_validator(mode="after")
    def validate_moe_routing(self) -> "DFlash2SpeculatorConfig":
        if (
            self.draft_ffn_type == "moe"
            and self.num_experts_per_tok > self.num_experts
        ):
            raise ValueError(
                "num_experts_per_tok cannot exceed num_experts: "
                f"{self.num_experts_per_tok} > {self.num_experts}."
            )
        return self