1use super::*;
4
5impl Gpu<'_> {
6 #[allow(clippy::too_many_arguments)]
8 pub(super) fn router_b(
9 &self,
10 enc: &Enc,
11 moe: &Moe,
12 nb: usize,
13 x: &Buf,
14 logits: &Buf,
15 idx: &Buf,
16 w: &Buf,
17 k: usize,
18 ) {
19 let c = &self.p.cfg;
20 let rows = c.num_experts as u32;
21 let cols = c.hidden_size as u32;
22 let nbu = nb as u32;
23
24 self.dispatch(
25 enc,
26 &self.pipes.bf16_matvec_b,
27 |e| {
28 self.bind(e, 0, &self.dense, moe.router.0);
29 self.bind(e, 1, x, 0);
30 self.bind(e, 2, logits, 0);
31 set_bytes(e, 3, &rows);
32 set_bytes(e, 4, &cols);
33 set_bytes(e, 5, &nbu);
34 },
35 (nb * c.num_experts).div_ceil(4),
36 128,
37 true,
38 );
39
40 let k = k as u32;
41 let renorm = c.norm_topk_prob as u32;
42
43 self.dispatch(
44 enc,
45 &self.pipes.topk_softmax_b,
46 |e| {
47 self.bind(e, 0, logits, 0);
48 self.bind(e, 1, idx, 0);
49 self.bind(e, 2, w, 0);
50 set_bytes(e, 3, &rows);
51 set_bytes(e, 4, &k);
52 set_bytes(e, 5, &renorm);
53 },
54 nb,
55 c.num_experts.max(32),
56 true,
57 );
58 }
59
60 pub(super) fn experts_b(&self, enc: &Enc, moe: &Moe, slot_row: usize, nb: usize, part: u32) {
66 let c = &self.p.cfg;
67 let s = &self.scratch;
68 let l = super::super::lowbit::Layout::four_bit(&self.p.manifest.experts);
69 let inter = self.p.manifest.experts.inter as u32;
70 let h = c.hidden_size as u32;
71 let k = c.num_experts_per_tok;
72 let shared = !self.skips("shared");
73
74 if self.skips("experts") {
75 if part == 0 {
76 self.zero(enc, &s.moe_out, nb as u32 * h);
77 }
78
79 return;
80 }
81
82 let n_max = (k * nb) as u32;
83 let mp = MoeBParams {
84 inter,
85 hidden: h,
86 n_u: n_max,
87 gate_w: l.gate_w as u32,
88 up_w: l.up_w as u32,
89 down_w: l.down_w as u32,
90 gate_s: l.gate_s as u32,
91 gate_b: l.gate_b as u32,
92 up_s: l.up_s as u32,
93 up_b: l.up_b as u32,
94 down_s: l.down_s as u32,
95 down_b: l.down_b as u32,
96 layer: slot_row as u32,
97 shared: shared as u32,
98 nb: nb as u32,
99 sh_gate_w: moe.sg.w as u32,
100 sh_gate_s: moe.sg.s as u32,
101 sh_gate_b: moe.sg.b as u32,
102 sh_up_w: moe.su.w as u32,
103 sh_up_s: moe.su.s as u32,
104 sh_up_b: moe.su.b as u32,
105 sh_down_w: moe.sd.w as u32,
106 sh_down_s: moe.sd.s as u32,
107 sh_down_b: moe.sd.b as u32,
108 sh_gate_vec: moe.shared_gate.0 as u32,
109 part,
110 low_gate_w: self.low_bit_store.map_or(0, |s| s.gate_w as u32),
111 low_up_w: self.low_bit_store.map_or(0, |s| s.up_w as u32),
112 low_down_w: self.low_bit_store.map_or(0, |s| s.down_w as u32),
113 low_gate_s: self.low_bit_store.map_or(0, |s| s.gate_s as u32),
114 low_gate_b: self.low_bit_store.map_or(0, |s| s.gate_b as u32),
115 low_up_s: self.low_bit_store.map_or(0, |s| s.up_s as u32),
116 low_up_b: self.low_bit_store.map_or(0, |s| s.up_b as u32),
117 low_down_s: self.low_bit_store.map_or(0, |s| s.down_s as u32),
118 low_down_b: self.low_bit_store.map_or(0, |s| s.down_b as u32),
119 };
120 let n_exp = (n_max + 1) as usize;
121 let rows2 = 2 * inter as usize;
122
123 self.dispatch(
124 enc,
125 &self.pipes.moe_gate_up_b[nb - 1],
126 |e| {
127 self.bind(e, 0, &s.hc.h1.xe, 0);
128 self.bind(e, 1, &s.hc.h1.xo, 0);
129 self.bind(e, 2, &s.hc.h1.xsum, 0);
130 self.bind(e, 3, &s.gate_e, 0);
131 set_bytes(e, 4, &mp);
132 self.bind(e, 5, &self.slot_tab, 0);
133 self.bind(e, 6, &self.dense, 0);
134 },
135 (n_exp * rows2).div_ceil(4),
136 128,
137 true,
138 );
139 self.dispatch(
140 enc,
141 &self.pipes.moe_act_b,
142 |e| {
143 self.bind(e, 0, &s.gate_e, 0);
144 self.bind(e, 1, &s.hx.xe, 0);
145 self.bind(e, 2, &s.hx.xo, 0);
146 self.bind(e, 3, &s.hx.xsum, 0);
147 set_bytes(e, 4, &mp);
148 self.bind(e, 5, &self.slot_tab, 0);
149 },
150 n_exp * nb * inter as usize / 2,
151 256,
152 false,
153 );
154 self.dispatch(
155 enc,
156 &self.pipes.moe_down_b[nb - 1],
157 |e| {
158 self.bind(e, 0, &s.hx.xe, 0);
159 self.bind(e, 1, &s.hx.xo, 0);
160 self.bind(e, 2, &s.hx.xsum, 0);
161 self.bind(e, 3, &s.y_e, 0);
162 set_bytes(e, 4, &mp);
163 self.bind(e, 5, &self.slot_tab, 0);
164 self.bind(e, 6, &self.dense, 0);
165 },
167 (n_exp * (h as usize / 2)).div_ceil(4),
168 128,
169 true,
170 );
171 self.dispatch(
172 enc,
173 &self.pipes.moe_combine_b,
174 |e| {
175 self.bind(e, 0, &s.y_e, 0);
176 self.bind(e, 1, &self.wmap, 0);
177 self.bind(e, 2, &s.moe_out, 0);
178 set_bytes(e, 3, &mp);
179 self.bind(e, 4, &self.dense, 0);
180 self.bind(e, 5, &s.hc.mixed, 0);
181 self.bind(e, 6, &self.slot_tab, 0);
182 },
183 nb * (h as usize).div_ceil(256),
184 256,
185 true,
186 );
187 }
188}