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