Skip to main content

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