Skip to main content

cherenkov/qwen4_exp/gpu/prefill/
deltanet.rs

1//! Prefill DeltaNet convolution, scan, and output gate.
2
3use super::*;
4
5impl Gpu<'_> {
6    /// Gated DeltaNet over t rows (sequential scan). pf.mixed -> pf.mix_out.
7    pub(super) fn pf_deltanet(&self, enc: &Enc, d: &Delta, t: usize, pf: &PrefillScratch) {
8        let c = &self.p.cfg;
9        let conv_dim = d.qkv.out;
10        let nbu = t as u32;
11
12        self.qmm(enc, &d.qkv, &pf.mixed, &pf.qkv, t);
13        self.qmm(enc, &d.z, &pf.mixed, &pf.z, t);
14        self.qmm(enc, &d.a, &pf.mixed, &pf.a, t);
15        self.qmm(enc, &d.b, &pf.mixed, &pf.b, t);
16
17        let cp = ConvParams {
18            channels: conv_dim,
19            ksize: c.linear_conv_kernel_dim as u32,
20        };
21        let snap = 0u32;
22
23        self.dispatch(
24            enc,
25            &self.pipes.conv_b,
26            |e| {
27                self.bind(e, 0, &pf.qkv, 0);
28                self.bind(e, 1, &self.dense, d.conv.0);
29                self.bind(e, 2, &d.hist, 0);
30                set_bytes(e, 3, &cp);
31                set_bytes(e, 4, &nbu);
32                set_bytes(e, 5, &snap);
33                self.bind(e, 6, &d.mid_hist, 0);
34            },
35            conv_dim as usize,
36            256,
37            false,
38        );
39
40        let p = DeltaPrepParams {
41            n_k: c.linear_num_key_heads as u32,
42            n_v: c.linear_num_value_heads as u32,
43            d_k: c.linear_key_head_dim as u32,
44            d_v: c.linear_value_head_dim as u32,
45            eps: 1e-6,
46            nb: nbu,
47            snap_after: 0,
48        };
49
50        self.dispatch(
51            enc,
52            &self.pipes.delta_norms,
53            |e| {
54                self.bind(e, 0, &pf.qkv, 0);
55                self.bind(e, 1, &pf.kqn, 0);
56                set_bytes(e, 2, &p);
57            },
58            (t * c.linear_num_key_heads).div_ceil(4),
59            128,
60            true,
61        );
62        self.dispatch(
63            enc,
64            &self.pipes.delta_gates,
65            |e| {
66                self.bind(e, 0, &pf.a, 0);
67                self.bind(e, 1, &pf.b, 0);
68                self.bind(e, 2, &self.dense, d.a_log.0);
69                self.bind(e, 3, &self.dense, d.dt_bias.0);
70                self.bind(e, 4, &pf.gbuf, 0);
71                set_bytes(e, 5, &p);
72            },
73            t * c.linear_num_value_heads,
74            96,
75            false,
76        );
77
78        let rows = c.linear_num_value_heads * c.linear_value_head_dim;
79
80        self.dispatch(
81            enc,
82            &self.pipes.delta_scan2,
83            |e| {
84                self.bind(e, 0, &pf.qkv, 0);
85                self.bind(e, 1, &pf.kqn, 0);
86                self.bind(e, 2, &pf.gbuf, 0);
87                self.bind(e, 3, &d.state, 0);
88                self.bind(e, 4, &pf.delta_y, 0);
89                self.bind(e, 5, &d.mid, 0);
90                set_bytes(e, 6, &p);
91            },
92            rows.div_ceil(4),
93            128,
94            true,
95        );
96        self.dispatch(
97            enc,
98            &self.pipes.gate_norm_sigmoid_b,
99            |e| {
100                self.bind(e, 0, &pf.delta_y, 0);
101                self.bind(e, 1, &pf.z, 0);
102                self.bind(e, 2, &self.dense, d.norm.0);
103                set_bytes(e, 3, &p);
104            },
105            t * c.linear_num_value_heads,
106            c.linear_value_head_dim,
107            true,
108        );
109        self.qmm(enc, &d.o, &pf.delta_y, &pf.mix_out, t);
110    }
111}