联邦学习(Federated Learning)的核心是 “数据不动模型动”:多个客户端在本地用私有数据训练模型,仅上传模型参数 / 梯度到服务器,服务器聚合后将新模型分发给客户端,反复迭代直至收敛。本示例模拟 3 个客户端 + 1 个服务器的联邦学习场景,任务为文本分类(基于情感分析数据)。

环境准备

先安装依赖:pip install flwr torch transformers datasets
pip install torch transformers datasets numpy
import sys
import numpy as np
import torch
import torch.nn.functional as F
from torch.utils.data import DataLoader
from transformers import DistilBertForSequenceClassification, DistilBertTokenizer
from datasets import load_dataset
import flwr as fl
from flwr.client import FlowerClient, start_client
from flwr.server import start_server
from flwr.server.strategy import FedAvg
from flwr.common import Context, Parameters, ndarrays_to_parameters, parameters_to_ndarrays


# -------------------------- 1. 模型与分词器定义 --------------------------
def get_model():
    """加载基础模型(替代千问0.5,这里用distilbert演示)"""
    model = DistilBertForSequenceClassification.from_pretrained(
        "distilbert-base-uncased", num_labels=2  # 二分类任务
    )
    return model

# 全局加载分词器(所有客户端共享)
tokenizer = DistilBertTokenizer.from_pretrained("distilbert-base-uncased")


# -------------------------- 2. 客户端数据准备(非IID分布) --------------------------
def prepare_client_data(num_clients=3):
    """拆分IMDb数据集,模拟非IID分布的客户端数据"""
    dataset = load_dataset("imdb")["train"]  # 加载电影评论数据集(正/负面分类)
    
    client_datasets = []
    for i in range(num_clients):
        # 每个客户端侧重不同类别(非IID特性)
        target_label = i % 2  # 客户端0/2侧重0类,客户端1侧重1类
        mask = dataset["label"] == target_label
        client_data = dataset.select(np.where(mask)[0])  # 筛选目标类别数据
        client_datasets.append(client_data)
    
    return client_datasets

# 预先生成3个客户端的数据集(全局变量,供客户端使用)
client_datasets = prepare_client_data(num_clients=3)


# -------------------------- 3. 数据预处理函数 --------------------------
def preprocess_function(examples):
    """对文本数据进行分词、截断和填充"""
    return tokenizer(
        examples["text"],
        truncation=True,
        max_length=128,  # 固定长度,便于批量处理
        padding="max_length"
    )


# -------------------------- 4. 联邦学习客户端实现 --------------------------
class FedClient(FlowerClient):
    def __init__(self, client_id, model, dataset):
        self.client_id = client_id
        self.model = model
        self.dataset = dataset
        
        # 预处理数据并创建DataLoader
        self.processed_dataset = self.dataset.map(preprocess_function, batched=True)
        self.processed_dataset.set_format("torch", columns=["input_ids", "attention_mask", "label"])
        self.dataloader = DataLoader(self.processed_dataset, batch_size=8)  # 批量大小8

    def get_parameters(self, context: Context):
        """获取本地模型参数(转为numpy数组)"""
        return ndarrays_to_parameters([
            p.cpu().detach().numpy() for p in self.model.parameters()
        ])

    def fit(self, parameters, context: Context):
        """本地训练:用全局参数初始化→训练→返回新参数"""
        # 1. 用服务器传来的全局参数更新本地模型
        params = parameters_to_ndarrays(parameters)
        for p, new_p in zip(self.model.parameters(), params):
            p.data = torch.tensor(new_p).to(p.device)  # 适配设备(CPU/GPU)

        # 2. 本地训练(1个epoch,简化演示)
        optimizer = torch.optim.Adam(self.model.parameters(), lr=1e-5)  # 优化器
        self.model.train()
        for batch in self.dataloader:
            # 提取批次数据
            input_ids = batch["input_ids"]
            attention_mask = batch["attention_mask"]
            labels = batch["label"]
            
            # 前向传播+计算损失
            outputs = self.model(
                input_ids=input_ids,
                attention_mask=attention_mask,
                labels=labels
            )
            loss = outputs.loss
            
            # 反向传播+参数更新
            loss.backward()
            optimizer.step()
            optimizer.zero_grad()

        # 3. 返回训练后的参数和本地数据量
        return self.get_parameters(context), len(self.dataset), {}

    def evaluate(self, parameters, context: Context):
        """简化评估:返回0(实际场景需计算准确率等指标)"""
        return 0.0, len(self.dataset), {}


def client_fn(context: Context):
    """客户端工厂函数(Flower框架调用)"""
    client_id = int(context.node_config["client_id"])  # 获取客户端ID
    model = get_model()  # 加载模型
    dataset = client_datasets[client_id]  # 分配对应ID的本地数据
    return FedClient(client_id, model, dataset)


# -------------------------- 5. 联邦学习服务器实现 --------------------------
def start_fl_server():
    """启动联邦学习服务器(聚合客户端参数)"""
    # 服务器策略:使用FedAvg算法聚合参数
    strategy = FedAvg(
        fraction_fit=1.0,  # 每次训练使用所有客户端
        fraction_evaluate=0.0,  # 简化:不进行评估
    )

    # 启动服务器(监听本地8080端口,训练3轮)
    start_server(
        server_address="0.0.0.0:8080",
        strategy=strategy,
        config={"num_rounds": 3}
    )


# -------------------------- 6. 启动入口(区分服务器/客户端) --------------------------
if __name__ == "__main__":
    # 通过命令行参数区分角色:server 或 client+ID
    if len(sys.argv) < 2:
        print("使用方式:")
        print("  启动服务器:python script.py server")
        print("  启动客户端:python script.py client 0(或1、2)")
        sys.exit(1)

    role = sys.argv[1]
    if role == "server":
        print("启动联邦学习服务器...")
        start_fl_server()
    elif role == "client" and len(sys.argv) == 3:
        client_id = int(sys.argv[2])
        print(f"启动客户端 {client_id}...")
        start_client(
            server_address="localhost:8080",
            client_fn=client_fn,
            client_id=client_id
        )
    else:
        print("参数错误!请参考使用方式。")

更多推荐