[모두의 연구소] MNIST 데이터셋 예측 모델 - MNIST_Hybrid_Model(260713 ver.)

WonTerry·2026년 7월 10일

Deep Learning

목록 보기
7/27
post-thumbnail

전혀 이런 의도가 아니었는데...
Deep SAD 학습 중 이상하게 흘러가버린 코드...


"순수 MNIST 숫자 분류"라는 목적만 보면 대체로 쓸데없이 복잡한 편이고, 이 복잡함이 실제 이점으로 이어지는 시나리오는 따로 있습니다.

구조적으로 무엇을 하고 있나

인코더(특징 추출) → ① 디코더로 복원, ② 분류기로 라벨 예측, 두 손실을 가중합해서 동시에 학습시키는 멀티태스크 학습(오토인코더 + 분류기) 구조입니다.

이론적으로 기대할 수 있는 이점

  • 정규화 효과: 복원 손실이 보조 과제로 작용해 인코더가 라벨에만 과적합되지 않고 입력 구조 전반을 보존하는 특징을 학습하도록 유도할 수 있습니다.
  • 준지도학습 확장성: 복원 손실은 라벨이 없어도 계산 가능하므로, 라벨 있는 데이터가 적고 라벨 없는 데이터가 많은 상황에서는 이 구조가 실제로 유용해집니다.
  • 부가 기능(이상 탐지/노이즈 제거): 인코더-디코더가 있으면 재구성 오차 기반 이상치 탐지, 노이즈 제거, 잠재공간 시각화 같은 부가 기능을 "공짜로" 얻습니다. (사실 이 코드의 BaseNet, BaseTrainer, BaseADDataset, MNIST_LeNet 구조 자체가 원래 Deep SVDD(이상 탐지) 논문 구현체에서 그대로 가져온 것으로 보입니다. 즉 이상 탐지용으로 설계된 인코더에 분류기 헤드를 붙인 형태예요.)

순수 분류 성능 관점에서의 단점

  • 목적 충돌 가능성: "복원에 유리한 표현"과 "분류에 유리한 표현"은 항상 일치하지 않습니다. 압축·복원에 필요한 정보(픽셀 디테일)와 클래스 판별에 필요한 정보(경계, 형태)가 다를 수 있어, 두 손실을 동시에 최적화하면 순수 분류기 대비 정확도가 오히려 살짝 낮아지는 경우가 흔합니다.
  • 병목(rep_dim=32)이 작음: 복원을 감안해 설계된 32차원 병목은 일반적인 분류 전용 CNN의 특징 벡터(보통 128~256차원 이상)보다 훨씬 작아서, 표현력 면에서 손해를 볼 수 있습니다.
  • 계산량 증가: 디코더(ConvTranspose 3개 + interpolate)는 분류 정확도에는 전혀 기여하지 않으면서 forward/backward 연산량과 학습 시간을 늘립니다.
  • 하이퍼파라미터 부담: classification_loss_weight 같은 손실 가중치를 따로 튜닝해야 합니다.
  • MNIST는 이미 너무 쉬운 문제: 얕은 CNN 하나만으로도 99%대 정확도가 쉽게 나오는 데이터셋이라, 이 정도 과제의 이점(정규화, 준지도학습)을 체감하기 어렵습니다. 이 구조가 진가를 발휘하는 건 데이터가 훨씬 작거나 노이즈가 많거나 라벨이 부족한 상황입니다.

정리

목적이 하이브리드 구조일반 CNN 분류기
순수 MNIST 분류 정확도비슷하거나 오히려 살짝 낮을 수 있음더 간단하고 빠르며 보통 동등 이상
학습/추론 속도느림 (디코더 오버헤드)빠름
재구성/이상탐지/시각화 필요 시유용불가능
라벨 부족 + 비지도 신호 활용유리불가능

즉, "MNIST 숫자를 잘 분류하고 싶다"가 유일한 목적이라면 오토인코더 부분은 걷어내고 인코더 + 분류기만 남기는 게 더 합리적입니다. 반대로 재구성 결과를 시각적으로 확인하고 싶다거나, 나중에 이상 탐지/준지도학습으로 확장할 계획이 있다면 지금 구조가 그 확장성을 미리 확보해둔 것이라 의미가 있습니다.


MNIST_Hybrid_Model

import torch
import torch.nn as nn
import torch.optim as optim
from torchvision import datasets, transforms
from torch.utils.data import DataLoader
import numpy as np

from base.base_net import BaseNet
from networks.mnist_LeNet import MNIST_LeNet, MNIST_LeNet_Decoder
from base.base_dataset import BaseADDataset
from base.base_trainer import BaseTrainer

class MNIST_Hybrid_Model(BaseNet):
    def __init__(self, rep_dim=32):
        super().__init__()
        
        self.rep_dim = rep_dim
        # Autoencoder components
        self.encoder = MNIST_LeNet(rep_dim=rep_dim)
        self.decoder = MNIST_LeNet_Decoder(rep_dim=rep_dim)
        
        # Classifier component
        self.classifier = nn.Sequential(
            nn.Linear(rep_dim, 128),
            nn.ReLU(),
            nn.Dropout(0.5),
            nn.Linear(128, 10)  # 10 classes for MNIST
        )
        
    def forward(self, x):
        # Get encoded features
        code = self.encoder(x)
        
        # Reconstruct original image
        reconstructed = self.decoder(code)
        
        # Classify the encoded features
        class_output = self.classifier(code)
        
        return reconstructed, class_output


class MNIST_Dataset(BaseADDataset):
    def __init__(self, root):
        super().__init__(root)
        self.normal_classes = (0, 1, 2, 3, 4, 5, 6, 7, 8, 9)  # All digits are normal
        self.outlier_classes = ()

    def loaders(self, batch_size, shuffle_train=True, shuffle_test=False, num_workers=0):
        # Define transforms for the data
        transform = transforms.Compose([
            transforms.ToTensor(),
            transforms.Normalize((0.1307,), (0.3081,))  # MNIST mean and std
        ])

        # Load the datasets
        train_dataset = datasets.MNIST(self.root, train=True, download=True, transform=transform)
        test_dataset = datasets.MNIST(self.root, train=False, download=True, transform=transform)

        # Create data loaders
        train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=shuffle_train, num_workers=num_workers)
        test_loader = DataLoader(test_dataset, batch_size=batch_size, shuffle=shuffle_test, num_workers=num_workers)

        return train_loader, test_loader

class MNIST_Hybrid_Trainer(BaseTrainer):
    def __init__(self, optimizer_name='adam', lr=0.001, n_epochs=100, lr_milestones=(50,), batch_size=128,
                 weight_decay=1e-6, device='cuda', n_jobs_dataloader=0, classification_loss_weight=0.5):
        super().__init__(optimizer_name, lr, n_epochs, lr_milestones, batch_size, weight_decay, device, n_jobs_dataloader)
        self.classification_loss_weight = classification_loss_weight

    def train(self, dataset, net):
        # Set up the optimizer
        if self.optimizer_name == 'adam':
            optimizer = optim.Adam(net.parameters(), lr=self.lr, weight_decay=self.weight_decay)
        else:
            raise ValueError('Unsupported optimizer')

        # Load training data
        train_loader, _ = dataset.loaders(batch_size=self.batch_size, shuffle_train=True, num_workers=self.n_jobs_dataloader)

        # Training loop
        for epoch in range(self.n_epochs):
            net.train()
            total_recon_loss = 0.0
            total_class_loss = 0.0
            
            for data in train_loader:
                inputs, labels = data  # MNIST dataset returns (data, label)
                inputs = inputs.to(self.device)
                labels = labels.to(self.device)
                
                # Forward pass
                reconstructed, class_output = net(inputs)
                
                # Reconstruction loss (MSE)
                recon_loss = nn.MSELoss()(reconstructed, inputs)    
                # 디코더로 생성한 이미지와 원본 비교 (260713)
                # recon_loss가 크게 개선되지 않음 (260713)

                # Classification loss (Cross Entropy)
                class_loss = nn.CrossEntropyLoss()(class_output, labels)
                
                # Combined loss
                total_loss = recon_loss + self.classification_loss_weight * class_loss
                
                # Backward pass
                optimizer.zero_grad()
                total_loss.backward()
                optimizer.step()
                
                total_recon_loss += recon_loss.item()
                total_class_loss += class_loss.item()
            
            if epoch in self.lr_milestones:
                for param_group in optimizer.param_groups:
                    param_group['lr'] *= 0.1
            
            print(f'Epoch [{epoch+1}/{self.n_epochs}], '
                  f'Recon Loss: {total_recon_loss/len(train_loader):.4f}, '
                  f'Class Loss: {total_class_loss/len(train_loader):.4f}')
        
        return net

    def test(self, dataset, net):
        # Load test data
        _, test_loader = dataset.loaders(batch_size=self.batch_size, shuffle_test=False, num_workers=self.n_jobs_dataloader)

        # Testing loop
        net.eval()
        total_recon_loss = 0.0
        total_class_loss = 0.0
        correct = 0
        total = 0
        
        with torch.no_grad():
            for data in test_loader:
                inputs, labels = data  # MNIST dataset returns (data, label)
                inputs = inputs.to(self.device)
                labels = labels.to(self.device)
                
                # Forward pass
                reconstructed, class_output = net(inputs)
                
                # Reconstruction loss
                recon_loss = nn.MSELoss()(reconstructed, inputs)
                
                # Classification loss
                class_loss = nn.CrossEntropyLoss()(class_output, labels)
                
                # Accuracy calculation for classification
                _, predicted = torch.max(class_output.data, 1)
                total += labels.size(0)
                correct += (predicted == labels).sum().item()
                
                total_recon_loss += recon_loss.item()
                total_class_loss += class_loss.item()
        
        avg_recon_loss = total_recon_loss / len(test_loader)
        avg_class_loss = total_class_loss / len(test_loader)
        accuracy = 100 * correct / total
        
        print(f'Test Results - Recon Loss: {avg_recon_loss:.4f}, '
              f'Class Loss: {avg_class_loss:.4f}, Accuracy: {accuracy:.2f}%')
        
        return avg_recon_loss, avg_class_loss, accuracy


# Main training function
def main():
    # Set device
    # device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
    # print(f'Using device: {device}')

    # 맥 전용 코드 (MPS GPU용)
    device = torch.device("mps" if torch.backends.mps.is_available() else "cpu")
    print(f"사용 장치: {device}")

    # Create dataset and trainer
    dataset = MNIST_Dataset(root='./datasets')
    trainer = MNIST_Hybrid_Trainer(device=device, n_epochs=50, classification_loss_weight=0.3)  # Reduced epochs for demonstration

    # Create the hybrid model (autoencoder + classifier)
    hybrid_model = MNIST_Hybrid_Model(rep_dim=32).to(device)

    print("Training the MNIST hybrid model (autoencoder + classifier)...")
    trained_net = trainer.train(dataset, hybrid_model)
    
    print("Testing the MNIST hybrid model...")
    recon_loss, class_loss, accuracy = trainer.test(dataset, trained_net)
    
    print("Training and testing completed.")

    # Save the trained model
    model_save_path = 'mnist_hybrid_model_260710.pt'
    torch.save({
        'model_state_dict': trained_net.state_dict(),
        'rep_dim': trained_net.rep_dim,
        'test_recon_loss': recon_loss,
        'test_class_loss': class_loss,
        'test_accuracy': accuracy,
    }, model_save_path)
    print(f"모델이 저장되었습니다: {model_save_path}")

if __name__ == '__main__':
    main()



테스트 결과

"""
저장된 MNIST_Hybrid_Model을 불러와서
임의의 MNIST 샘플 5개에 대한 예측 결과를 그래픽으로 표현하는 코드.

사용 전 준비:
- MNIST_basic_260710.py 를 먼저 실행하여 'mnist_hybrid_model.pt' 가 생성되어 있어야 함.
- 두 파일(MNIST_basic_260710.py, MNIST_predict_visualize.py)이 같은 폴더에 있어야 함
  (MNIST_Hybrid_Model 클래스를 그대로 import 해서 재사용하기 때문).
"""

import torch
import torch.nn as nn
import matplotlib.pyplot as plt
from torchvision import datasets, transforms

from base.base_net import BaseNet
from networks.mnist_LeNet import MNIST_LeNet, MNIST_LeNet_Decoder

# 학습 스크립트에서 모델 클래스를 그대로 가져와 재사용
# from MNIST_basic_260710 import MNIST_Hybrid_Model

# 한번더 작성해준다. (260710)
# Hybrid model with both autoencoder and classifier capabilities
class MNIST_Hybrid_Model(BaseNet):
    def __init__(self, rep_dim=32):
        super().__init__()
        
        self.rep_dim = rep_dim
        # Autoencoder components
        self.encoder = MNIST_LeNet(rep_dim=rep_dim)
        self.decoder = MNIST_LeNet_Decoder(rep_dim=rep_dim)
        
        # Classifier component
        self.classifier = nn.Sequential(
            nn.Linear(rep_dim, 128),
            nn.ReLU(),
            nn.Dropout(0.5),
            nn.Linear(128, 10)  # 10 classes for MNIST
        )
        
    def forward(self, x):
        # Get encoded features
        code = self.encoder(x)
        
        # Reconstruct original image
        reconstructed = self.decoder(code)
        
        # Classify the encoded features
        class_output = self.classifier(code)
        
        return reconstructed, class_output


def load_model(model_path='mnist_hybrid_model_260710.pt', device='cpu'):
    """저장된 체크포인트로부터 모델을 복원한다."""
    checkpoint = torch.load(model_path, map_location=device)

    model = MNIST_Hybrid_Model(rep_dim=checkpoint['rep_dim'])
    model.load_state_dict(checkpoint['model_state_dict'])
    model.to(device)
    model.eval()

    print(f"모델 로드 완료: {model_path}")
    if 'test_accuracy' in checkpoint:
        print(f"  (학습 시 테스트 정확도: {checkpoint['test_accuracy']:.2f}%)")

    return model


def get_random_samples(n_samples=5, root='./datasets'):
    """MNIST 테스트셋에서 임의의 샘플 n개를 가져온다."""
    transform = transforms.Compose([
        transforms.ToTensor(),
        transforms.Normalize((0.1307,), (0.3081,))
    ])
    test_dataset = datasets.MNIST(root, train=False, download=True, transform=transform)

    indices = torch.randperm(len(test_dataset))[:n_samples]
    images, labels = [], []
    for idx in indices:
        img, label = test_dataset[idx]
        images.append(img)
        labels.append(label)

    images = torch.stack(images)  # (n_samples, 1, 28, 28)
    labels = torch.tensor(labels)
    return images, labels


def unnormalize(img_tensor, mean=0.1307, std=0.3081):
    """정규화된 텐서를 다시 0~1 범위의 이미지로 되돌린다 (시각화용)."""
    return img_tensor * std + mean


def visualize_predictions(model, images, labels, device='cpu', save_path='mnist_predictions.png'):
    """실제 이미지와 모델의 예측 결과를 함께 그래픽으로 표현한다."""
    n_samples = images.size(0)
    images_device = images.to(device)

    with torch.no_grad():
        reconstructed, class_output = model(images_device)
        probs = torch.softmax(class_output, dim=1)
        confidences, predicted = torch.max(probs, dim=1)

    fig, axes = plt.subplots(2, n_samples, figsize=(3 * n_samples, 6))

    for i in range(n_samples):
        true_label = labels[i].item()
        pred_label = predicted[i].item()
        confidence = confidences[i].item() * 100
        is_correct = (true_label == pred_label)

        # 위쪽 행: 원본 이미지
        img_show = unnormalize(images[i, 0]).clamp(0, 1).cpu().numpy()
        axes[0, i].imshow(img_show, cmap='gray')
        axes[0, i].set_title(f'real: {true_label}', fontsize=12)
        axes[0, i].axis('off')

        # 아래쪽 행: 예측 결과 텍스트 표시
        # 아래쪽 행: 모델이 복원한(reconstructed) 이미지
        recon_show = unnormalize(reconstructed[i, 0]).clamp(0, 1).detach().cpu().numpy()
        color = 'green' if is_correct else 'red'
        mark = '✓' if is_correct else '✗'
        axes[1, i].imshow(recon_show, cmap='gray')
        axes[1, i].set_title(
            f'predict: {pred_label} ({confidence:.1f}%) {mark}',
            fontsize=12, color=color
        )
        axes[1, i].axis('off')

    axes[0, 0].set_ylabel('origin', fontsize=12)
    axes[1, 0].set_ylabel('reconstructed', fontsize=12)

    fig.suptitle('MNIST predict (upper: origin / below: prediction)', fontsize=14)
    plt.tight_layout()
    plt.savefig(save_path, dpi=150, bbox_inches='tight')
    print(f"결과 이미지가 저장되었습니다: {save_path}")
    plt.show()


def main():
    device = torch.device("mps" if torch.backends.mps.is_available() else "cpu")
    print(f"사용 장치: {device}")

    # 1) 저장된 모델 불러오기
    model = load_model('mnist_hybrid_model_260710.pt', device=device)

    # 2) 임의의 MNIST 샘플 5개 불러오기
    images, labels = get_random_samples(n_samples=5, root='./data')

    # 3) 예측 및 시각화
    visualize_predictions(model, images, labels, device=device)


if __name__ == '__main__':
    main()


profile
Hello, I'm Terry! 👋 Enjoy every moment of your life! 🌱 My current interests are Signal processing, Machine learning, Python, Database, LLM & RAG, MCP & ADK, Multi-Agents, Physical AI, ROS2...

0개의 댓글