|
|
| import os |
|
|
| import pytest |
| import torch |
|
|
| from fla.ops.utils import chunk_global_cumsum, chunk_local_cumsum, mean_pooling |
| from fla.ops.utils.index import prepare_lens |
| from fla.ops.utils.pack import pack_sequence, unpack_sequence |
| from fla.utils import assert_close, device |
|
|
|
|
| def reversed_cumsum(x, dim=-1): |
| dtype = x.dtype |
| x = x.float() |
| c = x.cumsum(dim) |
| y = x + c.index_select(dim, x.new_tensor([c.shape[dim]-1], dtype=torch.long)) - c |
| return y.to(dtype) |
|
|
|
|
| @pytest.mark.parametrize( |
| ('B', 'T', 'H', 'D', 'dtype'), |
| [ |
| pytest.param(*test, id="B{}-T{}-H{}-D{}-{}".format(*test)) |
| for test in [ |
| (1, 63, 1, 30, torch.float), |
| (2, 500, 4, 60, torch.float), |
| (2, 1000, 5, 128, torch.float), |
| (3, 1024, 6, 500, torch.float), |
| (4, 2048, 8, 1024, torch.float), |
| ] |
| ], |
| ) |
| def test_global_cumsum( |
| B: int, |
| T: int, |
| H: int, |
| D: int, |
| dtype: torch.dtype, |
| ): |
| torch.manual_seed(42) |
| s = torch.randn(B, T, H, dtype=dtype).to(device) |
| ref = s.float().cumsum(1).to(dtype) |
| tri = chunk_global_cumsum(s) |
| assert_close('global_cumsum', ref, tri, 1e-3) |
|
|
| s = torch.randn(B, T, H, D, dtype=dtype).to(device) |
| ref = s.float().cumsum(1).to(dtype) |
| tri = chunk_global_cumsum(s) |
| assert_close('global_cumsum', ref, tri, 1e-3) |
|
|
|
|
| @pytest.mark.parametrize( |
| ('H', 'D', 'cu_seqlens', 'dtype'), |
| [ |
| pytest.param(*test, id="H{}-D{}-cu_seqlens{}-{}".format(*test)) |
| for test in [ |
| (2, 60, [0, 15], torch.float), |
| (3, 100, [0, 256, 500, 1000], torch.float), |
| (4, 256, [0, 15, 100, 300, 1200, 2000], torch.float), |
| (4, 500, [0, 1, 100, 300, 1200, 2048], torch.float16), |
| (2, 1024, [0, 200, 512, 1200, 2048], torch.float16), |
| ] |
| ], |
| ) |
| @pytest.mark.skipif( |
| os.getenv('SKIP_TEST_CHUNK_VARLEN') == '1', |
| reason='Skipping test_chunk_varlen because SKIP_TEST_CHUNK_VARLEN is set', |
| ) |
| def test_global_cumsum_varlen( |
| H: int, |
| D: int, |
| cu_seqlens: list[int], |
| dtype: torch.dtype, |
| ): |
| torch.manual_seed(42) |
| T = cu_seqlens[-1] |
| cu_seqlens = torch.tensor(cu_seqlens, dtype=torch.int32, device=device) |
|
|
| s = torch.randn(1, T, H, dtype=dtype).to(device) |
| ref = torch.cat([s[:, start:end].float().cumsum(1) for start, end in zip(cu_seqlens[:-1], cu_seqlens[1:], strict=False)], 1).to(dtype) |
| tri = chunk_global_cumsum(s, cu_seqlens=cu_seqlens) |
| assert_close('global_cumsum', ref, tri, 1e-3) |
|
|
| s = torch.randn(1, T, H, D, dtype=dtype).to(device) |
| ref = torch.cat([s[:, start:end].float().cumsum(1) for start, end in zip(cu_seqlens[:-1], cu_seqlens[1:], strict=False)], 1).to(dtype) |
| tri = chunk_global_cumsum(s, cu_seqlens=cu_seqlens) |
| assert_close('global_cumsum', ref, tri, 1e-3) |
|
|
|
|
| @pytest.mark.parametrize( |
| ('B', 'T', 'H', 'D', 'dtype'), |
| [ |
| pytest.param(*test, id="B{}-T{}-H{}-D{}-{}".format(*test)) |
| for test in [ |
| (1, 63, 1, 30, torch.float), |
| (2, 500, 4, 60, torch.float), |
| (2, 1000, 5, 128, torch.float), |
| (3, 1024, 6, 500, torch.float), |
| (4, 2048, 8, 1024, torch.float), |
| ] |
| ], |
| ) |
| def test_global_reversed_cumsum( |
| B: int, |
| T: int, |
| H: int, |
| D: int, |
| dtype: torch.dtype, |
| ): |
| torch.manual_seed(42) |
| s = torch.randn(B, T, H, dtype=dtype).to(device) |
| ref = reversed_cumsum(s, dim=(1)).to(dtype) |
| tri = chunk_global_cumsum(s, reverse=True) |
| assert_close('global_cumsum', ref, tri, 1e-3) |
|
|
| s = torch.randn(B, T, H, D, dtype=dtype).to(device) |
| ref = reversed_cumsum(s, dim=(1)).to(dtype) |
| tri = chunk_global_cumsum(s, reverse=True) |
| assert_close('global_cumsum', ref, tri, 1e-3) |
|
|
|
|
| @pytest.mark.parametrize( |
| ('H', 'D', 'cu_seqlens', 'dtype'), |
| [ |
| pytest.param(*test, id="H{}-D{}-cu_seqlens{}-{}".format(*test)) |
| for test in [ |
| (2, 60, [0, 15], torch.float), |
| (3, 100, [0, 256, 500, 1000], torch.float), |
| (4, 256, [0, 15, 100, 300, 1200, 2000], torch.float), |
| (4, 500, [0, 1, 100, 300, 1200, 2048], torch.float16), |
| (2, 1024, [0, 200, 512, 1200, 2048], torch.float16), |
| ] |
| ], |
| ) |
| @pytest.mark.skipif( |
| os.getenv('SKIP_TEST_CHUNK_VARLEN') == '1', |
| reason='Skipping test_chunk_varlen because SKIP_TEST_CHUNK_VARLEN is set', |
| ) |
| def test_global_reversed_cumsum_varlen( |
| H: int, |
| D: int, |
| cu_seqlens: list[int], |
| dtype: torch.dtype, |
| ): |
| torch.manual_seed(42) |
| T = cu_seqlens[-1] |
| cu_seqlens = torch.tensor(cu_seqlens, dtype=torch.int32, device=device) |
|
|
| s = torch.randn(1, T, H, dtype=dtype).to(device) |
| ref = torch.cat([reversed_cumsum(s[:, start:end], 1) for start, end in zip(cu_seqlens[:-1], cu_seqlens[1:], strict=False)], 1).to(dtype) |
| tri = chunk_global_cumsum(s, reverse=True, cu_seqlens=cu_seqlens) |
| assert_close('global_reversed_cumsum', ref, tri, 1e-3) |
|
|
| s = torch.randn(1, T, H, D, dtype=dtype).to(device) |
| ref = torch.cat([reversed_cumsum(s[:, start:end], 1) for start, end in zip(cu_seqlens[:-1], cu_seqlens[1:], strict=False)], 1).to(dtype) |
| tri = chunk_global_cumsum(s, reverse=True, cu_seqlens=cu_seqlens) |
| assert_close('global_reversed_cumsum', ref, tri, 1e-3) |
|
|
|
|
| @pytest.mark.parametrize( |
| ('B', 'T', 'H', 'C', 'D', 'dtype'), |
| [ |
| pytest.param(*test, id="B{}-T{}-H{}-C{}-D{}-{}".format(*test)) |
| for test in [ |
| (1, 63, 1, 16, 30, torch.float), |
| (2, 500, 4, 32, 60, torch.float), |
| (2, 1000, 5, 64, 128, torch.float), |
| (3, 1024, 6, 64, 500, torch.float), |
| (4, 2048, 8, 128, 1024, torch.float), |
| ] |
| ], |
| ) |
| def test_local_cumsum( |
| B: int, |
| T: int, |
| H: int, |
| C: int, |
| D: int, |
| dtype: torch.dtype, |
| ): |
| torch.manual_seed(42) |
| s = torch.randn(B, T, H, dtype=dtype).to(device) |
| ref = torch.cat([s[:, i:i+C, :].float().cumsum(1) for i in range(0, T, C)], 1) |
| tri = chunk_local_cumsum(s, chunk_size=C) |
| assert_close('local_cumsum', ref, tri, 1e-3) |
|
|
| s = torch.randn(B, T, H, D, dtype=dtype).to(device) |
| ref = torch.cat([s[:, i:i+C, :].float().cumsum(1) for i in range(0, T, C)], 1) |
| tri = chunk_local_cumsum(s, chunk_size=C) |
| assert_close('local_cumsum', ref, tri, 1e-3) |
|
|
|
|
| @pytest.mark.parametrize( |
| ('H', 'C', 'D', 'cu_seqlens', 'dtype'), |
| [ |
| pytest.param(*test, id="H{}-C{}-D{}-cu_seqlens{}-{}".format(*test)) |
| for test in [ |
| (2, 32, 60, [0, 15], torch.float), |
| (3, 64, 100, [0, 256, 500, 1000], torch.float), |
| (4, 64, 256, [0, 15, 100, 300, 1200, 2000], torch.float), |
| (4, 128, 500, [0, 1, 100, 300, 1200, 2048], torch.float16), |
| (2, 128, 1024, [0, 200, 512, 1200, 2048], torch.float16), |
| ] |
| ], |
| ) |
| @pytest.mark.skipif( |
| os.getenv('SKIP_TEST_CHUNK_VARLEN') == '1', |
| reason='Skipping test_chunk_varlen because SKIP_TEST_CHUNK_VARLEN is set', |
| ) |
| def test_local_cumsum_varlen( |
| H: int, |
| C: int, |
| D: int, |
| cu_seqlens: list[int], |
| dtype: torch.dtype, |
| ): |
| torch.manual_seed(42) |
| T = cu_seqlens[-1] |
| cu_seqlens = torch.tensor(cu_seqlens, dtype=torch.int32, device=device) |
|
|
| s = torch.randn(1, T, H, dtype=dtype).to(device) |
| ref = torch.cat([ |
| torch.cat([s[:, i:min(end, i+C), :].float().cumsum(1) for i in range(start, end, C)], 1) |
| for start, end in zip(cu_seqlens[:-1], cu_seqlens[1:], strict=False) |
| ], 1) |
| tri = chunk_local_cumsum(s, chunk_size=C, cu_seqlens=cu_seqlens) |
| assert_close('local_cumsum', ref, tri, 1e-3) |
|
|
| s = torch.randn(1, T, H, D, dtype=dtype).to(device) |
| ref = torch.cat([ |
| torch.cat([s[:, i:min(end, i+C), :].float().cumsum(1) for i in range(start, end, C)], 1) |
| for start, end in zip(cu_seqlens[:-1], cu_seqlens[1:], strict=False) |
| ], 1) |
| tri = chunk_local_cumsum(s, chunk_size=C, cu_seqlens=cu_seqlens) |
| assert_close('local_cumsum', ref, tri, 1e-3) |
|
|
|
|
| @pytest.mark.parametrize( |
| ('B', 'T', 'H', 'C', 'D', 'dtype'), |
| [ |
| pytest.param(*test, id="B{}-T{}-H{}-C{}-D{}-{}".format(*test)) |
| for test in [ |
| (1, 63, 1, 16, 30, torch.float), |
| (2, 500, 4, 32, 60, torch.float), |
| (2, 1000, 5, 64, 128, torch.float), |
| (3, 1024, 6, 64, 500, torch.float), |
| (4, 2048, 8, 128, 1024, torch.float), |
| ] |
| ], |
| ) |
| def test_mean_pooling( |
| B: int, |
| T: int, |
| H: int, |
| C: int, |
| D: int, |
| dtype: torch.dtype, |
| ): |
| torch.manual_seed(42) |
| x = torch.randn(B, T, H, D, dtype=dtype).to(device) |
| x.requires_grad = True |
| ref = torch.cat([x[:, i:i+C, :].float().mean(1, True) for i in range(0, T, C)], 1).to(dtype) |
| do = torch.randn_like(ref) |
| ref.backward(do) |
| ref_dx, x.grad = x.grad.clone(), None |
|
|
| tri = mean_pooling(x, chunk_size=C) |
| tri.backward(do) |
| tri_dx, x.grad = x.grad.clone(), None |
|
|
| assert_close('mean_pooling', ref, tri, 1e-3) |
| assert_close('mean_pooling', ref_dx, tri_dx, 1e-3) |
|
|
|
|
| @pytest.mark.parametrize( |
| ('H', 'C', 'D', 'cu_seqlens', 'dtype'), |
| [ |
| pytest.param(*test, id="H{}-C{}-D{}-cu_seqlens{}-{}".format(*test)) |
| for test in [ |
| (2, 32, 60, [0, 15], torch.float), |
| (3, 64, 100, [0, 256, 500, 1000], torch.float), |
| (4, 64, 256, [0, 15, 100, 300, 1200, 2000], torch.float), |
| (4, 128, 500, [0, 1, 100, 300, 1200, 2048], torch.float16), |
| (2, 128, 1024, [0, 200, 512, 1200, 2048], torch.float16), |
| ] |
| ], |
| ) |
| @pytest.mark.skipif( |
| os.getenv('SKIP_TEST_CHUNK_VARLEN') == '1', |
| reason='Skipping test_chunk_varlen because SKIP_TEST_CHUNK_VARLEN is set', |
| ) |
| def test_mean_pooling_varlen( |
| H: int, |
| C: int, |
| D: int, |
| cu_seqlens: list[int], |
| dtype: torch.dtype, |
| ): |
| torch.manual_seed(42) |
| T = cu_seqlens[-1] |
| cu_seqlens = torch.tensor(cu_seqlens, dtype=torch.int32, device=device) |
|
|
| x = torch.randn(1, T, H, D, dtype=dtype).to(device).requires_grad_(True) |
| ref = torch.cat([ |
| torch.cat([x[:, i:min(end, i+C), :].float().mean(1, True) for i in range(start, end, C)], 1) |
| for start, end in zip(cu_seqlens[:-1], cu_seqlens[1:], strict=False) |
| ], 1).to(dtype) |
| do = torch.randn_like(ref) |
| ref.backward(do) |
| ref_dx, x.grad = x.grad.clone(), None |
|
|
| tri = mean_pooling(x, chunk_size=C, cu_seqlens=cu_seqlens) |
| tri.backward(do) |
| tri_dx, x.grad = x.grad.clone(), None |
|
|
| torch.testing.assert_close(ref, tri.to(ref.dtype), rtol=1.6e-2, atol=3e-5) |
| torch.testing.assert_close(ref_dx, tri_dx.to(ref_dx.dtype), rtol=1.6e-2, atol=3e-5) |
|
|
|
|
| @pytest.mark.parametrize( |
| ('B', 'T', 'H', 'D', 'padding_side', 'dtype'), |
| [ |
| pytest.param(*test, id="B{}-T{}-H{}-D{}-padding_side{}-{}".format(*test)) |
| for test in [ |
| (1, 63, 1, 30, 'left', torch.float), |
| (2, 500, 4, 60, 'right', torch.float), |
| (2, 1000, 5, 128, 'left', torch.float), |
| (3, 1024, 6, 500, 'right', torch.float), |
| (4, 2048, 8, 1024, 'left', torch.float), |
| ] |
| ], |
| ) |
| def test_pack_sequence( |
| B: int, |
| T: int, |
| H: int, |
| D: int, |
| padding_side: str, |
| dtype: torch.dtype, |
| ): |
| torch.manual_seed(42) |
| x = torch.randn(B, T, H, D, dtype=dtype).to(device).requires_grad_(True) |
| cu_seqlens = torch.cat( |
| [torch.tensor([0])]+[torch.randint(0, T, (1,)).clamp(min=1) for _ in range(B)], |
| ).cumsum(-1).to(device) |
| lens = prepare_lens(cu_seqlens) |
|
|
| if padding_side == 'left': |
| ref = torch.cat([x[i, -length:] for i, length in enumerate(lens.tolist())], 0) |
| else: |
| ref = torch.cat([x[i, :length] for i, length in enumerate(lens.tolist())], 0) |
| dy = torch.randn_like(ref) |
| ref.backward(dy) |
| ref_dx, x.grad = x.grad.clone(), None |
|
|
| tri = pack_sequence(x, cu_seqlens, padding_side=padding_side) |
| tri.backward(dy) |
| tri_dx, x.grad = x.grad.clone(), None |
|
|
| assert_close('y', ref, tri, 1e-3) |
| assert_close('dx', ref_dx, tri_dx, 1e-3) |
|
|
|
|
| @pytest.mark.parametrize( |
| ('B', 'T', 'H', 'D', 'padding_side', 'dtype'), |
| [ |
| pytest.param(*test, id="B{}-T{}-H{}-D{}-padding_side{}-{}".format(*test)) |
| for test in [ |
| (1, 63, 1, 30, 'left', torch.float), |
| (2, 500, 4, 60, 'right', torch.float), |
| (2, 1000, 5, 128, 'left', torch.float), |
| (3, 1024, 6, 500, 'right', torch.float), |
| (4, 2048, 8, 1024, 'left', torch.float), |
| ] |
| ], |
| ) |
| def test_unpack_sequence( |
| B: int, |
| T: int, |
| H: int, |
| D: int, |
| padding_side: str, |
| dtype: torch.dtype, |
| ): |
| torch.manual_seed(42) |
| cu_seqlens = torch.cat( |
| [torch.tensor([0])]+[torch.randint(0, T, (1,)).clamp(min=1) for _ in range(B)], |
| ).cumsum(-1).to(device) |
| lens = prepare_lens(cu_seqlens) |
| desired_shape = (B, lens.max().item() + torch.randint(0, 10, (1,)).item(), H, D) |
|
|
| x = torch.randn(cu_seqlens[-1].item(), H, D, dtype=dtype).to(device).requires_grad_(True) |
| ref = torch.zeros(desired_shape, device=device, dtype=dtype) |
| dy = torch.randn_like(ref) |
| for i, (bos, eos) in enumerate(zip(cu_seqlens[:-1].tolist(), cu_seqlens[1:].tolist(), strict=False)): |
| length = eos - bos |
| if padding_side == 'left': |
| ref[i, -length:] = x[bos:eos] |
| else: |
| ref[i, :length] = x[bos:eos] |
| ref.backward(dy) |
| ref_dx, x.grad = x.grad.clone(), None |
|
|
| tri = unpack_sequence(x, cu_seqlens, padding_side=padding_side, desired_shape=desired_shape) |
| tri.backward(dy) |
| tri_dx, x.grad = x.grad.clone(), None |
|
|
| assert_close('y', ref, tri, 1e-3) |
| assert_close('dx', ref_dx, tri_dx, 1e-3) |
|
|