Skip to main content

poulpy_hal/delegates/
convolution.rs

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);