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

使用Flower框架实现联邦学习时客户端训练报错的解决方法

联邦学习Flower框架训练启动报错求助

我在Python中使用Flower框架实现联邦学习时,启动训练流程触发报错。

以下是我的实现代码:

NUM_CLIENTS = 10

# 加载数据函数
def load_datasets(num_clients: int, train_loader, test_loader):
    # 将训练集拆分为`num_clients`个分区,模拟不同本地数据集
    partition_size = len(train_loader) // num_clients
    lengths = [partition_size] * num_clients
    datasets = random_split(train_loader, lengths, torch.Generator().manual_seed(42))

    # 将每个分区拆分为训练/验证集并创建DataLoader
    trainloaders = []
    valloaders = []
    for ds in datasets:
        len_val = len(ds) // 10  # 10% 验证集
        len_train = len(ds) - len_val
        lengths = [len_train, len_val]
        ds_train, ds_val = random_split(ds, lengths, torch.Generator().manual_seed(42))
        trainloaders.append(DataLoader(ds_train, batch_size=32, shuffle=True))
        valloaders.append(DataLoader(ds_val, batch_size=32))
    testloader = DataLoader(test_loader, batch_size=32)
    return trainloaders, valloaders, testloader


trainloaders, valloaders, testloader = load_datasets(NUM_CLIENTS ,train_loader,test_loader)


# 传递给服务器启动的客户端函数
def client_fn(cid) -> CardiacClient:
    net = CardiacModel().to(DEVICE)
    trainloader = trainloaders[cid]
    valloader = valloaders[cid]
    return CardiacClient(cid, net, trainloader, valloader)

注:代码中的cid指代客户端ID。

问题分析与解决方案

从报错信息(TypeError: list indices must be integers or slices, not str)来看,核心问题有两个:

  1. cid类型不匹配:Flower的client_fn接收的cid是字符串类型,但直接用它索引列表trainloaders和valloaders会触发类型错误——列表仅支持整数索引。
  2. 数据加载逻辑错误:random_split只能作用于PyTorch的Dataset对象,你传入的train_loader是DataLoader,len(train_loader)返回的是批次数而非样本总数,会导致数据分区错误。

修正代码

  • 处理cid类型转换:
def client_fn(cid) -> CardiacClient:
    net = CardiacModel().to(DEVICE)
    # 将字符串cid转为整数
    client_id = int(cid)
    trainloader = trainloaders[client_id]
    valloader = valloaders[client_id]
    return CardiacClient(cid, net, trainloader, valloader)
  • 修正数据加载函数,传入Dataset而非DataLoader:
def load_datasets(num_clients: int, train_dataset, test_dataset):
    # 将训练集拆分为`num_clients`个分区,模拟不同本地数据集
    partition_size = len(train_dataset) // num_clients
    lengths = [partition_size] * num_clients
    datasets = random_split(train_dataset, lengths, torch.Generator().manual_seed(42))

    # 将每个分区拆分为训练/验证集并创建DataLoader
    trainloaders = []
    valloaders = []
    for ds in datasets:
        len_val = len(ds) // 10  # 10% 验证集
        len_train = len(ds) - len_val
        lengths = [len_train, len_val]
        ds_train, ds_val = random_split(ds, lengths, torch.Generator().manual_seed(42))
        trainloaders.append(DataLoader(ds_train, batch_size=32, shuffle=True))
        valloaders.append(DataLoader(ds_val, batch_size=32))
    testloader = DataLoader(test_dataset, batch_size=32)
    return trainloaders, valloaders, testloader

# 调用时传入Dataset对象
trainloaders, valloaders, testloader = load_datasets(NUM_CLIENTS, train_dataset, test_dataset)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.23 11:27:32