[Pytorch] Optimizer

김명섭·2024년 5월 21일

optimizer는 parameter를 optimize 하는 도구로서, 여러 종류가 있다.
각 종류에 다루기 이전에 pytorch, tensorflow, keras 의 optimizer에 대해서 다뤄보자.

pytorch

pytorch는 auto_grad 방식으로,
model은 torch.nn.Module을 상속받아 정의하고,
criterion는 torch.nn.CrossEntropyLoss 등으로 정의하고,
optimizer는 torch.optim.Adam 등으로 정의한다.
optimizer.zero_grad() 로 누적된 gradient 값을 0으로 만든다.
output = model(input)
loss = criterion(output, target) 로 계산한 후
loss.backward() 를 해서 gradient를 구한다.
opimizer.step()를 하면, gradient를 이용하여 w가 업데이트되고, optimizer에 gradient를 더해서 관리하는 듯 하다.
따라서, RNN 처럼 누적시켜서 업데이트를 할 계획이 아니라면 zero_grad를 사용해야한다.
이때, optimizer.zero_grad()말고 model.zero_grad()도 가능하다. 기본적으로 gradient를 누적하는 방식에서 초기값을 0으로 설정하는 것으로 하나의 optimizer가 여러개의 model의 parameter를 가지고 있을 땐, optimizer.zero_grad()가 그 모든 parameter를 다 업데이트 할테고,
model이 하나고 optimizer 여러개가 각 다른 부분을 업데이트 해야할 때, model.zero_grad()를 이용하면 최소한 그 model의 parameter를 업데이트 하는 optimizer들은 zero_grad()의 효과를 받는다.
optimizer는 parameter를 제네레이터뿐만 아니라 리스트, 딕셔너리 모두 받을 수 있기 때문에 list(model1.parameters()) + list(model2.parameters()) 로 list를 만들어 넣어줄 수 있고,
공식문서에 의하면,
optim.SGD([{'params' : model.base.parameters()}, {'params' : model.classifier.parameters(), 'lr' : 1e-3}], lr=1e-2, momentum=0.9) 이런식으로 정의할 수 있고, 이 경우 base의 lr은 1e-2이고, classifier의 lr은 정의해줬으므로 1e-3이다.

(pytorch custom optimizer)
https://www.geeksforgeeks.org/custom-optimizers-in-pytorch/

tensorflow

tensorflow는 비슷한 방식으로는 gradient tape 방식이 있다.
with tf.GradientTape as tape:
output = model(input, training=True)
loss = loss_fn(target, output) # chanel 처럼 pytorch랑 반대다
grads = tape.gradient(loss, model.trainable_variables)
optimizer.apply_gradients(zip(grads, model.trainable_variables))
방식으로 진행한다.

keras

keras는 model.compile(optimizer=optimizer, loss=loss_fn, metrics=[metrics]) 방식을 사용한다.
각, 문자열 방식과 객체 방식 모두 사용가능하다.
(자동 미분 방식이 있는지 모르겠다.)

jax/flax


[추가]
2가지 종류의 optimizer1과 optimizer2를 model.parameters()에 대해 정의하고, 번갈아가며 적용해도 아무 문제 없는지 확인 필요.
(매번 새롭게 정의하면 놓치고 있는게 있는지는 모르지만 작동한다.)

종류와 원리

AdaBelief 구현

import torch
from torch.optim import Optimizer

class AdaBelief(Optimizer):
    def __init__(self, params, lr=1e-3, betas=(0.9, 0.999), eps=1e-8,
                 weight_decay=0, amsgrad=False):
        if not 0.0 <= lr:
            raise ValueError("Invalid learning rate: {}".format(lr))
        if not 0.0 <= eps:
            raise ValueError("Invalid epsilon value: {}".format(eps))
        if not 0.0 <= betas[0] < 1.0:
            raise ValueError("Invalid beta parameter at index 0: {}".format(betas[0]))
        if not 0.0 <= betas[1] < 1.0:
            raise ValueError("Invalid beta parameter at index 1: {}".format(betas[1]))
        defaults = dict(lr=lr, betas=betas, eps=eps,
                        weight_decay=weight_decay, amsgrad=amsgrad)
        super(AdaBelief, self).__init__(params, defaults)

    def __setstate__(self, state):
        super(AdaBelief, self).__setstate__(state)
        for group in self.param_groups:
            group.setdefault('amsgrad', False)

    @torch.no_grad()
    def step(self, closure=None):
        loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            for p in group['params']:
                if p.grad is None:
                    continue
                grad = p.grad
                if grad.is_sparse:
                    raise RuntimeError('AdaBelief does not support sparse gradients')
                amsgrad = group['amsgrad']

                state = self.state[p]

                # State initialization
                if len(state) == 0:
                    state['step'] = 0
                    state['exp_avg'] = torch.zeros_like(p, memory_format=torch.preserve_format)
                    state['exp_avg_var'] = torch.zeros_like(p, memory_format=torch.preserve_format)
                    if amsgrad:
                        state['max_exp_avg_var'] = torch.zeros_like(p, memory_format=torch.preserve_format)

                exp_avg, exp_avg_var = state['exp_avg'], state['exp_avg_var']
                if amsgrad:
                    max_exp_avg_var = state['max_exp_avg_var']
                beta1, beta2 = group['betas']

                state['step'] += 1
                bias_correction1 = 1 - beta1 ** state['step']
                bias_correction2 = 1 - beta2 ** state['step']

                if group['weight_decay'] != 0:
                    grad = grad.add(p, alpha=group['weight_decay'])

                # Update first and second moment running average
                exp_avg.mul_(beta1).add_(grad, alpha=1 - beta1)
                grad_residual = grad - exp_avg
                exp_avg_var.mul_(beta2).addcmul_(grad_residual, grad_residual, value=1 - beta2)

                if amsgrad:
                    torch.max(max_exp_avg_var, exp_avg_var, out=max_exp_avg_var)
                    denom = (max_exp_avg_var.sqrt() / math.sqrt(bias_correction2)).add_(group['eps'])
                else:
                    denom = (exp_avg_var.sqrt() / math.sqrt(bias_correction2)).add_(group['eps'])

                step_size = group['lr'] / bias_correction1
                
                p.addcdiv_(exp_avg, denom, value=-step_size)

        return loss
profile
ML Engineer

0개의 댓글