cherenkov/qwen4_exp/gpu/prefill/
hyperconnection.rs1use 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 #[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}