File size: 3,173 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 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 |
import pytest
import torch
import torch.nn as nn
import torch.nn.functional as F
from fla.modules import FusedCrossEntropyLoss, FusedLinearCrossEntropyLoss
from fla.utils import assert_close, device, device_platform
@pytest.mark.parametrize("B", [2])
@pytest.mark.parametrize("T", [512, 1024])
@pytest.mark.parametrize("D", [1024, 2048])
@pytest.mark.parametrize("V", [32000, 100000])
@pytest.mark.parametrize("reduction", ['mean'])
@pytest.mark.parametrize("dtype", [torch.bfloat16])
@pytest.mark.skipif(
device_platform == 'intel',
reason="Intel Triton Failure",
)
def test_fused_cross_entropy(B: int, T: int, D: int, V: int, reduction: str, dtype: torch.dtype):
torch.manual_seed(42)
logits = torch.randn(B * T, V).to(device).to(dtype=dtype).requires_grad_()
target = torch.randint(0, V, (B, T)).to(device)
target = torch.cat((target[..., 1:], torch.full_like(target[..., :1], -100)), -1)
target = target.flatten()
ref = nn.CrossEntropyLoss(reduction=reduction)(logits, target).to(dtype=dtype)
do = torch.randn_like(ref).to(device).to(dtype=dtype)
ref.backward(do)
ref_d, logits.grad = logits.grad.clone(), None
tri = FusedCrossEntropyLoss(reduction=reduction)(logits, target).to(dtype=dtype)
tri.backward(do)
tri_d, logits.grad = logits.grad.clone(), None
assert_close(" o", ref, tri, ratio=1e-2)
assert_close("dl", ref_d, tri_d, ratio=1e-2)
@pytest.mark.parametrize("B", [2])
@pytest.mark.parametrize("T", [512, 1024])
@pytest.mark.parametrize("D", [1024, 2048])
@pytest.mark.parametrize("V", [32000, 100000])
@pytest.mark.parametrize("scale", [1., 0.5])
@pytest.mark.parametrize("reduction", ['mean'])
@pytest.mark.parametrize("dtype", [torch.bfloat16])
@pytest.mark.skipif(
device_platform == 'intel',
reason="Intel Triton Failure",
)
def test_fused_linear_cross_entropy(B: int, T: int, D: int, V: int, scale: float, reduction: str, dtype: torch.dtype):
torch.manual_seed(42)
x = torch.randn(B * T, D).to(device).to(dtype=dtype).requires_grad_()
target = torch.randint(0, V, (B, T)).to(device)
target = torch.cat((target[..., 1:], torch.full_like(target[..., :1], -100)), -1)
target = target.flatten()
weight = torch.randn(V, D).to(device).to(dtype=dtype).requires_grad_()
bias = torch.randn(V).to(device).to(dtype=dtype).requires_grad_()
logits = F.linear(x, weight, bias)
ref = FusedCrossEntropyLoss(logit_scale=scale, reduction=reduction)(logits, target)
do = torch.randn_like(ref).to(device).to(dtype=dtype)
ref.backward(do)
ref_dx, x.grad = x.grad.clone(), None
ref_dw, weight.grad = weight.grad.clone(), None
ref_db, bias.grad = bias.grad.clone(), None
tri = FusedLinearCrossEntropyLoss(logit_scale=scale, reduction=reduction)(x, target, weight, bias)
tri.backward(do)
tri_dx, x.grad = x.grad.clone(), None
tri_dw, weight.grad = weight.grad.clone(), None
tri_db, bias.grad = bias.grad.clone(), None
assert_close(" o", ref, tri, ratio=1e-2)
assert_close("dx", ref_dx, tri_dx, ratio=1e-2)
assert_close("dw", ref_dw, tri_dw, ratio=1e-2)
assert_close("db", ref_db, tri_db, ratio=1e-2)
|