마스크 생성

마스크 생성 (Mask generation)

마스크 생성은 이미지에 대해 의미론적으로 유의미한 마스크를 생성하는 작업입니다. 이 작업은 이미지 분할과 매우 유사하지만 많은 차이점이 존재합니다.

출처: 문서

본문

이미지 분할 모델은 라벨이 있는 데이터셋으로 훈련되고 훈련 중에 본 클래스로만 제한됩니다. 즉, 이미지가 주어지면 일련의 마스크와 해당 클래스를 반환합니다.

마스크 생성 모델은 대량의 데이터로 훈련되며 두 가지 모드로 작동합니다.

  • 프롬프팅 모드(Prompting mode): 이 모드에서 모델은 이미지와 프롬프트를 받습니다. 프롬프트는 객체 내부의 2D 포인트 위치(XY 좌표)이거나 객체를 감싸는 바운딩 박스일 수 있습니다. 프롬프팅 모드에서 모델은 프롬프트가 가리키는 객체 위의 마스크만 반환합니다.
  • Segment Everything 모드: segment everything에서는 이미지가 주어지면 모델이 이미지의 모든 마스크를 생성합니다. 이를 위해 포인트 그리드가 생성되어 추론을 위해 이미지 위에 겹쳐집니다.
  • 비디오 추론(Video Inference): 모델은 비디오와 비디오 프레임의 포인트 또는 박스 프롬프트를 받아 비디오 전체에서 추적합니다. 비디오 추론 방법에 대한 자세한 정보는 SAM 2 docs를 참고하세요.

마스크 생성 작업은 Segment Anything Model (SAM)과 Segment Anything Model 2 (SAM2)가 지원하며, 비디오 추론은 Segment Anything Model 2 (SAM2)가 지원합니다. SAM은 Vision Transformer 기반 이미지 인코더, 프롬프트 인코더, 양방향 트랜스포머 마스크 디코더로 구성된 강력한 모델입니다. 이미지와 프롬프트가 인코딩되고, 디코더가 이러한 임베딩을 받아 유효한 마스크를 생성합니다. 한편 SAM 2는 마스크를 추적하는 메모리 모듈을 추가해 SAM을 확장합니다.

SAM은 큰 데이터 커버리지를 가지므로 분할을 위한 강력한 파운데이션 모델 역할을 합니다. 100만 개의 이미지와 11억 개의 마스크를 가진 SA-1B 데이터셋으로 훈련되었습니다.

이 가이드에서 다룰 내용은 다음과 같습니다.

  • 배칭으로 segment everything 모드에서 추론하기,
  • 포인트 프롬프팅 모드에서 추론하기,
  • 박스 프롬프팅 모드에서 추론하기.

먼저 transformers를 설치하겠습니다.

pip install -q transformers

Mask Generation Pipeline

마스크 생성 모델을 추론하는 가장 쉬운 방법은 mask-generation pipeline을 사용하는 것입니다.

>>> from transformers import pipeline

>>> checkpoint = "facebook/sam2-hiera-base-plus"
>>> mask_generator = pipeline(model=checkpoint, task="mask-generation")

이미지를 살펴보겠습니다.

from PIL import Image
import requests

img_url = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/bee.jpg"
image = Image.open(requests.get(img_url, stream=True).raw).convert("RGB")

segment everything을 해 보겠습니다. points-per-batch는 segment everything 모드에서 포인트의 병렬 추론을 가능하게 합니다. 이는 추론을 더 빠르게 하지만 더 많은 메모리를 소비합니다. 또한 SAM은 이미지가 아닌 포인트에 대해서만 배칭을 가능하게 합니다. pred_iou_thresh는 IoU 신뢰도 임계값으로, 해당 임계값 이상의 마스크만 반환됩니다.

masks = mask_generator(image, points_per_batch=128, pred_iou_thresh=0.88)

masks는 다음과 같습니다.

{'masks': [tensor([[False, False, False,  ...,  True,  True,  True],
          [False, False, False,  ...,  True,  True,  True],
          [False, False, False,  ...,  True,  True,  True],
          ...,
          [False, False, False,  ..., False, False, False], .. 
 'scores': tensor([0.9874, 0.9793, 0.9780, 0.9776, ... 0.9016])}

다음과 같이 시각화할 수 있습니다.

import matplotlib.pyplot as plt

plt.imshow(image, cmap='gray')

for i, mask in enumerate(masks["masks"]):
    plt.imshow(mask, cmap='viridis', alpha=0.1, vmin=0, vmax=1)

plt.axis('off')
plt.show()

아래는 그레이스케일 원본 이미지에 다채로운 지도를 겹쳐 놓은 것입니다. 매우 인상적입니다.

모델 추론 (Model Inference)

포인트 프롬프팅 (Point Prompting)

pipeline 없이 모델을 사용할 수도 있습니다. 이를 위해 모델과 프로세서를 초기화합니다.

from transformers import SamModel, SamProcessor
from accelerate import Accelerator
import torch
device = Accelerator().device
model = SamModel.from_pretrained("facebook/sam-vit-base").to(device)
processor = SamProcessor.from_pretrained("facebook/sam-vit-base")

포인트 프롬프팅을 하려면 입력 포인트를 프로세서에 전달한 다음, 프로세서 출력을 가져와 추론을 위해 모델에 전달합니다. 모델 출력을 후처리하려면 출력과 프로세서의 초기 출력에서 가져온 original_sizes를 전달합니다. 프로세서가 이미지를 리사이즈하고 출력을 외삽해야 하므로 이것들을 전달해야 합니다.

input_points = [[[2592, 1728]]] # point location of the bee

inputs = processor(image, input_points=input_points, return_tensors="pt").to(device)
with torch.no_grad():
    outputs = model(**inputs)
masks = processor.image_processor.post_process_masks(outputs.pred_masks.cpu(), inputs["original_sizes"].cpu())

masks 출력의 세 마스크를 시각화할 수 있습니다.

import matplotlib.pyplot as plt
import numpy as np

fig, axes = plt.subplots(1, 4, figsize=(15, 5))

axes[0].imshow(image)
axes[0].set_title('Original Image')
mask_list = [masks[0][0][0].numpy(), masks[0][0][1].numpy(), masks[0][0][2].numpy()]

for i, mask in enumerate(mask_list, start=1):
    overlayed_image = np.array(image).copy()

    overlayed_image[:,:,0] = np.where(mask == 1, 255, overlayed_image[:,:,0])
    overlayed_image[:,:,1] = np.where(mask == 1, 0, overlayed_image[:,:,1])
    overlayed_image[:,:,2] = np.where(mask == 1, 0, overlayed_image[:,:,2])
    
    axes[i].imshow(overlayed_image)
    axes[i].set_title(f'Mask {i}')
for ax in axes:
    ax.axis('off')

plt.show()

박스 프롬프팅 (Box Prompting)

포인트 프롬프팅과 비슷한 방식으로 박스 프롬프팅도 할 수 있습니다. [x_min, y_min, x_max, y_max] 형태의 입력 박스를 이미지와 함께 processor에 전달하면 됩니다. 프로세서 출력을 가져와 직접 모델에 전달한 다음, 다시 후처리하면 됩니다.

# bounding box around the bee
box = [2350, 1600, 2850, 2100]

inputs = processor(
        image,
        input_boxes=[[[box]]],
        return_tensors="pt"
    ).to(model.device)

with torch.no_grad():
    outputs = model(**inputs)

mask = processor.image_processor.post_process_masks(
    outputs.pred_masks.cpu(),
    inputs["original_sizes"].cpu(),
)[0][0][0].numpy()

아래와 같이 벌을 감싸는 바운딩 박스를 시각화할 수 있습니다.

import matplotlib.patches as patches

fig, ax = plt.subplots()
ax.imshow(image)

rectangle = patches.Rectangle((2350, 1600), 500, 500, linewidth=2, edgecolor='r', facecolor='none')
ax.add_patch(rectangle)
ax.axis("off")
plt.show()

아래에서 추론 출력을 볼 수 있습니다.

fig, ax = plt.subplots()
ax.imshow(image)
ax.imshow(mask, cmap='viridis', alpha=0.4)

ax.axis("off")
plt.show()

마스크 생성을 위한 파인튜닝

이미지 매팅(matting)을 위해 MicroMat 데이터셋의 작은 부분에서 SAM2.1을 파인튜닝하겠습니다. DICE loss를 사용하려면 monai 라이브러리를 설치하고, 훈련 중 마스크를 기록하기 위해 trackio가 필요합니다.

pip install -q datasets monai trackio

이제 데이터셋을 로드해 살펴보겠습니다.

from datasets import load_dataset

dataset = load_dataset("merve/MicroMat-mini", split="train")
dataset
# Dataset({
#    features: ['image', 'mask', 'prompt', 'image_id', 'object_id', 'sample_idx', 'granularity', 
# 'image_path', 'mask_path', 'prompt_path'],  num_rows: 94
#})

image, mask, prompt 컬럼이 필요합니다. train과 test로 나눕니다.

dataset = dataset.train_test_split(test_size=0.1)
train_ds = dataset["train"]
val_ds = dataset["test"]

샘플을 하나 살펴보겠습니다.

train_ds[0]
 {'image': <PIL.PngImagePlugin.PngImageFile image mode=RGB size=2040x1356>,
 'mask': <PIL.PngImagePlugin.PngImageFile image mode=L size=2040x1356>,
 'prompt': '{"point": [[137, 1165, 1], [77, 1273, 0], [58, 1351, 0]], "bbox": [0, 701, 251, 1356]}',
 'image_id': '0034',
 'object_id': '34',
 'sample_idx': 1,
 'granularity': 'fine',
 'image_path': '/content/MicroMat-mini/img/0034.png',
 'mask_path': '/content/MicroMat-mini/mask/0034_34.png',
 'prompt_path': '/content/MicroMat-mini/prompt/0034_34.json'}

프롬프트는 딕셔너리 문자열이므로 아래와 같이 바운딩 박스를 얻을 수 있습니다.

import json

json.loads(train_ds["prompt"][0])["bbox"]
# [0, 701, 251, 1356]

예시 이미지, 프롬프트, 마스크를 시각화합니다.

import matplotlib.pyplot as plt
import numpy as np

def show_mask(mask, ax):
    color = np.array([0.12, 0.56, 1.0, 0.6])
    mask = np.array(mask)
    h, w = mask.shape
    mask_image = mask.reshape(h, w, 1) * color.reshape(1, 1, 4)
    ax.imshow(mask_image)
    x0, y0, x1, y1 = eval(train_ds["prompt"][0])["bbox"]
    ax.add_patch(
        plt.Rectangle((x0, y0), x1 - x0, y1 - y0,
                      fill=False, edgecolor="lime", linewidth=2))

example = train_ds[0]
image = np.array(example["image"])
ground_truth_mask = np.array(example["mask"])

fig, ax = plt.subplots()
ax.imshow(image)
show_mask(ground_truth_mask, ax)
ax.set_title("Ground truth mask")
ax.set_axis_off()

plt.show() 

이제 데이터를 로드하기 위한 데이터셋을 정의할 수 있습니다. SAMDataset은 우리 데이터셋을 감싸고 각 샘플을 SAM 프로세서가 기대하는 방식으로 형식화합니다. 따라서 원시 이미지와 마스크 대신, 훈련 준비가 된 처리된 이미지, 바운딩 박스, 그라운드 트루스 마스크를 얻을 수 있습니다.

기본적으로 프로세서는 이미지를 리사이즈하므로 이미지와 마스크 외에도 원본 크기도 반환합니다. 또한 마스크는 [0, 255] 값을 가지므로 이진화해야 합니다.

from torch.utils.data import Dataset
import torch

class SAMDataset(Dataset):
  def __init__(self, dataset, processor):
    self.dataset = dataset
    self.processor = processor

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

  def __getitem__(self, idx):
    item = self.dataset[idx]
    image = item["image"]
    prompt = eval(item["prompt"])["bbox"]
    inputs = self.processor(image, input_boxes=[[prompt]], return_tensors="pt")
    inputs["ground_truth_mask"] = (np.array(item["mask"]) > 0).astype(np.float32)
    inputs["original_image_size"] = torch.tensor(image.size[::-1])

    return inputs

프로세서와 데이터셋을 그것으로 초기화할 수 있습니다.

from transformers import Sam2Processor

processor = Sam2Processor.from_pretrained("facebook/sam2.1-hiera-small")
train_dataset = SAMDataset(dataset=train_ds, processor=processor)

다양한 크기의 그라운드 트루스 마스크를 같은 모양의 재구성된 마스크 배치로 바꿔줄 데이터 콜레이터를 정의해야 합니다. 최근접 이웃 보간을 사용해 재구성합니다. 배치의 나머지 요소에 대해서도 배칭된 텐서를 만듭니다. 마스크가 모두 같은 크기라면 이 단계를 건너뛰어도 됩니다.

import torch.nn.functional as F

def collate_fn(batch, target_hw=(256, 256)):

    pixel_values = torch.cat([item["pixel_values"] for item in batch], dim=0)
    original_sizes = torch.stack([item["original_sizes"] for item in batch])
    input_boxes = torch.cat([item["input_boxes"] for item in batch], dim=0)
    ground_truth_masks = torch.cat([
        F.interpolate(
            torch.as_tensor(x["ground_truth_mask"]).unsqueeze(0).unsqueeze(0).float(),
            size=(256, 256),
            mode="nearest"
        )
        for x in batch
    ], dim=0).long()

    return {
        "pixel_values": pixel_values,
        "original_sizes": original_sizes,
        "input_boxes": input_boxes,
        "ground_truth_mask": ground_truth_masks,
        "original_image_size": torch.stack([item["original_image_size"] for item in batch]),
    }

from torch.utils.data import DataLoader
train_dataloader = DataLoader(
    train_dataset,
    batch_size=4,
    shuffle=True,
    collate_fn=collate_fn,
)

데이터 로더가 무엇을 생성하는지 살펴보겠습니다.

batch = next(iter(train_dataloader))
for k,v in batch.items():
  print(k,v.shape)

# pixel_values torch.Size([4, 3, 1024, 1024])
# original_sizes torch.Size([4, 1, 2])
# input_boxes torch.Size([4, 1, 4])
# ground_truth_mask torch.Size([4, 1, 256, 256])
#original_image_size torch.Size([4, 2])

이제 모델을 로드하고 마스크 디코더만 훈련하도록 비전과 프롬프트 인코더를 동결하겠습니다.

from transformers import Sam2Model

model = Sam2Model.from_pretrained("facebook/sam2.1-hiera-small")

for name, param in model.named_parameters():
  if name.startswith("vision_encoder") or name.startswith("prompt_encoder"):
    param.requires_grad_(False)

이제 옵티마이저와 loss 함수를 정의할 수 있습니다.

from torch.optim import Adam
import monai

optimizer = Adam(model.mask_decoder.parameters(), lr=1e-5, weight_decay=0)
seg_loss = monai.losses.DiceCELoss(sigmoid=True, squared_pred=True, reduction='mean')

훈련 전에 모델이 어떻게 수행되는지 확인해 보겠습니다.

import matplotlib.pyplot as plt

item = val_ds[1]
img = item["image"]
bbox = json.loads(item["prompt"])["bbox"]
inputs = processor(images=img, input_boxes=[[bbox]], return_tensors="pt").to(model.device)

with torch.no_grad():
  outputs = model(**inputs)

masks = processor.post_process_masks(outputs.pred_masks.cpu(), inputs["original_sizes"])[0]
preds = masks.squeeze(0)
mask = (preds[0] > 0).cpu().numpy()

overlay = np.asarray(img, dtype=np.uint8).copy()
overlay[mask] = 0.55 * overlay[mask] + 0.45 * np.array([0, 255, 0], dtype=np.float32)

plt.imshow(overlay)
plt.axis("off")
plt.show()

SAM2 result after training

훈련 중간에 모델 개선을 모니터링할 수 있도록 예측을 trackio에 기록해야 합니다.

from PIL import Image
import trackio
import json

@torch.no_grad()
def predict_fn(img, bbox):

  inputs = processor(images=img, input_boxes=[[bbox]], return_tensors="pt").to(model.device)

  with torch.no_grad():
      outputs = model(**inputs)

  masks = processor.post_process_masks(outputs.pred_masks.cpu(), inputs["original_sizes"])[0]
  return masks

def log_eval_masks_trackio(dataset, indices, step, predict_fn,  project=None, sample_cap=8):
    logs = {"eval/step": int(step)}
    for idx in indices[:sample_cap]:
        item = dataset[idx] 
        img = item["image"]
        bbox = json.loads(item["prompt"])["bbox"]
        preds = predict_fn(img, bbox)
        preds = preds.squeeze(0)
        mask = (preds[0] > 0).cpu().numpy()  

        overlay = np.asarray(img, dtype=np.uint8).copy()
        overlay[mask] = 0.55 * overlay[mask] + 0.45 * np.array([0, 255, 0], dtype=np.float32)
        logs[f"{idx}/overlay"] = trackio.Image(overlay, caption="overlay")
        
    trackio.log(logs)

이제 훈련 루프를 작성하고 훈련할 수 있습니다!

loss와 평가 마스크를 trackio로 기록하는 방식을 주목하세요.

from tqdm import tqdm
from statistics import mean
import trackio
import torch

num_epochs = 30

device = torch.accelerator.current_accelerator().type if torch.accelerator.is_available() else "cpu"
model.to(device)

model.train()
trackio.init(project="mask-eval")
for epoch in range(num_epochs):
    epoch_losses = []
    for batch in tqdm(train_dataloader):
      outputs = model(pixel_values=batch["pixel_values"].to(device),
                      input_boxes=batch["input_boxes"].to(device),
                      multimask_output=False)

      predicted_masks = outputs.pred_masks.squeeze(1)
      ground_truth_masks = batch["ground_truth_mask"].float().to(device)
      loss = seg_loss(predicted_masks, ground_truth_masks)

      optimizer.zero_grad()
      loss.backward()

      optimizer.step()
      epoch_losses.append(loss.item())
      
    log_eval_masks_trackio(dataset=val_ds, indices=[0, 3, 6, 9], step=epoch, predict_fn=predict_fn, project="mask-eval")
    print(f'Epoch: {epoch}')
    print(f'Mean loss: {mean(epoch_losses)}')
    trackio.log({"loss": mean(epoch_losses)})

trackio.finish()

훈련된 모델을 테스트해 보겠습니다.

import matplotlib.pyplot as plt

item = val_ds[1]
img = item["image"]
bbox = json.loads(item["prompt"])["bbox"]

inputs = processor(images=img, input_boxes=[[bbox]], return_tensors="pt").to(model.device)

with torch.no_grad():
  outputs = model(**inputs)

preds = processor.post_process_masks(outputs.pred_masks.cpu(), inputs["original_sizes"])[0]

preds = preds.squeeze(0)
mask = (preds[0] > 0).cpu().numpy()

overlay = np.asarray(img, dtype=np.uint8).copy()
overlay[mask] = 0.55 * overlay[mask] + 0.45 * np.array([0, 255, 0], dtype=np.float32)

plt.imshow(overlay)
plt.axis("off")
plt.show()

작은 데이터셋에서 20 에폭만 훈련했는데도 큰 개선이 있었습니다!

SAM2 result after training

더 알아보기 (Learn more)