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

使用PyTorch Lightning训练BERT模型时CPU内存溢出问题咨询

问题分析与解决办法

导致CPU内存溢出的核心原因及对应修复方案:


1. DataModule未正确初始化数据集

你的BertDataModule中self.train_dataset初始化为None,且未实现setup()方法来完成数据集的拆分与赋值。train_dataloader()直接调用未初始化的数据集对象,会触发异常的数据加载逻辑,进而导致内存泄漏或重复加载。

修复代码:

class BertDataModule(pl.LightningDataModule):
    def __init__(self, dataset, train_split, batch_size, data_collator):
        super().__init__()
        self.train_dataset = None
        self.val_dataset = None
        self.dataset = dataset
        self.collator = data_collator
        self.batch_size = batch_size
        self.train_split = train_split

    def setup(self, stage=None):
        # 完成训练/验证集拆分
        splits = self.dataset.train_test_split(test_size=1 - self.train_split)
        self.train_dataset = splits['train']
        self.val_dataset = splits['test']

    def train_dataloader(self):
        return torch.utils.data.DataLoader(
            self.train_dataset, 
            batch_size=self.batch_size, 
            collate_fn=self.collator, 
            num_workers=8,  # 下调worker数量
            pin_memory=True  # 优化GPU数据传输,减少CPU内存占用
        )

    # 补全val_dataloader方法,使用self.val_dataset
    def val_dataloader(self):
        return torch.utils.data.DataLoader(
            self.val_dataset, 
            batch_size=self.batch_size, 
            collate_fn=self.collator, 
            num_workers=4,
            pin_memory=True
        )

2. num_workers设置过高

num_workers=30会启动30个数据加载子进程,每个进程都会缓存部分数据集,叠加后CPU内存占用会急剧上升——即使单进程内存占用不高,多进程的内存副本总和也会轻松超过阈值。

修复方案:
根据容器CPU核心数设置合理值,一般建议为核心数的1~2倍(比如8或16),同时启用pin_memory=True优化GPU数据传输,减少CPU端的临时内存占用。


3. 未启用数据集流式加载

load_dataset('text', data_path)默认会将整个33GB数据集一次性加载到内存,再加上多进程的内存副本,实际占用内存会远大于33GB,最终触发OOM。

修复代码:

# 启用流式加载,避免一次性加载全量数据
dataset = load_dataset('text', data_path, streaming=True)

4. 验证逻辑的潜在内存缓存

val_check_interval=50000意味着每50000步才执行一次验证,但如果验证数据集未正确初始化或加载,后台可能会持续缓存未处理的数据,导致内存持续增长。

修复方案:
确保val_dataloader()使用拆分后的验证集,且设置合理的batch_size,避免加载过多验证数据到内存。


5. 临时张量的内存泄漏

DataCollatorForLanguageModeling在处理每个batch时会生成大量临时张量,若Python垃圾回收不及时,加上多进程的叠加效应,会导致CPU内存持续占用。

修复方案:
在训练循环中定期触发垃圾回收,比如在LightningModule的on_epoch_end钩子中添加:

import gc

class BertModel(pl.LightningModule):
    # ... 其他代码 ...
    def on_epoch_end(self):
        gc.collect()

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.16 00:53:18