Segformer 학습 1

조훈·2025년 6월 17일
post-thumbnail

1. Segformer란?


  • SegFormer는 2021년 NVIDIA에서 발표한 semantic segmentation 모델로,
    Transformer 기반 백본lightweight decoder를 조합한 구조

  • 기존 segmentation 모델들과 달리, 복잡한 디코더나 파라미터 조정 없이도
    높은 정확도와 빠른 속도를 동시에 달성한 것이 특징


a ) 특징 요약

  • Hierarchical Transformer encoder
    : 다양한 해상도의 정보를 동시에 추출할 수 있어 전역 문맥을 잘 파악함.
  • MLP-based decoder
    : 기존의 복잡한 디코더 대신 간단한 MLP 구조를 사용해 연산 효율이 좋음.
  • Position embedding 없음
    : 이미지 크기에 유연하게 대응 가능.


b ) YOLO와 SegFormer 비교

항목YOLOv8 (segmentation 기준)SegFormer
주요 목적Object detection + instance segmentationSemantic segmentation
기반 구조CNN 중심 + 일부 Transformer 요소완전 Transformer 기반
출력 방식객체 단위(box/mask)로 분할픽셀 단위로 클래스 분류
속도빠름 (실시간 감지에 적합)비교적 느림 (큰 입력에선 더 무거움)
정확도작은 객체나 빠른 장면에 강점복잡한 배경이나 장면 이해에 강점
활용 분야자율주행 객체 인식, CCTV, 로봇비전 등거리 단위 예측, 장면 분할, 레이아웃 해석 등

간단하게 정리하면?

  • YOLO는 "이게 무슨 객체인지" 빠르게 잡아내는 데 강함 ( 실시간성 )
  • SegFormer는 "이 픽셀이 어떤 의미를 가지는지" 더 정교하게 구분함 ( 정확도 )

그래서 YOLO는 instance-level 문제에,
SegFormer는 pixel-level 문제에 더 적합하다고 볼 수 있습니다.


2. Segformer 데이터셋 구성

SegFormerDataset/
├── images/
│   └── train/
│       ├── img1.jpg
│       └── ...
├── masks/
│   └── train/
│       ├── img1.png 
│       └── …
  • mask.png의 각 픽셀은 Class_Index
  • Yolo와 다른 점은 0 : Background 라는 점
    0: 'background',
    1: 'Driving Area',
    2: 'Parking Line',

3. Segformer 학습

< Segformer_train.py >

import os
import torch
import numpy as np
import matplotlib.pyplot as plt
from PIL import Image
from sklearn.metrics import (
    precision_score, recall_score, f1_score, average_precision_score,
    precision_recall_curve
)
from torch.utils.data import Dataset
from transformers import (
    SegformerForSemanticSegmentation,
    SegformerImageProcessor,
    TrainingArguments,
    Trainer,
    TrainerCallback,
)
import torch.nn.functional as F

train_image_dir="/home/elicer/dataset/images/train"
train_mask_dir="/home/elicer/dataset/masks/train"
val_image_dir="/home/elicer/dataset/images/val"
val_mask_dir="/home/elicer/dataset/masks/val"
output_dir="/home/elicer/jhj/Segformer_Train/Final/result"
logging_dir="/home/elicer/jhj/Segformer_Train/Final/log"
plot_dir="/home/elicer/jhj/Segformer_Train/Final/plot"

os.makedirs(output_dir, exist_ok=True)
os.makedirs(logging_dir, exist_ok=True)
os.makedirs(plot_dir, exist_ok=True)

class SegFormerDataset(Dataset):
    def __init__(self, image_dir, mask_dir, processor, size=(512, 512)):
        self.image_dir = image_dir
        self.mask_dir = mask_dir
        self.processor = processor
        self.size = size
        self.image_files = sorted(os.listdir(image_dir))

    def __len__(self):
        return len(self.image_files)

    def __getitem__(self, idx):
        image_path = os.path.join(self.image_dir, self.image_files[idx])
        mask_path = os.path.join(self.mask_dir, self.image_files[idx].replace(".jpg", ".png"))

        image = Image.open(image_path).convert("RGB")
        mask = Image.open(mask_path)

        image = image.resize(self.size)
        mask = mask.resize(self.size, resample=Image.NEAREST)

        encoding = self.processor(images=image, return_tensors="pt")
        encoding["labels"] = torch.from_numpy(np.array(mask)).long()

        return {k: v.squeeze() for k, v in encoding.items()}

class MetricsPlotCallback(TrainerCallback):
    def __init__(self, save_dir, num_classes=7, class_names=None):
        self.save_dir = save_dir
        os.makedirs(save_dir, exist_ok=True)
        self.epochs = []
        self.metrics_history = {
            "precision": [],
            "recall": [],
            "f1": [],
            "mIoU": [],
            "mAP": []
        }
        self.num_classes = num_classes
        self.class_names = class_names if class_names else [f"Class {i}" for i in range(num_classes)]

        self.trainer = None

    def on_evaluate(self, args, state, control, metrics=None, **kwargs):
        epoch = int(state.epoch) if state.epoch is not None else len(self.epochs) + 1
        self.epochs.append(epoch)

        self.metrics_history["precision"].append(metrics.get("eval_precision", 0))
        self.metrics_history["recall"].append(metrics.get("eval_recall", 0))
        self.metrics_history["f1"].append(metrics.get("eval_f1", 0))
        self.metrics_history["mIoU"].append(metrics.get("eval_mIoU", 0))
        self.metrics_history["mAP"].append(metrics.get("eval_mAP", 0))

        plt.figure(figsize=(10, 6))
        for key, values in self.metrics_history.items():
            plt.plot(self.epochs, values, label=key)
        plt.xlabel("Epoch")
        plt.ylabel("Score")
        plt.title("SegFormer Metrics Over Epochs")
        plt.legend()
        plt.grid(True)
        save_path = os.path.join(self.save_dir, f"metrics_epoch_{epoch}.png")
        plt.savefig(save_path)
        plt.close()

        print(f"Generating mask PR curves for epoch {epoch}...")

        trainer = self.trainer
        if trainer is None:
            print("Trainer instance not set in callback; skipping PR curve.")
            return

        model = trainer.model
        model.eval()
        eval_dataset = trainer.eval_dataset
        device = trainer.args.device

        all_preds = []
        all_labels = []

        with torch.no_grad():
            for batch in eval_dataset:
                inputs = batch["pixel_values"].unsqueeze(0).to(device) if batch["pixel_values"].ndim == 3 else batch["pixel_values"].to(device)
                labels = batch["labels"].unsqueeze(0).to(device) if batch["labels"].ndim == 2 else batch["labels"].to(device)

                outputs = model(inputs)
                logits = outputs.logits  # (N, C, H, W)

                # Resize logits if needed
                if logits.shape[2:] != labels.shape[1:]:
                    logits = torch.nn.functional.interpolate(logits, size=labels.shape[1:], mode='bilinear', align_corners=False)

                probs = torch.softmax(logits, dim=1)
                all_preds.append(probs.cpu())
                all_labels.append(labels.cpu())

        all_preds = torch.cat(all_preds, dim=0).numpy()  # (N, C, H, W)
        all_labels = torch.cat(all_labels, dim=0).numpy()  # (N, H, W)

        plt.figure(figsize=(12, 8))
        for cls in range(self.num_classes):
            y_true = (all_labels == cls).astype(np.uint8).flatten()
            y_scores = all_preds[:, cls, :, :].reshape(-1)

            if np.sum(y_true) == 0:
                print(f"Class {cls} ({self.class_names[cls]}) skipped (no ground truth).")
                continue

            precision, recall, _ = precision_recall_curve(y_true, y_scores)
            ap = average_precision_score(y_true, y_scores)

            plt.plot(recall, precision, label=f"{self.class_names[cls]} (AP={ap:.2f})")

        plt.xlabel("Recall")
        plt.ylabel("Precision")
        plt.title(f"Mask PR Curves by Class - Epoch {epoch}")
        plt.legend(loc="lower left")
        plt.grid(True)

        pr_curve_path = os.path.join(self.save_dir, f"mask_pr_curve_epoch_{epoch}.png")
        plt.savefig(pr_curve_path)
        plt.close()
        print(f"Saved mask PR curve plot: {pr_curve_path}")

def compute_metrics(eval_pred):
    logits, labels = eval_pred

    if isinstance(logits, tuple):
        logits = logits[0]

    logits_tensor = torch.from_numpy(logits)
    labels_tensor = torch.from_numpy(labels)

    # Resize logits to match labels
    target_size = labels_tensor.shape[-2:]  # (H, W)
    logits_resized = F.interpolate(logits_tensor, size=target_size, mode="bilinear", align_corners=False)

    preds = torch.argmax(logits_resized, dim=1).numpy()
    labels = labels_tensor.numpy()

    preds_flat = preds.flatten()
    labels_flat = labels.flatten()

    num_classes = logits.shape[1]

    f1 = f1_score(labels_flat, preds_flat, average="macro", zero_division=0)
    precision = precision_score(labels_flat, preds_flat, average="macro", zero_division=0)
    recall = recall_score(labels_flat, preds_flat, average="macro", zero_division=0)

    ious = []
    for cls in range(num_classes):
        pred_inds = preds_flat == cls
        label_inds = labels_flat == cls
        intersection = np.logical_and(pred_inds, label_inds).sum()
        union = np.logical_or(pred_inds, label_inds).sum()
        if union == 0:
            iou = float('nan')
        else:
            iou = intersection / union
        ious.append(iou)
    miou = np.nanmean(ious)

    preds_onehot = np.eye(num_classes)[preds_flat]
    labels_onehot = np.eye(num_classes)[labels_flat]

    try:
        map_score = average_precision_score(labels_onehot, preds_onehot, average="macro")
    except ValueError:
        map_score = float('nan')

    return {
        "eval_f1": f1,
        "eval_precision": precision,
        "eval_recall": recall,
        "eval_mIoU": miou,
        "eval_mAP": map_score,
    }

id2label = {
    0: 'background',
    1: 'Driving Area',
    2: 'Parking Area',
}
label2id = {v: k for k, v in id2label.items()}

processor = SegformerImageProcessor(do_resize=False, do_normalize=True)
model = SegformerForSemanticSegmentation.from_pretrained(
    "nvidia/segformer-b0-finetuned-ade-512-512",
    num_labels=len(id2label),
    id2label=id2label,
    label2id=label2id,
    ignore_mismatched_sizes=True
)

train_dataset = SegFormerDataset(
    image_dir=train_image_dir,
    mask_dir=train_mask_dir,
    processor=processor,
    size=(512, 512)
)

val_dataset = SegFormerDataset(
    image_dir=val_image_dir,
    mask_dir=val_mask_dir,
    processor=processor,
    size=(512, 512)
)

training_args = TrainingArguments(
    output_dir=output_dir,
    per_device_train_batch_size=16,
    num_train_epochs=10,
    logging_dir=logging_dir,
    save_strategy="steps",
    save_steps=100,
    logging_steps=50,
    evaluation_strategy="epoch",
    #eval_steps=100,
    learning_rate=5e-5,
    save_total_limit=3,
    remove_unused_columns=False,
    fp16=torch.cuda.is_available(),
)

metrics_callback = MetricsPlotCallback(
    save_dir=plot_dir,
    num_classes=len(id2label),
    class_names=[id2label[i] for i in range(len(id2label))]
)

trainer = Trainer(
    model=model,
    args=training_args,
    train_dataset=train_dataset,
    eval_dataset=val_dataset,
    compute_metrics=compute_metrics,
    callbacks=[metrics_callback],
)

metrics_callback.trainer = trainer

trainer.train(resume_from_checkpoint=True)
trainer.save_model()

metrics = trainer.evaluate()
print(metrics)

a ) id2label 설정

# 4. id2label, label2id 설정
id2label = {
    0: 'background',
    1: 'Driving Area',
    2: 'Parking Area',
}
label2id = {v: k for k, v in id2label.items()}

b ) pretrained model 선정 및 index 부여

# 5. Processor 및 모델 불러오기
processor = SegformerImageProcessor(do_resize=False, do_normalize=True)
model = SegformerForSemanticSegmentation.from_pretrained(
    "nvidia/segformer-b0-finetuned-ade-512-512",
    num_labels=len(id2label),
    id2label=id2label,
    label2id=label2id,
    ignore_mismatched_sizes=True
)
  • pretrained model은 주로 b0 ~ b2를 사용하는데 b0가 가벼운모델, b2가 무거운 모델이다.

c ) 데이터셋 경로 설정

# 6. 데이터셋 생성
train_dataset = SegFormerDataset(
    image_dir=train_image_dir,
    mask_dir=train_mask_dir,
    processor=processor,
    size=(512, 512)
)

val_dataset = SegFormerDataset(
    image_dir=val_image_dir,
    mask_dir=val_mask_dir,
    processor=processor,
    size=(512, 512)
)

d ) Hyperparameter 설정

# 7. 학습 인자 설정
training_args = TrainingArguments(
    output_dir=output_dir,
    per_device_train_batch_size=16,
    num_train_epochs=10,
    logging_dir=logging_dir,
    save_strategy="steps",
    save_steps=100,
    logging_steps=50,
    evaluation_strategy="epoch",
    #eval_steps=100,
    learning_rate=5e-5,
    save_total_limit=3,
    remove_unused_columns=False,
    fp16=torch.cuda.is_available(),
)
  • num_train_epochs : epoch 수 조절
  • save_strategy / save_steps : 체크포인트 생성 기준 (epochs, steps 등 ) / steps로 설정 시 step수 설정

e ) 평가 지표

class MetricsPlotCallback
  • 학습 도중 evaluation_strategy에 따라 precision, recall, f1 score, mIoU, mAP 등 평가지표를 plot 이미지로 저장

4. 결과

#데이터셋 클래스 분포


a ) 10GB 테스트 데이터셋으로 학습한 Segformer 모델의 평가지표


b ) 55GB 테스트 데이터셋으로 학습한 Segformer 모델의 평가지표


c ) YOLOv8-seg 모델의 최고 성능 평가지표

적은량의 데이터지만 YOLOv8-seg보다 높은 성능을 보임
-> 과적합이 존재할 수도 있어 더 많은 데이터로 검증 할 예정

0개의 댓글