Skip to main content

cherenkov/qwen4_exp/gpu/
experts.rs

1//! Router and resident/fetched expert compute dispatch.
2
3use super::*;
4
5impl Gpu<'_> {
6    /// Router over `nb` rows of `x`: softmax over experts, top-k.
7    #[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    /// Routed experts of `nb` rows (the union published in slot table row
61    /// `slot_row` at execution time) plus the shared expert, summed per
62    /// row into scratch.moe_out. Part 0 covers the records resident when
63    /// the table was published (plus the shared expert); part 1 the ones
64    /// fetched meanwhile. Input: scratch.hc.mixed (prepped in h1).
65    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                // Two output rows per simdgroup.
166            },
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}