1#![allow(non_snake_case)]
18
19use teeny_macros::kernel;
20use teeny_triton::triton::{
21 types::{AddOffsets, Comparison},
22 *,
23};
24
25#[kernel]
41pub fn rprop_step<T: Triton, const BLOCK_SIZE: i32>(
42 params_ptr: T::Pointer<f32>,
43 grad_ptr: T::Pointer<f32>,
44 prev_grad_ptr: T::Pointer<f32>,
45 step_size_ptr: T::Pointer<f32>,
46 n_elements: i32,
47 eta_plus: f32,
48 eta_minus: f32,
49 step_min: f32,
50 step_max: f32,
51) where
52 T::I32Tensor: types::Tensor<i32, 1>,
53 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
54 T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
55{
56 let pid = T::program_id(Axis::X);
57 let block_start = pid * BLOCK_SIZE;
58 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
59 let mask = offsets.lt(n_elements);
60
61 let p = T::load(
62 params_ptr.add_offsets(offsets),
63 Some(mask),
64 None,
65 &[],
66 None,
67 None,
68 None,
69 false,
70 );
71 let g = T::load(
72 grad_ptr.add_offsets(offsets),
73 Some(mask),
74 None,
75 &[],
76 None,
77 None,
78 None,
79 false,
80 );
81 let prev_g = T::load(
82 prev_grad_ptr.add_offsets(offsets),
83 Some(mask),
84 None,
85 &[],
86 None,
87 None,
88 None,
89 false,
90 );
91 let step_size = T::load(
92 step_size_ptr.add_offsets(offsets),
93 Some(mask),
94 None,
95 &[],
96 None,
97 None,
98 None,
99 false,
100 );
101
102 let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
103 let ones = T::full(&[BLOCK_SIZE], 1.0_f32);
104 let neg_ones = T::full(&[BLOCK_SIZE], -1.0_f32);
105 let eta_plus_t = T::full(&[BLOCK_SIZE], eta_plus);
106 let eta_minus_t = T::full(&[BLOCK_SIZE], eta_minus);
107 let step_min_t = T::full(&[BLOCK_SIZE], step_min);
108 let step_max_t = T::full(&[BLOCK_SIZE], step_max);
109
110 let prod = g * prev_g;
111 let sign_pos = T::gt(prod, zeros); let sign_neg = T::lt(prod, zeros); let step_after_pos = T::where_(sign_pos, step_size * eta_plus_t, step_size);
116 let step_scaled = T::where_(sign_neg, step_after_pos * eta_minus_t, step_after_pos);
117 let step_clamped = T::clamp(step_scaled, step_min_t, step_max_t);
118
119 let g_masked = T::where_(sign_neg, zeros, g);
121
122 let g_pos = T::gt(g_masked, zeros);
124 let g_neg = T::lt(g_masked, zeros);
125 let g_sign = T::where_(g_pos, ones, T::where_(g_neg, neg_ones, zeros));
126 let p_new = p - g_sign * step_clamped;
127
128 T::store(
129 params_ptr.add_offsets(offsets),
130 p_new,
131 Some(mask),
132 &[],
133 None,
134 None,
135 );
136 T::store(
137 step_size_ptr.add_offsets(offsets),
138 step_clamped,
139 Some(mask),
140 &[],
141 None,
142 None,
143 );
144 T::store(
145 prev_grad_ptr.add_offsets(offsets),
146 g_masked,
147 Some(mask),
148 &[],
149 None,
150 None,
151 );
152}