File size: 1,599 Bytes
c335050 | 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 37 38 39 40 41 42 43 44 45 |
import pytest
import torch
import torch.nn.functional as F
from fla.modules import FusedKLDivLoss
from fla.utils import assert_close, device, device_platform
@pytest.mark.parametrize("B", [2])
@pytest.mark.parametrize("T", [16, 32])
@pytest.mark.parametrize("D", [1024, 2048])
@pytest.mark.parametrize("V", [32000, 100000])
@pytest.mark.parametrize("reduction", ["batchmean"])
@pytest.mark.parametrize("dtype", [torch.float32, torch.float16])
@pytest.mark.skipif(
device_platform == 'intel',
reason="Intel Triton Failure",
)
def test_fused(B: int, T: int, D: int, V: int, reduction: str, dtype: torch.dtype):
torch.manual_seed(42)
x = torch.randn(B * T, D).to(device).to(dtype=dtype).requires_grad_()
x_weight = torch.randn(V, D).to(device).to(dtype=dtype).requires_grad_()
target_x = torch.randn(B * T, D).to(device).to(dtype=dtype)
target_weight = torch.randn(V, D).to(device).to(dtype=dtype)
ref = F.kl_div(
F.linear(x, x_weight).log_softmax(-1),
F.linear(target_x, target_weight).softmax(-1),
reduction=reduction,
).to(dtype)
do = torch.randn_like(ref).to(device)
ref.backward(do)
ref_dx, x.grad = x.grad.clone(), None
ref_dw, x_weight.grad = x_weight.grad.clone(), None
tri = FusedKLDivLoss(reduction)(x, target_x, x_weight, target_weight).to(dtype=dtype)
tri.backward(do)
tri_dx, x.grad = x.grad.clone(), None
tri_dw, x_weight.grad = x_weight.grad.clone(), None
assert_close(" o", ref, tri, 1e-2)
assert_close(" dx", ref_dx, tri_dx, 1e-2)
assert_close(" dw", ref_dw, tri_dw, 1e-2)
|