Skip to main content

tfhe/integer/server_key/radix_parallel/
modulus_switch_compression.rs

1use crate::integer::ciphertext::{
2    BaseRadixCiphertext, BaseSignedRadixCiphertext, CompressedModulusSwitchedRadixCiphertext,
3    CompressedModulusSwitchedRadixCiphertextGeneric,
4    CompressedModulusSwitchedSignedRadixCiphertext,
5};
6use crate::integer::{RadixCiphertext, ServerKey, SignedRadixCiphertext};
7use crate::shortint::Ciphertext;
8use rayon::prelude::*;
9
10impl ServerKey {
11    /// Compresses a ciphertext to have a smaller serialization size
12    ///
13    /// See [`CompressedModulusSwitchedRadixCiphertext#example`] for usage
14    pub fn switch_modulus_and_compress_parallelized(
15        &self,
16        ct: &RadixCiphertext,
17    ) -> CompressedModulusSwitchedRadixCiphertext {
18        CompressedModulusSwitchedRadixCiphertext(
19            self.switch_modulus_and_compress_generic_parallelized(&ct.blocks),
20        )
21    }
22    /// Decompresses a compressed ciphertext
23    /// This operation costs a PBS
24    ///
25    /// See [`CompressedModulusSwitchedRadixCiphertext#example`] for usage
26    pub fn decompress_parallelized(
27        &self,
28        compressed_ct: &CompressedModulusSwitchedRadixCiphertext,
29    ) -> RadixCiphertext {
30        BaseRadixCiphertext {
31            blocks: self.decompress_generic_parallelized(&compressed_ct.0),
32        }
33    }
34
35    /// Compresses a signed ciphertext to have a smaller serialization size
36    ///
37    /// See [`CompressedModulusSwitchedSignedRadixCiphertext#example`] for usage
38    pub fn switch_modulus_and_compress_signed_parallelized(
39        &self,
40        ct: &SignedRadixCiphertext,
41    ) -> CompressedModulusSwitchedSignedRadixCiphertext {
42        CompressedModulusSwitchedSignedRadixCiphertext(
43            self.switch_modulus_and_compress_generic_parallelized(&ct.blocks),
44        )
45    }
46    /// Decompresses a signed compressed ciphertext
47    /// This operation costs a PBS
48    ///
49    /// See [`CompressedModulusSwitchedSignedRadixCiphertext#example`] for usage
50    pub fn decompress_signed_parallelized(
51        &self,
52        compressed_ct: &CompressedModulusSwitchedSignedRadixCiphertext,
53    ) -> SignedRadixCiphertext {
54        BaseSignedRadixCiphertext {
55            blocks: self.decompress_generic_parallelized(&compressed_ct.0),
56        }
57    }
58
59    #[allow(clippy::int_plus_one)]
60    fn switch_modulus_and_compress_generic_parallelized(
61        &self,
62        blocks: &[Ciphertext],
63    ) -> CompressedModulusSwitchedRadixCiphertextGeneric {
64        assert!(
65            self.message_modulus().0 <= self.carry_modulus().0,
66            "Compression does not support message_modulus > carry_modulus"
67        );
68        assert!(
69            self.key.max_noise_level.get() >= self.message_modulus().0 + 1,
70            "Compression does not support max_noise_level < message_modulus + 1"
71        );
72
73        let len = blocks.len();
74
75        let (paired_blocks, last_block) = if len.is_multiple_of(2) {
76            (blocks, None)
77        } else {
78            (&blocks[..len - 1], Some(blocks.last().unwrap()))
79        };
80
81        let paired_blocks = paired_blocks
82            .par_chunks_exact(2)
83            .map(|pair| {
84                let mut packed = pair[0].clone();
85
86                let scaled = self
87                    .key
88                    .unchecked_scalar_mul(&pair[1], self.message_modulus().0 as u8);
89
90                self.key.unchecked_add_assign(&mut packed, &scaled);
91
92                self.key.switch_modulus_and_compress(&packed)
93            })
94            .collect();
95
96        let last_block = last_block.map(|a| self.key.switch_modulus_and_compress(a));
97
98        CompressedModulusSwitchedRadixCiphertextGeneric {
99            paired_blocks,
100            last_block,
101        }
102    }
103
104    fn decompress_generic_parallelized(
105        &self,
106        compressed_ct: &CompressedModulusSwitchedRadixCiphertextGeneric,
107    ) -> Vec<Ciphertext> {
108        let message_extract = self
109            .key
110            .generate_lookup_table(|x| x % self.message_modulus().0);
111
112        let carry_extract = self
113            .key
114            .generate_lookup_table(|x| x / self.message_modulus().0);
115
116        let mut blocks: Vec<Ciphertext> = compressed_ct
117            .paired_blocks
118            .par_iter()
119            .flat_map(|a| {
120                [
121                    self.key
122                        .decompress_and_apply_lookup_table(a, &message_extract),
123                    self.key
124                        .decompress_and_apply_lookup_table(a, &carry_extract),
125                ]
126            })
127            .collect();
128
129        if let Some(last_block) = compressed_ct.last_block.as_ref() {
130            blocks.push(
131                self.key
132                    .decompress_and_apply_lookup_table(last_block, &message_extract),
133            );
134        }
135
136        blocks
137    }
138}