1#![allow(non_snake_case)]
18
19use teeny_macros::kernel;
20use teeny_triton::triton::{
21 types::{AddOffsets, Comparison},
22 *,
23};
24
25#[kernel]
40pub fn adam_step<T: Triton, const BLOCK_SIZE: i32>(
41 params_ptr: T::Pointer<f32>,
42 grad_ptr: T::Pointer<f32>,
43 exp_avg_ptr: T::Pointer<f32>,
44 exp_avg_sq_ptr: T::Pointer<f32>,
45 n_elements: i32,
46 step_size: f32,
47 bias_corr2_sqrt: f32,
48 beta1: f32,
49 beta2: f32,
50 eps: f32,
51 weight_decay: f32,
52) where
53 T::I32Tensor: types::Tensor<i32, 1>,
54 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
55 T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
56{
57 let pid = T::program_id(Axis::X);
58 let block_start = pid * BLOCK_SIZE;
59 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
60 let mask = offsets.lt(n_elements);
61
62 let p = T::load(
63 params_ptr.add_offsets(offsets),
64 Some(mask),
65 None,
66 &[],
67 None,
68 None,
69 None,
70 false,
71 );
72 let g = T::load(
73 grad_ptr.add_offsets(offsets),
74 Some(mask),
75 None,
76 &[],
77 None,
78 None,
79 None,
80 false,
81 );
82 let exp_avg = T::load(
83 exp_avg_ptr.add_offsets(offsets),
84 Some(mask),
85 None,
86 &[],
87 None,
88 None,
89 None,
90 false,
91 );
92 let exp_avg_sq = T::load(
93 exp_avg_sq_ptr.add_offsets(offsets),
94 Some(mask),
95 None,
96 &[],
97 None,
98 None,
99 None,
100 false,
101 );
102
103 let beta1_t = T::full(&[BLOCK_SIZE], beta1);
104 let beta2_t = T::full(&[BLOCK_SIZE], beta2);
105 let one_m_beta1 = T::full(&[BLOCK_SIZE], 1.0_f32 - beta1);
106 let one_m_beta2 = T::full(&[BLOCK_SIZE], 1.0_f32 - beta2);
107 let eps_t = T::full(&[BLOCK_SIZE], eps);
108 let step_size_t = T::full(&[BLOCK_SIZE], step_size);
109 let bc2sqrt_t = T::full(&[BLOCK_SIZE], bias_corr2_sqrt);
110 let wd_t = T::full(&[BLOCK_SIZE], weight_decay);
111
112 let g_eff = g + wd_t * p;
113
114 let exp_avg_new = beta1_t * exp_avg + one_m_beta1 * g_eff;
115 let exp_avg_sq_new = beta2_t * exp_avg_sq + one_m_beta2 * g_eff * g_eff;
116
117 let denom = T::sqrt_rn(exp_avg_sq_new) / bc2sqrt_t + eps_t;
118 let p_new = p - step_size_t * exp_avg_new / denom;
119
120 T::store(
121 params_ptr.add_offsets(offsets),
122 p_new,
123 Some(mask),
124 &[],
125 None,
126 None,
127 );
128 T::store(
129 exp_avg_ptr.add_offsets(offsets),
130 exp_avg_new,
131 Some(mask),
132 &[],
133 None,
134 None,
135 );
136 T::store(
137 exp_avg_sq_ptr.add_offsets(offsets),
138 exp_avg_sq_new,
139 Some(mask),
140 &[],
141 None,
142 None,
143 );
144}
145
146#[kernel]
158pub fn adamw_step<T: Triton, const BLOCK_SIZE: i32>(
159 params_ptr: T::Pointer<f32>,
160 grad_ptr: T::Pointer<f32>,
161 exp_avg_ptr: T::Pointer<f32>,
162 exp_avg_sq_ptr: T::Pointer<f32>,
163 n_elements: i32,
164 step_size: f32,
165 bias_corr2_sqrt: f32,
166 beta1: f32,
167 beta2: f32,
168 eps: f32,
169 weight_decay: f32,
170 lr: f32,
171) where
172 T::I32Tensor: types::Tensor<i32, 1>,
173 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
174 T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
175{
176 let pid = T::program_id(Axis::X);
177 let block_start = pid * BLOCK_SIZE;
178 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
179 let mask = offsets.lt(n_elements);
180
181 let p = T::load(
182 params_ptr.add_offsets(offsets),
183 Some(mask),
184 None,
185 &[],
186 None,
187 None,
188 None,
189 false,
190 );
191 let g = T::load(
192 grad_ptr.add_offsets(offsets),
193 Some(mask),
194 None,
195 &[],
196 None,
197 None,
198 None,
199 false,
200 );
201 let exp_avg = T::load(
202 exp_avg_ptr.add_offsets(offsets),
203 Some(mask),
204 None,
205 &[],
206 None,
207 None,
208 None,
209 false,
210 );
211 let exp_avg_sq = T::load(
212 exp_avg_sq_ptr.add_offsets(offsets),
213 Some(mask),
214 None,
215 &[],
216 None,
217 None,
218 None,
219 false,
220 );
221
222 let beta1_t = T::full(&[BLOCK_SIZE], beta1);
223 let beta2_t = T::full(&[BLOCK_SIZE], beta2);
224 let one_m_beta1 = T::full(&[BLOCK_SIZE], 1.0_f32 - beta1);
225 let one_m_beta2 = T::full(&[BLOCK_SIZE], 1.0_f32 - beta2);
226 let eps_t = T::full(&[BLOCK_SIZE], eps);
227 let step_size_t = T::full(&[BLOCK_SIZE], step_size);
228 let bc2sqrt_t = T::full(&[BLOCK_SIZE], bias_corr2_sqrt);
229 let wd_decay = T::full(&[BLOCK_SIZE], 1.0_f32 - lr * weight_decay);
230
231 let p_decayed = p * wd_decay;
233
234 let exp_avg_new = beta1_t * exp_avg + one_m_beta1 * g;
235 let exp_avg_sq_new = beta2_t * exp_avg_sq + one_m_beta2 * g * g;
236
237 let denom = T::sqrt_rn(exp_avg_sq_new) / bc2sqrt_t + eps_t;
238 let p_new = p_decayed - step_size_t * exp_avg_new / denom;
239
240 T::store(
241 params_ptr.add_offsets(offsets),
242 p_new,
243 Some(mask),
244 &[],
245 None,
246 None,
247 );
248 T::store(
249 exp_avg_ptr.add_offsets(offsets),
250 exp_avg_new,
251 Some(mask),
252 &[],
253 None,
254 None,
255 );
256 T::store(
257 exp_avg_sq_ptr.add_offsets(offsets),
258 exp_avg_sq_new,
259 Some(mask),
260 &[],
261 None,
262 None,
263 );
264}