Skip to main content

cherenkov/qwen4_exp/gpu/prefill/
hyperconnection.rs

1//! Prefill grouped residual reads and injection.
2
3use super::*;
4
5impl Gpu<'_> {
6    pub(super) fn pf_inject(&self, enc: &Enc, hyper: &Buf, out: &Buf, inj: &Buf, t: usize) {
7        let gp = self.group_params(false, 0.0);
8        let nbu = t as u32;
9
10        self.dispatch(
11            enc,
12            &self.pipes.inject_b,
13            |e| {
14                self.bind(e, 0, hyper, 0);
15                self.bind(e, 1, out, 0);
16                self.bind(e, 2, inj, 0);
17                set_bytes(e, 3, &gp);
18                set_bytes(e, 4, &nbu);
19            },
20            t * self.p.cfg.hc_hidden(),
21            256,
22            false,
23        );
24    }
25
26    /// Gated residual read over t rows: pf.mixed (block input), pf.inj.
27    #[allow(clippy::too_many_arguments)]
28    pub(super) fn pf_hc_read(
29        &self,
30        enc: &Enc,
31        hc: &Hc,
32        t: usize,
33        pf: &PrefillScratch,
34        hyper: &Buf,
35        pending: Option<&Buf>,
36        inject: bool,
37    ) {
38        let c = &self.p.cfg;
39        let h = c.hidden_size as u32;
40
41        if let Some(out) = pending {
42            self.pf_inject(enc, hyper, out, &pf.inj, t);
43        }
44
45        self.group_norm_b(
46            enc,
47            hyper,
48            0,
49            hc.norm,
50            &pf.normed,
51            h,
52            c.hc_count as u32,
53            0.0,
54            t,
55        );
56        self.qmm(enc, &hc.down, &pf.normed, &pf.d, t);
57
58        let n = (t as u32) * hc.down.out;
59        let div = c.hc_count as f32;
60
61        self.dispatch(
62            enc,
63            &self.pipes.silu_rows,
64            |e| {
65                self.bind(e, 0, &pf.d, 0);
66                set_bytes(e, 1, &n);
67                set_bytes(e, 2, &div);
68            },
69            n as usize,
70            256,
71            false,
72        );
73        self.qmm(enc, &hc.up, &pf.d, &pf.u, t);
74
75        if inject && let Some(q) = &hc.inject {
76            self.qmm(enc, q, &pf.normed, &pf.inj, t);
77        }
78
79        let gp = self.group_params(false, 0.0);
80        let nbu = t as u32;
81
82        self.dispatch(
83            enc,
84            &self.pipes.hc_mix_b,
85            |e| {
86                self.bind(e, 0, &pf.u, 0);
87                self.bind(e, 1, &pf.normed, 0);
88                self.bind(e, 2, &pf.mixed, 0);
89                set_bytes(e, 3, &gp);
90                set_bytes(e, 4, &nbu);
91            },
92            t * h as usize,
93            256,
94            false,
95        );
96    }
97}