Branch Log · Open in interactive viewer →

3. Metrics

신경망에서 proximal operator 및 natural gradient를 적용해 볼 것이다. 이를 위해선 output space에서 metric을 정의한 뒤, 신경망의 weight space로 pull back해야 한다.


3.4 Function Space Distance

출력 공간 점만 입력으로 하는 g 함수를, (파라미터를 출력으로 보내는) f 함수를 활용한 합성을 통해 parameter space로 Pullback할 수 있다.

f*g(𝐱1,,𝐱K)=g(f(𝐱1),,f(𝐱K))

그렇다면 (출력 공간의) 비유사도 함수 ρ 를 pull back하면 어떨까?

f*ρ(𝐱,𝐱)=ρ(f(𝐱),f(𝐱)) 𝐆𝐱=𝐱2f*ρ(𝐱,𝐱0)|𝐱=𝐱0=𝐉𝐳𝐱[𝐳2ρ(𝐳,𝐳0)|𝐳=𝐳0]𝐉𝐳𝐱+aρza𝐱2[f(𝐱)]a=0=𝐉𝐳𝐱𝐆𝐳𝐉𝐳𝐱

Note: 일차 미분이 ddxρ(f(x))=ρ(f(x))·f(x) 이며, 이차 미분 시 f 가 두 번 곱해진다.

분해 형태가 GN Hessian(Lecture 2)과 똑같지만, 결정적 차이가 존재한다. 바로 오른쪽 항이 정확히 0이다. 다시 말해 GN Hessian은 오른쪽 항을 drop하는 근사였지만, pullback metric은 근사가 아니라 등식이다.


3.4.1 Example: Rosenbrock Function

비유사도 함수를 (제곱) 유클리드 거리 ρeuc(𝐳,𝐳)=12𝐳𝐳2 로 둔 Rosenbrock 예제에서, pull back을 적용해 보자.

f*ρeuc(𝐱,𝐱)=12f(𝐱)f(𝐱)2

유클리드 거리이므로 𝐆𝐳=𝐈 이며, 따라서 𝐆𝐱=𝐉𝐳𝐱𝐉𝐳𝐱 이다. 업데이트 궤적을 보면 출력 공간에서 원형 그릇을 곧장 내려가게 된다.

parameter space output space
pullback metric proximal update in parameter space pullback metric proximal update in output space

역함수 없이 3.1절의 '반칙'을 재현한 셈이다.


3.4.2 Generalization to Neural Networks

신경망 f(𝐰,𝐱) 에서는 가중치를 정해도 출력이 입력에 따라 달라진다. 즉 가중치 하나가 정하는 것은 출력 점 하나가 아니라 입력 출력 대응(함수) 전체다. 따라서, 𝐰𝐰 와 얼마나 다른가는, 두 함수를 비교하는 문제로 바라보아야 한다.

ρpull(𝐰,𝐰)=𝔼𝐱[ρ(f(𝐰,𝐱),f(𝐰,𝐱))]

샘플이 유한한 경우는 다음과 같이 식을 작성한다.

ρpull(𝐰,𝐰)=1Ni=1Nρ(f(𝐰,𝐱(i)), f(𝐰,𝐱(i)))

다음은 두 가중치 𝐰𝐰 가 정하는 함수에서, 유한한 샘플링 𝐱(i) 에서 거리를 측정한 그림이다.

function space distance

𝐆𝐰=𝐰2ρpull(𝐰,𝐰)|𝐰=𝐰=𝔼𝐱[𝐰2ρ(f(𝐰,𝐱),f(𝐰,𝐱))]=𝔼𝐱[𝐉𝐳𝐰𝐆𝐳𝐉𝐳𝐰]

이후부터는 pullback metric이라고 지칭할 것이다. (정식 명칭은 아니다.)


3.5 Connection to Gauss-Newton Hessian

Pullback metric의 분해식이 Gauss-Newton Hessian과 같은 형태인 건 우연이 아니다.

행렬 정의 가운데 행렬
Gauss-Newton Hessian 𝐆=𝔼𝐱[𝐉𝐳𝐰𝐇𝐳𝐉𝐳𝐰] 출력 손실의 Hessian 𝐇𝐳
Pullback metric 𝐆=𝔼𝐱[𝐉𝐳𝐰𝐆𝐳𝐉𝐳𝐰] 출력 공간 metric 𝐆𝐳

즉, 𝐆𝐳=𝐇𝐳 이면 동일하다. 예를 들어 '제곱 오차 손실 + 유클리드 거리'는 𝐇𝐳=𝐆𝐳=𝐈 이므로, 둘 다 고전적인 Gauss-Newton matrix 𝔼[𝐉𝐳𝐰𝐉𝐳𝐰] 로 동일하다.


3.5.1 Bregman Divergence

Bregman divergence는 𝐆𝐳=𝐇𝐳 (pullback metric = GN Hessian) 조건을 만족하는 비유사도를, 임의의 볼록 손실에서 획득할 수 있는 방법이다.

Dϕ(𝐳,𝐳)=ϕ(𝐳)ϕ(𝐳)ϕ(𝐳)(𝐳𝐳)

bregman divergence

ϕ 가 볼록이므로 항상 0 이며, 거리가 벌어지는 속도는 ϕ 의 곡률이 결정한다.

𝐳2Dϕ(𝐳,𝐳)|𝐳=𝐳=𝐳2[ϕ(𝐳)ϕ(𝐳)ϕ(𝐳)(𝐳𝐳)]|𝐳=𝐳=𝐳2[ϕ(𝐳)]|𝐳=𝐳=2ϕ(𝐳)

따라서 손실 이 볼록이면 ϕ= 로 골라 D 을 비유사도로 쓰면 된다. 그러면 자동으로 𝐆𝐳=2=𝐇𝐳 가 되어, pullback metric = GN Hessian이 성립한다.

ϕ(𝐳)=12𝐳2Dϕ(𝐳,𝐳)=12𝐳𝐳2 ϕ(𝐳)=logZ(𝐳)Dϕ(𝐳,𝐳)=DKL(p𝐳p𝐳)

Note: Z 의 유래

p(t=k𝐳)=ezkjezj


3.6 Fisher Information Matrix for Neural Networks

Limitations of the Empirical Fisher Approximation 논문(2019)

이번에는 KL divergence를 비유사도 ρ 로 삼고( 𝐆𝐳=𝐅𝐳 , 3.3절 참조 ), pullback metric을 구해 보자.

𝔼𝐱[𝐉𝐳𝐰𝐅𝐳𝐉𝐳𝐰]

score 벡터 𝒟𝐳=𝐳logp(𝐭𝐳) 는, 로짓 𝐳 가 변할 때 target의 로그 확률 변화를 나타낸다. (그리고 로짓 𝐳 는 가중치 𝐰 와 입력 𝐱 에 의해 결정된다.)

𝒟𝐳 : 3.5.1절 Bregman에서의 D 와 다른 의미의 기호이므로 착오 주의

𝐅𝐰=𝔼𝐱~pdata[𝐉𝐳𝐰𝐅𝐳𝐉𝐳𝐰]=𝔼𝐱~pdata[𝐉𝐳𝐰𝔼𝐭~r(·𝐱)[𝒟𝐳𝒟𝐳]𝐉𝐳𝐰]=𝔼𝐱~pdata𝐭~r(·𝐱)[𝐉𝐳𝐰𝒟𝐳𝒟𝐳𝐉𝐳𝐰]=𝔼𝐱~pdata𝐭~r(·𝐱)[𝒟𝐰𝒟𝐰]

마지막 수식은 chain rule( logp 를 가중치로 미분한 gradient는, 출력으로 미분한 gradient에 𝐉 을 곱한 것 )로 얻는다. 즉, backprop 한 번(VJP)+outer product로 샘플 하나를 얻는다. (샘플을 여러 개를 얻어 기댓값을 추정해야 한다.)

𝐱~pdata, 𝐭~r(·𝐱) : 데이터에서 입력 𝐱 를 샘플링하고, 그 입력에서의 모델 예측 분포에서 타깃 𝐭 를 샘플링한다.

모델의 예측 분포에서 샘플링한 타깃 t 은 데이터셋의 정답(입력-정답 쌍)이 아니라는 점에 주목해야 한다. (실제 정답과 무관하며 가중치 업데이트 목적으로는 활용되지 않는다)

Note: backprop의 목적은 통계 𝔼[𝒟𝐰𝒟𝐰]=𝐅𝐰 를 얻는 것이다.

그러므로 정답과 무관하게, '가중치가 흔들리면 모델의 예측 분포가 (방향별로) 얼마나 민감하게 변하나'를 알 수 있다.

Note: Empirical Fisher(Lecture 5)와 혼동 주의

true Fisher와 달리, empirical Fisher는 정답을 보고 샘플별 오차를 반영한다.

𝐅emp=𝔼(𝐱,𝐭)~pdata[𝒟𝐰𝒟𝐰]

True Fisher 𝐅 Empirical Fisher 𝐅emp
타깃의 출처 모델의 예측 분포에서 샘플링 학습 데이터의 실제 타깃(정답)
Hessian과의 관계 GN Hessian과 연관 (KL의 Bregman 구조) (주의) Hessian의 근사로 해석할 수 없음

3.6.1 Relationship to Other Metrics

Fisher metric은 편리하지만, 신경망을 최적화할 때 Fisher가 유일한 정답인 것은 아니다. 대부분의 알고리즘에서 여러 다른 출력 공간 metric을 사용해도 된다.

relationships between curvature matrices and metrics

Matrix 1 Matrix 2 Condition
Hessian 𝐇 GN Hessian 𝐆 선형화한 네트워크이거나, 출력이 손실을 최소화하는 지점(최적점)이면 등호 (Lecture 2)
GN Hessian 𝐆 Pullback metric 𝐆 ρ 가 볼록 손실의 Bregman divergence일 때 (3.5절)
GN Hessian 𝐆 Fisher 𝐅 (𝐳)=logZ(𝐳)𝐳T(𝐭) 형태일 때
(e.g., softmax + cross-entropy) (3.5절)
Pullback metric 𝐆 Fisher 𝐅 ρ 가 KL divergence일 때 (3.3, 3.6절)
GN/pullback 고전적 GN matrix 𝔼[𝐉𝐉] 제곱 오차 손실 / 유클리드 출력 metric일 때 (3.5절)

세 번째 설명의 기호: softmax의 경우, T 는 one-hot, Z=jezj


3.7 Invariance and Differential Geometry

경사 하강법에서 지금까지 살펴본 metric이 필요한 이유는, 최적화에서 일종의 차원 오류(type error)를 저지르기 때문이다. 다음 선형 회귀 예제에서, 단위(차원)을 함께 살펴보자.

$$ \overbrace{y}^{\text{output: \}}=w1{\/min}} \overbrace{x_1}^{\text{input: min}} + \underbrace{w_2}_{\text{\/ft}}x2input: ft+b{\}}

w_1 \leftarrow \overbrace{w_1}^{\text{\/min}}α\^2/\mathrm{min}^2}\, \overbrace{\frac{\mathrm{d}h}{\mathrm{d}w_1}}^{\text{min/\}},w2w2{\/ft}} - \underbrace{\alpha}_{\2/ft2}dhdw2{ft/\}} $$

주목할 부분은 두 가지다.

(1) 학습률은 하나인데, 좌표마다 다른 단위를 요구한다.

(2) 결국 파라미터 단위와, 실제 업데이트 단위가 일치하지 않는다.

Note: 이 때문에 경사 하강법은, 입력의 아핀 변환에 불변이 아니다.

차원 오류를 해결하려면, w2 기준으로는 업데이트에 $\$^2/\mathrm{ft} ^2.(pullbackmetric\mathbf{G}^{-1}$ 행렬에서, 대각 성분이 하는 역할이다.)


3.7.1 Vector, Covector, Riemannian Metric

차원 오류를 바로잡으려면, 각 객체(이동량, 미분, metric 등)가 다른 공간(parameter, output space)에 어떻게 옮겨지는지 알아야 한다.

Note: 후술에서 언급(3.7절 이후)하는 벡터는 모두 tangent vector를 의미한다.

Vectors
(빨간 화살표)
Covectors
(초록 등고선 다발)
Riemannian Metric
(파랑 타원=길이가 동일한 벡터의 등고선)
vectors covectors Riemannian metric

Note: 기하학적 직관


3.7.2 Examples

실제로 pull back과 push forward가 가능한지(어느 공간에서 계산해도 값이 보존되는지) 살펴보자.

(1) covector ω 로 '현재 위치에서 𝐯 만큼 움직이면, 비용이 얼마나 변하는가' 측정

파라미터 공간에서 읽기 출력 공간에서 읽기
(1) covector의 backprop 𝐉ω 을 획득한다. (pulling back)
(2) 파라미터 공간에서, 비용이 얼마나 변하는가를 읽는다: (𝐉ω)𝐯
(1) 파라미터를 𝐯 만큼 움직이면 출력이 𝐉𝐯 만큼 움직인다. (pushing forward)
(2) 출력 공간에서, 비용이 얼마나 변하는가를 읽는다: ω(𝐉𝐯)

pulling back covector

(𝐉ω)𝐯=ω𝐉𝐯

로, 두 값이 같다.

Note: 경사 하강법의 차원 오류는, covector의 계수를 𝐯 자리에 삽입하면서 발생한다.

(2) metric 𝐆𝐳 으로 두 벡터 𝐯,𝐯 의 내적 계산

파라미터 공간에서 읽기 출력 공간에서 읽기
(1) metric을 𝐉𝐆𝐳𝐉 으로 파라미터 공간에 가져온다.
(2) 내적한다: 𝐯(𝐉𝐆𝐳𝐉)𝐯
(1) 두 벡터를 출력 공간으로 보낸다.
(2) 출력 공간의 metric으로 내적한다: (𝐉𝐯)𝐆𝐳(𝐉𝐯)

pulling back metric

앞서 (3.4절에서) 살펴본 𝐉𝐆𝐳𝐉 항이 보이며, 이것이 바로 pullback metric의 정체다.


3.7.3 Summary

경사 하강법은 vector를 대입해야 할 자리에, covector의 계수 𝒥 를 대입한다.

(1) covector의 계수에서 각 성분의 단위가 다르기 때문에, 비교를 위한 metric이 필요하다.

(2) metric과의 내적이 𝒥 가 되는 대리 벡터 𝐮 를 구하면, 차원 오류를 해결할 수 있다.

𝐆𝐰𝐮=𝒥 𝐮=𝐆𝐰1𝒥

이러한 대리 벡터로의 변환을 musical isomorphism이라 부르며, 비용의 미분(covector)에 대한 변환 버전을 natural gradient라 부른다.

~𝒥(𝐰)=𝐆𝐰1𝒥(𝐰)

Note: natural의 의미

𝐆𝐰 를 좌표와 무관하게 정의하면(출력 공간에서 pull back한 metric, Fisher), 업데이트도 좌표 선택에 불변(invariance)이다.