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

如何优化机器翻译模型训练以解决GPU内存溢出问题?

问题描述

我正在使用PyTorch训练一个基于《Attention is All You Need》论文的标准机器翻译Transformer模型。本地PC上用标准超参数、batch size=128句对时,模型运行正常但速度慢,符合预期。但在搭载Tesla K80 GPU的AWS p2.xlarge实例上运行时,很快因GPU内存溢出崩溃。尝试各种释放内存的方法后,只能将batch size降至8,仍偶尔出现以下错误:

File "C:\Projects\MT004.venv\Lib\site-packages\torch\autograd\graph.py", line 744, in _engine_run_backward
return Variable._execution_engine.run_backward( # Calls into the C++ engine to run the backward pass
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
torch.cuda.OutOfMemoryError: CUDA out of memory. Tried to allocate 1.95 GiB. GPU

试过SpaCy分词器和XLM-R分词器,用XLM-R时batch size只能设为2,仍偶尔崩溃。程序在loss.backward()处崩溃,相关代码如下:

def train_epoch(src_train_sent, tgt_train_sent, model, optimizer):
    model.train()
    losses = 0

    torch.cuda.empty_cache()  # Clear cache before forward pass

    train_dataloader = SrcTgtIterable(src_train_sent, tgt_train_sent, batch_size=BATCH_SIZE, collate_fn=collate_fn)

    for src, tgt in train_dataloader:
        src = src.to(DEVICE)
        tgt = tgt.to(DEVICE)

        tgt_input = tgt[:-1, :]

        src_mask, tgt_mask, src_padding_mask, tgt_padding_mask = create_mask(src, tgt_input)

        logits = model(src, tgt_input, src_mask, tgt_mask, src_padding_mask, tgt_padding_mask, src_padding_mask)

        optimizer.zero_grad()

        tgt_out = tgt[1:, :].long()
        loss = loss_fn(logits.reshape(-1, logits.shape[-1]), tgt_out.reshape(-1))

        # Delete unnecessary variables before backward pass
        del src, tgt_input, src_mask, tgt_mask, src_padding_mask, tgt_padding_mask, logits, tgt_out
        torch.cuda.empty_cache()  # Clear cache after deleting variables

        loss.backward()

        optimizer.step()
        losses += loss.item()

        # Free GPU memory
        del loss
        torch.cuda.empty_cache()  # Clear cache after each batch

目前没有足够EC2配额升级服务器,请问哪里操作有误?有哪些优化建议?


操作错误分析
  • 无效的内存释放操作:反向传播依赖前向计算生成的中间张量,这些张量被autograd计算图持有,手动delPython变量无法释放GPU上的计算图内存,提前调用torch.cuda.empty_cache()也无济于事,反而增加额外开销。
  • 未限制输入序列长度:Transformer自注意力层的内存复杂度为O(n²)(n为序列token数),如果不对长句子做截断,单句就能占用大量GPU内存,尤其是XLM-R分词器生成的token数更多,内存压力呈指数级上升。

优化建议

1. 强制序列长度截断

在collate_fn中对超过固定长度(如128或256)的句子进行截断,确保src和tgt的序列长度一致。这是最有效的内存优化手段,直接降低自注意力层的内存消耗。

2. 梯度累积模拟大batch

无需直接调大batch size,通过累积小batch的梯度实现等效大batch的训练效果。例如batch size设为8,累积4次梯度再执行一次optimizer.step(),等效于batch size 32:

def train_epoch(src_train_sent, tgt_train_sent, model, optimizer, accumulate_steps=4):
    model.train()
    losses = 0
    optimizer.zero_grad()

    train_dataloader = SrcTgtIterable(src_train_sent, tgt_train_sent, batch_size=BATCH_SIZE, collate_fn=collate_fn)

    for idx, (src, tgt) in enumerate(train_dataloader):
        src = src.to(DEVICE)
        tgt = tgt.to(DEVICE)

        tgt_input = tgt[:-1, :]
        src_mask, tgt_mask, src_padding_mask, tgt_padding_mask = create_mask(src, tgt_input)

        logits = model(src, tgt_input, src_mask, tgt_mask, src_padding_mask, tgt_padding_mask, src_padding_mask)
        tgt_out = tgt[1:, :].long()
        loss = loss_fn(logits.reshape(-1, logits.shape[-1]), tgt_out.reshape(-1))
        loss = loss / accumulate_steps  # 平均梯度避免溢出

        loss.backward()

        if (idx + 1) % accumulate_steps == 0:
            optimizer.step()
            optimizer.zero_grad()

        losses += loss.item() * accumulate_steps

3. 启用混合精度训练

使用PyTorch的torch.cuda.amp自动混合精度,将部分张量从FP32转为FP16,内存占用减少约50%,同时提升训练速度:

from torch.cuda.amp import GradScaler, autocast

scaler = GradScaler()

def train_epoch(src_train_sent, tgt_train_sent, model, optimizer, accumulate_steps=4):
    model.train()
    losses = 0
    optimizer.zero_grad()

    train_dataloader = SrcTgtIterable(src_train_sent, tgt_train_sent, batch_size=BATCH_SIZE, collate_fn=collate_fn)

    for idx, (src, tgt) in enumerate(train_dataloader):
        src = src.to(DEVICE)
        tgt = tgt.to(DEVICE)

        tgt_input = tgt[:-1, :]
        src_mask, tgt_mask, src_padding_mask, tgt_padding_mask = create_mask(src, tgt_input)

        with autocast():
            logits = model(src, tgt_input, src_mask, tgt_mask, src_padding_mask, tgt_padding_mask, src_padding_mask)
            tgt_out = tgt[1:, :].long()
            loss = loss_fn(logits.reshape(-1, logits.shape[-1]), tgt_out.reshape(-1))
            loss = loss / accumulate_steps

        scaler.scale(loss).backward()

        if (idx + 1) % accumulate_steps == 0:
            scaler.step(optimizer)
            scaler.update()
            optimizer.zero_grad()

        losses += loss.item() * accumulate_steps

4. 精简模型结构

  • 减少encoder/decoder层数:将标准6层改为4层;
  • 降低隐藏层维度:从512降至256,内存占用会大幅降低;
  • 启用batch_first=True:修改nn.Transformer的参数为batch_first=True,减少内存拷贝开销,代码更直观。

5. 数据加载优化

  • 在collate_fn中确保张量连续:对拼接后的张量调用.contiguous(),避免内存碎片化;
  • 开启pin_memory:在dataloader中设置pin_memory=True,加快CPU到GPU的数据传输速度,减少内存碎片。

6. 其他细节优化

  • 用nvidia-smi检查GPU进程:kill掉无关进程释放内存;
  • 关闭不必要的梯度计算:对不需要训练的层用torch.no_grad()包裹(训练阶段作用有限)。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.22 03:04:54