cherenkov/qwen4_exp/gpu/prefill/
ple.rs1use super::*;
4
5impl Gpu<'_> {
6 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}