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

Kaggle环境下PyTorch+PyTorch Lightning训练时系统内存持续上涨问题求助

PyTorch Lightning训练内存泄漏排查与修复方案

1. 首要泄漏点:训练/验证步骤返回张量持有计算图

你在training_step、validation_step中返回的predictions、labels均为带梯度的张量,直接在*_epoch_end中用列表累加时,会持续持有全量批次的计算图无法释放,每迭代一个batch内存就上涨一份,最终触发内存阈值。
修复方案:
推荐直接使用TorchMetrics封装的指标工具,无需手动攒全量预测结果,避免内存占用:

from torchmetrics import AUROC

class G2NetLightningModule(pl.LightningModule):
    def __init__(self, pretrained=True):
        super().__init__()
        self.model = timm.create_model(CFG.model_name, pretrained=pretrained, in_chans=1)
        n_features = self.model.classifier.in_features
        self.model.classifier = nn.Linear(n_features, CFG.target_size)
        self.criterion = nn.BCEWithLogitsLoss()
        # 初始化二分类AUC指标
        self.train_auc = AUROC(task="binary")
        self.val_auc = AUROC(task="binary")
    
    def training_step(self, batch, batch_idx):        
        x, y = batch
        y_preds = self.model(x).view(-1)
        loss = self.criterion(y_preds, y)
        # 更新指标,无需保留全量预测值
        self.train_auc.update(y_preds.sigmoid(), y.long())
        self.log('train_loss_step', loss, on_step=True, prog_bar=True, logger=True)
        return loss
    
    def on_train_epoch_end(self):
        # 输出 epoch 指标后重置,释放内存
        self.log('train_auc_epoch', self.train_auc.compute(), prog_bar=True, logger=True)
        self.train_auc.reset()
    
    def validation_step(self, batch, batch_idx):
        x, y = batch
        y_preds = self.model(x).view(-1)
        loss = self.criterion(y_preds, y)
        self.val_auc.update(y_preds.sigmoid(), y.long())
        self.log('val_loss', loss, prog_bar=True, logger=True)
    
    def on_validation_epoch_end(self):
        self.log('val_auc_epoch', self.val_auc.compute(), prog_bar=True, logger=True)
        self.val_auc.reset()

如果要保留手动计算的逻辑,需要在step返回时就把张量移到CPU并切断梯度:

return dict(loss=loss, 
            predictions=y_preds.detach().cpu(), 
            labels=y.detach().cpu())

2. 数据加载配置优化

Kaggle内核CPU内存仅16G左右,不合理的Dataloader配置也会导致内存持续上涨:

  • CFG.num_workers建议设置为1~2,不要超过4,多进程加载数据会额外占用内存
  • 可以在Dataloader参数中添加persistent_workers=False,每个epoch结束后自动释放数据加载进程的内存
  • 若CQT1992v2变换为GPU实现,建议改为CPU实现,避免数据加载过程产生显存碎片

3. 其他优化点

  • 跑多fold训练时,每个fold结束后调用torch.cuda.empty_cache()清空显存缓存
  • 避免在step中打印/上报多余的张量数据,减少日志内存占用

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.07 15:00:03