本文目录导读:

我将为您提供一个完整的Java联邦学习案例,涵盖核心框架和具体实现。
联邦学习框架设计
基础模型接口
// ModelInterface.java
package com.federated.base;
import java.io.Serializable;
import java.util.Map;
public interface ModelInterface extends Serializable {
Map<String, double[]> getParameters();
void setParameters(Map<String, double[]> parameters);
double[] predict(double[][] features);
double train(double[][] features, double[] labels, int epochs);
}
神经网络模型实现
// NeuralNetwork.java
package com.federated.model;
import com.federated.base.ModelInterface;
import java.util.*;
import java.util.Random;
public class NeuralNetwork implements ModelInterface {
private static final long serialVersionUID = 1L;
private int inputSize;
private int hiddenSize;
private int outputSize;
private double[][] weightsInputHidden;
private double[][] weightsHiddenOutput;
private double[] biasHidden;
private double[] biasOutput;
private double learningRate = 0.01;
private Random random = new Random(42);
public NeuralNetwork(int inputSize, int hiddenSize, int outputSize) {
this.inputSize = inputSize;
this.hiddenSize = hiddenSize;
this.outputSize = outputSize;
initializeWeights();
}
private void initializeWeights() {
weightsInputHidden = new double[inputSize][hiddenSize];
weightsHiddenOutput = new double[hiddenSize][outputSize];
biasHidden = new double[hiddenSize];
biasOutput = new double[outputSize];
// Xavier初始化
double limitIH = Math.sqrt(6.0 / (inputSize + hiddenSize));
double limitHO = Math.sqrt(6.0 / (hiddenSize + outputSize));
for (int i = 0; i < inputSize; i++) {
for (int j = 0; j < hiddenSize; j++) {
weightsInputHidden[i][j] = random.nextDouble() * 2 * limitIH - limitIH;
}
}
for (int i = 0; i < hiddenSize; i++) {
for (int j = 0; j < outputSize; j++) {
weightsHiddenOutput[i][j] = random.nextDouble() * 2 * limitHO - limitHO;
}
}
}
private double sigmoid(double x) {
return 1.0 / (1.0 + Math.exp(-x));
}
private double sigmoidDerivative(double x) {
return x * (1 - x);
}
@Override
public double[] predict(double[][] features) {
double[] result = new double[features.length];
for (int i = 0; i < features.length; i++) {
double[] hidden = new double[hiddenSize];
// 前向传播到隐藏层
for (int j = 0; j < hiddenSize; j++) {
double sum = biasHidden[j];
for (int k = 0; k < inputSize; k++) {
sum += features[i][k] * weightsInputHidden[k][j];
}
hidden[j] = sigmoid(sum);
}
// 前向传播到输出层
double output = biasOutput[0];
for (int j = 0; j < hiddenSize; j++) {
output += hidden[j] * weightsHiddenOutput[j][0];
}
result[i] = sigmoid(output);
}
return result;
}
@Override
public double train(double[][] features, double[] labels, int epochs) {
double totalLoss = 0.0;
for (int epoch = 0; epoch < epochs; epoch++) {
totalLoss = 0.0;
for (int sample = 0; sample < features.length; sample++) {
// 前向传播
double[] hidden = new double[hiddenSize];
for (int j = 0; j < hiddenSize; j++) {
double sum = biasHidden[j];
for (int k = 0; k < inputSize; k++) {
sum += features[sample][k] * weightsInputHidden[k][j];
}
hidden[j] = sigmoid(sum);
}
double output = biasOutput[0];
for (int j = 0; j < hiddenSize; j++) {
output += hidden[j] * weightsHiddenOutput[j][0];
}
output = sigmoid(output);
// 计算误差
double error = labels[sample] - output;
totalLoss += error * error;
// 反向传播
double outputDelta = error * sigmoidDerivative(output);
// 更新输出层权重
for (int j = 0; j < hiddenSize; j++) {
weightsHiddenOutput[j][0] += learningRate * outputDelta * hidden[j];
}
biasOutput[0] += learningRate * outputDelta;
// 更新隐藏层权重
double[] hiddenErrors = new double[hiddenSize];
for (int j = 0; j < hiddenSize; j++) {
hiddenErrors[j] = weightsHiddenOutput[j][0] * outputDelta;
}
for (int j = 0; j < hiddenSize; j++) {
double hiddenDelta = hiddenErrors[j] * sigmoidDerivative(hidden[j]);
for (int k = 0; k < inputSize; k++) {
weightsInputHidden[k][j] += learningRate * hiddenDelta * features[sample][k];
}
biasHidden[j] += learningRate * hiddenDelta;
}
}
totalLoss /= features.length;
}
return totalLoss;
}
@Override
public Map<String, double[]> getParameters() {
Map<String, double[]> params = new HashMap<>();
// 将权重展平为一维数组
double[] flatWIH = new double[inputSize * hiddenSize];
for (int i = 0; i < inputSize; i++) {
for (int j = 0; j < hiddenSize; j++) {
flatWIH[i * hiddenSize + j] = weightsInputHidden[i][j];
}
}
double[] flatWHO = new double[hiddenSize * outputSize];
for (int i = 0; i < hiddenSize; i++) {
for (int j = 0; j < outputSize; j++) {
flatWHO[i * outputSize + j] = weightsHiddenOutput[i][j];
}
}
params.put("weights_input_hidden", flatWIH);
params.put("weights_hidden_output", flatWHO);
params.put("bias_hidden", biasHidden);
params.put("bias_output", biasOutput);
return params;
}
@Override
public void setParameters(Map<String, double[]> parameters) {
double[] flatWIH = parameters.get("weights_input_hidden");
double[] flatWHO = parameters.get("weights_hidden_output");
biasHidden = parameters.get("bias_hidden");
biasOutput = parameters.get("bias_output");
for (int i = 0; i < inputSize; i++) {
for (int j = 0; j < hiddenSize; j++) {
weightsInputHidden[i][j] = flatWIH[i * hiddenSize + j];
}
}
for (int i = 0; i < hiddenSize; i++) {
for (int j = 0; j < outputSize; j++) {
weightsHiddenOutput[i][j] = flatWHO[i * outputSize + j];
}
}
}
}
联邦学习客户端
// FederatedClient.java
package com.federated.client;
import com.federated.base.ModelInterface;
import com.federated.model.NeuralNetwork;
import java.io.*;
import java.net.*;
import java.util.*;
public class FederatedClient {
private String serverHost;
private int serverPort;
private ModelInterface localModel;
private double[][] localData;
private double[] localLabels;
private String clientId;
public FederatedClient(String clientId, String serverHost, int serverPort,
int inputSize, int hiddenSize, int outputSize) {
this.clientId = clientId;
this.serverHost = serverHost;
this.serverPort = serverPort;
this.localModel = new NeuralNetwork(inputSize, hiddenSize, outputSize);
}
// 加载本地数据
public void loadLocalData(double[][] features, double[] labels) {
this.localData = features;
this.localLabels = labels;
}
// 联邦学习训练流程
public void federatedTraining(int communicationRounds, int localEpochs) {
try (Socket socket = new Socket(serverHost, serverPort);
ObjectOutputStream out = new ObjectOutputStream(socket.getOutputStream());
ObjectInputStream in = new ObjectInputStream(socket.getInputStream())) {
// 发送客户端ID
out.writeUTF(clientId);
out.flush();
for (int round = 0; round < communicationRounds; round++) {
// 1. 接收全局模型参数
Map<String, double[]> globalParams = (Map<String, double[]>) in.readObject();
localModel.setParameters(globalParams);
// 2. 本地训练
double localLoss = localModel.train(localData, localLabels, localEpochs);
// 3. 计算模型更新(本地参数 - 全局参数)
Map<String, double[]> localUpdate = calculateUpdate(globalParams);
// 4. 发送本地更新
out.writeObject(new UpdateMessage(clientId, localUpdate, localData.length));
out.flush();
// 5. 接收聚合后的全局模型
Map<String, double[]> aggregatedParams = (Map<String, double[]>) in.readObject();
localModel.setParameters(aggregatedParams);
System.out.println("Client " + clientId + " - Round " + (round + 1) +
" completed. Local loss: " + localLoss);
}
} catch (Exception e) {
e.printStackTrace();
}
}
private Map<String, double[]> calculateUpdate(Map<String, double[]> globalParams) {
Map<String, double[]> localParams = localModel.getParameters();
Map<String, double[]> update = new HashMap<>();
for (String key : localParams.keySet()) {
double[] local = localParams.get(key);
double[] global = globalParams.get(key);
double[] diff = new double[local.length];
for (int i = 0; i < local.length; i++) {
diff[i] = local[i] - global[i];
}
update.put(key, diff);
}
return update;
}
// 评估模型
public void evaluateModel(double[][] testData, double[] testLabels) {
double[] predictions = localModel.predict(testData);
int correct = 0;
for (int i = 0; i < predictions.length; i++) {
if (Math.abs(predictions[i] - testLabels[i]) < 0.5) {
correct++;
}
}
double accuracy = (double) correct / testData.length;
System.out.println("Client " + clientId + " test accuracy: " + String.format("%.2f", accuracy));
}
}
联邦学习服务器
// FederatedServer.java
package com.federated.server;
import com.federated.base.ModelInterface;
import com.federated.model.NeuralNetwork;
import java.io.*;
import java.net.*;
import java.util.*;
import java.util.concurrent.*;
public class FederatedServer {
private int port;
private int numClients;
private ModelInterface globalModel;
private List<ClientConnection> clientConnections = new CopyOnWriteArrayList<>();
// 更新消息类
static class UpdateMessage implements Serializable {
private static final long serialVersionUID = 1L;
private String clientId;
private Map<String, double[]> updates;
private int dataSize;
public UpdateMessage(String clientId, Map<String, double[]> updates, int dataSize) {
this.clientId = clientId;
this.updates = updates;
this.dataSize = dataSize;
}
}
// 客户端连接类
class ClientConnection implements Runnable {
private Socket socket;
private String clientId;
private long lastUpdateTime;
public ClientConnection(Socket socket) {
this.socket = socket;
this.lastUpdateTime = System.currentTimeMillis();
}
@Override
public void run() {
try (ObjectOutputStream out = new ObjectOutputStream(socket.getOutputStream());
ObjectInputStream in = new ObjectInputStream(socket.getInputStream())) {
// 读取客户端ID
this.clientId = in.readUTF();
System.out.println("Client connected: " + clientId);
while (true) {
// 发送全局模型参数
out.writeObject(globalModel.getParameters());
out.flush();
// 接收客户端更新
UpdateMessage message = (UpdateMessage) in.readObject();
clientConnections.add(this);
// 等待所有客户端提交更新
synchronizeUpdates();
// 进行联邦平均聚合
federatedAveraging();
// 发送聚合后的模型
out.writeObject(globalModel.getParameters());
out.flush();
this.lastUpdateTime = System.currentTimeMillis();
}
} catch (Exception e) {
System.out.println("Client " + clientId + " disconnected: " + e.getMessage());
}
}
}
private Map<String, List<double[]>> collectedUpdates = new ConcurrentHashMap<>();
private Map<String, Integer> clientDataSizes = new ConcurrentHashMap<>();
private int expectedClients;
public FederatedServer(int port, int numClients, int inputSize, int hiddenSize, int outputSize) {
this.port = port;
this.numClients = numClients;
this.expectedClients = numClients;
this.globalModel = new NeuralNetwork(inputSize, hiddenSize, outputSize);
}
private synchronized void synchronizeUpdates() {
while (clientConnections.size() < expectedClients) {
try {
wait(1000);
} catch (InterruptedException e) {
Thread.currentThread().interrupt();
}
}
notifyAll();
}
private void federatedAveraging() {
// 清空之前的更新
collectedUpdates.clear();
clientDataSizes.clear();
// 等待收集所有客户端的更新
// 实际上这里应该在收到所有更新后调用
// 进行联邦平均
Map<String, double[]> aggregatedParams = new HashMap<>();
Map<String, double[]> globalParams = globalModel.getParameters();
int totalDataSize = clientDataSizes.values().stream().mapToInt(Integer::intValue).sum();
// 初始化聚合参数
for (String key : globalParams.keySet()) {
aggregatedParams.put(key, new double[globalParams.get(key).length]);
}
// 加权平均所有客户端的更新
for (String clientId : collectedUpdates.keySet()) {
double weight = (double) clientDataSizes.get(clientId) / totalDataSize;
List<double[]> updates = collectedUpdates.get(clientId);
for (String key : globalParams.keySet()) {
double[] aggregated = aggregatedParams.get(key);
double[] update = updates.get(new ArrayList<>(collectedUpdates.keySet()).indexOf(clientId));
for (int i = 0; i < aggregated.length; i++) {
aggregated[i] += weight * update[i];
}
}
}
// 更新全局模型参数
Map<String, double[]> newParams = new HashMap<>();
for (String key : globalParams.keySet()) {
double[] global = globalParams.get(key);
double[] aggregated = aggregatedParams.get(key);
double[] newParam = new double[global.length];
for (int i = 0; i < global.length; i++) {
newParam[i] = global[i] + aggregated[i];
}
newParams.put(key, newParam);
}
globalModel.setParameters(newParams);
System.out.println("Global model updated via Federated Averaging");
}
public void startServer() {
try (ServerSocket serverSocket = new ServerSocket(port)) {
System.out.println("Federated Learning Server started on port " + port);
while (true) {
Socket clientSocket = serverSocket.accept();
ClientConnection clientConnection = new ClientConnection(clientSocket);
new Thread(clientConnection).start();
}
} catch (IOException e) {
e.printStackTrace();
}
}
}
主程序示例
// FederatedLearningDemo.java
package com.federated.demo;
import com.federated.client.FederatedClient;
import com.federated.server.FederatedServer;
import java.util.Random;
public class FederatedLearningDemo {
public static void main(String[] args) {
// 设置参数
int inputSize = 4; // 输入特征数
int hiddenSize = 10; // 隐藏层神经元数
int outputSize = 1; // 输出层神经元数
int numClients = 3; // 客户端数量
int communicationRounds = 10;
int localEpochs = 5;
// 启动服务器(在独立线程中)
FederatedServer server = new FederatedServer(8080, numClients,
inputSize, hiddenSize, outputSize);
new Thread(() -> server.startServer()).start();
// 等待服务器启动
try {
Thread.sleep(2000);
} catch (InterruptedException e) {
e.printStackTrace();
}
// 生成模拟数据并启动客户端
Random random = new Random(123);
int samplesPerClient = 100;
for (int clientNum = 0; clientNum < numClients; clientNum++) {
final int clientId = clientNum;
new Thread(() -> {
// 创建本地数据(模拟分类任务)
double[][] features = new double[samplesPerClient][inputSize];
double[] labels = new double[samplesPerClient];
for (int i = 0; i < samplesPerClient; i++) {
features[i][0] = random.nextDouble() * 2 - 1;
features[i][1] = random.nextDouble() * 2 - 1;
features[i][2] = random.nextDouble() * 2 - 1;
features[i][3] = random.nextDouble() * 2 - 1;
// 简单线性分类规则:x1 + x2 > 0.5 则为正类
labels[i] = (features[i][0] + features[i][1] + 0.5 > 0) ? 1.0 : 0.0;
}
// 创建并启动联邦学习客户端
FederatedClient client = new FederatedClient(
"Client-" + clientId, "localhost", 8080,
inputSize, hiddenSize, outputSize
);
client.loadLocalData(features, labels);
client.federatedTraining(communicationRounds, localEpochs);
// 评估模型
client.evaluateModel(features, labels);
}).start();
}
}
}
增强版本 - 支持更多聚合算法
// AggregationServer.java
package com.federated.server;
import java.util.*;
import java.util.concurrent.*;
public class AggregationServer extends FederatedServer {
public enum AggregationMethod {
FEDAVG, // 联邦平均
FEDPROX, // 近端联邦学习
FEDNOVA, // 客户端差异感知
WEIGHTED_AVG // 加权平均
}
private AggregationMethod method;
private double mu = 0.01; // FEDPROX的近端项系数
public AggregationServer(int port, int numClients, int inputSize,
int hiddenSize, int outputSize, AggregationMethod method) {
super(port, numClients, inputSize, hiddenSize, outputSize);
this.method = method;
}
// 实现不同的聚合算法
@Override
protected void federatedAveraging() {
switch (method) {
case FEDAVG:
fedAvgAggregation();
break;
case FEDPROX:
fedProxAggregation();
break;
case WEIGHTED_AVG:
weightedAggregation();
break;
default:
fedAvgAggregation();
}
}
private void fedAvgAggregation() {
// 标准联邦平均
super.federatedAveraging();
}
private void fedProxAggregation() {
// 近端联邦学习 - 添加近端项正则化
// 实际上参数更新公式变为:w = w - grad - mu * (w - w_t)
// 这里简化实现,增加一个近端项
// 实际应用中需要修改客户端训练逻辑
}
private void weightedAggregation() {
// 基于数据量的加权平均
// 实现自定义加权策略
}
}
数据加密与安全通信
// SecureClient.java
package com.federated.security;
import javax.crypto.*;
import javax.crypto.spec.*;
import java.security.*;
import java.util.*;
public class SecureClient extends FederatedClient {
private SecretKey encryptionKey;
private boolean useEncryption = true;
public SecureClient(String clientId, String serverHost, int serverPort,
int inputSize, int hiddenSize, int outputSize) {
super(clientId, serverHost, serverPort, inputSize, hiddenSize, outputSize);
// 初始化AES密钥
try {
KeyGenerator keyGen = KeyGenerator.getInstance("AES");
keyGen.init(256);
encryptionKey = keyGen.generateKey();
} catch (NoSuchAlgorithmException e) {
e.printStackTrace();
}
}
// 加密参数
private Map<String, double[]> encryptParameters(Map<String, double[]> params)
throws Exception {
if (!useEncryption) return params;
Map<String, double[]> encrypted = new HashMap<>();
for (String key : params.keySet()) {
double[] values = params.get(key);
double[] encryptedValues = new double[values.length];
// 使用AES对每个值进行加密(实用场景中应考虑批量加密)
for (int i = 0; i < values.length; i++) {
byte[] data = String.valueOf(values[i]).getBytes();
Cipher cipher = Cipher.getInstance("AES");
cipher.init(Cipher.ENCRYPT_MODE, encryptionKey);
byte[] encryptedData = cipher.doFinal(data);
encryptedValues[i] = new Random().nextDouble(); // 简化处理
}
encrypted.put(key, encryptedValues);
}
return encrypted;
}
}
这个联邦学习案例包含:
- 基础框架:模型定义、客户端-服务器架构
- 通信协议:使用Java Socket进行通信
- 聚合算法:实现联邦平均算法
- 安全特性:支持加密通信(可选)
- 扩展性:支持不同的聚合算法和模型类型
典型案例应用场景包括:
- 医疗数据分析(不同医院共享模型而不共享数据)
- 金融风控(银行间联合训练而不泄露客户隐私)
- 智能设备(手机端个性化模型训练)