1use crate::{
2 api::{CnvPVecAlloc, CnvPVecBytesOf, Convolution},
3 layouts::{
4 Backend, CnvDftAccTerm, CnvPVecL, CnvPVecLBackendMut, CnvPVecLBackendRef, CnvPVecR, CnvPVecRBackendMut,
5 CnvPVecRBackendRef, 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> $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) -> CnvPVecL<BE::OwnedBuf, BE> {
23 CnvPVecL::alloc(self.n(), cols, size)
24 }
25
26 fn cnv_pvec_right_alloc(&self, cols: usize, size: usize) -> CnvPVecR<BE::OwnedBuf, BE> {
27 CnvPVecR::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, res_size: usize, cnv_offset: usize, a_size: usize, b_size: usize) -> usize {
71 <BE as HalConvolutionImpl<BE>>::cnv_by_const_apply_tmp_bytes(self, res_size, cnv_offset, 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_apply_dft(
88 &self,
89 cnv_offset: usize,
90 res: &mut VecZnxDftBackendMut<'_, BE>,
91 res_col: usize,
92 a: &CnvPVecLBackendRef<'_, BE>,
93 a_col: usize,
94 b: &CnvPVecRBackendRef<'_, BE>,
95 b_col: usize,
96 scratch: &mut ScratchArena<'_, BE>,
97 ) {
98 <BE as HalConvolutionImpl<BE>>::cnv_apply_dft(self, cnv_offset, res, res_col, a, a_col, b, b_col, scratch)
99 },
100 fn cnv_prepare_left_lazy_tmp_bytes(&self, res_size: usize, a_size: usize) -> usize {
101 <BE as HalConvolutionImpl<BE>>::cnv_prepare_left_lazy_tmp_bytes(self, res_size, a_size)
102 },
103 fn cnv_prepare_left_lazy(
104 &self,
105 res: &mut CnvPVecLBackendMut<'_, BE>,
106 a: &VecZnxBackendRef<'_, BE>,
107 mask: i64,
108 scratch: &mut ScratchArena<'_, BE>,
109 ) {
110 <BE as HalConvolutionImpl<BE>>::cnv_prepare_left_lazy(self, res, a, mask, scratch)
111 },
112 fn cnv_prepare_right_lazy_tmp_bytes(&self, res_size: usize, a_size: usize) -> usize {
113 <BE as HalConvolutionImpl<BE>>::cnv_prepare_right_lazy_tmp_bytes(self, res_size, a_size)
114 },
115 fn cnv_prepare_right_lazy(
116 &self,
117 res: &mut CnvPVecRBackendMut<'_, BE>,
118 a: &VecZnxBackendRef<'_, BE>,
119 mask: i64,
120 scratch: &mut ScratchArena<'_, BE>,
121 ) {
122 <BE as HalConvolutionImpl<BE>>::cnv_prepare_right_lazy(self, res, a, mask, scratch)
123 },
124 fn cnv_apply_dft_lazy_tmp_bytes(&self, cnv_offset: usize, res_size: usize, a_size: usize, b_size: usize) -> usize {
125 <BE as HalConvolutionImpl<BE>>::cnv_apply_dft_lazy_tmp_bytes(self, cnv_offset, res_size, a_size, b_size)
126 },
127 fn cnv_apply_dft_lazy(
128 &self,
129 cnv_offset: usize,
130 res: &mut VecZnxDftBackendMut<'_, BE>,
131 res_col: usize,
132 a: &CnvPVecLBackendRef<'_, BE>,
133 a_col: usize,
134 b: &CnvPVecRBackendRef<'_, BE>,
135 b_col: usize,
136 scratch: &mut ScratchArena<'_, BE>,
137 ) {
138 <BE as HalConvolutionImpl<BE>>::cnv_apply_dft_lazy(self, cnv_offset, res, res_col, a, a_col, b, b_col, scratch)
139 },
140 fn cnv_apply_dft_accumulate(
141 &self,
142 cnv_offset: usize,
143 res: &mut VecZnxDftBackendMut<'_, BE>,
144 res_col: usize,
145 a: &CnvPVecLBackendRef<'_, BE>,
146 a_col: usize,
147 b: &CnvPVecRBackendRef<'_, BE>,
148 b_col: usize,
149 scratch: &mut ScratchArena<'_, BE>,
150 ) {
151 <BE as HalConvolutionImpl<BE>>::cnv_apply_dft_accumulate(self, cnv_offset, res, res_col, a, a_col, b, b_col, scratch)
152 },
153 fn cnv_accumulate_dft_tmp_bytes(&self, cnv_offset: usize, res_size: usize, a_size: usize, b_size: usize) -> usize {
154 <BE as HalConvolutionImpl<BE>>::cnv_accumulate_dft_tmp_bytes(self, cnv_offset, res_size, a_size, b_size)
155 },
156 fn cnv_accumulate_dft<'a>(
157 &self,
158 cnv_offset: usize,
159 res: &mut VecZnxDftBackendMut<'_, BE>,
160 res_col: usize,
161 terms: &[CnvDftAccTerm<'a, BE>],
162 scratch: &mut ScratchArena<'_, BE>,
163 ) where
164 BE: 'a,
165 {
166 <BE as HalConvolutionImpl<BE>>::cnv_accumulate_dft(self, cnv_offset, res, res_col, terms, scratch)
167 },
168 fn cnv_pairwise_apply_dft_tmp_bytes(&self, cnv_offset: usize, res_size: usize, a_size: usize, b_size: usize) -> usize {
169 <BE as HalConvolutionImpl<BE>>::cnv_pairwise_apply_dft_tmp_bytes(self, cnv_offset, res_size, a_size, b_size)
170 },
171 fn cnv_pairwise_apply_dft(
172 &self,
173 cnv_offset: usize,
174 res: &mut VecZnxDftBackendMut<'_, BE>,
175 res_col: usize,
176 a: &CnvPVecLBackendRef<'_, BE>,
177 b: &CnvPVecRBackendRef<'_, BE>,
178 i: usize,
179 j: usize,
180 scratch: &mut ScratchArena<'_, BE>,
181 ) {
182 <BE as HalConvolutionImpl<BE>>::cnv_pairwise_apply_dft(self, cnv_offset, res, res_col, a, b, i, j, scratch)
183 },
184 fn cnv_prepare_self_tmp_bytes(&self, res_size: usize, a_size: usize) -> usize {
185 <BE as HalConvolutionImpl<BE>>::cnv_prepare_self_tmp_bytes(self, res_size, a_size)
186 },
187 fn cnv_prepare_self(
188 &self,
189 left: &mut CnvPVecLBackendMut<'_, BE>,
190 right: &mut CnvPVecRBackendMut<'_, BE>,
191 a: &VecZnxBackendRef<'_, BE>,
192 mask: i64,
193 scratch: &mut ScratchArena<'_, BE>,
194 ) {
195 <BE as HalConvolutionImpl<BE>>::cnv_prepare_self(self, left, right, a, mask, scratch)
196 }
197);