diff --git a/mindspore/core/utils/crypto.h b/mindspore/core/utils/crypto.h index 686c2adcc92..a1e388154b9 100644 --- a/mindspore/core/utils/crypto.h +++ b/mindspore/core/utils/crypto.h @@ -23,10 +23,10 @@ typedef unsigned char Byte; namespace mindspore { -constexpr size_t MAX_BLOCK_SIZE = 512 * 1024 * 1024; // Maximum ciphertext segment, units is Byte -constexpr size_t RESERVED_BYTE_PER_BLOCK = 50; // Reserved byte per block to save addition info -constexpr unsigned int GCM_MAGIC_NUM = 0x7F3A5ED8; // Magic number -constexpr unsigned int CBC_MAGIC_NUM = 0x7F3A5ED9; // Magic number +constexpr size_t MAX_BLOCK_SIZE = 64 * 1024 * 1024; // Maximum ciphertext segment, units is Byte +constexpr size_t RESERVED_BYTE_PER_BLOCK = 50; // Reserved byte per block to save addition info +constexpr unsigned int GCM_MAGIC_NUM = 0x7F3A5ED8; // Magic number +constexpr unsigned int CBC_MAGIC_NUM = 0x7F3A5ED9; // Magic number constexpr size_t Byte16 = 16; MS_CORE_API std::unique_ptr Encrypt(size_t *encrypt_len, const Byte *plain_data, size_t plain_len, diff --git a/mindspore/lite/src/common/decrypt.cc b/mindspore/lite/src/common/decrypt.cc index 466efe70712..b440f88d53e 100644 --- a/mindspore/lite/src/common/decrypt.cc +++ b/mindspore/lite/src/common/decrypt.cc @@ -43,9 +43,9 @@ std::unique_ptr Decrypt(const std::string &lib_path, size_t *, const Byt } #else namespace { -constexpr size_t MAX_BLOCK_SIZE = 512 * 1024 * 1024; // Maximum ciphertext segment, units is Byte -constexpr size_t Byte16 = 16; // Byte16 -constexpr unsigned int MAGIC_NUM = 0x7F3A5ED8; // Magic number +constexpr size_t MAX_BLOCK_SIZE = 64 * 1024 * 1024; // Maximum ciphertext segment, units is Byte +constexpr size_t Byte16 = 16; // Byte16 +constexpr unsigned int MAGIC_NUM = 0x7F3A5ED8; // Magic number DynamicLibraryLoader loader; } // namespace int32_t ByteToInt(const Byte *byteArray, size_t length) { @@ -253,7 +253,7 @@ std::unique_ptr Decrypt(const std::string &lib_path, size_t *decrypt_len } std::vector block_buf; std::vector int_buf(sizeof(int32_t)); - std::vector decrypt_block_buf(MAX_BLOCK_SIZE); + auto decrypt_data = std::make_unique(data_size); int32_t decrypt_block_len; @@ -303,13 +303,19 @@ std::unique_ptr Decrypt(const std::string &lib_path, size_t *decrypt_len } block_buf.assign(model_data + offset, model_data + offset + block_size); offset += block_buf.size(); - if (!(BlockDecrypt(decrypt_block_buf.data(), &decrypt_block_len, reinterpret_cast(block_buf.data()), - block_buf.size(), key, static_cast(key_len), dec_mode, tag))) { - MS_LOG(ERROR) << "Failed to decrypt data, please check if dec_key or dec_mode is valid"; + Byte *decrypt_block_buf = static_cast(malloc(MAX_BLOCK_SIZE * sizeof(Byte))); + if (decrypt_block_buf == nullptr) { + MS_LOG(ERROR) << "decrypt_block_buf is nullptr."; return nullptr; } - memcpy(decrypt_data.get() + *decrypt_len, decrypt_block_buf.data(), static_cast(decrypt_block_len)); - + if (!(BlockDecrypt(decrypt_block_buf, &decrypt_block_len, reinterpret_cast(block_buf.data()), + block_buf.size(), key, static_cast(key_len), dec_mode, tag))) { + MS_LOG(ERROR) << "Failed to decrypt data, please check if dec_key or dec_mode is valid"; + free(decrypt_block_buf); + return nullptr; + } + memcpy(decrypt_data.get() + *decrypt_len, decrypt_block_buf, static_cast(decrypt_block_len)); + free(decrypt_block_buf); *decrypt_len += static_cast(decrypt_block_len); } ret = loader.Close(); diff --git a/mindspore/python/mindspore/train/serialization.py b/mindspore/python/mindspore/train/serialization.py index a962b769c52..cede450c49f 100644 --- a/mindspore/python/mindspore/train/serialization.py +++ b/mindspore/python/mindspore/train/serialization.py @@ -69,6 +69,7 @@ SLICE_SIZE = 512 * 1024 PROTO_LIMIT_SIZE = 1024 * 1024 * 2 TOTAL_SAVE = 1024 * 1024 PARAMETER_SPLIT_SIZE = 1024 * 1024 * 1024 +ENCRYPT_BLOCK_SIZE = 64 * 1024 def _special_process_par(par, new_par): @@ -222,7 +223,7 @@ def _exec_save(ckpt_file_name, data_list, enc_key=None, enc_mode="AES-GCM"): else: plain_data += checkpoint_list.SerializeToString() - max_block_size = SLICE_SIZE * 1024 + max_block_size = ENCRYPT_BLOCK_SIZE * 1024 while len(plain_data) >= max_block_size: cipher_data += _encrypt(plain_data[0: max_block_size], max_block_size, enc_key, len(enc_key), enc_mode)