서론 이전 기사에서 DDPM 모델을 MNIST 데이터셋에 적용해 기본적인 동작을 검증했습니다. 그러나 실제 이미지 생성 작업에서는 더 복잡한 텍스처와 세부 정보가 요구됩니다. 특히 RGB 채널을 가진 CIFAR-10 데이터셋은 고해상도 이미지를 처리하는 데 적합한 테스트베드입니다. 이에 따라 U-Net이라는 강력한 구조를 도입하여 DDPM을 개선합니다.
U-Net은 의료 영상 분할을 위한 설계로, 인코더-디코더 구조와 skip connection이 특징입니다. 이러한 특성은 확산 모델의 생성 과정에 매우 효과적입니다. 본 기사에서는 다음과 같은 내용을 다룹니다:
- DDPM용 U-Net 구현
- CIFAR-10 데이터셋 적응
- 전체 학습 및 추론 프로세스 구현
- 생성 결과 시각화
모듈 설계: DDPM용 U-Net 구조 U-Net 선택 이유 U-Net의 주요 장점은 다음과 같습니다:
- 다중 해상도 처리 능력: 하향 변환을 통한 의미 정보 추출 및 상향 변환으로 세부 정보 복원
- 크로스 레이어 연결: 저수준 공간 구조 정보 보존
- 생성 작업과의 호환성: 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)