Skip to main content

cherenkov/qwen4_exp/gpu/
deltanet.rs

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