1use crate::{
2 api::{CnvPVecAlloc, CnvPVecBytesOf, Convolution},
3 layouts::{
4 Backend, CnvDftAccTerm, CnvPVecLBackendMut, CnvPVecLBackendRef, CnvPVecLOwned, CnvPVecRBackendMut, CnvPVecRBackendRef,
5 CnvPVecROwned, Module, ScratchArena, VecZnxBackendRef, VecZnxBigBackendMut, VecZnxDftBackendMut,
6 },
7 oep::{HalConvolutionImpl, HalVecZnxDftImpl},
8};
9
10macro_rules! impl_convolution_delegate {
11 ($trait:ty, $($body:item),+ $(,)?) => {
12 impl<BE: Backend<ZnxWord = i64>> $trait for Module<BE>
13 where
14 BE: HalConvolutionImpl<BE> + HalVecZnxDftImpl<BE>,
15 {
16 $($body)+
17 }
18 };
19}
20
21impl<BE: Backend> CnvPVecAlloc<BE> for Module<BE> {
22 fn cnv_pvec_left_alloc(&self, cols: usize, size: usize) -> CnvPVecLOwned<BE> {
23 CnvPVecLOwned::<BE>::alloc(self.n(), cols, size)
24 }
25
26 fn cnv_pvec_right_alloc(&self, cols: usize, size: usize) -> CnvPVecROwned<BE> {
27 CnvPVecROwned::<BE>::alloc(self.n(), cols, size)
28 }
29}
30
31impl<BE: Backend> CnvPVecBytesOf for Module<BE> {
32 fn bytes_of_cnv_pvec_left(&self, cols: usize, size: usize) -> usize {
33 BE::bytes_of_cnv_pvec_left(self.n(), cols, size)
34 }
35
36 fn bytes_of_cnv_pvec_right(&self, cols: usize, size: usize) -> usize {
37 BE::bytes_of_cnv_pvec_right(self.n(), cols, size)
38 }
39}
40
41impl_convolution_delegate!(
42 Convolution<BE>,
43 fn cnv_prepare_left_tmp_bytes(&self, res_size: usize, a_size: usize) -> usize {
44 <BE as HalConvolutionImpl<BE>>::cnv_prepare_left_tmp_bytes(self, res_size, a_size)
45 },
46 fn cnv_prepare_left(
47 &self,
48 res: &mut CnvPVecLBackendMut<'_, BE>,
49 a: &VecZnxBackendRef<'_, BE>,
50 mask: i64,
51 scratch: &mut ScratchArena<'_, BE>,
52 ) {
53 <BE as HalConvolutionImpl<BE>>::cnv_prepare_left(self, res, a, mask, scratch);
54 },
55 fn cnv_prepare_right_tmp_bytes(&self, res_size: usize, a_size: usize) -> usize {
56 <BE as HalConvolutionImpl<BE>>::cnv_prepare_right_tmp_bytes(self, res_size, a_size)
57 },
58 fn cnv_prepare_right(
59 &self,
60 res: &mut CnvPVecRBackendMut<'_, BE>,
61 a: &VecZnxBackendRef<'_, BE>,
62 mask: i64,
63 scratch: &mut ScratchArena<'_, BE>,
64 ) {
65 <BE as HalConvolutionImpl<BE>>::cnv_prepare_right(self, res, a, mask, scratch);
66 },
67 fn cnv_apply_dft_tmp_bytes(&self, cnv_offset: usize, res_size: usize, a_size: usize, b_size: usize) -> usize {
68 <BE as HalConvolutionImpl<BE>>::cnv_apply_dft_tmp_bytes(self, cnv_offset, res_size, a_size, b_size)
69 },
70 fn cnv_by_const_apply_tmp_bytes(&self, cnv_offset: usize, res_size: usize, a_size: usize, b_size: usize) -> usize {
71 <BE as HalConvolutionImpl<BE>>::cnv_by_const_apply_tmp_bytes(self, cnv_offset, res_size, a_size, b_size)
72 },
73 fn cnv_by_const_apply(
74 &self,
75 cnv_offset: usize,
76 res: &mut VecZnxBigBackendMut<'_, BE>,
77 res_col: usize,
78 a: &VecZnxBackendRef<'_, BE>,
79 a_col: usize,
80 b: &VecZnxBackendRef<'_, BE>,
81 b_col: usize,
82 b_coeff: usize,
83 scratch: &mut ScratchArena<'_, BE>,
84 ) {
85 <BE as HalConvolutionImpl<BE>>::cnv_by_const_apply(self, cnv_offset, res, res_col, a, a_col, b, b_col, b_coeff, scratch)
86 },
87 fn cnv_by_const_apply_add(
88 &self,
89 cnv_offset: usize,
90 res: &mut VecZnxBigBackendMut<'_, BE>,
91 res_col: usize,
92 a: &VecZnxBackendRef<'_, BE>,
93 a_col: usize,
94 b: &VecZnxBackendRef<'_, BE>,
95 b_col: usize,
96 b_coeff: usize,
97 scratch: &mut ScratchArena<'_, BE>,
98 ) {
99 <BE as HalConvolutionImpl<BE>>::cnv_by_const_apply_add(
100 self, cnv_offset, res, res_col, a, a_col, b, b_col, b_coeff, scratch,
101 )
102 },
103 fn cnv_apply_dft(
104 &self,
105 cnv_offset: usize,
106 res: &mut VecZnxDftBackendMut<'_, BE>,
107 res_col: usize,
108 a: &CnvPVecLBackendRef<'_, BE>,
109 a_col: usize,
110 b: &CnvPVecRBackendRef<'_, BE>,
111 b_col: usize,
112 scratch: &mut ScratchArena<'_, BE>,
113 ) {
114 <BE as HalConvolutionImpl<BE>>::cnv_apply_dft(self, cnv_offset, res, res_col, a, a_col, b, b_col, scratch)
115 },
116 fn cnv_prepare_left_lazy_tmp_bytes(&self, res_size: usize, a_size: usize) -> usize {
117 <BE as HalConvolutionImpl<BE>>::cnv_prepare_left_lazy_tmp_bytes(self, res_size, a_size)
118 },
119 fn cnv_prepare_left_lazy(
120 &self,
121 res: &mut CnvPVecLBackendMut<'_, BE>,
122 a: &VecZnxBackendRef<'_, BE>,
123 mask: i64,
124 scratch: &mut ScratchArena<'_, BE>,
125 ) {
126 <BE as HalConvolutionImpl<BE>>::cnv_prepare_left_lazy(self, res, a, mask, scratch)
127 },
128 fn cnv_prepare_right_lazy_tmp_bytes(&self, res_size: usize, a_size: usize) -> usize {
129 <BE as HalConvolutionImpl<BE>>::cnv_prepare_right_lazy_tmp_bytes(self, res_size, a_size)
130 },
131 fn cnv_prepare_right_lazy(
132 &self,
133 res: &mut CnvPVecRBackendMut<'_, BE>,
134 a: &VecZnxBackendRef<'_, BE>,
135 mask: i64,
136 scratch: &mut ScratchArena<'_, BE>,
137 ) {
138 <BE as HalConvolutionImpl<BE>>::cnv_prepare_right_lazy(self, res, a, mask, scratch)
139 },
140 fn cnv_apply_dft_lazy_tmp_bytes(&self, cnv_offset: usize, res_size: usize, a_size: usize, b_size: usize) -> usize {
141 <BE as HalConvolutionImpl<BE>>::cnv_apply_dft_lazy_tmp_bytes(self, cnv_offset, res_size, a_size, b_size)
142 },
143 fn cnv_apply_dft_lazy(
144 &self,
145 cnv_offset: usize,
146 res: &mut VecZnxDftBackendMut<'_, BE>,
147 res_col: usize,
148 a: &CnvPVecLBackendRef<'_, BE>,
149 a_col: usize,
150 b: &CnvPVecRBackendRef<'_, BE>,
151 b_col: usize,
152 scratch: &mut ScratchArena<'_, BE>,
153 ) {
154 <BE as HalConvolutionImpl<BE>>::cnv_apply_dft_lazy(self, cnv_offset, res, res_col, a, a_col, b, b_col, scratch)
155 },
156 fn cnv_apply_dft_accumulate(
157 &self,
158 cnv_offset: usize,
159 res: &mut VecZnxDftBackendMut<'_, BE>,
160 res_col: usize,
161 a: &CnvPVecLBackendRef<'_, BE>,
162 a_col: usize,
163 b: &CnvPVecRBackendRef<'_, BE>,
164 b_col: usize,
165 scratch: &mut ScratchArena<'_, BE>,
166 ) {
167 <BE as HalConvolutionImpl<BE>>::cnv_apply_dft_accumulate(self, cnv_offset, res, res_col, a, a_col, b, b_col, scratch)
168 },
169 fn cnv_accumulate_dft_tmp_bytes(&self, cnv_offset: usize, res_size: usize, a_size: usize, b_size: usize) -> usize {
170 <BE as HalConvolutionImpl<BE>>::cnv_accumulate_dft_tmp_bytes(self, cnv_offset, res_size, a_size, b_size)
171 },
172 fn cnv_accumulate_dft<'a>(
173 &self,
174 cnv_offset: usize,
175 res: &mut VecZnxDftBackendMut<'_, BE>,
176 res_col: usize,
177 terms: &[CnvDftAccTerm<'a, BE>],
178 scratch: &mut ScratchArena<'_, BE>,
179 ) where
180 BE: 'a,
181 {
182 <BE as HalConvolutionImpl<BE>>::cnv_accumulate_dft(self, cnv_offset, res, res_col, terms, scratch)
183 },
184 fn cnv_pairwise_apply_dft_tmp_bytes(&self, cnv_offset: usize, res_size: usize, a_size: usize, b_size: usize) -> usize {
185 <BE as HalConvolutionImpl<BE>>::cnv_pairwise_apply_dft_tmp_bytes(self, cnv_offset, res_size, a_size, b_size)
186 },
187 fn cnv_pairwise_apply_dft(
188 &self,
189 cnv_offset: usize,
190 res: &mut VecZnxDftBackendMut<'_, BE>,
191 res_col: usize,
192 a: &CnvPVecLBackendRef<'_, BE>,
193 b: &CnvPVecRBackendRef<'_, BE>,
194 i: usize,
195 j: usize,
196 scratch: &mut ScratchArena<'_, BE>,
197 ) {
198 <BE as HalConvolutionImpl<BE>>::cnv_pairwise_apply_dft(self, cnv_offset, res, res_col, a, b, i, j, scratch)
199 },
200 fn cnv_prepare_self_tmp_bytes(&self, res_size: usize, a_size: usize) -> usize {
201 <BE as HalConvolutionImpl<BE>>::cnv_prepare_self_tmp_bytes(self, res_size, a_size)
202 },
203 fn cnv_prepare_self(
204 &self,
205 left: &mut CnvPVecLBackendMut<'_, BE>,
206 right: &mut CnvPVecRBackendMut<'_, BE>,
207 a: &VecZnxBackendRef<'_, BE>,
208 mask: i64,
209 scratch: &mut ScratchArena<'_, BE>,
210 ) {
211 <BE as HalConvolutionImpl<BE>>::cnv_prepare_self(self, left, right, a, mask, scratch)
212 }
213);