File size: 798 Bytes
d7228c8 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 | """
Unit tests for CustomAdamW and CustomLion optimizers.
"""
import torch
import pytest
from slm.optimizer.adamw import CustomAdamW
from slm.optimizer.lion import CustomLion
def test_custom_adamw_step():
weights = torch.tensor([1.0, 2.0, 3.0], requires_grad=True)
optimizer = CustomAdamW([weights], lr=0.1, weight_decay=0.01)
loss = (weights ** 2).sum()
loss.backward()
optimizer.step()
# Weights should decrease after step
assert weights[0].item() < 1.0
assert weights[1].item() < 2.0
def test_custom_lion_step():
weights = torch.tensor([1.0, 2.0, 3.0], requires_grad=True)
optimizer = CustomLion([weights], lr=0.01, weight_decay=0.01)
loss = (weights ** 2).sum()
loss.backward()
optimizer.step()
assert weights[0].item() < 1.0
|