UNet과 DDPM의 결합: CIFAR-10 데이터셋을 활용한 이미지 생성 모델 구축

서론 이전 기사에서 DDPM 모델을 MNIST 데이터셋에 적용해 기본적인 동작을 검증했습니다. 그러나 실제 이미지 생성 작업에서는 더 복잡한 텍스처와 세부 정보가 요구됩니다. 특히 RGB 채널을 가진 CIFAR-10 데이터셋은 고해상도 이미지를 처리하는 데 적합한 테스트베드입니다. 이에 따라 U-Net이라는 강력한 구조를 도입하여 DDPM을 개선합니다.

U-Net은 의료 영상 분할을 위한 설계로, 인코더-디코더 구조와 skip connection이 특징입니다. 이러한 특성은 확산 모델의 생성 과정에 매우 효과적입니다. 본 기사에서는 다음과 같은 내용을 다룹니다:

  • DDPM용 U-Net 구현
  • CIFAR-10 데이터셋 적응
  • 전체 학습 및 추론 프로세스 구현
  • 생성 결과 시각화

모듈 설계: DDPM용 U-Net 구조 U-Net 선택 이유 U-Net의 주요 장점은 다음과 같습니다:

  1. 다중 해상도 처리 능력: 하향 변환을 통한 의미 정보 추출 및 상향 변환으로 세부 정보 복원
  2. 크로스 레이어 연결: 저수준 공간 구조 정보 보존
  3. 생성 작업과의 호환성: DDPM과 자연스럽게 결합 가능

모델 구현 시간 임베딩 모듈 확산 과정의 시간 단계를 인식하기 위한 임베딩 함수:

def time_positional_encoding(n, d):
    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
    embedding = torch.arange(d, device=device).float() / d
    embedding = 1.0 / (10000 ** embedding)
    pos = torch.arange(n, device=device).float().unsqueeze(1)
    emb = pos * embedding.unsqueeze(0)
    emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=1)
    return emb

잔차 블록 + 시간 정보 통합 시간 정보와 이미지 특징을 결합하여 모델의 시간 인식 능력을 강화하는 블록:

class ResidualConvBlock(nn.Module):
    def __init__(self, in_channels, out_channels, time_emb_dim):
        super().__init__()
        self.time_mlp = nn.Sequential(
            nn.Linear(time_emb_dim, out_channels),
            nn.ReLU()
        )

        def get_group_norm(c):
            if c < 8:
                return nn.GroupNorm(1, c)  # 3채널(RGB) 대상
            else:
                return nn.GroupNorm(8, c)

        self.block1 = nn.Sequential(
            nn.Conv2d(in_channels, out_channels, 3, padding=1),
            get_group_norm(out_channels),
            nn.ReLU()
        )

        self.block2 = nn.Sequential(
            nn.Conv2d(out_channels, out_channels, 3, padding=1),
            get_group_norm(out_channels),
            nn.ReLU()
        )

        self.residual_conv = nn.Conv2d(in_channels, out_channels, 1) if in_channels != out_channels else nn.Identity()

    def forward(self, x, t):
        h = self.block1(x)
        time_emb = self.time_mlp(t).view(t.shape[0], -1, 1, 1)
        h = h + time_emb
        h = self.block2(h)
        return h + self.residual_conv(x)

U-Net 주체 구조 완성된 인코더-디코더 구조:

class DiffusionUNet(nn.Module):
    def __init__(self, in_channels=3, base_channels=64, time_emb_dim=256):
        super().__init__()
        self.time_mlp = nn.Sequential(
            nn.Linear(1, time_emb_dim),
            nn.ReLU(),
            nn.Linear(time_emb_dim, time_emb_dim)
        )

        self.conv0 = ResidualConvBlock(in_channels, base_channels, time_emb_dim)
        self.down1 = ResidualConvBlock(base_channels, base_channels * 2, time_emb_dim)
        self.down2 = ResidualConvBlock(base_channels * 2, base_channels * 4, time_emb_dim)
        self.mid = ResidualConvBlock(base_channels * 4, base_channels * 4, time_emb_dim)
        self.up1 = ResidualConvBlock(base_channels * 4, base_channels * 2, time_emb_dim)
        self.up2 = ResidualConvBlock(base_channels * 2, base_channels, time_emb_dim)
        self.out = nn.Conv2d(base_channels, in_channels, 1)

        self.pool = nn.MaxPool2d(2)
        self.upsample = nn.Upsample(scale_factor=2, mode='nearest')

    def forward(self, x, t):
        t = self.time_mlp(t.unsqueeze(-1).float())
        x0 = self.conv0(x, t)
        x1 = self.down1(self.pool(x0), t)
        x2 = self.down2(self.pool(x1), t)
        xm = self.mid(self.pool(x2), t)
        x = self.up1(self.upsample(xm) + x2, t)
        x = self.up2(self.upsample(x) + x1, t)
        x = self.out(self.upsample(x) + x0)
        return x

CIFAR-10 데이터 준비 및 학습 프로세스 데이터 전처리

transform = transforms.Compose([
    transforms.ToTensor(),
    transforms.Lambda(lambda x: x * 2 - 1)  # [-1, 1] 범위 정규화
])
dataset = torchvision.datasets.CIFAR10(root='./data', train=True, transform=transform, download=True)
loader = DataLoader(dataset, batch_size=128, shuffle=True)

Beta 스케줄 및 확산 과정

T = 1000
betas = torch.linspace(1e-4, 0.02, T)
alphas = 1 - betas
alphas_cumprod = torch.cumprod(alphas, dim=0)
sqrt_alphas_cumprod = torch.sqrt(alphas_cumprod).to(device)
sqrt_one_minus_alphas_cumprod = torch.sqrt(1 - alphas_cumprod).to(device)

손실 함수 및 학습 루프

def noise_prediction_loss(model, x_0, t):
    noise = torch.randn_like(x_0)
    x_t = q_sample(x_0, t, noise)
    pred = model(x_t, t.float())
    return F.mse_loss(pred, noise)

model = DiffusionUNet().to(device)
optimizer = torch.optim.Adam(model.parameters(), lr=1e-4)

epochs = 50
for epoch in range(epochs):
    pbar = tqdm(loader)
    for x, _ in pbar:
        x = x.to(device)
        t = torch.randint(0, T, (x.size(0),), device=device).long()
        loss = noise_prediction_loss(model, x, t)

        optimizer.zero_grad()
        loss.backward()
        optimizer.step()

        pbar.set_description(f"Epoch {epoch+1} Loss: {loss.item():.4f}")

샘플링 프로세스

@torch.no_grad()
def generate_samples(model, image_size=32, channels=3, num_samples=16):
    model.eval()
    x = torch.randn(num_samples, channels, image_size, image_size).to(device)
    for t in reversed(range(T)):
        t_tensor = torch.full((num_samples,), t, device=device).float()
        noise_pred = model(x, t_tensor)

        beta = betas[t]
        alpha = alphas[t]
        alpha_hat = alphas_cumprod[t]

        if t > 0:
            noise = torch.randn_like(x)
        else:
            noise = torch.zeros_like(x)

        x = (1 / alpha**0.5) * (x - ((1 - alpha) / (1 - alpha_hat)**0.5) * noise_pred) + (beta**0.5) * noise
    return x

결과 시각화

samples = generate_samples(model)
samples = (samples.clamp(-1, 1) + 1) / 2  # [0, 1] 범위로 변환

grid = torchvision.utils.make_grid(samples, nrow=4)
plt.figure(figsize=(8, 8))
plt.axis("off")
plt.imshow(grid.permute(1, 2, 0).cpu().numpy())
plt.show()

요약

모듈 상태
모델 구조 U-Net 완성 ✅
데이터셋 CIFAR-10 적응 ✅
손실 함수 노이즈 예측 ✅
시각화 샘플 출력 ✅

다음 기사 예고 다음 기사에서는 다음 주제를 다룰 예정입니다:

  • 샘플링 속도 최적화(DDIM)
  • 생성 품질 평가(FID, IS)
  • 조건부 생성(Class-Conditional DDPM)

태그: DDPM UNet CIFAR-10 이미지 생성 GAN

8월 28일 05:13에 게시됨