File size: 5,427 Bytes
b29a6e1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
from transformers.configuration_utils import PretrainedConfig


class SarvamMLAConfig(PretrainedConfig):
    model_type = "sarvam_mla"

    base_model_pp_plan = {
        "embed_tokens": (["input_ids"], ["inputs_embeds"]),
        "layers": (["hidden_states", "attention_mask"], ["hidden_states"]),
        "norm": (["hidden_states"], ["hidden_states"]),
    }

    base_model_tp_plan = {
        "layers.*.self_attn.q_proj": "colwise",
        "layers.*.self_attn.kv_b_proj": "colwise",
        "layers.*.self_attn.o_proj": "rowwise",
    }

    def __init__(
        self,
        vocab_size: int = 262144,
        hidden_size: int = 4096,
        num_hidden_layers: int = 32,
        intermediate_size: int = 16384,
        moe_intermediate_size: int = 2048,
        num_experts: int = 128,
        num_experts_per_tok: int = 8,
        num_shared_experts: int = 1,
        first_k_dense_replace: int = 1,
        num_attention_heads: int = 64,
        qk_rope_head_dim: int = 64,
        qk_nope_head_dim: int = 128,
        kv_lora_rank: int = 512,
        v_head_dim: int = 128,
        max_position_embeddings: int = 4096,
        rope_theta: float = 10000.0,
        rope_scaling: dict = None,
        attention_dropout: float = 0.0,
        output_dropout: float = 0.0,
        rms_norm_eps: float = 1e-6,
        hidden_act: str = "silu",
        use_cache: bool = True,
        use_qk_norm: bool = True,
        moe_router_enable_expert_bias: bool = True,
        routed_scaling_factor: float = 2.5,
        output_router_logits: bool = False,
        tie_word_embeddings: bool = False,
        pad_token_id: int = 0,
        eos_token_id: int = 1,
        embedding_dropout: float = 0.0,
        initializer_range: float = 0.006,
        attn_implementation: str = "eager",
        **kwargs,
    ):
        # core geometry
        self.vocab_size = vocab_size
        self.hidden_size = hidden_size
        self.num_hidden_layers = num_hidden_layers
        self.intermediate_size = intermediate_size
        self.num_attention_heads = num_attention_heads
        self.max_position_embeddings = max_position_embeddings

        # MLA geometry
        self.qk_rope_head_dim = qk_rope_head_dim
        self.qk_nope_head_dim = qk_nope_head_dim
        self.kv_lora_rank = kv_lora_rank
        self.v_head_dim = v_head_dim
        # convenient derived dim
        self.q_head_dim = qk_rope_head_dim + qk_nope_head_dim
        # vLLM MLA expects "head size" = Lkv + R, not hidden_size/num_heads.
        self.head_dim = int(self.kv_lora_rank + self.qk_rope_head_dim)

        # MoE
        self.moe_intermediate_size = moe_intermediate_size
        self.num_experts = num_experts
        self.num_experts_per_tok = num_experts_per_tok
        self.num_shared_experts = num_shared_experts
        self.first_k_dense_replace = first_k_dense_replace

        # Router
        self.moe_router_enable_expert_bias = moe_router_enable_expert_bias
        self.routed_scaling_factor = routed_scaling_factor
        self.output_router_logits = output_router_logits

        # dropouts / norms / init
        self.attention_dropout = attention_dropout
        self.output_dropout = output_dropout
        self.embedding_dropout = embedding_dropout
        self.rms_norm_eps = rms_norm_eps
        self.initializer_range = initializer_range
        self.hidden_act = hidden_act

        # rope / cache
        self.rope_theta = rope_theta
        self.use_cache = use_cache
        self.use_qk_norm = use_qk_norm
        self.rope_scaling = rope_scaling
        self.default_theta = 10000.0
        
        if self.rope_scaling is None:
            self.rope_scaling = {
                'beta_fast': 32,
                'beta_slow': 1,
                'factor': 40,
                'mscale': 1.0,
                'mscale_all_dim': 1.0,
                'original_max_position_embeddings': 4096,
                'rope_type': 'deepseek_yarn',
            }

        self.attn_implementation = attn_implementation
        self._attn_implementation = attn_implementation
        
        if "_attn_implementation" in kwargs:
            self._attn_implementation = kwargs.pop("_attn_implementation")
            if hasattr(self, "attn_implementation"):
                self.attn_implementation = self._attn_implementation

        super().__init__(
            pad_token_id=pad_token_id,
            eos_token_id=eos_token_id,
            tie_word_embeddings=tie_word_embeddings,
            **kwargs,
        )

    def convert_rope_params_to_dict(self, ignore_keys_at_rope_validation: set | None = None, **kwargs):
        rope_scaling = kwargs.pop("rope_scaling", None)
        self.rope_parameters = rope_scaling or self.rope_parameters
        self.rope_parameters = self.rope_parameters if self.rope_parameters is not None else {}

        # Standardize and validate the correctness of rotary position embeddings parameters
        self.rope_parameters.setdefault("rope_theta", kwargs.pop("rope_theta", self.default_theta))
        self.standardize_rope_params()
        self.validate_rope(ignore_keys=ignore_keys_at_rope_validation)

        # Convert to float because RoPE fn expect a float. Models on the hub were saved as int
        for key in ["beta_fast", "beta_slow", "factor"]:
            if key in self.rope_parameters:
                self.rope_parameters[key] = float(self.rope_parameters[key])
        return kwargs