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

PyTorch联邦学习中DataLoader引发内存崩溃的解决办法?

联邦学习LSTM实现中Colab内存耗尽崩溃问题

我正尝试基于LADPU数据集实现联邦学习以预测太阳能光伏发电量。预处理数据集后,按唯一的METER_FID值划分为10个客户端的数据集片段模拟联邦环境,接着从这些分区生成序列作为LSTM模型输入,用START_READ和END_READ构成的序列预测INTERVAL_READ。但在划分训练集、验证集并创建PyTorch DataLoader对象后,Colab会话因耗尽所有可用内存持续崩溃。已尝试减小batch size并限制为6个CSV文件,问题仍存在。我是该领域新手,参考联邦学习相关教程和文档,无法定位问题,求帮助。

# 数据分区代码
NUM_CLIENTS = 10

def partition_data(df, num_clients):
    np.random.seed(42) 
    unique_ids = df['METER_FID'].unique()
    np.random.shuffle(unique_ids) 
    partitions = np.array_split(unique_ids, num_clients)  
    partitioned_dfs = [df[df['METER_FID'].isin(ids)] for ids in partitions]
    return partitioned_dfs

partitioned_dfs = partition_data(df, NUM_CLIENTS)

# 生成LSTM输入序列
def create_sequences_efficiently(df, sequence_length=5):
    sequences, targets = [], []

    df = df.sort_values('INTERVAL_TIME')
    for i in range(sequence_length, len(df)):
        sequence = df[['START_READ', 'END_READ']].values[i-sequence_length:i]  
        target = df['INTERVAL_READ'].values[i]  
        sequences.append(sequence)
        targets.append(target)
    return np.array(sequences), np.array(targets)

# 创建DataLoader
from torch.utils.data import DataLoader, TensorDataset
import torch
from sklearn.model_selection import train_test_split


def create_loaders(partition, sequence_length=5, batch_size=32):
    sequences, targets = create_sequences_efficiently(partition, sequence_length)
    # 转换为PyTorch张量
    sequence_tensor = torch.tensor(sequences, dtype=torch.float32)
    target_tensor = torch.tensor(targets, dtype=torch.float32)
    # 创建数据集和加载器
    dataset = TensorDataset(sequence_tensor, target_tensor)
    return DataLoader(dataset, batch_size=batch_size, shuffle=True)

# 为每个分区创建加载器
train_loaders, val_loaders = [], []
for partition_df in partitioned_dfs:
    # 划分训练集和验证集
    train_df, val_df = train_test_split(partition_df, test_size=0.2, random_state=42)
    # 创建加载器
    train_loader = create_loaders(train_df, sequence_length=5, batch_size=16)
    val_loader = create_loaders(val_df, sequence_length=5, batch_size=16)
    train_loaders.append(train_loader)
    val_loaders.append(val_loader)

内存优化解决方案

1. 改用自定义Dataset按需生成序列

当前代码会一次性生成所有序列并转为张量,直接把全量数据加载到内存。换成自定义Dataset类,在取数时才生成对应序列,避免预加载全部数据:

class SequenceDataset(torch.utils.data.Dataset):
    def __init__(self, df, sequence_length=5):
        # 提前排序并重置索引,避免重复排序
        self.df = df.sort_values('INTERVAL_TIME').reset_index(drop=True)
        self.sequence_length = sequence_length
        # 提前提取numpy数组,减少重复索引操作
        self.features = self.df[['START_READ', 'END_READ']].values
        self.targets = self.df['INTERVAL_READ'].values

    def __len__(self):
        return len(self.df) - self.sequence_length

    def __getitem__(self, idx):
        # 按需截取序列和目标值
        sequence = self.features[idx:idx+self.sequence_length]
        target = self.targets[idx+self.sequence_length]
        return torch.tensor(sequence, dtype=torch.float32), torch.tensor(target, dtype=torch.float32)

修改create_loaders函数适配新Dataset:

def create_loaders(partition, sequence_length=5, batch_size=16):
    dataset = SequenceDataset(partition, sequence_length)
    return DataLoader(dataset, batch_size=batch_size, shuffle=True)

2. 清理分区数据的内存引用

原分区代码生成的partitioned_dfs是原DataFrame的切片,会保留对原数据的引用,导致内存冗余。给每个分区加上.copy():

partitioned_dfs = [df[df['METER_FID'].isin(ids)].copy() for ids in partitions]

3. 按需创建客户端DataLoader,而非一次性生成所有

不要提前把10个客户端的所有加载器都创建出来,而是在联邦训练循环中,为当前参与的客户端创建加载器,训练完成后及时释放:

# 替换原有的批量创建逻辑
def get_client_loaders(client_df, sequence_length=5, batch_size=16):
    train_df, val_df = train_test_split(client_df, test_size=0.2, random_state=42)
    train_dataset = SequenceDataset(train_df, sequence_length)
    val_dataset = SequenceDataset(val_df, sequence_length)
    return DataLoader(train_dataset, batch_size=batch_size, shuffle=True), DataLoader(val_dataset, batch_size=batch_size)

# 联邦训练时的示例逻辑
import gc

for client_idx, client_df in enumerate(partitioned_dfs):
    train_loader, val_loader = get_client_loaders(client_df)
    # 执行当前客户端的训练/验证逻辑
    # ...(训练代码)
    
    # 训练完成后手动释放内存
    del train_loader, val_loader
    gc.collect()
    torch.cuda.empty_cache()  # 使用GPU时清理缓存

4. 降低数据类型的内存占用

检查原DataFrame的数值类型,将float64转为float32(张量已经转成float32,但原DataFrame可能还是高占用类型):

df[['START_READ', 'END_READ', 'INTERVAL_READ']] = df[['START_READ', 'END_READ', 'INTERVAL_READ']].astype('float32')

5. Colab环境内存管理

  • 若条件允许,切换到Colab Pro/Pro+获取更大内存配额
  • 定期执行gc.collect()和torch.cuda.empty_cache()清理内存
  • 分区完成后删除原DataFrame,释放内存:del df

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.27 22:46:08