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