본문 바로가기
실전 개발 노트/개발 가이드

PyTorch U-Net 이미지 세그멘테이션: 구조·Dice Loss·반려동물 실습

by 쑥쑥자라나라 2026. 8. 1.
728x90

사진 분류는 “고양이인가?”에 답하지만 이미지 세그멘테이션은 고양이의 몸이 어느 픽셀까지인지 답한다. 귀 끝과 꼬리처럼 작은 경계를 살리려면 넓은 문맥과 정확한 위치 정보가 모두 필요하다. U-Net은 이미지를 줄이며 문맥을 모으고, 다시 키우면서 인코더의 같은 해상도 특징을 스킵 연결로 가져와 이 문제를 푼다.

세그멘테이션 결과는 클래스 하나가 아니라 입력과 같은 공간에 놓인 픽셀별 예측 마스크다.

1. 정답은 사진 한 장마다 붙은 trimap이다

Oxford-IIIT Pet 데이터셋은 37개 고양이·강아지 품종의 이미지 7,349장과 품종, 머리 영역, 픽셀 단위 trimap을 제공한다. trimap은 반려동물 몸, 배경, 몸의 경계와 목줄 같은 모호한 영역으로 구성된다. 이 글에서는 세 영역을 모두 예측하는 3클래스 문제로 다룬다.

from torchvision.datasets import OxfordIIITPet

train_data = OxfordIIITPet(
    root="data",
    split="trainval",
    target_types="segmentation",
    download=True
)

image, mask = train_data[0]
print(image.size, mask.size)
데이터를 보기 전에 클래스 값을 확인한다.

마스크 파일의 값은 손실 함수가 기대하는 0부터 C-1 범위로 바꿔야 한다. 값의 의미를 추측하지 말고 데이터셋 문서와 실제 unique() 결과를 확인한 뒤 매핑을 고정한다.

2. U-Net의 U자는 해상도가 내려갔다 올라오는 경로다

인코더가 모은 같은 해상도 특징을 디코더로 건너 보내면 위치와 경계 정보를 복원하기 쉬워진다.

구간역할공간 해상도확인할 실패
인코더모서리·질감에서 객체 문맥까지 특징 추출점차 감소작은 객체 정보 소실
병목가장 압축된 문맥 표현가장 작음모델 용량 과다와 과적합
디코더특징을 입력 크기의 마스크로 복원점차 증가흐린 경계와 체커보드 패턴
스킵 연결인코더의 위치 정보를 같은 단계 디코더에 결합양쪽 크기가 같아야 함패딩 차이로 텐서 크기 불일치

원 논문의 U-Net은 수축 경로와 대칭적인 확장 경로를 사용한다. 구현에서는 업샘플링 뒤 인코더 특징을 채널 방향으로 이어 붙인다. 스킵 연결은 단순히 층을 건너뛰는 지름길이 아니라, 깊은 층에서 잃기 쉬운 위치 정보를 출력 경계에 다시 제공하는 통로다.

3. 이미지와 마스크는 같은 기하 변환을 받아야 한다

이미지만 좌우 반전하고 마스크는 그대로 두면 픽셀 정답이 어긋난다. 반대로 색상 밝기처럼 사진에만 적용해야 하는 변환을 마스크에도 적용하면 클래스 번호가 깨진다. Torchvision transforms v2는 이미지와 Mask를 함께 다루는 방식을 제공한다.

이미지·마스크 읽기같은 크기 조절같은 기하 증강이미지만 정규화클래스 값 변환
마스크 크기 조절에는 최근접 보간을 쓴다.

선형·삼차 보간은 클래스 0과 2 사이에 1.3 같은 중간값을 새로 만든다. 사진에는 부드러운 보간을 쓸 수 있지만, 클래스 마스크는 InterpolationMode.NEAREST로 원래 번호를 보존한다.

728x90

4. PyTorch U-Net 최소 구현

import torch
from torch import nn
import torch.nn.functional as F

class DoubleConv(nn.Module):
    def __init__(self, in_ch, out_ch):
        super().__init__()
        self.block = nn.Sequential(
            nn.Conv2d(in_ch, out_ch, 3, padding=1),
            nn.BatchNorm2d(out_ch),
            nn.ReLU(inplace=True),
            nn.Conv2d(out_ch, out_ch, 3, padding=1),
            nn.BatchNorm2d(out_ch),
            nn.ReLU(inplace=True),
        )

    def forward(self, x):
        return self.block(x)

class UNetSmall(nn.Module):
    def __init__(self, num_classes=3):
        super().__init__()
        self.enc1 = DoubleConv(3, 32)
        self.enc2 = DoubleConv(32, 64)
        self.bridge = DoubleConv(64, 128)
        self.pool = nn.MaxPool2d(2)
        self.up2 = nn.ConvTranspose2d(128, 64, 2, stride=2)
        self.dec2 = DoubleConv(128, 64)
        self.up1 = nn.ConvTranspose2d(64, 32, 2, stride=2)
        self.dec1 = DoubleConv(64, 32)
        self.head = nn.Conv2d(32, num_classes, 1)

    def forward(self, x):
        e1 = self.enc1(x)
        e2 = self.enc2(self.pool(e1))
        b = self.bridge(self.pool(e2))

        d2 = self.up2(b)
        d2 = self.dec2(torch.cat([d2, e2], dim=1))
        d1 = self.up1(d2)
        d1 = self.dec1(torch.cat([d1, e1], dim=1))
        return self.head(d1)

padding=1인 3×3 합성곱과 2의 배수 입력 크기를 사용해 인코더·디코더 특징 크기를 맞춘 예제다. 실제 데이터 크기가 홀수이거나 여러 번 다운샘플링하면 한두 픽셀 차이가 날 수 있다. torch.cat() 전에 shape를 출력하고, 필요하면 입력 패딩 또는 F.interpolate()로 명시적으로 맞춘다.

5. 픽셀 정확도 하나로는 작은 경계를 평가하기 어렵다

배경이 이미지 대부분을 차지하면 모든 픽셀을 배경으로 예측해도 픽셀 정확도가 높게 나올 수 있다. Cross Entropy는 각 픽셀의 클래스 분류를 안정적으로 학습하고, Dice는 예측 영역과 정답 영역의 겹침을 본다. 둘을 함께 쓰면 클래스별 판별과 영역 겹침을 동시에 관찰할 수 있다.

def multiclass_dice_loss(logits, target, eps=1e-6):
    probs = logits.softmax(dim=1)
    one_hot = F.one_hot(
        target, num_classes=logits.shape[1]
    ).permute(0, 3, 1, 2).float()

    dims = (0, 2, 3)
    intersection = (probs * one_hot).sum(dims)
    denominator = probs.sum(dims) + one_hot.sum(dims)
    dice = (2 * intersection + eps) / (denominator + eps)
    return 1 - dice.mean()

ce_loss = nn.CrossEntropyLoss()
logits = model(images)                 # N, C, H, W
loss = ce_loss(logits, masks) + multiclass_dice_loss(logits, masks)

CrossEntropyLoss에 클래스 인덱스를 넣을 때 logit은 N×C×H×W, 정답은 N×H×Wlong 텐서여야 한다. 정답에 채널 차원을 억지로 남기거나 one-hot 정답과 클래스 인덱스를 섞으면 shape 오류나 잘못된 손실이 생긴다.

6. 학습 결과는 평균 점수와 실패 마스크를 함께 본다

model.eval()
with torch.inference_mode():
    logits = model(images.to(device))
    prediction = logits.argmax(dim=1).cpu()

# 반드시 같은 샘플의 원본·정답·예측을 나란히 저장
for image, target, pred in zip(images.cpu(), masks.cpu(), prediction):
    save_triplet(image, target, pred)

검증 세트에서 클래스별 Dice와 IoU를 계산하고, 최악의 샘플을 따로 저장한다. 평균값만 보면 귀 끝을 놓치거나 배경의 쿠션을 몸으로 오인하는 패턴이 가려진다. 품종, 밝기, 자세, 털색, 객체 크기별로 실패를 나누면 다음 데이터 증강과 모델 변경의 근거가 된다.

  1. 원본 이미지·정답 마스크·예측 마스크를 같은 크기로 저장한다.
  2. 배경·몸·경계의 클래스별 Dice와 IoU를 기록한다.
  3. 가장 낮은 점수와 가장 높은 점수의 샘플을 모두 본다.
  4. 임계값이나 후처리 전후를 같은 검증 세트에서 비교한다.
  5. 테스트 세트는 모델과 하이퍼파라미터 선택을 끝낸 뒤 평가한다.

7. 실수하기 쉬운 부분

  • 분류 라벨 37개와 세그멘테이션 trimap 3개를 같은 목표로 착각한다.
  • 이미지와 마스크에 서로 다른 랜덤 크롭·반전을 적용한다.
  • 마스크를 bilinear로 확대해 존재하지 않던 클래스 번호를 만든다.
  • CrossEntropyLoss 전에 softmax를 적용하거나 정답 dtype을 float로 둔다.
  • 스킵 연결 텐서의 높이·너비를 확인하지 않고 cat한다.
  • 배경이 많은 데이터에서 픽셀 정확도만 보고 좋은 모델이라고 판단한다.
  • 좋은 예측 이미지만 골라 보여 주고 경계·가림·작은 객체 실패를 숨긴다.
핵심 요약
  1. U-Net은 인코더로 문맥을 모으고 디코더와 스킵 연결로 위치를 복원한다.
  2. 세그멘테이션 변환은 이미지와 마스크의 공간 정렬을 반드시 유지해야 한다.
  3. Cross Entropy는 픽셀별 클래스, Dice는 영역 겹침을 관찰하는 데 유용하다.
  4. 배경 비율이 크면 정확도가 과대평가될 수 있어 클래스별 Dice·IoU가 필요하다.
  5. 평균 점수와 함께 실패 마스크를 저장해야 다음 개선 방향이 보인다.

참고한 1차 자료

728x90