forked from huawei/mindspore2022
signds optimize
This commit is contained in:
parent
c77fc8b094
commit
640905e5fe
|
|
@ -210,7 +210,14 @@ ResultCode UpdateModelKernel::VerifyUpdateModel(const schema::RequestUpdateModel
|
|||
feature_map[weight_full_name] = weight_size;
|
||||
}
|
||||
|
||||
if (!LocalMetaStore::GetInstance().verifyAggregationFeatureMap(feature_map)) {
|
||||
bool verifyFeatureMapIsSuccess;
|
||||
if (ps::PSContext::instance()->encrypt_type() == ps::kDSEncryptType && update_model_req->sign() != 0) {
|
||||
MS_ERROR_IF_NULL_W_RET_VAL(update_model_req->index_array(), ResultCode::kSuccessAndReturn);
|
||||
verifyFeatureMapIsSuccess = VerifySignDSFeatureMap(feature_map, update_model_req);
|
||||
} else {
|
||||
verifyFeatureMapIsSuccess = LocalMetaStore::GetInstance().verifyAggregationFeatureMap(feature_map);
|
||||
}
|
||||
if (!verifyFeatureMapIsSuccess) {
|
||||
auto next_req_time = LocalMetaStore::GetInstance().value<uint64_t>(kCtxIterationNextRequestTimestamp);
|
||||
std::string reason = "Verify model feature map failed, retry later at time: " + std::to_string(next_req_time);
|
||||
BuildUpdateModelRsp(fbb, schema::ResponseCode_OutOfTime, reason, std::to_string(next_req_time));
|
||||
|
|
@ -249,13 +256,42 @@ ResultCode UpdateModelKernel::VerifyUpdateModel(const schema::RequestUpdateModel
|
|||
return ResultCode::kSuccess;
|
||||
}
|
||||
|
||||
bool UpdateModelKernel::VerifySignDSFeatureMap(const std::unordered_map<std::string, size_t> &model,
|
||||
const schema::RequestUpdateModel *update_model_req) {
|
||||
auto &aggregation_feature_map_ = LocalMetaStore::GetInstance().aggregation_feature_map();
|
||||
if (model.size() > aggregation_feature_map_.size()) {
|
||||
return false;
|
||||
}
|
||||
auto index_array = update_model_req->index_array();
|
||||
size_t index_array_size = index_array->size();
|
||||
size_t array_size_upper = 100;
|
||||
if (index_array_size == 0 || index_array_size > array_size_upper) {
|
||||
return false;
|
||||
}
|
||||
for (const auto &weight : model) {
|
||||
std::string weight_name = weight.first;
|
||||
if (aggregation_feature_map_.count(weight_name) == 0) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
ResultCode UpdateModelKernel::UpdateModel(const schema::RequestUpdateModel *update_model_req,
|
||||
const std::shared_ptr<FBBuilder> &fbb, const DeviceMeta &device_meta) {
|
||||
MS_ERROR_IF_NULL_W_RET_VAL(update_model_req, ResultCode::kSuccessAndReturn);
|
||||
MS_ERROR_IF_NULL_W_RET_VAL(update_model_req->fl_id(), ResultCode::kSuccessAndReturn);
|
||||
std::string update_model_fl_id = update_model_req->fl_id()->str();
|
||||
size_t data_size = device_meta.data_size();
|
||||
const auto &feature_map = ParseFeatureMap(update_model_req);
|
||||
|
||||
std::map<std::string, std::vector<float>> weight_map;
|
||||
std::map<std::string, UploadData> feature_map;
|
||||
if (ps::PSContext::instance()->encrypt_type() == ps::kDSEncryptType) {
|
||||
feature_map = ParseSignDSFeatureMap(update_model_req, data_size, &weight_map);
|
||||
} else {
|
||||
feature_map = ParseFeatureMap(update_model_req);
|
||||
}
|
||||
|
||||
if (feature_map.empty()) {
|
||||
std::string reason = "Feature map is empty.";
|
||||
BuildUpdateModelRsp(fbb, schema::ResponseCode_RequestError, reason, "");
|
||||
|
|
@ -313,6 +349,58 @@ std::map<std::string, UploadData> UpdateModelKernel::ParseFeatureMap(
|
|||
return feature_map;
|
||||
}
|
||||
|
||||
std::map<std::string, UploadData> UpdateModelKernel::ParseSignDSFeatureMap(
|
||||
const schema::RequestUpdateModel *update_model_req, size_t data_size,
|
||||
std::map<std::string, std::vector<float>> *weight_map) {
|
||||
auto fbs_feature_map = update_model_req->feature_map();
|
||||
std::map<std::string, UploadData> feature_map;
|
||||
auto sign = update_model_req->sign();
|
||||
if (sign == 0) {
|
||||
for (uint32_t i = 0; i < fbs_feature_map->size(); i++) {
|
||||
std::string weight_full_name = fbs_feature_map->Get(i)->weight_fullname()->str();
|
||||
float *weight_data = const_cast<float *>(fbs_feature_map->Get(i)->data()->data());
|
||||
size_t weight_size = fbs_feature_map->Get(i)->data()->size() * sizeof(float);
|
||||
UploadData upload_data;
|
||||
upload_data[kNewWeight].addr = weight_data;
|
||||
upload_data[kNewWeight].size = weight_size;
|
||||
feature_map[weight_full_name] = upload_data;
|
||||
}
|
||||
return feature_map;
|
||||
}
|
||||
|
||||
const auto &iter_to_model = ModelStore::GetInstance().iteration_to_model();
|
||||
size_t latest_iter_num = iter_to_model.rbegin()->first;
|
||||
std::map<std::string, AddressPtr> feature_maps_store = ModelStore::GetInstance().GetModelByIterNum(latest_iter_num);
|
||||
auto index_array = update_model_req->index_array();
|
||||
size_t index_store = 0;
|
||||
size_t index_array_j = 0;
|
||||
float signds_grad = sign * ps::PSContext::instance()->sign_global_lr();
|
||||
for (size_t i = 0; i < fbs_feature_map->size(); i++) {
|
||||
std::string weight_full_name = fbs_feature_map->Get(i)->weight_fullname()->str();
|
||||
AddressPtr iter_feature_map_data_ptr = feature_maps_store[weight_full_name];
|
||||
size_t iter_feature_num = iter_feature_map_data_ptr->size / sizeof(float);
|
||||
auto &weight_item = (*weight_map)[weight_full_name];
|
||||
weight_item.resize(iter_feature_num);
|
||||
float *iter_feature_map_data = reinterpret_cast<float *>(iter_feature_map_data_ptr->addr);
|
||||
for (size_t j = 0; j < iter_feature_num; j++) {
|
||||
float reconstruct_weight = iter_feature_map_data[j];
|
||||
if (index_array_j < index_array->size() && index_store == static_cast<size_t>(index_array->Get(index_array_j))) {
|
||||
reconstruct_weight += signds_grad;
|
||||
index_array_j++;
|
||||
}
|
||||
reconstruct_weight *= data_size;
|
||||
index_store++;
|
||||
weight_item[j] = reconstruct_weight;
|
||||
}
|
||||
size_t weight_size = iter_feature_num * sizeof(float);
|
||||
UploadData upload_data;
|
||||
upload_data[kNewWeight].addr = weight_item.data();
|
||||
upload_data[kNewWeight].size = weight_size;
|
||||
feature_map[weight_full_name] = upload_data;
|
||||
}
|
||||
return feature_map;
|
||||
}
|
||||
|
||||
ResultCode UpdateModelKernel::CountForAggregation(const std::string &req_fl_id) {
|
||||
std::string count_reason = "";
|
||||
if (!DistributedCountService::GetInstance().Count(kCountForAggregation, req_fl_id, &count_reason)) {
|
||||
|
|
|
|||
|
|
@ -18,6 +18,7 @@
|
|||
#define MINDSPORE_CCSRC_FL_SERVER_KERNEL_UPDATE_MODEL_KERNEL_H_
|
||||
|
||||
#include <map>
|
||||
#include <unordered_map>
|
||||
#include <memory>
|
||||
#include <string>
|
||||
#include <vector>
|
||||
|
|
@ -25,6 +26,7 @@
|
|||
#include "fl/server/kernel/round/round_kernel.h"
|
||||
#include "fl/server/kernel/round/round_kernel_factory.h"
|
||||
#include "fl/server/executor.h"
|
||||
#include "fl/server/model_store.h"
|
||||
#ifdef ENABLE_ARMOUR
|
||||
#include "fl/armour/cipher/cipher_meta_storage.h"
|
||||
#endif
|
||||
|
|
@ -58,6 +60,11 @@ class UpdateModelKernel : public RoundKernel {
|
|||
|
||||
void RunAggregation();
|
||||
ResultCode CountForAggregation(const std::string &req_fl_id);
|
||||
std::map<std::string, UploadData> ParseSignDSFeatureMap(const schema::RequestUpdateModel *update_model_req,
|
||||
size_t data_size,
|
||||
std::map<std::string, std::vector<float>> *weight_map);
|
||||
bool VerifySignDSFeatureMap(const std::unordered_map<std::string, size_t> &model,
|
||||
const schema::RequestUpdateModel *update_model_req);
|
||||
ResultCode CountForUpdateModel(const std::shared_ptr<FBBuilder> &fbb,
|
||||
const schema::RequestUpdateModel *update_model_req);
|
||||
sigVerifyResult VerifySignature(const schema::RequestUpdateModel *update_model_req);
|
||||
|
|
|
|||
|
|
@ -616,30 +616,57 @@ public class SecureProtocol {
|
|||
}
|
||||
|
||||
/**
|
||||
* SignDS model weights.
|
||||
* select num indexes from inputList, and put them into outputList.
|
||||
*
|
||||
* @param builder the FlatBufferBuilder object used for serialization model weights.
|
||||
* @param trainDataSize tne size of train data set.
|
||||
* @return the serialized model weights after adding masks.
|
||||
* @param secureRandom cryptographically strong random number generator.
|
||||
* @param inputList select index from inputList.
|
||||
* @param outputList put random index into outputList.
|
||||
* @param num the number of select indexes.
|
||||
*/
|
||||
private static void randomSelect(SecureRandom secureRandom, List<Integer> inputList, List<Integer> outputList, int num) {
|
||||
if (num <= 0) {
|
||||
LOGGER.severe(Common.addTag("[SignDS] The number to be selected is set incorrectly!"));
|
||||
return;
|
||||
}
|
||||
if (inputList.isEmpty()) {
|
||||
LOGGER.severe(Common.addTag("[SignDS] The input List is empty!"));
|
||||
return;
|
||||
}
|
||||
if (inputList.size() < num) {
|
||||
LOGGER.severe(Common.addTag("[SignDS] The size of inputList is small than num!"));
|
||||
return;
|
||||
}
|
||||
for (int i = inputList.size(); i > inputList.size() - num; i--) {
|
||||
int randomIndex = secureRandom.nextInt(i);
|
||||
int randomSelectTopkIndex = inputList.get(randomIndex);
|
||||
inputList.set(randomIndex, inputList.get(i - 1));
|
||||
inputList.set(i - 1, randomSelectTopkIndex);
|
||||
outputList.add(randomSelectTopkIndex);
|
||||
}
|
||||
}
|
||||
|
||||
public int[] signDSModel(FlatBufferBuilder builder, int trainDataSize, Map<String, float[]> trainedMap) {
|
||||
/**
|
||||
* SignDS alg.
|
||||
*
|
||||
* @param trainedMap trained model.
|
||||
* @param sign random sign value.
|
||||
* @return index list.
|
||||
*/
|
||||
public int[] signDSModel(Map<String, float[]> trainedMap, boolean sign) {
|
||||
Map<String, float[]> mapBeforeTrain = modelMap;
|
||||
int layerNum = updateFeatureName.size();
|
||||
int[] featuresMap = new int[layerNum];
|
||||
SecureRandom secureRandom = Common.getSecureRandom();
|
||||
boolean sign = secureRandom.nextBoolean();
|
||||
List<String> nonTopkKeyList = new ArrayList<>();
|
||||
List<String> topkKeyList = new ArrayList<>();
|
||||
Map<String, Float> allUpdateMap = new HashMap<>();
|
||||
List<Integer> nonTopkKeyList = new ArrayList<>();
|
||||
List<Integer> topkKeyList = new ArrayList<>();
|
||||
Map<Integer, Float> allUpdateMap = new HashMap<>();
|
||||
int index = 0;
|
||||
for (int i = 0; i < layerNum; i++) {
|
||||
String key = updateFeatureName.get(i);
|
||||
float[] dataAfterTrain = trainedMap.get(key);
|
||||
float[] dataBeforeTrain = mapBeforeTrain.get(key);
|
||||
for (int j = 0; j < dataAfterTrain.length; j++) {
|
||||
float updateData = dataAfterTrain[j] - dataBeforeTrain[j];
|
||||
String ij = Integer.toString(i) + ',' + j;
|
||||
allUpdateMap.put(ij, updateData);
|
||||
allUpdateMap.put(index++, updateData);
|
||||
}
|
||||
}
|
||||
int inputDim = allUpdateMap.size();
|
||||
|
|
@ -667,8 +694,7 @@ public class SecureProtocol {
|
|||
LOGGER.severe("[SignDS] topkDim or signDimOut is ERROR! please check");
|
||||
return new int[0];
|
||||
}
|
||||
|
||||
List<Map.Entry<String, Float>> allUpdateList = new ArrayList<>(allUpdateMap.entrySet());
|
||||
List<Map.Entry<Integer, Float>> allUpdateList = new ArrayList<>(allUpdateMap.entrySet());
|
||||
if (sign) {
|
||||
allUpdateList.sort((o1, o2) -> Float.compare(o2.getValue(), o1.getValue()));
|
||||
} else {
|
||||
|
|
@ -680,42 +706,11 @@ public class SecureProtocol {
|
|||
for (int i = topkDim; i < allUpdateList.size(); i++) {
|
||||
nonTopkKeyList.add(allUpdateList.get(i).getKey());
|
||||
}
|
||||
List<String> outputDimensionIJStringList = new ArrayList<>();
|
||||
for (int i = topkKeyList.size(); i > topkKeyList.size() - numInter; i--) {
|
||||
int randomIndex = secureRandom.nextInt(i);
|
||||
String randomChoiceTopkIJString = topkKeyList.get(randomIndex);
|
||||
topkKeyList.set(randomIndex, topkKeyList.get(i - 1));
|
||||
topkKeyList.set(i - 1, randomChoiceTopkIJString);
|
||||
outputDimensionIJStringList.add(randomChoiceTopkIJString);
|
||||
}
|
||||
for (int i = nonTopkKeyList.size(); i > nonTopkKeyList.size() - numOuter; i--) {
|
||||
int randomIndex = secureRandom.nextInt(i);
|
||||
String randomChoiceNonTopkIJString = nonTopkKeyList.get(randomIndex);
|
||||
nonTopkKeyList.set(randomIndex, nonTopkKeyList.get(i - 1));
|
||||
nonTopkKeyList.set(i - 1, randomChoiceNonTopkIJString);
|
||||
outputDimensionIJStringList.add(randomChoiceNonTopkIJString);
|
||||
}
|
||||
float signValue = sign ? 1f * signGlobalLr : -1f * signGlobalLr;
|
||||
for (String ijString : outputDimensionIJStringList) {
|
||||
String[] ij = ijString.split(",");
|
||||
int iKeyIndex = Integer.parseInt(ij[0]);
|
||||
int jDataIndex = Integer.parseInt(ij[1]);
|
||||
String key = updateFeatureName.get(iKeyIndex);
|
||||
float[] dataBeforeTrain = mapBeforeTrain.get(key);
|
||||
dataBeforeTrain[jDataIndex] += signValue;
|
||||
mapBeforeTrain.put(key, dataBeforeTrain);
|
||||
}
|
||||
for (int i = 0; i < layerNum; i++) {
|
||||
String key = updateFeatureName.get(i);
|
||||
float[] dataBeforeTrain = mapBeforeTrain.get(key);
|
||||
for (int j = 0; j < dataBeforeTrain.length; j++) {
|
||||
dataBeforeTrain[j] *= trainDataSize;
|
||||
}
|
||||
int featureName = builder.createString(key);
|
||||
int weight = FeatureMap.createDataVector(builder, dataBeforeTrain);
|
||||
int featureMap = FeatureMap.createFeatureMap(builder, featureName, weight);
|
||||
featuresMap[i] = featureMap;
|
||||
}
|
||||
return featuresMap;
|
||||
List<Integer> outputDimensionIndexList = new ArrayList<>();
|
||||
randomSelect(secureRandom, topkKeyList, outputDimensionIndexList, numInter);
|
||||
randomSelect(secureRandom, nonTopkKeyList, outputDimensionIndexList, numOuter);
|
||||
outputDimensionIndexList.sort(Integer::compare);
|
||||
LOGGER.info(Common.addTag("[SignDS] outputDimension size is " + outputDimensionIndexList.size()));
|
||||
return outputDimensionIndexList.stream().mapToInt(i -> i).toArray();
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -38,6 +38,7 @@ import java.util.HashMap;
|
|||
import java.util.List;
|
||||
import java.util.Map;
|
||||
import java.util.logging.Logger;
|
||||
import java.security.SecureRandom;
|
||||
|
||||
import static com.mindspore.flclient.LocalFLParameter.ALBERT;
|
||||
import static com.mindspore.flclient.LocalFLParameter.LENET;
|
||||
|
|
@ -211,6 +212,8 @@ public class UpdateModel {
|
|||
private int idOffset = 0;
|
||||
private int timestampOffset = 0;
|
||||
private int signDataOffset = 0;
|
||||
private int sign = 0;
|
||||
private int indexArrayOffset = 0;
|
||||
private int iteration = 0;
|
||||
private EncryptLevel encryptLevel = EncryptLevel.NOT_ENCRYPT;
|
||||
private float uploadLossOffset = 0.0f;
|
||||
|
|
@ -291,6 +294,8 @@ public class UpdateModel {
|
|||
} else {
|
||||
trainedMap = getFeatureMap();
|
||||
}
|
||||
long startTime;
|
||||
long endTime;
|
||||
switch (encryptLevel) {
|
||||
case PW_ENCRYPT:
|
||||
int[] fmOffsetsPW = secureProtocol.pwMaskModel(builder, trainDataSize, trainedMap);
|
||||
|
|
@ -303,6 +308,7 @@ public class UpdateModel {
|
|||
LOGGER.info(Common.addTag("[Encrypt] pairwise mask model ok!"));
|
||||
return this;
|
||||
case DP_ENCRYPT:
|
||||
startTime = System.currentTimeMillis();
|
||||
int[] fmOffsetsDP = secureProtocol.dpMaskModel(builder, trainDataSize, trainedMap);
|
||||
if (fmOffsetsDP == null || fmOffsetsDP.length == 0) {
|
||||
LOGGER.severe("[Encrypt] the return fmOffsetsDP from <secureProtocol.dpMaskModel> is " +
|
||||
|
|
@ -313,21 +319,44 @@ public class UpdateModel {
|
|||
}
|
||||
this.fmOffset = RequestUpdateModel.createFeatureMapVector(builder, fmOffsetsDP);
|
||||
LOGGER.info(Common.addTag("[Encrypt] DP mask model ok!"));
|
||||
endTime = System.currentTimeMillis();
|
||||
LOGGER.info(Common.addTag("dp time is:" + (endTime - startTime) + "ms"));
|
||||
return this;
|
||||
case SIGNDS:
|
||||
int[] fmOffsetsSignDS = secureProtocol.signDSModel(builder, trainDataSize, trainedMap);
|
||||
if (fmOffsetsSignDS == null || fmOffsetsSignDS.length == 0) {
|
||||
startTime = System.currentTimeMillis();
|
||||
// signds alg return indexArray, and package indexArray into flatbuffer.
|
||||
SecureRandom secureRandom = Common.getSecureRandom();
|
||||
boolean signBool = secureRandom.nextBoolean();
|
||||
this.sign = signBool ? 1 : -1;
|
||||
int[] indexArray = secureProtocol.signDSModel(trainedMap, signBool);
|
||||
if (indexArray == null || indexArray.length == 0) {
|
||||
LOGGER.severe("[Encrypt] the return fmOffsetsSignDS from <secureProtocol.signDSModel> is " +
|
||||
"null, please check");
|
||||
retCode = ResponseCode.RequestError;
|
||||
status = FLClientStatus.FAILED;
|
||||
throw new IllegalArgumentException();
|
||||
}
|
||||
this.fmOffset = RequestUpdateModel.createFeatureMapVector(builder, fmOffsetsSignDS);
|
||||
this.indexArrayOffset = RequestUpdateModel.createIndexArrayVector(builder, indexArray);
|
||||
|
||||
// only package featureName into flatbuffer.
|
||||
int compFeatureSize = updateFeatureName.size();
|
||||
int[] fmOffsetsSignds = new int[compFeatureSize];
|
||||
for (int i = 0; i < compFeatureSize; i++) {
|
||||
String key = updateFeatureName.get(i);
|
||||
float[] data = new float[0];
|
||||
int featureName = builder.createString(key);
|
||||
int weight = FeatureMap.createDataVector(builder, data);
|
||||
int featureMap = FeatureMap.createFeatureMap(builder, featureName, weight);
|
||||
fmOffsetsSignds[i] = featureMap;
|
||||
}
|
||||
this.fmOffset = RequestUpdateModel.createFeatureMapVector(builder, fmOffsetsSignds);
|
||||
LOGGER.info(Common.addTag("[Encrypt] SignDS mask model ok!"));
|
||||
endTime = System.currentTimeMillis();
|
||||
LOGGER.info(Common.addTag("signds time is:" + (endTime - startTime) + "ms"));
|
||||
return this;
|
||||
case NOT_ENCRYPT:
|
||||
default:
|
||||
startTime = System.currentTimeMillis();
|
||||
int featureSize = updateFeatureName.size();
|
||||
int[] fmOffsets = new int[featureSize];
|
||||
for (int i = 0; i < featureSize; i++) {
|
||||
|
|
@ -344,6 +373,8 @@ public class UpdateModel {
|
|||
fmOffsets[i] = featureMap;
|
||||
}
|
||||
this.fmOffset = RequestUpdateModel.createFeatureMapVector(builder, fmOffsets);
|
||||
endTime = System.currentTimeMillis();
|
||||
LOGGER.info(Common.addTag("not encrypt time is:" + (endTime - startTime) + "ms"));
|
||||
return this;
|
||||
}
|
||||
}
|
||||
|
|
@ -389,6 +420,8 @@ public class UpdateModel {
|
|||
RequestUpdateModel.addFeatureMap(builder, this.fmOffset);
|
||||
RequestUpdateModel.addSignature(builder, this.signDataOffset);
|
||||
RequestUpdateModel.addUploadLoss(builder, this.uploadLossOffset);
|
||||
RequestUpdateModel.addSign(builder, this.sign);
|
||||
RequestUpdateModel.addIndexArray(builder, this.indexArrayOffset);
|
||||
int root = RequestUpdateModel.endRequestUpdateModel(builder);
|
||||
builder.finish(root);
|
||||
return builder.sizedByteArray();
|
||||
|
|
|
|||
|
|
@ -92,6 +92,8 @@ table RequestUpdateModel{
|
|||
timestamp:string;
|
||||
signature:[ubyte];
|
||||
upload_loss:float;
|
||||
sign:int;
|
||||
index_array:[int];
|
||||
}
|
||||
|
||||
table ResponseUpdateModel{
|
||||
|
|
|
|||
Loading…
Reference in New Issue