CoolFace
Modelpublic

kernels-community/flash-mla

sourceHugging Facemitupdated 23d agoView on Hugging Face
5likes1kdownloads
test_flash_mla.py70 linesDownload Raw Back to tests
1import torch2import random3import torch.nn.functional as F4 5import flash_mla6 7# TODO: revise to use the same test as the original code8 9 10def test_flash_mla():11    # b = 12812    # s_q = 409613    # mean_sk = 819214    # h_q = 1615    # h_kv = 116    # d = 57617    # dv = 51218 19    b = 1620    s_q = 1621    mean_sk = 1622    h_q = 1623    h_kv = 124    d = 57625    dv = 51226 27 28    causal = True29    varlen = False30 31    print(f"{b=}, {s_q=}, {mean_sk=}, {h_q=}, {h_kv=}, {d=}, {dv=}, {causal=}, {varlen=}")32 33    cache_seqlens = torch.full((b,), mean_sk, dtype=torch.int32)34    if varlen:35        for i in range(b):36            cache_seqlens[i] = max(random.normalvariate(mean_sk, mean_sk / 2), s_q)37    total_seqlens = cache_seqlens.sum().item()38    mean_seqlens = cache_seqlens.float().mean().int().item()39    max_seqlen = cache_seqlens.max().item()40    # TODO: avoid triton from original code41    # max_seqlen_pad = triton.cdiv(max_seqlen, 256) * 25642    print(f"{total_seqlens=}, {mean_seqlens=}, {max_seqlen=}")43    max_seqlen_pad = max_seqlen + 255 & ~255  # round up to multiple of 25644    q = torch.randn(b, s_q, h_q, d)45    block_size = 6446    block_table = torch.arange(b * max_seqlen_pad // block_size, dtype=torch.int32).view(47        b, max_seqlen_pad // block_size48    )49    blocked_k = torch.randn(block_table.numel(), block_size, h_kv, d)50    print(blocked_k.shape)51    for i in range(b):52        blocked_k.view(b, max_seqlen_pad, h_kv, d)[i, cache_seqlens[i].item() :] = float(53            "nan"54        )55    blocked_v = blocked_k[..., :dv]56    print(blocked_k.shape, blocked_v.shape)57 58    cache_seqlens = cache_seqlens.to("cuda")59 60    tile_scheduler_metadata, num_splits = flash_mla.get_mla_metadata(61        seqlens_k=cache_seqlens,62        #63        s_q=s_q * h_q // h_kv,64        h_kv=h_kv,65    )66    print(tile_scheduler_metadata, num_splits)67 68    # TODO: update to expect the correct output69    assert False70