Linear Layer Backpropagation
linear layer는 input X (N×D)를 입력으로 받고 weight matrix W (D×M)이라 하자
이때 layer의 결과물로 output Y (N×M)가 계산되어 나온다.
실제 예시를 들기 위해 아래처럼 N=2,D=2,M=3이라고 가정한다.
X=(x1,1x2,1x1,2x2,2)W=(w1,1w2,1w1,2w2,2w1,3w2,3)
Y=XW=(x1,1w1,1+x1,2w2,1x2,1w1,1+x2,2w2,1x1,1w1,2+x1,2w2,2x2,1w1,2+x2,2w2,2x1,1w1,3+x1,2w2,3x2,1w1,3+x2,2w2,3)
forward 과정이후 Loss가 계산되고 Loss에 대응되는 back gradient ∂Y∂L가 역전파로 들어오게 된다.
∂Y∂L의 크기는 output Y와 동일하게 N×M이고 원소는 다음과 같은 식으로 표현가능하다.
∂Y∂L=(∂y1,1∂L∂y2,1∂L∂y1,2∂L∂y2,2∂L∂y1,3∂L∂y2,3∂L)
딥러닝에서 weight를 업데이트하기 위한 ∂W∂L과 그다음 back propagation의 gradient 값인
∂X∂L는 체인룰(chain-rule)을 사용해서 다음과 같이 계산이 가능하다.
∂X∂L=∂Y∂L∂X∂Y∂W∂L=∂Y∂L∂W∂Y
먼저 ∂X∂L에 대해서 표현하면 다음과 같다.
X=(x1,1x2,1x1,2x2,2)⟹∂X∂L=(∂x1,1∂L∂x2,1∂L∂x1,2∂L∂x2,2∂L)
여기서 원소 ∂x_1,1∂L에 대한 계산식은 다음과 같다.
∂x1,1∂L=i=1∑Nj=1∑M∂yi,j∂L∂x1,1∂yi,j=∂Y∂L∂x1,1∂Y
위 식에서 ∂x_1,1∂Y를 계산하면
∂x1,1∂Y=(w1,10w1,20w1,30)
∂x_1,1∂Y 값과 위의 ∂Y∂L를 이용해서
∂x_1,1∂L를 계산하면 다음과 같다.
∂x1,1∂L=∂Y∂L∂x1,1∂Y=(∂y1,1∂L∂y2,1∂L∂y1,2∂L∂y2,2∂L∂y1,3∂L∂y2,3∂L)(w1,10w1,20w1,30)=∂y1,1∂Lw1,1+∂y1,2∂Lw1,2+∂y1,3∂Lw1,3
마찬가지로 x_1,2,x_2,1,x_2,2에 대해서도 똑같이 계산하면 다음과 같다.
∂x1,2∂L=∂Y∂L∂x1,2∂Y=(∂y1,1∂L∂y2,1∂L∂y1,2∂L∂y2,2∂L∂y1,3∂L∂y2,3∂L)(w2,10w2,20w2,30)=∂y1,1∂Lw2,1+∂y1,2∂Lw2,2+∂y1,3∂Lw2,3
∂x2,1∂L=∂Y∂L∂x2,1∂Y=(∂y1,1∂L∂y2,1∂L∂y1,2∂L∂y2,2∂L∂y1,3∂L∂y2,3∂L)(0w1,10w1,20w1,3)=∂y2,1∂Lw1,1+∂y2,2∂Lw1,2+∂y2,3∂Lw1,3
∂x2,2∂L=∂Y∂L∂x2,2∂Y=(∂y1,1∂L∂y2,1∂L∂y1,2∂L∂y2,2∂L∂y1,3∂L∂y2,3∂L)(0w2,10w2,20w2,3)=∂y2,1∂Lw2,1+∂y2,2∂Lw2,2+∂y2,3∂Lw2,3
∂X∂L를 위에서 계산된 값으로 표현 하면 다음과 같다.
∂X∂L=(∂y1,1∂Lw1,1+∂y1,2∂Lw1,2+∂y1,3∂Lw1,3∂y2,1∂Lw1,1+∂y2,2∂Lw1,2+∂y2,3∂Lw1,3∂y1,1∂Lw2,1+∂y1,2∂Lw2,2+∂y1,3∂Lw2,3∂y2,1∂Lw2,1+∂y2,2∂Lw2,2+∂y2,3∂Lw2,3)
계산된 값에서 규칙성을 토대로 2개의 matrix로 분리가 가능하다. 이를 분리하면 다음과 같다.
∂X∂L=(∂y1,1∂L∂y2,1∂L∂y1,2∂L∂y2,2∂L∂y1,3∂L∂y2,3∂L)w1,1w1,2w1,3w2,1w2,2w2,3=∂Y∂LWT
마찬가지로 weight W에 대해서도 반복해서 구할 수 있다.
W=(w1,1w2,1w1,2w2,2w1,3w2,3)⟹∂W∂L=(∂w1,1∂L∂w2,1∂L∂w1,2∂L∂w2,2∂L∂w1,3∂L∂w2,3∂L)
첫번째 원소 ∂w_1,1∂L에 대한 계산식은 다음과 같이 체인룰로 표현 가능하다.
∂w1,1∂L=i=1∑Nj=1∑M∂yi,j∂L∂w1,1∂yi,j=∂Y∂L∂w1,1∂Y
여기에 Y를 w_1,1에 대해 부분 적분한 식을 대입해서 정리하면 다음과 같다.
∂w1,1∂L=∂Y∂L∂w1,1∂Y=(∂y1,1∂L∂y2,1∂L∂y1,2∂L∂y2,2∂L∂y1,3∂L∂y2,3∂L)(x1,1x2,10000)=∂y1,1∂Lx1,1+∂y2,1∂Lx2,1
이를 다른 원소에 대해서도 반복계산하여 ∂W∂L으로 표현하면 다음과 같다.
∂W∂L=(∂y1,1∂Lx1,1+∂y2,1∂Lx2,1∂y1,1∂Lx1,2+∂y2,1∂Lx2,2∂y1,2∂Lx1,1+∂y2,2∂Lx2,1∂y1,2∂Lx1,2+∂y2,2∂Lx2,2∂y1,3∂Lx1,1+∂y2,3∂Lx2,1∂y1,3∂Lx1,2+∂y2,3∂Lx2,2)
이를 각각의 matrix로 분리해서 표현하면 다음과 같이 표현된다.
∂W∂L=(x1,1x1,2x2,1x2,2)(∂y1,1∂L∂y2,1∂L∂y1,2∂L∂y2,2∂L∂y1,3∂L∂y2,3∂L)=XT∂Y∂L
따라서 정리하면 Weight W와 input X에 대한 gradient는 다음과 같이 간단한 공식으로 표현이 가능하다.
∂W∂L=XT∂Y∂L∂X∂L=∂Y∂LWT
bias 포함
X=(x1,1x2,1x1,2x2,2)W=(w1,1w2,1w1,2w2,2w1,3w2,3)B=(b1,1b1,2b1,3)
Y=XW+B=(x1,1w1,1+x1,2w2,1+b1,1x2,1w1,1+x2,2w2,1+b1,1x1,1w1,2+x1,2w2,2+b1,2x2,1w1,2+x2,2w2,2+b1,2x1,1w1,3+x1,2w2,3+b1,3x2,1w1,3+x2,2w2,3+b1,3)
위의 식Y=XW+B에서 보통 B앞에 broadcast matrix C(N×1)가 생략되었지만 포함되어있다. 이를 식으로 표현하면 다음과 같다.
Y=XW+(11)(b1,1b1,2b1,3)
Bias에 대한 편미분은 다음과 같고 각 원소에 대해서 chain rule을 적용하여 계산한다.
∂B∂L=(∂b1,1∂L∂b1,2∂L∂b1,3∂L)
∂b1,1∂L=∂Y∂L∂b1,1∂Y=(∂y1,1∂L∂y2,1∂L∂y1,2∂L∂y2,2∂L∂y1,3∂L∂y2,3∂L)(110000)=∂y1,1∂L+∂y2,1∂L
모든 원소에 대해서 계산하여 구하면 ∂B∂L는 다음과 같이 계산된다. 그리고 이를 분리하면 다음과 같다.
∂B∂L=(∂y1,1∂L+∂y2,1∂L∂y1,2∂L+∂y2,2∂L∂y1,3∂L+∂y2,3∂L)=(11)(∂y1,1∂L∂y2,1∂L∂y1,2∂L∂y2,2∂L∂y1,3∂L∂y2,3∂L)
위 식에서 앞의 1로 이루어진 matrix는 X와 W,B를 확장해보면 broatcast matrix CT가 된다는걸 알 수 있다. 따라서 다시 표현 하면 다음과 같다.
∂B∂L=CT∂Y∂L
Reference
[참조] cs231n