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 = HardtanhBackward)]
30pub fn hardtanh_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 min_val: f32,
35 max_val: f32,
36) where
37 T::I32Tensor: types::Tensor<i32, 1>,
38 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
39 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
40{
41 let pid = T::program_id(Axis::X);
42 let block_start = pid * BLOCK_SIZE;
43 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
44 let in_bounds = offsets.lt(n_elements);
45
46 let x = T::load(
47 x_ptr.add_offsets(offsets),
48 Some(in_bounds),
49 None,
50 &[],
51 None,
52 None,
53 None,
54 false,
55 );
56 let lo = T::full(&[BLOCK_SIZE], D::from_f64(min_val as f64));
57 let hi = T::full(&[BLOCK_SIZE], D::from_f64(max_val as f64));
58 let y = T::clamp(x, lo, hi);
59 T::store(
60 y_ptr.add_offsets(offsets),
61 y,
62 Some(in_bounds),
63 &[],
64 None,
65 None,
66 );
67}
68
69#[kernel]
71pub fn hardtanh_backward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
72 dy_ptr: T::Pointer<D>,
73 x_ptr: T::Pointer<D>,
74 dx_ptr: T::Pointer<D>,
75 n_elements: i32,
76 min_val: f32,
77 max_val: f32,
78) where
79 T::I32Tensor: types::Tensor<i32, 1>,
80 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
81 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
82{
83 let pid = T::program_id(Axis::X);
84 let block_start = pid * BLOCK_SIZE;
85 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
86 let in_bounds = offsets.lt(n_elements);
87
88 let dy = T::load(
89 dy_ptr.add_offsets(offsets),
90 Some(in_bounds),
91 None,
92 &[],
93 None,
94 None,
95 None,
96 false,
97 );
98 let x = T::load(
99 x_ptr.add_offsets(offsets),
100 Some(in_bounds),
101 None,
102 &[],
103 None,
104 None,
105 None,
106 false,
107 );
108
109 let lo = T::full(&[BLOCK_SIZE], D::from_f64(min_val as f64));
111 let hi = T::full(&[BLOCK_SIZE], D::from_f64(max_val as f64));
112 let in_range = T::gt(T::minimum(x - lo, hi - x), T::zeros_like(x));
113 let dx = T::where_(in_range, dy, T::zeros_like(dy));
114 T::store(
115 dx_ptr.add_offsets(offsets),
116 dx,
117 Some(in_bounds),
118 &[],
119 None,
120 None,
121 );
122}
123
124#[kernel(backward = Relu6Backward)]
128pub fn relu6_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
129 x_ptr: T::Pointer<D>,
130 y_ptr: T::Pointer<D>,
131 n_elements: i32,
132) where
133 T::I32Tensor: types::Tensor<i32, 1>,
134 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
135 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
136{
137 let pid = T::program_id(Axis::X);
138 let block_start = pid * BLOCK_SIZE;
139 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
140 let in_bounds = offsets.lt(n_elements);
141
142 let x = T::load(
143 x_ptr.add_offsets(offsets),
144 Some(in_bounds),
145 None,
146 &[],
147 None,
148 None,
149 None,
150 false,
151 );
152 let lo = T::zeros_like(x);
153 let hi = T::full(&[BLOCK_SIZE], D::from_f64(6.0));
154 let y = T::clamp(x, lo, hi);
155 T::store(
156 y_ptr.add_offsets(offsets),
157 y,
158 Some(in_bounds),
159 &[],
160 None,
161 None,
162 );
163}
164
165#[kernel]
167pub fn relu6_backward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
168 dy_ptr: T::Pointer<D>,
169 x_ptr: T::Pointer<D>,
170 dx_ptr: T::Pointer<D>,
171 n_elements: i32,
172) where
173 T::I32Tensor: types::Tensor<i32, 1>,
174 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
175 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
176{
177 let pid = T::program_id(Axis::X);
178 let block_start = pid * BLOCK_SIZE;
179 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
180 let in_bounds = offsets.lt(n_elements);
181
182 let dy = T::load(
183 dy_ptr.add_offsets(offsets),
184 Some(in_bounds),
185 None,
186 &[],
187 None,
188 None,
189 None,
190 false,
191 );
192 let x = T::load(
193 x_ptr.add_offsets(offsets),
194 Some(in_bounds),
195 None,
196 &[],
197 None,
198 None,
199 None,
200 false,
201 );
202
203 let six = T::full(&[BLOCK_SIZE], D::from_f64(6.0));
204 let in_range = T::gt(T::minimum(x, six - x), T::zeros_like(x));
206 let dx = T::where_(in_range, dy, T::zeros_like(dy));
207 T::store(
208 dx_ptr.add_offsets(offsets),
209 dx,
210 Some(in_bounds),
211 &[],
212 None,
213 None,
214 );
215}
216
217#[kernel(backward = HardsigmoidBackward)]
221pub fn hardsigmoid_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
222 x_ptr: T::Pointer<D>,
223 y_ptr: T::Pointer<D>,
224 n_elements: i32,
225) where
226 T::I32Tensor: types::Tensor<i32, 1>,
227 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
228 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
229{
230 let pid = T::program_id(Axis::X);
231 let block_start = pid * BLOCK_SIZE;
232 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
233 let in_bounds = offsets.lt(n_elements);
234
235 let x = T::load(
236 x_ptr.add_offsets(offsets),
237 Some(in_bounds),
238 None,
239 &[],
240 None,
241 None,
242 None,
243 false,
244 );
245 let three = T::full(&[BLOCK_SIZE], D::from_f64(3.0));
246 let six = T::full(&[BLOCK_SIZE], D::from_f64(6.0));
247 let lo = T::zeros_like(x);
248 let hi = T::full(&[BLOCK_SIZE], D::from_f64(1.0));
249 let y = T::clamp((x + three) / six, lo, hi);
250 T::store(
251 y_ptr.add_offsets(offsets),
252 y,
253 Some(in_bounds),
254 &[],
255 None,
256 None,
257 );
258}
259
260#[kernel]
262pub fn hardsigmoid_backward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
263 dy_ptr: T::Pointer<D>,
264 x_ptr: T::Pointer<D>,
265 dx_ptr: T::Pointer<D>,
266 n_elements: i32,
267) where
268 T::I32Tensor: types::Tensor<i32, 1>,
269 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
270 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
271{
272 let pid = T::program_id(Axis::X);
273 let block_start = pid * BLOCK_SIZE;
274 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
275 let in_bounds = offsets.lt(n_elements);
276
277 let dy = T::load(
278 dy_ptr.add_offsets(offsets),
279 Some(in_bounds),
280 None,
281 &[],
282 None,
283 None,
284 None,
285 false,
286 );
287 let x = T::load(
288 x_ptr.add_offsets(offsets),
289 Some(in_bounds),
290 None,
291 &[],
292 None,
293 None,
294 None,
295 false,
296 );
297
298 let three = T::full(&[BLOCK_SIZE], D::from_f64(3.0));
299 let in_range = T::lt(T::abs(x), three); let sixth = T::full(&[BLOCK_SIZE], D::from_f64(1.0 / 6.0));
301 let dx = T::where_(in_range, dy * sixth, T::zeros_like(dy));
302 T::store(
303 dx_ptr.add_offsets(offsets),
304 dx,
305 Some(in_bounds),
306 &[],
307 None,
308 None,
309 );
310}
311
312#[kernel(backward = HardswishBackward)]
316pub fn hardswish_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
317 x_ptr: T::Pointer<D>,
318 y_ptr: T::Pointer<D>,
319 n_elements: i32,
320) where
321 T::I32Tensor: types::Tensor<i32, 1>,
322 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
323 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
324{
325 let pid = T::program_id(Axis::X);
326 let block_start = pid * BLOCK_SIZE;
327 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
328 let in_bounds = offsets.lt(n_elements);
329
330 let x = T::load(
331 x_ptr.add_offsets(offsets),
332 Some(in_bounds),
333 None,
334 &[],
335 None,
336 None,
337 None,
338 false,
339 );
340 let three = T::full(&[BLOCK_SIZE], D::from_f64(3.0));
341 let six = T::full(&[BLOCK_SIZE], D::from_f64(6.0));
342 let lo = T::zeros_like(x);
343 let hi = T::full(&[BLOCK_SIZE], D::from_f64(1.0));
344 let hs = T::clamp((x + three) / six, lo, hi);
345 let y = x * hs;
346 T::store(
347 y_ptr.add_offsets(offsets),
348 y,
349 Some(in_bounds),
350 &[],
351 None,
352 None,
353 );
354}
355
356#[kernel]
361pub fn hardswish_backward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
362 dy_ptr: T::Pointer<D>,
363 x_ptr: T::Pointer<D>,
364 dx_ptr: T::Pointer<D>,
365 n_elements: i32,
366) where
367 T::I32Tensor: types::Tensor<i32, 1>,
368 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
369 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
370{
371 let pid = T::program_id(Axis::X);
372 let block_start = pid * BLOCK_SIZE;
373 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
374 let in_bounds = offsets.lt(n_elements);
375
376 let dy = T::load(
377 dy_ptr.add_offsets(offsets),
378 Some(in_bounds),
379 None,
380 &[],
381 None,
382 None,
383 None,
384 false,
385 );
386 let x = T::load(
387 x_ptr.add_offsets(offsets),
388 Some(in_bounds),
389 None,
390 &[],
391 None,
392 None,
393 None,
394 false,
395 );
396
397 let three = T::full(&[BLOCK_SIZE], D::from_f64(3.0));
398 let neg_three = T::full(&[BLOCK_SIZE], D::from_f64(-3.0));
399 let six = T::full(&[BLOCK_SIZE], D::from_f64(6.0));
400 let two = T::full(&[BLOCK_SIZE], D::from_f64(2.0));
401
402 let x_le_neg3 = T::le(x, neg_three);
403 let x_ge_3 = T::ge(x, three);
404 let dx_mid = dy * (two * x + three) / six;
405
406 let dx_not_lo = T::where_(x_ge_3, dy, dx_mid);
408 let dx = T::where_(x_le_neg3, T::zeros_like(dy), dx_not_lo);
409 T::store(
410 dx_ptr.add_offsets(offsets),
411 dx,
412 Some(in_bounds),
413 &[],
414 None,
415 None,
416 );
417}
418
419#[kernel(backward = HardshrinkBackward)]
423pub fn hardshrink_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
424 x_ptr: T::Pointer<D>,
425 y_ptr: T::Pointer<D>,
426 n_elements: i32,
427 lambda: f32,
428) where
429 T::I32Tensor: types::Tensor<i32, 1>,
430 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
431 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
432{
433 let pid = T::program_id(Axis::X);
434 let block_start = pid * BLOCK_SIZE;
435 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
436 let in_bounds = offsets.lt(n_elements);
437
438 let x = T::load(
439 x_ptr.add_offsets(offsets),
440 Some(in_bounds),
441 None,
442 &[],
443 None,
444 None,
445 None,
446 false,
447 );
448 let lam = T::full(&[BLOCK_SIZE], D::from_f64(lambda as f64));
449 let outside = T::gt(T::abs(x), lam);
450 let y = T::where_(outside, x, T::zeros_like(x));
451 T::store(
452 y_ptr.add_offsets(offsets),
453 y,
454 Some(in_bounds),
455 &[],
456 None,
457 None,
458 );
459}
460
461#[kernel]
463pub fn hardshrink_backward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
464 dy_ptr: T::Pointer<D>,
465 x_ptr: T::Pointer<D>,
466 dx_ptr: T::Pointer<D>,
467 n_elements: i32,
468 lambda: f32,
469) where
470 T::I32Tensor: types::Tensor<i32, 1>,
471 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
472 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
473{
474 let pid = T::program_id(Axis::X);
475 let block_start = pid * BLOCK_SIZE;
476 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
477 let in_bounds = offsets.lt(n_elements);
478
479 let dy = T::load(
480 dy_ptr.add_offsets(offsets),
481 Some(in_bounds),
482 None,
483 &[],
484 None,
485 None,
486 None,
487 false,
488 );
489 let x = T::load(
490 x_ptr.add_offsets(offsets),
491 Some(in_bounds),
492 None,
493 &[],
494 None,
495 None,
496 None,
497 false,
498 );
499 let lam = T::full(&[BLOCK_SIZE], D::from_f64(lambda as f64));
500 let outside = T::gt(T::abs(x), lam);
501 let dx = T::where_(outside, dy, T::zeros_like(dy));
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 HardtanhOp<D: Float> {
513 pub forward: HardtanhForward<D>,
514 pub backward: HardtanhBackward<D>,
515}
516
517pub struct Relu6Op<D: Float> {
518 pub forward: Relu6Forward<D>,
519 pub backward: Relu6Backward<D>,
520}
521
522pub struct HardsigmoidOp<D: Float> {
523 pub forward: HardsigmoidForward<D>,
524 pub backward: HardsigmoidBackward<D>,
525}
526
527pub struct HardswishOp<D: Float> {
528 pub forward: HardswishForward<D>,
529 pub backward: HardswishBackward<D>,
530}
531
532pub struct HardshrinkOp<D: Float> {
533 pub forward: HardshrinkForward<D>,
534 pub backward: HardshrinkBackward<D>,
535}