5장 — 잡음을 맞히는 신경망: DDPM과 잡음 제거 스코어 매칭
자주 하는 실수와 요약
자주 하는 실수
| 실수 | 나온 문제 | 바로잡는 법 |
|---|---|---|
| 답이 둘로 갈릴 때 더 그럴듯한 쪽 하나를 고름 | 1 | 제곱 오차를 가장 작게 하는 답은 확률로 섞은 평균이다 |
| 조건부 타깃을 처음 비중(반반)으로 섞음 | 2 | 그 자리를 본 뒤의 확률로 섞는다. 가까운 볼에서 왔을 확률이 크다 |
| 잡음 제거 스코어 매칭 손실이 0까지 내려가야 학습이 끝난다고 봄 | 3 | 손실의 최솟값은 그 자리에서 타깃이 흔들리는 정도의 평균이고, 잡음 수준마다 다르다 |
| 크기가 다른 오차를 그대로 더함 | 4 | 맞힐 값의 크기를 고르게 한 뒤 더한다. 잡음을 맞히면 모든 수준에서 크기가 1이다 |
| 잡음 예측을 원래 그림으로 바꿀 때 α로 나누는 것을 빠뜨림 | 5 | x̂₀ = (xₜ − σₜ ε̂)/αₜ |
| 잡음 예측에서 스코어를 −ε̂/σ² 로 씀 | 5 | 스코어는 −ε̂/σₜ. −(xₜ − x₀)/σ² 는 α = 1 일 때의 식이다 |
| 잡음 오차가 작으면 원래 그림의 어림도 늘 정확하다고 봄 | 6 | 오차가 σₜ/αₜ 배로 바뀐다. 짙은 잡음에서는 크게 부푼다 |
| 비중만 바뀌었는데 결과의 순위가 그대로일 거라 봄 | 7, 9 | 같은 점수표도 비중이 바뀌면 1등이 바뀐다 |
| 역방향 한 걸음의 KL 에서 분산을 σₜ² 로 씀 | 8 | 한 걸음 되돌린 분포의 분산은 βₜ 다. σₜ² 는 처음부터 쌓인 잡음이다 |
| FID 와 로그우도 가운데 한쪽만으로 모델을 판단함 | 9 | 무엇을 재느냐에 따라 1등이 다르다. 단순한 손실은 옅은 잡음의 비중을 줄인다 |
| 줄어든 양의 합만으로 처음 값을 구함(끝 값을 확인하지 않음) | 10 | 끝 값이 0일 때만 합 = 처음 값. 중간에서 끊으면 아래쪽 경계뿐이다 |
| 잡음 섞기 전의 바늘로 모든 수준의 피셔 발산을 계산함 | 11 | 잡음 수준마다 두 분포를 다시 적고 그 바늘을 쓴다 |
| 짙은 잡음에서 바늘이 맞으면 좋은 모델이라 봄 | 12 | 봉우리 모양은 옅은 잡음의 바늘에 있다. 로그우도는 그쪽이 정한다 |
요약
신경망은 잡음 섞인 분포의 바늘을 정답으로 받을 수 없다. 대신 데이터에 잡음을 섞어 본 기록 한 줄마다 「떠나온 점 쪽 바늘」 −(xₜ − x₀)/σₜ²를 타깃으로 주면, 제곱 오차를 가장 작게 하는 답은 그 자리에 선 사람들의 타깃의 평균이고, 그 평균이 트위디 공식에 따라 참 스코어다. 이것이 잡음 제거 스코어 매칭이다. 섞은 잡음 ε를 맞히든, 원래 그림 x₀를 맞히든, 바늘을 맞히든 같은 정보이고(x̂₀ = (xₜ − σₜ ε̂)/αₜ, s = −ε̂/σₜ), 잡음을 맞히면 모든 잡음 수준에서 맞힐 값의 크기가 1로 고르다. 잡음을 천 걸음에 나눠 섞는 사슬을 VAE처럼 보면, 데이터 x₀가 그림이고 나머지 사슬이 잠재 변수이며 인코더는 고정된 잡음 섞기다. 그 ELBO는 걸음마다의 잡음 맞히기 오차에 가중치 λₜ를 붙인 합이 되고(DDPM), 가중치를 모두 같게 한 단순한 손실이 그림의 질을 높였다. 같은 잡음을 섞은 두 분포의 KL은 잡음 분산을 따라 피셔 발산의 절반 빠르기로 줄어드므로, 모든 잡음 수준의 스코어 매칭을 알맞은 가중치로 합하면 KL로 돌아온다. 디퓨전 손실들의 차이는 잡음 수준마다의 가중치다.
flowchart LR A["정답 바늘을 모른다"] --> B["잡음을 섞어 본 기록<br/>타깃 −(xₜ − x₀)/σ²"] B --> C["제곱 오차의 답 = 그 자리의 평균<br/>= 참 스코어 (잡음 제거 스코어 매칭)"] C --> D["타깃 크기가 1/σ 로 오르내림"] D --> E["섞은 잡음 ε 를 맞힌다 (잡음 예측)<br/>x̂₀ 와 s 는 식 한 줄"] F["VAE 의 ELBO"] --> G["인코더 = 고정된 잡음 섞기<br/>사슬의 ELBO = Σ λₜ ‖ε − εθ‖² (DDPM)"] E --> G G --> H["단순한 손실: λₜ 를 모두 같게"] C --> I["피셔 발산의 합 = KL<br/>(모든 잡음 수준을 합치면)"] I --> H
막힌 곳
이제 신경망에게 무엇을 맞히게 할지 안다. 그림과 잡음 수준 t를 받아 그림과 같은 크기의 잡음 εθ(xt, t)를 내놓게 하고, 단순한 제곱 오차로 배우면 된다.
그런데 이 장의 장난감은 모두 숫자 하나짜리였다. 가로세로 32픽셀 컬러 그림이면 입력도 출력도 32 × 32 × 3 = 3072개의 숫자다. 모든 입력과 모든 출력을 잇는 신경망은 매개변수가 너무 많고, 이웃만 보는 합성곱은 짙은 잡음 속에서 그림 전체의 모양을 가늠하기에 너무 좁게 본다. 게다가 같은 신경망이 잡음 수준 t마다 다른 일을 해야 하는데, 숫자 하나인 t를 그림 크기의 신경망 어디에 어떻게 넣어야 할까?