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