使用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

