이 문서는 딥러닝 기반 얼굴 인식에서 특징 구별을 강화하기 위해 제안된 중심 손실(Center Loss) 함수의 원리를 이해하고, PyTorch를 활용한 구현 과정을 상세히 설명합니다.
개방 집합 문제(Open-set Problem)와 특징 학습
패턴 인식 분야에는 크게 두 가지 유형의 문제가 있습니다. 훈련 세트와 테스트 세트의 클래스가 완전히 동일한 '폐쇄 집합(closed-set) 문제'와, 훈련 세트에 없는 미지의 클래스가 테스트 세트에 포함될 수 있는 '개방 집합(open-set) 문제'입니다.
예를 들어, 0부터 9까지의 숫자만 학습한 시스템이 알파벳 문자를 입력받았을 때, 이를 훈련 데이터의 특정 숫자와 유사하다고 잘못 분류하는 대신 '알 수 없음'으로 거부할 수 있어야 합니다. 특히 얼굴 인식과 같이 새로운 인물이 지속적으로 등장할 수 있는 시나리오에서는 이러한 개방 집합 문제 해결 능력이 중요합니다.
Softmax 손실의 한계
일반적인 이미지 분류 작업에서 널리 사용되는 Softmax 손실 함수는 각 클래스에 대한 예측 확률을 최대화하는 데 중점을 둡니다. MNIST 데이터셋을 예로 들어보면, Softmax 손실만으로 학습된 모델의 최종 특징 분포는 다음과 같을 수 있습니다 (2차원 특징 벡터로 시각화한 경우).

위 그림의 왼쪽은 훈련 세트, 오른쪽은 테스트 세트의 특징 분포입니다. 분류 성능 자체는 양호하지만, 각 클래스 간의 경계가 다소 모호하여 클래스 내 특징들이 넓게 퍼져 있음을 알 수 있습니다. 이는 새로운, 이전에 보지 못한 데이터를 만났을 때 잘못 분류될 가능성을 높입니다.
중심 손실 함수(Center Loss)의 도입
Softmax 손실의 한계를 극복하고 클래스 내 특징의 응집력을 높여 특징 구별 능력을 향상시키기 위해, 논문 저자들은 중심 손실 함수를 제안했습니다. 이 아이디어는 각 클래스의 특징들이 해당 클래스의 '중심(center)'으로 더욱 가깝게 모이도록 유도하여 클래스 내 분산을 줄이는 것입니다.
Softmax 함수 및 교차 엔트로피 손실
먼저 Softmax 함수와 교차 엔트로피 손실에 대해 간략히 살펴보겠습니다.
Softmax 함수는 신경망의 출력 값을 확률 분포로 변환하여, 각 클래스에 속할 확률의 합이 1이 되도록 만듭니다:

여기서 \(z_i\)는 \(i\)번째 출력 노드의 값이고, \(K\)는 총 클래스 수입니다.
Softmax 기반의 분류 문제에서는 주로 교차 엔트로피 손실(Cross-Entropy Loss)을 사용합니다. 모델이 올바른 클래스를 예측할 확률을 최대화하는 방향으로 학습을 진행하며, 이는 다음 공식으로 표현됩니다:

여기서 \(m\)은 미니 배치 내 샘플 수, \(n\)은 클래스 수, \(x_i\)는 \(i\)번째 샘플의 심층 특징 벡터, \(y_i\)는 \(i\)번째 샘플의 실제 클래스 레이블입니다.
중심 손실의 핵심 아이디어
중심 손실은 분류 정확도를 유지하면서도 각 클래스 내 특징들의 응집도를 높이는 것을 목표로 합니다. 이를 위해, 각 샘플의 특징 벡터와 해당 클래스의 중심 벡터 간의 유클리드 거리 제곱을 최소화하도록 학습합니다.

여기서 \(x_i\)는 \(i\)번째 샘플의 특징 벡터이고, \(c_{y_i}\)는 \(y_i\) 클래스의 중심 벡터입니다. \(m\)은 미니 배치 내 샘플의 수입니다. 이 식은 본질적으로 각 클래스 특징들이 해당 클래스의 평균적인 위치(중심)에 가까워지도록 유도하는 군집화 문제와 유사합니다.
유클리드 거리 대신 유클리드 거리의 제곱을 사용하는 이유는 다음과 같습니다:
- 제곱 연산은 계산이 더 빠르며, 제곱근 연산을 생략할 수 있습니다.
- 거리의 차이를 더 민감하게 반영하여 이상치를 쉽게 식별할 수 있습니다.
- 분산, 공분산 등 다른 제곱 항과 결합하기 용이합니다.
최종 손실 함수는 Softmax 손실과 중심 손실을 결합한 형태로 정의됩니다:

여기서 \(\lambda\)는 두 손실 함수의 중요도를 조절하는 하이퍼파라미터입니다. \(\frac{1}{2}\)는 주로 미분 시 제곱 항을 상쇄하여 계산을 간편하게 하기 위해 사용됩니다.
중심 손실 함수 PyTorch 구현
중심 손실 함수의 PyTorch 구현은 다음과 같습니다. 클래스별 중심 벡터는 학습 가능한 파라미터로 정의되며, 각 배치에서 손실이 계산됩니다.
import torch
import torch.nn as nn
class FeatureCenterLoss(nn.Module):
"""
중심 손실 함수 구현.
참고: Wen et al. A Discriminative Feature Learning Approach for Deep Face Recognition. ECCV 2016.
"""
def __init__(self, class_count=10, feature_dims=2, use_accelerator=True):
super(FeatureCenterLoss, self).__init__()
self.num_classes = class_count
self.feature_dim = feature_dims
self.use_gpu = use_accelerator
# 클래스 중심 벡터를 학습 가능한 파라미터로 초기화
# torch.randn은 평균 0, 분산 1의 정규 분포에서 값을 추출합니다.
if self.use_gpu:
self.class_centroids = nn.Parameter(torch.randn(self.num_classes, self.feature_dim).cuda())
else:
self.class_centroids = nn.Parameter(torch.randn(self.num_classes, self.feature_dim))
def forward(self, features, labels):
"""
순방향 전파.
Args:
features: 특징 행렬 (배치_크기, 특징_차원)
labels: 실제 레이블 (배치_크기)
"""
batch_size = features.size(0)
# 각 특징 벡터의 제곱 합 (x^2) 계산 및 확장
sq_feature_sum = torch.pow(features, 2).sum(dim=1, keepdim=True).expand(batch_size, self.num_classes)
# 각 중심 벡터의 제곱 합 (c^2) 계산 및 확장
sq_centroid_sum = torch.pow(self.class_centroids, 2).sum(dim=1, keepdim=True).expand(self.num_classes, batch_size).t()
# 유클리드 거리 제곱 계산: ||x - c||^2 = ||x||^2 + ||c||^2 - 2x^T c
distance_matrix = sq_feature_sum + sq_centroid_sum
distance_matrix.addmm_(1, -2, features, self.class_centroids.t())
# 각 샘플에 해당하는 클래스 중심과의 거리만 마스킹
label_range = torch.arange(self.num_classes).long()
if self.use_gpu:
label_range = label_range.cuda()
# 레이블을 확장하여 마스크 생성
expanded_labels = labels.unsqueeze(1).expand(batch_size, self.num_classes)
mask = expanded_labels.eq(label_range.expand(batch_size, self.num_classes))
# 마스크를 적용하여 해당 클래스의 거리만 남김
masked_distances = distance_matrix * mask.float()
# 작은 값으로 클리핑하여 수치적 안정성 확보 후, 손실 계산 (배치 내 평균)
loss = masked_distances.clamp(min=1e-12, max=1e+12).sum() / batch_size
return loss
이 구현은 미니 배치 내의 각 샘플에 대해 해당 샘플의 특징 벡터와 실제 클래스에 해당하는 중심 벡터 간의 유클리드 거리 제곱을 계산합니다. 계산된 거리들을 모두 더한 후 배치 크기로 나누어 평균 손실을 산출합니다.
코드 실행 예시
간단한 PyTorch 텐서를 사용하여 `FeatureCenterLoss`의 동작을 확인해볼 수 있습니다.
import torch
from center_loss_module import FeatureCenterLoss # 위에서 정의한 클래스
# FeatureCenterLoss 인스턴스 생성: 10개 클래스, 2차원 특징 벡터, GPU 사용 안 함
center_loss_fn = FeatureCenterLoss(class_count=10, feature_dims=2, use_accelerator=False)
torch.manual_seed(0) # 재현성을 위한 랜덤 시드 설정
# 가상의 특징 벡터 (10개 샘플, 각 2차원)
sample_features = torch.randn(10, 2)
# 가상의 레이블 (0-9 사이의 정수, 10개 샘플)
sample_labels = torch.randint(0, 10, (10,))
print("샘플 특징 벡터:\n", sample_features)
print("샘플 레이블:\n", sample_labels)
# 중심 손실 계산
calculated_loss = center_loss_fn(sample_features, sample_labels)
print("계산된 중심 손실:", calculated_loss)
샘플 특징 벡터:
tensor([[-1.1258, -1.1524],
[-0.2506, -0.4339],
[ 0.5988, -1.5551],
[-0.3414, 1.8530],
[ 0.4681, -0.1577],
[ 1.4437, 0.2660],
[ 1.3894, 1.5863],
[ 0.9463, -0.8437],
[ 0.9318, 1.2590],
[ 2.0050, 0.0537]])
샘플 레이블:
tensor([2, 9, 1, 8, 8, 3, 6, 9, 1, 7])
계산된 중심 손실: tensor(2.9281, grad_fn=<DivBackward0>)
모델 아키텍처: LeNet++
중심 손실 논문에서 제안된 모델은 LeNet++이라고 불리며, LeNet 구조를 기반으로 더 깊은 계층과 PReLU 활성화 함수를 사용하는 변형입니다. MNIST (28x28x1) 데이터셋에 대한 이 모델의 구조는 다음과 같습니다.

여기서는 3개의 컨볼루션 블록이 각각 컨볼루션 계층 2개와 PReLU 활성화 함수, 그리고 최대 풀링(Max Pooling) 계층으로 구성됩니다. 마지막에는 특징 벡터를 추출하는 선형 계층과 최종 분류를 위한 선형 계층이 연결됩니다.
import torch
import torch.nn as nn
from torch.nn import functional as F
class ModifiedConvNet(nn.Module):
"""
Center Loss 논문에서 설명된 LeNet++ 모델 구조.
"""
def __init__(self, num_classes):
super(ModifiedConvNet, self).__init__()
# 첫 번째 컨볼루션 블록: (28x28x1) -> (28x28x32) -> (14x14x32)
self.conv_block1_1 = nn.Conv2d(1, 32, kernel_size=5, stride=1, padding=2)
self.act_block1_1 = nn.PReLU()
self.conv_block1_2 = nn.Conv2d(32, 32, kernel_size=5, stride=1, padding=2)
self.act_block1_2 = nn.PReLU()
# 두 번째 컨볼루션 블록: (14x14x32) -> (14x14x64) -> (7x7x64)
self.conv_block2_1 = nn.Conv2d(32, 64, kernel_size=5, stride=1, padding=2)
self.act_block2_1 = nn.PReLU()
self.conv_block2_2 = nn.Conv2d(64, 64, kernel_size=5, stride=1, padding=2)
self.act_block2_2 = nn.PReLU()
# 세 번째 컨볼루션 블록: (7x7x64) -> (7x7x128) -> (3x3x128) (정수 나눗셈)
self.conv_block3_1 = nn.Conv2d(64, 128, kernel_size=5, stride=1, padding=2)
self.act_block3_1 = nn.PReLU()
self.conv_block3_2 = nn.Conv2d(128, 128, kernel_size=5, stride=1, padding=2)
self.act_block3_2 = nn.PReLU()
# 특징 추출을 위한 완전 연결 계층 (embedding layer)
# 최종 컨볼루션 출력 크기: 128 * 3 * 3 = 1152
self.feature_layer = nn.Linear(128 * 3 * 3, 2) # 2차원 특징 벡터 추출
self.act_feature = nn.PReLU()
# 최종 분류를 위한 완전 연결 계층
self.classifier_layer = nn.Linear(2, num_classes)
def forward(self, x_input):
# 첫 번째 블록
x_input = self.act_block1_1(self.conv_block1_1(x_input))
x_input = self.act_block1_2(self.conv_block1_2(x_input))
x_input = F.max_pool2d(x_input, kernel_size=2) # 크기 절반으로 축소
# 두 번째 블록
x_input = self.act_block2_1(self.conv_block2_1(x_input))
x_input = self.act_block2_2(self.conv_block2_2(x_input))
x_input = F.max_pool2d(x_input, kernel_size=2) # 크기 절반으로 축소
# 세 번째 블록
x_input = self.act_block3_1(self.conv_block3_1(x_input))
x_input = self.act_block3_2(self.conv_block3_2(x_input))
x_input = F.max_pool2d(x_input, kernel_size=2) # 크기 절반으로 축소
# 특징 벡터 추출을 위해 평탄화
flat_features = x_input.view(-1, 128 * 3 * 3)
extracted_features = self.act_feature(self.feature_layer(flat_features)) # 2차원 특징
# 분류 결과 생성
classification_output = self.classifier_layer(extracted_features)
return extracted_features, classification_output
# 모델 팩토리 함수 (모델 인스턴스 생성 간소화)
model_registry = {
'custom_cnn': ModifiedConvNet,
}
def create_model_instance(name, num_classes):
if name not in model_registry:
raise KeyError(f"알 수 없는 모델: {name}")
return model_registry[name](num_classes)
if __name__ == '__main__':
# 모델 인스턴스 생성 및 출력 예시
sample_model = create_model_instance('custom_cnn', 10)
print(sample_model)
torch.manual_seed(0)
dummy_input = torch.randn(1, 1, 28, 28) # 단일 28x28 흑백 이미지
features_out, logits_out = sample_model(dummy_input)
print("추출된 특징 벡터:", features_out)
print("분류 로짓:", logits_out)
데이터셋 준비 및 로딩
MNIST 데이터셋을 PyTorch의 `torchvision` 라이브러리를 사용하여 로드하고 전처리하는 과정을 설명합니다. 이미지 전처리에는 텐서 변환과 정규화가 포함됩니다.
import torch
import torchvision
from torch.utils.data import DataLoader
from torchvision import transforms # transforms 모듈을 직접 임포트
class MNISTDataLoader:
def __init__(self, batch_size, use_gpu, worker_count):
# 데이터 전처리 파이프라인 정의
data_transform = transforms.Compose([
transforms.ToTensor(), # PIL Image나 numpy.ndarray를 Tensor로 변환하고 [0, 1]로 정규화
transforms.Normalize((0.1307,), (0.3081,)) # (평균, 표준편차)로 정규화
])
# GPU 사용 시 pin_memory 활성화로 데이터 전송 속도 향상
pin_memory_enabled = True if use_gpu else False
# MNIST 훈련 데이터셋 로드
training_data = torchvision.datasets.MNIST(
root='./data/mnist', # 데이터 저장 경로
train=True, # 훈련 세트 로드
download=True, # 데이터셋이 없으면 다운로드
transform=data_transform # 정의된 전처리 적용
)
self.training_loader = DataLoader(
training_data,
batch_size=batch_size,
shuffle=True, # 훈련 시 데이터 섞기
num_workers=worker_count,
pin_memory=pin_memory_enabled,
)
# MNIST 테스트 데이터셋 로드
test_data = torchvision.datasets.MNIST(
root='./data/mnist',
train=False, # 테스트 세트 로드
download=True,
transform=data_transform
)
self.test_loader = DataLoader(
test_data,
batch_size=batch_size,
shuffle=False, # 테스트 시 데이터는 섞지 않음
num_workers=worker_count,
pin_memory=pin_memory_enabled,
)
self.num_classes = 10 # MNIST는 10개 클래스 (0-9)
# 데이터셋 팩토리 함수
dataset_registry = {
'mnist_data': MNISTDataLoader,
}
def prepare_dataset(name, batch_size, use_gpu, worker_count):
if name not in dataset_registry:
raise KeyError(f"알 수 없는 데이터셋: {name}")
return dataset_registry[name](batch_size, use_gpu, worker_count)
if __name__ == '__main__':
# 데이터 로더 생성 및 샘플 확인 예시
data_source = prepare_dataset('mnist_data', batch_size=10, use_gpu=False, worker_count=0)
print("데이터 로더 객체:", data_source)
for i, (input_images, target_labels) in enumerate(data_source.training_loader):
print("입력 이미지 텐서 형태:", input_images.shape) # (batch_size, channels, height, width)
print("타겟 레이블 텐서:", target_labels)
break # 첫 번째 배치만 확인
데이터 로더 객체: <__main__.MNISTDataLoader object at 0x...>
입력 이미지 텐서 형태: torch.Size([10, 1, 28, 28])
타겟 레이블 텐서: tensor([6, 9, 8, 3, 0, 3, 8, 2, 7, 3])
유틸리티 함수
훈련 과정에서 유용하게 사용되는 몇 가지 보조 함수들입니다. 디렉토리 생성, 평균 값 추적, 체크포인트 저장, 로그 기록 등의 기능을 제공합니다.
import os
import sys
import errno
import shutil
import os.path as osp
import torch
def create_directory_if_not_exists(path):
"""
지정된 경로에 디렉토리가 없으면 생성합니다.
"""
if not osp.exists(path):
try:
os.makedirs(path)
except OSError as e:
if e.errno != errno.EEXIST: # 디렉토리가 이미 존재하는 오류가 아니면 예외 발생
raise
class MetricAverager(object):
"""
값의 현재 상태와 평균을 계산하고 저장합니다.
"""
def __init__(self):
self.reset()
def reset(self):
self.current_value = 0
self.average_value = 0
self.total_sum = 0
self.item_count = 0
def update(self, val, n=1):
self.current_value = val
self.total_sum += val * n
self.item_count += n
self.average_value = self.total_sum / self.item_count
def store_model_state(state_dict, is_best, file_path='checkpoint.pth.tar'):
"""
모델의 상태를 체크포인트 파일로 저장합니다.
가장 좋은 모델인 경우, 별도의 'best_model.pth.tar'로도 저장합니다.
"""
create_directory_if_not_exists(osp.dirname(file_path))
torch.save(state_dict, file_path)
if is_best:
shutil.copy(file_path, osp.join(osp.dirname(file_path), 'best_model.pth.tar'))
class FileAndConsoleLogger(object):
"""
콘솔 출력 내용을 외부 텍스트 파일에도 기록합니다.
"""
def __init__(self, log_file_path=None):
self.console_output = sys.stdout
self.file_output = None
if log_file_path is not None:
create_directory_if_not_exists(os.path.dirname(log_file_path))
self.file_output = open(log_file_path, 'w')
def __del__(self):
self.close()
def __enter__(self):
pass
def __exit__(self, *args):
self.close()
def write(self, message):
self.console_output.write(message)
if self.file_output is not None:
self.file_output.write(message)
def flush(self):
self.console_output.flush()
if self.file_output is not None:
self.file_output.flush()
os.fsync(self.file_output.fileno()) # 파일 버퍼를 디스크에 동기화
def close(self):
self.console_output.close()
if self.file_output is not None:
self.file_output.close()
if __name__ == '__main__':
# 로거 및 Averager 사용 예시
# 'test_logs' 디렉토리에 'log_example.txt'로 로그를 저장하도록 표준 출력 리디렉션
log_dir = 'test_logs'
create_directory_if_not_exists(log_dir)
sys.stdout = FileAndConsoleLogger(osp.join(log_dir, 'log_example.txt'))
metric_tracker = MetricAverager()
metric_tracker.update(1, 10) # 값 1을 10번 추가
print(f"평균 값: {metric_tracker.average_value}")
print(f"현재 값: {metric_tracker.current_value}")
print(f"누적 합: {metric_tracker.total_sum}")
print(f"항목 수: {metric_tracker.item_count}")
sys.stdout.write('이것은 로그 메시지입니다.\n')
# 표준 출력 복원 (선택 사항, 스크립트 종료 시 자동으로 닫힘)
# sys.stdout = sys.__stdout__
실행하면 `test_logs` 디렉토리가 생성되고 그 안에 `log_example.txt` 파일에 출력 내용이 기록됩니다.
훈련 및 평가 과정
주요 훈련 스크립트는 `argparse` 모듈을 사용하여 명령줄 인수를 처리하고, `matplotlib`을 사용하여 훈련 중 특징 분포를 시각화합니다. 또한, PyTorch의 옵티마이저와 스케줄러를 활용하여 모델 파라미터와 중심 손실 파라미터를 업데이트합니다.
특징 분포 시각화 함수
훈련 과정에서 모델이 학습하는 2차원 특징 벡터의 분포를 시각화하여 각 클래스 특징들이 어떻게 군집되는지 확인할 수 있습니다.
import matplotlib.pyplot as plt
import os
import os.path as osp
import torch
import numpy as np # concatenate를 위해 numpy 필요
def visualize_features(feature_vectors, corresponding_labels, num_classes, epoch_num, output_prefix, output_dir):
"""
2차원 평면에 특징 벡터를 플로팅합니다.
Args:
feature_vectors: 특징 행렬 (샘플_수, 특징_차원).
corresponding_labels: 각 샘플의 레이블 (샘플_수).
num_classes: 총 클래스 수.
epoch_num: 현재 에폭 번호 (파일 이름에 사용).
output_prefix: 저장될 디렉토리의 접두사 (예: 'train' 또는 'test').
output_dir: 시각화 이미지가 저장될 기본 디렉토리.
"""
colors = ['C0', 'C1', 'C2', 'C3', 'C4', 'C5', 'C6', 'C7', 'C8', 'C9'] # 10가지 색상
plt.figure(figsize=(8, 8)) # 플롯 크기 설정
for label_idx in range(num_classes):
# 특정 레이블에 해당하는 특징 벡터만 추출하여 플롯
plt.scatter(
feature_vectors[corresponding_labels == label_idx, 0], # X 좌표
feature_vectors[corresponding_labels == label_idx, 1], # Y 좌표
c=colors[label_idx], # 클래스별 색상
s=1, # 점의 크기
label=str(label_idx) # 범례에 사용될 레이블
)
plt.legend(loc='upper right') # 범례 표시 위치
plt.title(f"Features Distribution at Epoch {epoch_num+1}")
plt.xlabel("Feature Dimension 1")
plt.ylabel("Feature Dimension 2")
# 이미지 저장 디렉토리 생성
save_path = osp.join(output_dir, output_prefix)
if not osp.exists(save_path):
os.makedirs(save_path)
image_filename = osp.join(save_path, f'epoch_{epoch_num+1}.png')
plt.savefig(image_filename, bbox_inches='tight') # 이미지 저장, 여백 잘라내기
plt.close() # 현재 플롯 닫기 (메모리 절약)
if __name__ == '__main__':
# 시각화 함수 사용 예시
torch.manual_seed(0)
sample_features = torch.randn(100, 2) # 100개 샘플의 2차원 특징
sample_labels = torch.randint(0, 10, (100,)) # 100개 샘플의 0-9 레이블
visualize_features(sample_features, sample_labels, 10, 0, 'sample_plots', 'output_visuals')
print("특징 시각화 이미지가 'output_visuals/sample_plots' 디렉토리에 저장되었습니다.")
실행하면 `output_visuals/sample_plots` 디렉토리에 `epoch_1.png` 파일이 생성됩니다.
훈련 루프 (`train` 함수)
모델 훈련 함수는 교차 엔트로피 손실과 중심 손실을 결합하여 사용합니다. 각 배치에서 특징과 로짓을 계산하고, 두 손실을 합산하여 역전파를 수행합니다. 옵티마이저는 모델 파라미터와 중심 파라미터를 별도로 업데이트합니다.
import torch
import torch.nn as nn
import numpy as np
import time
import datetime
import argparse
from torch.optim import lr_scheduler
# 다른 모듈에서 정의된 클래스와 함수 임포트 (가정)
from models import create_model_instance as create_network_model
from datasets import prepare_dataset as get_dataset
from center_loss_module import FeatureCenterLoss
from utils import MetricAverager, FileAndConsoleLogger, store_model_state
from visuals import visualize_features # 위에서 정의한 visualize_features 함수
# ArgumentParser 설정 (main 함수 내에서 초기화될 args 객체를 가정)
parser = argparse.ArgumentParser("Center Loss Example")
parser.add_argument('-d', '--dataset', type=str, default='mnist_data', choices=['mnist_data'])
parser.add_argument('-j', '--workers', default=4, type=int, help="데이터 로딩 워커 수")
parser.add_argument('--batch-size', type=int, default=128, help="미니 배치 크기")
parser.add_argument('--lr-model', type=float, default=0.001, help="모델 학습률")
parser.add_argument('--lr-cent', type=float, default=0.5, help="중심 손실 학습률")
parser.add_argument('--weight-cent', type=float, default=1.0, help="중심 손실 가중치")
parser.add_argument('--max-epoch', type=int, default=100, help="최대 훈련 에폭 수")
parser.add_argument('--stepsize', type=int, default=20, help="학습률 감소 주기 (에폭 단위)")
parser.add_argument('--gamma', type=float, default=0.5, help="학습률 감소 비율")
parser.add_argument('--model-type', type=str, default='custom_cnn', help="사용할 모델 종류")
parser.add_argument('--eval-freq', type=int, default=10, help="평가 주기 (에폭 단위)")
parser.add_argument('--print-freq', type=int, default=50, help="로그 출력 주기 (배치 단위)")
parser.add_argument('--gpu-id', type=str, default='0', help="사용할 GPU ID")
parser.add_argument('--seed', type=int, default=1, help="랜덤 시드")
parser.add_argument('--use-cpu', action='store_true', help="CPU만 사용 여부")
parser.add_argument('--log-dir', type=str, default='training_logs', help="로그 및 결과 저장 디렉토리")
parser.add_argument('--plot-features', action='store_true', help="에폭마다 특징 분포 시각화 여부")
args = parser.parse_args([]) # 스크립트 실행 시 인수가 없으면 기본값 사용 (실제 실행 시에는 주석 처리)
def train_epoch(network_model, xent_criterion, center_criterion,
model_optimizer, center_optimizer,
data_loader, gpu_available, num_classes, current_epoch):
network_model.train() # 모델을 훈련 모드로 설정
xent_loss_tracker = MetricAverager() # 교차 엔트로피 손실 추적
center_loss_tracker = MetricAverager() # 중심 손실 추적
total_loss_tracker = MetricAverager() # 총 손실 추적
if args.plot_features:
all_features_batch, all_labels_batch = [], []
for batch_idx, (batch_data, batch_labels) in enumerate(data_loader):
if gpu_available:
batch_data, batch_labels = batch_data.cuda(), batch_labels.cuda()
# 순방향 전파: 특징 벡터와 최종 분류 로짓 얻기
extracted_features, classification_outputs = network_model(batch_data)
# 손실 계산
loss_xent = xent_criterion(classification_outputs, batch_labels) # 분류 손실
loss_cent = center_criterion(extracted_features, batch_labels) # 중심 손실
weighted_loss_cent = loss_cent * args.weight_cent # 중심 손실에 가중치 적용
total_combined_loss = loss_xent + weighted_loss_cent # 총 손실
# 옵티마이저의 기존 기울기 초기화
model_optimizer.zero_grad()
center_optimizer.zero_grad()
# 역전파를 통해 기울기 계산
total_combined_loss.backward()
# 모델 파라미터 업데이트
model_optimizer.step()
# 중심 손실 함수의 기울기 조정 후 중심 파라미터 업데이트
# args.weight_cent는 중심 손실의 학습에 영향을 주지 않도록 기울기 스케일링
for param in center_criterion.parameters():
if param.grad is not None:
param.grad.data *= (1.0 / args.weight_cent)
center_optimizer.step()
# 손실 추적기 업데이트
total_loss_tracker.update(total_combined_loss.item(), batch_labels.size(0))
xent_loss_tracker.update(loss_xent.item(), batch_labels.size(0))
center_loss_tracker.update(weighted_loss_cent.item(), batch_labels.size(0))
if args.plot_features:
if gpu_available:
all_features_batch.append(extracted_features.data.cpu().numpy())
all_labels_batch.append(batch_labels.data.cpu().numpy())
else:
all_features_batch.append(extracted_features.data.numpy())
all_labels_batch.append(batch_labels.data.numpy())
# 일정 주기마다 훈련 진행 상황 출력
if (batch_idx + 1) % args.print_freq == 0:
print(f"Batch {batch_idx+1}/{len(data_loader)}\t "
f"Total Loss {total_loss_tracker.current_value:.6f} ({total_loss_tracker.average_value:.6f}) "
f"XentLoss {xent_loss_tracker.current_value:.6f} ({xent_loss_tracker.average_value:.6f}) "
f"CenterLoss {center_loss_tracker.current_value:.6f} ({center_loss_tracker.average_value:.6f})")
if args.plot_features:
# 모든 특징과 레이블을 결합하여 시각화
combined_features = np.concatenate(all_features_batch, axis=0)
combined_labels = np.concatenate(all_labels_batch, axis=0)
visualize_features(combined_features, combined_labels, num_classes, current_epoch,
output_prefix='train_plots', output_dir=args.log_dir)
평가 루프 (`test` 함수)
모델 평가 함수는 훈련된 모델의 성능을 측정합니다. `torch.no_grad()` 컨텍스트를 사용하여 불필요한 기울기 계산을 비활성화하고, 정확도와 오류율을 계산합니다. 훈련과 마찬가지로 특징 시각화를 포함할 수 있습니다.
def evaluate_model(network_model, data_loader, gpu_available, num_classes, current_epoch):
network_model.eval() # 모델을 평가 모드로 설정
correct_predictions, total_samples = 0, 0
if args.plot_features:
all_features_batch, all_labels_batch = [], []
with torch.no_grad(): # 기울기 계산 비활성화 (메모리 절약, 속도 향상)
for batch_data, batch_labels in data_loader:
if gpu_available:
batch_data, batch_labels = batch_data.cuda(), batch_labels.cuda()
extracted_features, classification_outputs = network_model(batch_data)
# 예측값 가져오기 (가장 높은 확률을 가진 클래스의 인덱스)
predicted_labels = classification_outputs.data.max(1)[1]
total_samples += batch_labels.size(0)
correct_predictions += (predicted_labels == batch_labels.data).sum().item()
if args.plot_features:
if gpu_available:
all_features_batch.append(extracted_features.data.cpu().numpy())
all_labels_batch.append(batch_labels.data.cpu().numpy())
else:
all_features_batch.append(extracted_features.data.numpy())
all_labels_batch.append(batch_labels.data.numpy())
if args.plot_features:
# 모든 특징과 레이블을 결합하여 시각화
combined_features = np.concatenate(all_features_batch, axis=0)
combined_labels = np.concatenate(all_labels_batch, axis=0)
visualize_features(combined_features, combined_labels, num_classes, current_epoch,
output_prefix='test_plots', output_dir=args.log_dir)
accuracy = correct_predictions * 100.0 / total_samples
error_rate = 100.0 - accuracy
return accuracy, error_rate
메인 실행 함수 (`main` 함수)
전체 훈련 및 평가 과정을 조율하는 메인 함수입니다. 데이터셋 로드, 모델 및 손실 함수 초기화, 옵티마이저 및 스케줄러 설정, 그리고 훈련 루프 실행을 담당합니다.
import torch
import torch.nn as nn
import torch.optim as optim
import torch.backends.cudnn as cudnn
from torch.optim import lr_scheduler
import os
import sys
import time
import datetime
import argparse
import os.path as osp
# 위에 정의된 함수 및 클래스 임포트
# from models import create_model_instance
# from datasets import prepare_dataset
# from center_loss_module import FeatureCenterLoss
# from utils import MetricAverager, FileAndConsoleLogger, store_model_state
# from visuals import visualize_features
# from train_eval_functions import train_epoch, evaluate_model # 위에서 정의한 train_epoch, evaluate_model
# ArgumentParser 설정
parser = argparse.ArgumentParser("Center Loss Example")
# Dataset
parser.add_argument('-d', '--dataset', type=str, default='mnist_data', choices=['mnist_data'])
parser.add_argument('-j', '--workers', default=4, type=int, help="데이터 로딩 워커 수 (기본값: 4)")
# Optimization
parser.add_argument('--batch-size', type=int, default=128, help="미니 배치 크기 (기본값: 128)")
parser.add_argument('--lr-model', type=float, default=0.001, help="모델 학습률 (기본값: 0.001)")
parser.add_argument('--lr-cent', type=float, default=0.5, help="중심 손실 학습률 (기본값: 0.5)")
parser.add_argument('--weight-cent', type=float, default=1.0, help="중심 손실 가중치 (기본값: 1.0)")
parser.add_argument('--max-epoch', type=int, default=100, help="최대 훈련 에폭 수 (기본값: 100)")
parser.add_argument('--stepsize', type=int, default=20, help="학습률 감소 주기 (에폭 단위, 기본값: 20)")
parser.add_argument('--gamma', type=float, default=0.5, help="학습률 감소 비율 (기본값: 0.5)")
# Model
parser.add_argument('--model-type', type=str, default='custom_cnn', help="사용할 모델 종류 (기본값: custom_cnn)")
# Misc
parser.add_argument('--eval-freq', type=int, default=10, help="평가 주기 (에폭 단위, 기본값: 10)")
parser.add_argument('--print-freq', type=int, default=50, help="로그 출력 주기 (배치 단위, 기본값: 50)")
parser.add_argument('--gpu-id', type=str, default='0', help="사용할 GPU ID (기본값: '0')")
parser.add_argument('--seed', type=int, default=1, help="랜덤 시드 (기본값: 1)")
parser.add_argument('--use-cpu', action='store_true', help="CPU만 사용 여부")
parser.add_argument('--log-dir', type=str, default='training_logs', help="로그 및 결과 저장 디렉토리 (기본값: 'training_logs')")
parser.add_argument('--plot-features', action='store_true', help="에폭마다 특징 분포 시각화 여부")
def main():
global args # 전역 변수 args 사용 (visualize_features 등에서 접근)
args = parser.parse_args()
torch.manual_seed(args.seed) # PyTorch 랜덤 시드 설정
os.environ['CUDA_VISIBLE_DEVICES'] = args.gpu_id # GPU 환경 변수 설정
gpu_available = torch.cuda.is_available() # GPU 사용 가능 여부 확인
if args.use_cpu:
gpu_available = False # CPU만 사용하도록 강제
# 로그 파일 설정
sys.stdout = FileAndConsoleLogger(osp.join(args.log_dir, f'log_{args.dataset}.txt'))
if gpu_available:
print(f"현재 GPU 사용 중: {args.gpu_id}")
cudnn.benchmark = True # 컨볼루션 연산 최적화 (입력 크기 고정 시 유용)
torch.cuda.manual_seed_all(args.seed) # 모든 GPU에 대한 랜덤 시드 설정
else:
print("현재 CPU 사용 중")
print(f"데이터셋 생성 중: {args.dataset}")
data_source = get_dataset(
name=args.dataset,
batch_size=args.batch_size,
use_gpu=gpu_available,
worker_count=args.workers,
)
training_loader, test_loader = data_source.training_loader, data_source.test_loader
print(f"모델 생성 중: {args.model-type}")
network_model = create_network_model(name=args.model_type, num_classes=data_source.num_classes)
if gpu_available:
network_model = nn.DataParallel(network_model).cuda() # 다중 GPU 사용 시 모델 병렬화
# 손실 함수 정의
xent_loss_fn = nn.CrossEntropyLoss()
center_loss_fn = FeatureCenterLoss(num_classes=data_source.num_classes, feature_dims=2, use_accelerator=gpu_available)
# 옵티마이저 정의
# 모델 파라미터 옵티마이저: SGD, L2 정규화(weight_decay), 모멘텀 포함
model_optimizer = optim.SGD(network_model.parameters(),
lr=args.lr_model,
weight_decay=5e-04,
momentum=0.9)
# 중심 손실 파라미터 옵티마이저: SGD
center_optimizer = optim.SGD(center_loss_fn.parameters(), lr=args.lr_cent)
# 학습률 스케줄러 정의
model_scheduler = None
if args.stepsize > 0:
model_scheduler = lr_scheduler.StepLR(model_optimizer, step_size=args.stepsize, gamma=args.gamma)
start_time = time.time() # 훈련 시작 시간 기록
for epoch_idx in range(args.max_epoch):
print(f"==> 에폭 {epoch_idx+1}/{args.max_epoch}")
# 훈련 함수 호출
train_epoch(network_model, xent_loss_fn, center_loss_fn,
model_optimizer, center_optimizer,
training_loader, gpu_available, data_source.num_classes, epoch_idx)
if model_scheduler is not None:
model_scheduler.step() # 학습률 업데이트
# 일정 주기마다 또는 마지막 에폭에 모델 평가
if args.eval_freq > 0 and (epoch_idx + 1) % args.eval_freq == 0 or (epoch_idx + 1) == args.max_epoch:
print("==> 모델 평가 시작")
accuracy, error_rate = evaluate_model(network_model, test_loader, gpu_available, data_source.num_classes, epoch_idx)
print(f"정확도 (%): {accuracy:.2f}\t 오류율 (%): {error_rate:.2f}")
# 전체 훈련 시간 계산 및 출력
elapsed_time = round(time.time() - start_time)
formatted_time = str(datetime.timedelta(seconds=elapsed_time))
print(f"훈련 완료. 총 소요 시간 (시:분:초): {formatted_time}")
if __name__ == '__main__':
main()