PyTorch 전이 학습 (Transfer Learning)

전이 학습(Transfer Learning)은 기존에 학습된 모델의 가중치를 새로운 문제에 활용하는 방법입니다. 일반적으로 이미지 분류, 객체 탐지 등에서 많이 사용되며, 특히 작은 데이터셋을 사용할 때 유용합니다.

PyTorch에서는 사전 학습된 모델을 쉽게 활용할 수 있도록 torchvision.models 모듈에서 다양한 모델을 제공합니다.


1. 전이 학습의 개념

전이 학습은 다음 두 가지 방법 중 하나로 사용할 수 있습니다:

  1. Feature Extraction (특징 추출):
  2. 사전 학습된 모델의 가중치를 고정(freeze)하고, 마지막 분류 계층만 새롭게 학습.
  3. 일반적으로 새로운 데이터셋의 클래스에 맞게 마지막 fully connected 계층만 수정.

  4. Fine-tuning (미세 조정):

  5. 사전 학습된 모델의 일부 또는 전체 가중치를 초기화하지 않고 학습.
  6. 학습률(Learning Rate)을 조정하여 기존의 가중치와 새로 학습하는 가중치가 잘 조화를 이루도록 함.

2. PyTorch에서 전이 학습 구현

2.1. 라이브러리 임포트

import torch
import torch.nn as nn
import torch.optim as optim
from torchvision import datasets, transforms, models

2.2. 데이터 준비

데이터셋을 불러오고 전처리합니다. 예제에서는 CIFAR-10 데이터셋을 사용합니다.

# 데이터 변환 및 전처리
transform = transforms.Compose([
    transforms.Resize((224, 224)),  # 사전 학습된 모델은 일반적으로 224x224 크기를 사용
    transforms.ToTensor(),
    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])  # ImageNet 평균과 표준편차
])

# CIFAR-10 데이터셋
train_dataset = datasets.CIFAR10(root='./data', train=True, transform=transform, download=True)
test_dataset = datasets.CIFAR10(root='./data', train=False, transform=transform, download=True)

train_loader = torch.utils.data.DataLoader(train_dataset, batch_size=32, shuffle=True)
test_loader = torch.utils.data.DataLoader(test_dataset, batch_size=32, shuffle=False)

2.3. 사전 학습된 모델 불러오기

torchvision.models에서 사전 학습된 모델을 가져옵니다.

# Pre-trained ResNet18 모델 불러오기
model = models.resnet18(pretrained=True)  # ImageNet으로 학습된 가중치

2.4. Feature Extraction

특징 추출 방식으로 학습하려면 모델의 모든 가중치를 고정합니다.

# 모든 가중치 고정
for param in model.parameters():
    param.requires_grad = False

# 마지막 분류 계층 수정 (CIFAR-10은 클래스가 10개)
num_features = model.fc.in_features  # ResNet18의 마지막 계층 입력 크기
model.fc = nn.Linear(num_features, 10)  # 새로운 분류 계층 정의

2.5. Fine-tuning

Fine-tuning 방식을 사용하려면 특정 계층만 학습 가능하도록 설정합니다.

# Conv 레이어는 고정하지 않고, 마지막 분류 계층만 변경
num_features = model.fc.in_features
model.fc = nn.Linear(num_features, 10)  # CIFAR-10 클래스에 맞게 수정

# 필요한 레이어만 학습할 수 있도록 requires_grad 설정
for name, param in model.named_parameters():
    if "fc" in name:  # fc 계층만 학습
        param.requires_grad = True
    else:
        param.requires_grad = False

2.6. 손실 함수 및 옵티마이저

criterion = nn.CrossEntropyLoss()  # 분류 문제
optimizer = optim.Adam(model.fc.parameters(), lr=0.001)  # 마지막 계층만 학습

2.7. 학습 루프

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = model.to(device)

epochs = 5
for epoch in range(epochs):
    model.train()
    running_loss = 0.0

    for images, labels in train_loader:
        images, labels = images.to(device), labels.to(device)

        # Forward pass
        outputs = model(images)
        loss = criterion(outputs, labels)

        # Backward pass
        optimizer.zero_grad()
        loss.backward()
        optimizer.step()

        running_loss += loss.item()

    print(f"Epoch [{epoch+1}/{epochs}], Loss: {running_loss/len(train_loader):.4f}")

2.8. 평가

model.eval()
correct = 0
total = 0

with torch.no_grad():
    for images, labels in test_loader:
        images, labels = images.to(device), labels.to(device)
        outputs = model(images)
        _, predicted = torch.max(outputs, 1)
        total += labels.size(0)
        correct += (predicted == labels).sum().item()

print(f"Test Accuracy: {100 * correct / total:.2f}%")

3. 전이 학습의 장점

  1. 적은 데이터로도 학습 가능:
  2. 사전 학습된 모델의 가중치를 활용해 데이터 부족 문제를 해결.

  3. 효율적인 학습:

  4. 초기 가중치를 사용하므로 학습 속도가 빠름.

  5. 일반화 능력 향상:

  6. ImageNet과 같은 대규모 데이터셋에서 학습된 모델은 일반화 성능이 뛰어남.

4. 활용 가능한 사전 학습 모델

PyTorch는 다양한 사전 학습된 모델을 제공합니다: - ResNet: models.resnet18, models.resnet50 - VGG: models.vgg16, models.vgg19 - EfficientNet: models.efficientnet_b0 - DenseNet: models.densenet121 - MobileNet: models.mobilenet_v2

이 모델들은 pretrained=True 옵션으로 쉽게 사용할 수 있습니다.


5. 요약

  • 전이 학습은 기존 모델의 가중치를 활용하여 새로운 문제를 해결하는 강력한 방법.
  • PyTorch에서 전이 학습은 사전 학습된 모델(torchvision.models)과 간단한 코드로 구현 가능.
  • Feature Extraction 또는 Fine-tuning을 통해 데이터셋 크기와 문제 유형에 따라 유연하게 적용 가능.