RTX2070 8G显存微调SegFormer MIT-B0时CUDA内存不足排查
SegFormer MIT-B0微调CUDA内存不足问题解决
问题背景
在配备8GB显存的RTX 2070显卡上,基于卫星图像微调SegFormer MIT-B0语义分割模型实现稻田分割时,第一个epoch启动即触发CUDA out of memory错误。即使将batch size降至2,在Google Colab环境下仍复现该问题。
相关代码
主代码
import os import torch import argparse from tqdm import tqdm import albumentations as A from data import RiceDataset from torch import nn from torch.utils.data import DataLoader from transformers import SegformerImageProcessor, SegformerForSemanticSegmentation def train_step(model: nn.Module, dataloader: torch.utils.data.DataLoader, optimizer: torch.optim.Optimizer, device: torch.device): model.train() loss = 0.0 for i, (images, masks) in enumerate(dataloader): images, masks = images.to(device), masks.to(device) print(images.shape) print(masks.shape) optimizer.zero_grad() z = model(images, masks) loss += z.loss loss = loss / len(dataloader) return loss def eval_step(model: nn.Module, dataloader: torch.utils.data.DataLoader, device: torch.device): model.eval() loss = 0.0 with torch.inference_mode(): for i, (images, masks) in enumerate(dataloader): images, masks = images.to(device), masks.to(device) print(images.shape) print(masks.shape) z = model(images, masks) loss += z.loss loss = loss / len(dataloader) return loss def train_loop(dataset_loc: str = None, num_epochs: int = 1, batch_size: int = 4, num_workers: int = 10, model_path: str = None): train_images = os.path.join(dataset_loc, "images/train") train_masks = os.path.join(dataset_loc, "masks/train") list_of_train_images = os.listdir(train_images) val_images = os.path.join(dataset_loc, "images/val") val_masks = os.path.join(dataset_loc, "masks/val") list_of_val_images = os.listdir(val_images) train_transform = A.Compose([ A.HorizontalFlip(p=0.3), A.VerticalFlip(p=0.3), A.RandomRotate90(p=0.3), ]) model_checkpoint = "nvidia/mit-b0" processor = SegformerImageProcessor.from_pretrained(model_checkpoint) train_dataset = RiceDataset(images=list_of_train_images, image_folder=train_images, mask_folder=train_masks, transform=train_transform, processor=processor) val_dataset = RiceDataset(images=list_of_val_images, image_folder=val_images, mask_folder=val_masks, processor=processor) print(f'Train images: {len(train_dataset)}\nValidation images: {len(val_dataset)}') train_dataloader = DataLoader(train_dataset, batch_size=batch_size, num_workers=num_workers, shuffle=True) val_dataloader = DataLoader(val_dataset, batch_size=batch_size, num_workers=num_workers, shuffle=False) device = torch.device("cuda" if torch.cuda.is_available() else "cpu") id2label = {0: "outer", 1: "rice_paddy"} label2id = {label: id for id, label in id2label.items()} num_labels = len(id2label) model = SegformerForSemanticSegmentation.from_pretrained( model_checkpoint, num_labels=num_labels, id2label=id2label, label2id=label2id, ignore_mismatched_sizes=True, reshape_last_stage=True ) model.to(device) optimizer = torch.optim.AdamW(model.parameters(), lr=0.00006) scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=10) best_val_loss = float("inf") best_model_state_dict = None for epoch in tqdm(range(num_epochs)): train_loss = train_step(model=model, dataloader=train_dataloader, optimizer=optimizer, device=device) val_loss = eval_step(model=model, dataloader=val_dataloader, device=device) print( f"Epoch: {epoch+1} | " f"train_loss: {train_loss:.4f} | " f"val_loss: {val_loss:.4f} | " ) scheduler.step() if val_loss < best_val_loss: best_val_loss = val_loss best_model_state_dict = model.state_dict() if best_model_state_dict is not None: if model_path.endswith(".pth") or model_path.endswith(".pt"): torch.save(best_model_state_dict, model_path) else: torch.save(best_model_state_dict, model_path + ".pth") print(f"Best validation loss: {best_val_loss:.4f}") print("DONE")
数据集代码
class RiceDataset(torch.utils.data.Dataset): def __init__(self, images, image_folder, mask_folder, processor, transform=None): self.images = images self.image_folder = image_folder self.mask_folder = mask_folder self.processor = processor self.transform = transform def __len__(self): return len(self.images) def __getitem__(self, idx): image_name = self.images[idx] image_path = os.path.join(self.image_folder, image_name) mask_name = os.path.splitext(image_name)[0] + '.png' mask_path = os.path.join(self.mask_folder, mask_name) image = Image.open(image_path).convert("RGB") mask = Image.open(mask_path).convert("L") image_np = np.array(image) mask_np = np.array(mask) if self.transform: augmented = self.transform(image=image_np, mask=mask_np) image_np = augmented['image'] mask_np = augmented['mask'] encoded_inputs = self.processor(image_np, mask_np) return encoded_inputs["pixel_values"][0], encoded_inputs["labels"][0]
针对性解决方案
1. 缩小图像预处理尺寸
SegFormer默认处理器会将图像缩放到1024x1024,这对8GB显存压力过大。初始化处理器时指定更小的尺寸:
processor = SegformerImageProcessor.from_pretrained(model_checkpoint, size={"height": 512, "width": 512})
这能直接减少每个batch的内存占用,是最有效的内存优化手段之一。
2. 修复数据集代码逻辑错误
当前数据集的__getitem__方法存在逻辑漏洞:当没有transform时,encoded_inputs未定义会直接报错,同时预处理逻辑不统一。修改后:
def __getitem__(self, idx): image_name = self.images[idx] image_path = os.path.join(self.image_folder, image_name) mask_name = os.path.splitext(image_name)[0] + '.png' mask_path = os.path.join(self.mask_folder, mask_name) image = Image.open(image_path).convert("RGB") mask = Image.open(mask_path).convert("L") image_np = np.array(image) mask_np = np.array(mask) if self.transform: augmented = self.transform(image=image_np, mask=mask_np) image_np = augmented['image'] mask_np = augmented['mask'] # 统一执行processor预处理 encoded_inputs = self.processor(image_np, mask_np) return encoded_inputs["pixel_values"][0], encoded_inputs["labels"][0]
3. 优化训练流程的内存使用
- 梯度累积:若batch size=1仍OOM,用梯度累积模拟大batch(比如累积4次更新一次):
def train_step(model: nn.Module, dataloader: torch.utils.data.DataLoader, optimizer: torch.optim.Optimizer, device: torch.device, accumulate_steps: int = 4): model.train() loss = 0.0 optimizer.zero_grad() for i, (images, masks) in enumerate(dataloader): images, masks = images.to(device), masks.to(device) z = model(images, masks) step_loss = z.loss / accumulate_steps step_loss.backward() loss += z.loss.item() if (i + 1) % accumulate_steps == 0: optimizer.step() optimizer.zero_grad() if len(dataloader) % accumulate_steps != 0: optimizer.step() loss = loss / len(dataloader) return loss - 减少num_workers:
num_workers=10会导致CPU加载数据过多,挤占GPU内存,建议调整为2-4(根据CPU核心数)。 - 关闭训练时的打印:注释掉
print(images.shape)和print(masks.shape),减少内存额外开销。
4. 启用混合精度训练
利用PyTorch自动混合精度降低内存占用:
from torch.cuda.amp import GradScaler, autocast # 在train_loop中初始化scaler scaler = GradScaler() # 修改train_step def train_step(model: nn.Module, dataloader: torch.utils.data.DataLoader, optimizer: torch.optim.Optimizer, device: torch.device, scaler: GradScaler): model.train() loss = 0.0 for i, (images, masks) in enumerate(dataloader): images, masks = images.to(device), masks.to(device) optimizer.zero_grad() with autocast(): z = model(images, masks) step_loss = z.loss scaler.scale(step_loss).backward() scaler.step(optimizer) scaler.update() loss += step_loss.item() loss = loss / len(dataloader) return loss
验证阶段也可添加with autocast():减少内存占用。
5. 精简模型加载参数
初始化SegformerForSemanticSegmentation时,reshape_last_stage=True可能带来额外内存开销,若无需调整最后阶段形状可尝试移除该参数;同时确保模型所有参数正确移至GPU。
内容的提问来源于stack exchange,提问作者Miras Sagirbay
相关产品推荐
相关产品推荐

