如何优化机器翻译模型训练以解决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

