학습 Target은 직접 만든 Noise입니다
forward는 batch마다 t를 무작위로 고른 뒤 get_losses를 호출합니다. 그 안에서 image와 같은 shape의 noise를 만들고, noisy image를 model에 넣어 estimated noise와 비교합니다.
1
2
3
4
5
6
7
8
9
10
| def get_losses(self, x, t, y):
noise = torch.randn_like(x)
perturbed_x = self.perturb_x(x, t, noise)
estimated_noise = self.model(perturbed_x, t, y)
if self.loss_type == "l1":
loss = F.l1_loss(estimated_noise, noise)
elif self.loss_type == "l2":
loss = F.mse_loss(estimated_noise, noise)
return loss
|
따라서 모델 입력에는 noisy image뿐 아니라 noise 수준을 알려 주는 t가 필요합니다. Class conditioning을 쓰면 y도 함께 넘깁니다. 원문의 간단한 조건부 layer는 nn.Embedding(num_classes, out_channels)로 class bias를 만들고 feature에 더하지만, 실제 U-Net 전체 위치와 shape는 생략돼 있습니다.
또한 원문 forward는 height를 img_size[0]과 비교한 뒤 width도 같은 img_size[0]과 비교합니다. 직사각형 img_size=(H,W)를 지원하려면 width 검사가 의도와 맞는지 확인해야 합니다.