联邦学习代码训练轮次≥97时触发RuntimeError报错求助
联邦学习乳腺癌分类代码报错排查与解决
问题描述
基于PyTorch实现的联邦学习二分类代码,在训练轮次num_rounds ≤96时正常运行,但设置num_rounds ≥97时触发以下RuntimeError:
RuntimeError: all elements of input should be between 0 and 1
代码实现如下:
import torch import torch.nn as nn import torch.optim as optim from torch.utils.data import DataLoader, Dataset import numpy as np from sklearn.datasets import load_breast_cancer from sklearn.model_selection import train_test_split # Define the deep neural network model class DNN(nn.Module): def __init__(self, input_size, hidden_size, output_size): super(DNN, self).__init__() self.fc1 = nn.Linear(input_size, hidden_size) self.relu = nn.ReLU() self.fc2 = nn.Linear(hidden_size, output_size) self.sigmoid = nn.Sigmoid() def forward(self, x): out = self.fc1(x) out = self.relu(out) out = self.fc2(out) out = self.sigmoid(out) return out # Load the breast cancer dataset data = load_breast_cancer() X = data.data y = data.target X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2, random_state=42) # Define the number of training rounds and the number of clients num_rounds = 100 num_clients = 2 batch_size = 10 # Split the data into equal chunks for each client X_splits = np.array_split(X_train, num_clients) y_splits = np.array_split(y_train, num_clients) # Define the loss function and optimizer criterion = nn.BCELoss() # Perform federated learning global_model = DNN(X_train.shape[1], 16, 1) optimizer = optim.SGD(global_model.parameters(), lr=.01) for i in range(num_rounds): local_models = [] for j in range(num_clients): # Create a local model by copying the current global model local_model = DNN(X_train.shape[1], 16, 1) local_model.load_state_dict(global_model.state_dict()) # Create a dataloader for the local client's data local_X = torch.tensor(X_splits[j], dtype=torch.float32) local_y = torch.tensor(y_splits[j], dtype=torch.float32) local_dataset = torch.utils.data.TensorDataset(local_X, local_y) local_dataloader = DataLoader(local_dataset, batch_size=batch_size, shuffle=True) # Train the local model local_optimizer = optim.SGD(local_model.parameters(), lr=0.1) for inputs, labels in local_dataloader: local_optimizer.zero_grad() outputs = local_model(inputs) loss = criterion(outputs, labels.view(-1, 1)) loss.backward() local_optimizer.step() # Add the trained local model to the list of local models local_models.append(local_model) # Aggregate the local models to create a global model with torch.no_grad(): for global_param, local_params in zip(global_model.parameters(), zip(*[local_model.parameters() for local_model in local_models])): global_param.data += torch.stack(local_params).sum(0) / num_clients # Evaluate the global model on the train dataset global_model.eval() with torch.no_grad(): global_outputs = global_model(torch.tensor(X_train, dtype=torch.float32)) global_loss = criterion(global_outputs, torch.tensor(y_train, dtype=torch.float32).view(-1, 1)) global_pred = (global_outputs > 0.5).int().numpy().flatten() accuracy = np.mean(global_pred == y_train) print(f"Round {i}, train accuracy:{accuracy}")
报错原因分析
核心问题出在联邦学习的模型聚合逻辑错误:
- 代码中聚合步骤使用了
global_param.data += torch.stack(local_params).sum(0) / num_clients,这是在原有全局参数的基础上累加本地模型参数的平均值,而非替换为本地模型参数的平均值。 - 随着训练轮次增加,全局模型的参数会不断累积增大,导致模型前向传播时,全连接层输出的数值绝对值越来越大。尽管sigmoid函数理论输出范围是(0,1),但当输入数值过大时,会因浮点数值精度问题出现超出[0,1]范围的异常值,最终触发BCELoss的输入合法性检查报错。
解决方案
修正模型聚合逻辑,将全局参数替换为所有本地模型参数的平均值(符合FedAvg算法的标准实现):
将原聚合代码:
with torch.no_grad(): for global_param, local_params in zip(global_model.parameters(), zip(*[local_model.parameters() for local_model in local_models])): global_param.data += torch.stack(local_params).sum(0) / num_clients
修改为:
with torch.no_grad(): for global_param, local_params in zip(global_model.parameters(), zip(*[local_model.parameters() for local_model in local_models])): # 替换为本地模型参数的平均值,而非累加 global_param.data = torch.stack(local_params).sum(0) / num_clients
额外优化建议:
- 在训练前对数据集做标准化处理(乳腺癌数据集特征尺度差异较大),可以提升模型稳定性和收敛速度:
from sklearn.preprocessing import StandardScaler scaler = StandardScaler() X_train = scaler.fit_transform(X_train) X_test = scaler.transform(X_test) - 本地训练时可以设置固定的训练epoch数,而非遍历完整个数据集一次,避免本地模型训练过度偏离全局模型。
内容的提问来源于stack exchange,提问作者Mustakim Pallab
相关产品推荐
相关产品推荐

