新版本PySyft中sy.FederatedDataLoader替代方案咨询
解决PySyft新版本中
FederatedDataLoader不存在的问题 在新版PySyft(0.7.x及以后版本)中,sy.FederatedDataLoader已被移除,取而代之的是直接使用PyTorch原生的torch.utils.data.DataLoader,同时需要调整VirtualMachine的使用方式——必须先获取VM对应的客户端实例,再将数据联邦化到这些客户端上。
关键修改点
- 为每个
VirtualMachine创建客户端对象:westside_client = westside.get_client()、grapevine_client = grapevine.get_client() - 用
torch.utils.data.DataLoader替换sy.FederatedDataLoader,参数用法和原生DataLoader完全一致
修改后的完整代码
import torch import torch.nn as nn import torch.nn.functional as F import torch.optim as optim from torchvision import datasets, transforms import logging import syft as sy # 创建虚拟机后必须获取客户端实例 westside = sy.VirtualMachine(name="westside") westside_client = westside.get_client() grapevine = sy.VirtualMachine(name="grapevine") grapevine_client = grapevine.get_client() # 超参数定义不变 args = { 'use_cuda': True, 'batch_size': 64, 'test_batch_size': 1000, 'lr': 0.01, 'log_interval': 100, 'epochs': 10 } use_cuda = args['use_cuda'] and torch.cuda.is_available() device = torch.device('cuda' if use_cuda else 'cpu') # CNN网络定义不变 class Net(nn.Module): def __init__(self): super(Net, self).__init__() self.conv = nn.Sequential( nn.Conv2d(in_channels=1, out_channels=32, kernel_size=3, stride=1), nn.ReLU(), nn.Conv2d(in_channels=32, out_channels=64, kernel_size=3, stride=1), nn.ReLU() ) self.fc = nn.Sequential( nn.Linear(in_features=64*12*12, out_features=128), nn.ReLU(), nn.Linear(in_features=128, out_features=10), ) def forward(self, x): x = self.conv(x) x = F.max_pool2d(x,2) x = x.view(-1, 64*12*12) x = self.fc(x) x = F.log_softmax(x, dim=1) return x # 替换为PyTorch原生DataLoader,传入联邦化后的数据集 train_dataset = datasets.MNIST('../data', train=True, download=True, transform=transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ])) # 将数据联邦化到客户端(注意这里传的是client对象,不是VM) federated_train_dataset = train_dataset.federate((grapevine_client, westside_client)) # 使用原生DataLoader federated_train_loader = torch.utils.data.DataLoader( federated_train_dataset, batch_size=args['batch_size'], shuffle=True )
原理说明
新版PySyft重构了核心架构,更贴近PyTorch原生生态:
- 联邦化后的数据集(
FederatedDataset)实现了PyTorch的Dataset接口,因此可以直接用原生DataLoader加载 VirtualMachine本身是服务端实例,必须通过get_client()获取客户端对象,才能作为数据联邦化的目标节点
内容的提问来源于stack exchange,提问作者onapte
相关产品推荐
相关产品推荐

