Flatten 레이어 자동계산

PyTorch는 nn.Flatten() 이후 nn.Linear 계층을 정의할 때, 입력 텐서의 차원을 자동으로 계산하지 않습니다. 따라서 nn.Linear의 입력 크기와 출력 크기를 명시적으로 지정해야 합니다. 그러나 PyTorch를 사용하면 입력 크기를 자동으로 계산하는 방법을 간단히 구현할 수 있습니다.


자동으로 크기 계산하는 방법

1. Forward Pass에서 텐서 크기 계산

모델의 forward 메서드에서 입력 데이터가 컨볼루션과 풀링 레이어를 거친 후 남은 텐서의 크기를 계산하여 동적으로 nn.Linear를 초기화할 수 있습니다. 이를 위해 임시 텐서를 사용해 크기를 계산하는 방법이 있습니다.

아래는 이를 구현한 예제입니다.

import torch
import torch.nn as nn

class CNNModel(nn.Module):
    def __init__(self, input_channels=1, num_classes=10):
        super(CNNModel, self).__init__()
        # Convolutional layers
        self.conv_layers = nn.Sequential(
            nn.Conv2d(input_channels, 32, kernel_size=3, stride=1, padding=1),
            nn.ReLU(),
            nn.MaxPool2d(kernel_size=2, stride=2),

            nn.Conv2d(32, 64, kernel_size=3, stride=1, padding=1),
            nn.ReLU(),
            nn.MaxPool2d(kernel_size=2, stride=2)
        )

        # Placeholder for fully connected layers
        self.fc_layers = None  # 초기화되지 않은 상태
        self.num_classes = num_classes

    def forward(self, x):
        # Convolutional layers forward pass
        x = self.conv_layers(x)

        # Fully connected layers (initialize dynamically if not already initialized)
        if self.fc_layers is None:
            num_features = x.shape[1] * x.shape[2] * x.shape[3]  # 채널 * 높이 * 너비
            self.fc_layers = nn.Sequential(
                nn.Flatten(),
                nn.Linear(num_features, 128),
                nn.ReLU(),
                nn.Linear(128, self.num_classes)
            ).to(x.device)  # GPU로 이동

        # Fully connected layers forward pass
        x = self.fc_layers(x)
        return x

동작 방식

  1. 초기 Conv 레이어 실행:
  2. forward 메서드에서 입력 데이터를 self.conv_layers에 통과시켜 텐서의 크기를 계산합니다.

  3. 동적 Linear 초기화:

  4. 텐서의 출력 크기를 기반으로 nn.Linear 계층을 동적으로 초기화합니다.
  5. 이를 통해 nn.Linear 계층의 입력 크기를 미리 알 필요가 없습니다.

  6. 한 번만 초기화:

  7. self.fc_layers가 초기화되지 않은 경우에만 Linear 계층을 생성하므로 효율적입니다.

테스트

# 모델 생성
model = CNNModel(input_channels=1, num_classes=10)

# 임의의 입력 데이터 (흑백 이미지: 1채널, 28x28)
x = torch.rand(64, 1, 28, 28)  # 배치 크기: 64
output = model(x)

# 출력 크기 확인
print("Output shape:", output.shape)  # (64, 10)

장점

  • 모델의 입력 데이터 크기에 따라 nn.Linear를 동적으로 설정할 수 있어 재사용성이 높아짐.
  • 데이터셋마다 다른 크기의 입력 데이터를 사용할 때 유용.

주의점

  • Forward pass를 통해 계산하므로 __init__ 메서드에서 계층 구조를 완전히 정의하지 못할 수 있습니다. 이는 일부 모델 디버깅 도구에서 불편할 수 있습니다.
  • 복잡한 네트워크에서는 크기를 정확히 계산하여 __init__에서 모든 계층을 명시적으로 설정하는 것이 더 명확할 수 있습니다.