1#![allow(non_snake_case)]
18
19use core::ops::{BitAnd, BitOr};
20
21use teeny_core::dtype::Num;
22use teeny_macros::kernel;
23use teeny_triton::triton::{
24 types::{AddOffsets, Comparison, Tensor},
25 *,
26};
27
28#[kernel]
35pub fn constant_pad2d_forward<
36 T: Triton,
37 D: Num,
38 const PT: i32,
39 const PB: i32,
40 const PL: i32,
41 const PR: i32,
42 const BLOCK_OW: i32,
43>(
44 input_ptr: T::Pointer<D>,
45 output_ptr: T::Pointer<D>,
46 _B: i32,
47 C: i32,
48 H: i32,
49 W: i32,
50 OH: i32,
51 OW: i32,
52 value: f32,
53) where
54 T::I32Tensor: Tensor<i32, 1>,
55 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
56 T::BoolTensor: BitAnd<Output = T::BoolTensor>,
57 T::BoolTensor: BitOr<Output = T::BoolTensor>,
58 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
59{
60 let pid = T::program_id(Axis::X);
61 let num_ow_tiles = T::cdiv(OW, BLOCK_OW);
62
63 let ow_tile = pid % num_ow_tiles;
64 let rest = pid / num_ow_tiles;
65 let oh = rest % OH;
66 let bc = rest / OH;
67 let c = bc % C;
68 let b = bc / C;
69
70 let ow_start = ow_tile * BLOCK_OW;
71 let ow_range = T::arange(0, BLOCK_OW) + ow_start;
72 let ow_mask = ow_range.lt(OW);
73
74 let ih = oh - PT;
75 let value_vec = T::cast::<f32, D>(T::full::<f32>(&[BLOCK_OW], value), None, false);
76 let out_offsets = ow_range + ((b * C + c) * OH + oh) * OW;
77
78 if ih < 0 || ih >= H {
79 T::store(
80 output_ptr.add_offsets(out_offsets),
81 value_vec,
82 Some(ow_mask),
83 &[],
84 None,
85 None,
86 );
87 return;
88 }
89
90 let iw_range = ow_range - PL;
91 let w_in_bounds = iw_range.ge(0) & iw_range.lt(W);
92 let combined_mask = ow_mask & w_in_bounds;
93 let in_bc_base = (b * C + c) * H * W + ih * W;
94
95 let tile = T::load(
96 input_ptr.add_offsets(iw_range + in_bc_base),
97 Some(combined_mask),
98 Some(value_vec),
99 &[],
100 None,
101 None,
102 None,
103 false,
104 );
105 let result = T::where_(combined_mask, tile, value_vec);
106 T::store(
107 output_ptr.add_offsets(out_offsets),
108 result,
109 Some(ow_mask),
110 &[],
111 None,
112 None,
113 );
114}
115
116#[kernel]
118pub fn constant_pad2d_backward<
119 T: Triton,
120 D: Num,
121 const PT: i32,
122 const PB: i32,
123 const PL: i32,
124 const PR: i32,
125 const BLOCK_OW: i32,
126>(
127 dy_ptr: T::Pointer<D>,
128 dx_ptr: T::Pointer<D>,
129 _B: i32,
130 C: i32,
131 H: i32,
132 W: i32,
133 OH: i32,
134 OW: i32,
135) where
136 T::I32Tensor: Tensor<i32, 1>,
137 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
138 T::BoolTensor: BitAnd<Output = T::BoolTensor>,
139 T::BoolTensor: BitOr<Output = T::BoolTensor>,
140 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
141{
142 let pid = T::program_id(Axis::X);
143 let num_ow_tiles = T::cdiv(OW, BLOCK_OW);
144
145 let ow_tile = pid % num_ow_tiles;
146 let rest = pid / num_ow_tiles;
147 let oh = rest % OH;
148 let bc = rest / OH;
149 let c = bc % C;
150 let b = bc / C;
151
152 let ow_start = ow_tile * BLOCK_OW;
153 let ow_range = T::arange(0, BLOCK_OW) + ow_start;
154 let ow_mask = ow_range.lt(OW);
155
156 let ih = oh - PT;
157 if ih < 0 || ih >= H {
158 return;
159 }
160 let iw_range = ow_range - PL;
161
162 let w_in_bounds = iw_range.ge(0) & iw_range.lt(W);
163
164 let dy_bc_base = ((b * C + c) * OH + oh) * OW;
165 let dx_bc_base = (b * C + c) * H * W + ih * W;
166
167 let store_mask = ow_mask & w_in_bounds;
168
169 let dy_offsets = ow_range + dy_bc_base;
170 let dy_tile = T::load(
171 dy_ptr.add_offsets(dy_offsets),
172 Some(ow_mask),
173 Some(T::zeros::<D>(&[BLOCK_OW])),
174 &[],
175 None,
176 None,
177 None,
178 false,
179 );
180
181 let dx_offsets = iw_range + dx_bc_base;
182 T::store(
183 dx_ptr.add_offsets(dx_offsets),
184 dy_tile,
185 Some(store_mask),
186 &[],
187 None,
188 None,
189 );
190}
191
192pub struct ConstantPad2dOp<'a, T: Num> {
193 pub forward: ConstantPad2dForward<T>,
194 pub backward: ConstantPad2dBackward<T>,
195 _marker: core::marker::PhantomData<&'a ()>,
196}