#pragma once #include #include "core/DeviceVector.h" #include "core/ElementWise.h" #include "core/NTT.h" #include "core/Parameter.h" namespace cheddar { template class ModSwitchHandler { private: using Dv = DeviceVector; // The order mattters here const int level_; const int num_aux_; const int beta_; const Parameter ¶m_; const ElementWiseHandler &elem_handler_; const NTTHandler &ntt_handler_; public: ModSwitchHandler(const Parameter ¶m, int level, const ElementWiseHandler &elem_handler, const NTTHandler &ntt_handler); // diable copying (or moving also) ModSwitchHandler(const ModSwitchHandler &) = delete; ModSwitchHandler &operator=(const ModSwitchHandler &) = delete; // for forwarding purposes ModSwitchHandler(ModSwitchHandler &&) = default; void PseudoModUp(DvView &dst, const DvConstView &src, const DvConstView &p_prod) const; void ModUp(std::vector> &dst, const DvConstView &src) const; void ModDown(DvView &dst, const DvConstView &src) const; void Rescale(DvView &dst, const DvConstView &src) const; void ModDownAndRescale(DvView &dst, const DvConstView &src) const; private: // ModUp constants Dv mod_up1_; std::vector>> mod_up2_; // ModDown constants Dv mod_down1_; DeviceVector> mod_down2_; Dv inv_prime_prod_; // Rescale constants int rescale_pad_start_; int rescale_pad_end_; int rescale_restore_start_; int rescale_restore_end_; Dv rescale1_; DeviceVector> rescale2_; Dv rescale_inv_prime_prod_; Dv rescale_padding_; // ModDownAndRescale constants Dv mod_down_rescale1_; DeviceVector> mod_down_rescale2_; Dv mod_down_rescale_inv_prime_prod_; Dv mod_down_rescale_padding_; Dv entire_padding_; // heuristic CUDA kernel block number; static constexpr int block_dim_ = 256; static inline bool cm_populated_ = false; void PopulateModSwitchConstants(Dv &const1, DeviceVector> &const2, const std::vector &src_primes, const std::vector &dst_primes, int restore_start, int restore_end); void PopulateModDownEpilogueConstants(Dv &inv_p_prod, Dv &padding, const std::vector &src_primes, const std::vector &dst_primes, int restore_start, int restore_end); enum class ModDownType { ModDown, Rescale, ModDownAndRescale }; void ModDownWorker(DvView &dst, const DvConstView &src, ModDownType type) const; }; } // namespace cheddar