Skip to main content

cherenkov/qwen4_exp/gpu/prefill/
ple.rs

1//! Prefill PLE n-gram gather, gating, and dilated convolution.
2
3use super::*;
4
5impl Gpu<'_> {
6    /// PLE block over t rows. Nothing else is live between blocks, so its
7    /// buffers alias the stream-wide scratch: e -> mix_out, key -> normed
8    /// (and the output, once the key is normed), keyn -> u, value -> mixed,
9    /// query -> qkv, gated -> fh, gvn -> mtp_hyper.
10    pub(super) fn pf_ple(
11        &self,
12        enc: &Enc,
13        pl: &Ple,
14        base: usize,
15        t: usize,
16        pf: &PrefillScratch,
17    ) -> Result<()> {
18        let c = &self.p.cfg;
19        let h = c.hidden_size as u32;
20        let hh = c.hc_hidden();
21
22        anyhow::ensure!(
23            c.ple_embed_dim <= c.hidden_size,
24            "PLE embedding wider than the hidden size"
25        );
26
27        let (ple_e, ple_key, ple_keyn, ple_value, ple_query, ple_gated, ple_gvn, ple_out) = (
28            &pf.mix_out,
29            &pf.normed,
30            &pf.u,
31            &pf.mixed,
32            &pf.qkv,
33            &pf.fh,
34            &pf.mtp_hyper,
35            &pf.normed,
36        );
37        let dim = self.p.manifest.ngram.dim;
38        let e = unsafe {
39            std::slice::from_raw_parts_mut(
40                ple_e.contents().cast::<f32>().as_ptr(),
41                pf.rows * c.ple_embed_dim,
42            )
43        };
44        let t_gather = std::time::Instant::now();
45
46        self.ngram_prefetch_join();
47
48        for b in 0..t {
49            let ids = self.ngram_ids_at(pl, base + b);
50
51            anyhow::ensure!(
52                ids.len() * dim == c.ple_embed_dim,
53                "n-gram head layout mismatch"
54            );
55
56            let eb = &mut e[b * c.ple_embed_dim..(b + 1) * c.ple_embed_dim];
57
58            for (hi, &id) in ids.iter().enumerate() {
59                self.p.ngram_row(id, &mut eb[hi * dim..(hi + 1) * dim]);
60            }
61        }
62
63        self.ngram_gather_s
64            .set(self.ngram_gather_s.get() + t_gather.elapsed().as_secs_f64());
65        self.qmm(enc, &pl.key, ple_e, ple_key, t);
66        self.qmm(enc, &pl.value, ple_e, ple_value, t);
67
68        let groups = c.hc_count as u32;
69
70        self.group_norm_b(enc, ple_key, 0, pl.norm_key, ple_keyn, h, groups, 0.0, t);
71        self.group_norm_b(
72            enc,
73            &pf.hyper,
74            0,
75            pl.norm_query,
76            ple_query,
77            h,
78            groups,
79            0.0,
80            t,
81        );
82
83        let gp = self.group_params(true, 0.0);
84
85        self.dispatch(
86            enc,
87            &self.pipes.ple_gate_b,
88            |e| {
89                self.bind(e, 0, ple_keyn, 0);
90                self.bind(e, 1, ple_query, 0);
91                self.bind(e, 2, ple_value, 0);
92                self.bind(e, 3, ple_gated, 0);
93                set_bytes(e, 4, &gp);
94            },
95            t * c.hc_count,
96            256,
97            true,
98        );
99        self.group_norm_b(enc, ple_gated, 0, pl.norm_conv, ple_gvn, h, groups, 0.0, t);
100
101        let cp = PleConvParams {
102            channels: hh as u32,
103            ksize: pl.kernel,
104            dilation: pl.dilation,
105            span: pl.span,
106            filled: base as u32,
107            nb: t as u32,
108        };
109
110        self.dispatch(
111            enc,
112            &self.pipes.ple_conv_b,
113            |e| {
114                self.bind(e, 0, ple_gated, 0);
115                self.bind(e, 1, ple_gvn, 0);
116                self.bind(e, 2, &self.dense, pl.conv.0);
117                self.bind(e, 3, &pl.hist, 0);
118                self.bind(e, 4, ple_out, 0);
119                set_bytes(e, 5, &cp);
120            },
121            hh,
122            256,
123            false,
124        );
125        self.add(enc, &pf.hyper, ple_out, t * hh);
126
127        Ok(())
128    }
129}