In [11]:
import numpy as np

#Gradient Descent

In [12]:
def df_w(W):
  """
  Thực hiện tính gradient của dw1 và dw2
  Arguments:
  W -- np.array [w1, w2]
  Returns:
  dW -- np.array [dw1, dw2], array chứa giá trị đạo hàm theo w1 và w2
  """
  w1, w2 = W
  dw1 = 0.2 * w1
  dw2 = 4 * w2
  dW = np.array([dw1, dw2])
  return dW


def sgd(W, dW, lr):
  """
  Thực hiện thuật toán Gradient Descent để update w1 và w2
  Arguments:
  W -- np.array [w1, w2]
  dW -- np.array [dw1, dw2], array chứa giá trị đạo hàm theo w1 và w2
  lr -- float: learning rate
  Returns:
  W -- np.array [w1, w2] w1 và w2 sau khi update
  """
  W = W - lr * dW
  return W


def train_p1(optimizer, lr, epochs):
  """
  Thực hiện tìm điểm minium của function (1) dựa vào thuật toán được
  truyền vào từ optimizer
  Arguments:
  optimize: function thực hiện thuật toán optimization cụ thể
  lr -- float: learning rate
  epochs -- int: số lượng lần (epoch) lặp để tìm minium
  Returns:
  results -- list: list các cặp điểm [w1, w2] sau mỗi epoch (mỗi lần cập nhật)
  """

  # initial point
  W = np.array([-5, -2], dtype=np.float32)
  # list of results
  results = [W]

  # Tạo vòng lặp theo số epochs
  # Tìm gradient dW gồm dw1 và dw2
  # dùng thuật toán optimization cập nhật w1, w2
  # append cặp [w1, w2] vào list results
  for i in range(epochs):
    dW = df_w(W)
    W = optimizer(W, dW, lr)
    results.append(W)
    print(f'Epoch {i + 1}: w1 = {W[0]}, w2 = {W[1]}')
  return results

In [13]:
train_sgd = train_p1(sgd, 0.4, 30)

Epoch 1: w1 = -4.6, w2 = 1.2000000000000002
Epoch 2: w1 = -4.231999999999999, w2 = -0.7200000000000002
Epoch 3: w1 = -3.893439999999999, w2 = 0.43200000000000016
Epoch 4: w1 = -3.5819647999999993, w2 = -0.2592000000000001
Epoch 5: w1 = -3.2954076159999994, w2 = 0.1555200000000001
Epoch 6: w1 = -3.0317750067199993, w2 = -0.09331200000000006
Epoch 7: w1 = -2.7892330061823993, w2 = 0.05598720000000004
Epoch 8: w1 = -2.5660943656878072, w2 = -0.03359232000000004
Epoch 9: w1 = -2.360806816432783, w2 = 0.020155392000000022
Epoch 10: w1 = -2.1719422711181604, w2 = -0.012093235200000017
Epoch 11: w1 = -1.9981868894287076, w2 = 0.007255941120000012
Epoch 12: w1 = -1.838331938274411, w2 = -0.0043535646720000085
Epoch 13: w1 = -1.691265383212458, w2 = 0.0026121388032000056
Epoch 14: w1 = -1.5559641525554613, w2 = -0.0015672832819200039
Epoch 15: w1 = -1.4314870203510244, w2 = 0.0009403699691520025
Epoch 16: w1 = -1.3169680587229424, w2 = -0.0005642219814912016
Epoch 17: w1 = -1.211610614025107, w

#Gradient Descent + Momentum

In [14]:
def df_w(W):
  w1, w2 = W
  dw1 = 0.2 * w1
  dw2 = 4 * w2
  dW = np.array([dw1, dw2])
  return dW


def sgd_momentum(W, dW, lr, beta, V):
  V = beta * V + (1 - beta) * dW
  W = W - lr * V
  return W, V


def train_p2(optimizer, lr, beta, epochs):
  W = np.array([-5, -2], dtype=np.float32)
  V = np.zeros_like(W)
  results = [W]
  for i in range(epochs):
      dW = df_w(W)
      W, V = optimizer(W, dW, lr, beta, V)
      results.append(W)
      print(f'Epoch {i + 1}: w1 = {W[0]}, w2 = {W[1]}')
  return results



In [15]:
train_sgd_momentum = train_p2(sgd_momentum, 0.6, 0.5, 30)

Epoch 1: w1 = -4.7, w2 = 0.3999999999999999
Epoch 2: w1 = -4.268, w2 = 1.12
Epoch 3: w1 = -3.7959199999999997, w2 = 0.13600000000000012
Epoch 4: w1 = -3.3321248, w2 = -0.5192
Epoch 5: w1 = -2.900299712, w2 = -0.22376000000000013
Epoch 6: w1 = -2.5103691852799996, w2 = 0.19247199999999992
Epoch 7: w1 = -2.1647817708031996, w2 = 0.16962160000000004
Epoch 8: w1 = -1.8621011573166075, w2 = -0.04534951999999995
Epoch 9: w1 = -1.599034781134315, w2 = -0.09841565599999999
Epoch 10: w1 = -1.3715595061751098, w2 = -0.0068499368000000255
Epoch 11: w1 = -1.1755282983250006, w2 = 0.04715284695999999
Epoch 12: w1 = -1.006980996500446, w2 = 0.01757082248800001
Epoch 13: w1 = -0.8622884857981419, w2 = -0.018305176733599993
Epoch 14: w1 = -0.7382049212991013, w2 = -0.01427696426408
Epoch 15: w1 = -0.6318708437716349, w2 = 0.004869499087575998
Epoch 16: w1 = -0.5407915543816036, w2 = 0.0085993318583128
Epoch 17: w1 = -0.4628044164236918, w2 = 0.00014505001370584102
Epoch 18: w1 = -0.39604258245931434, 

#RMSProp

In [16]:
def df_w(W):
  w1, w2 = W
  dw1 = 0.2 * w1
  dw2 = 4 * w2
  dW = np.array([dw1, dw2])
  return dW


def RMSProp(W, dW, lr, S, gamma):
  epsilon = 1e-6
  S = gamma * S + (1 - gamma) * dW ** 2
  adapt_lr = lr / np.sqrt(S + epsilon)
  W = W - adapt_lr * dW
  return W, S


def train_p3(optimizer, lr, epoch):
  W = np.array([-5, -2], dtype=np.float32)
  S = np.array([0, 0], dtype=np.float32)
  results = [W]
  for i in range(epoch):
    dW = df_w(W)
    W, S = optimizer(W, dW, lr, S, 0.9)
    results.append(W)
    print(f'Epoch {i + 1}: w1 = {W[0]}, w2 = {W[1]}')
  return results

In [17]:
train_RMSProp = train_p3(RMSProp, 0.3, 30)

Epoch 1: w1 = -4.051321445330401, w2 = -1.0513167760653601
Epoch 2: w1 = -3.435197540710313, w2 = -0.59152342591607
Epoch 3: w1 = -2.9589369293489796, w2 = -0.32943940499816177
Epoch 4: w1 = -2.5654628900149308, w2 = -0.1775648185723558
Epoch 5: w1 = -2.22920552377513, w2 = -0.09163256127358084
Epoch 6: w1 = -1.9362675156207105, w2 = -0.044944986580951356
Epoch 7: w1 = -1.6781768574274967, w2 = -0.020814229601575286
Epoch 8: w1 = -1.4493498477990567, w2 = -0.009035585595074875
Epoch 9: w1 = -1.245881993508816, w2 = -0.003645905472988451
Epoch 10: w1 = -1.0649030085077547, w2 = -0.0013535098945501255
Epoch 11: w1 = -0.9042022597717997, w2 = -0.00045644443087383875
Epoch 12: w1 = -0.7619964948529878, w2 = -0.0001375629281105624
Epoch 13: w1 = -0.6367784991349715, w2 = -3.62601019486888e-05
Epoch 14: w1 = -0.5272152373016314, w2 = -8.113374556116922e-06
Epoch 15: w1 = -0.4320785049217716, w2 = -1.47473411837664e-06
Epoch 16: w1 = -0.3501985066951055, w2 = -2.0278399084030024e-07
Epoch 17:

#Adam

In [18]:
def df_w(W):
  w1, w2 = W
  dw1 = 0.2 * w1
  dw2 = 4 * w2
  dW = np.array([dw1, dw2])
  return dW


def adam(W, dW, lr, V, S, t, beta1=0.9, beta2=0.999):
  epsilon = 1e-6
  V = beta1 * V + (1 - beta1) * dW
  S = beta2 * S + (1 - beta2) * (dW ** 2)
  v_corr = V / (1 - beta1**t)
  s_corr = S / (1 - beta2**t)
  W = W - lr * (v_corr / (np.sqrt(s_corr) + epsilon))
  return W, V, S


def train_p4(optimizer, lr, epochs):
  W = np.array([-5, -2], dtype=np.float32)
  V = np.array([0, 0], dtype=np.float32)
  S = np.array([0, 0], dtype=np.float32)
  results = [W]
  for i in range(epochs):
    dW = df_w(W)
    W, V, S = optimizer(W, dW, lr, V, S, i+1)
    results.append(W)
    print(f'Epoch {i + 1}: w1 = {W[0]}, w2 = {W[1]}')
  return results

In [19]:
train_adam = train_p4(adam, 0.2, 30)

Epoch 1: w1 = -4.8000001999998, w2 = -1.8000000249999968
Epoch 2: w1 = -4.600254779434054, w2 = -1.6008245063697515
Epoch 3: w1 = -4.400948476628311, w2 = -1.4031726206945152
Epoch 4: w1 = -4.2022776366594705, w2 = -1.2078782223488431
Epoch 5: w1 = -4.004450327821214, w2 = -1.015927446346848
Epoch 6: w1 = -3.807686378997748, w2 = -0.8284730661322335
Epoch 7: w1 = -3.6122173226091405, w2 = -0.6468415893870743
Epoch 8: w1 = -3.4182862261081466, w2 = -0.4725276521059605
Epoch 9: w1 = -3.2261473934546006, w2 = -0.3071693439456018
Epoch 10: w1 = -3.036065916693978, w2 = -0.15249855183024877
Epoch 11: w1 = -2.848317056874701, w2 = -0.010263256257146358
Epoch 12: w1 = -2.663185433233414, w2 = 0.11787552325788148
Epoch 13: w1 = -2.4809640000598776, w2 = 0.23046161354014214
Epoch 14: w1 = -2.301952792136848, w2 = 0.32635870212860313
Epoch 15: w1 = -2.126457422346911, w2 = 0.404841946592144
Epoch 16: w1 = -1.9547873191379472, w2 = 0.4656496111781283
Epoch 17: w1 = -1.7872536971852042, w2 = 0.508

#ADOPT

In [20]:
# mypy: allow-untyped-decorators
# mypy: allow-untyped-defs
from typing import cast, List, Optional, Tuple, Union

import torch
from torch import Tensor

from torch.optim.optimizer import (
    _capturable_doc,
    _default_to_fused_or_foreach,
    _device_dtype_check_for_fused,
    _differentiable_doc,
    _disable_dynamo_if_unsupported,
    _foreach_doc,
    _fused_doc,
    _get_capturable_supported_devices,
    _get_scalar_dtype,
    _get_value,
    _maximize_doc,
    _stack_if_compiling,
    _use_grad_for_differentiable,
    _view_as_real,
    DeviceDict,
    Optimizer,
    ParamsT,
)


__all__ = ["ADOPT", "adopt"]


class ADOPT(Optimizer):
    def __init__(
        self,
        params: ParamsT,
        lr: Union[float, Tensor] = 1e-3,
        betas: Tuple[float, float] = (0.9, 0.9999),
        eps: float = 1e-6,
        weight_decay: float = 0.0,
        decoupled: bool = False,
        *,
        foreach: Optional[bool] = None,
        maximize: bool = False,
        capturable: bool = False,
        differentiable: bool = False,
        fused: Optional[bool] = None,
    ):
        if isinstance(lr, Tensor):
            if foreach and not capturable:
                raise ValueError(
                    "lr as a Tensor is not supported for capturable=False and foreach=True"
                )
            if lr.numel() != 1:
                raise ValueError("Tensor lr must be 1-element")
        if not 0.0 <= lr:
            raise ValueError(f"Invalid learning rate: {lr}")
        if not 0.0 <= eps:
            raise ValueError(f"Invalid epsilon value: {eps}")
        if not 0.0 <= betas[0] < 1.0:
            raise ValueError(f"Invalid beta parameter at index 0: {betas[0]}")
        if not 0.0 <= betas[1] < 1.0:
            raise ValueError(f"Invalid beta parameter at index 1: {betas[1]}")
        if not 0.0 <= weight_decay:
            raise ValueError(f"Invalid weight_decay value: {weight_decay}")

        defaults = dict(
            lr=lr,
            betas=betas,
            eps=eps,
            weight_decay=weight_decay,
            decoupled=decoupled,
            maximize=maximize,
            foreach=foreach,
            capturable=capturable,
            differentiable=differentiable,
            fused=fused,
        )
        super().__init__(params, defaults)

        if fused:
            # TODO: support fused
            raise RuntimeError("`fused` is not currently supported")

            if differentiable:
                raise RuntimeError("`fused` does not support `differentiable`")
            self._step_supports_amp_scaling = True
            # TODO(crcrpar): [low prec params & their higher prec copy]
            # Support AMP with FP16/BF16 model params which would need
            # higher prec copy of params to do update math in higher prec to
            # alleviate the loss of information.
            if foreach:
                raise RuntimeError("`fused` and `foreach` cannot be `True` together.")

    def __setstate__(self, state):
        super().__setstate__(state)
        for group in self.param_groups:
            group.setdefault("maximize", False)
            group.setdefault("foreach", None)
            group.setdefault("capturable", False)
            group.setdefault("differentiable", False)
            fused = group.setdefault("fused", None)
            for p in group["params"]:
                p_state = self.state.get(p, [])
                if len(p_state) != 0 and not torch.is_tensor(p_state["step"]):
                    step_val = float(p_state["step"])
                    p_state["step"] = (
                        torch.tensor(
                            step_val,
                            dtype=_get_scalar_dtype(is_fused=fused),
                            device=p.device,
                        )
                        if group["capturable"] or group["fused"]
                        else torch.tensor(step_val, dtype=_get_scalar_dtype())
                    )

    def _init_group(
        self,
        group,
        params_with_grad,
        grads,
        exp_avgs,
        exp_avg_sqs,
        state_steps,
    ):
        has_complex = False
        for p in group["params"]:
            if p.grad is not None:
                has_complex |= torch.is_complex(p)
                params_with_grad.append(p)
                if p.grad.is_sparse:
                    raise RuntimeError(
                        "ADOPT does not support sparse gradients"
                    )
                grads.append(p.grad)

                state = self.state[p]
                # Lazy state initialization
                if len(state) == 0:
                    if group["fused"]:
                        _device_dtype_check_for_fused(p)
                    # note(crcrpar): [special device hosting for step]
                    # Deliberately host `step` on CPU if both capturable and fused are off.
                    # This is because kernel launches are costly on CUDA and XLA.
                    state["step"] = (
                        torch.zeros(
                            (),
                            dtype=_get_scalar_dtype(is_fused=group["fused"]),
                            device=p.device,
                        )
                        if group["capturable"] or group["fused"]
                        else torch.tensor(0.0, dtype=_get_scalar_dtype())
                    )
                    # Exponential moving average of gradient values
                    state["exp_avg"] = torch.zeros_like(
                        p, memory_format=torch.preserve_format
                    )
                    # Exponential moving average of squared gradient values
                    state["exp_avg_sq"] = torch.zeros_like(
                        p, memory_format=torch.preserve_format
                    )

                exp_avgs.append(state["exp_avg"])
                exp_avg_sqs.append(state["exp_avg_sq"])

                if group["differentiable"] and state["step"].requires_grad:
                    raise RuntimeError(
                        "`requires_grad` is not supported for `step` in differentiable mode"
                    )

                # Foreach without capturable does not support a tensor lr
                if (
                    group["foreach"]
                    and torch.is_tensor(group["lr"])
                    and not group["capturable"]
                ):
                    raise RuntimeError(
                        "lr as a Tensor is not supported for capturable=False and foreach=True"
                    )

                state_steps.append(state["step"])
        return has_complex

    @_use_grad_for_differentiable
    def step(self, closure=None):
        """Perform a single optimization step.

        Args:
            closure (Callable, optional): A closure that reevaluates the model
                and returns the loss.
        """
        self._cuda_graph_capture_health_check()

        loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            params_with_grad: List[Tensor] = []
            grads: List[Tensor] = []
            exp_avgs: List[Tensor] = []
            exp_avg_sqs: List[Tensor] = []
            state_steps: List[Tensor] = []
            beta1, beta2 = group["betas"]

            has_complex = self._init_group(
                group,
                params_with_grad,
                grads,
                exp_avgs,
                exp_avg_sqs,
                state_steps,
            )

            adopt(
                params_with_grad,
                grads,
                exp_avgs,
                exp_avg_sqs,
                state_steps,
                has_complex=has_complex,
                beta1=beta1,
                beta2=beta2,
                lr=group["lr"],
                weight_decay=group["weight_decay"],
                decoupled=group["decoupled"],
                eps=group["eps"],
                maximize=group["maximize"],
                foreach=group["foreach"],
                capturable=group["capturable"],
                differentiable=group["differentiable"],
                fused=group["fused"],
                grad_scale=getattr(self, "grad_scale", None),
                found_inf=getattr(self, "found_inf", None),
            )

        return loss


def _single_tensor_adopt(
    params: List[Tensor],
    grads: List[Tensor],
    exp_avgs: List[Tensor],
    exp_avg_sqs: List[Tensor],
    state_steps: List[Tensor],
    grad_scale: Optional[Tensor],
    found_inf: Optional[Tensor],
    *,
    has_complex: bool,
    beta1: float,
    beta2: float,
    lr: Union[float, Tensor],
    weight_decay: float,
    decoupled: bool,
    eps: float,
    maximize: bool,
    capturable: bool,
    differentiable: bool,
):
    assert grad_scale is None and found_inf is None

    if torch.jit.is_scripting():
        # this assert is due to JIT being dumb and not realizing that the ops below
        # have overloads to handle both float and Tensor lrs, so we just assert it's
        # a float since most people using JIT are using floats
        assert isinstance(lr, float)

    for i, param in enumerate(params):
        grad = grads[i] if not maximize else -grads[i]
        exp_avg = exp_avgs[i]
        exp_avg_sq = exp_avg_sqs[i]
        step_t = state_steps[i]

        # If compiling, the compiler will handle cudagraph checks, see note [torch.compile x capturable]
        if not torch._utils.is_compiling() and capturable:
            capturable_supported_devices = _get_capturable_supported_devices()
            assert (
                param.device.type == step_t.device.type
                and param.device.type in capturable_supported_devices
            ), f"If capturable=True, params and state_steps must be on supported devices: {capturable_supported_devices}."

        # update step
        step_t += 1

        if weight_decay != 0:
            if decoupled:
                param.add_(param, alpha=-lr*weight_decay)
            else:
                grad = grad.add(param, alpha=weight_decay)

        if torch.is_complex(param):
            grad = torch.view_as_real(grad)
            if exp_avg is not None:
                exp_avg = torch.view_as_real(exp_avg)
            if exp_avg_sq is not None:
                exp_avg_sq = torch.view_as_real(exp_avg_sq)
            param = torch.view_as_real(param)

        step = step_t if capturable or differentiable else _get_value(step_t)
        if step == 1:
            exp_avg_sq.addcmul_(grad, grad.conj())
            continue

        denom = torch.clamp(exp_avg_sq.sqrt(), eps)
        if step == 2:
            exp_avg.addcdiv_(grad, denom)
        else:
            exp_avg.mul_(beta1).addcdiv_(grad, denom, value=1 - beta1)

        param.add_(exp_avg, alpha=-lr)
        exp_avg_sq.mul_(beta2).addcmul_(grad, grad.conj(), value=1 - beta2)


def _multi_tensor_adopt(
    params: List[Tensor],
    grads: List[Tensor],
    exp_avgs: List[Tensor],
    exp_avg_sqs: List[Tensor],
    state_steps: List[Tensor],
    grad_scale: Optional[Tensor],
    found_inf: Optional[Tensor],
    *,
    has_complex: bool,
    beta1: float,
    beta2: float,
    lr: Union[float, Tensor],
    weight_decay: float,
    decoupled: bool,
    eps: float,
    maximize: bool,
    capturable: bool,
    differentiable: bool,
):
    if len(params) == 0:
        return

    if isinstance(lr, Tensor) and not capturable:
        raise RuntimeError(
            "lr as a Tensor is not supported for capturable=False and foreach=True"
        )

    # If compiling, the compiler will handle cudagraph checks, see note [torch.compile x capturable]
    if not torch._utils.is_compiling() and capturable:
        capturable_supported_devices = _get_capturable_supported_devices(
            supports_xla=False
        )
        assert all(
            p.device.type == step.device.type
            and p.device.type in capturable_supported_devices
            for p, step in zip(params, state_steps)
        ), f"If capturable=True, params and state_steps must be on supported devices: {capturable_supported_devices}."

    assert grad_scale is None and found_inf is None

    assert not differentiable, "_foreach ops don't support autograd"

    grouped_tensors = Optimizer._group_tensors_by_device_and_dtype(
        [params, grads, exp_avgs, exp_avg_sqs, state_steps]  # type: ignore[list-item]
    )
    for (
        device_params_,
        device_grads_,
        device_exp_avgs_,
        device_exp_avg_sqs_,
        device_state_steps_,
    ), _ in grouped_tensors.values():
        device_params = cast(List[Tensor], device_params_)
        device_grads = cast(List[Tensor], device_grads_)
        device_exp_avgs = cast(List[Tensor], device_exp_avgs_)
        device_exp_avg_sqs = cast(List[Tensor], device_exp_avg_sqs_)
        device_state_steps = cast(List[Tensor], device_state_steps_)

        # Handle complex parameters
        if has_complex:
            _view_as_real(
                device_params, device_grads, device_exp_avgs, device_exp_avg_sqs
            )

        if maximize:
            device_grads = torch._foreach_neg(device_grads)  # type: ignore[assignment]

        # Update steps
        # If steps are on CPU, foreach will fall back to the slow path, which is a for-loop calling t.add(1) over
        # and over. 1 will then be wrapped into a Tensor over and over again, which is slower than if we just
        # wrapped it once now. The alpha is required to assure we go to the right overload.
        if not torch._utils.is_compiling() and device_state_steps[0].is_cpu:
            torch._foreach_add_(
                device_state_steps, torch.tensor(1.0, device="cpu"), alpha=1.0
            )
        else:
            torch._foreach_add_(device_state_steps, 1)

        if weight_decay != 0:
            if decoupled:
                torch._foreach_add_(device_params, device_params, alpha=-lr*weight_decay)
            else:
                # Re-use the intermediate memory (device_grads) already allocated for maximize
                if maximize:
                    torch._foreach_add_(device_grads, device_params, alpha=weight_decay)
                else:
                    device_grads = torch._foreach_add(  # type: ignore[assignment]
                        device_grads, device_params, alpha=weight_decay
                    )

        if device_state_steps[0] == 1:
            torch._foreach_addcmul_(device_exp_avg_sqs, device_grads, device_grads)
            continue

        exp_avg_sq_sqrt = torch._foreach_sqrt(device_exp_avg_sqs)
        exp_avg_sq_sqrt = torch._foreach_maximum(exp_avg_sq_sqrt, eps)

        if device_state_steps[0] == 2:
            torch._foreach_addcdiv_(device_exp_avgs, device_grads, exp_avg_sq_sqrt)
        else:
            torch._foreach_mul_(device_exp_avgs, beta1)
            torch._foreach_addcdiv_(
                device_exp_avgs, device_grads, exp_avg_sq_sqrt, value=1 - beta1
            )

        torch._foreach_add_(device_params, device_exp_avgs, alpha=-lr)
        torch._foreach_mul_(device_exp_avg_sqs, beta2)
        torch._foreach_addcmul_(
            device_exp_avg_sqs, device_grads, device_grads, value=1 - beta2
        )


@_disable_dynamo_if_unsupported(single_tensor_fn=_single_tensor_adopt)
def adopt(
    params: List[Tensor],
    grads: List[Tensor],
    exp_avgs: List[Tensor],
    exp_avg_sqs: List[Tensor],
    state_steps: List[Tensor],
    # kwonly args with defaults are not supported by functions compiled with torchscript issue #70627
    # setting this as kwarg for now as functional API is compiled by torch/distributed/optim
    foreach: Optional[bool] = None,
    capturable: bool = False,
    differentiable: bool = False,
    fused: Optional[bool] = None,
    grad_scale: Optional[Tensor] = None,
    found_inf: Optional[Tensor] = None,
    has_complex: bool = False,
    *,
    beta1: float,
    beta2: float,
    lr: Union[float, Tensor],
    weight_decay: float,
    decoupled: bool,
    eps: float,
    maximize: bool,
):
    r"""Functional API that performs ADOPT algorithm computation.

    """
    # Respect when the user inputs False/True for foreach or fused. We only want to change
    # the default when neither have been user-specified. Note that we default to foreach
    # and pass False to use_fused. This is not a mistake--we want to give the fused impl
    # bake-in time before making it the default, even if it is typically faster.
    if fused is None and foreach is None:
        _, foreach = _default_to_fused_or_foreach(
            params, differentiable, use_fused=False
        )
        # Do not flip on foreach for the unsupported case where lr is a Tensor and capturable=False.
        if foreach and isinstance(lr, Tensor) and not capturable:
            foreach = False
    if fused is None:
        fused = False
    if foreach is None:
        foreach = False

    # this check is slow during compilation, so we skip it
    # if it's strictly needed we can add this check back in dynamo
    if not torch._utils.is_compiling() and not all(
        isinstance(t, torch.Tensor) for t in state_steps
    ):
        raise RuntimeError(
            "API has changed, `state_steps` argument must contain a list of singleton tensors"
        )

    if foreach and torch.jit.is_scripting():
        raise RuntimeError("torch.jit.script not supported with foreach optimizers")
    if fused and torch.jit.is_scripting():
        raise RuntimeError("torch.jit.script not supported with fused optimizers")

    if fused and not torch.jit.is_scripting():
        func = _fused_adopt
    elif foreach and not torch.jit.is_scripting():
        func = _multi_tensor_adopt
    else:
        func = _single_tensor_adopt

    func(
        params,
        grads,
        exp_avgs,
        exp_avg_sqs,
        state_steps,
        has_complex=has_complex,
        beta1=beta1,
        beta2=beta2,
        lr=lr,
        weight_decay=weight_decay,
        decoupled=decoupled,
        eps=eps,
        maximize=maximize,
        capturable=capturable,
        differentiable=differentiable,
        grad_scale=grad_scale,
        found_inf=found_inf,
    )

In [22]:
import torch
from torch import nn

def main():
    # Tạo một mô hình ví dụ
    model = nn.Linear(10, 2)
    # Sử dụng ADOPT làm optimizer với learning rate lr = 0.2
    optimizer = ADOPT(model.parameters(), lr=0.2)

    # Tạo một loss function
    criterion = nn.MSELoss()

    # Dữ liệu giả lập
    inputs = torch.randn(5, 10)
    targets = torch.randn(5, 2)

    # Tiến hành huấn luyện trong 30 vòng lặp
    for epoch in range(30):  # Chỉnh sửa số lượng epoch thành 30
        optimizer.zero_grad()  # Đặt gradient về 0
        outputs = model(inputs)  # Tiến hành truyền qua mô hình
        loss = criterion(outputs, targets)  # Tính toán loss
        loss.backward()  # Lan truyền gradient
        optimizer.step()  # Cập nhật tham số

        print(f"Epoch {epoch+1}, Loss: {loss.item()}")

if __name__ == "__main__":
    main()

Epoch 1, Loss: 0.9868853688240051
Epoch 2, Loss: 0.9868853688240051
Epoch 3, Loss: 0.4614774286746979
Epoch 4, Loss: 0.7535218000411987
Epoch 5, Loss: 1.0737407207489014
Epoch 6, Loss: 1.1585156917572021
Epoch 7, Loss: 1.2143940925598145
Epoch 8, Loss: 1.4108127355575562
Epoch 9, Loss: 1.6695811748504639
Epoch 10, Loss: 1.8374748229980469
Epoch 11, Loss: 1.864385962486267
Epoch 12, Loss: 1.777327299118042
Epoch 13, Loss: 1.593823790550232
Epoch 14, Loss: 1.2887804508209229
Epoch 15, Loss: 0.86248379945755
Epoch 16, Loss: 0.44315099716186523
Epoch 17, Loss: 0.2067890465259552
Epoch 18, Loss: 0.17975711822509766
Epoch 19, Loss: 0.22670495510101318
Epoch 20, Loss: 0.22927939891815186
Epoch 21, Loss: 0.20042526721954346
Epoch 22, Loss: 0.21419289708137512
Epoch 23, Loss: 0.28603464365005493
Epoch 24, Loss: 0.3671693205833435
Epoch 25, Loss: 0.41217952966690063
Epoch 26, Loss: 0.40928274393081665
Epoch 27, Loss: 0.3710615038871765
Epoch 28, Loss: 0.31296306848526
Epoch 29, Loss: 0.239037066