Skip to main content

teeny_kernels/nn/optim/
radam.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/// RAdam step — adaptive variance phase (rho_t > 5).
26///
27/// ```text
28/// exp_avg    = beta1 * exp_avg    + (1 - beta1) * g
29/// exp_avg_sq = beta2 * exp_avg_sq + (1 - beta2) * g²
30/// denom      = sqrt(exp_avg_sq / bias_correction2) + eps
31/// p         -= step_size * exp_avg / denom
32/// ```
33///
34/// `step_size = lr * rect / bias_correction1` is precomputed on host, where
35/// `rect = sqrt((rho_t-4)(rho_t-2)*rho_inf / ((rho_inf-4)(rho_inf-2)*rho_t))`.
36///
37/// Grid: `[ceil(n_elements / BLOCK_SIZE), 1, 1]`.
38#[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/// RAdam step — SGD fallback phase (rho_t <= 5).
146///
147/// ```text
148/// exp_avg    = beta1 * exp_avg    + (1 - beta1) * g
149/// exp_avg_sq = beta2 * exp_avg_sq + (1 - beta2) * g²   (state updated but unused)
150/// p         -= step_size * exp_avg
151/// ```
152///
153/// `step_size = lr / bias_correction1` is precomputed on host.
154///
155/// Grid: `[ceil(n_elements / BLOCK_SIZE), 1, 1]`.
156#[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}