Skip to main content

teeny_kernels/nn/optim/
rmsprop.rs

1/*
2 * Copyright (c) 2026 Teenygrad.
3 *
4 * Licensed under the Apache License, Version 2.0 (the "License");
5 * you may not use this file except in compliance with the License.
6 * You may obtain a copy of the License at
7 *
8 *   http://www.apache.org/licenses/LICENSE-2.0
9 *
10 * Unless required by applicable law or agreed to in writing, software
11 * distributed under the License is distributed on an "AS IS" BASIS,
12 * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13 * See the License for the specific language governing permissions and
14 * limitations under the License.
15 */
16
17#![allow(non_snake_case)]
18
19use teeny_macros::kernel;
20use teeny_triton::triton::{
21    types::{AddOffsets, Comparison},
22    *,
23};
24
25/// RMSprop step (no momentum, not centred).
26///
27/// ```text
28/// square_avg = alpha * square_avg + (1 - alpha) * g²
29/// p          = p - lr * g / (sqrt(square_avg) + eps)
30/// ```
31///
32/// Grid: `[ceil(n_elements / BLOCK_SIZE), 1, 1]`.
33#[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/// RMSprop step with momentum.
113///
114/// ```text
115/// square_avg = alpha * square_avg + (1 - alpha) * g²
116/// buf        = momentum * buf + g / sqrt(square_avg + eps)
117/// p          = p - lr * buf
118/// ```
119///
120/// Grid: `[ceil(n_elements / BLOCK_SIZE), 1, 1]`.
121#[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}