tfhe/integer/server_key/radix_parallel/
modulus_switch_compression.rs1use 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 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 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 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 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}