You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

基于联邦平均的MNIST数据集CNN模型准确率异常波动问题求助

联邦平均CNN模型准确率异常波动问题修复

问题描述

我实现了一个基于联邦平均(Federated Averaging)的CNN模型,在MNIST数据集上训练。训练流程为多个客户端在本地数据上训练若干轮后,对模型参数进行平均。但出现异常:客户端数量为2时全局模型准确率达97%,数量为3时仅11%。由于客户端数据是IID分布,预期1个客户端时准确率较低,随客户端数量增加逐步提升。

错误分析

  • 联邦平均计算逻辑错误:原代码中,全局模型参数初始值未清零就直接累加客户端参数,最终得到的是(初始参数 + 客户端1参数 + 客户端2参数 + ... + 客户端N参数)/N,而非正确的所有客户端参数的平均值(客户端1参数 + 客户端2参数 + ... + 客户端N参数)/N。这会导致初始随机参数持续干扰全局模型更新,当客户端数量变化时,干扰程度不同,引发准确率波动。
  • 数据分配未做Shuffle:MNIST原始训练数据是按标签排序存储的,直接按切片分配给客户端会导致每个客户端的数据分布严重不均(非IID),这会让联邦平均后的模型无法有效泛化,尤其是客户端数量为3时,每个客户端的标签覆盖范围极小,导致模型失效。
  • 重复的Tensor转换:自定义MNISTDataset中再次使用transforms.ToTensor(),而原始train_dataset已经做过该转换,重复转换会导致数据值异常(比如两次归一化),影响模型训练。

修复方案

  1. 分配数据前先打乱训练数据的索引,确保每个客户端拿到IID分布的数据。
  2. 调整联邦平均逻辑:每个全局epoch开始时,先将全局模型参数重置为0,再累加所有客户端训练后的参数,最后除以客户端数量得到平均参数。
  3. 移除MNISTDataset中重复的ToTensor转换,只保留必要的数据处理步骤。

修复后的代码

import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import DataLoader, Dataset
from torchvision import datasets, transforms

# Define the CNN architecture
class CNN(nn.Module):
    def __init__(self):
        super(CNN, self).__init__()
        self.conv1 = nn.Conv2d(1, 32, kernel_size=3)
        self.relu1 = nn.ReLU()
        self.conv2 = nn.Conv2d(32, 64, kernel_size=3)
        self.relu2 = nn.ReLU()
        self.pool = nn.MaxPool2d(2)
        self.fc1 = nn.Linear(64 * 12 * 12, 128)
        self.relu3 = nn.ReLU()
        self.fc2 = nn.Linear(128, 10)

    def forward(self, x):
        x = self.relu1(self.conv1(x))
        x = self.pool(self.relu2(self.conv2(x)))
        x = x.view(-1, 64 * 12 * 12)
        x = self.relu3(self.fc1(x))
        x = self.fc2(x)
        return x

# Define the dataset class
class MNISTDataset(Dataset):
    def __init__(self, data, targets, transform=None):
        self.data = data
        self.targets = targets
        self.transform = transform

    def __len__(self):
        return len(self.data)

    def __getitem__(self, index):
        x = self.data[index]
        y = self.targets[index]

        if self.transform:
            x = transforms.ToPILImage()(x)
            x = self.transform(x)

        return x, y

# Load MNIST dataset
train_dataset = datasets.MNIST(
    './data',
    train=True,
    download=True,
    transform=transforms.Compose([
        transforms.ToTensor(),
    ])
)

test_dataset = datasets.MNIST(
    './data',
    train=False,
    download=True,
    transform=transforms.Compose([
        transforms.ToTensor(),
    ])
)

# Define the number of clients
num_clients = 3

# Shuffle the training data indices first to ensure IID distribution
shuffled_indices = torch.randperm(len(train_dataset))
data_per_client = len(train_dataset) // num_clients
client_datasets = []

for i in range(num_clients):
    start_index = i * data_per_client
    end_index = (i + 1) * data_per_client
    client_indices = shuffled_indices[start_index:end_index]
    data = train_dataset.data[client_indices]
    targets = train_dataset.targets[client_indices]
    # Remove duplicate ToTensor transform
    client_dataset = MNISTDataset(data, targets, transform=None)
    client_datasets.append(client_dataset)

# Define the federated learning parameters
num_epochs = 3
learning_rate = 0.01

# Initialize the global model
global_model = CNN()

# Define the loss function
criterion = nn.CrossEntropyLoss()

# Train the global model using federated averaging
for epoch in range(num_epochs):
    # Reset global model parameters to 0 before averaging
    for param in global_model.parameters():
        param.data.zero_()
    
    for client_dataset in client_datasets:
        # Create data loader for each client
        client_loader = DataLoader(client_dataset, batch_size=64, shuffle=True)
        # Initialize the local model with global state
        local_model = CNN()
        local_model.load_state_dict(global_model.state_dict())

        # Define the optimizer for the local model
        local_optimizer = optim.Adam(local_model.parameters(), lr=learning_rate)

        # Train the local model
        local_model.train()
        for inputs, labels in client_loader:
            local_optimizer.zero_grad()
            outputs = local_model(inputs)
            loss = criterion(outputs, labels)
            loss.backward()
            local_optimizer.step()

        # Accumulate local model parameters to global model
        for global_param, local_param in zip(global_model.parameters(), local_model.parameters()):
            global_param.data += local_param.data

    # Average the accumulated parameters by number of clients
    for global_param in global_model.parameters():
        global_param.data /= num_clients

# Evaluate the global model on testing data
test_loader = DataLoader(test_dataset, batch_size=64, shuffle=False)
total_correct = 0
total_samples = 0

global_model.eval()
with torch.no_grad():
    for inputs, labels in test_loader:
        outputs = global_model(inputs)
        _, predicted = torch.max(outputs, 1)
        total_samples += labels.size(0)
        total_correct += (predicted == labels).sum().item()

accuracy = 100.0 * total_correct / total_samples
print(f"Global Model Accuracy: {accuracy}%") 

内容的提问来源于stack exchange,提问作者Hashitiba

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.19 12:13:20