1#![allow(non_snake_case)]
26
27use teeny_core::dtype::Float;
28use teeny_macros::kernel;
29use teeny_triton::triton::{
30 types::{AddOffsets, Comparison},
31 *,
32};
33
34#[kernel]
40pub fn rms_norm_forward<T: Triton, D: Float, const BLOCK_N: i32>(
41 x_ptr: T::Pointer<D>,
42 y_ptr: T::Pointer<D>,
43 weight_ptr: T::Pointer<D>,
44 rrms_ptr: T::Pointer<D>,
45 _M: i32,
46 N: i32,
47 eps: f32,
48) where
49 T::I32Tensor: types::Tensor<i32, 1>,
50 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
51 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
52{
53 let row = T::program_id(Axis::X);
54 let row_start = row * N;
55 let row_idx = T::arange(0, 1) + row;
56
57 let zeros = T::zeros::<D>(&[BLOCK_N]);
58 let zero_1 = T::zeros::<D>(&[1]);
59 let n_inv = T::cast::<f32, D>(T::full::<f32>(&[1], 1.0f32 / (N as f32)), None, false);
60
61 let mut sq_sum = zero_1;
63 let mut n_start: i32 = 0;
64 while n_start < N {
65 let col_offs = T::arange(0, BLOCK_N) + n_start;
66 let mask = col_offs.lt(N);
67 let x_tile = T::load(
68 x_ptr.add_offsets(col_offs + row_start),
69 Some(mask),
70 Some(zeros),
71 &[],
72 None,
73 None,
74 None,
75 false,
76 );
77 sq_sum = sq_sum + T::sum(x_tile * x_tile, None, true);
78 n_start += BLOCK_N;
79 }
80 let eps_t = T::cast::<f32, D>(T::full::<f32>(&[1], eps), None, false);
81 let rrms_1 = T::rsqrt(sq_sum * n_inv + eps_t);
82 let rrms = T::broadcast_to(rrms_1, &[BLOCK_N]);
83
84 T::store(rrms_ptr.add_offsets(row_idx), rrms_1, None, &[], None, None);
85
86 n_start = 0;
88 while n_start < N {
89 let col_offs = T::arange(0, BLOCK_N) + n_start;
90 let mask = col_offs.lt(N);
91 let x_tile = T::load(
92 x_ptr.add_offsets(col_offs + row_start),
93 Some(mask),
94 Some(zeros),
95 &[],
96 None,
97 None,
98 None,
99 false,
100 );
101 let gamma = T::load(
102 weight_ptr.add_offsets(col_offs),
103 Some(mask),
104 Some(zeros),
105 &[],
106 None,
107 None,
108 None,
109 false,
110 );
111 let y_tile = x_tile * rrms * gamma;
112 T::store(
113 y_ptr.add_offsets(col_offs + row_start),
114 y_tile,
115 Some(mask),
116 &[],
117 None,
118 None,
119 );
120 n_start += BLOCK_N;
121 }
122}
123
124#[cfg(feature = "training")]
135#[kernel]
136pub fn rms_norm_backward<T: Triton, D: Float, const BLOCK_N: i32>(
137 dy_ptr: T::Pointer<D>,
138 x_ptr: T::Pointer<D>,
139 dx_ptr: T::Pointer<D>,
140 weight_ptr: T::Pointer<D>,
141 dweight_ptr: T::Pointer<D>,
142 rrms_ptr: T::Pointer<D>,
143 _M: i32,
144 N: i32,
145) where
146 T::I32Tensor: types::Tensor<i32, 1>,
147 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
148 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
149{
150 let row = T::program_id(Axis::X);
151 let row_start = row * N;
152 let row_idx = T::arange(0, 1) + row;
153
154 let zeros = T::zeros::<D>(&[BLOCK_N]);
155 let zero_1 = T::zeros::<D>(&[1]);
156 let n_inv = T::cast::<f32, D>(T::full::<f32>(&[1], 1.0f32 / (N as f32)), None, false);
157
158 let rrms_1 = T::load(
159 rrms_ptr.add_offsets(row_idx),
160 None,
161 None,
162 &[],
163 None,
164 None,
165 None,
166 false,
167 );
168 let rrms = T::broadcast_to(rrms_1, &[BLOCK_N]);
169
170 let mut dot = zero_1;
172 let mut n_start: i32 = 0;
173 while n_start < N {
174 let col_offs = T::arange(0, BLOCK_N) + n_start;
175 let mask = col_offs.lt(N);
176 let x_tile = T::load(
177 x_ptr.add_offsets(col_offs + row_start),
178 Some(mask),
179 Some(zeros),
180 &[],
181 None,
182 None,
183 None,
184 false,
185 );
186 let dy_tile = T::load(
187 dy_ptr.add_offsets(col_offs + row_start),
188 Some(mask),
189 Some(zeros),
190 &[],
191 None,
192 None,
193 None,
194 false,
195 );
196 let gamma = T::load(
197 weight_ptr.add_offsets(col_offs),
198 Some(mask),
199 Some(zeros),
200 &[],
201 None,
202 None,
203 None,
204 false,
205 );
206 dot = dot + T::sum(dy_tile * gamma * x_tile, None, true);
207 n_start += BLOCK_N;
208 }
209 let rrms_sq = T::broadcast_to(rrms_1 * rrms_1, &[BLOCK_N]);
210 let scale = T::broadcast_to(dot * n_inv, &[BLOCK_N]);
211
212 n_start = 0;
214 while n_start < N {
215 let col_offs = T::arange(0, BLOCK_N) + n_start;
216 let mask = col_offs.lt(N);
217 let x_tile = T::load(
218 x_ptr.add_offsets(col_offs + row_start),
219 Some(mask),
220 Some(zeros),
221 &[],
222 None,
223 None,
224 None,
225 false,
226 );
227 let dy_tile = T::load(
228 dy_ptr.add_offsets(col_offs + row_start),
229 Some(mask),
230 Some(zeros),
231 &[],
232 None,
233 None,
234 None,
235 false,
236 );
237 let gamma = T::load(
238 weight_ptr.add_offsets(col_offs),
239 Some(mask),
240 Some(zeros),
241 &[],
242 None,
243 None,
244 None,
245 false,
246 );
247 let dw_old = T::load(
248 dweight_ptr.add_offsets(col_offs),
249 Some(mask),
250 Some(zeros),
251 &[],
252 None,
253 None,
254 None,
255 false,
256 );
257
258 let dx_tile = rrms * gamma * (dy_tile - x_tile * rrms_sq * scale);
259 T::store(
260 dx_ptr.add_offsets(col_offs + row_start),
261 dx_tile,
262 Some(mask),
263 &[],
264 None,
265 None,
266 );
267 T::store(
268 dweight_ptr.add_offsets(col_offs),
269 dw_old + dy_tile * x_tile * rrms,
270 Some(mask),
271 &[],
272 None,
273 None,
274 );
275 n_start += BLOCK_N;
276 }
277}