1use super::*;
4
5impl Gpu<'_> {
6 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}