GeForce GTX 3060Ti训练大语言模型调参仍显存不足求助
解决Transformer训练显存不足问题
一、优化模型结构,降低显存占用
- 减少模型层数:当前设置
n_layer=20,对于显存有限的GPU来说层数过多,建议先下调至4-6层,后续再根据剩余显存逐步调整。 - 启用梯度检查点:在Transformer Block中添加梯度检查点,减少反向传播时的显存占用,修改
Block类的forward方法:def forward(self, x): x = x + torch.utils.checkpoint.checkpoint(self.sa, self.ln1(x)) x = x + torch.utils.checkpoint.checkpoint(self.ffwd, self.ln2(x)) return x - 简化输出层逻辑:确保训练时仅保留必要的前向传播计算,关闭生成无关的冗余逻辑。
二、优化数据加载与训练流程
- 修复目标张量生成逻辑:当前自回归任务的目标张量生成错误,会导致无效显存占用,修改训练循环中的数据处理部分:
xb_input = xb[:, :-1].to(device) # 输入取前block_size-1个token yb_target = xb[:, 1:].to(device) # 目标取后block_size-1个token logits, loss = model(xb_input, yb_target) - 下调评估迭代次数:当前
eval_iters=200会在评估阶段占用大量显存,建议降至50以内,减少评估时的显存消耗。 - 统一数据加载方式:代码中同时存在
DataLoader和get_batch两种数据加载逻辑,会导致冗余显存占用,建议统一使用DataLoader处理训练与评估数据。
三、启用PyTorch显存优化工具
- 开启混合精度训练:使用
torch.cuda.amp自动混合精度,降低浮点运算的显存占用,修改训练循环:from torch.cuda.amp import GradScaler, autocast scaler = GradScaler() accumulation_steps = 8 optimizer = torch.optim.AdamW(model.parameters(), lr=learning_rate) for iter in range(max_iters): if iter % eval_interval == 0 or iter == max_iters - 1: losses = estimate_loss() print(f"step {iter}: train loss {losses['train']:.4f}, val loss {losses['val']:.4f}") total_loss = 0.0 optimizer.zero_grad(set_to_none=True) for i, xb in enumerate(data_loader): xb_input = xb[:, :-1].to(device) yb_target = xb[:, 1:].to(device) with autocast(): logits, loss = model(xb_input, yb_target) loss = loss / accumulation_steps scaler.scale(loss).backward() total_loss += loss.item() if (i + 1) % accumulation_steps == 0: scaler.step(optimizer) scaler.update() optimizer.zero_grad(set_to_none=True) total_loss /= len(data_loader) print(f"Iteration {iter}: average training loss {total_loss:.4f}") save_model(model, pre_trained_model_path) - 主动清理显存:在训练循环的合适节点,删除无用张量并清理显存:
del xb_input, yb_target, logits, loss torch.cuda.empty_cache()
四、优化数据预处理
- 压缩词汇表大小:当前基于全量文本生成的词汇表可能过大,导致词嵌入层占用大量显存,可过滤低频词限制词汇表规模:
from collections import Counter word_counts = Counter(words) # 仅保留出现次数≥2的词汇 vocab = sorted([word for word, cnt in word_counts.items() if cnt >= 2], key=lambda x: word_counts[x], reverse=True) vocab_size = len(vocab) - 避免预加载全量数据:修改
CustomDataset,不预先将全量数据加载到内存,而是按需读取分片:class CustomDataset(Dataset): def __init__(self, file_path, block_size, stoi): self.file_path = file_path self.block_size = block_size self.stoi = stoi with open(file_path, 'r', encoding='utf-8') as f: self.words = re.findall(r'\w+|[^\w\s]', f.read()) def __len__(self): return len(self.words) - self.block_size def __getitem__(self, idx): chunk = self.words[idx:idx+self.block_size] return torch.tensor([self.stoi[word] for word in chunk], dtype=torch.long)
内容的提问来源于stack exchange,提问作者Auto
相关产品推荐
相关产品推荐

