bupalinyu commited on
Commit
cac6d3c
·
verified ·
1 Parent(s): 5ddfa99

Upload configuration_bailing_moe_v3.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. configuration_bailing_moe_v3.py +173 -0
configuration_bailing_moe_v3.py ADDED
@@ -0,0 +1,173 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Bailing MoE V2 model configuration"""
2
+
3
+ from transformers.configuration_utils import PretrainedConfig
4
+
5
+
6
+ class BailingMoeV3Config(PretrainedConfig):
7
+
8
+ def __init__(
9
+ self,
10
+ vocab_size=157184,
11
+ hidden_size=2048,
12
+ intermediate_size=5120,
13
+ num_hidden_layers=20,
14
+ num_attention_heads=16,
15
+ num_key_value_heads=4,
16
+ hidden_act="silu",
17
+ use_qkv_bias=False, # bailing only
18
+ use_bias=False, # bailing only
19
+ rms_norm_eps=1e-06,
20
+ tie_word_embeddings=False, # PretrainedConfig key, here change default value.
21
+ embedding_dropout=0.0,
22
+ attention_dropout=0.0,
23
+ output_dropout=0.0,
24
+ initializer_range=0.02,
25
+ max_position_embeddings=32768,
26
+ rope_theta=600000.0,
27
+ use_cache=True,
28
+ max_window_layers=20,
29
+ rope_scaling=None,
30
+ pad_token_id=156892,
31
+ eos_token_id=156892,
32
+ num_experts=256,
33
+ num_shared_experts=1,
34
+ num_experts_per_tok=8,
35
+ n_group=8,
36
+ topk_group=4,
37
+ moe_intermediate_size=512,
38
+ moe_shared_expert_intermediate_size=512,
39
+ first_k_dense_replace=1,
40
+ head_dim=128,
41
+ output_router_logits=False,
42
+ use_qk_norm=True,
43
+ num_nextn_predict_layers=0,
44
+ mtp_loss_scaling_factor=0,
45
+ moe_router_enable_expert_bias=True,
46
+ routed_scaling_factor=1.0,
47
+ layer_group_size=5,
48
+ kv_lora_rank=512,
49
+ q_lora_rank=None,
50
+ qk_rope_head_dim=64,
51
+ v_head_dim=128,
52
+ qk_nope_head_dim=128,
53
+ rope_interleave=True,
54
+ score_function="sigmoid",
55
+ scoring_func="sigmoid",
56
+ seq_aux=True,
57
+ topk_method="noaux_tc",
58
+ router_dtype="fp32",
59
+ gated_attention_proj_granularity_type=None,
60
+ no_kda_lora=False,
61
+ kda_safe_gate=False,
62
+ kda_lower_bound=None,
63
+ short_conv_kernel_size=4,
64
+ pregate_enabled=False,
65
+ pregate_hidden=512,
66
+ pregate_inference=False,
67
+ pregate_use_prev_topk=True,
68
+ pregate_init_router=False,
69
+ pregate_use_prev_token=True,
70
+ pregate_shallow_hidden=1024,
71
+ pregate_shallow_layers=5,
72
+ pregate_shallow_loss_weight=1.5,
73
+ pregate_start_layer=7,
74
+ # v6: cross-token pre-gate. pregate_N at position t is trained /
75
+ # consumed to route layer N+1 at position t+1 (one-token-ahead),
76
+ # matching the MLX single-sync fast path. False keeps the v5
77
+ # same-token semantics.
78
+ pregate_cross_token=False,
79
+ # Top-k OPD (on-policy distillation) for the pre-gate modules.
80
+ # strategy: "none" (full-vocab KL, legacy) | "union" | "only_stu" |
81
+ # "only_tch" | "intersection". The support is always the top-8 expert
82
+ # sets selected by the student pre-gate / teacher router, so the KL is
83
+ # sparse over <=16 experts instead of all 128.
84
+ pregate_opd_strategy="none",
85
+ pregate_opd_topk=8,
86
+ pregate_opd_temperature=1.0,
87
+ pregate_opd_weight=1.0,
88
+ pregate_opd_ce_weight=1.0,
89
+ pregate_opd_weight_mode="teacher_p",
90
+ **kwargs,
91
+ ):
92
+ self.num_hidden_layers = num_hidden_layers
93
+ self.vocab_size = vocab_size
94
+ self.hidden_size = hidden_size
95
+ self.intermediate_size = intermediate_size
96
+ self.num_attention_heads = num_attention_heads
97
+ self.num_key_value_heads = num_key_value_heads
98
+ self.hidden_act = hidden_act
99
+ self.use_qkv_bias = use_qkv_bias
100
+ self.use_bias = use_bias
101
+ self.rms_norm_eps = rms_norm_eps
102
+ self.embedding_dropout = embedding_dropout
103
+ self.attention_dropout = attention_dropout
104
+ self.output_dropout = output_dropout
105
+ self.num_nextn_predict_layers = num_nextn_predict_layers
106
+ self.mtp_loss_scaling_factor = mtp_loss_scaling_factor
107
+ self.initializer_range = initializer_range
108
+ self.max_position_embeddings = max_position_embeddings
109
+ self.rope_theta = rope_theta
110
+ self.use_cache = use_cache
111
+ self.max_window_layers = max_window_layers
112
+ self.head_dim = head_dim or self.hidden_size // self.num_attention_heads
113
+ self.rope_scaling = rope_scaling
114
+ self.use_qk_norm = use_qk_norm
115
+ self.moe_router_enable_expert_bias = moe_router_enable_expert_bias
116
+ self.routed_scaling_factor = routed_scaling_factor
117
+
118
+ # MoE configs
119
+ self.num_experts = num_experts
120
+ self.num_shared_experts = num_shared_experts
121
+ self.num_experts_per_tok = num_experts_per_tok
122
+ self.n_group = n_group
123
+ self.topk_group = topk_group
124
+ self.moe_intermediate_size = moe_intermediate_size
125
+ self.moe_shared_expert_intermediate_size = moe_shared_expert_intermediate_size
126
+ self.first_k_dense_replace = first_k_dense_replace
127
+ self.output_router_logits = output_router_logits
128
+
129
+ # Linear configs
130
+ self.layer_group_size = layer_group_size
131
+ # mla
132
+ self.kv_lora_rank = kv_lora_rank
133
+ self.q_lora_rank = q_lora_rank
134
+ self.qk_rope_head_dim = qk_rope_head_dim
135
+
136
+ self.score_function = score_function
137
+ self.scoring_func = scoring_func
138
+ self.seq_aux = seq_aux
139
+ self.topk_method = topk_method
140
+ self.v_head_dim = v_head_dim
141
+ self.qk_nope_head_dim = qk_nope_head_dim
142
+ self.qk_head_dim = qk_nope_head_dim + qk_rope_head_dim
143
+ self.rope_interleave = rope_interleave
144
+ self.router_dtype = router_dtype
145
+ self.gated_attention_proj_granularity_type = gated_attention_proj_granularity_type
146
+ self.no_kda_lora = no_kda_lora
147
+ self.kda_safe_gate = kda_safe_gate
148
+ self.kda_lower_bound = kda_lower_bound
149
+ self.short_conv_kernel_size = short_conv_kernel_size
150
+ # Pre-gated MoE (arXiv 2308.12066): layer N's pre-gate selects the
151
+ # experts for MoE layer N+1. pregate_enabled turns the modules on;
152
+ # pregate_inference makes the deployed model use the pre-gate as the
153
+ # router (the first MoE layer keeps its original gate, zero lead).
154
+ self.pregate_enabled = pregate_enabled
155
+ self.pregate_hidden = pregate_hidden
156
+ self.pregate_inference = pregate_inference
157
+ self.pregate_use_prev_topk = pregate_use_prev_topk
158
+ self.pregate_init_router = pregate_init_router
159
+ self.pregate_use_prev_token = pregate_use_prev_token
160
+ self.pregate_shallow_hidden = pregate_shallow_hidden
161
+ self.pregate_shallow_layers = pregate_shallow_layers
162
+ self.pregate_shallow_loss_weight = pregate_shallow_loss_weight
163
+ self.pregate_start_layer = pregate_start_layer
164
+ self.pregate_cross_token = pregate_cross_token
165
+ self.pregate_opd_strategy = pregate_opd_strategy
166
+ self.pregate_opd_topk = int(pregate_opd_topk)
167
+ self.pregate_opd_temperature = float(pregate_opd_temperature)
168
+ self.pregate_opd_weight = float(pregate_opd_weight)
169
+ self.pregate_opd_ce_weight = float(pregate_opd_ce_weight)
170
+ self.pregate_opd_weight_mode = pregate_opd_weight_mode
171
+ super().__init__(
172
+ pad_token_id=pad_token_id, eos_token_id=eos_token_id, tie_word_embeddings=tie_word_embeddings, **kwargs
173
+ )