Skip to main content

teeny_kernels/nn/optim/
adam.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/// Adam step.
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_corr2_sqrt + eps
31/// p          = p - step_size * exp_avg / denom
32/// ```
33///
34/// Scalars precomputed on host:
35/// - `step_size = lr / bias_correction1`   where `bias_correction1 = 1 - beta1^t`
36/// - `bias_corr2_sqrt = sqrt(1 - beta2^t)`
37///
38/// Grid: `[ceil(n_elements / BLOCK_SIZE), 1, 1]`.
39#[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/// AdamW step (decoupled weight decay).
147///
148/// ```text
149/// p          = p * (1 - lr * weight_decay)
150/// exp_avg    = beta1 * exp_avg    + (1 - beta1) * g
151/// exp_avg_sq = beta2 * exp_avg_sq + (1 - beta2) * g²
152/// denom      = sqrt(exp_avg_sq) / bias_corr2_sqrt + eps
153/// p          = p - step_size * exp_avg / denom
154/// ```
155///
156/// Grid: `[ceil(n_elements / BLOCK_SIZE), 1, 1]`.
157#[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    // Decoupled weight decay applied directly to params
232    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}