Linear Layer Backpropagation

linear layer는 input XX (N×DN \times D)를 입력으로 받고 weight matrix WW (D×MD \times M)이라 하자 이때 layer의 결과물로 output YY (N×MN \times M)가 계산되어 나온다.

실제 예시를 들기 위해 아래처럼 N=2,  D=2,  M=3N=2, \; D=2, \; M=3이라고 가정한다.

X=(x1,1x1,2x2,1x2,2)    W=(w1,1w1,2w1,3w2,1w2,2w2,3) X = \begin{pmatrix} x_{1,1} & x_{1,2} \\ x_{2,1} & x_{2,2} \end{pmatrix} \;\; W = \begin{pmatrix} w_{1,1} & w_{1,2} & w_{1,3} \\ w_{2,1} & w_{2,2} & w_{2,3} \end{pmatrix} Y=XW=(x1,1w1,1+x1,2w2,1x1,1w1,2+x1,2w2,2x1,1w1,3+x1,2w2,3x2,1w1,1+x2,2w2,1x2,1w1,2+x2,2w2,2x2,1w1,3+x2,2w2,3) \begin{split} Y &= XW \\ &= \begin{pmatrix} x_{1,1}w_{1,1} + x_{1,2}w_{2,1} & x_{1,1}w_{1,2} + x_{1,2}w_{2,2} & x_{1,1}w_{1,3} + x_{1,2}w_{2,3} \\ x_{2,1}w_{1,1} + x_{2,2}w_{2,1} & x_{2,1}w_{1,2} + x_{2,2}w_{2,2} & x_{2,1}w_{1,3} + x_{2,2}w_{2,3} \end{pmatrix} \end{split}

forward 과정이후 Loss가 계산되고 Loss에 대응되는 back gradient LY\frac{\partial L}{\partial Y}가 역전파로 들어오게 된다.

LY\frac{\partial L}{\partial Y}의 크기는 output YY와 동일하게 N×MN \times M이고 원소는 다음과 같은 식으로 표현가능하다.

LY=(Ly1,1Ly1,2Ly1,3Ly2,1Ly2,2Ly2,3) \frac{\partial L}{\partial Y} = \begin{pmatrix} \frac{\partial L}{\partial y_{1,1}} & \frac{\partial L}{\partial y_{1,2}} & \frac{\partial L}{\partial y_{1,3}} \\ \frac{\partial L}{\partial y_{2,1}} & \frac{\partial L}{\partial y_{2,2}} & \frac{\partial L}{\partial y_{2,3}} \end{pmatrix}

딥러닝에서 weight를 업데이트하기 위한 LW\frac{\partial L}{\partial W}과 그다음 back propagation의 gradient 값인 LX\frac{\partial L}{\partial X}는 체인룰(chain-rule)을 사용해서 다음과 같이 계산이 가능하다.

LX=LYYXLW=LYYW \frac{\partial L}{\partial X} = \frac{\partial L}{\partial Y} \frac{\partial Y}{\partial X} \quad\quad \frac{\partial L}{\partial W} = \frac{\partial L}{\partial Y} \frac{\partial Y}{\partial W}

먼저 LX\frac{\partial L}{\partial X}에 대해서 표현하면 다음과 같다.

X=(x1,1x1,2x2,1x2,2)    LX=(Lx1,1Lx1,2Lx2,1Lx2,2) X = \begin{pmatrix} x_{1,1} & x_{1,2} \\ x_{2,1} & x_{2,2} \end{pmatrix} \implies \frac{\partial L} {\partial X} = \begin{pmatrix} \frac{\partial L}{\partial x_{1,1}} & \frac{\partial L}{\partial x_{1,2}} \\ \frac{\partial L}{\partial x_{2,1}} & \frac{\partial L}{\partial x_{2,2}} \end{pmatrix}

여기서 원소 Lx_1,1\frac{\partial L}{\partial x\_{1,1}}에 대한 계산식은 다음과 같다.

Lx1,1=i=1Nj=1MLyi,jyi,jx1,1=LYYx1,1 \frac{\partial L}{\partial x_{1,1}} = \sum^{N}_{i=1} \sum^{M}_{j=1} \frac{\partial L}{\partial y_{i,j}} \frac{\partial y_{i,j}}{\partial x_{1,1}} = \frac{\partial L}{\partial Y} \frac{\partial Y}{\partial x_{1,1}}

위 식에서 Yx_1,1\frac{\partial Y}{\partial x\_{1,1}}를 계산하면

Yx1,1=(w1,1w1,2w1,3000) \frac{\partial Y}{\partial x_{1,1}} = \begin{pmatrix} w_{1,1} & w_{1,2} & w_{1,3} \\ 0 & 0 & 0 \end{pmatrix}

Yx_1,1\frac{\partial Y}{\partial x\_{1,1}} 값과 위의 LY\frac{\partial L}{\partial Y}를 이용해서 Lx_1,1\frac{\partial L}{\partial x\_{1,1}}를 계산하면 다음과 같다.

Lx1,1=LYYx1,1=(Ly1,1Ly1,2Ly1,3Ly2,1Ly2,2Ly2,3)(w1,1w1,2w1,3000)=Ly1,1w1,1+Ly1,2w1,2+Ly1,3w1,3 \begin{split} \frac{\partial L}{\partial x_{1,1}} &= \frac{\partial L}{\partial Y} \frac{\partial Y}{\partial x_{1,1}} \\ &= \begin{pmatrix} \frac{\partial L}{\partial y_{1,1}} & \frac{\partial L}{\partial y_{1,2}} & \frac{\partial L}{\partial y_{1,3}} \\ \frac{\partial L}{\partial y_{2,1}} & \frac{\partial L}{\partial y_{2,2}} & \frac{\partial L}{\partial y_{2,3}} \end{pmatrix} \begin{pmatrix} w_{1,1} & w_{1,2} & w_{1,3} \\ 0 & 0 & 0 \end{pmatrix} \\ &= \frac{\partial L}{\partial y_{1,1}} w_{1,1} + \frac{\partial L}{\partial y_{1,2}} w_{1,2} + \frac{\partial L}{\partial y_{1,3}} w_{1,3} \end{split}

마찬가지로 x_1,2,  x_2,1,  x_2,2x\_{1,2}, \; x\_{2,1}, \; x\_{2,2}에 대해서도 똑같이 계산하면 다음과 같다.

Lx1,2=LYYx1,2=(Ly1,1Ly1,2Ly1,3Ly2,1Ly2,2Ly2,3)(w2,1w2,2w2,3000)=Ly1,1w2,1+Ly1,2w2,2+Ly1,3w2,3 \begin{split} \frac{\partial L}{\partial x_{1,2}} &= \frac{\partial L}{\partial Y} \frac{\partial Y}{\partial x_{1,2}} \\ &= \begin{pmatrix} \frac{\partial L}{\partial y_{1,1}} & \frac{\partial L}{\partial y_{1,2}} & \frac{\partial L}{\partial y_{1,3}} \\ \frac{\partial L}{\partial y_{2,1}} & \frac{\partial L}{\partial y_{2,2}} & \frac{\partial L}{\partial y_{2,3}} \end{pmatrix} \begin{pmatrix} w_{2,1} & w_{2,2} & w_{2,3} \\ 0 & 0 & 0 \end{pmatrix} \\ &= \frac{\partial L}{\partial y_{1,1}} w_{2,1} + \frac{\partial L}{\partial y_{1,2}} w_{2,2} + \frac{\partial L}{\partial y_{1,3}} w_{2,3} \end{split} Lx2,1=LYYx2,1=(Ly1,1Ly1,2Ly1,3Ly2,1Ly2,2Ly2,3)(000w1,1w1,2w1,3)=Ly2,1w1,1+Ly2,2w1,2+Ly2,3w1,3 \begin{split} \frac{\partial L}{\partial x_{2,1}} &= \frac{\partial L}{\partial Y} \frac{\partial Y}{\partial x_{2,1}} \\ &= \begin{pmatrix} \frac{\partial L}{\partial y_{1,1}} & \frac{\partial L}{\partial y_{1,2}} & \frac{\partial L}{\partial y_{1,3}} \\ \frac{\partial L}{\partial y_{2,1}} & \frac{\partial L}{\partial y_{2,2}} & \frac{\partial L}{\partial y_{2,3}} \end{pmatrix} \begin{pmatrix} 0 & 0 & 0 \\ w_{1,1} & w_{1,2} & w_{1,3} \end{pmatrix} \\ &= \frac{\partial L}{\partial y_{2,1}} w_{1,1} + \frac{\partial L}{\partial y_{2,2}} w_{1,2} + \frac{\partial L}{\partial y_{2,3}} w_{1,3} \end{split} Lx2,2=LYYx2,2=(Ly1,1Ly1,2Ly1,3Ly2,1Ly2,2Ly2,3)(000w2,1w2,2w2,3)=Ly2,1w2,1+Ly2,2w2,2+Ly2,3w2,3 \begin{split} \frac{\partial L}{\partial x_{2,2}} &= \frac{\partial L}{\partial Y} \frac{\partial Y}{\partial x_{2,2}} \\ &= \begin{pmatrix} \frac{\partial L}{\partial y_{1,1}} & \frac{\partial L}{\partial y_{1,2}} & \frac{\partial L}{\partial y_{1,3}} \\ \frac{\partial L}{\partial y_{2,1}} & \frac{\partial L}{\partial y_{2,2}} & \frac{\partial L}{\partial y_{2,3}} \end{pmatrix} \begin{pmatrix} 0 & 0 & 0 \\ w_{2,1} & w_{2,2} & w_{2,3} \end{pmatrix} \\ &= \frac{\partial L}{\partial y_{2,1}} w_{2,1} + \frac{\partial L}{\partial y_{2,2}} w_{2,2} + \frac{\partial L}{\partial y_{2,3}} w_{2,3} \end{split}

LX\frac{\partial L}{\partial X}를 위에서 계산된 값으로 표현 하면 다음과 같다.

LX=(Ly1,1w1,1+Ly1,2w1,2+Ly1,3w1,3Ly1,1w2,1+Ly1,2w2,2+Ly1,3w2,3Ly2,1w1,1+Ly2,2w1,2+Ly2,3w1,3Ly2,1w2,1+Ly2,2w2,2+Ly2,3w2,3) \frac{\partial L}{\partial X} = \begin{pmatrix} \frac{\partial L}{\partial y_{1,1}} w_{1,1} + \frac{\partial L}{\partial y_{1,2}} w_{1,2} + \frac{\partial L}{\partial y_{1,3}} w_{1,3} & \frac{\partial L}{\partial y_{1,1}} w_{2,1} + \frac{\partial L}{\partial y_{1,2}} w_{2,2} + \frac{\partial L}{\partial y_{1,3}} w_{2,3} \\ \frac{\partial L}{\partial y_{2,1}} w_{1,1} + \frac{\partial L}{\partial y_{2,2}} w_{1,2} + \frac{\partial L}{\partial y_{2,3}} w_{1,3} & \frac{\partial L}{\partial y_{2,1}} w_{2,1} + \frac{\partial L}{\partial y_{2,2}} w_{2,2} + \frac{\partial L}{\partial y_{2,3}} w_{2,3} \end{pmatrix}

계산된 값에서 규칙성을 토대로 2개의 matrix로 분리가 가능하다. 이를 분리하면 다음과 같다.

LX=(Ly1,1Ly1,2Ly1,3Ly2,1Ly2,2Ly2,3)(w1,1w2,1w1,2w2,2w1,3w2,3)=LYWT \begin{split} \frac{\partial L}{\partial X} &= \begin{pmatrix} \frac{\partial L}{\partial y_{1,1}} & \frac{\partial L}{\partial y_{1,2}} & \frac{\partial L}{\partial y_{1,3}} \\ \frac{\partial L}{\partial y_{2,1}} & \frac{\partial L}{\partial y_{2,2}} & \frac{\partial L}{\partial y_{2,3}} \end{pmatrix} \begin{pmatrix} w_{1,1} & w_{2,1} \\ w_{1,2} & w_{2,2} \\ w_{1,3} & w_{2,3} \end{pmatrix} \\ &= \frac{\partial L}{\partial Y} W^{T} \end{split}

마찬가지로 weight WW에 대해서도 반복해서 구할 수 있다.

W=(w1,1w1,2w1,3w2,1w2,2w2,3)    LW=(Lw1,1Lw1,2Lw1,3Lw2,1Lw2,2Lw2,3) W = \begin{pmatrix} w_{1,1} & w_{1,2} & w_{1,3} \\ w_{2,1} & w_{2,2} & w_{2,3} \end{pmatrix} \implies \frac{\partial L}{\partial W} = \begin{pmatrix} \frac{\partial L}{\partial w_{1,1}} & \frac{\partial L}{\partial w_{1,2}} & \frac{\partial L}{\partial w_{1,3}} \\ \frac{\partial L}{\partial w_{2,1}} & \frac{\partial L}{\partial w_{2,2}} & \frac{\partial L}{\partial w_{2,3}} \end{pmatrix}

첫번째 원소 Lw_1,1\frac{\partial L}{\partial w\_{1,1}}에 대한 계산식은 다음과 같이 체인룰로 표현 가능하다.

Lw1,1=i=1Nj=1MLyi,jyi,jw1,1=LYYw1,1 \frac{\partial L}{\partial w_{1,1}} = \sum^{N}_{i=1} \sum^{M}_{j=1} \frac{\partial L}{\partial y_{i,j}} \frac{\partial y_{i,j}}{\partial w_{1,1}} = \frac{\partial L}{\partial Y} \frac{\partial Y}{\partial w_{1,1}}

여기에 YYw_1,1w\_{1,1}에 대해 부분 적분한 식을 대입해서 정리하면 다음과 같다.

Lw1,1=LYYw1,1=(Ly1,1Ly1,2Ly1,3Ly2,1Ly2,2Ly2,3)(x1,100x2,100)=Ly1,1x1,1+Ly2,1x2,1 \begin{split} \frac{\partial L}{\partial w_{1,1}} &= \frac{\partial L}{\partial Y} \frac{\partial Y}{\partial w_{1,1}} \\ &= \begin{pmatrix} \frac{\partial L}{\partial y_{1,1}} & \frac{\partial L}{\partial y_{1,2}} & \frac{\partial L}{\partial y_{1,3}} \\ \frac{\partial L}{\partial y_{2,1}} & \frac{\partial L}{\partial y_{2,2}} & \frac{\partial L}{\partial y_{2,3}} \end{pmatrix} \begin{pmatrix} x_{1,1} & 0 & 0 \\ x_{2,1} & 0 & 0 \end{pmatrix} \\ &= \frac{\partial L}{\partial y_{1,1}} x_{1,1} + \frac{\partial L}{\partial y_{2,1}} x_{2,1} \end{split}

이를 다른 원소에 대해서도 반복계산하여 LW\frac{\partial L}{\partial W}으로 표현하면 다음과 같다.

LW=(Ly1,1x1,1+Ly2,1x2,1Ly1,2x1,1+Ly2,2x2,1Ly1,3x1,1+Ly2,3x2,1Ly1,1x1,2+Ly2,1x2,2Ly1,2x1,2+Ly2,2x2,2Ly1,3x1,2+Ly2,3x2,2) \frac{\partial L}{\partial W} = \begin{pmatrix} \frac{\partial L}{\partial y_{1,1}} x_{1,1} + \frac{\partial L}{\partial y_{2,1}} x_{2,1} & \frac{\partial L}{\partial y_{1,2}} x_{1,1} + \frac{\partial L}{\partial y_{2,2}} x_{2,1} & \frac{\partial L}{\partial y_{1,3}} x_{1,1} + \frac{\partial L}{\partial y_{2,3}} x_{2,1} \\ \frac{\partial L}{\partial y_{1,1}} x_{1,2} + \frac{\partial L}{\partial y_{2,1}} x_{2,2} & \frac{\partial L}{\partial y_{1,2}} x_{1,2} + \frac{\partial L}{\partial y_{2,2}} x_{2,2} & \frac{\partial L}{\partial y_{1,3}} x_{1,2} + \frac{\partial L}{\partial y_{2,3}} x_{2,2} \end{pmatrix}

이를 각각의 matrix로 분리해서 표현하면 다음과 같이 표현된다.

LW=(x1,1x2,1x1,2x2,2)(Ly1,1Ly1,2Ly1,3Ly2,1Ly2,2Ly2,3)=XTLY \begin{split} \frac{\partial L}{\partial W} &= \begin{pmatrix} x_{1,1} & x_{2,1} \\ x_{1,2} & x_{2,2} \end{pmatrix} \begin{pmatrix} \frac{\partial L}{\partial y_{1,1}} & \frac{\partial L}{\partial y_{1,2}} & \frac{\partial L}{\partial y_{1,3}} \\ \frac{\partial L}{\partial y_{2,1}} & \frac{\partial L}{\partial y_{2,2}} & \frac{\partial L}{\partial y_{2,3}} \end{pmatrix} \\ &= X^{T} \frac{\partial L}{\partial Y} \end{split}

따라서 정리하면 Weight WW와 input XX에 대한 gradient는 다음과 같이 간단한 공식으로 표현이 가능하다.

LW=XTLYLX=LYWT \frac{\partial L}{\partial W} = X^{T} \frac{\partial L}{\partial Y} \quad\quad \frac{\partial L}{\partial X} = \frac{\partial L}{\partial Y} W^{T}

bias 포함

X=(x1,1x1,2x2,1x2,2)    W=(w1,1w1,2w1,3w2,1w2,2w2,3)    B=(b1,1b1,2b1,3) X = \begin{pmatrix} x_{1,1} & x_{1,2} \\ x_{2,1} & x_{2,2} \end{pmatrix} \;\; W = \begin{pmatrix} w_{1,1} & w_{1,2} & w_{1,3} \\ w_{2,1} & w_{2,2} & w_{2,3} \end{pmatrix} \;\; B = \begin{pmatrix} b_{1,1} & b_{1,2} & b_{1,3} \end{pmatrix} Y=XW+B=(x1,1w1,1+x1,2w2,1+b1,1x1,1w1,2+x1,2w2,2+b1,2x1,1w1,3+x1,2w2,3+b1,3x2,1w1,1+x2,2w2,1+b1,1x2,1w1,2+x2,2w2,2+b1,2x2,1w1,3+x2,2w2,3+b1,3) \begin{split} Y &= XW + B \\ &= \begin{pmatrix} x_{1,1}w_{1,1} + x_{1,2}w_{2,1} + b_{1,1} & x_{1,1}w_{1,2} + x_{1,2}w_{2,2} + b_{1,2} & x_{1,1}w_{1,3} + x_{1,2}w_{2,3} + b_{1,3} \\ x_{2,1}w_{1,1} + x_{2,2}w_{2,1} + b_{1,1} & x_{2,1}w_{1,2} + x_{2,2}w_{2,2} + b_{1,2} & x_{2,1}w_{1,3} + x_{2,2}w_{2,3} + b_{1,3} \end{pmatrix} \end{split}

위의 식Y=XW+BY=XW+B에서 보통 B앞에 broadcast matrix C(N×1)C(N \times 1)가 생략되었지만 포함되어있다. 이를 식으로 표현하면 다음과 같다.

Y=XW+(11)(b1,1b1,2b1,3) Y = XW + \begin{pmatrix} 1 \\ 1 \end{pmatrix} \begin{pmatrix} b_{1,1} & b_{1,2} & b_{1,3} \end{pmatrix}

Bias에 대한 편미분은 다음과 같고 각 원소에 대해서 chain rule을 적용하여 계산한다.

LB=(Lb1,1Lb1,2Lb1,3) \frac{\partial L}{\partial B} = \begin{pmatrix} \frac{\partial L}{\partial b_{1,1}} & \frac{\partial L}{\partial b_{1,2}} & \frac{\partial L}{\partial b_{1,3}} \end{pmatrix} Lb1,1=LYYb1,1=(Ly1,1Ly1,2Ly1,3Ly2,1Ly2,2Ly2,3)(100100)=Ly1,1+Ly2,1 \begin{split} \frac{\partial L}{\partial b_{1,1}} &= \frac{\partial L}{\partial Y} \frac{\partial Y}{\partial b_{1,1}} \\ &= \begin{pmatrix} \frac{\partial L}{\partial y_{1,1}} & \frac{\partial L}{\partial y_{1,2}} & \frac{\partial L}{\partial y_{1,3}} \\ \frac{\partial L}{\partial y_{2,1}} & \frac{\partial L}{\partial y_{2,2}} & \frac{\partial L}{\partial y_{2,3}} \end{pmatrix} \begin{pmatrix} 1 & 0 & 0 \\ 1 & 0 & 0 \end{pmatrix} \\ &= \frac{\partial L}{\partial y_{1,1}} + \frac{\partial L}{\partial y_{2,1}} \end{split}

모든 원소에 대해서 계산하여 구하면 LB\frac{\partial L}{\partial B}는 다음과 같이 계산된다. 그리고 이를 분리하면 다음과 같다.

LB=(Ly1,1+Ly2,1Ly1,2+Ly2,2Ly1,3+Ly2,3)=(11)(Ly1,1Ly1,2Ly1,3Ly2,1Ly2,2Ly2,3) \begin{split} \frac{\partial L}{\partial B} &= \begin{pmatrix} \frac{\partial L}{\partial y_{1,1}} + \frac{\partial L}{\partial y_{2,1}} & \frac{\partial L}{\partial y_{1,2}} + \frac{\partial L}{\partial y_{2,2}} & \frac{\partial L}{\partial y_{1,3}} + \frac{\partial L}{\partial y_{2,3}} \end{pmatrix} \\ &= \begin{pmatrix} 1 & 1 \end{pmatrix} \begin{pmatrix} \frac{\partial L}{\partial y_{1,1}} & \frac{\partial L}{\partial y_{1,2}} & \frac{\partial L}{\partial y_{1,3}} \\ \frac{\partial L}{\partial y_{2,1}} & \frac{\partial L}{\partial y_{2,2}} & \frac{\partial L}{\partial y_{2,3}} \end{pmatrix} \end{split}

위 식에서 앞의 1로 이루어진 matrix는 X와 W,B를 확장해보면 broatcast matrix CTC^T가 된다는걸 알 수 있다. 따라서 다시 표현 하면 다음과 같다.

LB=CTLY \frac{\partial L}{\partial B} = C^{T} \frac{\partial L}{\partial Y}

Reference

[참조] cs231n