Fix an issue of federated learning

This commit is contained in:
jin-xiulang 2022-01-10 16:40:03 +08:00
parent f69699f1ec
commit a340f45a9d
2 changed files with 2 additions and 54 deletions

View File

@ -441,53 +441,6 @@ public class FLLiteClient {
return featureMap;
}
/**
* Obtain the weight of the model before training.
*
* @param map a map to store the weight of the model.
* @return the weight.
*/
public static synchronized Map<String, float[]> getOldMapCopy(Map<String, float[]> map) {
if (mapBeforeTrain == null) {
Map<String, float[]> copyMap = new TreeMap<>();
for (String key : map.keySet()) {
float[] data = map.get(key);
int dataLen = data.length;
float[] weights = new float[dataLen];
if ((key.indexOf("Default") < 0) && (key.indexOf("nhwc") < 0) && (key.indexOf("moment") < 0) &&
(key.indexOf("learning") < 0)) {
for (int j = 0; j < dataLen; j++) {
float weight = data[j];
weights[j] = weight;
}
copyMap.put(key, weights);
}
}
mapBeforeTrain = copyMap;
} else {
for (String key : map.keySet()) {
float[] data = map.get(key);
float[] copyData = mapBeforeTrain.get(key);
int dataLen = data.length;
if ((key.indexOf("Default") < 0) && (key.indexOf("nhwc") < 0) && (key.indexOf("moment") < 0) &&
(key.indexOf("learning") < 0)) {
for (int j = 0; j < dataLen; j++) {
copyData[j] = data[j];
}
}
}
}
return mapBeforeTrain;
}
private void getOldFeatureMap() {
EncryptLevel encryptLevel = localFLParameter.getEncryptLevel();
if (encryptLevel == EncryptLevel.DP_ENCRYPT) {
Map<String, float[]> featureMap = getFeatureMap();
oldFeatureMap = getOldMapCopy(featureMap);
}
}
public void updateDpNormClip() {
EncryptLevel encryptLevel = localFLParameter.getEncryptLevel();
if (encryptLevel == EncryptLevel.DP_ENCRYPT) {
@ -554,12 +507,7 @@ public class FLLiteClient {
return curStatus;
case DP_ENCRYPT:
// get the feature map before train
getOldFeatureMap();
if (oldFeatureMap.isEmpty()) {
LOGGER.severe(Common.addTag("[Encrypt] the return map in getOldFeatureMapis empty "));
retCode = ResponseCode.RequestError;
return FLClientStatus.FAILED;
}
oldFeatureMap = getFeatureMap();
curStatus = secureProtocol.setDPParameter(iteration, dpEps, dpDelta, dpNormClipAdapt, oldFeatureMap);
retCode = ResponseCode.SUCCEED;
if (curStatus != FLClientStatus.SUCCESS) {

View File

@ -154,7 +154,7 @@ public class CertVerify {
LOGGER.severe(Common.addTag("[verifyChain] catch Exception: " + e.getMessage()));
return false;
}
LOGGER.severe(Common.addTag("verifyChain success!"));
LOGGER.info(Common.addTag("verifyChain success!"));
return true;
}