Skip to main content

poulpy_hal/delegates/
convolution.rs

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