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

新版本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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.21 15:15:35