基于联邦平均的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已经做过该转换,重复转换会导致数据值异常(比如两次归一化),影响模型训练。
修复方案
- 分配数据前先打乱训练数据的索引,确保每个客户端拿到IID分布的数据。
- 调整联邦平均逻辑:每个全局epoch开始时,先将全局模型参数重置为0,再累加所有客户端训练后的参数,最后除以客户端数量得到平均参数。
- 移除
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
相关产品推荐
相关产品推荐

