
SegFormer는 2021년 NVIDIA에서 발표한 semantic segmentation 모델로,
Transformer 기반 백본과 lightweight decoder를 조합한 구조
기존 segmentation 모델들과 달리, 복잡한 디코더나 파라미터 조정 없이도
높은 정확도와 빠른 속도를 동시에 달성한 것이 특징
| 항목 | YOLOv8 (segmentation 기준) | SegFormer |
|---|---|---|
| 주요 목적 | Object detection + instance segmentation | Semantic segmentation |
| 기반 구조 | CNN 중심 + 일부 Transformer 요소 | 완전 Transformer 기반 |
| 출력 방식 | 객체 단위(box/mask)로 분할 | 픽셀 단위로 클래스 분류 |
| 속도 | 빠름 (실시간 감지에 적합) | 비교적 느림 (큰 입력에선 더 무거움) |
| 정확도 | 작은 객체나 빠른 장면에 강점 | 복잡한 배경이나 장면 이해에 강점 |
| 활용 분야 | 자율주행 객체 인식, CCTV, 로봇비전 등 | 거리 단위 예측, 장면 분할, 레이아웃 해석 등 |
그래서 YOLO는 instance-level 문제에,
SegFormer는 pixel-level 문제에 더 적합하다고 볼 수 있습니다.
SegFormerDataset/
├── images/
│ └── train/
│ ├── img1.jpg
│ └── ...
├── masks/
│ └── train/
│ ├── img1.png
│ └── …
0: 'background',
1: 'Driving Area',
2: 'Parking Line',
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)
# 4. id2label, label2id 설정
id2label = {
0: 'background',
1: 'Driving Area',
2: 'Parking Area',
}
label2id = {v: k for k, v in id2label.items()}
# 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
)
# 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)
)
# 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(),
)
class MetricsPlotCallback




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