Skip to main content

poulpy_core/layouts/
scratch_views.rs

1use poulpy_hal::layouts::{
2    Backend, ScalarZnx, SvpPPolReborrowBackendMut, SvpPPolReborrowBackendRef, VmpPMatReborrowBackendMut,
3    VmpPMatReborrowBackendRef, mat_znx_backend_mut_from_mut, mat_znx_backend_ref_from_mut, vec_znx_backend_mut_from_mut,
4    vec_znx_backend_ref_from_mut, vec_znx_backend_ref_from_ref,
5};
6
7use crate::{
8    GetDistribution, GetDistributionMut,
9    dist::Distribution,
10    layouts::{
11        Base2K, GGLWE, GGLWEBackendMut, GGLWEBackendRef, GGLWEInfos, GGLWEPrepared, GGLWEPreparedBackendMut,
12        GGLWEPreparedBackendRef, GGLWEPreparedToBackendMut, GGLWEPreparedToBackendRef, GGLWEToBackendMut, GGLWEToBackendRef,
13        GGSW, GGSWBackendMut, GGSWBackendRef, GGSWInfos, GGSWPrepared, GGSWPreparedBackendMut, GGSWPreparedBackendRef,
14        GGSWPreparedToBackendMut, GGSWPreparedToBackendRef, GGSWToBackendMut, GGSWToBackendRef, GLWE, GLWEBackendMut,
15        GLWEBackendRef, GLWEPlaintext, GLWESecret, GLWESecretBackendMut, GLWESecretBackendRef, GLWESecretPrepared,
16        GLWESecretPreparedBackendMut, GLWESecretPreparedBackendRef, GLWESecretPreparedToBackendMut,
17        GLWESecretPreparedToBackendRef, GLWESecretTensor, GLWESecretTensorBackendMut, GLWESecretTensorBackendRef,
18        GLWESecretTensorToBackendMut, GLWESecretTensorToBackendRef, GLWESecretToBackendMut, GLWESecretToBackendRef, GLWETensor,
19        GLWEToBackendMut, GLWEToBackendRef, LWE, LWEBackendMut, LWEBackendRef, LWEPlaintext, LWEPlaintextBackendMut,
20        LWEPlaintextBackendRef, LWEPlaintextToBackendMut, LWEPlaintextToBackendRef, LWEToBackendMut, LWEToBackendRef, Rank,
21        SetBase2k, SetGGLWEInfos, SetK, TorusPrecision,
22    },
23};
24
25/// Defines a nominal mutable scratch view over a backend-borrowed layout.
26///
27/// The wrapper gives projected backend buffer types a distinct identity for
28/// trait coherence while forwarding the common layout metadata and deref
29/// surface to the wrapped value.
30#[macro_export]
31macro_rules! view_wrapper {
32    ($(#[$meta:meta])* $name:ident, $inner:ty) => {
33        $(#[$meta])*
34        pub struct $name<'a, BE: ::poulpy_hal::layouts::Backend + 'a> {
35            inner: $inner,
36        }
37
38        impl<'a, BE: ::poulpy_hal::layouts::Backend + 'a> $name<'a, BE> {
39            pub fn from_inner(inner: $inner) -> Self {
40                Self { inner }
41            }
42
43            pub fn into_inner(self) -> $inner {
44                self.inner
45            }
46        }
47
48        impl<'a, BE: ::poulpy_hal::layouts::Backend + 'a> ::core::ops::Deref for $name<'a, BE> {
49            type Target = $inner;
50
51            fn deref(&self) -> &Self::Target {
52                &self.inner
53            }
54        }
55
56        impl<'a, BE: ::poulpy_hal::layouts::Backend + 'a> ::core::ops::DerefMut for $name<'a, BE> {
57            fn deref_mut(&mut self) -> &mut Self::Target {
58                &mut self.inner
59            }
60        }
61
62        impl<'a, BE: ::poulpy_hal::layouts::Backend + 'a> $crate::layouts::LWEInfos for $name<'a, BE> {
63            fn base2k(&self) -> $crate::layouts::Base2K {
64                $crate::layouts::LWEInfos::base2k(&self.inner)
65            }
66
67            fn n(&self) -> $crate::layouts::Degree {
68                $crate::layouts::LWEInfos::n(&self.inner)
69            }
70
71            fn max_size(&self) -> usize {
72                $crate::layouts::LWEInfos::max_size(&self.inner)
73            }
74
75            fn size(&self) -> usize {
76                $crate::layouts::LWEInfos::size(&self.inner)
77            }
78
79            fn k(&self) -> $crate::layouts::TorusPrecision {
80                $crate::layouts::LWEInfos::k(&self.inner)
81            }
82        }
83    };
84}
85
86view_wrapper!(LWEViewMut, LWE<BE::BufMut<'a>, BE::ZnxWord>);
87view_wrapper!(LWEPlaintextViewMut, LWEPlaintext<BE::BufMut<'a>, BE::ZnxWord>);
88view_wrapper!(GLWEViewRef, GLWE<BE::BufRef<'a>, BE::ZnxWord>);
89view_wrapper!(GLWEViewMut, GLWE<BE::BufMut<'a>, BE::ZnxWord>);
90view_wrapper!(GLWEPlaintextViewMut, GLWEPlaintext<BE::BufMut<'a>, BE::ZnxWord>);
91view_wrapper!(GLWETensorViewMut, GLWETensor<BE::BufMut<'a>, BE::ZnxWord>);
92view_wrapper!(GLWESecretViewMut, GLWESecret<BE::BufMut<'a>, BE::ZnxWord>);
93view_wrapper!(GLWESecretTensorViewMut, GLWESecretTensor<BE::BufMut<'a>, BE::ZnxWord>);
94view_wrapper!(GLWESecretPreparedViewMut, GLWESecretPrepared<BE::BufMut<'a>, BE>);
95view_wrapper!(GGLWEViewMut, GGLWE<BE::BufMut<'a>, BE::ZnxWord>);
96view_wrapper!(GGLWEPreparedViewMut, GGLWEPrepared<BE::BufMut<'a>, BE>);
97view_wrapper!(GGSWViewMut, GGSW<BE::BufMut<'a>, BE::ZnxWord>);
98view_wrapper!(GGSWPreparedViewMut, GGSWPrepared<BE::BufMut<'a>, BE>);
99
100impl<'a, BE: Backend + 'a> GGLWEViewMut<'a, BE> {
101    pub fn at_view(&self, row: usize, col: usize) -> GLWEViewRef<'_, BE> {
102        GLWEViewRef::from_inner(crate::layouts::gglwe_at_backend_ref_from_mut::<BE>(&self.inner, row, col))
103    }
104
105    pub fn at_view_mut(&mut self, row: usize, col: usize) -> GLWEViewMut<'_, BE> {
106        GLWEViewMut::from_inner(crate::layouts::gglwe_at_backend_mut_from_mut::<BE>(&mut self.inner, row, col))
107    }
108}
109
110macro_rules! impl_set_lwe_infos {
111    ($name:ident) => {
112        impl<'a, BE: Backend + 'a> SetBase2k for $name<'a, BE> {
113            fn set_base2k(&mut self, base2k: Base2K) {
114                self.inner.set_base2k(base2k);
115            }
116        }
117    };
118}
119
120impl_set_lwe_infos!(LWEViewMut);
121impl_set_lwe_infos!(GLWEViewMut);
122impl_set_lwe_infos!(GLWEPlaintextViewMut);
123
124impl<'a, BE: Backend + 'a> crate::layouts::IntPolyInfos for GLWEPlaintextViewMut<'a, BE> {
125    fn encoded_k(&self) -> crate::layouts::TorusPrecision {
126        self.inner.encoded_k()
127    }
128}
129
130impl<'a, BE: Backend + 'a> crate::layouts::IntPolyInfos for LWEPlaintextViewMut<'a, BE> {
131    fn encoded_k(&self) -> crate::layouts::TorusPrecision {
132        self.inner.encoded_k()
133    }
134}
135
136impl<'a, BE: Backend + 'a> SetK for GLWEViewMut<'a, BE> {
137    fn set_k(&mut self, k: TorusPrecision) {
138        self.inner.set_k(k);
139    }
140}
141
142impl<'a, BE: Backend + 'a> SetBase2k for LWEPlaintextViewMut<'a, BE> {
143    fn set_base2k(&mut self, base2k: Base2K) {
144        self.inner.base2k = base2k;
145    }
146}
147
148/// Forwards [`GLWEInfos`](crate::layouts::GLWEInfos) through a nominal
149/// backend view wrapper generated by [`view_wrapper!`](crate::view_wrapper).
150#[macro_export]
151macro_rules! impl_glwe_infos {
152    ($name:ident) => {
153        impl<'a, BE: ::poulpy_hal::layouts::Backend + 'a> $crate::layouts::GLWEInfos for $name<'a, BE> {
154            fn rank(&self) -> $crate::layouts::Rank {
155                $crate::layouts::GLWEInfos::rank(&self.inner)
156            }
157        }
158    };
159}
160
161impl_glwe_infos!(GLWEViewMut);
162impl_glwe_infos!(GLWEViewRef);
163impl_glwe_infos!(GLWEPlaintextViewMut);
164impl_glwe_infos!(GLWETensorViewMut);
165impl_glwe_infos!(GLWESecretViewMut);
166impl_glwe_infos!(GLWESecretTensorViewMut);
167impl_glwe_infos!(GLWESecretPreparedViewMut);
168impl_glwe_infos!(GGLWEViewMut);
169impl_glwe_infos!(GGLWEPreparedViewMut);
170impl_glwe_infos!(GGSWViewMut);
171impl_glwe_infos!(GGSWPreparedViewMut);
172
173macro_rules! impl_dist {
174    ($name:ident) => {
175        impl<'a, BE: Backend + 'a> GetDistribution for $name<'a, BE> {
176            fn dist(&self) -> &Distribution {
177                self.inner.dist()
178            }
179        }
180
181        impl<'a, BE: Backend + 'a> GetDistributionMut for $name<'a, BE> {
182            fn dist_mut(&mut self) -> &mut Distribution {
183                self.inner.dist_mut()
184            }
185        }
186    };
187}
188
189impl_dist!(GLWESecretTensorViewMut);
190impl_dist!(GLWESecretPreparedViewMut);
191
192impl<'a, BE: Backend + 'a> GetDistribution for GLWESecretViewMut<'a, BE> {
193    fn dist(&self) -> &Distribution {
194        self.inner.dist()
195    }
196}
197
198impl<'a, BE: Backend + 'a> GGLWEInfos for GGLWEViewMut<'a, BE> {
199    fn k_aux(&self) -> crate::layouts::TorusPrecision {
200        self.inner.k_aux()
201    }
202
203    fn dnum(&self) -> crate::layouts::Dnum {
204        self.inner.dnum()
205    }
206
207    fn dsize(&self) -> crate::layouts::Dsize {
208        self.inner.dsize()
209    }
210
211    fn rank_in(&self) -> Rank {
212        self.inner.rank_in()
213    }
214
215    fn rank_out(&self) -> Rank {
216        self.inner.rank_out()
217    }
218}
219
220impl<'a, BE: Backend + 'a> GGLWEInfos for GGLWEPreparedViewMut<'a, BE> {
221    fn k_aux(&self) -> crate::layouts::TorusPrecision {
222        self.inner.k_aux()
223    }
224
225    fn dnum(&self) -> crate::layouts::Dnum {
226        self.inner.dnum()
227    }
228
229    fn dsize(&self) -> crate::layouts::Dsize {
230        self.inner.dsize()
231    }
232
233    fn rank_in(&self) -> Rank {
234        self.inner.rank_in()
235    }
236
237    fn rank_out(&self) -> Rank {
238        self.inner.rank_out()
239    }
240}
241
242impl<'a, BE: Backend + 'a> SetGGLWEInfos for GGLWEViewMut<'a, BE> {
243    fn set_dsize(&mut self, dsize: usize) {
244        self.inner.dsize = dsize.into();
245    }
246}
247
248impl<'a, BE: Backend + 'a> GGSWInfos for GGSWViewMut<'a, BE> {
249    fn k_aux(&self) -> crate::layouts::TorusPrecision {
250        self.inner.k_aux()
251    }
252
253    fn dnum(&self) -> crate::layouts::Dnum {
254        self.inner.dnum()
255    }
256
257    fn dsize(&self) -> crate::layouts::Dsize {
258        self.inner.dsize()
259    }
260}
261
262impl<'a, BE: Backend + 'a> GGSWInfos for GGSWPreparedViewMut<'a, BE> {
263    fn k_aux(&self) -> crate::layouts::TorusPrecision {
264        self.inner.k_aux()
265    }
266
267    fn dnum(&self) -> crate::layouts::Dnum {
268        self.inner.dnum()
269    }
270
271    fn dsize(&self) -> crate::layouts::Dsize {
272        self.inner.dsize()
273    }
274}
275
276impl<'a, BE: Backend + 'a> LWEToBackendRef<BE> for LWEViewMut<'a, BE> {
277    fn to_backend_ref(&self) -> LWEBackendRef<'_, BE> {
278        LWE {
279            base2k: self.inner.base2k,
280            k: self.inner.k,
281            body: vec_znx_backend_ref_from_mut::<BE>(&self.inner.body),
282            mask: vec_znx_backend_ref_from_mut::<BE>(&self.inner.mask),
283        }
284    }
285}
286
287impl<'a, BE: Backend + 'a> LWEToBackendMut<BE> for LWEViewMut<'a, BE> {
288    fn to_backend_mut(&mut self) -> LWEBackendMut<'_, BE> {
289        let base2k = self.inner.base2k;
290        let k = self.inner.k;
291        let body = vec_znx_backend_mut_from_mut::<BE>(&mut self.inner.body);
292        let mask = vec_znx_backend_mut_from_mut::<BE>(&mut self.inner.mask);
293        LWE { base2k, k, body, mask }
294    }
295}
296
297impl<'a, BE: Backend + 'a> LWEPlaintextToBackendRef<BE> for LWEPlaintextViewMut<'a, BE> {
298    fn to_backend_ref(&self) -> LWEPlaintextBackendRef<'_, BE> {
299        LWEPlaintext {
300            base2k: self.inner.base2k,
301            k: self.inner.k,
302            data: vec_znx_backend_ref_from_mut::<BE>(&self.inner.data),
303        }
304    }
305}
306
307impl<'a, BE: Backend + 'a> LWEPlaintextToBackendMut<BE> for LWEPlaintextViewMut<'a, BE> {
308    fn to_backend_mut(&mut self) -> LWEPlaintextBackendMut<'_, BE> {
309        LWEPlaintext {
310            base2k: self.inner.base2k,
311            k: self.inner.k,
312            data: vec_znx_backend_mut_from_mut::<BE>(&mut self.inner.data),
313        }
314    }
315}
316
317macro_rules! impl_glwe_to_backend {
318    ($name:ident) => {
319        impl<'a, BE: Backend + 'a> GLWEToBackendRef<BE> for $name<'a, BE> {
320            fn to_backend_ref(&self) -> GLWEBackendRef<'_, BE> {
321                GLWE {
322                    base2k: self.inner.base2k,
323                    k: self.inner.k,
324                    data: vec_znx_backend_ref_from_mut::<BE>(&self.inner.data),
325                }
326            }
327        }
328
329        impl<'a, BE: Backend + 'a> GLWEToBackendMut<BE> for $name<'a, BE> {
330            fn to_backend_mut(&mut self) -> GLWEBackendMut<'_, BE> {
331                GLWE {
332                    base2k: self.inner.base2k,
333                    k: self.inner.k,
334                    data: vec_znx_backend_mut_from_mut::<BE>(&mut self.inner.data),
335                }
336            }
337        }
338    };
339}
340
341impl_glwe_to_backend!(GLWEViewMut);
342impl_glwe_to_backend!(GLWEPlaintextViewMut);
343impl_glwe_to_backend!(GLWETensorViewMut);
344
345impl<'a, BE: Backend + 'a> GLWEToBackendRef<BE> for GLWEViewRef<'a, BE> {
346    fn to_backend_ref(&self) -> GLWEBackendRef<'_, BE> {
347        GLWE {
348            base2k: self.inner.base2k,
349            k: self.inner.k,
350            data: vec_znx_backend_ref_from_ref::<BE>(&self.inner.data),
351        }
352    }
353}
354
355impl<'a, BE: Backend + 'a> GLWESecretToBackendRef<BE> for GLWESecretViewMut<'a, BE> {
356    fn to_backend_ref(&self) -> GLWESecretBackendRef<'_, BE> {
357        GLWESecret {
358            dist: self.inner.dist,
359            data: ScalarZnx::from_data(
360                BE::view_ref_mut(&self.inner.data.data),
361                self.inner.data.n(),
362                self.inner.data.cols(),
363            ),
364        }
365    }
366}
367
368impl<'a, BE: Backend + 'a> GLWESecretToBackendMut<BE> for GLWESecretViewMut<'a, BE> {
369    fn to_backend_mut(&mut self) -> GLWESecretBackendMut<'_, BE> {
370        let n = self.inner.data.n();
371        let cols = self.inner.data.cols();
372        GLWESecret {
373            dist: self.inner.dist,
374            data: ScalarZnx::from_data(BE::view_mut_ref(&mut self.inner.data.data), n, cols),
375        }
376    }
377}
378
379impl<'a, BE: Backend + 'a> GLWESecretTensorToBackendRef<BE> for GLWESecretTensorViewMut<'a, BE> {
380    fn to_backend_ref(&self) -> GLWESecretTensorBackendRef<'_, BE> {
381        GLWESecretTensor {
382            dist: self.inner.dist,
383            rank: self.inner.rank,
384            data: ScalarZnx::from_data(
385                BE::view_ref_mut(&self.inner.data.data),
386                self.inner.data.n(),
387                self.inner.data.cols(),
388            ),
389        }
390    }
391}
392
393impl<'a, BE: Backend + 'a> GLWESecretTensorToBackendMut<BE> for GLWESecretTensorViewMut<'a, BE> {
394    fn to_backend_mut(&mut self) -> GLWESecretTensorBackendMut<'_, BE> {
395        let n = self.inner.data.n();
396        let cols = self.inner.data.cols();
397        GLWESecretTensor {
398            dist: self.inner.dist,
399            rank: self.inner.rank,
400            data: ScalarZnx::from_data(BE::view_mut_ref(&mut self.inner.data.data), n, cols),
401        }
402    }
403}
404
405impl<'a, BE: Backend + 'a> GLWESecretPreparedToBackendRef<BE> for GLWESecretPreparedViewMut<'a, BE> {
406    fn to_backend_ref(&self) -> GLWESecretPreparedBackendRef<'_, BE> {
407        GLWESecretPrepared {
408            dist: self.inner.dist,
409            data: self.inner.data.reborrow_backend_ref(),
410        }
411    }
412}
413
414impl<'a, BE: Backend + 'a> GLWESecretPreparedToBackendMut<BE> for GLWESecretPreparedViewMut<'a, BE> {
415    fn to_backend_mut(&mut self) -> GLWESecretPreparedBackendMut<'_, BE> {
416        GLWESecretPrepared {
417            dist: self.inner.dist,
418            data: self.inner.data.reborrow_backend_mut(),
419        }
420    }
421}
422
423impl<'a, BE: Backend + 'a> GGLWEToBackendRef<BE> for GGLWEViewMut<'a, BE> {
424    fn to_backend_ref(&self) -> GGLWEBackendRef<'_, BE> {
425        GGLWEBackendRef::from_inner(GGLWE {
426            base2k: self.inner.base2k,
427            k_aux: self.inner.k_aux,
428            dsize: self.inner.dsize,
429            data: mat_znx_backend_ref_from_mut::<BE>(&self.inner.data),
430        })
431    }
432}
433
434impl<'a, BE: Backend + 'a> GGLWEToBackendMut<BE> for GGLWEViewMut<'a, BE> {
435    fn to_backend_mut(&mut self) -> GGLWEBackendMut<'_, BE> {
436        GGLWEBackendMut::from_inner(GGLWE {
437            base2k: self.inner.base2k,
438            k_aux: self.inner.k_aux,
439            dsize: self.inner.dsize,
440            data: mat_znx_backend_mut_from_mut::<BE>(&mut self.inner.data),
441        })
442    }
443}
444
445impl<'a, BE: Backend + 'a> GGLWEPreparedToBackendRef<BE> for GGLWEPreparedViewMut<'a, BE> {
446    fn to_backend_ref(&self) -> GGLWEPreparedBackendRef<'_, BE> {
447        GGLWEPrepared {
448            base2k: self.inner.base2k,
449            k_aux: self.inner.k_aux,
450            dsize: self.inner.dsize,
451            dnum: self.inner.dnum,
452            stride: self.inner.stride,
453            data: self.inner.data.reborrow_backend_ref(),
454        }
455    }
456}
457
458impl<'a, BE: Backend + 'a> GGLWEPreparedToBackendMut<BE> for GGLWEPreparedViewMut<'a, BE> {
459    fn to_backend_mut(&mut self) -> GGLWEPreparedBackendMut<'_, BE> {
460        GGLWEPrepared {
461            base2k: self.inner.base2k,
462            k_aux: self.inner.k_aux,
463            dsize: self.inner.dsize,
464            dnum: self.inner.dnum,
465            stride: self.inner.stride,
466            data: self.inner.data.reborrow_backend_mut(),
467        }
468    }
469}
470
471impl<'a, BE: Backend + 'a> GGSWToBackendRef<BE> for GGSWViewMut<'a, BE> {
472    fn to_backend_ref(&self) -> GGSWBackendRef<'_, BE> {
473        GGSWBackendRef::from_inner(GGSW {
474            base2k: self.inner.base2k,
475            k_aux: self.inner.k_aux,
476            dsize: self.inner.dsize,
477            data: mat_znx_backend_ref_from_mut::<BE>(&self.inner.data),
478        })
479    }
480}
481
482impl<'a, BE: Backend + 'a> GGSWToBackendMut<BE> for GGSWViewMut<'a, BE> {
483    fn to_backend_mut(&mut self) -> GGSWBackendMut<'_, BE> {
484        GGSWBackendMut::from_inner(GGSW {
485            base2k: self.inner.base2k,
486            k_aux: self.inner.k_aux,
487            dsize: self.inner.dsize,
488            data: mat_znx_backend_mut_from_mut::<BE>(&mut self.inner.data),
489        })
490    }
491}
492
493impl<'a, BE: Backend + 'a> GGSWPreparedToBackendRef<BE> for GGSWPreparedViewMut<'a, BE> {
494    fn to_backend_ref(&self) -> GGSWPreparedBackendRef<'_, BE> {
495        GGSWPrepared {
496            base2k: self.inner.base2k,
497            k_aux: self.inner.k_aux,
498            dsize: self.inner.dsize,
499            data: self.inner.data.reborrow_backend_ref(),
500        }
501    }
502}
503
504impl<'a, BE: Backend + 'a> GGSWPreparedToBackendMut<BE> for GGSWPreparedViewMut<'a, BE> {
505    fn to_backend_mut(&mut self) -> GGSWPreparedBackendMut<'_, BE> {
506        GGSWPrepared {
507            base2k: self.inner.base2k,
508            k_aux: self.inner.k_aux,
509            dsize: self.inner.dsize,
510            data: self.inner.data.reborrow_backend_mut(),
511        }
512    }
513}