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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.23 10:04:50