1#![allow(non_snake_case)]
18
19use teeny_core::dtype::Float;
20use teeny_macros::kernel;
21use teeny_triton::triton::{
22 types::{AddOffsets, Comparison},
23 *,
24};
25
26#[kernel(backward = LeakyReluBackward)]
30pub fn leaky_relu_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
31 x_ptr: T::Pointer<D>,
32 y_ptr: T::Pointer<D>,
33 n_elements: i32,
34 negative_slope: f32,
35) where
36 T::I32Tensor: types::Tensor<i32, 1>,
37 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
38 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
39{
40 let pid = T::program_id(Axis::X);
41 let block_start = pid * BLOCK_SIZE;
42 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
43 let in_bounds = offsets.lt(n_elements);
44
45 let x = T::load(
46 x_ptr.add_offsets(offsets),
47 Some(in_bounds),
48 None,
49 &[],
50 None,
51 None,
52 None,
53 false,
54 );
55 let slope = T::full(&[BLOCK_SIZE], D::from_f64(negative_slope as f64));
56 let x_pos = T::gt(x, T::zeros_like(x));
57 let y = T::where_(x_pos, x, slope * x);
58 T::store(
59 y_ptr.add_offsets(offsets),
60 y,
61 Some(in_bounds),
62 &[],
63 None,
64 None,
65 );
66}
67
68#[kernel]
70pub fn leaky_relu_backward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
71 dy_ptr: T::Pointer<D>,
72 x_ptr: T::Pointer<D>,
73 dx_ptr: T::Pointer<D>,
74 n_elements: i32,
75 negative_slope: f32,
76) where
77 T::I32Tensor: types::Tensor<i32, 1>,
78 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
79 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
80{
81 let pid = T::program_id(Axis::X);
82 let block_start = pid * BLOCK_SIZE;
83 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
84 let in_bounds = offsets.lt(n_elements);
85
86 let dy = T::load(
87 dy_ptr.add_offsets(offsets),
88 Some(in_bounds),
89 None,
90 &[],
91 None,
92 None,
93 None,
94 false,
95 );
96 let x = T::load(
97 x_ptr.add_offsets(offsets),
98 Some(in_bounds),
99 None,
100 &[],
101 None,
102 None,
103 None,
104 false,
105 );
106 let slope = T::full(&[BLOCK_SIZE], D::from_f64(negative_slope as f64));
107 let x_pos = T::gt(x, T::zeros_like(x));
108 let dx = T::where_(x_pos, dy, slope * dy);
109 T::store(
110 dx_ptr.add_offsets(offsets),
111 dx,
112 Some(in_bounds),
113 &[],
114 None,
115 None,
116 );
117}
118
119#[kernel(backward = ThresholdBackward)]
123pub fn threshold_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
124 x_ptr: T::Pointer<D>,
125 y_ptr: T::Pointer<D>,
126 n_elements: i32,
127 threshold: f32,
128 value: f32,
129) where
130 T::I32Tensor: types::Tensor<i32, 1>,
131 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
132 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
133{
134 let pid = T::program_id(Axis::X);
135 let block_start = pid * BLOCK_SIZE;
136 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
137 let in_bounds = offsets.lt(n_elements);
138
139 let x = T::load(
140 x_ptr.add_offsets(offsets),
141 Some(in_bounds),
142 None,
143 &[],
144 None,
145 None,
146 None,
147 false,
148 );
149 let thr = T::full(&[BLOCK_SIZE], D::from_f64(threshold as f64));
150 let val = T::full(&[BLOCK_SIZE], D::from_f64(value as f64));
151 let above = T::gt(x, thr);
152 let y = T::where_(above, x, val);
153 T::store(
154 y_ptr.add_offsets(offsets),
155 y,
156 Some(in_bounds),
157 &[],
158 None,
159 None,
160 );
161}
162
163#[kernel]
165pub fn threshold_backward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
166 dy_ptr: T::Pointer<D>,
167 x_ptr: T::Pointer<D>,
168 dx_ptr: T::Pointer<D>,
169 n_elements: i32,
170 threshold: f32,
171) where
172 T::I32Tensor: types::Tensor<i32, 1>,
173 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
174 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
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 in_bounds = offsets.lt(n_elements);
180
181 let dy = T::load(
182 dy_ptr.add_offsets(offsets),
183 Some(in_bounds),
184 None,
185 &[],
186 None,
187 None,
188 None,
189 false,
190 );
191 let x = T::load(
192 x_ptr.add_offsets(offsets),
193 Some(in_bounds),
194 None,
195 &[],
196 None,
197 None,
198 None,
199 false,
200 );
201 let thr = T::full(&[BLOCK_SIZE], D::from_f64(threshold as f64));
202 let above = T::gt(x, thr);
203 let dx = T::where_(above, dy, T::zeros_like(dy));
204 T::store(
205 dx_ptr.add_offsets(offsets),
206 dx,
207 Some(in_bounds),
208 &[],
209 None,
210 None,
211 );
212}
213
214#[kernel(backward = SoftsignBackward)]
218pub fn softsign_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
219 x_ptr: T::Pointer<D>,
220 y_ptr: T::Pointer<D>,
221 n_elements: i32,
222) where
223 T::I32Tensor: types::Tensor<i32, 1>,
224 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
225 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
226{
227 let pid = T::program_id(Axis::X);
228 let block_start = pid * BLOCK_SIZE;
229 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
230 let in_bounds = offsets.lt(n_elements);
231
232 let x = T::load(
233 x_ptr.add_offsets(offsets),
234 Some(in_bounds),
235 None,
236 &[],
237 None,
238 None,
239 None,
240 false,
241 );
242 let one = T::full(&[BLOCK_SIZE], D::from_f64(1.0));
243 let d = one + T::abs(x);
244 let y = x / d;
245 T::store(
246 y_ptr.add_offsets(offsets),
247 y,
248 Some(in_bounds),
249 &[],
250 None,
251 None,
252 );
253}
254
255#[kernel]
257pub fn softsign_backward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
258 dy_ptr: T::Pointer<D>,
259 x_ptr: T::Pointer<D>,
260 dx_ptr: T::Pointer<D>,
261 n_elements: i32,
262) where
263 T::I32Tensor: types::Tensor<i32, 1>,
264 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
265 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
266{
267 let pid = T::program_id(Axis::X);
268 let block_start = pid * BLOCK_SIZE;
269 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
270 let in_bounds = offsets.lt(n_elements);
271
272 let dy = T::load(
273 dy_ptr.add_offsets(offsets),
274 Some(in_bounds),
275 None,
276 &[],
277 None,
278 None,
279 None,
280 false,
281 );
282 let x = T::load(
283 x_ptr.add_offsets(offsets),
284 Some(in_bounds),
285 None,
286 &[],
287 None,
288 None,
289 None,
290 false,
291 );
292 let one = T::full(&[BLOCK_SIZE], D::from_f64(1.0));
293 let d = one + T::abs(x);
294 let dx = dy / (d * d);
295 T::store(
296 dx_ptr.add_offsets(offsets),
297 dx,
298 Some(in_bounds),
299 &[],
300 None,
301 None,
302 );
303}
304
305#[kernel(backward = SoftshrinkBackward)]
309pub fn softshrink_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
310 x_ptr: T::Pointer<D>,
311 y_ptr: T::Pointer<D>,
312 n_elements: i32,
313 lambda: f32,
314) where
315 T::I32Tensor: types::Tensor<i32, 1>,
316 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
317 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
318{
319 let pid = T::program_id(Axis::X);
320 let block_start = pid * BLOCK_SIZE;
321 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
322 let in_bounds = offsets.lt(n_elements);
323
324 let x = T::load(
325 x_ptr.add_offsets(offsets),
326 Some(in_bounds),
327 None,
328 &[],
329 None,
330 None,
331 None,
332 false,
333 );
334 let lam = T::full(&[BLOCK_SIZE], D::from_f64(lambda as f64));
335 let neg_lam = T::full(&[BLOCK_SIZE], D::from_f64(-(lambda as f64)));
336 let x_gt_lam = T::gt(x, lam);
337 let x_lt_neg = T::lt(x, neg_lam);
338 let y_upper = x - lam;
339 let y_lower = x + lam;
340 let y_mid = T::where_(x_lt_neg, y_lower, T::zeros_like(x));
341 let y = T::where_(x_gt_lam, y_upper, y_mid);
342 T::store(
343 y_ptr.add_offsets(offsets),
344 y,
345 Some(in_bounds),
346 &[],
347 None,
348 None,
349 );
350}
351
352#[kernel]
354pub fn softshrink_backward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
355 dy_ptr: T::Pointer<D>,
356 x_ptr: T::Pointer<D>,
357 dx_ptr: T::Pointer<D>,
358 n_elements: i32,
359 lambda: f32,
360) where
361 T::I32Tensor: types::Tensor<i32, 1>,
362 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
363 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
364{
365 let pid = T::program_id(Axis::X);
366 let block_start = pid * BLOCK_SIZE;
367 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
368 let in_bounds = offsets.lt(n_elements);
369
370 let dy = T::load(
371 dy_ptr.add_offsets(offsets),
372 Some(in_bounds),
373 None,
374 &[],
375 None,
376 None,
377 None,
378 false,
379 );
380 let x = T::load(
381 x_ptr.add_offsets(offsets),
382 Some(in_bounds),
383 None,
384 &[],
385 None,
386 None,
387 None,
388 false,
389 );
390 let lam = T::full(&[BLOCK_SIZE], D::from_f64(lambda as f64));
391 let outside = T::gt(T::abs(x), lam);
392 let dx = T::where_(outside, dy, T::zeros_like(dy));
393 T::store(
394 dx_ptr.add_offsets(offsets),
395 dx,
396 Some(in_bounds),
397 &[],
398 None,
399 None,
400 );
401}
402
403#[kernel(backward = SoftplusBackward)]
408pub fn softplus_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
409 x_ptr: T::Pointer<D>,
410 y_ptr: T::Pointer<D>,
411 n_elements: i32,
412 beta: f32,
413 threshold: f32,
414) where
415 T::I32Tensor: types::Tensor<i32, 1>,
416 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
417 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
418{
419 let pid = T::program_id(Axis::X);
420 let block_start = pid * BLOCK_SIZE;
421 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
422 let in_bounds = offsets.lt(n_elements);
423
424 let x = T::load(
425 x_ptr.add_offsets(offsets),
426 Some(in_bounds),
427 None,
428 &[],
429 None,
430 None,
431 None,
432 false,
433 );
434 let beta_t = T::full(&[BLOCK_SIZE], D::from_f64(beta as f64));
435 let inv_beta = T::full(&[BLOCK_SIZE], D::from_f64(1.0 / beta as f64));
436 let thr = T::full(&[BLOCK_SIZE], D::from_f64(threshold as f64));
437 let one = T::full(&[BLOCK_SIZE], D::from_f64(1.0));
438 let bx = beta_t * x;
439 let above_thr = T::gt(bx, thr);
440 let y_safe = inv_beta * T::log(one + T::exp(bx));
441 let y = T::where_(above_thr, x, y_safe);
442 T::store(
443 y_ptr.add_offsets(offsets),
444 y,
445 Some(in_bounds),
446 &[],
447 None,
448 None,
449 );
450}
451
452#[kernel]
455pub fn softplus_backward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
456 dy_ptr: T::Pointer<D>,
457 x_ptr: T::Pointer<D>,
458 dx_ptr: T::Pointer<D>,
459 n_elements: i32,
460 beta: f32,
461 threshold: f32,
462) where
463 T::I32Tensor: types::Tensor<i32, 1>,
464 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
465 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
466{
467 let pid = T::program_id(Axis::X);
468 let block_start = pid * BLOCK_SIZE;
469 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
470 let in_bounds = offsets.lt(n_elements);
471
472 let dy = T::load(
473 dy_ptr.add_offsets(offsets),
474 Some(in_bounds),
475 None,
476 &[],
477 None,
478 None,
479 None,
480 false,
481 );
482 let x = T::load(
483 x_ptr.add_offsets(offsets),
484 Some(in_bounds),
485 None,
486 &[],
487 None,
488 None,
489 None,
490 false,
491 );
492 let beta_t = T::full(&[BLOCK_SIZE], D::from_f64(beta as f64));
493 let neg_beta = T::full(&[BLOCK_SIZE], D::from_f64(-(beta as f64)));
494 let thr = T::full(&[BLOCK_SIZE], D::from_f64(threshold as f64));
495 let one = T::full(&[BLOCK_SIZE], D::from_f64(1.0));
496 let bx = beta_t * x;
497 let neg_bx = neg_beta * x;
498 let above_thr = T::gt(bx, thr);
499 let dx_safe = dy * (one / (one + T::exp(neg_bx)));
501 let dx = T::where_(above_thr, dy, dx_safe);
502 T::store(
503 dx_ptr.add_offsets(offsets),
504 dx,
505 Some(in_bounds),
506 &[],
507 None,
508 None,
509 );
510}
511
512pub struct LeakyReluOp<D: Float> {
513 pub forward: LeakyReluForward<D>,
514 pub backward: LeakyReluBackward<D>,
515}
516
517pub struct ThresholdOp<D: Float> {
518 pub forward: ThresholdForward<D>,
519 pub backward: ThresholdBackward<D>,
520}
521
522pub struct SoftsignOp<D: Float> {
523 pub forward: SoftsignForward<D>,
524 pub backward: SoftsignBackward<D>,
525}
526
527pub struct SoftshrinkOp<D: Float> {
528 pub forward: SoftshrinkForward<D>,
529 pub backward: SoftshrinkBackward<D>,
530}
531
532pub struct SoftplusOp<D: Float> {
533 pub forward: SoftplusForward<D>,
534 pub backward: SoftplusBackward<D>,
535}