Java联邦学习案例

wen java案例 2

本文目录导读:

Java联邦学习案例

  1. 联邦学习框架设计
  2. 增强版本 - 支持更多聚合算法
  3. 数据加密与安全通信

我将为您提供一个完整的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;
    }
}

这个联邦学习案例包含:

  1. 基础框架:模型定义、客户端-服务器架构
  2. 通信协议:使用Java Socket进行通信
  3. 聚合算法:实现联邦平均算法
  4. 安全特性:支持加密通信(可选)
  5. 扩展性:支持不同的聚合算法和模型类型

典型案例应用场景包括:

  • 医疗数据分析(不同医院共享模型而不共享数据)
  • 金融风控(银行间联合训练而不泄露客户隐私)
  • 智能设备(手机端个性化模型训练)

抱歉,评论功能暂时关闭!