오차역전파 (Backpropagation)

ys10·2024년 5월 10일

오차역전파법:

오차역전파법이란 내가 뽑고 싶은 아웃풋과, 모델의 아웃풋의 오차를 구해놓은 뒤에, 그 결과를 뒤로 보내면서 각 노드의 가중치를 갱신해나가는 것이다.

우리가 어떠한 수식으로 계산을 하면서 값을 다음 층으로 보내는 것을 Foward Pass(순전파) 라고 한다.

그리고 이 과정에서 앞 노드의 값이 변함에 따라, 다음 노드의 값이 변화는 정도인 미분이 가능하다. (편미분)

나는 오차를 역순으로 보내며 가중치를 갱신하는 것이라고 이해했다.

계산 그래프의 역전파

계산 그래프에서는 연전파를 따라 온 신호에 노드의 국소적 미분을 곱한 후 다음 노드로 전달한다.

연쇄법칙

그냥 머 당연한 소리이다.

덧셈 노드의 역전파

뭐 더하기니까 그대로 신호를 하류로 보낼 것이다.

곱셈 노드의 역전파

z = xy
∂z/∂x = y
∂z/∂y = x
여기서는 상류에서 온 값에 각각의 곱셈 노드의 순전파 때 온 입력 신호를 서로 바꾼 값을 곱해 하류로 전달한다.

그러므로 순방향 입력 신호의 값을 저장해 둘 필요가 있다.

단순한 계층 구현하기

곱셈 계층 구현

class MulLayer:
    def __init__(self):
        self.x = None
        self.y = None

    def forward(self, x, y):
        self.x = x
        self.y = y
        out = x * y
        return out

    def backward(self, dout):
        dx = dout * self.y
        dy = dout * self.x
        return dx, dy

여기서는 위에 설명했던데로 foward pass할 때 받은 x,y를 저장해두고 backward할 때에 x와 y를 바꿔서 곱해준다.

덧셈 계층 구현

class AddLayer:
    def __init__(self):
        pass

    def forward(self, x, y):
        out = x + y
        return out

    def backward(self, dout):
        dx = dout * 1
        dy = dout * 1
        return dx, dy

그냥 더하기이다. 역전파할 때도 그대로..

활성화 함수 계층 구현하기

이제부터는 계산그래프를 신경망에 적용한다.

Relu 함수의 구현:

Relu는 단순하게 x가 0보다 작거나 같으면 0, 크면 x를 반환하는 함수였다.

그러니, 곱셈이랑 비슷하게 순전파 때의 x의 값을 저장해 놓았다가, 미분을 하여, 1혹은 0을 곱하여 뒤로 보내면 된다.

class Relu:
    def __init__(self):
        self.mask=None
    
    def foward(self,x):
        self.mask=(x<=0) #x가 0보다 작거나 같으면 맞으니까 1 저장
        out = x.copy()
        out[self.mask]=0
        return out

    def backward(self,dout):
        dout[self.mask] = 0
        return dout

시그모이드 함수의 구현:

뭐 계산그래프 단계가 뭐라뭐라 써있었는데 그냥
1단계: / 보내기
y = 1/x
∂y/∂x = -1/x^2 = -y²
상류에서 흘러온 값에 -y^2(순전파의 출력을 제곱하고 마이너스)을 곱해서 하류로 전달 : -∂L/∂y*y²

3단계: exp함수
∂y/∂x = exp(x)
상류의 값에 순전파 때의 출력(이 경우엔 exp(-x))을 곱해 하류로 전달 : -∂L/∂yy²exp(-x)
만 좀 중요한 것 같다.

class Sigmoid:
   def __init__(self):
       self.out = None

   def forward(self, x):
       out = 1 / (1 + np.exp(-x))
       self.out = out

       return out

   def backward(self, dout):
       dx = dout * (1.0 - self.out) * self.out

       return dx

아이구 쓰는 중인데 줌을 키셨네.. 허허

profile
hyu infosys24

0개의 댓글