1#![allow(non_snake_case)]
18
19use teeny_macros::kernel;
20use teeny_triton::triton::{
21 types::{AddOffsets, Comparison},
22 *,
23};
24
25#[kernel]
34pub fn adagrad_step<T: Triton, const BLOCK_SIZE: i32>(
35 params_ptr: T::Pointer<f32>,
36 grad_ptr: T::Pointer<f32>,
37 sum_ptr: T::Pointer<f32>,
38 n_elements: i32,
39 lr: f32,
40 eps: f32,
41 weight_decay: f32,
42) where
43 T::I32Tensor: types::Tensor<i32, 1>,
44 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
45 T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
46{
47 let pid = T::program_id(Axis::X);
48 let block_start = pid * BLOCK_SIZE;
49 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
50 let mask = offsets.lt(n_elements);
51
52 let p = T::load(
53 params_ptr.add_offsets(offsets),
54 Some(mask),
55 None,
56 &[],
57 None,
58 None,
59 None,
60 false,
61 );
62 let g = T::load(
63 grad_ptr.add_offsets(offsets),
64 Some(mask),
65 None,
66 &[],
67 None,
68 None,
69 None,
70 false,
71 );
72 let sum = T::load(
73 sum_ptr.add_offsets(offsets),
74 Some(mask),
75 None,
76 &[],
77 None,
78 None,
79 None,
80 false,
81 );
82
83 let lr_t = T::full(&[BLOCK_SIZE], lr);
84 let eps_t = T::full(&[BLOCK_SIZE], eps);
85 let wd_t = T::full(&[BLOCK_SIZE], weight_decay);
86
87 let g_eff = g + wd_t * p;
88 let sum_new = sum + g_eff * g_eff;
89 let p_new = p - lr_t * g_eff / (T::sqrt_rn(sum_new) + eps_t);
90
91 T::store(
92 params_ptr.add_offsets(offsets),
93 p_new,
94 Some(mask),
95 &[],
96 None,
97 None,
98 );
99 T::store(
100 sum_ptr.add_offsets(offsets),
101 sum_new,
102 Some(mask),
103 &[],
104 None,
105 None,
106 );
107}
108
109#[kernel]
120pub fn adadelta_step<T: Triton, const BLOCK_SIZE: i32>(
121 params_ptr: T::Pointer<f32>,
122 grad_ptr: T::Pointer<f32>,
123 square_avg_ptr: T::Pointer<f32>,
124 acc_delta_ptr: T::Pointer<f32>,
125 n_elements: i32,
126 lr: f32,
127 rho: f32,
128 eps: f32,
129 weight_decay: f32,
130) where
131 T::I32Tensor: types::Tensor<i32, 1>,
132 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
133 T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
134{
135 let pid = T::program_id(Axis::X);
136 let block_start = pid * BLOCK_SIZE;
137 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
138 let mask = offsets.lt(n_elements);
139
140 let p = T::load(
141 params_ptr.add_offsets(offsets),
142 Some(mask),
143 None,
144 &[],
145 None,
146 None,
147 None,
148 false,
149 );
150 let g = T::load(
151 grad_ptr.add_offsets(offsets),
152 Some(mask),
153 None,
154 &[],
155 None,
156 None,
157 None,
158 false,
159 );
160 let square_avg = T::load(
161 square_avg_ptr.add_offsets(offsets),
162 Some(mask),
163 None,
164 &[],
165 None,
166 None,
167 None,
168 false,
169 );
170 let acc_delta = T::load(
171 acc_delta_ptr.add_offsets(offsets),
172 Some(mask),
173 None,
174 &[],
175 None,
176 None,
177 None,
178 false,
179 );
180
181 let lr_t = T::full(&[BLOCK_SIZE], lr);
182 let rho_t = T::full(&[BLOCK_SIZE], rho);
183 let one_m_rho = T::full(&[BLOCK_SIZE], 1.0_f32 - rho);
184 let eps_t = T::full(&[BLOCK_SIZE], eps);
185 let wd_t = T::full(&[BLOCK_SIZE], weight_decay);
186
187 let g_eff = g + wd_t * p;
188 let square_avg_new = rho_t * square_avg + one_m_rho * g_eff * g_eff;
189 let std = T::sqrt_rn(square_avg_new + eps_t);
190 let delta = T::sqrt_rn(acc_delta + eps_t) / std * g_eff;
191 let p_new = p - lr_t * delta;
192 let acc_delta_new = rho_t * acc_delta + one_m_rho * delta * delta;
193
194 T::store(
195 params_ptr.add_offsets(offsets),
196 p_new,
197 Some(mask),
198 &[],
199 None,
200 None,
201 );
202 T::store(
203 square_avg_ptr.add_offsets(offsets),
204 square_avg_new,
205 Some(mask),
206 &[],
207 None,
208 None,
209 );
210 T::store(
211 acc_delta_ptr.add_offsets(offsets),
212 acc_delta_new,
213 Some(mask),
214 &[],
215 None,
216 None,
217 );
218}