Update flashMLA implementation on MACA (#1)

This commit is contained in:
zhan3916 2026-07-06 14:38:08 +08:00 committed by GitHub
parent d67000b06d
commit d4d889febb
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
617 changed files with 296690 additions and 4379 deletions

View File

@ -1,8 +1,8 @@
MetaX-MACA/FlashMLA是deepseek-ai/FlashMLA算法在MXMACA软件栈及FlashAttention-22.6.3版的实现。MetaX-MACA/FlashMLA以下简称“本软件”适用MIT许可证。本软件亦包含第三方开源组件后者采用的开源许可证将在下文列出。
MetaX-MACA/FlashMLA is the implementation of the deepseek-ai/FlashMLA algorithm on the MXMACA software stack and FlashAttention-2 (version 2.6.3). MetaX-MACA/FlashMLA (“This software”) is licensed under MIT. This software also contains third-party open source components, the open source licenses of which are listed below.
Copyright © 2025 MetaX Integrated Circuits (Shanghai) Co., Ltd.
版权所有©2025 沐曦集成电路(上海)股份有限公司。
Copyright © 2025-2026 MetaX Integrated Circuits (Shanghai) Co., Ltd.
版权所有©2026 沐曦集成电路(上海)股份有限公司。
MIT License
Permission is hereby granted, free of charge, to any person obtaining a copy of this software and associated documentation files (the “Software”), to deal in the Software without restriction, including without limitation the rights to use, copy, modify, merge, publish, distribute, sublicense, and/or sell copies of the Software, and to permit persons to whom the Software is furnished to do so, subject to the following conditions:
@ -18,7 +18,7 @@ This software also contains code from deepseek-ai /FlashMLAhttps://github.com
deepseek-ai /FlashMLA
MIT License
Copyright (c) 2025 DeepSeek
Copyright (c) 2025-2026 DeepSeek
Permission is hereby granted, free of charge, to any person obtaining a copy
of this software and associated documentation files (the "Software"), to deal
@ -101,4 +101,3 @@ SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.

View File

@ -1,520 +0,0 @@
# MLA Triton kernel is from: https://github.com/monellz/vllm/commit/feebaa7c063be6bfb590a876741aeef1c5f58cf8#diff-7b2e1c9032522f7266051b9887246a65753871dfb3625a258fee40109fe6e87a
import argparse
import math
import random
import flashinfer
import torch
import triton
import triton.language as tl
# pip install flashinfer-python
from flash_mla import flash_mla_with_kvcache, get_mla_metadata
def scaled_dot_product_attention(query, key, value, h_q, h_kv, is_causal=False):
query = query.float()
key = key.float()
value = value.float()
key = key.repeat_interleave(h_q // h_kv, dim=0)
value = value.repeat_interleave(h_q // h_kv, dim=0)
attn_weight = query @ key.transpose(-2, -1) / math.sqrt(query.size(-1))
if is_causal:
s_q = query.shape[-2]
s_k = key.shape[-2]
attn_bias = torch.zeros(s_q, s_k, dtype=query.dtype)
temp_mask = torch.ones(s_q, s_k, dtype=torch.bool).tril(diagonal=s_k - s_q)
attn_bias.masked_fill_(temp_mask.logical_not(), float("-inf"))
attn_bias.to(query.dtype)
attn_weight += attn_bias
lse = attn_weight.logsumexp(dim=-1)
attn_weight = torch.softmax(attn_weight, dim=-1, dtype=torch.float32)
return attn_weight @ value, lse
@torch.inference_mode()
def run_torch_mla(q, block_table, blocked_k, max_seqlen_pad, block_size, b, s_q, cache_seqlens, h_q, h_kv, d, dv, causal, dtype):
for i in range(b):
blocked_k.view(b, max_seqlen_pad, h_kv, d)[i, cache_seqlens[i].item():] = float("nan")
blocked_v = blocked_k[..., :dv]
def ref_mla():
out = torch.empty(b, s_q, h_q, dv, dtype=torch.float32)
lse = torch.empty(b, h_q, s_q, dtype=torch.float32)
for i in range(b):
begin = i * max_seqlen_pad
end = begin + cache_seqlens[i]
O, LSE = scaled_dot_product_attention(
q[i].transpose(0, 1),
blocked_k.view(-1, h_kv, d)[begin:end].transpose(0, 1),
blocked_v.view(-1, h_kv, dv)[begin:end].transpose(0, 1),
h_q, h_kv,
is_causal=causal,
)
out[i] = O.transpose(0, 1)
lse[i] = LSE
return out, lse
out_torch, lse_torch = ref_mla()
t = triton.testing.do_bench(ref_mla)
return out_torch, lse_torch, t
@torch.inference_mode()
def run_flash_mla(q, block_table, blocked_k, max_seqlen_pad, block_size, b, s_q, cache_seqlens, h_q, h_kv, d, dv, causal, dtype):
for i in range(b):
blocked_k.view(b, max_seqlen_pad, h_kv, d)[i, cache_seqlens[i].item():] = float("nan")
blocked_v = blocked_k[..., :dv]
tile_scheduler_metadata, num_splits = get_mla_metadata(cache_seqlens, s_q * h_q // h_kv, h_kv)
def flash_mla():
return flash_mla_with_kvcache(
q, blocked_k, block_table, cache_seqlens, dv,
tile_scheduler_metadata, num_splits, causal=causal,
)
out_flash, lse_flash = flash_mla()
t = triton.testing.do_bench(flash_mla)
return out_flash, lse_flash, t
@torch.inference_mode()
def run_flash_infer(q, block_table, blocked_k, max_seqlen_pad, block_size, b, s_q, cache_seqlens, h_q, h_kv, d, dv, causal, dtype):
for i in range(b):
blocked_k.view(b, max_seqlen_pad, h_kv, d)[i, cache_seqlens[i].item():] = float("nan")
assert d > dv, "mla with rope dim should be larger than no rope dim"
q_nope, q_pe = q[..., :dv].contiguous(), q[..., dv:].contiguous()
blocked_k_nope, blocked_k_pe = blocked_k[..., :dv].contiguous(), blocked_k[..., dv:].contiguous()
kv_indptr = [0]
kv_indices = []
for i in range(b):
seq_len = cache_seqlens[i]
assert seq_len > 0
num_blocks = (seq_len + block_size - 1) // block_size
kv_indices.extend(block_table[i, :num_blocks])
kv_indptr.append(kv_indptr[-1] + num_blocks)
for seq_len in cache_seqlens[1:]:
kv_indptr.append((seq_len + block_size - 1) // block_size + kv_indptr[-1])
q_indptr = torch.arange(0, b + 1).int() * s_q
kv_indptr = torch.tensor(kv_indptr, dtype=torch.int32)
kv_indices = torch.tensor(kv_indices, dtype=torch.int32)
mla_wrapper = flashinfer.mla.BatchMLAPagedAttentionWrapper(
torch.empty(128 * 1024 * 1024, dtype=torch.int8),
backend="fa3"
)
mla_wrapper.plan(
q_indptr,
kv_indptr,
kv_indices,
cache_seqlens,
h_q,
dv,
d-dv,
block_size,
causal,
1 / math.sqrt(d),
q.dtype,
blocked_k.dtype,
)
def flash_infer():
output, lse = mla_wrapper.run(q_nope.view(-1, h_q, dv), q_pe.view(-1, h_q, d-dv), blocked_k_nope, blocked_k_pe, return_lse=True)
return output.view(b, -1, h_q, dv), lse.view(b, h_q, 1)
out_flash, lse_flash = flash_infer()
t = triton.testing.do_bench(flash_infer)
return out_flash, lse_flash, t
@triton.jit
def _mla_attn_kernel(
Q_nope,
Q_pe,
Kv_c_cache,
K_pe_cache,
Req_to_tokens,
B_seq_len,
O,
sm_scale,
stride_q_nope_bs,
stride_q_nope_h,
stride_q_pe_bs,
stride_q_pe_h,
stride_kv_c_bs,
stride_k_pe_bs,
stride_req_to_tokens_bs,
stride_o_b,
stride_o_h,
stride_o_s,
BLOCK_H: tl.constexpr,
BLOCK_N: tl.constexpr,
NUM_KV_SPLITS: tl.constexpr,
PAGE_SIZE: tl.constexpr,
HEAD_DIM_CKV: tl.constexpr,
HEAD_DIM_KPE: tl.constexpr,
):
cur_batch = tl.program_id(1)
cur_head_id = tl.program_id(0)
split_kv_id = tl.program_id(2)
cur_batch_seq_len = tl.load(B_seq_len + cur_batch)
offs_d_ckv = tl.arange(0, HEAD_DIM_CKV)
cur_head = cur_head_id * BLOCK_H + tl.arange(0, BLOCK_H)
offs_q_nope = cur_batch * stride_q_nope_bs + cur_head[:, None] * stride_q_nope_h + offs_d_ckv[None, :]
q_nope = tl.load(Q_nope + offs_q_nope)
offs_d_kpe = tl.arange(0, HEAD_DIM_KPE)
offs_q_pe = cur_batch * stride_q_pe_bs + cur_head[:, None] * stride_q_pe_h + offs_d_kpe[None, :]
q_pe = tl.load(Q_pe + offs_q_pe)
e_max = tl.zeros([BLOCK_H], dtype=tl.float32) - float("inf")
e_sum = tl.zeros([BLOCK_H], dtype=tl.float32)
acc = tl.zeros([BLOCK_H, HEAD_DIM_CKV], dtype=tl.float32)
kv_len_per_split = tl.cdiv(cur_batch_seq_len, NUM_KV_SPLITS)
split_kv_start = kv_len_per_split * split_kv_id
split_kv_end = tl.minimum(split_kv_start + kv_len_per_split, cur_batch_seq_len)
for start_n in range(split_kv_start, split_kv_end, BLOCK_N):
offs_n = start_n + tl.arange(0, BLOCK_N)
kv_page_number = tl.load(
Req_to_tokens + stride_req_to_tokens_bs * cur_batch + offs_n // PAGE_SIZE,
mask=offs_n < split_kv_end,
other=0,
)
kv_loc = kv_page_number * PAGE_SIZE + offs_n % PAGE_SIZE
offs_k_c = kv_loc[None, :] * stride_kv_c_bs + offs_d_ckv[:, None]
k_c = tl.load(Kv_c_cache + offs_k_c, mask=offs_n[None, :] < split_kv_end, other=0.0)
qk = tl.dot(q_nope, k_c.to(q_nope.dtype))
offs_k_pe = kv_loc[None, :] * stride_k_pe_bs + offs_d_kpe[:, None]
k_pe = tl.load(K_pe_cache + offs_k_pe, mask=offs_n[None, :] < split_kv_end, other=0.0)
qk += tl.dot(q_pe, k_pe.to(q_pe.dtype))
qk *= sm_scale
qk = tl.where(offs_n[None, :] < split_kv_end, qk, float("-inf"))
v_c = tl.trans(k_c)
n_e_max = tl.maximum(tl.max(qk, 1), e_max)
re_scale = tl.exp(e_max - n_e_max)
p = tl.exp(qk - n_e_max[:, None])
acc *= re_scale[:, None]
acc += tl.dot(p.to(v_c.dtype), v_c)
e_sum = e_sum * re_scale + tl.sum(p, 1)
e_max = n_e_max
offs_o = cur_batch * stride_o_b + cur_head[:, None] * stride_o_h + split_kv_id * stride_o_s + offs_d_ckv[None, :]
tl.store(O + offs_o, acc / e_sum[:, None])
offs_o_1 = cur_batch * stride_o_b + cur_head * stride_o_h + split_kv_id * stride_o_s + HEAD_DIM_CKV
tl.store(O + offs_o_1, e_max + tl.log(e_sum))
def _mla_attn(
q_nope,
q_pe,
kv_c_cache,
k_pe_cache,
attn_logits,
req_to_tokens,
b_seq_len,
num_kv_splits,
sm_scale,
page_size,
):
batch_size, head_num = q_nope.shape[0], q_nope.shape[1]
head_dim_ckv = q_nope.shape[-1]
head_dim_kpe = q_pe.shape[-1]
BLOCK_H = 16
BLOCK_N = 64
grid = (
triton.cdiv(head_num, BLOCK_H),
batch_size,
num_kv_splits,
)
_mla_attn_kernel[grid](
q_nope,
q_pe,
kv_c_cache,
k_pe_cache,
req_to_tokens,
b_seq_len,
attn_logits,
sm_scale,
# stride
q_nope.stride(0),
q_nope.stride(1),
q_pe.stride(0),
q_pe.stride(1),
kv_c_cache.stride(-2),
k_pe_cache.stride(-2),
req_to_tokens.stride(0),
attn_logits.stride(0),
attn_logits.stride(1),
attn_logits.stride(2),
BLOCK_H=BLOCK_H,
BLOCK_N=BLOCK_N,
NUM_KV_SPLITS=num_kv_splits,
PAGE_SIZE=page_size,
HEAD_DIM_CKV=head_dim_ckv,
HEAD_DIM_KPE=head_dim_kpe,
)
@triton.jit
def _mla_softmax_reducev_kernel(
Logits,
B_seq_len,
O,
stride_l_b,
stride_l_h,
stride_l_s,
stride_o_b,
stride_o_h,
NUM_KV_SPLITS: tl.constexpr,
HEAD_DIM_CKV: tl.constexpr,
):
cur_batch = tl.program_id(0)
cur_head = tl.program_id(1)
cur_batch_seq_len = tl.load(B_seq_len + cur_batch)
offs_d_ckv = tl.arange(0, HEAD_DIM_CKV)
e_sum = 0.0
e_max = -float("inf")
acc = tl.zeros([HEAD_DIM_CKV], dtype=tl.float32)
offs_l = cur_batch * stride_l_b + cur_head * stride_l_h + offs_d_ckv
offs_l_1 = cur_batch * stride_l_b + cur_head * stride_l_h + HEAD_DIM_CKV
for split_kv_id in range(0, NUM_KV_SPLITS):
kv_len_per_split = tl.cdiv(cur_batch_seq_len, NUM_KV_SPLITS)
split_kv_start = kv_len_per_split * split_kv_id
split_kv_end = tl.minimum(split_kv_start + kv_len_per_split, cur_batch_seq_len)
if split_kv_end > split_kv_start:
logits = tl.load(Logits + offs_l + split_kv_id * stride_l_s)
logits_1 = tl.load(Logits + offs_l_1 + split_kv_id * stride_l_s)
n_e_max = tl.maximum(logits_1, e_max)
old_scale = tl.exp(e_max - n_e_max)
acc *= old_scale
exp_logic = tl.exp(logits_1 - n_e_max)
acc += exp_logic * logits
e_sum = e_sum * old_scale + exp_logic
e_max = n_e_max
tl.store(
O + cur_batch * stride_o_b + cur_head * stride_o_h + offs_d_ckv,
acc / e_sum,
)
def _mla_softmax_reducev(
logits,
o,
b_seq_len,
num_kv_splits,
):
batch_size, head_num, head_dim_ckv = o.shape[0], o.shape[1], o.shape[2]
grid = (batch_size, head_num)
_mla_softmax_reducev_kernel[grid](
logits,
b_seq_len,
o,
logits.stride(0),
logits.stride(1),
logits.stride(2),
o.stride(0),
o.stride(1),
NUM_KV_SPLITS=num_kv_splits,
HEAD_DIM_CKV=head_dim_ckv,
num_warps=4,
num_stages=2,
)
def mla_decode_triton(
q_nope,
q_pe,
kv_c_cache,
k_pe_cache,
o,
req_to_tokens,
b_seq_len,
attn_logits,
num_kv_splits,
sm_scale,
page_size,
):
assert num_kv_splits == attn_logits.shape[2]
_mla_attn(
q_nope,
q_pe,
kv_c_cache,
k_pe_cache,
attn_logits,
req_to_tokens,
b_seq_len,
num_kv_splits,
sm_scale,
page_size,
)
_mla_softmax_reducev(
attn_logits,
o,
b_seq_len,
num_kv_splits,
)
@torch.inference_mode()
def run_flash_mla_triton(q, block_table, blocked_k, max_seqlen_pad, block_size, b, s_q, cache_seqlens, h_q, h_kv, d, dv, causal, dtype):
for i in range(b):
blocked_k.view(b, max_seqlen_pad, h_kv, d)[i, cache_seqlens[i].item():] = float("nan")
blocked_v = blocked_k[..., :dv]
assert d > dv, "mla with rope dim should be larger than no rope dim"
q_nope, q_pe = q[..., :dv].contiguous(), q[..., dv:].contiguous()
blocked_k_nope, blocked_k_pe = blocked_k[..., :dv].contiguous(), blocked_k[..., dv:].contiguous()
def flash_mla_triton():
num_kv_splits = 32
o = torch.empty([b * s_q, h_q, dv])
attn_logits = torch.empty([b * s_q, h_q, num_kv_splits, dv + 1])
mla_decode_triton(q_nope.view(-1, h_q, dv), q_pe.view(-1, h_q, d-dv), blocked_k_nope.view(-1, dv), blocked_k_pe.view(-1, d-dv), o, block_table, cache_seqlens, attn_logits, num_kv_splits, 1 / math.sqrt(d), block_size)
return o.view([b, s_q, h_q, dv])
out_flash = flash_mla_triton()
t = triton.testing.do_bench(flash_mla_triton)
return out_flash, None, t
FUNC_TABLE = {
"torch": run_torch_mla,
"flash_mla": run_flash_mla,
"flash_infer": run_flash_infer,
"flash_mla_triton": run_flash_mla_triton,
}
def compare_ab(baseline, target, b, s_q, cache_seqlens, h_q, h_kv, d, dv, causal, dtype):
print(f"comparing {baseline} vs {target}: {b=}, {s_q=}, mean_seqlens={cache_seqlens.float().mean()}, {h_q=}, {h_kv=}, {d=}, {dv=}, {causal=}, {dtype=}")
device = torch.device("cuda:0")
torch.set_default_dtype(dtype)
torch.set_default_device(device)
torch.cuda.set_device(device)
torch.manual_seed(0)
random.seed(0)
assert baseline in FUNC_TABLE
assert target in FUNC_TABLE
baseline_func = FUNC_TABLE[baseline]
target_func = FUNC_TABLE[target]
total_seqlens = cache_seqlens.sum().item()
mean_seqlens = cache_seqlens.float().mean().int().item()
max_seqlen = cache_seqlens.max().item()
max_seqlen_pad = triton.cdiv(max_seqlen, 256) * 256
# print(f"{total_seqlens=}, {mean_seqlens=}, {max_seqlen=}")
q = torch.randn(b, s_q, h_q, d)
block_size = 64
block_table = torch.arange(b * max_seqlen_pad // block_size, dtype=torch.int32).view(b, max_seqlen_pad // block_size)
blocked_k = torch.randn(block_table.numel(), block_size, h_kv, d)
out_a, lse_a, perf_a = baseline_func(q, block_table, blocked_k, max_seqlen_pad, block_size, b, s_q, cache_seqlens, h_q, h_kv, d, dv, causal, dtype)
out_b, lse_b, perf_b = target_func(q, block_table, blocked_k, max_seqlen_pad, block_size, b, s_q, cache_seqlens, h_q, h_kv, d, dv, causal, dtype)
torch.testing.assert_close(out_b.float(), out_a.float(), atol=1e-2, rtol=1e-2), "out"
if target not in ["flash_infer", "flash_mla_triton"]:
# flash_infer has a different lse return value
# flash_mla_triton doesn't return lse
torch.testing.assert_close(lse_b.float(), lse_a.float(), atol=1e-2, rtol=1e-2), "lse"
FLOPS = s_q * total_seqlens * h_q * (d + dv) * 2
bytes = (total_seqlens * h_kv * d + b * s_q * h_q * d + b * s_q * h_q * dv) * (torch.finfo(dtype).bits // 8)
print(f"perf {baseline}: {perf_a:.3f} ms, {FLOPS / 10 ** 9 / perf_a:.0f} TFLOPS, {bytes / 10 ** 6 / perf_a:.0f} GB/s")
print(f"perf {target}: {perf_b:.3f} ms, {FLOPS / 10 ** 9 / perf_b:.0f} TFLOPS, {bytes / 10 ** 6 / perf_b:.0f} GB/s")
return bytes / 10 ** 6 / perf_a, bytes / 10 ** 6 / perf_b
def compare_a(target, b, s_q, cache_seqlens, h_q, h_kv, d, dv, causal, dtype):
print(f"{target}: {b=}, {s_q=}, mean_seqlens={cache_seqlens.float().mean()}, {h_q=}, {h_kv=}, {d=}, {dv=}, {causal=}, {dtype=}")
torch.set_default_dtype(dtype)
device = torch.device("cuda:0")
torch.set_default_device(device)
torch.cuda.set_device(device)
torch.manual_seed(0)
random.seed(0)
assert target in FUNC_TABLE
target_func = FUNC_TABLE[target]
total_seqlens = cache_seqlens.sum().item()
mean_seqlens = cache_seqlens.float().mean().int().item()
max_seqlen = cache_seqlens.max().item()
max_seqlen_pad = triton.cdiv(max_seqlen, 256) * 256
# print(f"{total_seqlens=}, {mean_seqlens=}, {max_seqlen=}")
q = torch.randn(b, s_q, h_q, d)
block_size = 64
block_table = torch.arange(b * max_seqlen_pad // block_size, dtype=torch.int32).view(b, max_seqlen_pad // block_size)
blocked_k = torch.randn(block_table.numel(), block_size, h_kv, d)
out_b, lse_b, perf_b = target_func(q, block_table, blocked_k, max_seqlen_pad, block_size, b, s_q, cache_seqlens, h_q, h_kv, d, dv, causal, dtype)
FLOPS = s_q * total_seqlens * h_q * (d + dv) * 2
bytes = (total_seqlens * h_kv * d + b * s_q * h_q * d + b * s_q * h_q * dv) * (torch.finfo(dtype).bits // 8)
print(f"perf {target}: {perf_b:.3f} ms, {FLOPS / 10 ** 9 / perf_b:.0f} TFLOPS, {bytes / 10 ** 6 / perf_b:.0f} GB/s")
return bytes / 10 ** 6 / perf_b
available_targets = [
"torch",
"flash_mla",
"flash_infer",
"flash_mla_triton",
]
shape_configs = [
{"b": batch, "s_q": 1, "cache_seqlens": torch.tensor([seqlen + 2 * i for i in range(batch)], dtype=torch.int32, device="cuda"), "h_q": head, "h_kv": 1, "d": 512+64, "dv": 512, "causal": True, "dtype": torch.bfloat16}
for batch in [128] for seqlen in [1024, 2048, 4096, 8192, 8192*2, 8192*4] for head in [128]
]
def get_args():
parser = argparse.ArgumentParser()
parser.add_argument("--baseline", type=str, default="torch")
parser.add_argument("--target", type=str, default="flash_mla")
parser.add_argument("--all", action="store_true")
parser.add_argument("--one", action="store_true")
parser.add_argument("--compare", action="store_true")
args = parser.parse_args()
return args
if __name__ == "__main__":
args = get_args()
benchmark_type = "all" if args.all else f"{args.baseline}_vs_{args.target}" if args.compare else args.target
with open(f"{benchmark_type}_perf.csv", "w") as fout:
fout.write("name,batch,seqlen,head,bw\n")
for shape in shape_configs:
if args.all:
for target in available_targets:
perf = compare_a(target, shape["b"], shape["s_q"], shape["cache_seqlens"], shape["h_q"], shape["h_kv"], shape["d"], shape["dv"], shape["causal"], shape["dtype"])
fout.write(f'{target},{shape["b"]},{shape["cache_seqlens"].float().mean().cpu().item():.0f},{shape["h_q"]},{perf:.0f}\n')
elif args.compare:
perfa, prefb = compare_ab(args.baseline, args.target, shape["b"], shape["s_q"], shape["cache_seqlens"], shape["h_q"], shape["h_kv"], shape["d"], shape["dv"], shape["causal"], shape["dtype"])
fout.write(f'{args.baseline},{shape["b"]},{shape["cache_seqlens"].float().mean().cpu().item():.0f},{shape["h_q"]},{perfa:.0f}\n')
fout.write(f'{args.target},{shape["b"]},{shape["cache_seqlens"].float().mean().cpu().item():.0f},{shape["h_q"]},{prefb:.0f}\n')
elif args.one:
perf = compare_a(args.target, shape["b"], shape["s_q"], shape["cache_seqlens"], shape["h_q"], shape["h_kv"], shape["d"], shape["dv"], shape["causal"], shape["dtype"])
fout.write(f'{args.target},{shape["b"]},{shape["cache_seqlens"].float().mean().cpu().item():.0f},{shape["h_q"]},{perf:.0f}\n')

View File

@ -1,29 +0,0 @@
import argparse
import matplotlib.pyplot as plt
import pandas as pd
def parse_args():
parser = argparse.ArgumentParser(description='Visualize benchmark results')
parser.add_argument('--file', type=str, default='all_perf.csv',
help='Path to the CSV file with benchmark results (default: all_perf.csv)')
return parser.parse_args()
args = parse_args()
file_path = args.file
df = pd.read_csv(file_path)
names = df['name'].unique()
for name in names:
subset = df[df['name'] == name]
plt.plot(subset['seqlen'], subset['bw'], label=name)
plt.title('bandwidth')
plt.xlabel('seqlen')
plt.ylabel('bw (GB/s)')
plt.legend()
plt.savefig(f'{file_path.split(".")[0].split("/")[-1]}_bandwidth_vs_seqlen.png')

View File

@ -10,122 +10,90 @@
#include "flash_mla.h"
#include "static_switch.h"
#include "run_mha.h"
#include "run_mla.h"
#include "host_utils.h"
#define CHECK_DEVICE(x) TORCH_CHECK(x.is_cuda(), #x " must be on CUDA")
#define CHECK_SHAPE(x, ...) TORCH_CHECK(x.sizes() == torch::IntArrayRef({__VA_ARGS__}), #x " must have shape (" #__VA_ARGS__ ")")
#define CHECK_CONTIGUOUS(x) TORCH_CHECK(x.is_contiguous(), #x " must be contiguous")
// Find the number of splits that maximizes the occupancy. For example, if we have
// batch * n_heads = 48 and we have 108 SMs, having 2 splits (efficiency = 0.89) is
// better than having 3 splits (efficiency = 0.67). However, we also don't want too many
// splits as that would incur more HBM reads/writes.
// So we find the best efficiency, then find the smallest number of splits that gets 85%
// of the best efficiency.
int num_splits_heuristic(int batch_nheads_mblocks, int num_SMs, int num_n_blocks, int max_splits, float balance_weight) {
// If we have enough to almost fill the SMs, then just use 1 split
// if (batch_nheads_mblocks >= 0.9f * num_SMs) { return 1; }
max_splits = std::min({max_splits, num_SMs, num_n_blocks});
float max_efficiency = 0.f;
std::vector<float> efficiency;
efficiency.reserve(max_splits);
auto ceildiv = [](int a, int b) { return (a + b - 1) / b; };
// Some splits are not eligible. For example, if we have 64 blocks and choose 11 splits,
// we'll have 6 * 10 + 4 blocks. If we choose 12 splits, we'll have 6 * 11 + (-2) blocks
// (i.e. it's 11 splits anyway).
// So we check if the number of blocks per split is the same as the previous num_splits.
auto is_split_eligible = [&ceildiv, &num_n_blocks](int num_splits) {
return num_splits == 1 || ceildiv(num_n_blocks, num_splits) != ceildiv(num_n_blocks, num_splits - 1);
};
for (int num_splits = 1; num_splits <= max_splits; num_splits++) {
if (!is_split_eligible(num_splits)) {
efficiency.push_back(0.f);
inline int int64_stride_to_int(int64_t orig_stride) {
if (orig_stride > std::numeric_limits<int>::max()) {
TORCH_CHECK(false, "[Sparse TopK Attention] Stride exceeds int32 limit: ", orig_stride);
}
return static_cast<int>(orig_stride);
}
// Note: should match the kernel dispatch tile size
inline std::pair<int, int> get_tile_size(int arch, int seqlen_q, bool is_sparse_attn) {
int block_m = 0;
int block_n = 0;
if (is_sparse_attn) {
// xcore1500 use the same kernel with xcore1000 in sparse decode now
block_m = 64, block_n = 16;
} else {
if (arch >= 1500) {
// only support blockM=64 in xcore1500 dense decode now
block_m = 64, block_n = 32;
} else {
float n_waves = float(batch_nheads_mblocks * num_splits) / num_SMs;
float eff = n_waves / ceil(n_waves);
// printf("num_splits = %d, eff = %f\n", num_splits, eff);
if (eff > max_efficiency) { max_efficiency = eff; }
efficiency.push_back(eff);
if (seqlen_q >= 64) {
block_m = 64, block_n = 16;
} else if (seqlen_q >= 32) {
block_m = 32, block_n = 16;
} else {
block_m = 16, block_n = 16;
}
}
}
for (int num_splits = 1; num_splits <= max_splits; num_splits++) {
if (!is_split_eligible(num_splits)) { continue; }
if (efficiency[num_splits - 1] >= balance_weight * max_efficiency) {
// printf("num_splits chosen = %d\n", num_splits);
return num_splits;
}
}
return 1;
return {block_m, block_n};
}
void compute_params_numsplits(mcFlashAttn::Flash_fwd_mla_params &params, const int num_splits){
auto num_heads = params.h;
auto batch_size = params.b;
auto max_seqlen_k = params.seqlen_k;
auto max_seqlen_q = params.seqlen_q;
auto dprops = at::cuda::getCurrentDeviceProperties();
struct DecodingAttnImplMeta {
int num_sm_parts;
int fixed_overhead_num_blocks;
int k_block_size;
};
const int block_n = 16;
const int num_n_blocks = (max_seqlen_k + block_n - 1) / block_n;
const int block_m = max_seqlen_q >= 64 ? 64 : 32;
const int num_m_blocks = (max_seqlen_q + block_m - 1) / block_m;
params.num_splits = num_splits;
if (num_splits < 1) {
const int AP_nums = dprops->multiProcessorCount;
int block_nums_per_AP = 1;
// TODO: fine tune balance_weight later
float balance_weight = batch_size == 128 ? 0.95 : 0.9;
params.num_splits = num_splits_heuristic(batch_size * num_heads * num_m_blocks, AP_nums * block_nums_per_AP,
num_n_blocks, 128, balance_weight);
}
}
std::vector<at::Tensor>
get_mla_metadata(
at::Tensor &seqlens_k,
const int num_heads_per_head_k,
const int num_heads_k
DecodingAttnImplMeta get_attn_impl_meta(
int arch,
int sm_count,
int num_q_tokens_per_head_k,
int h_k,
int block_m,
int block_n,
std::optional<int> h_q_,
bool is_fp8_kvcache,
bool is_sparse_attn
) {
// This should match the logic in the MLA kernel.
static constexpr int block_size_m = 64;
static constexpr int block_size_n = 64;
static constexpr int fixed_overhead_num_blocks = 5;
CHECK_DEVICE(seqlens_k);
TORCH_CHECK(seqlens_k.is_contiguous());
TORCH_CHECK(seqlens_k.dtype() == torch::kInt32);
int batch_size = seqlens_k.size(0);
int *seqlens_k_ptr = seqlens_k.data_ptr<int>();
auto options = seqlens_k.options();
auto dprops = at::cuda::getCurrentDeviceProperties();
int sm_count = dprops->multiProcessorCount;
int num_sm_parts = sm_count / num_heads_k / mctlass::ceil_div(num_heads_per_head_k, block_size_m);
auto tile_scheduler_metadata = torch::empty({num_sm_parts, TileSchedulerMetaDataSize}, options);
auto num_splits = torch::empty({batch_size + 1}, options);
int *tile_scheduler_metadata_ptr = tile_scheduler_metadata.data_ptr<int>();
int *num_splits_ptr = num_splits.data_ptr<int>();
at::cuda::CUDAGuard device_guard{(char)seqlens_k.get_device()};
auto stream = at::cuda::getCurrentCUDAStream().stream();
Mla_metadata_params params = {};
params.seqlens_k_ptr = seqlens_k_ptr;
params.tile_scheduler_metadata_ptr = tile_scheduler_metadata_ptr;
params.num_splits_ptr = num_splits_ptr;
params.batch_size = batch_size;
params.block_size_n = block_size_n;
params.fixed_overhead_num_blocks = fixed_overhead_num_blocks;
params.num_sm_parts = num_sm_parts;
// get_mla_metadata_func(params, stream);
return {tile_scheduler_metadata, num_splits};
if (is_sparse_attn) {
if (is_fp8_kvcache) {
TORCH_CHECK(false, "Sparse fp8 MLA is not supported.");
} else {
// Sparse BF16 MLA
TORCH_CHECK(h_q_.has_value());
int h_q = h_q_.value();
TORCH_CHECK(h_q % h_k == 0, "h_k must be divisible by h_q.");
int s_q = num_q_tokens_per_head_k * h_k / h_q;
// BF16/FP16 + Sparse MLA
return {
std::max((sm_count/2) / h_k / (mctlass::ceil_div(h_q/h_k, 2*64) * s_q), 1),
5,
block_n // block_n
};
}
} else {
TORCH_CHECK(!is_fp8_kvcache, "FP8 KV Cache is not supported.");
// Dense BF16/FP8 MLA
return {
std::max(sm_count / h_k / mctlass::ceil_div(num_q_tokens_per_head_k, block_m), 1),
5,
block_n,
};
}
}
std::vector<at::Tensor>
mha_fwd_kvcache_mla(
fwd_kvcache_mla(
at::Tensor &q, // batch_size x seqlen_q x num_heads x head_size
const at::Tensor &kcache, // num_blocks x page_block_size x num_heads_k x head_size
c10::optional<const at::Tensor> &vcache_, // num_blocks x page_block_size x num_heads_k x head_size_v
@ -135,17 +103,23 @@ mha_fwd_kvcache_mla(
const float softmax_scale,
bool is_causal,
const at::Tensor &tile_scheduler_metadata, // num_sm_parts x TileSchedulerMetaDataSize
const at::Tensor &num_splits // batch_size + 1
const at::Tensor &num_splits, // batch_size + 1
bool is_fp8_kvcache, // fp8 kvcache=False
c10::optional<const at::Tensor> &indices, // None, or batch_size x seqlen_q x topk
c10::optional<const at::Tensor> &indices_all_valid_per_q, // batch_size x seqlen_q x 1, per-query flag indicating whether all top-k indices for each query token are valid.
int const cp_world_size, // context parallelism (cp) world size
int const cp_rank, // cp rank
c10::optional<const at::Tensor> &cp_tot_seqused_k_ // b. total seqused_k in cp world
) {
auto dprops = at::cuda::getCurrentDeviceProperties();
bool is_sm90 = dprops->major == 9 && dprops->minor == 0;
// TORCH_CHECK(is_sm90);
auto dprops = flash::mcGetCurrentDeviceProperties();
int arch = dprops.major * 100 + dprops.minor;
at::Tensor vcache = vcache_.has_value() ? vcache_.value() : kcache;
auto q_dtype = q.dtype();
TORCH_CHECK(q_dtype == torch::kBFloat16 || q_dtype == torch::kFloat16);
TORCH_CHECK(kcache.dtype() == q_dtype, "query and key must have the same dtype");
TORCH_CHECK(!is_fp8_kvcache, "flash mla with kvcache api not support fp8 now");
CHECK_DEVICE(q); CHECK_DEVICE(kcache); CHECK_DEVICE(vcache);
@ -157,6 +131,14 @@ mha_fwd_kvcache_mla(
TORCH_CHECK(block_table.dtype() == torch::kInt32, "block_table must have dtype torch.int32");
TORCH_CHECK(block_table.stride(-1) == 1, "block_table must have contiguous last dimension");
bool is_sparse_attn = indices.has_value();
int topk = is_sparse_attn ? indices->size(-1) : -1;
TORCH_CHECK(!is_sparse_attn || indices->dtype() == torch::kInt32, "indices must have dtype int32");
TORCH_CHECK(!is_sparse_attn || indices->stride(-1) == 1, "indices must have contiguous last dimension");
TORCH_CHECK(!is_sparse_attn || indices_all_valid_per_q->dtype() == torch::kBool, "indices_all_valid_per_q must have dtype bool");
TORCH_CHECK(!is_sparse_attn || indices_all_valid_per_q->stride(-1) == 1, "indices_all_valid_per_q must have contiguous last dimension");
const auto sizes = q.sizes();
const int batch_size = sizes[0];
const int seqlen_q_ori = sizes[1];
@ -176,6 +158,9 @@ mha_fwd_kvcache_mla(
const int ngroups = num_heads_ori / num_heads_k;
const int seqlen_q = seqlen_q_ori * ngroups;
const int num_heads = num_heads_k;
if (is_sparse_attn){
TORCH_CHECK(num_heads_ori >= 64 || seqlen_q_ori == 1, "sparse decoding head q must greter than 64 when seqlen q > 1");
}
q = q.view({batch_size, seqlen_q_ori, num_heads_k, ngroups, head_size}).transpose(2, 3)
.reshape({batch_size, seqlen_q, num_heads, head_size});
@ -191,6 +176,15 @@ mha_fwd_kvcache_mla(
CHECK_CONTIGUOUS(seqlens_k);
CHECK_SHAPE(seqlens_k, batch_size);
if (cp_tot_seqused_k_.has_value()) {
auto cp_tot_seqused_k = cp_tot_seqused_k_.value();
TORCH_CHECK(cp_tot_seqused_k.dtype() == torch::kInt32, "seqused_k must have dtype int32");
CHECK_DEVICE(cp_tot_seqused_k); CHECK_CONTIGUOUS(cp_tot_seqused_k);
CHECK_SHAPE(cp_tot_seqused_k, batch_size);
}
at::cuda::CUDAGuard device_guard{(char)q.get_device()};
auto opts = q.options();
@ -202,13 +196,15 @@ mha_fwd_kvcache_mla(
// Set the sizes.
params.b = batch_size;
params.seqlen_q = seqlen_q;
params.seqlen_k = seqlens_k.max().cpu().item<int>();
// params.seqlen_k = seqlens_k.max().cpu().item<int>();
params.cu_seqlens_k = seqlens_k.data_ptr<int>();
params.is_seqlens_k_cumulative = false; // seqlens_k always has value
params.h = num_heads;
params.h_h_k_ratio = num_heads / num_heads_k;
params.ngroups = ngroups;
params.is_causal = is_causal;
params.is_sparse_attn = is_sparse_attn;
params.topk = topk;
params.d = head_size;
params.d_v = head_size_v;
params.scale_softmax = softmax_scale;
@ -233,34 +229,49 @@ mha_fwd_kvcache_mla(
params.v_head_stride = vcache.stride(-2);
params.o_head_stride = out.stride(-2);
// indices ptr
params.indices_ptr = is_sparse_attn ? indices->data_ptr<int32_t>() : nullptr;
params.indices_batch_stride = is_sparse_attn ? indices->stride(0) : 0;
params.indices_row_stride = is_sparse_attn ? indices->stride(1) : 0;
params.indices_all_valid_per_q_ptr = is_sparse_attn ? indices_all_valid_per_q->data_ptr<bool>() : nullptr;
params.indices_all_valid_per_q_batch_stride = is_sparse_attn ? indices_all_valid_per_q->stride(0) : 0;
params.indices_all_valid_per_q_row_stride = is_sparse_attn ? indices_all_valid_per_q->stride(1) : 0;
params.block_table = block_table.data_ptr<int>();
params.block_table_batch_stride = block_table.stride(0);
params.page_block_size = page_block_size;
params.arch = arch;
params.cp_world_size = cp_world_size;
params.cp_rank = cp_rank;
params.cp_tot_seqused_k = cp_tot_seqused_k_.has_value() ? cp_tot_seqused_k_->data_ptr<int>() : nullptr;
TORCH_CHECK(cp_world_size > 0, "cp_world_size must be positive, required by downstream unified code path. Use 1 if CP is not enabled.");
TORCH_CHECK(cp_world_size != 1 || cp_rank == 0, "When context parallelism is disabled, cp_rank must be zero");
TORCH_CHECK(cp_world_size == 1 || cp_tot_seqused_k_.has_value(), "cp_tot_seqused_k_ must be provided when context parallelism is enabled.");
TORCH_CHECK(num_splits.dtype() == torch::kInt32, "num_splits must have dtype int32");
// printf("num_splits%d",num_splits);
CHECK_DEVICE(num_splits);
CHECK_CONTIGUOUS(num_splits);
TORCH_CHECK(tile_scheduler_metadata.dtype() == torch::kInt32, "tile_scheduler_metadata must have dtype int32");
TORCH_CHECK(tile_scheduler_metadata.size(1) == TileSchedulerMetaDataSize);
CHECK_DEVICE(tile_scheduler_metadata);
CHECK_CONTIGUOUS(tile_scheduler_metadata);
// params.tile_scheduler_metadata_ptr = tile_scheduler_metadata.data_ptr<int>();
// params.num_sm_parts = tile_scheduler_metadata.size(0);
TORCH_CHECK(num_splits.dtype() == torch::kInt32, "num_splits must have dtype int32");
CHECK_DEVICE(num_splits);
CHECK_CONTIGUOUS(num_splits);
// params.num_splits_ptr = num_splits.data_ptr<int>();
const int max_num_splits = 128;
// TODO: enable get_mla_mate_data for load balance
compute_params_numsplits(params, 0);
TORCH_CHECK(params.num_splits <= max_num_splits, "num_splits must less than or equal to 128");
at::Tensor softmax_lse_accum = torch::empty({params.num_splits, batch_size, num_heads, seqlen_q}, opts.dtype(torch::kFloat32));
at::Tensor out_accum = torch::empty({params.num_splits, batch_size, num_heads, seqlen_q, head_size_v}, opts.dtype(torch::kFloat32));
params.tile_scheduler_metadata_ptr = tile_scheduler_metadata.data_ptr<int>();
params.num_sm_parts = tile_scheduler_metadata.size(0);
params.num_splits_ptr = num_splits.data_ptr<int>();
at::Tensor softmax_lse_accum = torch::empty({batch_size + params.num_sm_parts, num_heads, seqlen_q}, opts.dtype(at::kFloat));
at::Tensor out_accum = torch::empty({batch_size + params.num_sm_parts, num_heads, seqlen_q, head_size_v}, opts.dtype(at::kFloat));
params.softmax_lseaccum_ptr = softmax_lse_accum.data_ptr();
params.oaccum_ptr = out_accum.data_ptr();
auto stream = at::cuda::getCurrentCUDAStream().stream();
TORCH_CHECK(head_size == 576);
params.is_bf16 = q_dtype == torch::kBFloat16;
run_mha_fwd(params,stream, /*force_split_kernel*/true);
run_mla_fwd(params, stream);
out = out.view({batch_size, seqlen_q_ori, ngroups, num_heads_k, head_size_v}).transpose(2, 3)
.reshape({batch_size, seqlen_q_ori, num_heads_ori, head_size_v});
softmax_lse = softmax_lse.view({batch_size, num_heads_k, seqlen_q_ori, ngroups}).transpose(2, 3)
@ -268,9 +279,148 @@ mha_fwd_kvcache_mla(
return {out, softmax_lse};
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.doc() = "FlashAttention";
//FlashMLA
m.def("get_mla_metadata", &get_mla_metadata);
m.def("fwd_kvcache_mla", &mha_fwd_kvcache_mla);
std::vector<at::Tensor> sparse_prefill_fwd(
const at::Tensor &q,
const at::Tensor &kv,
const at::Tensor &indices,
float sm_scale,
int d_v,
const at::Tensor &indices_all_valid_per_q
) {
auto dprops = flash::mcGetCurrentDeviceProperties();
int arch = dprops.major * 100 + dprops.minor;
CHECK_DEVICE(q);
CHECK_DEVICE(kv);
CHECK_DEVICE(indices);
CHECK_DEVICE(indices_all_valid_per_q);
TORCH_CHECK(q.dtype() == torch::kBFloat16);
TORCH_CHECK(kv.dtype() == torch::kBFloat16);
TORCH_CHECK(indices.dtype() == torch::kInt32);
TORCH_CHECK(indices_all_valid_per_q.dtype() == torch::kBool);
int s_q = q.size(0);
int s_kv = kv.size(0);
int h_q = q.size(1);
int h_kv = kv.size(1);
int d_qk = q.size(2);
int topk = indices.size(2);
TORCH_CHECK(h_q % 64 == 0 && h_q >= 64);
CHECK_SHAPE(q, s_q, h_q, d_qk);
CHECK_SHAPE(kv, s_kv, h_kv, d_qk);
CHECK_SHAPE(indices, s_q, h_kv, topk);
CHECK_SHAPE(indices_all_valid_per_q, s_q, 1);
TORCH_CHECK(q.stride(-1) == 1);
TORCH_CHECK(kv.stride(-1) == 1);
TORCH_CHECK(indices.stride(-1) == 1);
at::cuda::CUDAGuard device_guard{(char)q.get_device()};
auto opts = q.options();
at::Tensor out = torch::empty({s_q, h_q, d_v}, opts);
CHECK_CONTIGUOUS(out);
at::Tensor buf_attn_score, max_logits, lse, p_sum;
max_logits = torch::empty({s_q, h_q}, opts.dtype(torch::kFloat));
lse = torch::empty({s_q, h_q}, opts.dtype(torch::kFloat));
CHECK_CONTIGUOUS(max_logits);
CHECK_CONTIGUOUS(lse);
SparsePrefillParams params = {
s_q, s_kv, h_q, h_kv, d_qk, d_v, topk,
sm_scale, sm_scale * 1.44269504f,
arch,
(mctlass::bfloat16_t*)q.data_ptr(),
(mctlass::bfloat16_t*)kv.data_ptr(),
(int*)indices.data_ptr(),
(bool*)indices_all_valid_per_q.data_ptr(),
int64_stride_to_int(q.stride(0)), int64_stride_to_int(q.stride(1)),
int64_stride_to_int(kv.stride(0)), int64_stride_to_int(kv.stride(1)),
int64_stride_to_int(indices.stride(0)), int64_stride_to_int(indices.stride(1)),
int64_stride_to_int(out.stride(0)),int64_stride_to_int(out.stride(1)),
(mctlass::bfloat16_t*)out.data_ptr(),
(float*)max_logits.data_ptr(),
(float*)lse.data_ptr(),
at::cuda::getCurrentCUDAStream().stream()
};
run_mla_fwd(params);
return {out, max_logits, lse};
}
std::vector<at::Tensor>
get_mla_decoding_metadata(
at::Tensor &seqlens_k,
const int num_q_tokens_per_head_k,
const int h_k,
const std::optional<int> h_q,
const bool is_fp8_kvcache,
const std::optional<int> topk
) {
auto dprops = flash::mcGetCurrentDeviceProperties();
int arch = dprops.major * 100 + dprops.minor;
// This should match the logic in the MLA kernel.
const int seqlen_q = num_q_tokens_per_head_k * h_k;
bool is_sparse_attn = topk.has_value();
const auto [block_size_m, block_size_n] = get_tile_size(arch, seqlen_q, is_sparse_attn);
CHECK_DEVICE(seqlens_k);
TORCH_CHECK(seqlens_k.is_contiguous());
TORCH_CHECK(seqlens_k.dtype() == torch::kInt32);
if (is_sparse_attn)
TORCH_CHECK(h_q.has_value(), "num_heads_q must be provided when topk is provided");
CHECK_DEVICE(seqlens_k);
TORCH_CHECK(seqlens_k.is_contiguous());
TORCH_CHECK(seqlens_k.dtype() == torch::kInt32);
int batch_size = seqlens_k.size(0);
int *seqlens_k_ptr = seqlens_k.data_ptr<int>();
auto options = seqlens_k.options();
int sm_count = dprops.multiProcessorCount;
const char* val = std::getenv("FMLA_SM");
if(val != nullptr){
sm_count = std::stoi(val);
}
DecodingAttnImplMeta attn_impl_meta = get_attn_impl_meta(arch, sm_count, num_q_tokens_per_head_k, h_k, block_size_m, block_size_n, h_q, is_fp8_kvcache, is_sparse_attn);
if(std::getenv("FMLA_LOG")){
printf("block_size_m %d, num_q_tokens_per_head_k %d, h_k %d, seqlen_q %d, sm_count %d, sm_parts %d \n",
block_size_m, num_q_tokens_per_head_k, h_k, seqlen_q, sm_count, attn_impl_meta.num_sm_parts);
}
auto tile_scheduler_metadata = torch::empty({attn_impl_meta.num_sm_parts, TileSchedulerMetaDataSize}, options);
auto num_splits = torch::empty({batch_size + 1}, options);
int *tile_scheduler_metadata_ptr = tile_scheduler_metadata.data_ptr<int>();
int *num_splits_ptr = num_splits.data_ptr<int>();
at::cuda::CUDAGuard device_guard{(char)seqlens_k.get_device()};
auto stream = at::cuda::getCurrentCUDAStream().stream();
GetDecodingMetadataParams params = {};
params.seqlens_k_ptr = seqlens_k_ptr;
params.tile_scheduler_metadata_ptr = tile_scheduler_metadata_ptr;
params.num_splits_ptr = num_splits_ptr;
params.batch_size = batch_size;
params.block_size_n = attn_impl_meta.k_block_size;
params.fixed_overhead_num_blocks = attn_impl_meta.fixed_overhead_num_blocks;
params.num_sm_parts = attn_impl_meta.num_sm_parts;
params.topk = is_sparse_attn ? topk.value() : -1;
run_get_mla_metadata_kernel(params, stream);
return {tile_scheduler_metadata, num_splits};
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.doc() = "FlashMLA";
m.def("get_mla_metadata", &get_mla_decoding_metadata);
m.def("fwd_kvcache_mla", &fwd_kvcache_mla);
m.def("sparse_prefill_fwd", &sparse_prefill_fwd);
}

View File

@ -1,18 +0,0 @@
// Adapted from Dao-AILab/flash-attention (https://github.com/Dao-AILab/flash-attention/tree/v2.6.3)
#include "flash_mla.h"
#include "flash_run_fwd_template_impl.h"
#include <mctlass/numeric_types.h>
template void run_flash_splitkv_fwd_template<
576,
16,
16,
4,
true,
true,
cutlass::bfloat16_t,
false,
512,
2
>(Flash_fwd_mla_params &params, cudaStream_t stream);

View File

@ -1,18 +0,0 @@
// Adapted from Dao-AILab/flash-attention (https://github.com/Dao-AILab/flash-attention/tree/v2.6.3)
#include "flash_mla.h"
#include "flash_run_fwd_template_impl.h"
#include <mctlass/numeric_types.h>
template void run_flash_splitkv_fwd_template<
576,
32,
16,
4,
true,
true,
cutlass::bfloat16_t,
false,
512,
2
>(Flash_fwd_mla_params &params, cudaStream_t stream);

View File

@ -1,18 +0,0 @@
// Adapted from Dao-AILab/flash-attention (https://github.com/Dao-AILab/flash-attention/tree/v2.6.3)
#include "flash_mla.h"
#include "flash_run_fwd_template_impl.h"
#include <mctlass/numeric_types.h>
template void run_flash_splitkv_fwd_template<
576,
32,
16,
4,
true,
true,
cutlass::half_t,
false,
512,
2
>(Flash_fwd_mla_params &params, cudaStream_t stream);

View File

@ -7,7 +7,8 @@
#include <cuda.h>
#include <vector>
#include "mctlass/bfloat16.h"
#include "static_switch.h"
constexpr int maxValidBlockSizeM = 128;
namespace mcFlashAttn {
@ -30,6 +31,10 @@ struct Qkv_params {
index_t q_head_stride;
index_t k_head_stride;
index_t v_head_stride;
index_t indices_batch_stride;
index_t indices_row_stride;
index_t indices_all_valid_per_q_batch_stride;
index_t indices_all_valid_per_q_row_stride;
// The number of heads.
int h, h_k;
@ -62,6 +67,8 @@ struct Flash_fwd_mla_params : public Qkv_params {
// The dimensions.
int b, seqlen_q, seqlen_k, seqlen_knew, d, d_v, seqlen_q_rounded, seqlen_k_rounded, d_rounded, rotary_dim, total_q;
int ngroups;
bool is_sparse_attn = false;
int topk;
// The scaling factors for the kernel.
float scale_softmax;
@ -71,32 +78,15 @@ struct Flash_fwd_mla_params : public Qkv_params {
int * __restrict__ cu_seqlens_q;
int * __restrict__ cu_seqlens_k;
int * __restrict__ leftpad_k;
int *__restrict__ indices_ptr; // [batch, s_q, topk]
// If provided, the actual length of each k sequence.
int * __restrict__ seqused_k;
int *__restrict__ blockmask;
// The K_new and V_new matrices.
void * __restrict__ knew_ptr;
void * __restrict__ vnew_ptr;
// The stride between rows of the Q, K and V matrices.
index_t knew_batch_stride;
index_t vnew_batch_stride;
index_t knew_row_stride;
index_t vnew_row_stride;
index_t knew_head_stride;
index_t vnew_head_stride;
// kv cache dequant
index_t kscale_batch_stride;
index_t vscale_batch_stride;
index_t kscale_row_stride;
index_t vscale_row_stride;
index_t kscale_head_stride;
index_t vscale_head_stride;
// The cos and sin matrices for rotary embedding.
void * __restrict__ rotary_cos_ptr;
void * __restrict__ rotary_sin_ptr;
@ -115,65 +105,45 @@ struct Flash_fwd_mla_params : public Qkv_params {
void *__restrict__ k_scale_ptr;
void *__restrict__ v_scale_ptr;
// The dropout probability (probability of keeping an activation).
float p_dropout;
// uint32_t p_dropout_in_uint;
// uint16_t p_dropout_in_uint16_t;
uint8_t p_dropout_in_uint8_t;
// Scale factor of 1 / (1 - p_dropout).
float rp_dropout;
float scale_softmax_rp_dropout;
// Local window size
int window_size_left, window_size_right;
// ratio of softcapping attention
// S = exp2(log2(e) * softcap * tanh(S * softmax_scale / softcap))
// only value > 0.0 will take effect
float softcap;
// Random state.
// at::PhiloxCudaState philox_args;
// the RNG seed and offset .
uint64_t rng_state_seed = 0;
uint64_t rng_state_offset = 0;
bool is_bf16;
bool is_fp8 = false;
bool is_causal;
bool* indices_all_valid_per_q_ptr;
// If is_seqlens_k_cumulative, then seqlen_k is cu_seqlens_k[bidb + 1] - cu_seqlens_k[bidb].
// Otherwise it's cu_seqlens_k[bidb], i.e., we use cu_seqlens_k to store the sequence lengths of K.
bool is_seqlens_k_cumulative;
bool is_rotary_interleaved;
int num_splits; // For split-KV version
void * __restrict__ alibi_slopes_ptr;
index_t alibi_slopes_batch_stride;
// attn_mask support for bert model Jira[C500-21935]
bool has_attn_mask;
void * __restrict__ attn_mask_ptr = nullptr;
index_t attn_mask_batch_stride = 0;
index_t attn_mask_nheads_stride = 0;
index_t attn_mask_row_stride = 0;
index_t attn_mask_col_stride = 1;
index_t attn_mask_batch_shape = 1;
index_t attn_mask_nheads_shape = 1;
index_t attn_mask_row_shape = 1;
index_t attn_mask_col_shape = 1;
bool unpadded_lse; // For varlen paths: LSE is in [nheads, total_seqlen_q] format instead of [b, nheads, seqlen_q].
bool seqlenq_ngroups_swapped; // q has been transposed from (b, 1, (nheads_kv ngroups), d) to (b, ngroups, nheads_kv, d).
int d_value;
int d_value_rounded;
int arch;
bool is_support_splitkv = false;
int *__restrict__ tile_scheduler_metadata_ptr;
int num_sm_parts;
int *__restrict__ num_splits_ptr;
// fp8 params
float* __restrict__ descale_q_ptr = nullptr;
float* __restrict__ descale_k_ptr = nullptr;
// CP (Context Parallelism) parameters
int cp_world_size;
int cp_rank;
int *__restrict__ cp_tot_seqused_k;
cudaStream_t stream;
};
@ -187,12 +157,46 @@ struct Flash_launch_params {
Flash_launch_params():
is_balance(false),rowblock_parallel(0),block_type(0),performance_mode(false){}
};
}
struct SparsePrefillParams {
int s_q, s_kv, h_q, h_kv, d_qk, d_v, topk;
float sm_scale, sm_scale_div_log2;
int arch;
// Input tensors
mctlass::bfloat16_t* __restrict__ q_ptr; // [s_q, h_q, d_qk]
mctlass::bfloat16_t* __restrict__ kv_ptr; // [s_kv, h_kv, d_qk]
int* __restrict__ indices_ptr; // [s_q, h_kv, topk]
bool* indices_all_valid_per_q_ptr; // [1]
// int stride_q_s_q; int stride_q_h_q;
// int stride_kv_s_kv; int stride_kv_h_kv;
int q_row_stride;int q_head_stride;
int k_row_stride;int k_head_stride;
int stride_indices_s_q; int stride_indices_h_kv;
int o_row_stride;int o_head_stride;
// Output tensors
mctlass::bfloat16_t* __restrict__ out_ptr; // [s_q, h_q, d_v]
float* __restrict__ max_logits; // [s_q, h_q]
float* __restrict__ lse_ptr; // [s_q, h_q]
cudaStream_t stream;
};
static constexpr int TileSchedulerMetaDataSize = 8;
// [begin_idx, begin_seqlen, end_idx, end_seqlen, begin_n_split_idx, _, _, _]
struct GetDecodingMetadataParams {
int *__restrict__ seqlens_k_ptr;
int *__restrict__ tile_scheduler_metadata_ptr;
int *__restrict__ num_splits_ptr;
int batch_size;
int block_size_n;
int fixed_overhead_num_blocks;
int num_sm_parts;
int topk;
};
////////////////////////////////////////////////////////////////////////////////////////////////////
struct Mla_metadata_params {
@ -205,4 +209,4 @@ struct Mla_metadata_params {
int num_sm_parts;
};
void get_mla_metadata_func(Mla_metadata_params &params, cudaStream_t stream);
void run_get_mla_metadata_kernel(GetDecodingMetadataParams &params, cudaStream_t stream);

View File

@ -21,41 +21,131 @@ template<
typename elem_type,
bool Is_splits = false,
int kHeadDimV = kHeadDim,
int Num_Stages = 1
int Num_Stages = 1,
Arch arch = Arch::xcore1000
>
void run_flash_splitkv_fwd_template(Flash_fwd_mla_params &params, cudaStream_t stream);
void run_flash_splitkv_fwd_mla_template(mcFlashAttn::Flash_fwd_mla_params &params, cudaStream_t stream);
template<
int kHeadDim,
int kBlockM,
int kBlockN,
int kNWarps,
bool Is_Q_in_regs,
bool Share_Q_K_smem,
typename elem_type,
bool Is_splits = false,
int kHeadDimV = kHeadDim,
int Num_Stages = 1,
Arch arch = Arch::xcore1000
>
void run_flash_splitkv_fwd_sparse_mla_template(mcFlashAttn::Flash_fwd_mla_params &params, cudaStream_t stream);
template<
int kHeadDim,
int kBlockM,
int kBlockN,
int kNWarps,
bool Is_Q_in_regs,
bool Share_Q_K_smem,
typename elem_type,
bool Is_splits = false,
int kHeadDimV = kHeadDim,
int Num_Stages = 1,
Arch arch = Arch::xcore1000
>
void run_flash_mla_sparse_prefill_template(SparsePrefillParams &params, cudaStream_t stream);
namespace mcFlashAttn {
template<int Headdim>
void run_mha_fwd_splitkv_dispatch(Flash_fwd_mla_params &params, const cudaStream_t stream);
template<int Headdim, Arch arch>
void run_mla_fwd_splitkv_dispatch(Flash_fwd_mla_params &params, const cudaStream_t stream);
template<int Headdim, Arch arch>
void run_flash_mla_sparse_prefill_dispatch(SparsePrefillParams &params, const cudaStream_t stream);
template<>
inline void run_mha_fwd_splitkv_dispatch<576>(Flash_fwd_mla_params &params, const cudaStream_t stream) {
inline void run_flash_mla_sparse_prefill_dispatch<576, Arch::xcore1000>(SparsePrefillParams &params, const cudaStream_t stream) {
constexpr static int HeaddimQK = 576;
constexpr static int HeaddimVO = 512;
constexpr static int Num_Stages = 2;
constexpr Arch arch = Arch::xcore1000;
assert(params.is_bf16 && "sparse prefill only support bf16");
constexpr static int kBlockM = 64;
constexpr static int kBlockN = 16;
constexpr static int kNWarps = 8;
run_flash_mla_sparse_prefill_template<HeaddimQK, kBlockM, kBlockN, kNWarps, true, true, mctlass::bfloat16_t, true, HeaddimVO, Num_Stages, arch>(params, stream);
}
template<>
inline void run_flash_mla_sparse_prefill_dispatch<576, Arch::xcore1500>(SparsePrefillParams &params, const cudaStream_t stream) {
constexpr static int HeaddimQK = 576;
constexpr static int HeaddimVO = 512;
constexpr static int Num_Stages = 2;
constexpr Arch arch = Arch::xcore1500;
assert(params.is_bf16 && "sparse prefill only support bf16");
constexpr static int kBlockM = 64;
constexpr static int kBlockN = 16;
constexpr static int kNWarps = 8;
run_flash_mla_sparse_prefill_template<HeaddimQK, kBlockM, kBlockN, kNWarps, true, true, mctlass::bfloat16_t, true, HeaddimVO, Num_Stages, arch>(params, stream);
}
template<>
inline void run_mla_fwd_splitkv_dispatch<576, Arch::xcore1000>(Flash_fwd_mla_params &params, const cudaStream_t stream) {
constexpr static int HeaddimQK = 576;
constexpr static int HeaddimVO = 512;
constexpr static int Num_Stages = 2;
constexpr Arch arch = Arch::xcore1000;
FP16_SWITCH(!params.is_bf16, [&] {
BOOL_SWITCH(params.num_splits > 1, Is_splits, [&] {
if (!params.is_sparse_attn){
if (params.seqlen_q >= 64) {
constexpr static int kBlockM = 64;
constexpr static int kBlockN = 16;
constexpr static int kNWarps = 8;
run_flash_splitkv_fwd_template<HeaddimQK, kBlockM, kBlockN, kNWarps, true, true, elem_type, Is_splits, HeaddimVO, Num_Stages>(params, stream);
} else if (params.seqlen_q >= 32) {
run_flash_splitkv_fwd_mla_template<HeaddimQK, kBlockM, kBlockN, kNWarps, true, true, elem_type, true, HeaddimVO, Num_Stages, arch>(params, stream);
}
else if (params.seqlen_q >= 32) {
constexpr static int kBlockM = 32;
constexpr static int kBlockN = 16;
constexpr static int kNWarps = 4;
run_flash_splitkv_fwd_template<HeaddimQK, kBlockM, kBlockN, kNWarps, true, true, elem_type, Is_splits, HeaddimVO, Num_Stages>(params, stream);
run_flash_splitkv_fwd_mla_template<HeaddimQK, kBlockM, kBlockN, kNWarps, true, true, elem_type, true, HeaddimVO, Num_Stages, arch>(params, stream);
} else {
constexpr static int kBlockM = 16;
constexpr static int kBlockN = 16;
constexpr static int kNWarps = 4;
run_flash_splitkv_fwd_template<HeaddimQK, kBlockM, kBlockN, kNWarps, true, true, elem_type, Is_splits, HeaddimVO, Num_Stages>(params, stream);
run_flash_splitkv_fwd_mla_template<HeaddimQK, kBlockM, kBlockN, kNWarps, true, true, elem_type, true, HeaddimVO, 1, arch>(params, stream);
}
});
}else{
constexpr static int kBlockM = 64;
constexpr static int kBlockN = 16;
constexpr static int kNWarps = 8;
run_flash_splitkv_fwd_sparse_mla_template<HeaddimQK, kBlockM, kBlockN, kNWarps, true, true, elem_type, true, HeaddimVO, Num_Stages, arch>(params, stream);
}
});
}
template<>
inline void run_mla_fwd_splitkv_dispatch<576, Arch::xcore1500>(Flash_fwd_mla_params &params, const cudaStream_t stream) {
constexpr static int HeaddimQK = 576;
constexpr static int HeaddimVO = 512;
constexpr static int Num_Stages = 2;
constexpr Arch arch = Arch::xcore1500;
FP16_SWITCH(!params.is_bf16, [&] {
if (!params.is_sparse_attn) {
// NOTE: only support blockM=64 in xcore1500 now
constexpr static int kBlockM = 64;
constexpr static int kBlockN = 32;
constexpr static int kNWarps = 8;
run_flash_splitkv_fwd_mla_template<HeaddimQK, kBlockM, kBlockN, kNWarps, true, true, elem_type, true, HeaddimVO, Num_Stages, arch>(params, stream);
} else {
constexpr static int kBlockM = 64;
constexpr static int kBlockN = 16;
constexpr static int kNWarps = 8;
run_flash_splitkv_fwd_sparse_mla_template<HeaddimQK, kBlockM, kBlockN, kNWarps, true, true, elem_type, true, HeaddimVO, Num_Stages, arch>(params, stream);
}
});
}
} // namespace mcFlashAttn end

View File

@ -6,59 +6,143 @@
#include "flash_mla.h"
#include "static_switch.h"
#include "flash_fwd_split_kernel.h"
#include "feature/attn_mask.h"
#include "flash_dense_mla_decode_kernel.h"
#include "flash_sparse_mla_decode_kernel.h"
#include "flash_fwd_splitkv_mla_combine_kernel.h"
#include "xcore1000/sparse_prefill_kernel_64x16_8waves_xcore1000.h"
#include "print_parameter.h"
using namespace mcFlashAttn;
template<typename Kernel_traits, bool Is_causal, bool Is_local, bool Has_alibi, bool Is_even_MN, bool Is_even_K, bool Is_softcap, bool Split>
__global__ void flash_fwd_splitkv_kernel(const Flash_fwd_mla_params params, const int num_m_block) {
flash::compute_attn_splitkv<Kernel_traits, Is_causal, Is_local, Has_alibi, Is_even_MN, Is_even_K, Is_softcap, Split>(params, num_m_block);
template<typename Kernel_traits, bool Is_causal, bool Is_even_MN, bool Is_even_K, bool Is_enable_dcp>
__global__ void flash_fwd_splitkv_mla_kernel(const Flash_fwd_mla_params params, const int num_m_block) {
flash::compute_attn_1rowblock_splitkv_mla<Kernel_traits, Is_causal, Is_even_MN, Is_even_K, Is_enable_dcp>(params, num_m_block);
}
template<typename Kernel_traits, bool Is_causal, bool Is_even_TopK>
__global__ void flash_fwd_splitkv_sparse_mla_kernel(const Flash_fwd_mla_params params, const int num_m_block) {
flash::compute_attn_1rowblock_splitkv_sparse_mla<Kernel_traits, Is_causal, Is_even_TopK>(params, num_m_block);
}
template<typename Kernel_traits, bool Is_causal, bool Is_even_TopK>
__global__ void sparse_attn_global_fwd_kernel(const SparsePrefillParams params) {
flash::sparse_attn_fwd_kernel<Kernel_traits, Is_causal, Is_even_TopK>(params);
}
//flash-meta combine kernel
template<typename Kernel_traits, int kBlockM, int Log_max_splits, bool Is_even_K>
__global__ void flash_fwd_splitkv_combine_kernel(const Flash_fwd_mla_params params) {
__global__ void flash_fwd_splitkv_mla_combine_kernel(const Flash_fwd_mla_params params) {
static_assert(Log_max_splits >= 1);
flash::combine_attn_seqk_parallel<Kernel_traits, kBlockM, Log_max_splits, Is_even_K>(params);
flash::combine_attn_seqk_parallel_splitkv_mla<Kernel_traits, kBlockM, Log_max_splits, Is_even_K>(params);
}
template<typename Kernel_traits, bool Is_causal>
void run_flash_splitkv_fwd(Flash_fwd_mla_params &params, cudaStream_t stream) {
constexpr size_t smem_size = Kernel_traits::kSmemSize;
const int num_m_block = (params.seqlen_q + Kernel_traits::kBlockM - 1) / Kernel_traits::kBlockM;
dim3 grid(num_m_block, params.num_splits > 1 ? params.num_splits : params.b, params.num_splits > 1 ? params.b * params.h : params.h);
template<typename Kernel_traits, bool Is_causal, Arch arch>
void run_flash_splitkv_fwd_mla(Flash_fwd_mla_params &params, cudaStream_t stream) {
constexpr int max_smem_size = arch == Arch::xcore1000 ? 64 * 1024 : 128 * 1024;
constexpr int smem_size = std::min(Kernel_traits::kSmemSize, max_smem_size);
// const int num_m_block = (params.seqlen_q + Kernel_traits::kBlockM - 1) / Kernel_traits::kBlockM;
const int num_m_block = cute::ceil_div(params.seqlen_q, Kernel_traits::kBlockM);
// dim3 grid(num_m_block, params.num_splits > 1 ? params.num_splits : params.b, params.num_splits > 1 ? params.b * params.h : params.h);
dim3 grid(num_m_block, params.h, params.num_sm_parts);
static_assert(Kernel_traits::kHeadDim == 576 && Kernel_traits::kHeadDimV == 512);
const bool is_even_MN = params.cu_seqlens_q == nullptr && params.cu_seqlens_k == nullptr && params.seqlen_k % Kernel_traits::kBlockN == 0 && params.seqlen_q % Kernel_traits::kBlockM == 0;
const bool is_even_K = params.d == Kernel_traits::kHeadDim && params.d_v == Kernel_traits::kHeadDimV;
EVENK_SWITCH(is_even_K, IsEvenKConst, [&] {
LOCAL_SWITCH_AND_CONST_PRECOND((!Is_causal), (params.window_size_left >= 0 || params.window_size_right >= 0) && !Is_causal, Is_local, [&] {
BOOL_SWITCH(params.num_splits > 1, Split, [&] {
auto kernel = &flash_fwd_splitkv_kernel<Kernel_traits, Is_causal, false, false, false, IsEvenKConst, false, Split>;
if (smem_size >= 32 * 1024) {
CUDA_CHECK(cudaFuncSetAttribute(
kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size));
}
kernel<<<grid, Kernel_traits::kNThreads, smem_size, stream>>>(params, num_m_block);
CUDA_KERNEL_LAUNCH_CHECK();
});
BOOL_SWITCH(params.cp_world_size > 1 && params.cp_tot_seqused_k != nullptr, IsEnableDcp, [&] {
EVENK_SWITCH(is_even_K, IsEvenKConst, [&] {
if (std::getenv("MHA_PRINT_PARA")) {
shape_print(params, Is_causal, "mla_dense_decode");
}
if (std::getenv("MHA_DEBUG_PARA")){
debug_print(params, "mla_dense_decode");
}
auto kernel = &flash_fwd_splitkv_mla_kernel<Kernel_traits, Is_causal, false, IsEvenKConst, IsEnableDcp>;
if (smem_size >= 32 * 1024) {
CUDA_CHECK(cudaFuncSetAttribute(
kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size));
}
kernel<<<grid, Kernel_traits::kNThreads, smem_size, stream>>>(params, num_m_block);
CUDA_KERNEL_LAUNCH_CHECK();
});
});
if (params.num_splits > 1) {
// We want kBlockM to be as small as possible for more parallelism.
// With 128 threads we can load 512 elements at a time, so if headdim is divisible by 128, kBlockM = 4.
// If headdim is divisible by 64, then we set kBlockM = 8, etc.
//constexpr static int kBlockM = Kernel_traits::kHeadDim % 128 == 0 ? 4 : (Kernel_traits::kHeadDim % 64 == 0 ? 8 : 16);
constexpr static int kBlockM = Kernel_traits::kHeadDim % 128 == 0 ? 8 : (Kernel_traits::kHeadDim % 64 == 0 ? 16 : 32);
// We want kBlockM to be as small as possible for more parallelism.
// In MLA case head_dim_vo = 512, we will switch different kBlockM for different case to get better performance
COMBINE_BLOCKM_SWITCH(params.b, params.h, params.seqlen_q, kBlockM, [&] {
dim3 grid_combine((params.b * params.h * params.seqlen_q + kBlockM - 1) / kBlockM);
const int kNThreads = 256; /*Kernel_traits::kNThreads;*/
EVENK_SWITCH(is_even_K, IsEvenKConst, [&] {
NUMSPLITS_SWITCH(params.num_splits, kLogMaxSplits, [&] {
flash_fwd_splitkv_combine_kernel<Kernel_traits, kBlockM, kLogMaxSplits, IsEvenKConst><<<grid_combine, kNThreads, 0, stream>>>(params);
NUMSPLITS_SWITCH(params.num_sm_parts, kLogMaxSplits, [&] {
const int kNThreads = kBlockM == 1 ? 64 : (kBlockM == 2 ? 128 : 256);
flash_fwd_splitkv_mla_combine_kernel<Kernel_traits, kBlockM, kLogMaxSplits, IsEvenKConst><<<grid_combine, kNThreads, 0, stream>>>(params);
CUDA_KERNEL_LAUNCH_CHECK();
});
});
}
});
}
template<typename Kernel_traits, bool Is_causal, Arch arch>
void run_flash_splitkv_fwd_sparse_mla(Flash_fwd_mla_params &params, cudaStream_t stream) {
constexpr int max_smem_size = arch == Arch::xcore1000 ? 64 * 1024 : 128 * 1024;
constexpr int smem_size = std::min(Kernel_traits::kSmemSize, max_smem_size);
// const int num_m_block = (params.seqlen_q + Kernel_traits::kBlockM - 1) / Kernel_traits::kBlockM;
const int num_m_block = cute::ceil_div(params.seqlen_q, Kernel_traits::kBlockM);
// dim3 grid(num_m_block, params.num_splits > 1 ? params.num_splits : params.b, params.num_splits > 1 ? params.b * params.h : params.h);
dim3 grid(num_m_block, params.h, params.num_sm_parts);
static_assert(Kernel_traits::kHeadDim == 576 && Kernel_traits::kHeadDimV == 512);
// const bool is_even_MN = params.cu_seqlens_q == nullptr && params.cu_seqlens_k == nullptr && params.seqlen_k % Kernel_traits::kBlockN == 0 && params.seqlen_q % Kernel_traits::kBlockM == 0;
const bool is_even_K = true; // is_even_k is always true in mla case;
const bool is_even_topK = params.topk % Kernel_traits::kBlockN == 0;
BOOL_SWITCH(is_even_topK, IsEvenTopKConst, [&] {
if (std::getenv("MHA_PRINT_PARA")) {
shape_print(params, Is_causal, "mla_sparse_decode");
}
if (std::getenv("MHA_DEBUG_PARA")){
debug_print(params, "mla_sparse_decode");
}
auto kernel = &flash_fwd_splitkv_sparse_mla_kernel<Kernel_traits, Is_causal, IsEvenTopKConst>;
if (smem_size >= 32 * 1024) {
CUDA_CHECK(cudaFuncSetAttribute(
kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size));
}
kernel<<<grid, Kernel_traits::kNThreads, smem_size, stream>>>(params, num_m_block);
CUDA_KERNEL_LAUNCH_CHECK();
});
// We want kBlockM to be as small as possible for more parallelism.
// In MLA case head_dim_vo = 512, we will switch different kBlockM for different case to get better performance
COMBINE_BLOCKM_SWITCH(params.b, params.h, params.seqlen_q, kBlockM, [&] {
dim3 grid_combine((params.b * params.h * params.seqlen_q + kBlockM - 1) / kBlockM);
EVENK_SWITCH(is_even_K, IsEvenKConst, [&] {
NUMSPLITS_SWITCH(params.num_sm_parts, kLogMaxSplits, [&] {
const int kNThreads = kBlockM == 1 ? 64 : (kBlockM == 2 ? 128 : 256);
flash_fwd_splitkv_mla_combine_kernel<Kernel_traits, kBlockM, kLogMaxSplits, IsEvenKConst><<<grid_combine, kNThreads, 0, stream>>>(params);
CUDA_KERNEL_LAUNCH_CHECK();
});
});
});
}
template<typename Kernel_traits, bool Is_causal, Arch arch>
void run_sparse_prefill(SparsePrefillParams &params, cudaStream_t stream) {
constexpr int max_smem_size = arch == Arch::xcore1000 ? 64 * 1024 : 128 * 1024;
constexpr int smem_size = std::min(Kernel_traits::kSmemSize, max_smem_size);
dim3 grid((params.h_q/Kernel_traits::kBlockM)*params.s_q, 1, 1);
static_assert(Kernel_traits::kHeadDim == 576 && Kernel_traits::kHeadDimV == 512);
const bool is_even_topK = params.topk % Kernel_traits::kBlockN == 0;
// WARNING: Be aware of the correctness of this condition
BOOL_SWITCH(is_even_topK, IsEvenTopKConst, [&] {
if (std::getenv("MHA_PRINT_PARA")) {
shape_print(params, Is_causal, "mla_sparse_prefill");
}
if (std::getenv("MHA_DEBUG_PARA")){
debug_print(params, "mla_sparse_prefill");
}
auto kernel = &sparse_attn_global_fwd_kernel<Kernel_traits, /*is_causal*/false, IsEvenTopKConst>;
if (smem_size >= 32 * 1024) {
CUDA_CHECK(cudaFuncSetAttribute(
kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size));
}
kernel<<<grid, Kernel_traits::kNThreads, smem_size, stream>>>(params);
CUDA_KERNEL_LAUNCH_CHECK();
});
}

View File

@ -0,0 +1,69 @@
// Adapted from Dao-AILab/flash-attention (https://github.com/Dao-AILab/flash-attention/tree/v2.6.3)
#pragma once
#include <cuda.h>
#include "flash_fwd_launch_template.h"
#include "flash_mla.h"
#include "static_switch.h"
template<
int kHeadDim,
int kBlockM,
int kBlockN,
int kNWarps,
bool Is_Q_in_regs,
bool Share_Q_K_smem,
typename elem_type,
bool Is_splits = false,
int kHeadDimV = kHeadDim,
int Num_Stages = 1,
Arch arch = Arch::xcore1000
>
void run_flash_splitkv_fwd_mla_template(mcFlashAttn::Flash_fwd_mla_params &params, cudaStream_t stream){
using Kernel_traits = Flash_fwd_kernel_traits<kHeadDim, kBlockM, kBlockN, kNWarps, Is_Q_in_regs, Share_Q_K_smem, elem_type, Is_splits, kHeadDimV, Num_Stages>;
BOOL_SWITCH(params.is_causal, Is_causal, [&] {
run_flash_splitkv_fwd_mla<Kernel_traits, Is_causal, arch>(params, stream);
});
}
template<
int kHeadDim,
int kBlockM,
int kBlockN,
int kNWarps,
bool Is_Q_in_regs,
bool Share_Q_K_smem,
typename elem_type,
bool Is_splits = false,
int kHeadDimV = kHeadDim,
int Num_Stages = 1,
Arch arch = Arch::xcore1000
>
void run_flash_splitkv_fwd_sparse_mla_template(mcFlashAttn::Flash_fwd_mla_params &params, cudaStream_t stream){
using Kernel_traits = Flash_fwd_kernel_traits<kHeadDim, kBlockM, kBlockN, kNWarps, Is_Q_in_regs, Share_Q_K_smem, elem_type, Is_splits, kHeadDimV, Num_Stages>;
BOOL_SWITCH(params.is_causal, Is_causal, [&] {
run_flash_splitkv_fwd_sparse_mla<Kernel_traits, Is_causal, arch>(params, stream);
});
}
template<
int kHeadDim,
int kBlockM,
int kBlockN,
int kNWarps,
bool Is_Q_in_regs,
bool Share_Q_K_smem,
typename elem_type,
bool Is_splits = false,
int kHeadDimV = kHeadDim,
int Num_Stages = 1,
Arch arch = Arch::xcore1000
>
void run_flash_mla_sparse_prefill_template(SparsePrefillParams &params, cudaStream_t stream){
using Kernel_traits = Flash_fwd_kernel_traits<kHeadDim, kBlockM, kBlockN, kNWarps, Is_Q_in_regs, Share_Q_K_smem, elem_type, Is_splits, kHeadDimV, Num_Stages>;
run_sparse_prefill<Kernel_traits, false, arch>(params, stream);
}

View File

@ -1,28 +0,0 @@
// Adapted from Dao-AILab/flash-attention (https://github.com/Dao-AILab/flash-attention/tree/v2.6.3)
#pragma once
#include <cuda.h>
#include "flash_fwd_launch_template.h"
#include "flash_mla.h"
#include "static_switch.h"
template<
int kHeadDim,
int kBlockM,
int kBlockN,
int kNWarps,
bool Is_Q_in_regs,
bool Share_Q_K_smem,
typename elem_type,
bool Is_splits = false,
int kHeadDimV = kHeadDim,
int Num_Stages = 1
>
void run_flash_splitkv_fwd_template(mcFlashAttn::Flash_fwd_mla_params &params, cudaStream_t stream){
using Kernel_traits = Flash_fwd_kernel_traits<kHeadDim, kBlockM, kBlockN, kNWarps, Is_Q_in_regs, Share_Q_K_smem, elem_type, Is_splits, kHeadDimV, Num_Stages>;
BOOL_SWITCH(params.is_causal, Is_causal, [&] {
run_flash_splitkv_fwd<Kernel_traits,Is_causal>(params, stream);
});
}

View File

@ -1,76 +0,0 @@
// Adapted from Dao-AILab/flash-attention (https://github.com/Dao-AILab/flash-attention/tree/v2.6.3)
#pragma once
#include <cmath>
#include <cute/tensor.hpp>
#include <mctlass/mctlass.h>
#include <mctlass/array.h>
#include "utils.h"
namespace flash {
using namespace cute;
////////////////////////////////////////////////////////////////////////////////////////////////////
template <bool Is_causal>
struct Alibi {
const float alibi_slope;
const int max_seqlen_k, max_seqlen_q;
__forceinline__ __device__ Alibi(const float alibi_slope, const int max_seqlen_k, const int max_seqlen_q)
: alibi_slope(alibi_slope)
, max_seqlen_k(max_seqlen_k)
, max_seqlen_q(max_seqlen_q) {
};
template <typename Engine, typename Layout>
__forceinline__ __device__ void apply_alibi(Tensor<Engine, Layout> &tensor,
const int col_idx_offset_,
const int row_idx_offset,
const int warp_row_stride,
const int warp_col_stride = 16) {
// tensor has shape (nrow=(1, MMA_M), ncol=(4, MMA_N))
static_assert(Layout::rank == 2, "Only support 2D Tensor");
static_assert(decltype(size<0, 0>(tensor))::value == 1);
static_assert(decltype(size<1, 0>(tensor))::value == 4);
const int col_idx_offset = col_idx_offset_ + ((__lane_id() >> 4) << 2);
if constexpr (Is_causal) { // Simpler, we add the same bias vector to all rows
#pragma unroll
for (int nj = 0; nj < size<1, 1>(tensor); ++nj) {
const int col_idx_base = col_idx_offset + nj * warp_col_stride;
#pragma unroll
for (int j = 0; j < size<1, 0>(tensor); ++j) {
const int col_idx = col_idx_base + j;
#pragma unroll
for (int mi = 0; mi < size<0>(tensor); ++mi) {
tensor(mi, make_coord(j, nj)) += alibi_slope * col_idx;
}
}
}
} else { // Bias depends on both row_idx and col_idx
#pragma unroll
for (int mi = 0; mi < size<0, 1>(tensor); ++mi) {
const int row_idx = row_idx_offset + mi * warp_row_stride;
#pragma unroll
for (int nj = 0; nj < size<1, 1>(tensor); ++nj) {
const int col_idx_base = col_idx_offset + nj * warp_col_stride;
#pragma unroll
for (int j = 0; j < size<1, 0>(tensor); ++j) {
const int col_idx = col_idx_base + j;
tensor(make_coord(0, mi), make_coord(j, nj)) -= alibi_slope * abs(row_idx + max_seqlen_k - max_seqlen_q - col_idx);
}
}
}
}
}
};
} // namespace flash

View File

@ -1,166 +0,0 @@
// Adapted from Dao-AILab/flash-attention (https://github.com/Dao-AILab/flash-attention/tree/v2.6.3)
#pragma once
#include <cmath>
#include <cute/tensor.hpp>
#include <cutlass/cutlass.h>
#include <mctlass/array.h>
#include "utils.h"
#include "block_info.h"
namespace flash {
using namespace cute;
////////////////////////////////////////////////////////////////////////////////////////////////////
template<typename Params>
inline __host__ __device__ bool use_attn_mask_merge_ldg(Params &params) {
/*merged impl require
1. bias_col_shape % 4 == 0
2. bias_col_stride == 1
*/
if((params.attn_mask_col_shape % 4) == 0
&& params.attn_mask_col_stride == 1) {
return true;
} else {
return false;
}
}
template <bool mergeLdg=false, bool Is_even_MN=false,typename Engine, typename Layout, typename T>
inline __device__ void apply_attn_mask(Tensor<Engine, Layout> &tensor,
const int col_idx_offset_,
const int max_seqlen_k,
const int row_idx_offset_,
const int max_seqlen_q,
const int warp_row_stride,
const int warp_col_stride,
const float softmax_scale,
T *bias,
const int bias_row_stride,
const int bias_col_stride) {
// tensor has shape (nrow=(1, MMA_M), ncol=(4, MMA_N)
CUTE_STATIC_ASSERT_V((size<1, 0>(tensor)) == Int<4>{});
static_assert(Layout::rank == 2, "Only support 2D Tensor");
typedef __NATIVE_VECTOR__(2, int) VecType;
const int lane_id = threadIdx.x % 64;
const int row_idx_offset = row_idx_offset_;
const int col_idx_offset = col_idx_offset_ + (lane_id / 16) * 4;
if constexpr (mergeLdg) {
#pragma unroll
for (int mi = 0; mi < size<0, 1>(tensor); ++mi) {
const int row_idx = row_idx_offset + mi * warp_row_stride;
#pragma unroll
for (int nj = 0; nj < size<1, 1>(tensor); ++nj) {
const int col_idx_base = col_idx_offset + nj * warp_col_stride;
T bias_16[4];
if constexpr (Is_even_MN) {
uint64_t *bias_64 = reinterpret_cast<uint64_t *>(bias_16);
bias_64[0] = *((uint64_t *)(bias + row_idx * bias_row_stride + col_idx_base));
} else {
bool mask = row_idx < max_seqlen_q && col_idx_base < max_seqlen_k;
VecType *dst_ptr = reinterpret_cast<VecType *>(bias_16);
VecType *src_ptr = reinterpret_cast<VecType *>(bias + row_idx * bias_row_stride + col_idx_base);
*dst_ptr = __builtin_mxc_ldg_b64_predicator(src_ptr, 0, true, true, false, false,
mask, 1, MACA_ICMP_EQ);
}
#pragma unroll
for (int j = 0; j < size<1, 0>(tensor); ++j) {
if (row_idx < max_seqlen_q && col_idx_base < max_seqlen_k) {
tensor(make_coord(0, mi), make_coord(j, nj)) += bias_16[j] / softmax_scale;
}
}
}
}
} else {
#pragma unroll
for (int mi = 0; mi < size<0, 1>(tensor); ++mi) {
const int row_idx = row_idx_offset + mi * warp_row_stride;
bool row_mask = Is_even_MN || row_idx < max_seqlen_q;
#pragma unroll
for (int nj = 0; nj < size<1, 1>(tensor); ++nj) {
const int col_idx_base = col_idx_offset + nj * warp_col_stride;
#pragma unroll
for (int j = 0; j < size<1, 0>(tensor); ++j) {
/*naive impl will have ldg_u16*/
const int col_idx = col_idx_base + j;
if (row_mask && col_idx < max_seqlen_k) {
tensor(make_coord(0, mi), make_coord(j, nj)) += *(bias + row_idx * bias_row_stride + col_idx * bias_col_stride) / softmax_scale;
}
}
}
}
}
}
template <typename Tensor0, typename Tensor1>
inline __device__ void apply_attn_mask(Tensor0& acc_s, Tensor1& tSrMask, const float softmax_scale) {
// acc_s shape is (4, MMA_M, MMA_N)
CUTE_STATIC_ASSERT_V(size<0>(acc_s) == _4{});
CUTE_STATIC_ASSERT_V(size<0>(acc_s) == size<0>(tSrMask));
CUTE_STATIC_ASSERT_V(size<1>(acc_s) == size<1>(tSrMask));
CUTE_STATIC_ASSERT_V(size<2>(acc_s) == size<2>(tSrMask));
using T = typename Tensor1::value_type;
CONVERT_TENSOR_TYPE(T, float, tSrMask, rMask)
#pragma unroll
for (int m = 0; m < size<1>(acc_s); m++) {
#pragma unroll
for (int n = 0; n < size<2>(acc_s); n++) {
#pragma unroll
for (int i = 0; i < size<0>(acc_s); i++) {
acc_s(i, m, n) += rMask(i, m, n) / softmax_scale;
}
}
}
}
template <bool mergeLdg=false, bool Is_even_MN=false, typename Tensor0, typename Tensor1, typename Tensor2>
inline __device__ void load_attn_mask(Tensor0& tSgMask, Tensor1& tSrMask, Tensor2& tScMask, const int max_N, const int max_M) {
// load attn_mask bias from global -> register
// tSgMask shape is (4, MMA_M, MMA_N)
CUTE_STATIC_ASSERT_V(size<0>(tSgMask) == _4{});
CUTE_STATIC_ASSERT_V(size<0>(tSgMask) == size<0>(tSrMask));
CUTE_STATIC_ASSERT_V(size<1>(tSgMask) == size<1>(tSrMask));
CUTE_STATIC_ASSERT_V(size<2>(tSgMask) == size<2>(tSrMask));
typedef __NATIVE_VECTOR__(2, int) VecType;
if constexpr (mergeLdg) {
#pragma unroll
for (int m = 0; m < size<1>(tSgMask); m++) {
bool row_mask = Is_even_MN || get<0>(tScMask(0, m, 0)) < max_M;
#pragma unroll
for (int n = 0; n < size<2>(tSgMask); n++) {
bool col_mask = Is_even_MN || get<1>(tScMask(0, 0, n)) < max_N;
auto src_ptr = (VecType *)(tSgMask(_, m, n).data().get()); // gmem
auto dst_ptr = (VecType *)(tSrMask(_, m, n).data()); // rf
if constexpr (Is_even_MN) {
*dst_ptr = __builtin_mxc_ldg_b64(src_ptr, 0, -1, true, true, false, false);
} else{
*dst_ptr = __builtin_mxc_ldg_b64_predicator(src_ptr, 0, true, true, false, false,
row_mask && col_mask, 1, MACA_ICMP_EQ);
}
}
}
} else {
#pragma unroll
for (int m = 0; m < size<1>(tSgMask); m++) {
bool row_mask = Is_even_MN || get<0>(tScMask(0, m, 0)) < max_M;
#pragma unroll
for (int n = 0; n < size<2>(tSgMask); n++) {
int col_base_idx = get<1>(tScMask(0, 0, n));
#pragma unroll
for (int i = 0; i < size<0>(tSgMask); i++) {
if (row_mask && col_base_idx + i < max_N) {
tSrMask(i, m, n) = tSgMask(i, m, n);
}
}
}
}
}
}
} // namespace flash

View File

@ -1,206 +0,0 @@
// Adapted from Dao-AILab/flash-attention (https://github.com/Dao-AILab/flash-attention/tree/v2.6.3)
/******************************************************************************
* Copyright (c) 2024, Tri Dao.
******************************************************************************/
#pragma once
#include "philox.cuh"
#include "utils.h"
namespace flash {
struct Dropout {
const unsigned long long seed, offset;
const uint8_t p_dropout_in_uint8_t;
__forceinline__ __device__ Dropout(const unsigned long long seed, const unsigned long long offset,
const uint8_t p_dropout_in_uint8_t,
const int bid, const int hid, const int tid, const int nheads)
: seed(seed)
, offset(offset + (bid * nheads + hid) * 64 + tid % 64)
, p_dropout_in_uint8_t(p_dropout_in_uint8_t) {
}
template <bool encode_dropout_in_sign_bit=false, typename Engine, typename Layout>
__forceinline__ __device__ void apply_dropout(Tensor<Engine, Layout> &tensor_,
int block_row_start, int block_col_start, int block_row_stride) {
// convert shape from (4, MMA_M, MMA_N) to (8, MMA_M, MMA_N / 2)
Tensor tensor = make_tensor(tensor_.data(), flash::convert_layout_acc_dropout(tensor_.layout()));
using T = typename Engine::value_type;
auto encode_dropout = [](bool keep, T val) {
return keep ? val : (encode_dropout_in_sign_bit ? -val : T(0));
};
static_assert(decltype(size<2>(tensor))::value % 2 == 0);
const uint16_t p_dropout_8bit_in_uint16_t = uint16_t(p_dropout_in_uint8_t);
const uint32_t p_dropout_8bit_in_uint32_t = (uint32_t(p_dropout_8bit_in_uint16_t) << 16) | uint32_t(p_dropout_8bit_in_uint16_t);
// if (cute::thread0()) { printf("threshold2 = 0x%x\n", p_dropout_8bit_in_uint32_t); }
#pragma unroll
for (int m = 0; m < size<1>(tensor); ++m, block_row_start += block_row_stride) {
uint2 rowcol = make_uint2(block_row_start, block_col_start);
#pragma unroll
for (int n = 0; n < size<2>(tensor) / 2; ++n, ++rowcol.y) {
// if (cute::thread(32, 0)) { printf("m = %d, n = %d, row = %d, col = %d\n", m, n, int(rowcol.x), int(rowcol.y));}
uint4 random_uint4 = flash::philox(seed, reinterpret_cast<unsigned long long&>(rowcol), offset);
// if (cute::thread0()) { printf("philox = %u, %d, %d, %d\n", random_uint4.x, random_uint4.y, random_uint4.z, random_uint4.w);}
uint8_t (&rnd_8)[16] = reinterpret_cast<uint8_t (&)[16]>(random_uint4);
// Special implementation for 16-bit types: we duplicate the threshold to the
// low and high 16 bits of a 32-bit value, then use the f16x2 comparison instruction
// to get a mask. The low 16 bits of the mask will be either 0xffff or 0x0000,
// and the high 16 bits will be either 0xffff or 0x0000, depending on whether
// the random value is less than the threshold.
// We then do a bit-wise AND between the mask and the original value (in 32-bit).
// We're exploiting the fact that floating point comparison is equivalent to integer
// comparison, since we're comparing unsigned integers whose top 8-bits are zero.
if (!encode_dropout_in_sign_bit
&& (std::is_same<T, mctlass::half_t>::value || std::is_same<T, mctlass::bfloat16_t>::value)) {
uint16_t rnd_16[16];
#pragma unroll
for (int i = 0; i < 16; i++) { rnd_16[i] = uint16_t(rnd_8[i]); }
uint32_t (&rnd_32)[8] = reinterpret_cast<uint32_t (&)[8]>(rnd_16);
#pragma unroll
for (int j = 0; j < 2; j++) {
Tensor tensor_uint32 = recast<uint32_t>(tensor(_, m, n * 2 + j));
// if (cute::thread0()) { printf("random = 0x%x, 0x%x, 0x%x, 0x%x\n", rnd_32[j * 4 + 0], rnd_32[j * 4 + 1], rnd_32[j * 4 + 2], rnd_32[j * 4 + 3]); }
// if (cute::thread0()) { printf("tensor_uint32 = 0x%x, 0x%x, 0x%x, 0x%x\n", tensor_uint32(0), tensor_uint32(1), tensor_uint32(2), tensor_uint32(3)); }
#pragma unroll
for (int i = 0; i < 4; i++) {
uint32_t mask;
//asm volatile("set.le.u32.f16x2 %0, %1, %2;\n" : "=r"(mask) : "r"(rnd_32[j * 4 + i]), "r"(p_dropout_8bit_in_uint32_t));
tensor_uint32(i) &= mask;
}
// if (cute::thread0()) { printf("tensor_uint32 = 0x%x, 0x%x, 0x%x, 0x%x\n", tensor_uint32(0), tensor_uint32(1), tensor_uint32(2), tensor_uint32(3)); }
}
} else {
#pragma unroll
for (int j = 0; j < 2; j++) {
#pragma unroll
for (int i = 0; i < 8; i++) {
tensor(i, m, n * 2 + j) = encode_dropout(rnd_8[j * 8 + i] <= p_dropout_in_uint8_t, tensor(i, m, n * 2 + j));
}
Tensor tensor_uint32 = recast<uint32_t>(tensor(_, m, n * 2 + j));
// if (cute::thread0()) { printf("tensor_uint32 = 0x%x, 0x%x, 0x%x, 0x%x\n", tensor_uint32(0), tensor_uint32(1), tensor_uint32(2), tensor_uint32(3)); }
}
}
// // if ((threadIdx.x == 0) && (blockIdx.x == 0) && (blockIdx.y == 0)) {
// // printf("n = %d, ph Philox: %u, %u, %u, %u\n", n, rnd_8.x, rnd_8.y, rnd_8.z, rnd_8.w);
// // }
}
}
}
template <bool encode_dropout_in_sign_bit = false, int AtomLayoutNS = 1, typename Engine, typename Layout>
__forceinline__ __device__ void mc_apply_dropout(Tensor<Engine, Layout> &tensor,
int block_row_start, int block_col_start,
int block_row_stride,
int kBlockN, int n_block) {
using T = typename Engine::value_type;
auto encode_dropout = [](bool keep, T val) {
return keep ? val : (encode_dropout_in_sign_bit ? -val : T(0));
};
static_assert(decltype(size<0>(tensor))::value == 4);
int warp_id = threadIdx.x / 64;
#pragma unroll
for (int m = 0; m < size<1>(tensor); ++m, block_row_start += block_row_stride) {
// use block_col_offset to control rnd_8 apply
int block_col_offset = ((kBlockN * n_block) % 64) / 16;
int warp_col_offset = warp_id / block_row_stride;
int block_col_start_tmp = block_col_start;
#pragma unroll
for (int n = 0; n < size<2>(tensor); ++n, warp_col_offset += AtomLayoutNS) {
// when blockN=128, one block contain two 64 in col, so we need to update block_col_start_tmp
block_col_start_tmp += warp_col_offset / 4;
warp_col_offset %= 4;
// rnd_8 contains 16 nums for kBlockN 64 in 1 block, to process 64 in col
// if use kBlockN 32 in 2 block to process 64 in col, rnd_8 should apply in two differen block
uint2 rowcol = make_uint2(block_row_start, block_col_start_tmp);
uint4 random_uint4 = flash::philox(seed, reinterpret_cast<unsigned long long &>(rowcol), offset);
uint8_t (&rnd_8)[16] = reinterpret_cast<uint8_t (&)[16]>(random_uint4);
#pragma unroll
for (int i = 0; i < 4; ++i) {
int rng_idx = (block_col_offset + warp_col_offset) * 4 + i;
// This implementation is a native perf version,here assembly instruction will have ldp_u8.
tensor(i, m, n) = encode_dropout(rnd_8[rng_idx] <= p_dropout_in_uint8_t, tensor(i, m, n));
}
}
}
}
// it works for:
// blockN=64, waves layout: 4x1, 2x2, 4x2, 1x4, 2x4
// blockN=128, waves layout: 1x4, 2x4
template <bool encode_dropout_in_sign_bit = false, int AtomLayoutMS = 2, int AtomLayoutNS = 2, typename Engine, typename Layout>
__forceinline__ __device__ void mc_apply_dropout(Tensor<Engine, Layout> &tensor, int block_row_start, int block_col_start) {
using T = typename Engine::value_type;
auto encode_dropout = [](bool keep, T val) {
return keep ? val : (encode_dropout_in_sign_bit ? -val : T(0));
};
static_assert(decltype(size<0>(tensor))::value == 4);
const int wave_col = threadIdx.x / 64 / AtomLayoutMS;
if constexpr (AtomLayoutNS == 1) {
#pragma unroll
for (int m = 0; m < size<1>(tensor); ++m, block_row_start += AtomLayoutMS) {
uint2 rowcol = make_uint2(block_row_start, block_col_start);
uint4 random_uint4 = flash::philox(seed, reinterpret_cast<unsigned long long &>(rowcol), offset);
uint8_t (&rnd_8)[16] = reinterpret_cast<uint8_t (&)[16]>(random_uint4);
#pragma unroll
for (int n = 0; n < size<2>(tensor); ++n) {
#pragma unroll
for (int i = 0; i < 4; ++i) {
// w0|w0|w0|w0
tensor(i, m, n) = encode_dropout(rnd_8[n * 4 + i] <= p_dropout_in_uint8_t, tensor(i, m, n));
}
}
}
} else if constexpr (AtomLayoutNS == 2) {
#pragma unroll
for (int m = 0; m < size<1>(tensor); ++m, block_row_start += AtomLayoutMS) {
uint2 rowcol = make_uint2(block_row_start, block_col_start);
uint4 random_uint4 = flash::philox(seed, reinterpret_cast<unsigned long long &>(rowcol), offset);
uint8_t (&rnd_8)[16] = reinterpret_cast<uint8_t (&)[16]>(random_uint4);
#pragma unroll
for (int n = 0; n < size<2>(tensor); ++n) {
#pragma unroll
for (int i = 0; i < 4; ++i) {
// e.g., w0|w2|w0|w2
if (wave_col == 0) {
tensor(i, m, n) = encode_dropout(rnd_8[n * 8 + i] <= p_dropout_in_uint8_t, tensor(i, m, n));
} else {
tensor(i, m, n) = encode_dropout(rnd_8[n * 8 + i + 4] <= p_dropout_in_uint8_t, tensor(i, m, n));
}
}
}
}
} else if constexpr (AtomLayoutNS == 4) {
#pragma unroll
for (int m = 0; m < size<1>(tensor); ++m, block_row_start += AtomLayoutMS) {
#pragma unroll
for (int n = 0; n < size<2>(tensor); ++n, block_col_start += 1) {
uint2 rowcol = make_uint2(block_row_start, block_col_start);
uint4 random_uint4 = flash::philox(seed, reinterpret_cast<unsigned long long &>(rowcol), offset);
uint8_t (&rnd_8)[16] = reinterpret_cast<uint8_t (&)[16]>(random_uint4);
#pragma unroll
for (int i = 0; i < 4; ++i) {
// e.g., w0|w2|w4|w6
if (wave_col == 0) {
tensor(i, m, n) = encode_dropout(rnd_8[i] <= p_dropout_in_uint8_t, tensor(i, m, n));
} else if (wave_col == 1) {
tensor(i, m, n) = encode_dropout(rnd_8[i + 4] <= p_dropout_in_uint8_t, tensor(i, m, n));
} else if (wave_col == 2) {
tensor(i, m, n) = encode_dropout(rnd_8[i + 8] <= p_dropout_in_uint8_t, tensor(i, m, n));
} else {
tensor(i, m, n) = encode_dropout(rnd_8[i + 12] <= p_dropout_in_uint8_t, tensor(i, m, n));
}
}
}
}
}
}
};
} // namespace flash

View File

@ -7,6 +7,7 @@
#pragma once
#include <cute/tensor.hpp>
#include "utils.h"
namespace flash {
@ -37,113 +38,38 @@ __forceinline__ __device__ void apply_mask(Tensor<Engine, Layout> &tensor, const
}
}
// [warp_col_stride] tiled mma 4x1: 16, tiled mma 2x2: 32
template <bool HasWSLeft=true, typename Engine, typename Layout>
__forceinline__ __device__ void apply_mask_local(Tensor<Engine, Layout> &tensor, const int col_idx_offset_,
const int max_seqlen_k, const int row_idx_offset,
const int max_seqlen_q, const int warp_row_stride,
const int window_size_left, const int window_size_right,
const int warp_col_stride = 16) {
// tensor has shape (nrow=(1, MMA_M), ncol=(4, MMA_N))
static_assert(Layout::rank == 2, "Only support 2D Tensor");
static_assert(decltype(size<0, 0>(tensor))::value == 1);
static_assert(decltype(size<1, 0>(tensor))::value == 4);
const int col_idx_offset = col_idx_offset_ + ((__lane_id() >> 4) << 2);
#pragma unroll
for (int mi = 0; mi < size<0, 1>(tensor); ++mi) {
const int row_idx = row_idx_offset + mi * warp_row_stride;
const int col_idx_limit_left = std::max(0, row_idx + max_seqlen_k - max_seqlen_q - window_size_left);
const int col_idx_limit_right = std::min(max_seqlen_k, row_idx + 1 + max_seqlen_k - max_seqlen_q + window_size_right);
#pragma unroll
for (int nj = 0; nj < size<1, 1>(tensor); ++nj) {
const int col_idx_base = col_idx_offset + nj * warp_col_stride;
#pragma unroll
for (int j = 0; j < size<1, 0>(tensor); ++j) {
const int col_idx = col_idx_base + j;
if (col_idx >= col_idx_limit_right || (HasWSLeft && col_idx < col_idx_limit_left)) {
tensor(make_coord(0, mi), make_coord(j, nj)) = -INFINITY;
}
}
}
// if (cute::thread0()) {
// printf("mi = %d, i = %d, row_idx = %d, max_seqlen_k = %d\n", mi, i, row_idx, max_seqlen_k);
// print(tensor(make_coord(i, mi), _));
// // print(tensor(_, j + nj * size<1, 0>(tensor)));
// }
}
}
template <typename Engine, typename Layout>
__forceinline__ __device__ void apply_mask_causal(Tensor<Engine, Layout> &tensor, const int col_idx_offset_,
const int max_seqlen_k, const int row_idx_offset,
const int max_seqlen_q, const int warp_row_stride,
const int warp_col_stride = 16) {
// Causal masking is equivalent to local masking with window_size_left = infinity and window_size_right = 0
apply_mask_local</*HasWSLeft=*/false>(tensor, col_idx_offset_, max_seqlen_k, row_idx_offset,
max_seqlen_q, warp_row_stride, -1, 0, warp_col_stride);
}
template <typename Engine0, typename Layout0, typename Engine1, typename Layout1>
__forceinline__ __device__ void apply_mask_causal_w_idx(
Tensor<Engine0, Layout0> &tensor, Tensor<Engine1, Layout1> const &idx_rowcol,
const int col_idx_offset_, const int max_seqlen_k, const int row_idx_offset)
{
// tensor has shape (ncol=(2, MMA_M), nrow=(2, MMA_N))
static_assert(Layout0::rank == 2, "Only support 2D Tensor");
static_assert(Layout1::rank == 2, "Only support 2D Tensor");
CUTE_STATIC_ASSERT_V(size<0>(tensor) == size<0>(idx_rowcol));
CUTE_STATIC_ASSERT_V(size<1>(tensor) == size<1>(idx_rowcol));
#pragma unroll
for (int mi = 0; mi < size<0>(tensor); ++mi) {
const int col_idx_limit = std::min(max_seqlen_k, 1 + row_idx_offset + get<0>(idx_rowcol(mi, 0)));
#pragma unroll
for (int ni = 0; ni < size<1, 1>(tensor); ++ni) {
if (col_idx_offset_ + get<1>(idx_rowcol(0, ni)) >= col_idx_limit) {
tensor(mi, ni) = -INFINITY;
}
}
// if (cute::thread0()) {
// printf("ni = %d, j = %d, col_idx = %d, max_seqlen_k = %d\n", ni, j, col_idx, max_seqlen_k);
// print(tensor(_, make_coord(j, ni)));
// // print(tensor(_, j + ni * size<1, 0>(tensor)));
// }
}
}
template <bool Is_causal, bool Is_local, bool Has_alibi>
template <bool Is_causal, bool Is_enable_dcp = false>
struct Mask {
const int max_seqlen_k, max_seqlen_q, ngroups;
const int window_size_left, window_size_right;
const float alibi_slope;
// CP (Context Parallelism) parameters
const int tot_seqlen_k, cp_world_size, cp_rank;
__forceinline__ __device__ Mask(const int max_seqlen_k, const int max_seqlen_q, const int ngroups,
const int window_size_left, const int window_size_right,
const float alibi_slope=0.f)
__forceinline__ __device__ Mask(const int max_seqlen_k, const int max_seqlen_q, const int ngroups, const int tot_seqlen_k = 0, const int cp_world_size = 1, const int cp_rank = 0)
: max_seqlen_k(max_seqlen_k)
, max_seqlen_q(max_seqlen_q)
, ngroups(ngroups)
, window_size_left(window_size_left)
, window_size_right(window_size_right)
, alibi_slope(!Has_alibi ? 0.0 : alibi_slope) {
, tot_seqlen_k(tot_seqlen_k)
, cp_world_size(cp_world_size)
, cp_rank(cp_rank) {
};
// Causal_mask: whether this particular iteration needs causal masking
template <bool Causal_mask=false, bool Is_even_MN=true, typename Engine, typename Layout>
template <bool Causal_mask=false, bool Is_even_MN=true, int Elem_per_thread = 4, typename Engine, typename Layout>
__forceinline__ __device__ void apply_mask(Tensor<Engine, Layout> &tensor_,
const int col_idx_offset_,
const int row_idx_offset,
const int warp_row_stride) {
static_assert(!(Causal_mask && Is_local), "Cannot be both causal and local");
static_assert(Layout::rank == 3, "Only support 3D Tensor");
static_assert(decltype(size<0>(tensor_))::value == 4, "First dimension must be 4");
static constexpr bool Need_masking = Has_alibi || Causal_mask || Is_local || !Is_even_MN;
// if (cute::thread0()) { printf("Has_alibi = %d, Causal_mask=%d, Is_local=%d, Is_even_MN = %d, Need_masking = %d\n", Has_alibi, Causal_mask, Is_local, Is_even_MN, Need_masking); }
static_assert(decltype(size<0>(tensor_))::value == Elem_per_thread, "The tensor_ first dimension not match the Elem_per_thread");
static constexpr bool Need_masking = Causal_mask || !Is_even_MN;
if constexpr (Need_masking) {
// Reshape tensor_ from (MMA=4, MMA_M, MMA_N) to (nrow=(2, MMA_M), ncol=(2, MMA_N))
Tensor tensor = make_tensor(tensor_.data(), flash::convert_layout_acc_rowcol(tensor_.layout()));
// Do we need both row and column indices, or just column incides?
static constexpr bool Col_idx_only = !(Has_alibi && !Is_causal) && !Is_local && !Causal_mask;
static constexpr bool Col_idx_only = !Causal_mask;
const int col_idx_offset = col_idx_offset_ + ((__lane_id() >> 4) << 2);
if constexpr (Col_idx_only) {
#pragma unroll
@ -155,9 +81,6 @@ struct Mask {
#pragma unroll
for (int mi = 0; mi < size<0>(tensor); ++mi) {
// No causal, no local
if constexpr (Has_alibi) {
tensor(mi, make_coord(j, nj)) += alibi_slope * col_idx;
}
if constexpr (!Is_even_MN) {
if (col_idx >= max_seqlen_k) { tensor(mi, make_coord(j, nj)) = -INFINITY; }
}
@ -171,38 +94,26 @@ struct Mask {
#pragma unroll
for (int i = 0; i < size<0, 0>(tensor); ++i) {
const int row_idx = row_idx_base + i * 16;
const int col_idx_limit_left = std::max(0, row_idx + max_seqlen_k - max_seqlen_q - window_size_left);
const int col_idx_limit_right = std::min(max_seqlen_k, row_idx / ngroups + 1 + max_seqlen_k - max_seqlen_q / ngroups + window_size_right);
const int col_idx_limit_right = std::min(tot_seqlen_k, row_idx / ngroups + 1 + tot_seqlen_k - max_seqlen_q / ngroups);
#pragma unroll
for (int nj = 0; nj < size<1, 1>(tensor); ++nj) {
const int col_idx_base = col_idx_offset + nj * 16;
#pragma unroll
for (int j = 0; j < size<1, 0>(tensor); ++j) {
const int col_idx = col_idx_base + j;
if constexpr (Has_alibi) {
if constexpr (Is_causal) {
tensor(make_coord(i, mi), make_coord(j, nj)) += alibi_slope * col_idx;
} else {
tensor(make_coord(i, mi), make_coord(j, nj)) -= alibi_slope * abs(row_idx + max_seqlen_k - max_seqlen_q - col_idx);
if constexpr (Is_enable_dcp) {
// casusal mask with dcp
const int actual_col_idx = col_idx * cp_world_size + cp_rank + 1;
// actual_col_idx start from 1 to tot_seqlen_k
if (actual_col_idx > col_idx_limit_right || col_idx >= max_seqlen_k) {
tensor(make_coord(i, mi), make_coord(j, nj)) = -INFINITY;
}
}
if constexpr (Causal_mask) {
} else {
// casusal mask without dcp
if (col_idx >= col_idx_limit_right) {
tensor(make_coord(i, mi), make_coord(j, nj)) = -INFINITY;
}
}
if constexpr (Is_local) {
if (col_idx >= col_idx_limit_right || col_idx < col_idx_limit_left) {
tensor(make_coord(i, mi), make_coord(j, nj)) = -INFINITY;
}
}
if constexpr (!Causal_mask && !Is_local && !Is_even_MN) {
// Causal and Local already handles MN masking
if (col_idx >= max_seqlen_k) {
tensor(make_coord(i, mi), make_coord(j, nj)) = -INFINITY;
}
}
}
}
}
@ -211,6 +122,32 @@ struct Mask {
}
};
template <int kBlockTopK, bool Is_even_MN=true, typename Engine, typename Layout>
__forceinline__ __device__ void apply_sparse_attn_mask(Tensor<Engine, Layout> &tensor_,
const int col_idx_offset_,
int32_t* indices_smem_ptr,
bool is_indices_all_valid) {
Tensor tensor = make_tensor(tensor_.data(), flash::convert_layout_acc_rowcol(tensor_.layout()));
const int col_idx_offset = col_idx_offset_ + ((__lane_id() >> 4) << 2);
#pragma unroll
for (int nj = 0; nj < size<1, 1>(tensor); ++nj) {
const int col_idx_base = col_idx_offset + nj * 16;
#pragma unroll
for (int j = 0; j < size<1, 0>(tensor); ++j) {
const int col_idx = col_idx_base + j;
// const bool invalid_flag = !is_indices_all_valid && (CHECK_BIT(is_valid_indices[(col_idx % kBlockN) >> 5], col_idx % kBlockN) == false);
const bool invalid_flag = !is_indices_all_valid && indices_smem_ptr[col_idx % kBlockTopK] < 0;
#pragma unroll
for (int mi = 0; mi < size<0>(tensor); ++mi) {
if constexpr (!Is_even_MN) {
if (col_idx >= max_seqlen_k) { tensor(mi, make_coord(j, nj)) = -INFINITY; }
}
if (invalid_flag) tensor(mi, make_coord(j, nj)) = -INFINITY;
}
}
}
};
};
} // namespace flash

View File

@ -1,529 +0,0 @@
// Adapted from Dao-AILab/flash-attention (https://github.com/Dao-AILab/flash-attention/tree/v2.6.3)
/******************************************************************************
* Copyright (c) 2024, Tri Dao.
******************************************************************************/
#pragma once
#include <cute/algorithm/copy.hpp>
#include "utils.h"
////////////////////////////////////////////////////////////////////////////////////////////////////
namespace flash {
using namespace cute;
////////////////////////////////////////////////////////////////////////////////////////////////////
template <bool Is_even_K=true, bool Clear_OOB_K=true,
typename Engine0, typename Layout0, typename Engine1, typename Layout1,
typename Engine2, typename Layout2, typename Engine3, typename Layout3>
__forceinline__ __device__ void copy_rotary_interleaved(Tensor<Engine0, Layout0> const &S,
Tensor<Engine1, Layout1> &D,
Tensor<Engine2, Layout2> const &Cos,
Tensor<Engine2, Layout2> const &Sin,
Tensor<Engine3, Layout3> const &identity_MN,
const int max_MN, const int min_MN,
const int dim, const int rotary_dim) {
CUTE_STATIC_ASSERT_V(rank(S) == Int<3>{});
CUTE_STATIC_ASSERT_V(rank(D) == Int<3>{});
CUTE_STATIC_ASSERT_V(size<0>(S) == size<0>(D)); // MMA
CUTE_STATIC_ASSERT_V(size<1>(S) == size<1>(D)); // MMA_M
CUTE_STATIC_ASSERT_V(size<2>(S) == size<2>(D)); // MMA_K
CUTE_STATIC_ASSERT_V(size<1>(S) == size<1>(Cos)); // MMA_M
CUTE_STATIC_ASSERT_V(size<2>(S) == size<2>(Cos)); // MMA_K
CUTE_STATIC_ASSERT_V(size<1>(S) == size<1>(Sin)); // MMA_M
CUTE_STATIC_ASSERT_V(size<2>(S) == size<2>(Sin)); // MMA_K
CUTE_STATIC_ASSERT_V(size<0>(Cos) == size<0>(Sin)); // MMA_K
static_assert(decltype(size<0>(S))::value == decltype(size<0>(Cos))::value * 2);
static_assert(decltype(size<0>(Cos))::value % 2 == 0); // Since we do fast conversion from fp16/bf16 to fp32
Tensor rCos = make_fragment_like(Cos);
Tensor rSin = make_fragment_like(Sin);
Tensor rS = make_fragment_like(S);
typedef __NATIVE_VECTOR__(2, float) Float2;
#pragma unroll
for (int m = 0; m < size<1>(S); ++m) {
if (get<0>(identity_MN(0, m, 0)) >= min_MN && get<0>(identity_MN(0, m, 0)) < max_MN) {
#pragma unroll
for (int k = 0; k < size<2>(S); ++k) {
if (Is_even_K || get<1>(identity_MN(0, 0, k)) < dim) {
cute::copy(S(_, m, k), rS(_, m, k));
if (get<1>(identity_MN(0, 0, k)) < rotary_dim) {
cute::copy(Cos(_, m, k), rCos(_, m, k));
cute::copy(Sin(_, m, k), rSin(_, m, k));
// Tensor S_fp32 = convert_type<float>(rS(_, m, k));
// Tensor cos_fp32 = convert_type<float>(rCos(_, m, k));
// Tensor sin_fp32 = convert_type<float>(rSin(_, m, k));
using T = typename Engine0::value_type;
using T_rotary = typename Engine2::value_type;
CONVERT_TENSOR_TYPE(T, float, rS(_, m, k), S_fp32)
CONVERT_TENSOR_TYPE(T_rotary, float, rCos(_, m, k), cos_fp32)
CONVERT_TENSOR_TYPE(T_rotary, float, rSin(_, m, k), sin_fp32)
#pragma unroll
for (int i = 0; i < size<0>(rS) / 2; ++i) {
Float2 x_vec = {S_fp32(2 * i), S_fp32(2 * i + 1)};
Float2 real_vec = {cos_fp32(i), sin_fp32(i)};
Float2 imag_vec = {sin_fp32(i), cos_fp32(i)};
Float2 beta_vec = {0.0f, 0.0f};
real_vec = __builtin_mxc_pk_fma_f32(x_vec, real_vec, beta_vec);
imag_vec = __builtin_mxc_pk_fma_f32(x_vec, imag_vec, beta_vec);
S_fp32(2 * i) = real_vec[0] - real_vec[1];
S_fp32(2 * i + 1) = imag_vec[0] + imag_vec[1];
//float real = S_fp32(2 * i) * cos_fp32(i) - S_fp32(2 * i + 1) * sin_fp32(i);
//float imag = S_fp32(2 * i) * sin_fp32(i) + S_fp32(2 * i + 1) * cos_fp32(i);
//S_fp32(2 * i) = real;
//S_fp32(2 * i + 1) = imag;
}
// Idk but I need to copy for the convert_type to work
Tensor S_fp32_copy = make_fragment_like(S_fp32);
cute::copy(S_fp32, S_fp32_copy);
//Tensor S_og_type = convert_type<T>(S_fp32_copy);
CONVERT_TENSOR_TYPE(float, T, S_fp32_copy, S_og_type)
cute::copy(S_og_type, rS(_, m, k));
}
cute::copy(rS(_, m, k), D(_, m, k));
} else if (Clear_OOB_K) {
cute::clear(D(_, m, k));
}
}
}
}
}
////////////////////////////////////////////////////////////////////////////////////////////////////
template <bool Is_even_K=true, bool Clear_OOB_K=true,
typename Engine0, typename Layout0, typename Engine1, typename Layout1,
typename Engine2, typename Layout2, typename Engine3, typename Layout3>
__forceinline__ __device__ void copy_rotary_contiguous(Tensor<Engine0, Layout0> const &S,
Tensor<Engine1, Layout1> &D,
Tensor<Engine2, Layout2> const &Cos,
Tensor<Engine2, Layout2> const &Sin,
Tensor<Engine3, Layout3> const &identity_MN,
const int max_MN, const int min_MN,
const int dim, const int rotary_dim) {
CUTE_STATIC_ASSERT_V(rank(S) == Int<3>{});
CUTE_STATIC_ASSERT_V(rank(D) == Int<3>{});
CUTE_STATIC_ASSERT_V(size<0>(S) == size<0>(D)); // MMA
CUTE_STATIC_ASSERT_V(size<1>(S) == size<1>(D)); // MMA_M
CUTE_STATIC_ASSERT_V(size<2>(S) == size<2>(D)); // MMA_K
CUTE_STATIC_ASSERT_V(size<1>(S) == size<1>(Cos)); // MMA_M
CUTE_STATIC_ASSERT_V(size<2>(S) == size<2>(Cos)); // MMA_K
CUTE_STATIC_ASSERT_V(size<1>(S) == size<1>(Sin)); // MMA_M
CUTE_STATIC_ASSERT_V(size<2>(S) == size<2>(Sin)); // MMA_K
CUTE_STATIC_ASSERT_V(size<0>(S) == size<0>(Cos)); // MMA
CUTE_STATIC_ASSERT_V(size<0>(Cos) == size<0>(Sin));
static_assert(decltype(size<0>(Cos))::value % 2 == 0); // Since we do fast conversion from fp16/bf16 to fp32
Tensor rCos = make_fragment_like(Cos);
Tensor rSin = make_fragment_like(Sin);
Tensor rS = make_fragment_like(S);
Tensor rS_other = make_fragment_like(rS(_, 0, 0));
typedef __NATIVE_VECTOR__(2, float) Float2;
Float2 beta_vec = {0.0f, 0.0f};
const int rotary_dim_half = rotary_dim >> 1;
#pragma unroll
for (int m = 0; m < size<1>(S); ++m) {
if (get<0>(identity_MN(0, m, 0)) >= min_MN && get<0>(identity_MN(0, m, 0)) < max_MN) {
#pragma unroll
for (int k = 0; k < size<2>(S); ++k) {
if (Is_even_K || get<1>(identity_MN(0, 0, k)) < dim) {
cute::copy(S(_, m, k), rS(_, m, k));
if (get<1>(identity_MN(0, 0, k)) < rotary_dim) {
const bool is_left = get<1>(identity_MN(0, 0, k)) < rotary_dim_half;
Tensor gS_other = make_tensor(S(_, m, k).data() + (is_left ? rotary_dim_half : -rotary_dim_half), S(_, m, k).layout());
cute::copy(gS_other, rS_other);
// if (cute::thread0()) { print_tensor(rS(_, m, k)); print_tensor(rS_other); }
Tensor gCos = make_tensor(Cos(_, m, k).data() + (is_left ? 0 : -rotary_dim_half), Cos(_, m, k).layout());
Tensor gSin = make_tensor(Sin(_, m, k).data() + (is_left ? 0 : -rotary_dim_half), Sin(_, m, k).layout());
cute::copy(gCos, rCos(_, m, k));
cute::copy(gSin, rSin(_, m, k));
// if (cute::thread0()) { print_tensor(rCos(_, m, k)); print_tensor(rSin(_, m, k)); }
// Tensor S_fp32 = convert_type<float>(rS(_, m, k));
// Tensor S_other_fp32 = convert_type<float>(rS_other);
// Tensor cos_fp32 = convert_type<float>(rCos(_, m, k));
// Tensor sin_fp32 = convert_type<float>(rSin(_, m, k));
using T = typename Engine0::value_type;
using T_rotary = typename Engine2::value_type;
CONVERT_TENSOR_TYPE(T, float, rS(_,m,k), S_fp32)
CONVERT_TENSOR_TYPE(T, float, rS_other, S_other_fp32)
CONVERT_TENSOR_TYPE(T_rotary, float, rCos(_, m, k), cos_fp32)
CONVERT_TENSOR_TYPE(T_rotary, float, rSin(_, m, k), sin_fp32)
#pragma unroll
for (int i = 0; i < size<0>(rS); ++i) {
//S_fp32(i) = S_fp32(i) * cos_fp32(i) + S_other_fp32(i) * (is_left ? -sin_fp32(i) : sin_fp32(i));
Float2 x_vec = {S_fp32(i), S_other_fp32(i)};
Float2 alpha_vec = {cos_fp32(i), is_left ? -sin_fp32(i) : sin_fp32(i)};
Float2 y_vec = __builtin_mxc_pk_fma_f32(x_vec, alpha_vec, beta_vec);
S_fp32(i) = y_vec[0] + y_vec[1];
}
// Idk but I need to copy for the convert_type to work
Tensor S_fp32_copy = make_fragment_like(S_fp32);
cute::copy(S_fp32, S_fp32_copy);
//Tensor S_og_type = convert_type<T>(S_fp32_copy);
CONVERT_TENSOR_TYPE(float, T, S_fp32_copy, S_og_type)
cute::copy(S_og_type, rS(_, m, k));
// if (cute::thread0()) { print_tensor(rS(_, m, k)); }
}
cute::copy(rS(_, m, k), D(_, m, k));
} else if (Clear_OOB_K) {
cute::clear(D(_, m, k));
}
}
}
}
}
////////////////////////////////////////////////////////////////////////////////////////////////////
template <bool Is_even_K=true,
typename Engine0, typename Layout0, typename Engine1, typename Layout1,
typename Engine2, typename Layout2>
__forceinline__ __device__ void copy_rotary_interleaved_to_reg(Tensor<Engine0, Layout0> const &S,
uint32_t *D_ptr,
Tensor<Engine1, Layout1> const &Cos,
Tensor<Engine1, Layout1> const &Sin,
Tensor<Engine2, Layout2> const &identity_MN,
const int max_MN, const int min_MN,
const int dim, const int rotary_dim) {
CUTE_STATIC_ASSERT_V(rank(S) == Int<3>{});
CUTE_STATIC_ASSERT_V(size<1>(S) == size<1>(Cos)); // MMA_M
CUTE_STATIC_ASSERT_V(size<2>(S) == size<2>(Cos)); // MMA_K
CUTE_STATIC_ASSERT_V(size<1>(S) == size<1>(Sin)); // MMA_M
CUTE_STATIC_ASSERT_V(size<2>(S) == size<2>(Sin)); // MMA_K
CUTE_STATIC_ASSERT_V(size<0>(Cos) == size<0>(Sin)); // MMA_K
static_assert(decltype(size<0>(S))::value == decltype(size<0>(Cos))::value * 2);
static_assert(decltype(size<0>(Cos))::value % 2 == 0); // Since we do fast conversion from fp16/bf16 to fp32
Tensor rCos = make_fragment_like(Cos);
Tensor rSin = make_fragment_like(Sin);
Tensor rS = make_fragment_like(S);
typedef __NATIVE_VECTOR__(2, float) Float2;
typedef __NATIVE_VECTOR__(4, int) VecTypeB128;
typedef __NATIVE_VECTOR__(2, int) VecTypeB64;
#pragma unroll
for (int m = 0; m < size<1>(S); ++m) {
bool row_mask = get<0>(identity_MN(0, m, 0)) >= min_MN && get<0>(identity_MN(0, m, 0)) < max_MN;
#pragma unroll
for (int k = 0; k < size<2>(S); ++k) {
bool mask = row_mask && (Is_even_K || get<1>(identity_MN(0, 0, k)) < dim);
auto S_ptr = (VecTypeB128 *)(S(_, m, k).data().ptr_);
auto rS_ptr = (VecTypeB128 *)(rS(_, m, k).data());
rS_ptr[0] = __builtin_mxc_ldg_b128_predicator(S_ptr, 0, true, true, false, false,
mask, 1, MACA_ICMP_EQ);
bool rotary_mask = mask && (get<1>(identity_MN(0, 0, k)) < rotary_dim);
auto gCos_ptr = (VecTypeB64 *)(Cos(_, m, k).data().ptr_);
auto gSin_ptr = (VecTypeB64 *)(Sin(_, m, k).data().ptr_);
auto rCos_ptr = (VecTypeB64 *)(rCos(_, m, k).data());
auto rSin_ptr = (VecTypeB64 *)(rSin(_, m, k).data());
rCos_ptr[0] = __builtin_mxc_ldg_b64_predicator(gCos_ptr, 0, true, true, false, false,
rotary_mask, 1, MACA_ICMP_EQ);
rSin_ptr[0] = __builtin_mxc_ldg_b64_predicator(gSin_ptr, 0, true, true, false, false,
rotary_mask, 1, MACA_ICMP_EQ);
if (rotary_mask) {
using T = typename Engine0::value_type;
using T_rotary = typename Engine1::value_type;
CONVERT_TENSOR_TYPE(T, float, rS(_, m, k), S_fp32)
CONVERT_TENSOR_TYPE(T_rotary, float, rCos(_, m, k), cos_fp32)
CONVERT_TENSOR_TYPE(T_rotary, float, rSin(_, m, k), sin_fp32)
#pragma unroll
for (int i = 0; i < size<0>(rS) / 2; ++i) {
Float2 x_vec = {S_fp32(2 * i), S_fp32(2 * i + 1)};
Float2 real_vec = {cos_fp32(i), sin_fp32(i)};
Float2 imag_vec = {sin_fp32(i), cos_fp32(i)};
Float2 beta_vec = {0.0f, 0.0f};
real_vec = __builtin_mxc_pk_fma_f32(x_vec, real_vec, beta_vec);
imag_vec = __builtin_mxc_pk_fma_f32(x_vec, imag_vec, beta_vec);
S_fp32(2 * i) = real_vec[0] - real_vec[1];
S_fp32(2 * i + 1) = imag_vec[0] + imag_vec[1];
}
// Idk but I need to copy for the convert_type to work
Tensor S_fp32_copy = make_fragment_like(S_fp32);
cute::copy(S_fp32, S_fp32_copy);
//Tensor S_og_type = convert_type<T>(S_fp32_copy);
CONVERT_TENSOR_TYPE(float, T, S_fp32_copy, S_og_type)
cute::copy(S_og_type, rS(_, m, k));
}
const int idx = (m * size<2>(S) + k) << 2;
auto D = (VecTypeB128 *)(D_ptr + idx);
D[0] = rS_ptr[0];
}
}
}
////////////////////////////////////////////////////////////////////////////////////////////////////
template <bool Is_even_K=true,
typename Engine0, typename Layout0, typename Engine1, typename Layout1,
typename Engine2, typename Layout2>
__forceinline__ __device__ void copy_rotary_contiguous_to_reg(Tensor<Engine0, Layout0> const &S,
uint32_t *D_ptr,
Tensor<Engine1, Layout1> const &Cos,
Tensor<Engine1, Layout1> const &Sin,
Tensor<Engine2, Layout2> const &identity_MN,
const int max_MN, const int min_MN,
const int dim, const int rotary_dim) {
CUTE_STATIC_ASSERT_V(rank(S) == Int<3>{});
CUTE_STATIC_ASSERT_V(size<1>(S) == size<1>(Cos)); // MMA_M
CUTE_STATIC_ASSERT_V(size<2>(S) == size<2>(Cos)); // MMA_K
CUTE_STATIC_ASSERT_V(size<1>(S) == size<1>(Sin)); // MMA_M
CUTE_STATIC_ASSERT_V(size<2>(S) == size<2>(Sin)); // MMA_K
CUTE_STATIC_ASSERT_V(size<0>(S) == size<0>(Cos)); // MMA
CUTE_STATIC_ASSERT_V(size<0>(Cos) == size<0>(Sin));
static_assert(decltype(size<0>(Cos))::value % 2 == 0); // Since we do fast conversion from fp16/bf16 to fp32
Tensor rCos = make_fragment_like(Cos);
Tensor rSin = make_fragment_like(Sin);
Tensor rS = make_fragment_like(S);
Tensor rS_other = make_fragment_like(rS(_, 0, 0));
typedef __NATIVE_VECTOR__(2, float) Float2;
Float2 beta_vec = {0.0f, 0.0f};
const int rotary_dim_half = rotary_dim >> 1;
typedef __NATIVE_VECTOR__(4, int) VecType;
#pragma unroll
for (int m = 0; m < size<1>(S); ++m) {
bool row_mask = get<0>(identity_MN(0, m, 0)) >= min_MN && get<0>(identity_MN(0, m, 0)) < max_MN;
#pragma unroll
for (int k = 0; k < size<2>(S); ++k) {
bool mask = row_mask && (Is_even_K || get<1>(identity_MN(0, 0, k)) < dim);
auto S_ptr = (VecType *)(S(_, m, k).data().ptr_);
auto rS_ptr = (VecType *)(rS(_, m, k).data());
rS_ptr[0] = __builtin_mxc_ldg_b128_predicator(S_ptr, 0, true, true, false, false,
mask, 1, MACA_ICMP_EQ);
bool rotary_mask = mask && (get<1>(identity_MN(0, 0, k)) < rotary_dim);
const bool is_left = get<1>(identity_MN(0, 0, k)) < rotary_dim_half;
auto gS_ptr = (VecType *)(S(_, m, k).data().ptr_ + (is_left ? rotary_dim_half : -rotary_dim_half));
auto rS_other_ptr = (VecType *)(rS_other.data());
rS_other_ptr[0] = __builtin_mxc_ldg_b128_predicator(gS_ptr, 0, true, true, false, false,
rotary_mask, 1, MACA_ICMP_EQ);
auto gCos_ptr = (VecType *)(Cos(_, m, k).data().ptr_ + (is_left ? 0 : -rotary_dim_half));
auto gSin_ptr = (VecType *)(Sin(_, m, k).data().ptr_ + (is_left ? 0 : -rotary_dim_half));
auto rCos_ptr = (VecType *)(rCos(_, m, k).data());
auto rSin_ptr = (VecType *)(rSin(_, m, k).data());
rCos_ptr[0] = __builtin_mxc_ldg_b128_predicator(gCos_ptr, 0, true, true, false, false,
rotary_mask, 1, MACA_ICMP_EQ);
rSin_ptr[0] = __builtin_mxc_ldg_b128_predicator(gSin_ptr, 0, true, true, false, false,
rotary_mask, 1, MACA_ICMP_EQ);
if (rotary_mask) {
using T = typename Engine0::value_type;
using T_rotary = typename Engine1::value_type;
CONVERT_TENSOR_TYPE(T, float, rS(_, m, k), S_fp32)
CONVERT_TENSOR_TYPE(T, float, rS_other, S_other_fp32)
CONVERT_TENSOR_TYPE(T_rotary, float, rCos(_, m, k), cos_fp32)
CONVERT_TENSOR_TYPE(T_rotary, float, rSin(_, m, k), sin_fp32)
#pragma unroll
for (int i = 0; i < size<0>(rS); ++i) {
// S_fp32(i) = S_fp32(i) * cos_fp32(i) + S_other_fp32(i) * (is_left ? -sin_fp32(i) : sin_fp32(i));
Float2 x_vec = {S_fp32(i), S_other_fp32(i)};
Float2 alpha_vec = {cos_fp32(i), is_left ? -sin_fp32(i) : sin_fp32(i)};
Float2 y_vec = __builtin_mxc_pk_fma_f32(x_vec, alpha_vec, beta_vec);
S_fp32(i) = y_vec[0] + y_vec[1];
}
// Idk but I need to copy for the convert_type to work
Tensor S_fp32_copy = make_fragment_like(S_fp32);
cute::copy(S_fp32, S_fp32_copy);
//Tensor S_og_type = convert_type<T>(S_fp32_copy);
CONVERT_TENSOR_TYPE(float, T, S_fp32_copy, S_og_type)
cute::copy(S_og_type, rS(_, m, k));
}
const int idx = (m * size<2>(S) + k) * 4;
auto D = (VecType *)(D_ptr + idx);
D[0] = rS_ptr[0];
}
}
}
////////////////////////////////////////////////////////////////////////////////////////////////////
template <bool Is_even_K=true,
typename Engine0, typename Layout0, typename Engine1, typename Layout1,
typename Engine2, typename Layout2, typename Engine3, typename Layout3>
__forceinline__ __device__ void copy_rotary_interleaved_to_global(Tensor<Engine0, Layout0> const &S,
Tensor<Engine1, Layout1> &D,
Tensor<Engine2, Layout2> const &Cos,
Tensor<Engine2, Layout2> const &Sin,
Tensor<Engine3, Layout3> const &identity_MN,
const int max_MN, const int min_MN,
const int dim, const int rotary_dim) {
CUTE_STATIC_ASSERT_V(rank(S) == Int<3>{});
CUTE_STATIC_ASSERT_V(rank(D) == Int<3>{});
CUTE_STATIC_ASSERT_V(size<0>(S) == size<0>(D)); // MMA
CUTE_STATIC_ASSERT_V(size<1>(S) == size<1>(D)); // MMA_M
CUTE_STATIC_ASSERT_V(size<2>(S) == size<2>(D)); // MMA_K
CUTE_STATIC_ASSERT_V(size<1>(S) == size<1>(Cos)); // MMA_M
CUTE_STATIC_ASSERT_V(size<2>(S) == size<2>(Cos)); // MMA_K
CUTE_STATIC_ASSERT_V(size<1>(S) == size<1>(Sin)); // MMA_M
CUTE_STATIC_ASSERT_V(size<2>(S) == size<2>(Sin)); // MMA_K
CUTE_STATIC_ASSERT_V(size<0>(Cos) == size<0>(Sin)); // MMA_K
static_assert(decltype(size<0>(S))::value == decltype(size<0>(Cos))::value * 2);
static_assert(decltype(size<0>(Cos))::value % 2 == 0); // Since we do fast conversion from fp16/bf16 to fp32
Tensor rCos = make_fragment_like(Cos);
Tensor rSin = make_fragment_like(Sin);
Tensor rS = make_fragment_like(S);
typedef __NATIVE_VECTOR__(2, float) Float2;
typedef __NATIVE_VECTOR__(4, int) VecTypeB128;
typedef __NATIVE_VECTOR__(2, int) VecTypeB64;
#pragma unroll
for (int m = 0; m < size<1>(S); ++m) {
bool row_mask = get<0>(identity_MN(0, m, 0)) >= min_MN && get<0>(identity_MN(0, m, 0)) < max_MN;
#pragma unroll
for (int k = 0; k < size<2>(S); ++k) {
bool mask = row_mask && (Is_even_K || get<1>(identity_MN(0, 0, k)) < dim);
auto S_ptr = (VecTypeB128 *)(S(_, m, k).data().ptr_);
auto rS_ptr = (VecTypeB128 *)(rS(_, m, k).data());
rS_ptr[0] = __builtin_mxc_ldg_b128_predicator(S_ptr, 0, true, true, false, false,
mask, 1, MACA_ICMP_EQ);
bool rotary_mask = mask && (get<1>(identity_MN(0, 0, k)) < rotary_dim);
auto gCos_ptr = (VecTypeB64 *)(Cos(_, m, k).data().ptr_);
auto gSin_ptr = (VecTypeB64 *)(Sin(_, m, k).data().ptr_);
auto rCos_ptr = (VecTypeB64 *)(rCos(_, m, k).data());
auto rSin_ptr = (VecTypeB64 *)(rSin(_, m, k).data());
rCos_ptr[0] = __builtin_mxc_ldg_b64_predicator(gCos_ptr, 0, true, true, false, false,
rotary_mask, 1, MACA_ICMP_EQ);
rSin_ptr[0] = __builtin_mxc_ldg_b64_predicator(gSin_ptr, 0, true, true, false, false,
rotary_mask, 1, MACA_ICMP_EQ);
if (rotary_mask) {
using T = typename Engine0::value_type;
using T_rotary = typename Engine1::value_type;
CONVERT_TENSOR_TYPE(T, float, rS(_, m, k), S_fp32)
CONVERT_TENSOR_TYPE(T_rotary, float, rCos(_, m, k), cos_fp32)
CONVERT_TENSOR_TYPE(T_rotary, float, rSin(_, m, k), sin_fp32)
#pragma unroll
for (int i = 0; i < size<0>(rS) / 2; ++i) {
Float2 x_vec = {S_fp32(2 * i), S_fp32(2 * i + 1)};
Float2 real_vec = {cos_fp32(i), sin_fp32(i)};
Float2 imag_vec = {sin_fp32(i), cos_fp32(i)};
Float2 beta_vec = {0.0f, 0.0f};
real_vec = __builtin_mxc_pk_fma_f32(x_vec, real_vec, beta_vec);
imag_vec = __builtin_mxc_pk_fma_f32(x_vec, imag_vec, beta_vec);
S_fp32(2 * i) = real_vec[0] - real_vec[1];
S_fp32(2 * i + 1) = imag_vec[0] + imag_vec[1];
}
// Idk but I need to copy for the convert_type to work
Tensor S_fp32_copy = make_fragment_like(S_fp32);
cute::copy(S_fp32, S_fp32_copy);
//Tensor S_og_type = convert_type<T>(S_fp32_copy);
CONVERT_TENSOR_TYPE(float, T, S_fp32_copy, S_og_type)
cute::copy(S_og_type, rS(_, m, k));
}
auto D_ptr = (VecTypeB128 *)(D(_, m, k).data().ptr_);
__builtin_mxc_stg_b128_predicator(D_ptr, 0, rS_ptr[0], true, false, true, mask, 1, MACA_ICMP_EQ);
}
}
}
////////////////////////////////////////////////////////////////////////////////////////////////////
template <bool Is_even_K=true,
typename Engine0, typename Layout0, typename Engine1, typename Layout1,
typename Engine2, typename Layout2, typename Engine3, typename Layout3>
__forceinline__ __device__ void copy_rotary_contiguous_to_global(Tensor<Engine0, Layout0> const &S,
Tensor<Engine1, Layout1> &D,
Tensor<Engine2, Layout2> const &Cos,
Tensor<Engine2, Layout2> const &Sin,
Tensor<Engine3, Layout3> const &identity_MN,
const int max_MN, const int min_MN,
const int dim, const int rotary_dim) {
CUTE_STATIC_ASSERT_V(rank(S) == Int<3>{});
CUTE_STATIC_ASSERT_V(rank(D) == Int<3>{});
CUTE_STATIC_ASSERT_V(size<0>(S) == size<0>(D)); // MMA
CUTE_STATIC_ASSERT_V(size<1>(S) == size<1>(D)); // MMA_M
CUTE_STATIC_ASSERT_V(size<2>(S) == size<2>(D)); // MMA_K
CUTE_STATIC_ASSERT_V(size<1>(S) == size<1>(Cos)); // MMA_M
CUTE_STATIC_ASSERT_V(size<2>(S) == size<2>(Cos)); // MMA_K
CUTE_STATIC_ASSERT_V(size<1>(S) == size<1>(Sin)); // MMA_M
CUTE_STATIC_ASSERT_V(size<2>(S) == size<2>(Sin)); // MMA_K
CUTE_STATIC_ASSERT_V(size<0>(S) == size<0>(Cos)); // MMA
CUTE_STATIC_ASSERT_V(size<0>(Cos) == size<0>(Sin));
static_assert(decltype(size<0>(Cos))::value % 2 == 0); // Since we do fast conversion from fp16/bf16 to fp32
Tensor rCos = make_fragment_like(Cos);
Tensor rSin = make_fragment_like(Sin);
Tensor rS = make_fragment_like(S);
Tensor rS_other = make_fragment_like(rS(_, 0, 0));
typedef __NATIVE_VECTOR__(2, float) Float2;
Float2 beta_vec = {0.0f, 0.0f};
const int rotary_dim_half = rotary_dim >> 1;
typedef __NATIVE_VECTOR__(4, int) VecType;
#pragma unroll
for (int m = 0; m < size<1>(S); ++m) {
bool row_mask = get<0>(identity_MN(0, m, 0)) >= min_MN && get<0>(identity_MN(0, m, 0)) < max_MN;
#pragma unroll
for (int k = 0; k < size<2>(S); ++k) {
bool mask = row_mask && (Is_even_K || get<1>(identity_MN(0, 0, k)) < dim);
auto S_ptr = (VecType *)(S(_, m, k).data().ptr_);
auto rS_ptr = (VecType *)(rS(_, m, k).data());
rS_ptr[0] = __builtin_mxc_ldg_b128_predicator(S_ptr, 0, true, true, false, false,
mask, 1, MACA_ICMP_EQ);
bool rotary_mask = mask && (get<1>(identity_MN(0, 0, k)) < rotary_dim);
const bool is_left = get<1>(identity_MN(0, 0, k)) < rotary_dim_half;
auto gS_ptr = (VecType *)(S(_, m, k).data().ptr_ + (is_left ? rotary_dim_half : -rotary_dim_half));
auto rS_other_ptr = (VecType *)(rS_other.data());
rS_other_ptr[0] = __builtin_mxc_ldg_b128_predicator(gS_ptr, 0, true, true, false, false,
rotary_mask, 1, MACA_ICMP_EQ);
auto gCos_ptr = (VecType *)(Cos(_, m, k).data().ptr_ + (is_left ? 0 : -rotary_dim_half));
auto gSin_ptr = (VecType *)(Sin(_, m, k).data().ptr_ + (is_left ? 0 : -rotary_dim_half));
auto rCos_ptr = (VecType *)(rCos(_, m, k).data());
auto rSin_ptr = (VecType *)(rSin(_, m, k).data());
rCos_ptr[0] = __builtin_mxc_ldg_b128_predicator(gCos_ptr, 0, true, true, false, false,
rotary_mask, 1, MACA_ICMP_EQ);
rSin_ptr[0] = __builtin_mxc_ldg_b128_predicator(gSin_ptr, 0, true, true, false, false,
rotary_mask, 1, MACA_ICMP_EQ);
if (rotary_mask) {
using T = typename Engine0::value_type;
using T_rotary = typename Engine1::value_type;
CONVERT_TENSOR_TYPE(T, float, rS(_, m, k), S_fp32)
CONVERT_TENSOR_TYPE(T, float, rS_other, S_other_fp32)
CONVERT_TENSOR_TYPE(T_rotary, float, rCos(_, m, k), cos_fp32)
CONVERT_TENSOR_TYPE(T_rotary, float, rSin(_, m, k), sin_fp32)
#pragma unroll
for (int i = 0; i < size<0>(rS); ++i) {
// S_fp32(i) = S_fp32(i) * cos_fp32(i) + S_other_fp32(i) * (is_left ? -sin_fp32(i) : sin_fp32(i));
Float2 x_vec = {S_fp32(i), S_other_fp32(i)};
Float2 alpha_vec = {cos_fp32(i), is_left ? -sin_fp32(i) : sin_fp32(i)};
Float2 y_vec = __builtin_mxc_pk_fma_f32(x_vec, alpha_vec, beta_vec);
S_fp32(i) = y_vec[0] + y_vec[1];
}
// Idk but I need to copy for the convert_type to work
Tensor S_fp32_copy = make_fragment_like(S_fp32);
cute::copy(S_fp32, S_fp32_copy);
//Tensor S_og_type = convert_type<T>(S_fp32_copy);
CONVERT_TENSOR_TYPE(float, T, S_fp32_copy, S_og_type)
cute::copy(S_og_type, rS(_, m, k));
}
auto D_ptr = (VecType *)(D(_, m, k).data().ptr_);
__builtin_mxc_stg_b128_predicator(D_ptr, 0, rS_ptr[0], true, false, true, mask, 1, MACA_ICMP_EQ);
}
}
}
////////////////////////////////////////////////////////////////////////////////////////////////////
} // namespace flash

View File

@ -12,7 +12,6 @@
#include <mctlass/numeric_types.h>
#include "philox.cuh"
#include "utils.h"
namespace flash {
@ -41,7 +40,7 @@ __device__ __forceinline__ void quad_allreduce_(Tensor<Engine0, Layout0> &dst, T
CUTE_STATIC_ASSERT_V(size(dst) == size(src));
#pragma unroll
for (int i = 0; i < size(dst); i++){
dst(i) = Allreduce<64>::run(src(i), op);
dst(i) = Partialreduce::run(src(i), op);
}
}
@ -220,6 +219,50 @@ struct Softmax {
#pragma unroll
for (int mi = 0; mi < size<0>(scores); mi++) {
if constexpr(AddVec) {
Float2 x_vec = {row_sum(mi), 0.0f};
Float2 scale_vec = {1.0f, 1.0f};
#pragma unroll
for (int ni = 0; ni < size<1>(scores); ni += 2) {
Float2 beta_vec = {scores(mi, ni), scores(mi, ni + 1)};
x_vec = __builtin_mxc_pk_fma_f32(x_vec, scale_vec, beta_vec);
}
row_sum(mi) = x_vec[0] + x_vec[1];
}
else {
#pragma unroll
for (int ni = 0; ni < size<1>(scores); ni++) {
row_sum(mi) += scores(mi, ni);
}
}
}
}
};
template<bool Is_first, bool Check_inf=false, bool Syncthreads=false, bool AddVec=false, typename Tensor0, typename Tensor1, typename Tensor2>
__forceinline__ __device__ void softmax_rescale_o(Tensor0 &acc_s, Tensor1 &acc_o, Tensor2 &sRowMax, float softmax_scale_log2) {
// Reshape acc_s from (MMA=4, MMA_M, MMA_N) to (nrow=(2, MMA_M), ncol=(2, MMA_N))
Tensor scores = make_tensor(acc_s.data(), flash::convert_layout_acc_rowcol(acc_s.layout()));
MaxOp<float> max_op;
static_assert(decltype(size<0>(scores))::value == kNRows);
static_assert(decltype(size<1>(scores))::value % 2 == 0);
typedef __NATIVE_VECTOR__(2, float) Float2;
const int tidx = threadIdx.x;
const int wave_idx = tidx / 64;
const int lane_idx = tidx % 64;
const int wave_group_idx = wave_idx / 4;
const int row_offset = wave_idx % 4 * 16 + lane_idx % 16;
if constexpr (Is_first) {
flash::template thread_reduce_</*zero_init=*/true>(scores, row_max, max_op);
flash::template quad_allreduce_(row_max, row_max, max_op);
if (lane_idx / 16 == 0) {
sRowMax(wave_group_idx, row_offset) = row_max(0); //sts row_max
}
flash::sync_threads();
row_max(0) = max(row_max(0), sRowMax(wave_group_idx ^ 1, row_offset)); //lds row_max
flash::scale_apply_exp2(scores, row_max, softmax_scale_log2);
if constexpr(AddVec) {
#pragma unroll
for (int mi = 0; mi < size<0>(scores); mi++) {
Float2 x_vec = { 0.0f, 0.0f};
Float2 scale_vec = {1.0f, 1.0f};
#pragma unroll
@ -227,7 +270,159 @@ struct Softmax {
Float2 beta_vec = {scores(mi, ni), scores(mi, ni + 1)};
x_vec = __builtin_mxc_pk_fma_f32(x_vec, scale_vec, beta_vec);
}
row_sum(mi) += x_vec[0] + x_vec[1];
row_sum(mi) = x_vec[0] + x_vec[1];
}
} else {
SumOp<float> sum_op;
flash::thread_reduce_</*zero_init=*/true>(scores, row_sum, sum_op);
}
} else {
Tensor scores_max_prev = make_fragment_like(row_max);
cute::copy(row_max, scores_max_prev);
flash::template thread_reduce_</*zero_init=*/false>(scores, row_max, max_op);
flash::template quad_allreduce_(row_max, row_max, max_op);
if (lane_idx / 16 == 0) {
sRowMax(wave_group_idx, row_offset) = row_max(0); //sts row_max
}
flash::sync_threads();
row_max(0) = max(row_max(0), sRowMax(wave_group_idx ^ 1, row_offset)); //lds row_max
// Reshape acc_o from (MMA=4, MMA_M, MMA_K) to (nrow=(2, MMA_M), ncol=(2, MMA_K))
Tensor acc_o_rowcol = make_tensor(acc_o.data(), flash::convert_layout_acc_rowcol(acc_o.layout()));
static_assert(decltype(size<0>(acc_o_rowcol))::value == kNRows);
static_assert(decltype(size<1>(acc_o_rowcol))::value % 2 == 0);
#pragma unroll
for (int mi = 0; mi < size(row_max); ++mi) {
float scores_max_cur = !Check_inf
? row_max(mi)
: (row_max(mi) == -INFINITY ? 0.0f : row_max(mi));
float scores_scale = __builtin_exp2f((scores_max_prev(mi) - scores_max_cur) * softmax_scale_log2);
row_sum(mi) *= scores_scale;
Float2 scale_vec = {scores_scale , scores_scale};
Float2 beta_vec = {0.0f, 0.0f};
#pragma unroll
for (int ni = 0; ni < size<1>(acc_o_rowcol); ni += 2) {
Float2 x_vec = {acc_o_rowcol(mi, ni), acc_o_rowcol(mi, ni + 1)};
x_vec = __builtin_mxc_pk_fma_f32(x_vec, scale_vec, beta_vec);
acc_o_rowcol(mi, ni) = x_vec[0];
acc_o_rowcol(mi, ni + 1) = x_vec[1];
}
}
flash::scale_apply_exp2(scores, row_max, softmax_scale_log2);
#pragma unroll
for (int mi = 0; mi < size<0>(scores); mi++) {
if constexpr(AddVec) {
Float2 x_vec = {row_sum(mi), 0.0f};
Float2 scale_vec = {1.0f, 1.0f};
#pragma unroll
for (int ni = 0; ni < size<1>(scores); ni += 2) {
Float2 beta_vec = {scores(mi, ni), scores(mi, ni + 1)};
x_vec = __builtin_mxc_pk_fma_f32(x_vec, scale_vec, beta_vec);
}
row_sum(mi) = x_vec[0] + x_vec[1];
}
else {
#pragma unroll
for (int ni = 0; ni < size<1>(scores); ni++) {
row_sum(mi) += scores(mi, ni);
}
}
}
}
}
template<bool Is_first, typename Tensor0, typename Tensor1, typename Tensor2>
__forceinline__ __device__ void get_row_max(Tensor0 &acc_s, Tensor1 &scores_max_prev, Tensor2 &sRowMax,float softmax_scale_log2) {
Tensor scores = make_tensor(acc_s.data(), flash::convert_layout_acc_rowcol(acc_s.layout()));
MaxOp<float> max_op;
static_assert(decltype(size<0>(scores))::value == kNRows);
static_assert(decltype(size<1>(scores))::value % 2 == 0);
const int tidx = threadIdx.x;
const int wave_idx = tidx / 64;
const int lane_idx = tidx % 64;
const int wave_group_idx = wave_idx / 4;
const int row_offset = wave_idx % 4 * 16 + lane_idx % 16;
if constexpr (Is_first) {
flash::template thread_reduce_</*zero_init=*/true>(scores, row_max, max_op);
flash::template quad_allreduce_(row_max, row_max, max_op);
if (lane_idx / 16 == 0) {
sRowMax(wave_group_idx, row_offset) = row_max(0); //sts row_max
}
flash::sync_threads();
row_max(0) = max(row_max(0), sRowMax(wave_group_idx ^ 1, row_offset)); //lds row_max
} else {
cute::copy(row_max, scores_max_prev);
flash::template thread_reduce_</*zero_init=*/false>(scores, row_max, max_op);
flash::template quad_allreduce_(row_max, row_max, max_op);
if (lane_idx / 16 == 0) {
sRowMax(wave_group_idx, row_offset) = row_max(0); //sts row_max
}
flash::sync_threads();
row_max(0) = max(row_max(0), sRowMax(wave_group_idx ^ 1, row_offset)); //lds row_max
}
}
template<bool Is_first, bool Check_inf=false, bool AddVec=false, typename Tensor0, typename Tensor1, typename Tensor2>
__forceinline__ __device__ void softmax_rescale_o_without_row_max(Tensor0 &acc_s, Tensor1 &acc_o, Tensor2 &scores_max_prev, float softmax_scale_log2) {
// Reshape acc_s from (MMA=4, MMA_M, MMA_N) to (nrow=(2, MMA_M), ncol=(2, MMA_N))
Tensor scores = make_tensor(acc_s.data(), flash::convert_layout_acc_rowcol(acc_s.layout()));
static_assert(decltype(size<0>(scores))::value == kNRows);
static_assert(decltype(size<1>(scores))::value % 2 == 0);
typedef __NATIVE_VECTOR__(2, float) Float2;
if constexpr (Is_first) {
flash::scale_apply_exp2(scores, row_max, softmax_scale_log2);
if constexpr (AddVec) {
#pragma unroll
for (int mi = 0; mi < size<0>(scores); mi++) {
Float2 x_vec = {0.0f, 0.0f};
Float2 scale_vec = {1.0f, 1.0f};
#pragma unroll
for (int ni = 0; ni < size<1>(scores); ni += 2) {
Float2 beta_vec = {scores(mi, ni), scores(mi, ni + 1)};
x_vec = __builtin_mxc_pk_fma_f32(x_vec, scale_vec, beta_vec);
}
row_sum(mi) = x_vec[0] + x_vec[1];
}
}
else {
SumOp<float> sum_op;
flash::thread_reduce_</*zero_init=*/true>(scores, row_sum, sum_op);
}
} else {
// Reshape acc_o from (MMA=4, MMA_M, MMA_K) to (nrow=(2, MMA_M), ncol=(2, MMA_K))
Tensor acc_o_rowcol = make_tensor(acc_o.data(), flash::convert_layout_acc_rowcol(acc_o.layout()));
static_assert(decltype(size<0>(acc_o_rowcol))::value == kNRows);
static_assert(decltype(size<1>(acc_o_rowcol))::value % 2 == 0);
#pragma unroll
for (int mi = 0; mi < size(row_max); ++mi) {
float scores_max_cur = !Check_inf
? row_max(mi)
: (row_max(mi) == -INFINITY ? 0.0f : row_max(mi));
float scores_scale = __builtin_exp2f((scores_max_prev(mi) - scores_max_cur) * softmax_scale_log2);
row_sum(mi) *= scores_scale;
// #pragma unroll
// for (int ni = 0; ni < size<1>(acc_o_rowcol); ++ni) { acc_o_rowcol(mi, ni) *= scores_scale; }
Float2 scale_vec = {scores_scale , scores_scale};
Float2 beta_vec = {0.0f, 0.0f};
#pragma unroll
for (int ni = 0; ni < size<1>(acc_o_rowcol); ni += 2) {
Float2 x_vec = {acc_o_rowcol(mi, ni), acc_o_rowcol(mi, ni + 1)};
x_vec = __builtin_mxc_pk_fma_f32(x_vec, scale_vec, beta_vec);
acc_o_rowcol(mi, ni) = x_vec[0];
acc_o_rowcol(mi, ni + 1) = x_vec[1];
}
}
flash::scale_apply_exp2(scores, row_max, softmax_scale_log2);
#pragma unroll
for (int mi = 0; mi < size<0>(scores); mi++) {
if constexpr(AddVec) {
Float2 x_vec = {row_sum(mi), 0.0f};
Float2 scale_vec = {1.0f, 1.0f};
#pragma unroll
for (int ni = 0; ni < size<1>(scores); ni += 2) {
Float2 beta_vec = {scores(mi, ni), scores(mi, ni + 1)};
x_vec = __builtin_mxc_pk_fma_f32(x_vec, scale_vec, beta_vec);
}
row_sum(mi) = x_vec[0] + x_vec[1];
}
else {
#pragma unroll
@ -240,18 +435,17 @@ struct Softmax {
};
template<bool Is_dropout=false, bool Return_lse=true, bool Split=false, typename Tensor0>
__forceinline__ __device__ TensorT normalize_softmax_lse(Tensor0 &acc_o, float softmax_scale, float rp_dropout=1.0) {
__forceinline__ __device__ TensorT normalize_softmax_lse(Tensor0 &acc_o, float softmax_scale, float rp_dropout=1.0, float k_descale=1.0) {
flash::quadreduce_sum(row_sum);
TensorT lse = make_fragment_like(row_sum);
Tensor acc_o_rowcol = make_tensor(acc_o.data(), flash::convert_layout_acc_rowcol(acc_o.layout()));
static_assert(decltype(size<0>(acc_o_rowcol))::value == kNRows);
static_assert(decltype(size<1>(acc_o_rowcol))::value % 2 == 0);
typedef __NATIVE_VECTOR__(2, float) Float2;
#pragma unroll
for (int mi = 0; mi < size<0>(acc_o_rowcol); ++mi) {
float sum = row_sum(mi);
float inv_sum = (sum == 0.f || sum != sum) ? 1.f : 1.f / sum;
float inv_sum = (sum == 0.f || sum != sum) ? 1.f : k_descale / sum;
if (Return_lse)
lse(mi) = (sum == 0.f || sum != sum) ? (Split ? -INFINITY : INFINITY) : row_max(mi) * softmax_scale + __logf(sum);
float scale = !Is_dropout ? inv_sum : inv_sum * rp_dropout;
@ -271,6 +465,45 @@ struct Softmax {
}
return lse;
};
template<bool Is_dropout=false, bool Return_lse=true, bool Split=false, typename Tensor0, typename Tensor1>
__forceinline__ __device__ TensorT normalize_softmax_lse(Tensor0 &acc_o, Tensor1 &sRowSum, float softmax_scale, float rp_dropout=1.0) {
const int tidx = threadIdx.x;
const int wave_idx = tidx / 64;
const int lane_idx = tidx % 64;
const int wave_group_idx = wave_idx / 4;
const int row_offset = wave_idx % 4 * 16 + lane_idx % 16;
flash::quadreduce_sum(row_sum);
if (lane_idx / 16 == 0) {
sRowSum(wave_group_idx, row_offset) = row_sum(0); //sts row_max
}
flash::sync_threads();
row_sum(0) += sRowSum(wave_group_idx ^ 1, row_offset); //lds row_max
TensorT lse = make_fragment_like(row_sum);
Tensor acc_o_rowcol = make_tensor(acc_o.data(), flash::convert_layout_acc_rowcol(acc_o.layout()));
static_assert(decltype(size<0>(acc_o_rowcol))::value == kNRows);
static_assert(decltype(size<1>(acc_o_rowcol))::value % 2 == 0);
typedef __NATIVE_VECTOR__(2, float) Float2;
#pragma unroll
for (int mi = 0; mi < size<0>(acc_o_rowcol); ++mi) {
float sum = row_sum(mi);
float inv_sum = (sum == 0.f || sum != sum) ? 1.f : 1.f / sum;
if (Return_lse)
lse(mi) = (sum == 0.f || sum != sum) ? (Split ? -INFINITY : INFINITY) : row_max(mi) * softmax_scale + __logf(sum);
float scale = !Is_dropout ? inv_sum : inv_sum * rp_dropout;
Float2 scale_vec = {scale, scale};
Float2 beta_vec = {0.0f, 0.0f};
#pragma unroll
for (int ni = 0; ni < size<1>(acc_o_rowcol); ni += 2) {
Float2 x_vec = {acc_o_rowcol(mi, ni), acc_o_rowcol(mi, ni + 1)};
x_vec = __builtin_mxc_pk_fma_f32(x_vec, scale_vec, beta_vec);
acc_o_rowcol(mi, ni) = x_vec[0];
acc_o_rowcol(mi, ni + 1) = x_vec[1];
}
}
return lse;
};
};
} // namespace flash

View File

@ -0,0 +1,87 @@
// Adapted from Dao-AILab/flash-attention (https://github.com/Dao-AILab/flash-attention/tree/v2.6.3)
#pragma once
#include <cute/algorithm/copy.hpp>
#include <mctlass/mctlass.h>
#include <mctlass/array.h>
#include <mctlass/numeric_types.h>
#include "block_info.h"
#include "kernel_traits.h"
#include "utils.h"
#include "softmax.h"
#include "mask.h"
#include "xcore1000/flash_fwd_mla_kernel_k64_16x16_4waves_stage1_xcore1000.h"
// #include "xcore1000/flash_fwd_mla_kernel_k64_16x16_4waves_xcore1000.h"
#include "xcore1000/flash_fwd_mla_kernel_k64_32x16_4waves_xcore1000.h"
#include "xcore1000/flash_fwd_mla_kernel_k64_64x16_8waves_xcore1000.h"
#include "xcore1500/flash_fwd_mla_kernel_k64_64x16_8waves_xcore1500.h"
#include "xcore1500/flash_fwd_mla_kernel_k64_64x32_8waves_xcore1500.h"
#include "static_switch.h"
namespace flash {
using namespace cute;
template<typename Kernel_traits, bool Is_causal, bool Is_even_MN, bool Is_even_K, bool Is_enable_dcp, typename Params>
__forceinline__ __device__ void compute_attn_1rowblock_splitkv_mla(const Params &params, const int m_block_max) {
constexpr int kBlockN = Kernel_traits::kBlockN;
const int m_block = blockIdx.x;
const int bidh = blockIdx.y;
const int partition_idx = blockIdx.z;
extern __shared__ char shared_memory[];
int *tile_scheduler_metadata_ptr = params.tile_scheduler_metadata_ptr + partition_idx * TileSchedulerMetaDataSize;
// int4 tile_scheduler_metadata = __ldg(reinterpret_cast<int4 *>(tile_scheduler_metadata_ptr));
int4 tile_scheduler_metadata = *reinterpret_cast<int4 *>(tile_scheduler_metadata_ptr);
int begin_idx = tile_scheduler_metadata.x;
int sched_begin_block_idx = tile_scheduler_metadata.y;
int end_idx = tile_scheduler_metadata.z;
int sched_end_block_idx = tile_scheduler_metadata.w;
if (begin_idx >= params.b || begin_idx < 0) return;
// int begin_n_split_idx = __ldg(tile_scheduler_metadata_ptr + 4);
int begin_n_split_idx = tile_scheduler_metadata_ptr[4];
#pragma unroll 1
for (int batch_id = begin_idx; batch_id <= end_idx; ++batch_id) {
const int n_split_idx = batch_id == begin_idx ? begin_n_split_idx : 0;
const int seqlen_k = params.cu_seqlens_k[batch_id];
const int n_block_min = batch_id == begin_idx ? sched_begin_block_idx : 0;
const int n_block_max = batch_id == end_idx ? sched_end_block_idx : cute::ceil_div(seqlen_k, kBlockN);
// [n_block_min, n_block_max) need be calculated in kernel
if (n_block_max <= n_block_min) continue;
const bool NoSplit = __ldg(params.num_splits_ptr + batch_id + 1) - __ldg(params.num_splits_ptr + batch_id) == 1;
if (batch_id > begin_idx) {
__syncthreads(); // Barrier between two tiles.
}
#if defined(__MACA_ARCH__) && (__MACA_ARCH__ == 1500)
if constexpr (Kernel_traits::kBlockM == 64 && Kernel_traits::kBlockN == 16 && Kernel_traits::kNWarps == 8) {
compute_attn_1rowblock_splitkv_mla_k64_64x16_8waves_xcore1500<Kernel_traits, Is_causal, Is_even_MN, Is_even_K, Is_enable_dcp>(
params, batch_id, bidh, m_block, n_split_idx, n_block_min, n_block_max, NoSplit);
}else if constexpr (Kernel_traits::kBlockM == 64 && Kernel_traits::kBlockN == 32 && Kernel_traits::kNWarps == 8) {
compute_attn_1rowblock_splitkv_mla_k64_64x32_8waves_xcore1500<Kernel_traits, Is_causal, Is_even_MN, Is_even_K, Is_enable_dcp>(
params, batch_id, bidh, m_block, n_split_idx, n_block_min, n_block_max, NoSplit);
}
#else defined(__MACA_ARCH__) && (__MACA_ARCH__ == 1000)
if constexpr (Kernel_traits::kBlockM == 32 && Kernel_traits::kBlockN == 16 && Kernel_traits::kNWarps == 4) {
compute_attn_1rowblock_splitkv_mla_k64_32x16_4waves_xcore1000<Kernel_traits, Is_causal, Is_even_MN, Is_even_K, Is_enable_dcp>(
params, batch_id, bidh, m_block, n_split_idx, n_block_min, n_block_max, NoSplit);
} else if constexpr (Kernel_traits::kBlockM == 64 && Kernel_traits::kBlockN == 16 && Kernel_traits::kNWarps == 8) {
compute_attn_1rowblock_splitkv_mla_k64_64x16_8waves_xcore1000<Kernel_traits, Is_causal, Is_even_MN, Is_even_K, Is_enable_dcp>(
params, batch_id, bidh, m_block, n_split_idx, n_block_min, n_block_max, NoSplit);
} else if constexpr (Kernel_traits::kBlockM == 16 && Kernel_traits::kBlockN == 16 && Kernel_traits::kNWarps == 4) {
compute_attn_1rowblock_splitkv_mla_k64_16x16_4waves_xcore1000<Kernel_traits, Is_causal, Is_even_MN, Is_even_K, Is_enable_dcp>(
params, batch_id, bidh, m_block, n_split_idx, n_block_min, n_block_max, NoSplit);
}
#endif
}
}
}// namespace flash

View File

@ -1,686 +0,0 @@
// Adapted from Dao-AILab/flash-attention (https://github.com/Dao-AILab/flash-attention/tree/v2.6.3)
#pragma once
#include <cute/algorithm/copy.hpp>
#include <mctlass/mctlass.h>
#include <mctlass/array.h>
#include <mctlass/numeric_types.h>
#include "block_info.h"
#include "kernel_traits.h"
#include "utils.h"
#include "softmax.h"
#include "mask.h"
#include "rotary.h"
#include "attn_mask.h"
namespace flash {
using namespace cute;
template<typename Kernel_traits, bool Is_causal, bool Is_local, bool Has_alibi, bool Is_even_MN, bool Is_even_K, bool Is_softcap, bool Split, bool Append_KV, bool Is_page_attn, typename Params>
__forceinline__ __device__ void compute_attn_1rowblock_splitkv_k64_mla_V1x8(const Params &params, const int bidb, const int bidh, const int m_block, const int n_split_idx, const int num_n_splits) {
using Element = typename Kernel_traits::Element;
using ElementAccum = typename Kernel_traits::ElementAccum;
using index_t = typename Kernel_traits::index_t;
// Shared memory.
extern __shared__ char smem_[];
// The thread index.
const int tidx = threadIdx.x;
constexpr int kBlockM = Kernel_traits::kBlockM;
constexpr int kBlockN = Kernel_traits::kBlockN;
constexpr int kHeadDim = Kernel_traits::kHeadDim;
constexpr int kHeadDimV = Kernel_traits::kHeadDimV;
constexpr int kNWarps = Kernel_traits::kNWarps;
constexpr int kBlockKSmem = Kernel_traits::kBlockKSmem;
constexpr int kBlockKGmem = Kernel_traits::UseWarpsNx1 ? Kernel_traits::kBlockKSmem : 128;
constexpr int kAtomLayoutMS = Kernel_traits::kAtomLayoutMS;
constexpr int kAtomLayoutMO = Kernel_traits::kAtomLayoutMO;
static_assert(kBlockKSmem == 64);
using GmemTiledCopyO = std::conditional_t<
!Split,
typename Kernel_traits::GmemTiledCopyO,
typename Kernel_traits::GmemTiledCopyOaccum
>;
using ElementO = std::conditional_t<!Split, Element, ElementAccum>;
const BlockInfo</*Varlen=*/!Is_even_MN> binfo(params, bidb);
// if (threadIdx.x == 0 && blockIdx.y == 0 && blockIdx.z == 0) { printf("Is_even_MN = %d, is_cumulativ = %d, seqlen_k_cache = %d, actual_seqlen_k = %d\n", Is_even_MN, params.is_seqlens_k_cumulative, binfo.seqlen_k_cache, binfo.actual_seqlen_k); }
// if (threadIdx.x == 0 && blockIdx.y == 1 && blockIdx.z == 0) { printf("params.knew_ptr = %p, seqlen_k_cache + seqlen_knew = %d\n", params.knew_ptr, binfo.seqlen_k_cache + (params.knew_ptr == nullptr ? 0 : params.seqlen_knew)); }
if (m_block * kBlockM >= binfo.actual_seqlen_q) return;
const int n_blocks_per_split = ((binfo.actual_seqlen_k + kBlockN - 1) / kBlockN + num_n_splits - 1) / num_n_splits;
const int n_block_min = !Is_local
? n_split_idx * n_blocks_per_split
: std::max(n_split_idx * n_blocks_per_split, (m_block * kBlockM + binfo.actual_seqlen_k - binfo.actual_seqlen_q - params.window_size_left) / kBlockN);
int n_block_max = std::min(cute::ceil_div(binfo.actual_seqlen_k, kBlockN), (n_split_idx + 1) * n_blocks_per_split);
if (Is_causal || Is_local) {
n_block_max = std::min(n_block_max,
cute::ceil_div((m_block + 1) * kBlockM + binfo.actual_seqlen_k - binfo.actual_seqlen_q / params.ngroups + params.window_size_right, kBlockN));
}
if (n_block_min >= n_block_max) { // This also covers the case where n_block_max <= 0
// We exit early and write 0 to gOaccum and -inf to gLSEaccum.
// Otherwise we might read OOB elements from gK and gV,
// or get wrong results when we combine gOaccum from different blocks.
const index_t row_offset_o = binfo.q_offset(params.o_batch_stride, params.o_row_stride, bidb)
+ m_block * kBlockM * params.o_row_stride + bidh * params.o_head_stride;
const index_t row_offset_oaccum = (((n_split_idx * params.b + bidb) * params.h + bidh) * params.seqlen_q
+ m_block * kBlockM) * params.d_v;
const index_t row_offset_lseaccum = ((n_split_idx * params.b + bidb) * params.h + bidh) * params.seqlen_q + m_block * kBlockM;
Tensor gOaccum = make_tensor(make_gmem_ptr(reinterpret_cast<ElementO *>(Split ? params.oaccum_ptr : params.o_ptr) + (Split ? row_offset_oaccum : row_offset_o)),
Shape<Int<kBlockM>, Int<kHeadDimV>>{},
make_stride(Split ? kHeadDimV : params.o_row_stride, _1{}));
Tensor gLSEaccum = make_tensor(make_gmem_ptr(reinterpret_cast<ElementAccum *>(Split ? params.softmax_lseaccum_ptr : params.softmax_lse_ptr) + row_offset_lseaccum),
Shape<Int<kBlockM>>{}, Stride<_1>{});
GmemTiledCopyO gmem_tiled_copy_Oaccum;
auto gmem_thr_copy_Oaccum = gmem_tiled_copy_Oaccum.get_thread_slice(tidx);
Tensor tOgOaccum = gmem_thr_copy_Oaccum.partition_D(gOaccum);
Tensor tOrOaccum = make_tensor<ElementO>(shape(tOgOaccum));
clear(tOrOaccum);
// Construct identity layout for sO
Tensor cO = make_identity_tensor(make_shape(size<0>(gOaccum), size<1>(gOaccum))); // (BLK_M,BLK_K) -> (blk_m,blk_k)
// Repeat the partitioning with identity layouts
Tensor tOcO = gmem_thr_copy_Oaccum.partition_D(cO);
Tensor tOpO = make_tensor<bool>(make_shape(size<2>(tOgOaccum)));
if (!Is_even_K) {
#pragma unroll
for (int k = 0; k < size(tOpO); ++k) { tOpO(k) = get<1>(tOcO(0, 0, k)) < params.d_v; }
}
// Clear_OOB_K must be false since we don't want to write zeros to gmem
flash::copy<Is_even_MN, Is_even_K, /*Clear_OOB_MN=*/false, /*Clear_OOB_K=*/false>(
gmem_tiled_copy_Oaccum, tOrOaccum, tOgOaccum, tOcO, tOpO, binfo.actual_seqlen_q - m_block * kBlockM
);
#pragma unroll
for (int m = 0; m < size<1>(tOgOaccum); ++m) {
const int row = get<0>(tOcO(0, m, 0));
if (row < binfo.actual_seqlen_q - m_block * kBlockM && get<1>(tOcO(0, m, 0)) == 0) { gLSEaccum(row) = Split ? -INFINITY : INFINITY; }
}
return;
}
// We iterate over the blocks in reverse order. This is because the last block is the only one
// that needs masking when we read K and V from global memory. Moreover, iterating in reverse
// might save us 1 register (we just need n_block instead of both n_block and n_block_max).
const index_t row_offset_q = binfo.q_offset(params.q_batch_stride, params.q_row_stride, bidb)
+ m_block * kBlockM * params.q_row_stride + bidh * params.q_head_stride;
// We move K and V to the last block.
const int bidb_cache = params.cache_batch_idx == nullptr ? bidb : params.cache_batch_idx[bidb];
const int *block_table = params.block_table == nullptr ? nullptr : params.block_table + bidb * params.block_table_batch_stride;
const int block_table_idx = block_table == nullptr ? 0 : (n_block_max - 1) * kBlockN / params.page_block_size;
const int block_table_offset = block_table == nullptr ? 0 : (n_block_max - 1) * kBlockN - block_table_idx * params.page_block_size;
const index_t row_offset_k = block_table == nullptr
? binfo.k_offset(params.k_batch_stride, params.k_row_stride, bidb_cache)
+ (n_block_max - 1) * kBlockN * params.k_row_stride + (bidh / params.h_h_k_ratio) * params.k_head_stride
: (bidh / params.h_h_k_ratio) * params.k_head_stride;
const index_t row_offset_v = block_table == nullptr
? binfo.k_offset(params.v_batch_stride, params.v_row_stride, bidb_cache)
+ (n_block_max - 1) * kBlockN * params.v_row_stride + (bidh / params.h_h_k_ratio) * params.v_head_stride
: (bidh / params.h_h_k_ratio) * params.v_head_stride;
Tensor gQ = make_tensor(make_gmem_ptr(reinterpret_cast<Element *>(params.q_ptr) + row_offset_q),
Shape<Int<kBlockM>, Int<kHeadDim>>{},
make_stride(params.q_row_stride, _1{}));
Tensor gK = make_tensor(make_gmem_ptr(reinterpret_cast<Element *>(params.k_ptr) + row_offset_k),
Shape<Int<kBlockN>, Int<kHeadDim>>{},
make_stride(params.k_row_stride, _1{}));
// if (threadIdx.x == 0 && blockIdx.y == 0 && blockIdx.z == 0) { printf("k_ptr = %p, row_offset_k = %d, gK_ptr = %p\n", params.k_ptr, row_offset_k, gK.data()); }
Tensor gV = make_tensor(make_gmem_ptr(reinterpret_cast<Element *>(params.v_ptr) + row_offset_v),
Shape<Int<kBlockN>, Int<kHeadDimV>>{},
make_stride(params.v_row_stride, _1{}));
Tensor sQ = make_tensor(make_smem_ptr(reinterpret_cast<Element *>(smem_)),
typename Kernel_traits::SmemLayoutQ{});
//Tensor sK = make_tensor(sQ.data() + size(sQ), typename Kernel_traits::SmemLayoutKV{});
Tensor sK = make_tensor(sQ.data() + (Kernel_traits::Share_Q_K_smem ? 0 : size(sQ)),
typename Kernel_traits::SmemLayoutK{});
Tensor sV = make_tensor(sK.data() + size(sK), typename Kernel_traits::SmemLayoutVtNoSwizzle{});
Tensor sVt = make_tensor(sV.data(), typename Kernel_traits::SmemLayoutVtransposed{});
Tensor sVtNoSwizzle = make_tensor(sV.data(), typename Kernel_traits::SmemLayoutVtransposedNoSwizzle{});
typename Kernel_traits::GmemTiledCopyB128 gmem_tiled_copy_Q;
auto gmem_thr_copy_Q = gmem_tiled_copy_Q.get_thread_slice(tidx);
Tensor tQgQ = gmem_thr_copy_Q.partition_S(gQ);
Tensor tQsQ = gmem_thr_copy_Q.partition_D(sQ);
typename Kernel_traits::GmemTiledCopyB64 gmem_tiled_copy_KV;
auto gmem_thr_copy_KV = gmem_tiled_copy_KV.get_thread_slice(tidx);
Tensor tKgK = gmem_thr_copy_KV.partition_S(gK); // (KCPY, KCPY_N, KCPY_K)
Tensor tKsK = gmem_thr_copy_KV.partition_D(sK);
Tensor tVgV = gmem_thr_copy_KV.partition_S(gV); // (VCPY, VCPY_N, VCPY_K)
Tensor tVsV = gmem_thr_copy_KV.partition_D(sV);
Tensor tVrV = make_fragment_like(tVgV);
// wave0 and wave2 compute the same S, wave1 and wave3 compute the same S
int tidx_mma_s = tidx & 0x7F;
typename Kernel_traits::TiledMmaS tiled_mma_s;
auto thr_mma_s = tiled_mma_s.get_thread_slice(tidx_mma_s);
Tensor tSrQ = thr_mma_s.partition_fragment_A(sQ); // (MMA,MMA_M,MMA_K)
Tensor tSrK = thr_mma_s.partition_fragment_B(sK); // (MMA,MMA_N,MMA_K)
typename Kernel_traits::TiledMmaO tiled_mma_o;
auto thr_mma_o = tiled_mma_o.get_thread_slice(tidx);
Tensor tOrVt = thr_mma_o.partition_fragment_B(sVtNoSwizzle); // (MMA, MMA_K,MMA_N)
Tensor acc_o = partition_fragment_C(tiled_mma_o, Shape<Int<kBlockM>, Int<kHeadDimV>>{}); // MMA, MMA_M, MMA_K
//
// Copy Atom retiling
//
auto smem_tiled_copy_Q = make_tiled_copy_A(typename Kernel_traits::SmemCopyAtomB64{}, tiled_mma_s);
auto smem_thr_copy_Q = smem_tiled_copy_Q.get_thread_slice(tidx_mma_s);
Tensor tSsQ = smem_thr_copy_Q.partition_S(sQ);
auto smem_tiled_copy_K = make_tiled_copy_B(typename Kernel_traits::SmemCopyAtomB64{}, tiled_mma_s);
auto smem_thr_copy_K = smem_tiled_copy_K.get_thread_slice(tidx_mma_s);
Tensor tSsK = smem_thr_copy_K.partition_S(sK);
auto smem_tiled_copy_V = make_tiled_copy_B(typename Kernel_traits::SmemCopyAtomTransposed{}, tiled_mma_o);
auto smem_thr_copy_V = smem_tiled_copy_V.get_thread_slice(tidx);
Tensor tOsVt = smem_thr_copy_V.partition_S(sVtNoSwizzle);
// PREDICATES
// Construct identity layout for sQ and sK
Tensor cQ = make_identity_tensor(make_shape(size<0>(sQ), size<1>(sQ))); // (BLK_M,BLK_K) -> (blk_m,blk_k)
Tensor cKV = make_identity_tensor(make_shape(size<0>(sK), size<1>(sK))); // (BLK_N,BLK_K) -> (blk_n,blk_k)
// Repeat the partitioning with identity layouts
Tensor tQcQ = gmem_thr_copy_Q.partition_S(cQ); // (ACPY,ACPY_M,ACPY_K) -> (blk_m,blk_k)
Tensor tKVcKV = gmem_thr_copy_KV.partition_S(cKV); // (BCPY,BCPY_N,BCPY_K) -> (blk_n,blk_k)
// Prologue
// Read Q from gmem to smem, optionally apply rotary embedding.
Tensor tQrQ = make_fragment_like(tQgQ);
// We don't need to clear the sQ smem tiles since we'll only write out the valid outputs
flash::copy_b128<Is_even_MN, Is_even_K>(tQgQ, tQrQ, tQcQ, params.d, binfo.actual_seqlen_q - m_block * kBlockM);
cute::copy(tQrQ, tQsQ);
if constexpr (Kernel_traits::Is_Q_in_regs) {
flash::sync_threads();
cute::copy(smem_tiled_copy_Q, tSsQ, tSrQ);
flash::sync_threads();
}
int n_block = n_block_max - 1;
// We don't need to clear the sK smem tiles since we'll mask out the scores anyway.
Tensor tKrK = make_fragment_like(tKgK);
if constexpr (!Is_page_attn) {
flash::copy_b64<Is_even_MN, Is_even_K>(tKgK, tKrK, tKVcKV, params.d, binfo.actual_seqlen_k - n_block * kBlockN);
} else {
flash::copy_b64_page_one<Kernel_traits, Is_even_MN, Is_even_K>(gK, tKgK, tKrK, tKVcKV, params.d, n_block,
block_table, params.k_batch_stride, params.k_row_stride, params.page_block_size, binfo.actual_seqlen_k - n_block * kBlockN);
}
// flash::cp_async_wait<0>();
// __syncthreads();
// if (tidx == 0 && blockIdx.y == 0 && blockIdx.z == 0) { print(tKsK); }
// __syncthreads();
clear(acc_o);
flash::Softmax<size<1>(acc_o)> softmax;
const float alibi_slope = !Has_alibi ? 0.0f : reinterpret_cast<float *>(params.alibi_slopes_ptr)[bidb * params.alibi_slopes_batch_stride + bidh] / params.scale_softmax;
flash::Mask<Is_causal, Is_local, Has_alibi> mask(binfo.actual_seqlen_k, binfo.actual_seqlen_q, params.ngroups, params.window_size_left, params.window_size_right, alibi_slope);
// For performance reason, we separate out two kinds of iterations:
// those that need masking on S, and those that don't.
// We need masking on S for the very last block when K and V has length not multiple of kBlockN.
// We also need masking on S if it's causal, for the last ceil_div(kBlockM, kBlockN) blocks.
// We will have at least 1 "masking" iteration.
// If not even_N, then seqlen_k might end in the middle of a block. In that case we need to
// mask 2 blocks (e.g. when kBlockM == kBlockN), not just 1.
constexpr int n_masking_steps = (!Is_causal && !Is_local)
? 1
: ((Is_even_MN && Is_causal) ? cute::ceil_div(kBlockM, kBlockN) : cute::ceil_div(kBlockM, kBlockN) + 1);
#pragma unroll
for (int masking_step = 0; masking_step < n_masking_steps; ++masking_step, --n_block) {
Tensor acc_s = partition_fragment_C(tiled_mma_s, Shape<Int<kBlockM>, Int<kBlockN>>{}); // (MMA=4, MMA_M, MMA_N)
cute::copy(tKrK, tKsK);
clear(acc_s);
// Advance gV
if (masking_step > 0) {
if constexpr (!Is_page_attn) {
tVgV.data() = tVgV.data() + (-int(kBlockN * params.v_row_stride));
flash::copy_b64</*Is_even_MN=*/true, Is_even_K>(tVgV, tVrV, tKVcKV, params.d_v);
} else {
flash::copy_b64_page_one<Kernel_traits, /*Is_even_MN=*/true, Is_even_K>(gV, tVgV, tVrV, tKVcKV, params.d_v, n_block,
block_table, params.v_batch_stride, params.v_row_stride, params.page_block_size);
}
} else {
if constexpr (!Is_page_attn) {
// Clear the smem tiles to account for predicated off loads
flash::copy_b64<Is_even_MN, Is_even_K>(
tVgV, tVrV, tKVcKV, params.d_v, binfo.actual_seqlen_k - n_block * kBlockN
);
} else {
flash::copy_b64_page_one<Kernel_traits, Is_even_MN, Is_even_K>(gV, tVgV, tVrV, tKVcKV, params.d_v, n_block,
block_table, params.v_batch_stride, params.v_row_stride, params.page_block_size, binfo.actual_seqlen_k - n_block * kBlockN);
}
}
flash::sync_threads();
flash::gemm_opt</*A_in_regs=*/Kernel_traits::Is_Q_in_regs>(
acc_s, tSrQ, tSrK, tSsQ, tSsK, tiled_mma_s, smem_tiled_copy_Q, smem_tiled_copy_K,
smem_thr_copy_Q, smem_thr_copy_K
);
// if (cute::thread0()) { print(acc_s); }
if constexpr (Is_softcap){
flash::apply_softcap(acc_s, params.softcap);
}
mask.template apply_mask<Is_causal, Is_even_MN>(
acc_s, n_block * kBlockN, m_block * kBlockM + (tidx / 64) % kAtomLayoutMS * 16 + (tidx & 0xf), kAtomLayoutMS * 16
);
cute::copy(tVrV, tVsV);
if (n_block > n_block_min) {
// Advance gK
if constexpr (!Is_page_attn) {
tKgK.data() = tKgK.data() + (-int(kBlockN * params.k_row_stride));
flash::copy_b64</*Is_even_MN=*/true, Is_even_K>(tKgK, tKrK, tKVcKV, params.d);
} else {
flash::copy_b64_page_one<Kernel_traits, /*Is_even_MN=*/true, Is_even_K>(gK, tKgK, tKrK, tKVcKV, params.d, n_block - 1,
block_table, params.k_batch_stride, params.k_row_stride, params.page_block_size);
}
}
// We have key_padding_mask so we'll need to Check_inf
masking_step == 0
? softmax.template softmax_rescale_o</*Is_first=*/true, /*Check_inf=*/Is_causal || Is_local || !Is_even_MN, true, true>(acc_s, acc_o, params.scale_softmax_log2)
: softmax.template softmax_rescale_o</*Is_first=*/false, /*Check_inf=*/Is_causal || Is_local || !Is_even_MN, true, true>(acc_s, acc_o, params.scale_softmax_log2);
// if (cute::thread0()) { print(scores_max); print(scores_sum); print(scores); }
// Convert acc_s from fp32 to fp16/bf16
//Tensor rP = flash::convert_type<Element>(acc_s);
CONVERT_TENSOR_TYPE(ElementAccum, Element, acc_s, rP)
// Reshape rP from (MMA=4, MMA_M, MMA_N) to ((4, 2), MMA_M, MMA_N / 2)
// if using m16n8k16 or (4, MMA_M, MMA_N) if using m16n8k8.
//Tensor tOrP = make_tensor(rP.data(), flash::convert_layout_acc_Aregs<Kernel_traits::TiledMma>(rP.layout()));
Tensor tOrP = make_tensor(rP.data(), acc_s.layout());
flash::gemm_rs(acc_o, tOrP, tOrVt, tOsVt, tiled_mma_o, smem_tiled_copy_V, smem_thr_copy_V);
// This check is at the end of the loop since we always have at least 1 iteration
if (n_masking_steps > 1 && n_block <= n_block_min) {
--n_block;
break;
}
}
// These are the iterations where we don't need masking on S
for (; n_block >= n_block_min; --n_block) {
Tensor acc_s = partition_fragment_C(tiled_mma_s, Shape<Int<kBlockM>, Int<kBlockN>>{}); // (MMA=4, MMA_M, MMA_N)
cute::copy(tKrK, tKsK);
clear(acc_s);
// Advance gV
if constexpr (!Is_page_attn) {
tVgV.data() = tVgV.data() + (-int(kBlockN * params.v_row_stride));
flash::copy_b64</*Is_even_MN=*/true, Is_even_K>(tVgV, tVrV, tKVcKV, params.d_v);
} else {
flash::copy_b64_page_one<Kernel_traits, /*Is_even_MN=*/true, Is_even_K>(gV, tVgV, tVrV, tKVcKV, params.d_v, n_block,
block_table, params.v_batch_stride, params.v_row_stride, params.page_block_size);
}
flash::sync_threads();
flash::gemm_opt</*A_in_regs=*/Kernel_traits::Is_Q_in_regs>(
acc_s, tSrQ, tSrK, tSsQ, tSsK, tiled_mma_s, smem_tiled_copy_Q, smem_tiled_copy_K,
smem_thr_copy_Q, smem_thr_copy_K
);
cute::copy(tVrV, tVsV);
if (n_block > n_block_min) {
// Advance gK
if constexpr (!Is_page_attn) {
tKgK.data() = tKgK.data() + (-int(kBlockN * params.k_row_stride));
flash::copy_b64</*Is_even_MN=*/true, Is_even_K>(tKgK, tKrK, tKVcKV, params.d);
} else {
flash::copy_b64_page_one<Kernel_traits, /*Is_even_MN=*/true, Is_even_K>(gK, tKgK, tKrK, tKVcKV, params.d, n_block - 1,
block_table, params.k_batch_stride, params.k_row_stride, params.page_block_size);
}
}
if constexpr (Is_softcap){
flash::apply_softcap(acc_s, params.softcap);
}
mask.template apply_mask<Is_causal, Is_even_MN>(
acc_s, n_block * kBlockN, m_block * kBlockM + (tidx / 64) % kAtomLayoutMS * 16 + (tidx & 0xf), kAtomLayoutMS * 16
);
softmax.template softmax_rescale_o</*Is_first=*/false, /*Check_inf=*/Is_causal || Is_local || !Is_even_MN, true, true>(acc_s, acc_o, params.scale_softmax_log2);
//Tensor rP = flash::convert_type<Element>(acc_s);
CONVERT_TENSOR_TYPE(ElementAccum, Element, acc_s, rP)
// Reshape rP from (MMA=4, MMA_M, MMA_N) to ((4, 2), MMA_M, MMA_N / 2)
// if using m16n8k16 or (4, MMA_M, MMA_N) if using m16n8k8.
//Tensor tOrP = make_tensor(rP.data(), flash::convert_layout_acc_Aregs<Kernel_traits::TiledMma>(rP.layout()));
Tensor tOrP = make_tensor(rP.data(), acc_s.layout());
flash::gemm_rs(acc_o, tOrP, tOrVt, tOsVt, tiled_mma_o, smem_tiled_copy_V, smem_thr_copy_V);
}
// Epilogue
Tensor lse = softmax.template normalize_softmax_lse</*Is_dropout=*/false, /*Return_lse*/true, Split>(acc_o, params.scale_softmax);
// if (cute::thread0()) { print(lse); }
if constexpr (!Split) {
// use smem for O (mtreg->smem->mtreg->global)
Tensor sOaccum = make_tensor(make_smem_ptr(reinterpret_cast<ElementO *>(smem_)), typename Kernel_traits::SmemLayoutO{}); // (SMEM_M,SMEM_N)
// Partition sO to match the accumulator partitioning
using SmemTiledCopyO = typename Kernel_traits::SmemCopyAtomO;
auto smem_tiled_copy_Oaccum = make_tiled_copy_C(SmemTiledCopyO{}, tiled_mma_o);
auto smem_thr_copy_Oaccum = smem_tiled_copy_Oaccum.get_thread_slice(tidx);
//Tensor rO = flash::convert_type<ElementO>(acc_o);
CONVERT_TENSOR_TYPE(ElementAccum, ElementO, acc_o, rO)
Tensor taccOrOaccum = smem_thr_copy_Oaccum.retile_S(rO); // ((Atom,AtomNum), MMA_M, MMA_N)
Tensor taccOsOaccum = smem_thr_copy_Oaccum.partition_D(sOaccum); // ((Atom,AtomNum),PIPE_M,PIPE_N)
// sOaccum is larger than sQ, so we need to syncthreads here
// TODO: allocate enough smem for sOaccum
if constexpr (Kernel_traits::Share_Q_K_smem) { flash::sync_threads(); }
cute::copy(smem_tiled_copy_Oaccum, taccOrOaccum, taccOsOaccum);
const index_t row_offset_o = binfo.q_offset(params.o_batch_stride, params.o_row_stride, bidb)
+ m_block * kBlockM * params.o_row_stride + bidh * params.o_head_stride;
const index_t row_offset_oaccum = (((n_split_idx * params.b + bidb) * params.h + bidh) * params.seqlen_q
+ m_block * kBlockM) * params.d_v;
const index_t row_offset_lseaccum = ((n_split_idx * params.b + bidb) * params.h + bidh) * params.seqlen_q + m_block * kBlockM;
Tensor gOaccum = make_tensor(make_gmem_ptr(reinterpret_cast<ElementO *>(params.o_ptr) + (row_offset_o)),
Shape<Int<kBlockM>, Int<kHeadDimV>>{},
make_stride(params.o_row_stride, _1{}));
Tensor gLSEaccum = make_tensor(make_gmem_ptr(reinterpret_cast<ElementAccum *>(params.softmax_lse_ptr) + row_offset_lseaccum),
Shape<Int<kBlockM>>{}, Stride<_1>{});
// if (tidx == 0) { printf("row_offset_o = %d, bidh = %d, gOaccum = %p\n", row_offset_o, bidh, gOaccum.data()); }
GmemTiledCopyO gmem_tiled_copy_Oaccum;
auto gmem_thr_copy_Oaccum = gmem_tiled_copy_Oaccum.get_thread_slice(tidx);
Tensor tOsOaccum = gmem_thr_copy_Oaccum.partition_S(sOaccum); // ((Atom,AtomNum),ATOM_M,ATOM_N)
Tensor tOgOaccum = gmem_thr_copy_Oaccum.partition_D(gOaccum);
flash::sync_threads();
Tensor tOrOaccum = make_tensor<ElementO>(shape(tOgOaccum));
cute::copy(gmem_tiled_copy_Oaccum, tOsOaccum, tOrOaccum);
Tensor caccO = make_identity_tensor(Shape<Int<kBlockM>, Int<kHeadDimV>>{}); // (BLK_M,BLK_K) -> (blk_m,blk_k)
Tensor taccOcO = thr_mma_o.partition_C(caccO); // (MMA,MMA_M,MMA_K)
static_assert(decltype(size<0>(taccOcO))::value == 4);
// Convert to ((2, 2), MMA_M, MMA_K) then take only the row indices.
Tensor taccOcO_row = logical_divide(taccOcO, Shape<_4>{})(make_coord(0, _), _, 0);
CUTE_STATIC_ASSERT_V(size(lse) == size(taccOcO_row)); // MMA_M
if (get<1>(taccOcO_row(0)) == 0) {
#pragma unroll
for (int mi = 0; mi < size(lse); ++mi) {
const int row = get<0>(taccOcO_row(mi));
if (row < binfo.actual_seqlen_q - m_block * kBlockM) { gLSEaccum(row) = lse(mi); }
}
}
// Construct identity layout for sO
Tensor cO = make_identity_tensor(make_shape(size<0>(sOaccum), size<1>(sOaccum))); // (BLK_M,BLK_K) -> (blk_m,blk_k)
// Repeat the partitioning with identity layouts
Tensor tOcO = gmem_thr_copy_Oaccum.partition_D(cO); // (ACPY,ACPY_M,ACPY_K) -> (blk_m,blk_k)
// Clear_OOB_K must be false since we don't want to write zeros to gmem
flash::copy_reg_to_global<Is_even_MN, Is_even_K>(
tOrOaccum, tOgOaccum, tOcO, params.d_v, binfo.actual_seqlen_q - m_block * kBlockM
);
} else {
// don't use smem for O (mtreg->global)
const index_t row_offset_oaccum = (((n_split_idx * params.b + bidb) * params.h + bidh) * params.seqlen_q
+ m_block * kBlockM) * params.d_v;
const index_t row_offset_lseaccum = ((n_split_idx * params.b + bidb) * params.h + bidh) * params.seqlen_q + m_block * kBlockM;
Tensor gOaccum = make_tensor(make_gmem_ptr(reinterpret_cast<ElementO *>(params.oaccum_ptr) + row_offset_oaccum),
Shape<Int<kBlockM>, Int<kHeadDimV>>{},
make_stride(kHeadDimV, _1{}));
Tensor gLSEaccum = make_tensor(make_gmem_ptr(reinterpret_cast<ElementAccum *>(params.softmax_lseaccum_ptr) + row_offset_lseaccum),
Shape<Int<kBlockM>>{}, Stride<_1>{});
// if (tidx == 0) { printf("row_offset_o = %d, bidh = %d, gOaccum = %p\n", row_offset_o, bidh, gOaccum.data()); }
using GmemCopyAtomOaccum = typename Kernel_traits::SmemCopyAtomOaccum;
auto gmem_tiled_copy_Oaccum = make_tiled_copy_C(GmemCopyAtomOaccum{}, tiled_mma_o);
auto gmem_thr_copy_Oaccum = gmem_tiled_copy_Oaccum.get_thread_slice(tidx);
Tensor taccOrOaccum = gmem_thr_copy_Oaccum.retile_S(acc_o); // ((Atom,AtomNum), MMA_M, MMA_N)
Tensor taccOgOaccum = gmem_thr_copy_Oaccum.partition_D(gOaccum);
Tensor caccO = make_identity_tensor(Shape<Int<kBlockM>, Int<kHeadDimV>>{}); // (BLK_M,BLK_K) -> (blk_m,blk_k)
Tensor taccOcO = thr_mma_o.partition_C(caccO); // (MMA,MMA_M,MMA_K)
static_assert(decltype(size<0>(taccOcO))::value == 4);
// Convert to ((2, 2), MMA_M, MMA_K) then take only the row indices.
Tensor taccOcO_row = logical_divide(taccOcO, Shape<_4>{})(make_coord(0, _), _, 0);
CUTE_STATIC_ASSERT_V(size(lse) == size(taccOcO_row)); // MMA_M
if (get<1>(taccOcO_row(0)) == 0) {
#pragma unroll
for (int mi = 0; mi < size(lse); ++mi) {
const int row = get<0>(taccOcO_row(mi));
if (row < binfo.actual_seqlen_q - m_block * kBlockM) { gLSEaccum(row) = lse(mi); }
}
}
// Clear_OOB_K must be false since we don't want to write zeros to gmem
flash::copy_reg_to_global<Is_even_MN, Is_even_K>(
taccOrOaccum, taccOgOaccum, taccOcO, params.d_v, binfo.actual_seqlen_q - m_block * kBlockM
);
}
}
template<typename Kernel_traits, bool Is_causal, bool Is_local, bool Has_alibi, bool Is_even_MN, bool Is_even_K, bool Is_softcap, bool Split, bool Append_KV, bool Is_page_attn, typename Params>
__forceinline__ __device__ void compute_attn_splitkv(const Params &params, const int m_block_max) {
const int m_block = blockIdx.x;
// The block index for the batch.
const int bidb = Split ? blockIdx.z / params.h : blockIdx.y;
// The block index for the head.
const int bidh = Split ? blockIdx.z - bidb * params.h : blockIdx.z;
const int n_split_idx = Split ? blockIdx.y : 0;
const int num_n_splits = Split ? gridDim.y : 1;
compute_attn_1rowblock_splitkv_k64_mla_V1x8<Kernel_traits, Is_causal, Is_local, Has_alibi, Is_even_MN,Is_even_K, Is_softcap, Split, Append_KV, Is_page_attn>(
params, bidb, bidh, m_block, n_split_idx, num_n_splits);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
template<typename Kernel_traits, int kBlockM, int Log_max_splits, bool Is_even_K, typename Params>
__forceinline__ __device__ void combine_attn_seqk_parallel(const Params &params) {
using Element = typename Kernel_traits::Element;
using ElementAccum = typename Kernel_traits::ElementAccum;
using index_t = typename Kernel_traits::index_t;
constexpr int kMaxSplits = 1 << Log_max_splits;
constexpr int kHeadDimV = Kernel_traits::kHeadDimV;
constexpr int kNThreads = 256;/*Kernel_traits::kNThreads*/;
static_assert(kMaxSplits <= 128, "kMaxSplits must be <= 128");
static_assert(kBlockM == 4 || kBlockM == 8 || kBlockM == 16 || kBlockM == 32, "kBlockM must be 4, 8, 16 or 32");
static_assert(kNThreads == 128 || kNThreads == 256, "We assume that each block has 128 or 256 threads");
// Shared memory.
// kBlockM + 1 instead of kBlockM to reduce bank conflicts.
__shared__ ElementAccum sLSE[kMaxSplits][kBlockM + 1];
// The thread and block index.
const int tidx = threadIdx.x;
const int bidx = blockIdx.x;
const index_t lse_size = params.b * params.h * params.seqlen_q;
const index_t row_offset_lse = bidx * kBlockM;
Tensor gLSEaccum = make_tensor(make_gmem_ptr(reinterpret_cast<ElementAccum *>(params.softmax_lseaccum_ptr) + row_offset_lse),
Shape<Int<kMaxSplits>, Int<kBlockM>>{},
make_stride(lse_size, _1{}));
// LSE format is different depending on params.unpadded_lse and params.seqlenq_ngroups_swapped, see comment in get_lse_tile.
// This tensor's layout maps row_offset_lse to {bidb, bidh, lse_size}.
Tensor gLSE = make_tensor(make_gmem_ptr(reinterpret_cast<ElementAccum *>(params.softmax_lse_ptr) + row_offset_lse),
Shape<Int<kBlockM>>{}, Stride<_1>{});
// This layout maps row_offset_lse to {bidh, lse_size, bidb} or {bidh, bidb, lse_size}.
Layout flat_layout = make_layout(lse_size);
Layout orig_layout = make_layout(make_shape(params.seqlen_q, params.h, params.b));
auto transposed_stride = make_stride(params.b, params.seqlen_q * params.b, params.seqlen_q / params.ngroups);
Layout remapped_layout = make_layout(make_shape(params.seqlen_q, params.h, params.b), transposed_stride);
Layout final_layout = cute::composition(remapped_layout, cute::composition(orig_layout, flat_layout));
Tensor gLSE_unpadded = make_tensor(make_gmem_ptr(reinterpret_cast<ElementAccum *>(params.softmax_lse_ptr)), final_layout);
constexpr int kNLsePerThread = (kMaxSplits * kBlockM + kNThreads - 1) / kNThreads;
// Read the LSE values from gmem and store them in shared memory, then tranpose them.
constexpr int kRowsPerLoadLSE = kNThreads / kBlockM;
typedef __NATIVE_VECTOR__(1, ElementAccum) B32Type;
#pragma unroll
for (int l = 0; l < kNLsePerThread; ++l) {
const int row = l * kRowsPerLoadLSE + tidx / kBlockM;
const int col = tidx % kBlockM;
ElementAccum lse = (row < params.num_splits && col < lse_size - bidx * kBlockM) ? gLSEaccum(row, col) : -INFINITY;
if (row < kMaxSplits) { sLSE[row][col] = lse; }
}
flash::sync_threads();
Tensor lse_accum = make_tensor<ElementAccum>(Shape<Int<kNLsePerThread>>{});
constexpr int kRowsPerLoadTranspose = std::min(kRowsPerLoadLSE, kMaxSplits);
// To make sure that kMaxSplits is within 1 warp: we decide how many elements within kMaxSplits
// each thread should hold. If kMaxSplits = 16, then each thread holds 2 elements (128 threads,
// kBlockM rows, so each time we load we can load 128 / kBlockM rows).
// constexpr int kThreadsPerSplit = kMaxSplits / kRowsPerLoadTranspose;
// static_assert(kThreadsPerSplit <= 32);
//static_assert(kRowsPerLoadTranspose <= 32);
static_assert(kRowsPerLoadTranspose <= 64);
static_assert(kNLsePerThread * kRowsPerLoadTranspose <= kMaxSplits);
const int lse_base_row = tidx % kRowsPerLoadTranspose;
const int lse_base_col = tidx / kRowsPerLoadTranspose;
#pragma unroll
for (int l = 0; l < kNLsePerThread; ++l) {
const int row = l * kRowsPerLoadTranspose + lse_base_row;
const int col = lse_base_col;
lse_accum(l) = (row < kMaxSplits && col < kBlockM) ? sLSE[row][col] : -INFINITY;
// if (bidx == 0 && tidx < 32) { printf("tidx = %d, row = %d, col = %d, lse = %f\n", tidx, row, col, lse_accum(l)); }
}
// Compute the logsumexp of the LSE along the split dimension.
ElementAccum lse_max = lse_accum(0);
#pragma unroll
for (int l = 1; l < kNLsePerThread; ++l) { lse_max = max(lse_max, lse_accum(l)); }
MaxOp<float> max_op;
lse_max = Allreduce<kRowsPerLoadTranspose>::run(lse_max, max_op);
lse_max = lse_max == -INFINITY ? 0.0f : lse_max; // In case all local LSEs are -inf
float lse_sum = __expf(lse_accum(0) - lse_max);
#pragma unroll
for (int l = 1; l < kNLsePerThread; ++l) { lse_sum += __expf(lse_accum(l) - lse_max); }
SumOp<float> sum_op;
lse_sum = Allreduce<kRowsPerLoadTranspose>::run(lse_sum, sum_op);
// For the case where all local lse == -INFINITY, we want to set lse_logsum to INFINITY. Otherwise
// lse_logsum is log(0.0) = -INFINITY and we get NaN when we do lse_accum(l) - lse_logsum.
ElementAccum lse_logsum = (lse_sum == 0.f || lse_sum != lse_sum) ? INFINITY : __logf(lse_sum) + lse_max;
if (tidx % kRowsPerLoadTranspose == 0 && tidx / kRowsPerLoadTranspose < kBlockM) {
if (params.unpadded_lse) {
const index_t lse_offset = row_offset_lse + tidx / kRowsPerLoadTranspose;
if (lse_offset < lse_size) {
gLSE_unpadded(lse_offset) = lse_logsum;
}
} else {
gLSE(tidx / kRowsPerLoadTranspose) = lse_logsum;
}
}
// Store the scales exp(lse - lse_logsum) in shared memory.
#pragma unroll
for (int l = 0; l < kNLsePerThread; ++l) {
const int row = l * kRowsPerLoadTranspose + lse_base_row;
const int col = lse_base_col;
if (row < params.num_splits && col < kBlockM) { sLSE[row][col] = __expf(lse_accum(l) - lse_logsum); }
}
const index_t row_offset_oaccum = bidx * kBlockM * params.d_v;
Tensor gOaccum = make_tensor(make_gmem_ptr(reinterpret_cast<ElementAccum *>(params.oaccum_ptr) + row_offset_oaccum),
Shape<Int<kBlockM>, Int<kHeadDimV>>{},
Stride<Int<kHeadDimV>, _1>{});
constexpr int kBlockN = kNThreads / kBlockM;
using GmemLayoutAtomOaccum = Layout<Shape<Int<kBlockM>, Int<kBlockN>>, Stride<Int<kBlockN>, _1>>;
using GmemTiledCopyOaccum = decltype(
make_tiled_copy(Copy_Atom<DefaultCopy, ElementAccum>{},
GmemLayoutAtomOaccum{},
Layout<Shape < _1, _4>>{})); // Val layout, 4 vals per store
GmemTiledCopyOaccum gmem_tiled_copy_Oaccum;
auto gmem_thr_copy_Oaccum = gmem_tiled_copy_Oaccum.get_thread_slice(tidx);
Tensor tOgOaccum = gmem_thr_copy_Oaccum.partition_S(gOaccum);
Tensor tOrO = make_tensor<ElementAccum>(shape(tOgOaccum));
Tensor tOrOaccum = make_tensor<ElementAccum>(shape(tOgOaccum));
clear(tOrO);
flash::sync_threads();
typedef __NATIVE_VECTOR__(2, float) Float2;
// Predicates
Tensor cOaccum = make_identity_tensor(Shape<Int<kBlockM>, Int<kHeadDimV>>{});
// Repeat the partitioning with identity layouts
Tensor tOcOaccum = gmem_thr_copy_Oaccum.partition_S(cOaccum);
static_assert(decltype(size<0>(tOrOaccum))::value % 2 == 0);
// Load Oaccum in then scale and accumulate to O
for (int split = 0; split < params.num_splits; ++split) {
flash::copy_b128</*Is_even_MN=*/false, Is_even_K>(
tOgOaccum, tOrOaccum, tOcOaccum, params.d_v, lse_size - bidx * kBlockM
);
#pragma unroll
for (int m = 0; m < size<1>(tOrOaccum); ++m) {
int row = get<0>(tOcOaccum(0, m, 0));
ElementAccum lse_scale = sLSE[split][row];
Float2 lse_scale_vec = {lse_scale, lse_scale};
#pragma unroll
for (int k = 0; k < size<2>(tOrOaccum); ++k) {
#pragma unroll
for (int i = 0; i < size<0>(tOrOaccum); i += 2) {
Float2 x_vec = {tOrOaccum(i, m, k), tOrOaccum(i + 1, m, k)};
Float2 y_vec = {tOrO(i, m, k), tOrO(i + 1, m, k)};
y_vec = __builtin_mxc_pk_fma_f32(x_vec, lse_scale_vec, y_vec);
tOrO(i, m, k) = y_vec[0];
tOrO(i + 1, m, k) = y_vec[1];
}
}
}
tOgOaccum.data() = tOgOaccum.data() + lse_size * params.d_v;
}
//Tensor rO = flash::convert_type<Element>(tOrO);
CONVERT_TENSOR_TYPE(ElementAccum, Element, tOrO, rO)
const int q_head_offset = params.h * params.seqlen_q;
// Write to gO
#pragma unroll
for (int m = 0; m < size<1>(rO); ++m) {
const int idx = bidx * kBlockM + get<0>(tOcOaccum(0, m, 0));
const int batch_idx = idx / q_head_offset;
const int head_idx = (idx - batch_idx * q_head_offset) / params.seqlen_q;
// The index to the rows of Q
const int row = idx - batch_idx * q_head_offset - head_idx * params.seqlen_q;
auto o_ptr = reinterpret_cast<Element *>(params.o_ptr) + batch_idx * params.o_batch_stride
+ head_idx * params.o_head_stride + row * params.o_row_stride;
#pragma unroll
for (int k = 0; k < size<2>(rO); ++k) {
const int col = get<1>(tOcOaccum(0, m, k));
Tensor gO = make_tensor(make_gmem_ptr(o_ptr + col),
Shape<Int<decltype(size<0>(rO))::value>>{}, Stride<_1>{});
auto gO_ptr = reinterpret_cast<uint64_t *>(gO.data().ptr_);
auto rO_ptr = reinterpret_cast<uint64_t *>(rO(_, m, k).data().ptr_);
__builtin_mxc_stg_b64_predicator(gO_ptr, 0, rO_ptr[0], true, false, false, idx < lse_size && (Is_even_K || col < params.d_v), 1, MACA_ICMP_EQ);
}
}
}
} // namespace flash

View File

@ -1,5 +1,3 @@
// Adapted from Dao-AILab/flash-attention (https://github.com/Dao-AILab/flash-attention/tree/v2.6.3)
#pragma once
#include <cute/algorithm/copy.hpp>
@ -8,56 +6,25 @@
#include <mctlass/array.h>
#include <mctlass/numeric_types.h>
#include "block_info.h"
#include "kernel_traits.h"
#include "utils.h"
#include "softmax.h"
#include "mask.h"
#include "rotary.h"
#include "attn_mask.h"
#include "flash_fwd_split_kernel_k64_16x16_4waves.h"
#include "flash_fwd_split_kernel_k64_32x16_4waves.h"
#include "flash_fwd_split_kernel_k64_64x16_8waves.h"
namespace flash {
using namespace cute;
template<typename Kernel_traits, bool Is_causal, bool Is_local, bool Has_alibi, bool Is_even_MN, bool Is_even_K, bool Is_softcap, bool Split, typename Params>
__forceinline__ __device__ void compute_attn_splitkv(const Params &params, const int m_block_max) {
const int m_block = blockIdx.x;
// The block index for the batch.
const int bidb = Split ? blockIdx.z / params.h : blockIdx.y;
// The block index for the head.
const int bidh = Split ? blockIdx.z - bidb * params.h : blockIdx.z;
const int n_split_idx = Split ? blockIdx.y : 0;
const int num_n_splits = Split ? gridDim.y : 1;
if constexpr (Kernel_traits::kBlockM == 32 && Kernel_traits::kBlockN == 16 && Kernel_traits::kNWarps == 4) {
compute_attn_1rowblock_splitkv_k64_mla_32x16_4waves<Kernel_traits, Is_causal, Is_local, Has_alibi, Is_even_MN, Is_even_K, Is_softcap, Split>(
params, bidb, bidh, m_block, n_split_idx, num_n_splits);
} else if constexpr (Kernel_traits::kBlockM == 64 && Kernel_traits::kBlockN == 16 && Kernel_traits::kNWarps == 8) {
compute_attn_1rowblock_splitkv_k64_mla_64x16_8waves<Kernel_traits, Is_causal, Is_local, Has_alibi, Is_even_MN, Is_even_K, Is_softcap, Split>(
params, bidb, bidh, m_block, n_split_idx, num_n_splits);
} else if constexpr (Kernel_traits::kBlockM == 16 && Kernel_traits::kBlockN == 16 && Kernel_traits::kNWarps == 4) {
compute_attn_1rowblock_splitkv_k64_mla_16x16_4waves<Kernel_traits, Is_causal, Is_local, Has_alibi, Is_even_MN, Is_even_K, Is_softcap, Split>(
params, bidb, bidh, m_block, n_split_idx, num_n_splits);
}
}
////////////////////////////////////////////////////////////////////////////////////////////////////
template<typename Kernel_traits, int kBlockM, int Log_max_splits, bool Is_even_K, typename Params>
__forceinline__ __device__ void combine_attn_seqk_parallel(const Params &params) {
using Element = typename Kernel_traits::Element;
__forceinline__ __device__ void combine_attn_seqk_parallel_splitkv_mla(const Params &params) {
using ElementO = typename Kernel_traits::ElementO;
using ElementAccum = typename Kernel_traits::ElementAccum;
using index_t = typename Kernel_traits::index_t;
constexpr int kMaxSplits = 1 << Log_max_splits;
constexpr int kHeadDimV = Kernel_traits::kHeadDimV;
constexpr int kNThreads = 256;/*Kernel_traits::kNThreads*/;
constexpr int kNThreads = kBlockM == 1 ? 64 : (kBlockM == 2 ? 128 : 256);/*Kernel_traits::kNThreads*/;
static_assert(kMaxSplits <= 128, "kMaxSplits must be <= 128");
static_assert(kBlockM == 4 || kBlockM == 8 || kBlockM == 16 || kBlockM == 32, "kBlockM must be 4, 8, 16 or 32");
static_assert(kNThreads == 128 || kNThreads == 256, "We assume that each block has 128 or 256 threads");
// static_assert(kMaxSplits <= 128, "kMaxSplits must be <= 128");
static_assert(kBlockM == 1 || kBlockM == 2 || kBlockM == 4 || kBlockM == 8 || kBlockM == 16 || kBlockM == 32, "kBlockM must be 1, 2, 4, 8, 16 or 32");
static_assert(kNThreads == 64 || kNThreads == 128 || kNThreads == 256, "We assume that each block has 64, 128 or 256 threads");
// Shared memory.
// kBlockM + 1 instead of kBlockM to reduce bank conflicts.
@ -67,10 +34,20 @@ __forceinline__ __device__ void combine_attn_seqk_parallel(const Params &params)
const int tidx = threadIdx.x;
const int bidx = blockIdx.x;
const index_t lse_size = params.b * params.h * params.seqlen_q;
const int hs = params.h * params.seqlen_q;
const int batch_idx = (bidx * kBlockM) / hs;
const int hs_idx = (bidx * kBlockM) % hs;
const int split_offset = params.num_splits_ptr[batch_idx];
const int actual_num_splits = params.num_splits_ptr[batch_idx + 1] - split_offset;
FLASH_DEVICE_ASSERT(actual_num_splits <= kMaxSplits);
if (actual_num_splits == 1) return;
const index_t lse_size = hs;
const index_t row_offset_lseaccum = split_offset * lse_size + hs_idx;
const index_t row_offset_lse = bidx * kBlockM;
Tensor gLSEaccum = make_tensor(make_gmem_ptr(reinterpret_cast<ElementAccum *>(params.softmax_lseaccum_ptr) + row_offset_lse),
Tensor gLSEaccum = make_tensor(make_gmem_ptr(reinterpret_cast<ElementAccum *>(params.softmax_lseaccum_ptr) + row_offset_lseaccum),
Shape<Int<kMaxSplits>, Int<kBlockM>>{},
make_stride(lse_size, _1{}));
// LSE format is different depending on params.unpadded_lse and params.seqlenq_ngroups_swapped, see comment in get_lse_tile.
@ -96,7 +73,7 @@ __forceinline__ __device__ void combine_attn_seqk_parallel(const Params &params)
for (int l = 0; l < kNLsePerThread; ++l) {
const int row = l * kRowsPerLoadLSE + tidx / kBlockM;
const int col = tidx % kBlockM;
ElementAccum lse = (row < params.num_splits && col < lse_size - bidx * kBlockM) ? gLSEaccum(row, col) : -INFINITY;
ElementAccum lse = (row < actual_num_splits && col < lse_size - hs_idx) ? gLSEaccum(row, col) : -INFINITY;
if (row < kMaxSplits) { sLSE[row][col] = lse; }
}
@ -118,7 +95,6 @@ __forceinline__ __device__ void combine_attn_seqk_parallel(const Params &params)
const int row = l * kRowsPerLoadTranspose + lse_base_row;
const int col = lse_base_col;
lse_accum(l) = (row < kMaxSplits && col < kBlockM) ? sLSE[row][col] : -INFINITY;
// if (bidx == 0 && tidx < 32) { printf("tidx = %d, row = %d, col = %d, lse = %f\n", tidx, row, col, lse_accum(l)); }
}
// Compute the logsumexp of the LSE along the split dimension.
@ -151,10 +127,10 @@ __forceinline__ __device__ void combine_attn_seqk_parallel(const Params &params)
for (int l = 0; l < kNLsePerThread; ++l) {
const int row = l * kRowsPerLoadTranspose + lse_base_row;
const int col = lse_base_col;
if (row < params.num_splits && col < kBlockM) { sLSE[row][col] = __expf(lse_accum(l) - lse_logsum); }
if (row < actual_num_splits && col < kBlockM) { sLSE[row][col] = __expf(lse_accum(l) - lse_logsum); }
}
const index_t row_offset_oaccum = bidx * kBlockM * params.d_v;
const index_t row_offset_oaccum = (split_offset * hs + hs_idx) * params.d_v;
Tensor gOaccum = make_tensor(make_gmem_ptr(reinterpret_cast<ElementAccum *>(params.oaccum_ptr) + row_offset_oaccum),
Shape<Int<kBlockM>, Int<kHeadDimV>>{},
Stride<Int<kHeadDimV>, _1>{});
@ -170,7 +146,8 @@ __forceinline__ __device__ void combine_attn_seqk_parallel(const Params &params)
Tensor tOrO = make_tensor<ElementAccum>(shape(tOgOaccum));
Tensor tOrOaccum = make_tensor<ElementAccum>(shape(tOgOaccum));
clear(tOrO);
flash::sync_threads();
flash::sync_threads(); // first barrier
flash::barrier(); // second barrier
typedef __NATIVE_VECTOR__(2, float) Float2;
@ -180,9 +157,9 @@ __forceinline__ __device__ void combine_attn_seqk_parallel(const Params &params)
Tensor tOcOaccum = gmem_thr_copy_Oaccum.partition_S(cOaccum);
static_assert(decltype(size<0>(tOrOaccum))::value % 2 == 0);
// Load Oaccum in then scale and accumulate to O
for (int split = 0; split < params.num_splits; ++split) {
for (int split = 0; split < actual_num_splits; ++split) {
flash::copy_b128</*Is_even_MN=*/false, Is_even_K>(
tOgOaccum, tOrOaccum, tOcOaccum, params.d_v, lse_size - bidx * kBlockM
tOgOaccum, tOrOaccum, tOcOaccum, params.d_v, lse_size - hs_idx
);
#pragma unroll
for (int m = 0; m < size<1>(tOrOaccum); ++m) {
@ -201,21 +178,20 @@ __forceinline__ __device__ void combine_attn_seqk_parallel(const Params &params)
}
}
}
tOgOaccum.data() = tOgOaccum.data() + lse_size * params.d_v;
tOgOaccum.data() = tOgOaccum.data() + hs * params.d_v;
}
//Tensor rO = flash::convert_type<Element>(tOrO);
CONVERT_TENSOR_TYPE(ElementAccum, Element, tOrO, rO)
//Tensor rO = flash::convert_type<ElementO>(tOrO);
CONVERT_TENSOR_TYPE(ElementAccum, ElementO, tOrO, rO)
const int q_head_offset = params.h * params.seqlen_q;
// Write to gO
#pragma unroll
for (int m = 0; m < size<1>(rO); ++m) {
const int idx = bidx * kBlockM + get<0>(tOcOaccum(0, m, 0));
const int batch_idx = idx / q_head_offset;
const int head_idx = (idx - batch_idx * q_head_offset) / params.seqlen_q;
const int idx = hs_idx + get<0>(tOcOaccum(0, m, 0));
const int head_idx = idx / params.seqlen_q;
// The index to the rows of Q
const int row = idx - batch_idx * q_head_offset - head_idx * params.seqlen_q;
auto o_ptr = reinterpret_cast<Element *>(params.o_ptr) + batch_idx * params.o_batch_stride
const int row = idx % params.seqlen_q;
auto o_ptr = reinterpret_cast<ElementO *>(params.o_ptr) + batch_idx * params.o_batch_stride
+ head_idx * params.o_head_stride + row * params.o_row_stride;
#pragma unroll
for (int k = 0; k < size<2>(rO); ++k) {
@ -229,4 +205,4 @@ __forceinline__ __device__ void combine_attn_seqk_parallel(const Params &params)
}
}
} // namespace flash
}// namespace flash

View File

@ -0,0 +1,63 @@
#pragma once
#include <cute/algorithm/copy.hpp>
#include <mctlass/mctlass.h>
#include <mctlass/array.h>
#include <mctlass/numeric_types.h>
#include "block_info.h"
#include "kernel_traits.h"
#include "utils.h"
#include "softmax.h"
#include "mask.h"
#include "xcore1000/flash_fwd_sparse_mla_kernel_k64_64x16_8waves_xcore1000.h"
namespace flash {
using namespace cute;
template<typename Kernel_traits, bool Is_causal, bool Is_even_K, typename Params>
__forceinline__ __device__ void compute_attn_1rowblock_splitkv_sparse_mla(const Params &params, const int m_block_max) {
constexpr int kBlockN = Kernel_traits::kBlockN;
const int m_block = blockIdx.x;
const int bidh = blockIdx.y;
const int partition_idx = blockIdx.z;
extern __shared__ char shared_memory[];
int *tile_scheduler_metadata_ptr = params.tile_scheduler_metadata_ptr + partition_idx * TileSchedulerMetaDataSize;
int4 tile_scheduler_metadata = *reinterpret_cast<int4 *>(tile_scheduler_metadata_ptr);
int begin_idx = tile_scheduler_metadata.x;
int sched_begin_block_idx = tile_scheduler_metadata.y;
int end_idx = tile_scheduler_metadata.z;
int sched_end_block_idx = tile_scheduler_metadata.w;
if (begin_idx >= params.b || begin_idx < 0) return;
int begin_n_split_idx = tile_scheduler_metadata_ptr[4];
#pragma unroll 1
for (int batch_id = begin_idx; batch_id <= end_idx; ++batch_id) {
const int n_split_idx = batch_id == begin_idx ? begin_n_split_idx : 0;
const int seqlen_k = params.cu_seqlens_k[batch_id];
const int n_block_min = batch_id == begin_idx ? sched_begin_block_idx : 0;
const int n_block_max = batch_id == end_idx ? sched_end_block_idx : cute::ceil_div(params.topk, kBlockN);
// [n_block_min, n_block_max) need be calculated in kernel
if (n_block_max <= n_block_min) continue;
const bool NoSplit = __ldg(params.num_splits_ptr + batch_id + 1) - __ldg(params.num_splits_ptr + batch_id) == 1;
if (batch_id > begin_idx) {
__syncthreads(); // Barrier between two tiles.
}
// printf(
// "batch_id is %d, begin_idx is %d, end_idx is %d, params.b is %d, n_split_idx is %d, seqlen_k is %d, n_block_min is %d, n_block_max is %d, NoSplit is %d\n",
// batch_id, begin_idx, end_idx, params.b, n_split_idx, seqlen_k, n_block_min, n_block_max, NoSplit);
compute_attn_1rowblock_splitkv_sparse_mla_k64_64x16_8waves_xcore1000<Kernel_traits, Is_causal, Is_even_K>(
params, batch_id, bidh, m_block, n_split_idx, n_block_min, n_block_max, NoSplit);
}
}
}// namespace flash

View File

@ -17,33 +17,29 @@ using namespace cute;
template<int kHeadDim_, int kBlockM_, int kBlockN_, int kNWarps_, typename elem_type=mctlass::half_t>
struct Flash_kernel_traits {
#if defined(__MACA_ARCH__)
using Element = elem_type;
static constexpr bool Has_cp_async = false;
#else
using Element = mctlass::half_t;
static constexpr bool Has_cp_async = false;
#endif
using ElementAccum = float;
using index_t = int64_t;
#if defined(__MACA_ARCH__)
using MMA_Atom_Arch = std::conditional_t<
std::is_same_v<elem_type, mctlass::half_t>,
using MMA_Atom_Arch_16x16x16_fp16 = std::conditional_t<std::is_same_v<elem_type, mctlass::half_t>,
MMA_Atom<MACA_16x16x16_F32F16F16F32>,
MMA_Atom<MACA_16x16x16_F32BF16BF16F32>
>;
using MMA_Atom_Arch_16x16x32_fp16 = std::conditional_t<std::is_same_v<elem_type, mctlass::half_t>,
MMA_Atom<MACA_16x16x32_F32F16F16F32>,
MMA_Atom<MACA_16x16x32_F32BF16BF16F32>
>;
using MMA_Atom_Arch_16x16x32_i8 = MMA_Atom<MACA_16x16x32_I32I8I8I32>;
using ValLayoutMNK = Layout<Shape<_1, _1, _1>>;
#else
using MMA_Atom_Arch = MMA_Atom<SM75_16x8x8_F32F16F16F32_TN>;
using ValLayoutMNK = Layout<Shape<_1, _2, _2>>;
#endif
using SmemCopyAtom = Copy_Atom<DefaultCopy, elem_type>;
using SmemCopyAtomTransposed = Copy_Atom<DefaultCopy, elem_type>;
using SmemCopyB64 = Copy_Atom<UniversalCopy<uint64_t>, elem_type>;
using UniversalCopyAtom32 = Copy_Atom<UniversalCopy<uint32_t>, elem_type>;
using UniversalCopyAtomB32 = Copy_Atom<UniversalCopy<uint32_t>, elem_type>;
using UniversalCopyAtomB64 = Copy_Atom<UniversalCopy<uint64_t>, elem_type>;
using UniversalCopyAtomB128 = Copy_Atom<UniversalCopy<uint128_t>, elem_type>;
using LDSB64Trans4x16Atom = Copy_Atom<Copy_Traits<MACA_LDS_TRANS_4X16>, elem_type>;
};
// If Share_Q_K_smem is true, that forces Is_Q_in_regs to be true
@ -56,11 +52,17 @@ struct Flash_fwd_kernel_traits : public Base {
static constexpr bool Has_cp_async = Base::Has_cp_async;
using SmemCopyAtom = typename Base::SmemCopyAtom;
using SmemCopyAtomB64 = typename Base::SmemCopyB64;
using UniversalCopyAtom32 = typename Base::UniversalCopyAtom32;
using UniversalCopyAtomB32 = typename Base::UniversalCopyAtomB32;
using UniversalCopyAtomB64 = typename Base::UniversalCopyAtomB64;
using UniversalCopyAtomB128 = typename Base::UniversalCopyAtomB128;
using LDSB64Trans4x16Atom = typename Base::LDSB64Trans4x16Atom;
using SmemCopyAtomTransposed = typename Base::SmemCopyAtomTransposed;
using MMA_Atom_Arch_16x16x16_fp16 = typename Base::MMA_Atom_Arch_16x16x16_fp16;
using MMA_Atom_Arch_16x16x32_fp16 = typename Base::MMA_Atom_Arch_16x16x32_fp16;
using MMA_Atom_Arch_16x16x32_i8 = typename Base::MMA_Atom_Arch_16x16x32_i8;
using ElementO = Element;
static constexpr bool Share_Q_K_smem = Share_Q_K_smem_;
static constexpr bool Is_Q_in_regs = Is_Q_in_regs_ || Share_Q_K_smem;
static constexpr int Num_Stages = Num_Stages_;
@ -73,6 +75,8 @@ struct Flash_fwd_kernel_traits : public Base {
static constexpr int kBlockN = kBlockN_;
static constexpr int kHeadDim = kHeadDim_;
static constexpr int kHeadDimV = kHeadDimV_;
static constexpr int kHeadDimNope = kHeadDimV;
static constexpr int kHeadDimRope = kHeadDim - kHeadDimV;
static_assert(kHeadDim % 32 == 0);
static constexpr int kBlockKSmem = kHeadDim % 64 == 0 ? 64 : 32;
static constexpr int kBlockKSmemV = kHeadDimV % 64 == 0 ? 64 : 32;
@ -83,17 +87,55 @@ struct Flash_fwd_kernel_traits : public Base {
static constexpr int SShift_OPT = kBlockKSmem == 32 ? 3 : 4; // for bank conflict free
static constexpr int kAtomLayoutMS = std::min(kBlockM / 16, kNWarps);
static constexpr int kAtomLayoutMO = kAtomLayoutMS;
static constexpr int kBlockTopK = kNThreads * (16 / sizeof(int32_t));
using MMA_Atom_QK = MMA_Atom_Arch_16x16x16_fp16;
using MMA_Atom_PV = MMA_Atom_Arch_16x16x16_fp16;
using TiledMmaS = TiledMMA<
typename Base::MMA_Atom_Arch,
MMA_Atom_QK,
Layout<Shape<Int<kAtomLayoutMS>,_1,_1>>, // 2x1x1 or 4x1x1
typename Base::ValLayoutMNK>;
using TiledMmaS_16x16x32 = TiledMMA<
MMA_Atom_Arch_16x16x32_fp16,
Layout<Shape<Int<kAtomLayoutMS>,_1,_1>>, // 2x1x1 or 4x1x1
typename Base::ValLayoutMNK>;
using TiledMmaS_16x16x32_4x2 = TiledMMA<
MMA_Atom_Arch_16x16x32_fp16,
Layout<Shape<Int<kAtomLayoutMS>,Int<kNWarps / kAtomLayoutMS>,_1>>, // 4x2x1
typename Base::ValLayoutMNK>;
using TiledMmaO = TiledMMA<
typename Base::MMA_Atom_Arch,
MMA_Atom_PV,
Layout<Shape<Int<kAtomLayoutMO>,Int<kNWarps / kAtomLayoutMO>,_1>>, // 2x2x1 or 4x2x1
typename Base::ValLayoutMNK>;
using SmemLayoutAtomRowMax = decltype(
composition(Swizzle<0, 0, 0>{},
Layout<Shape<_1, Int<kBlockM>>,
Stride<Int<kBlockM>, _1>>{}));
using SmemLayoutRowMax = decltype(tile_to_shape(
SmemLayoutAtomRowMax{},
Shape<Int<kNWarps / kAtomLayoutMS>, Int<kBlockM>>{})); //rowmax 16 value per wave
using SmemLayoutAtomRowSum = decltype(
composition(Swizzle<0, 0, 0>{},
Layout<Shape<_1, Int<kBlockM>>,
Stride<Int<kBlockM>, _1>>{}));
using SmemLayoutRowSum = decltype(tile_to_shape(
SmemLayoutAtomRowSum{},
Shape<Int<kNWarps / kAtomLayoutMS>, Int<kBlockM>>{})); //rowsum 64 value per wave
using SmemLayoutAtomP = decltype(
composition(Swizzle<3, 2, 4>{},
Layout<Shape<Int<kBlockM>,Int<kBlockN>>,
Stride<Int<kBlockN>, _1>>{}));
using SmemLayoutP = decltype(tile_to_shape(
SmemLayoutAtomP{},
Shape<Int<kBlockM>,Int<kBlockN>>{})); //rowmax_wg0
using SmemLayoutAtomQ = decltype(
composition(Swizzle<kSwizzle, MBase, SShift>{},
// This has to be kBlockKSmem, using kHeadDim gives wrong results for d=128
@ -103,6 +145,10 @@ struct Flash_fwd_kernel_traits : public Base {
composition(Swizzle<4, 2, 4>{},
Layout<Shape<_16, Int<kBlockKSmem>>,
Stride<Int<kBlockKSmem>, _1>>{}));
using SmemLayoutQNoSwizzle = decltype(tile_to_shape(
Layout<Shape<_16, Int<kBlockKSmem>>,
Stride<Int<kBlockKSmem>, _1>>{},
Shape<Int<kBlockM>, Int<kHeadDim>>{}));
using SmemLayoutQ = decltype(tile_to_shape(
SmemLayoutAtomQ{},
Shape<Int<kBlockM>, Int<kHeadDim>>{}));
@ -122,31 +168,55 @@ struct Flash_fwd_kernel_traits : public Base {
SmemLayoutAtomQ{},
Shape<Int<kBlockN>, Int<kHeadDim>>{}));
using SmemLayoutAtomK = decltype(
using SmemLayoutAtomKNoswizzle = Layout<Shape<_16, Int<kBlockKSmem>>,
Stride<Int<kBlockKSmem>, _1>>;
using SmemLayoutAtomK424 = decltype(
composition(Swizzle<4, 2, 4>{},
Layout<Shape<_16, Int<kBlockKSmem>>,
Stride<Int<kBlockKSmem>, _1>>{}));
using SmemLayoutK = decltype(tile_to_shape(
SmemLayoutAtomK{},
using SmemLayoutAtomK242 = decltype(
composition(Swizzle<2, 4, 2>{},
Layout<Shape<_16, Int<kBlockKSmem>>,
Stride<Int<kBlockKSmem>, _1>>{}));
using SmemLayoutAtomK333 = decltype(
composition(Swizzle<3, 3, 3>{},
Layout<Shape<_16, Int<kBlockKSmem>>,
Stride<Int<kBlockKSmem>, _1>>{}));
using SmemLayoutK424 = decltype(tile_to_shape(
SmemLayoutAtomK424{},
Shape<Int<kBlockN>, Int<kHeadDim>, Int<Num_Stages>>{}));
using SmemLayoutK242 = decltype(tile_to_shape(
SmemLayoutAtomK242{},
Shape<Int<kBlockN>, Int<kHeadDim>, Int<Num_Stages>>{}));
using SmemLayoutK333 = decltype(tile_to_shape(
SmemLayoutAtomK333{},
Shape<Int<kBlockN>, Int<kHeadDim>, Int<Num_Stages>>{}));
using SmemLayoutKNoswizzle = decltype(tile_to_shape(
SmemLayoutAtomKNoswizzle{},
Shape<Int<kBlockN>, Int<kHeadDim>, Int<Num_Stages>>{}));
using SmemLayoutV = decltype(tile_to_shape(
SmemLayoutAtomQ{},
Shape<Int<kBlockN>, Int<kHeadDimV>>{}));
// This has to be kBlockN and not 8, otherwise we get wrong results for d=128
using SmemLayoutAtomVtransposedNoSwizzle = Layout<Shape<Int<kBlockKSmemV>, Int<kBlockN>>,
Stride<_1, Int<kBlockKSmemV>>>;
using SmemLayoutAtomVtransposed = decltype(
using SmemLayoutAtomVtransposedNoSwizzle = Layout<Shape<Int<kBlockKSmemV>, Int<kBlockN>, Int<Num_Stages>>,
Stride<_1, Int<kBlockKSmemV>, Int<kBlockN*kHeadDim>>>;
using SmemLayoutAtomVtransposed424 = decltype(
composition(Swizzle<4, 2, 4>{}, SmemLayoutAtomVtransposedNoSwizzle{}));
using SmemLayoutVtransposed = decltype(tile_to_shape(
SmemLayoutAtomVtransposed{},
Shape<Int<kHeadDimV>, Int<kBlockN>>{}));
using SmemLayoutVtransposed424 = decltype(tile_to_shape(
SmemLayoutAtomVtransposed424{},
Shape<Int<kHeadDimV>, Int<kBlockN>, Int<Num_Stages>>{}));
using SmemLayoutAtomVtransposed242 = decltype(
composition(Swizzle<2, 4, 2>{}, SmemLayoutAtomVtransposedNoSwizzle{}));
using SmemLayoutVtransposed242 = decltype(tile_to_shape(
SmemLayoutAtomVtransposed242{},
Shape<Int<kHeadDimV>, Int<kBlockN>, Int<Num_Stages>>{}));
// Maybe the VtransposeNoSwizzle just needs to have the right shape
// And the strides don't matter?
using SmemLayoutVtransposedNoSwizzle = decltype(tile_to_shape(
SmemLayoutAtomVtransposedNoSwizzle{},
Shape<Int<kHeadDimV>, Int<kBlockN>>{}));
Shape<Int<kHeadDimV>, Int<kBlockN>, Int<Num_Stages>>{}));
using SmemLayoutVtNoSwizzle = decltype(tile_to_shape(
Layout<Shape<_16, Int<kBlockKSmemV>>,
@ -160,15 +230,16 @@ struct Flash_fwd_kernel_traits : public Base {
using SmemLayoutO = decltype(tile_to_shape(
SmemLayoutAtomO{},
Shape<Int<kBlockM>, Int<kHeadDimV>>{}));
using SmemCopyAtomO = Copy_Atom<UniversalCopy<uint64_t>, Element>;
using SmemCopyAtomOb128 = Copy_Atom<UniversalCopy<uint128_t>, ElementO>;
using SmemCopyAtomO = Copy_Atom<UniversalCopy<uint64_t>, ElementO>;
using SmemCopyAtomOaccum = Copy_Atom<UniversalCopy<uint128_t>, ElementAccum>;
static constexpr int kSmemOSize = size(SmemLayoutO{}) * sizeof(ElementAccum);
static constexpr int kSmemQSize = size(SmemLayoutQ{}) * sizeof(Element);
static constexpr int kSmemKSize = size(SmemLayoutK{}) * sizeof(Element);
static constexpr int kSmemKSize = size(SmemLayoutK424{}) * sizeof(Element);
static constexpr int kSmemVSize = size(SmemLayoutV{}) * sizeof(Element);
static constexpr int kSmemKVSize = kSmemKSize + kSmemVSize;
static constexpr int kSmemSize = Share_Q_K_smem ? std::max(std::min((Is_Splits_ ? 2 : 1) * kSmemQSize, 64 * 1024), kSmemKSize) : kSmemQSize + kSmemKSize;
static constexpr int kRegSize = kSmemSize / sizeof(uint32_t) / kNThreads;
static constexpr int kSmemSize = Share_Q_K_smem ? std::max(std::max(kSmemQSize, kSmemKSize), kSmemOSize) : std::max(kSmemQSize + kSmemKSize, kSmemOSize);
static constexpr int kGmemElemsPerLoadB128 = sizeof(cute::uint128_t) / sizeof(Element);
static constexpr int kGmemElemsPerLoadB64 = sizeof(cute::uint64_t) / sizeof(Element);
@ -196,30 +267,37 @@ struct Flash_fwd_kernel_traits : public Base {
using GmemLayoutAtomB32 = Layout<Shape <Int<kNThreads / kGmemThreadsPerRowB32>, Int<kGmemThreadsPerRowB32>>,
Stride<Int<kGmemThreadsPerRowB32>, _1>>;
static constexpr int kGmemThreadsPerRowV = kBlockKSmemV / kGmemElemsPerLoadB128;
static_assert(kNThreads % kGmemThreadsPerRowV == 0, "kNThreads must be a multiple of kGmemThreadsPerRow");
using GmemLayoutAtomV = Layout<Shape <Int<kNThreads / kGmemThreadsPerRowV>, Int<kGmemThreadsPerRowV>>,
Stride<Int<kGmemThreadsPerRowV>, _1>>;
static constexpr int kGmemElemsPerLoadB128O = sizeof(cute::uint128_t) / sizeof(ElementO);
static constexpr int kGmemThreadsPerRowO = kBlockKSmemV / kGmemElemsPerLoadB128O;
static_assert(kNThreads % kGmemThreadsPerRowO == 0, "kNThreads must be a multiple of kGmemThreadsPerRowO");
static constexpr bool UseWarpsNx1 = kBlockM % (kNThreads / kGmemThreadsPerRowO) == 0;
using GmemLayoutAtomO = std::conditional_t<
UseWarpsNx1,
Layout<Shape <Int<kNThreads / kGmemThreadsPerRowO>, Int<kGmemThreadsPerRowO>>,
Stride<Int<kGmemThreadsPerRowO>, _1>>,
Layout<Shape<Int<kBlockM>, Shape<Int<kGmemThreadsPerRowO>, Int<kNThreads / kGmemThreadsPerRowO / kBlockM>>>,
Stride<Int<kGmemThreadsPerRowO>, Stride<_1, Int<kBlockM * kGmemThreadsPerRowO>>>>
>;
// We use CACHEGLOBAL instead of CACHEALWAYS for both Q and K/V, since we won't be reading
// from the same address by the same threadblock. This is slightly faster.
using GmemTiledCopyB128 = decltype(
make_tiled_copy(Copy_Atom<UniversalCopy<uint128_t>, Element>{},
GmemLayoutAtomB128{},
Layout<Shape<_1, _8>>{})); // Val layout, 8 vals per read
Layout<Shape<_1, Int<kGmemElemsPerLoadB128>>>{})); // Val layout, 8 vals per read
using GmemTiledCopyB64 = decltype(
make_tiled_copy(Copy_Atom<UniversalCopy<uint64_t>, Element>{},
GmemLayoutAtomB64{},
Layout<Shape<_1, _4>>{})); // Val layout, 4 vals per read
Layout<Shape<_1, Int<kGmemElemsPerLoadB64>>>{})); // Val layout, 4 vals per read
using GmemTiledCopyB32 = decltype(
make_tiled_copy(Copy_Atom<UniversalCopy<uint32_t>, Element>{},
GmemLayoutAtomB32{},
Layout<Shape<_1, _2>>{})); // Val layout, 2 vals per read
Layout<Shape<_1, Int<kGmemElemsPerLoadB32>>>{})); // Val layout, 2 vals per read
using GmemTiledCopyO = decltype(
make_tiled_copy(Copy_Atom<UniversalCopy<uint128_t>, Element>{},
GmemLayoutAtomV{},
Layout<Shape<_1, _8>>{})); // Val layout, 8 vals per store
make_tiled_copy(Copy_Atom<UniversalCopy<uint128_t>, ElementO>{},
GmemLayoutAtomO{},
Layout<Shape<_1, Int<kGmemElemsPerLoadB128O>>>{})); // Val layout, 8 vals per store
using GmemLayoutAtomOaccum = std::conditional_t<
kBlockKSmem == 32,

View File

@ -0,0 +1,470 @@
// Adapted from Dao-AILab/flash-attention (https://github.com/Dao-AILab/flash-attention/tree/v2.6.3)
#pragma once
#include <cute/algorithm/copy.hpp>
#include <mctlass/mctlass.h>
#include <mctlass/array.h>
#include <mctlass/numeric_types.h>
#include "block_info.h"
#include "kernel_traits.h"
#include "utils.h"
#include "softmax.h"
#include "mask.h"
namespace flash {
using namespace cute;
template<typename Kernel_traits, bool Split, bool Is_even_MN, bool Is_even_K, typename AccO, typename Softmax, typename Params>
__forceinline__ __device__ void store_16x16(const Params &params, const int bidb, const int bidh, const int m_block, const int n_split_idx, char* smem_, AccO acc_o, Softmax softmax ) {
using ElementAccum = typename Kernel_traits::ElementAccum;
using Element = typename Kernel_traits::Element;
using index_t = typename Kernel_traits::index_t;
using GmemTiledCopyO = std::conditional_t<
!Split,
typename Kernel_traits::GmemTiledCopyO,
typename Kernel_traits::GmemTiledCopyOaccum
>;
using ElementO = std::conditional_t<!Split, Element, ElementAccum>;
// The thread index.
const int tidx = threadIdx.x;
const int warp_idx = tidx / 64;
const int lane_idx = tidx % 64;
constexpr int kBlockM = Kernel_traits::kBlockM;
constexpr int kBlockN = Kernel_traits::kBlockN;
constexpr int kNWarps = Kernel_traits::kNWarps;
constexpr int kHeadDimV = Kernel_traits::kHeadDimV;
const BlockInfo</*Varlen=*/!Is_even_MN> binfo(params, bidb);
typename Kernel_traits::TiledMmaO tiled_mma_o;
auto thr_mma_o = tiled_mma_o.get_thread_slice(tidx);
Tensor lse = softmax.template normalize_softmax_lse</*Is_dropout=*/false, /*Return_lse*/true, Split>(acc_o, params.scale_softmax);
Tensor acc_o_view = make_tensor(acc_o.data(), make_layout(Shape<_4, Shape<_4, _2>>{},
Stride<_1, Shape<_4, _16>>{}));
Tensor acc_o_copy = make_fragment_like(acc_o_view);
#pragma unroll
for (int k = 0; k < size<1, 1>(acc_o_view); k++) {
#pragma unroll
for (int idx = 0; idx < 16; idx++) {
int row = idx / 4;
int col = idx % 4;
acc_o_copy(row, make_coord(col, k)) = acc_o_view(col, make_coord(row, k));
}
}
if constexpr (!Split) {
Tensor sOaccum = make_tensor(make_smem_ptr(reinterpret_cast<ElementO *>(smem_)), typename Kernel_traits::SmemLayoutO{}); // (SMEM_M,SMEM_N)
// Partition sO to match the accumulator partitioning
using SmemTiledCopyO = typename Kernel_traits::SmemCopyAtomO;
CONVERT_TENSOR_TYPE(ElementAccum, ElementO, acc_o_copy, rO)
int warp_offset = warp_idx * 16 * 64;
int thread_offset = lane_idx % 16 * 64 + lane_idx / 16 * 16;
ElementO *Osmem_ptr_sts = reinterpret_cast<ElementO *>(smem_) + warp_offset + thread_offset;
Tensor tOsO = make_tensor(make_smem_ptr(Osmem_ptr_sts), make_layout(Shape<_16, _2>{},
Stride<_1, Int<16*64*kNWarps>>{}));
Tensor tOrO = make_tensor(rO.data(), make_layout(Shape<_16, _2>{},
Stride<_1, _16>{}));
if constexpr (Kernel_traits::Share_Q_K_smem) { flash::sync_threads(); }
cute::copy(tOrO, tOsO);
const index_t row_offset_o = binfo.q_offset(params.o_batch_stride, params.o_row_stride, bidb)
+ m_block * kBlockM * params.o_row_stride + bidh * params.o_head_stride;
const index_t row_offset_lseaccum = ((n_split_idx * params.b + bidb) * params.h + bidh) * params.seqlen_q + m_block * kBlockM;
Tensor gOaccum = make_tensor(make_gmem_ptr(reinterpret_cast<ElementO *>(params.o_ptr) + (row_offset_o)),
Shape<Int<kBlockM>, Int<kHeadDimV>>{},
make_stride(params.o_row_stride, _1{}));
Tensor gLSEaccum = make_tensor(make_gmem_ptr(reinterpret_cast<ElementAccum *>(params.softmax_lse_ptr) + row_offset_lseaccum),
Shape<Int<kBlockM>>{}, Stride<_1>{});
GmemTiledCopyO gmem_tiled_copy_Oaccum;
auto gmem_thr_copy_Oaccum = gmem_tiled_copy_Oaccum.get_thread_slice(tidx);
Tensor tOsOaccum = gmem_thr_copy_Oaccum.partition_S(sOaccum); // ((Atom,AtomNum),ATOM_M,ATOM_N)
Tensor tOgOaccum = gmem_thr_copy_Oaccum.partition_D(gOaccum);
flash::sync_threads();
Tensor tOrOaccum = make_tensor<ElementO>(shape(tOgOaccum));
cute::copy(gmem_tiled_copy_Oaccum, tOsOaccum, tOrOaccum);
Tensor caccO = make_identity_tensor(Shape<Int<kBlockM>, Int<kHeadDimV>>{}); // (BLK_M,BLK_K) -> (blk_m,blk_k)
Tensor taccOcO = thr_mma_o.partition_C(caccO); // (MMA,MMA_M,MMA_K)
static_assert(decltype(size<0>(taccOcO))::value == 4);
// Convert to ((2, 2), MMA_M, MMA_K) then take only the row indices.
Tensor taccOcO_row = logical_divide(taccOcO, Shape<_4>{})(make_coord(0, _), _, 0);
CUTE_STATIC_ASSERT_V(size(lse) == size(taccOcO_row)); // MMA_M
if (get<1>(taccOcO_row(0)) == 0) {
#pragma unroll
for (int mi = 0; mi < size(lse); ++mi) {
const int row = get<0>(taccOcO_row(mi));
if (row < binfo.actual_seqlen_q - m_block * kBlockM) { gLSEaccum(row) = lse(mi); }
}
}
// Construct identity layout for sO
Tensor cO = make_identity_tensor(make_shape(size<0>(sOaccum), size<1>(sOaccum))); // (BLK_M,BLK_K) -> (blk_m,blk_k)
// Repeat the partitioning with identity layouts
Tensor tOcO = gmem_thr_copy_Oaccum.partition_D(cO); // (ACPY,ACPY_M,ACPY_K) -> (blk_m,blk_k)
// Clear_OOB_K must be false since we don't want to write zeros to gmem
flash::copy_reg_to_global<Is_even_MN, Is_even_K>(
tOrOaccum, tOgOaccum, tOcO, params.d_v, binfo.actual_seqlen_q - m_block * kBlockM
);
} else {
const int split_offset = params.num_splits_ptr[bidb];
Tensor sOaccum = make_tensor(make_smem_ptr(reinterpret_cast<ElementO *>(smem_)), typename Kernel_traits::SmemLayoutO{}); // (SMEM_M,SMEM_N)
const index_t row_offset_oaccum = (((split_offset + n_split_idx) * params.h + bidh) * params.seqlen_q
+ m_block * kBlockM) * params.d_v;
const index_t row_offset_lseaccum = ((split_offset + n_split_idx) * params.h + bidh) * params.seqlen_q + m_block * kBlockM;
Tensor gOaccum = make_tensor(make_gmem_ptr(reinterpret_cast<ElementO *>(params.oaccum_ptr) + row_offset_oaccum),
Shape<Int<kBlockM>, Int<kHeadDimV/2>>{},
make_stride(kHeadDimV, _1{}));
Tensor gLSEaccum = make_tensor(make_gmem_ptr(reinterpret_cast<ElementAccum *>(params.softmax_lseaccum_ptr) + row_offset_lseaccum),
Shape<Int<kBlockM>>{}, Stride<_1>{});
Tensor taccOrOaccum = make_tensor(acc_o_copy.data(), acc_o_copy.layout());
int warp_offset = warp_idx * 16 * 64;
int thread_offset = lane_idx % 16 * 64 + lane_idx / 16 * 16;
ElementO *accOsmem_ptr_sts = reinterpret_cast<ElementO *>(smem_) + warp_offset + thread_offset;
Tensor taccOsOaccum = make_tensor(make_smem_ptr(accOsmem_ptr_sts), make_layout(Shape<_4, _4>{},
Stride<_1, _4>{}));
if constexpr (Kernel_traits::Share_Q_K_smem) { flash::sync_threads(); }
int O_swizzle_row_sts = tidx % 4;
#pragma unroll
for (int i = 0; i < 4; i++) {
cute::copy(taccOrOaccum(_, make_coord(i, 0)), taccOsOaccum(_, O_swizzle_row_sts ^ i));
}
GmemTiledCopyO gmem_tiled_copy_Oaccum;
auto gmem_thr_copy_Oaccum = gmem_tiled_copy_Oaccum.get_thread_slice(tidx);
// Tensor tOsOaccum = gmem_thr_copy_Oaccum.partition_S(sOaccum); // ((Atom,AtomNum),ATOM_M,ATOM_N)
int O_swizzle_row_lds = tidx / 16 % 4;
int O_swizzle_col_lds = tidx % 16 % 4;
int O_swizzle_col_lds_new = O_swizzle_col_lds ^ O_swizzle_row_lds;
ElementO *accOsmem_ptr_lds = reinterpret_cast<ElementO *>(smem_) + (tidx + O_swizzle_col_lds_new - O_swizzle_col_lds) * 4;
Tensor tOsOaccum = make_tensor(make_smem_ptr(accOsmem_ptr_lds), make_layout(Shape<_4, _1, Int<kHeadDimV/2/64>>{},
Stride<_1, _0, Int<16*64>>{}));
Tensor tOgOaccum = gmem_thr_copy_Oaccum.partition_D(gOaccum);
flash::sync_threads();
Tensor tOrOaccum = make_tensor<ElementO>(shape(tOgOaccum));
cute::copy(gmem_tiled_copy_Oaccum, tOsOaccum, tOrOaccum);
flash::sync_threads();
Tensor caccO = make_identity_tensor(Shape<Int<kBlockM>, Int<kHeadDimV>>{}); // (BLK_M,BLK_K) -> (blk_m,blk_k)
Tensor taccOcO = thr_mma_o.partition_C(caccO); // (MMA,MMA_M,MMA_K)
static_assert(decltype(size<0>(taccOcO))::value == 4);
// Convert to ((2, 2), MMA_M, MMA_K) then take only the row indices.
Tensor taccOcO_row = logical_divide(taccOcO, Shape<_4>{})(make_coord(0, _), _, 0);
CUTE_STATIC_ASSERT_V(size(lse) == size(taccOcO_row)); // MMA_M
if (get<1>(taccOcO_row(0)) == 0) {
#pragma unroll
for (int mi = 0; mi < size(lse); ++mi) {
const int row = get<0>(taccOcO_row(mi));
if (row < binfo.actual_seqlen_q - m_block * kBlockM) { gLSEaccum(row) = lse(mi); }
}
}
Tensor cO = make_identity_tensor(make_shape(size<0>(sOaccum), size<1>(sOaccum))); // (BLK_M,BLK_K) -> (blk_m,blk_k)
// Repeat the partitioning with identity layouts
Tensor tOcaccO = gmem_thr_copy_Oaccum.partition_D(cO); // (ACPY,ACPY_M,ACPY_K) -> (blk_m,blk_k)
// Clear_OOB_K must be false since we don't want to write zeros to gmem
flash::copy_reg_to_global<Is_even_MN, Is_even_K>(
tOrOaccum, tOgOaccum, tOcaccO, params.d_v, binfo.actual_seqlen_q - m_block * kBlockM
);
#pragma unroll
for (int i = 0; i < 4; i++) {
cute::copy(taccOrOaccum(_, make_coord(i, 1)), taccOsOaccum(_, O_swizzle_row_sts ^ i));
}
flash::sync_threads();
cute::copy(gmem_tiled_copy_Oaccum, tOsOaccum, tOrOaccum);
tOgOaccum.data() = tOgOaccum.data() + (kHeadDimV/2);
flash::copy_reg_to_global<Is_even_MN, Is_even_K>(
tOrOaccum, tOgOaccum, tOcaccO, params.d_v, binfo.actual_seqlen_q - m_block * kBlockM
);
}
}
template<typename Kernel_traits, bool Is_causal, bool Is_even_MN, bool Is_even_K, bool Is_enable_dcp, typename Params>
__forceinline__ __device__ void compute_attn_1rowblock_splitkv_mla_k64_16x16_4waves_xcore1000(
const Params &params, const int bidb, const int bidh, const int m_block, const int n_split_idx, const int n_block_min, const int n_block_max, const bool NoSplit) {
using Element = typename Kernel_traits::Element;
using ElementAccum = typename Kernel_traits::ElementAccum;
using index_t = typename Kernel_traits::index_t;
// Shared memory.
extern __shared__ char smem_[];
// The thread index.
const int tidx = threadIdx.x;
const int warp_idx = tidx / 64;
const int lane_idx = tidx % 64;
constexpr int kBlockM = Kernel_traits::kBlockM;
constexpr int kBlockN = Kernel_traits::kBlockN;
constexpr int kHeadDim = Kernel_traits::kHeadDim;
constexpr int kHeadDimV = Kernel_traits::kHeadDimV;
constexpr int kNWarps = Kernel_traits::kNWarps;
constexpr int kBlockKSmem = Kernel_traits::kBlockKSmem;
constexpr int kAtomLayoutMS = Kernel_traits::kAtomLayoutMS;
constexpr int kAtomLayoutMO = Kernel_traits::kAtomLayoutMO;
constexpr int Num_Stages = Kernel_traits::Num_Stages;
constexpr int kSmemSize = Kernel_traits::kSmemSize;
static_assert(kBlockKSmem == 64);
const BlockInfo</*Varlen=*/!Is_even_MN> binfo(params, bidb);
if (m_block * kBlockM >= binfo.actual_seqlen_q) return;
const int actual_seqlen_k = params.cu_seqlens_k[bidb];
// We iterate over the blocks in reverse order. This is because the last block is the only one
// that needs masking when we read K and V from global memory. Moreover, iterating in reverse
// might save us 1 register (we just need n_block instead of both n_block and n_block_max).
const index_t row_offset_q = binfo.q_offset(params.q_batch_stride, params.q_row_stride, bidb)
+ m_block * kBlockM * params.q_row_stride + bidh * params.q_head_stride;
// We move K and V to the last block.
const int bidb_cache = params.cache_batch_idx == nullptr ? bidb : params.cache_batch_idx[bidb];
const int *block_table = params.block_table == nullptr ? nullptr : params.block_table + bidb * params.block_table_batch_stride;
const int block_table_idx = block_table == nullptr ? 0 : (n_block_max - 1) * kBlockN / params.page_block_size;
const int block_table_offset = block_table == nullptr ? 0 : (n_block_max - 1) * kBlockN - block_table_idx * params.page_block_size;
const index_t row_offset_k = block_table == nullptr
? binfo.k_offset(params.k_batch_stride, params.k_row_stride, bidb_cache)
+ (n_block_max - 1) * kBlockN * params.k_row_stride + (bidh / params.h_h_k_ratio) * params.k_head_stride
: (bidh / params.h_h_k_ratio) * params.k_head_stride;
const index_t row_offset_v = block_table == nullptr
? binfo.k_offset(params.v_batch_stride, params.v_row_stride, bidb_cache)
+ (n_block_max - 1) * kBlockN * params.v_row_stride + (bidh / params.h_h_k_ratio) * params.v_head_stride
: (bidh / params.h_h_k_ratio) * params.v_head_stride;
Tensor gQ = make_tensor(make_gmem_ptr(reinterpret_cast<Element *>(params.q_ptr) + row_offset_q),
Shape<Int<kBlockM>, Int<kHeadDim>>{},
make_stride(params.q_row_stride, _1{}));
Tensor gK = make_tensor(make_gmem_ptr(reinterpret_cast<Element *>(params.k_ptr) + row_offset_k),
Shape<Int<kBlockN>, Int<kHeadDim>>{},
make_stride(params.k_row_stride, _1{}));
Tensor gV = make_tensor(make_gmem_ptr(reinterpret_cast<Element *>(params.v_ptr) + row_offset_v),
Shape<Int<kBlockN>, Int<kHeadDimV>>{},
make_stride(params.v_row_stride, _1{}));
Tensor sQ = make_tensor(make_smem_ptr(reinterpret_cast<Element *>(smem_)),
typename Kernel_traits::SmemLayoutQ424{});
Tensor sK = make_tensor(sQ.data() + (Kernel_traits::Share_Q_K_smem ? 0 : size(sQ)),
typename Kernel_traits::SmemLayoutK424{});
Tensor sV = make_tensor(sK.data(), typename Kernel_traits::SmemLayoutVtNoSwizzle{});
Tensor sVt = make_tensor(sV.data(), typename Kernel_traits::SmemLayoutVtransposed424{});
Tensor sVtNoSwizzle = make_tensor(sV.data(), typename Kernel_traits::SmemLayoutVtransposedNoSwizzle{});
typename Kernel_traits::GmemTiledCopyB64 gmem_tiled_copy_Q;
auto gmem_thr_copy_Q = gmem_tiled_copy_Q.get_thread_slice(tidx);
Tensor tQgQ = gmem_thr_copy_Q.partition_S(gQ);
Tensor tQsQ = gmem_thr_copy_Q.partition_D(sQ);
typename Kernel_traits::GmemTiledCopyB64 gmem_tiled_copy_KV;
auto gmem_thr_copy_KV = gmem_tiled_copy_KV.get_thread_slice(tidx);
Tensor tKgK = gmem_thr_copy_KV.partition_S(gK); // (KCPY, KCPY_N, KCPY_K)
Tensor tKsK = gmem_thr_copy_KV.partition_D(sK);
// S is only 16x16 size, so all 4 waves compute the same S
int tidx_mma_s = tidx & 0x3F;
typename Kernel_traits::TiledMmaS tiled_mma_s;
auto thr_mma_s = tiled_mma_s.get_thread_slice(tidx_mma_s);
Tensor tSrQ = thr_mma_s.partition_fragment_A(sQ); // (MMA,MMA_M,MMA_K)
Tensor tSrK = thr_mma_s.partition_fragment_B(sK(_, _, 0)); // (MMA,MMA_N,MMA_K)
typename Kernel_traits::TiledMmaO tiled_mma_o;
auto thr_mma_o = tiled_mma_o.get_thread_slice(tidx);
// Tensor tOrVt = thr_mma_o.partition_fragment_B(sVt); // (MMA, MMA_K,MMA_N)
Tensor tOrVt = make_tensor<Element>(Shape<_4, Shape<_4, _2>, _1>{});
Tensor acc_o = partition_fragment_C(tiled_mma_o, Shape<Int<kBlockM>, Int<kHeadDimV>>{}); // MMA, MMA_M, MMA_K
//
// Copy Atom retiling
//
auto smem_tiled_copy_Q = make_tiled_copy_A(typename Kernel_traits::UniversalCopyAtomB64{}, tiled_mma_s);
auto smem_thr_copy_Q = smem_tiled_copy_Q.get_thread_slice(tidx_mma_s);
Tensor tSsQ = smem_thr_copy_Q.partition_S(sQ);
auto smem_tiled_copy_K = make_tiled_copy_B(typename Kernel_traits::UniversalCopyAtomB64{}, tiled_mma_s);
auto smem_thr_copy_K = smem_tiled_copy_K.get_thread_slice(tidx_mma_s);
Tensor tSsK = smem_thr_copy_K.partition_S(sK);
auto smem_tiled_copy_V = make_tiled_copy_B(typename Kernel_traits::SmemCopyAtomTransposed{}, tiled_mma_o);
auto smem_thr_copy_V = smem_tiled_copy_V.get_thread_slice(tidx);
int warp_offset = warp_idx / kAtomLayoutMO * 16 * 64;
int thread_offset = lane_idx / 16 * 4 * 64;
Element *Vtsmem_ptr_lds = reinterpret_cast<Element *>(sVt.data().get()) + warp_offset + thread_offset;
Tensor tOsVt = make_tensor(make_smem_ptr(Vtsmem_ptr_lds), make_layout(Shape<_4, _2, Int<Num_Stages>>{}, // MMA MMA_N NUM_STAGES
Stride<_1, Int<16*256>, Int<kBlockN*kHeadDim>>{}));
// PREDICATES
// Construct identity layout for sQ and sK
Tensor cQ = make_identity_tensor(make_shape(size<0>(sQ), size<1>(sQ))); // (BLK_M,BLK_K) -> (blk_m,blk_k)
Tensor cKV = make_identity_tensor(make_shape(size<0>(sK), size<1>(sK))); // (BLK_N,BLK_K) -> (blk_n,blk_k)
// Repeat the partitioning with identity layouts
Tensor tQcQ = gmem_thr_copy_Q.partition_S(cQ); // (ACPY,ACPY_M,ACPY_K) -> (blk_m,blk_k)
Tensor tKVcKV = gmem_thr_copy_KV.partition_S(cKV); // (BCPY,BCPY_N,BCPY_K) -> (blk_n,blk_k)
// Prologue
// Read Q from gmem to smem, optionally apply rotary embedding.
Tensor tQrQ = make_fragment_like(tQgQ);
// We don't need to clear the sQ smem tiles since we'll only write out the valid outputs
flash::copy_b64<Is_even_MN, Is_even_K>(tQgQ, tQrQ, tQcQ, params.d, binfo.actual_seqlen_q - m_block * kBlockM);
cute::copy(tQrQ, tQsQ);
if constexpr (Kernel_traits::Is_Q_in_regs) {
flash::sync_threads();
cute::copy(smem_tiled_copy_Q, tSsQ, tSrQ);
flash::sync_threads();
}
int n_block = n_block_max - 1;
// We don't need to clear the sK smem tiles since we'll mask out the scores anyway.
Tensor tKrK = make_fragment_like(tKgK);
int row_offset = tidx / 16 + n_block * kBlockN;
int virtual_page_idx = row_offset / params.page_block_size;
int page_offset = row_offset - virtual_page_idx * params.page_block_size;
int32_t page_idx = block_table[virtual_page_idx];
flash::copy_b64_page_one<Kernel_traits, Is_even_MN, Is_even_K>(gK, tKgK, tKrK, tKVcKV, params.d, n_block,
block_table, params.k_batch_stride, params.k_row_stride, params.page_block_size, page_idx, page_offset, actual_seqlen_k - n_block * kBlockN);
clear(acc_o);
flash::Softmax<size<1>(acc_o)> softmax;
flash::Mask<Is_causal, Is_enable_dcp> mask(actual_seqlen_k, binfo.actual_seqlen_q, params.ngroups, binfo.tot_seqlen_k, params.cp_world_size, params.cp_rank);
// For performance reason, we separate out two kinds of iterations:
// those that need masking on S, and those that don't.
// We need masking on S for the very last block when K and V has length not multiple of kBlockN.
// We also need masking on S if it's causal, for the last ceil_div(kBlockM, kBlockN) blocks.
// We will have at least 1 "masking" iteration.
// If not even_N, then seqlen_k might end in the middle of a block. In that case we need to
// mask 2 blocks (e.g. when kBlockM == kBlockN), not just 1.
constexpr int n_masking_steps = (!Is_causal)
? 1
: ((Is_even_MN && Is_causal) ? cute::ceil_div(kBlockM, kBlockN) : cute::ceil_div(kBlockM, kBlockN) + 1);
#pragma unroll
for (int masking_step = 0; masking_step < n_masking_steps; ++masking_step, --n_block) {
if (n_block > n_block_min) {
// prefetch load page index
row_offset -= kBlockN;
virtual_page_idx = row_offset / params.page_block_size;
page_offset = row_offset - virtual_page_idx * params.page_block_size;
page_idx = block_table[virtual_page_idx];
}
Tensor acc_s = partition_fragment_C(tiled_mma_s, Shape<Int<kBlockM>, Int<kBlockN>>{}); // (MMA=4, MMA_M, MMA_N)
flash::sync_threads();
cute::copy(tKrK, tKsK(_, _, _, 0));
clear(acc_s);
flash::sync_threads();
if (n_block > n_block_min) {
flash::copy_b64_page_one<Kernel_traits, /*Is_even_MN=*/true, Is_even_K>(gK, tKgK, tKrK, tKVcKV, params.d, n_block - 1,
block_table, params.k_batch_stride, params.k_row_stride, params.page_block_size, page_idx, page_offset);
}
flash::gemm</*A_in_regs=*/Kernel_traits::Is_Q_in_regs>(
acc_s, tSrQ, tSrK, tSsQ, tSsK(_, _, _, 0), tiled_mma_s, smem_tiled_copy_Q, smem_tiled_copy_K,
smem_thr_copy_Q, smem_thr_copy_K
);
mask.template apply_mask<Is_causal, Is_even_MN>(
acc_s, n_block * kBlockN, m_block * kBlockM + (tidx / 64) % kAtomLayoutMS * 16 + (tidx & 0xf), kAtomLayoutMS * 16
);
// We have key_padding_mask so we'll need to Check_inf
masking_step == 0
? softmax.template softmax_rescale_o</*Is_first=*/true, /*Check_inf=*/Is_causal || !Is_even_MN, true, true>(acc_s, acc_o, params.scale_softmax_log2)
: softmax.template softmax_rescale_o</*Is_first=*/false, /*Check_inf=*/Is_causal || !Is_even_MN, true, true>(acc_s, acc_o, params.scale_softmax_log2);
// Convert acc_s from fp32 to fp16/bf16
//Tensor rP = flash::convert_type<Element>(acc_s);
CONVERT_TENSOR_TYPE(ElementAccum, Element, acc_s, rP)
// Reshape rP from (MMA=4, MMA_M, MMA_N) to ((4, 2), MMA_M, MMA_N / 2)
// if using m16n8k16 or (4, MMA_M, MMA_N) if using m16n8k8.
//Tensor tOrP = make_tensor(rP.data(), flash::convert_layout_acc_Aregs<Kernel_traits::TiledMma>(rP.layout()));
// flash::sync_threads();
lds4x4_with_swizzle424(tOsVt(_, _, 0), tOrVt);
CUTE_STATIC_ASSERT_V(size<2>(tOrVt) == _1{}); // only support MMA_K = 1
Tensor tOrVt_permute_view = make_tensor(tOrVt.data(), make_layout(make_shape(size<0>(tOrVt), size<1, 0>(tOrVt), size<1, 1>(tOrVt))));
permute_4x4_b16(tOrVt_permute_view);
Tensor tOrP = make_tensor(rP.data(), acc_s.layout());
flash::gemm_rr(acc_o, tOrP, tOrVt, tiled_mma_o);
// This check is at the end of the loop since we always have at least 1 iteration
if (n_masking_steps > 1 && n_block <= n_block_min) {
--n_block;
break;
}
}
// These are the iterations where we don't need masking on S
for (; n_block >= n_block_min; --n_block) {
if (n_block > n_block_min) {
// prefetch load page index
row_offset -= kBlockN;
virtual_page_idx = row_offset / params.page_block_size;
page_offset = row_offset - virtual_page_idx * params.page_block_size;
page_idx = block_table[virtual_page_idx];
}
Tensor acc_s = partition_fragment_C(tiled_mma_s, Shape<Int<kBlockM>, Int<kBlockN>>{}); // (MMA=4, MMA_M, MMA_N)
flash::sync_threads();
cute::copy(tKrK, tKsK(_, _, _, 0));
clear(acc_s);
flash::sync_threads();
if (n_block > n_block_min) {
// Advance gK
flash::copy_b64_page_one<Kernel_traits, /*Is_even_MN=*/true, Is_even_K>(gK, tKgK, tKrK, tKVcKV, params.d, n_block - 1,
block_table, params.k_batch_stride, params.k_row_stride, params.page_block_size, page_idx, page_offset);
}
flash::gemm</*A_in_regs=*/Kernel_traits::Is_Q_in_regs>(
acc_s, tSrQ, tSrK, tSsQ, tSsK(_, _, _, 0), tiled_mma_s, smem_tiled_copy_Q, smem_tiled_copy_K,
smem_thr_copy_Q, smem_thr_copy_K
);
mask.template apply_mask<Is_causal, Is_even_MN>(
acc_s, n_block * kBlockN, m_block * kBlockM + (tidx / 64) % kAtomLayoutMS * 16 + (tidx & 0xf), kAtomLayoutMS * 16
);
softmax.template softmax_rescale_o</*Is_first=*/false, /*Check_inf=*/Is_causal || !Is_even_MN, true, true>(acc_s, acc_o, params.scale_softmax_log2);
//Tensor rP = flash::convert_type<Element>(acc_s);
CONVERT_TENSOR_TYPE(ElementAccum, Element, acc_s, rP)
// Reshape rP from (MMA=4, MMA_M, MMA_N) to ((4, 2), MMA_M, MMA_N / 2)
// if using m16n8k16 or (4, MMA_M, MMA_N) if using m16n8k8.
//Tensor tOrP = make_tensor(rP.data(), flash::convert_layout_acc_Aregs<Kernel_traits::TiledMma>(rP.layout()));
// flash::sync_threads();
lds4x4_with_swizzle424(tOsVt(_, _, 0), tOrVt);
CUTE_STATIC_ASSERT_V(size<2>(tOrVt) == _1{}); // only support MMA_K = 1
Tensor tOrVt_permute_view = make_tensor(tOrVt.data(), make_layout(make_shape(size<0>(tOrVt), size<1, 0>(tOrVt), size<1, 1>(tOrVt))));
permute_4x4_b16(tOrVt_permute_view);
Tensor tOrP = make_tensor(rP.data(), acc_s.layout());
flash::gemm_rr(acc_o, tOrP, tOrVt, tiled_mma_o);
}
// Epilogue
if (NoSplit) {
store_16x16<Kernel_traits, false, Is_even_MN, Is_even_K>(params, bidb, bidh, m_block, n_split_idx, smem_, acc_o, softmax);
}else{
store_16x16<Kernel_traits, true, Is_even_MN, Is_even_K>(params, bidb, bidh, m_block, n_split_idx, smem_, acc_o, softmax);
}
}
} // namespace flash

View File

@ -13,40 +13,16 @@
#include "utils.h"
#include "softmax.h"
#include "mask.h"
#include "rotary.h"
#include "attn_mask.h"
namespace flash {
using namespace cute;
template<typename Kernel_traits, bool Is_causal, bool Is_local, bool Has_alibi, bool Is_even_MN, bool Is_even_K, bool Is_softcap, bool Split, typename Params>
__forceinline__ __device__ void compute_attn_1rowblock_splitkv_k64_mla_16x16_4waves(const Params &params, const int bidb, const int bidh, const int m_block, const int n_split_idx, const int num_n_splits) {
template<typename Kernel_traits, bool Split, bool Is_even_MN, bool Is_even_K, typename AccO, typename Softmax, typename Params>
__forceinline__ __device__ void store_16x16(const Params &params, const int bidb, const int bidh, const int m_block, const int n_split_idx, char* smem_, AccO acc_o, Softmax softmax ) {
using ElementAccum = typename Kernel_traits::ElementAccum;
using Element = typename Kernel_traits::Element;
using ElementAccum = typename Kernel_traits::ElementAccum;
using index_t = typename Kernel_traits::index_t;
// Shared memory.
extern __shared__ char smem_[];
// The thread index.
const int tidx = threadIdx.x;
const int warp_idx = tidx / 64;
const int lane_idx = tidx % 64;
constexpr int kBlockM = Kernel_traits::kBlockM;
constexpr int kBlockN = Kernel_traits::kBlockN;
constexpr int kHeadDim = Kernel_traits::kHeadDim;
constexpr int kHeadDimV = Kernel_traits::kHeadDimV;
constexpr int kNWarps = Kernel_traits::kNWarps;
constexpr int kBlockKSmem = Kernel_traits::kBlockKSmem;
constexpr int kAtomLayoutMS = Kernel_traits::kAtomLayoutMS;
constexpr int kAtomLayoutMO = Kernel_traits::kAtomLayoutMO;
constexpr int Num_Stages = Kernel_traits::Num_Stages;
static_assert(kBlockKSmem == 64);
using GmemTiledCopyO = std::conditional_t<
!Split,
typename Kernel_traits::GmemTiledCopyO,
@ -54,300 +30,19 @@ __forceinline__ __device__ void compute_attn_1rowblock_splitkv_k64_mla_16x16_4wa
>;
using ElementO = std::conditional_t<!Split, Element, ElementAccum>;
// The thread index.
const int tidx = threadIdx.x;
const int warp_idx = tidx / 64;
const int lane_idx = tidx % 64;
constexpr int kBlockM = Kernel_traits::kBlockM;
constexpr int kBlockN = Kernel_traits::kBlockN;
constexpr int kNWarps = Kernel_traits::kNWarps;
constexpr int kHeadDimV = Kernel_traits::kHeadDimV;
const BlockInfo</*Varlen=*/!Is_even_MN> binfo(params, bidb);
// if (threadIdx.x == 0 && blockIdx.y == 0 && blockIdx.z == 0) { printf("Is_even_MN = %d, is_cumulativ = %d, seqlen_k_cache = %d, actual_seqlen_k = %d\n", Is_even_MN, params.is_seqlens_k_cumulative, binfo.seqlen_k_cache, binfo.actual_seqlen_k); }
// if (threadIdx.x == 0 && blockIdx.y == 1 && blockIdx.z == 0) { printf("params.knew_ptr = %p, seqlen_k_cache + seqlen_knew = %d\n", params.knew_ptr, binfo.seqlen_k_cache + (params.knew_ptr == nullptr ? 0 : params.seqlen_knew)); }
if (m_block * kBlockM >= binfo.actual_seqlen_q) return;
const int n_blocks_per_split = ((binfo.actual_seqlen_k + kBlockN - 1) / kBlockN + num_n_splits - 1) / num_n_splits;
const int n_block_min = !Is_local
? n_split_idx * n_blocks_per_split
: std::max(n_split_idx * n_blocks_per_split, (m_block * kBlockM + binfo.actual_seqlen_k - binfo.actual_seqlen_q - params.window_size_left) / kBlockN);
int n_block_max = std::min(cute::ceil_div(binfo.actual_seqlen_k, kBlockN), (n_split_idx + 1) * n_blocks_per_split);
if (Is_causal || Is_local) {
n_block_max = std::min(n_block_max,
cute::ceil_div((m_block + 1) * kBlockM + binfo.actual_seqlen_k - binfo.actual_seqlen_q / params.ngroups + params.window_size_right, kBlockN));
}
if (n_block_min >= n_block_max) { // This also covers the case where n_block_max <= 0
// We exit early and write 0 to gOaccum and -inf to gLSEaccum.
// Otherwise we might read OOB elements from gK and gV,
// or get wrong results when we combine gOaccum from different blocks.
const index_t row_offset_o = binfo.q_offset(params.o_batch_stride, params.o_row_stride, bidb)
+ m_block * kBlockM * params.o_row_stride + bidh * params.o_head_stride;
const index_t row_offset_oaccum = (((n_split_idx * params.b + bidb) * params.h + bidh) * params.seqlen_q
+ m_block * kBlockM) * params.d_v;
const index_t row_offset_lseaccum = ((n_split_idx * params.b + bidb) * params.h + bidh) * params.seqlen_q + m_block * kBlockM;
Tensor gOaccum = make_tensor(make_gmem_ptr(reinterpret_cast<ElementO *>(Split ? params.oaccum_ptr : params.o_ptr) + (Split ? row_offset_oaccum : row_offset_o)),
Shape<Int<kBlockM>, Int<kHeadDimV>>{},
make_stride(Split ? kHeadDimV : params.o_row_stride, _1{}));
Tensor gLSEaccum = make_tensor(make_gmem_ptr(reinterpret_cast<ElementAccum *>(Split ? params.softmax_lseaccum_ptr : params.softmax_lse_ptr) + row_offset_lseaccum),
Shape<Int<kBlockM>>{}, Stride<_1>{});
GmemTiledCopyO gmem_tiled_copy_Oaccum;
auto gmem_thr_copy_Oaccum = gmem_tiled_copy_Oaccum.get_thread_slice(tidx);
Tensor tOgOaccum = gmem_thr_copy_Oaccum.partition_D(gOaccum);
Tensor tOrOaccum = make_tensor<ElementO>(shape(tOgOaccum));
clear(tOrOaccum);
// Construct identity layout for sO
Tensor cO = make_identity_tensor(make_shape(size<0>(gOaccum), size<1>(gOaccum))); // (BLK_M,BLK_K) -> (blk_m,blk_k)
// Repeat the partitioning with identity layouts
Tensor tOcO = gmem_thr_copy_Oaccum.partition_D(cO);
Tensor tOpO = make_tensor<bool>(make_shape(size<2>(tOgOaccum)));
if (!Is_even_K) {
#pragma unroll
for (int k = 0; k < size(tOpO); ++k) { tOpO(k) = get<1>(tOcO(0, 0, k)) < params.d_v; }
}
// Clear_OOB_K must be false since we don't want to write zeros to gmem
flash::copy<Is_even_MN, Is_even_K, /*Clear_OOB_MN=*/false, /*Clear_OOB_K=*/false>(
gmem_tiled_copy_Oaccum, tOrOaccum, tOgOaccum, tOcO, tOpO, binfo.actual_seqlen_q - m_block * kBlockM
);
#pragma unroll
for (int m = 0; m < size<1>(tOgOaccum); ++m) {
const int row = get<0>(tOcO(0, m, 0));
if (row < binfo.actual_seqlen_q - m_block * kBlockM && get<1>(tOcO(0, m, 0)) == 0) { gLSEaccum(row) = Split ? -INFINITY : INFINITY; }
}
return;
}
// We iterate over the blocks in reverse order. This is because the last block is the only one
// that needs masking when we read K and V from global memory. Moreover, iterating in reverse
// might save us 1 register (we just need n_block instead of both n_block and n_block_max).
const index_t row_offset_q = binfo.q_offset(params.q_batch_stride, params.q_row_stride, bidb)
+ m_block * kBlockM * params.q_row_stride + bidh * params.q_head_stride;
// We move K and V to the last block.
const int bidb_cache = params.cache_batch_idx == nullptr ? bidb : params.cache_batch_idx[bidb];
const int *block_table = params.block_table == nullptr ? nullptr : params.block_table + bidb * params.block_table_batch_stride;
const int block_table_idx = block_table == nullptr ? 0 : (n_block_max - 1) * kBlockN / params.page_block_size;
const int block_table_offset = block_table == nullptr ? 0 : (n_block_max - 1) * kBlockN - block_table_idx * params.page_block_size;
const index_t row_offset_k = block_table == nullptr
? binfo.k_offset(params.k_batch_stride, params.k_row_stride, bidb_cache)
+ (n_block_max - 1) * kBlockN * params.k_row_stride + (bidh / params.h_h_k_ratio) * params.k_head_stride
: (bidh / params.h_h_k_ratio) * params.k_head_stride;
const index_t row_offset_v = block_table == nullptr
? binfo.k_offset(params.v_batch_stride, params.v_row_stride, bidb_cache)
+ (n_block_max - 1) * kBlockN * params.v_row_stride + (bidh / params.h_h_k_ratio) * params.v_head_stride
: (bidh / params.h_h_k_ratio) * params.v_head_stride;
Tensor gQ = make_tensor(make_gmem_ptr(reinterpret_cast<Element *>(params.q_ptr) + row_offset_q),
Shape<Int<kBlockM>, Int<kHeadDim>>{},
make_stride(params.q_row_stride, _1{}));
Tensor gK = make_tensor(make_gmem_ptr(reinterpret_cast<Element *>(params.k_ptr) + row_offset_k),
Shape<Int<kBlockN>, Int<kHeadDim>>{},
make_stride(params.k_row_stride, _1{}));
// if (threadIdx.x == 0 && blockIdx.y == 0 && blockIdx.z == 0) { printf("k_ptr = %p, row_offset_k = %d, gK_ptr = %p\n", params.k_ptr, row_offset_k, gK.data()); }
Tensor gV = make_tensor(make_gmem_ptr(reinterpret_cast<Element *>(params.v_ptr) + row_offset_v),
Shape<Int<kBlockN>, Int<kHeadDimV>>{},
make_stride(params.v_row_stride, _1{}));
Tensor sQ = make_tensor(make_smem_ptr(reinterpret_cast<Element *>(smem_)),
typename Kernel_traits::SmemLayoutQ424{});
Tensor sK = make_tensor(sQ.data() + (Kernel_traits::Share_Q_K_smem ? 0 : size(sQ)),
typename Kernel_traits::SmemLayoutK{});
Tensor sV = make_tensor(sK.data(), typename Kernel_traits::SmemLayoutVtNoSwizzle{});
Tensor sVt = make_tensor(sV.data(), typename Kernel_traits::SmemLayoutVtransposed{});
Tensor sVtNoSwizzle = make_tensor(sV.data(), typename Kernel_traits::SmemLayoutVtransposedNoSwizzle{});
typename Kernel_traits::GmemTiledCopyB64 gmem_tiled_copy_Q;
auto gmem_thr_copy_Q = gmem_tiled_copy_Q.get_thread_slice(tidx);
Tensor tQgQ = gmem_thr_copy_Q.partition_S(gQ);
Tensor tQsQ = gmem_thr_copy_Q.partition_D(sQ);
typename Kernel_traits::GmemTiledCopyB64 gmem_tiled_copy_KV;
auto gmem_thr_copy_KV = gmem_tiled_copy_KV.get_thread_slice(tidx);
Tensor tKgK = gmem_thr_copy_KV.partition_S(gK); // (KCPY, KCPY_N, KCPY_K)
Tensor tKsK = gmem_thr_copy_KV.partition_D(sK);
// S is only 16x16 size, so all 4 waves compute the same S
int tidx_mma_s = tidx & 0x3F;
typename Kernel_traits::TiledMmaS tiled_mma_s;
auto thr_mma_s = tiled_mma_s.get_thread_slice(tidx_mma_s);
Tensor tSrQ = thr_mma_s.partition_fragment_A(sQ); // (MMA,MMA_M,MMA_K)
Tensor tSrK = thr_mma_s.partition_fragment_B(sK(_, _, 0)); // (MMA,MMA_N,MMA_K)
typename Kernel_traits::TiledMmaO tiled_mma_o;
auto thr_mma_o = tiled_mma_o.get_thread_slice(tidx);
// Tensor tOrVt = thr_mma_o.partition_fragment_B(sVt); // (MMA, MMA_K,MMA_N)
Tensor tOrVt = make_tensor<Element>(Shape<_4, Shape<_4, _2>, _1>{});
Tensor acc_o = partition_fragment_C(tiled_mma_o, Shape<Int<kBlockM>, Int<kHeadDimV>>{}); // MMA, MMA_M, MMA_K
//
// Copy Atom retiling
//
auto smem_tiled_copy_Q = make_tiled_copy_A(typename Kernel_traits::SmemCopyAtomB64{}, tiled_mma_s);
auto smem_thr_copy_Q = smem_tiled_copy_Q.get_thread_slice(tidx_mma_s);
Tensor tSsQ = smem_thr_copy_Q.partition_S(sQ);
auto smem_tiled_copy_K = make_tiled_copy_B(typename Kernel_traits::SmemCopyAtomB64{}, tiled_mma_s);
auto smem_thr_copy_K = smem_tiled_copy_K.get_thread_slice(tidx_mma_s);
Tensor tSsK = smem_thr_copy_K.partition_S(sK);
auto smem_tiled_copy_V = make_tiled_copy_B(typename Kernel_traits::SmemCopyAtomTransposed{}, tiled_mma_o);
auto smem_thr_copy_V = smem_tiled_copy_V.get_thread_slice(tidx);
int warp_offset = warp_idx / kAtomLayoutMO * 16 * 64;
int thread_offset = lane_idx / 16 * 4 * 64;
Element *Vtsmem_ptr_lds = reinterpret_cast<Element *>(sVt.data().get()) + warp_offset + thread_offset;
Tensor tOsVt = make_tensor(make_smem_ptr(Vtsmem_ptr_lds), make_layout(Shape<_4, _2, Int<Num_Stages>>{}, // MMA MMA_N NUM_STAGES
Stride<_1, Int<16*256>, Int<kBlockN*kHeadDim>>{}));
// PREDICATES
// Construct identity layout for sQ and sK
Tensor cQ = make_identity_tensor(make_shape(size<0>(sQ), size<1>(sQ))); // (BLK_M,BLK_K) -> (blk_m,blk_k)
Tensor cKV = make_identity_tensor(make_shape(size<0>(sK), size<1>(sK))); // (BLK_N,BLK_K) -> (blk_n,blk_k)
// Repeat the partitioning with identity layouts
Tensor tQcQ = gmem_thr_copy_Q.partition_S(cQ); // (ACPY,ACPY_M,ACPY_K) -> (blk_m,blk_k)
Tensor tKVcKV = gmem_thr_copy_KV.partition_S(cKV); // (BCPY,BCPY_N,BCPY_K) -> (blk_n,blk_k)
// Prologue
// Read Q from gmem to smem, optionally apply rotary embedding.
Tensor tQrQ = make_fragment_like(tQgQ);
// We don't need to clear the sQ smem tiles since we'll only write out the valid outputs
flash::copy_b64<Is_even_MN, Is_even_K>(tQgQ, tQrQ, tQcQ, params.d, binfo.actual_seqlen_q - m_block * kBlockM);
cute::copy(tQrQ, tQsQ);
if constexpr (Kernel_traits::Is_Q_in_regs) {
flash::sync_threads();
cute::copy(smem_tiled_copy_Q, tSsQ, tSrQ);
flash::sync_threads();
}
int n_block = n_block_max - 1;
int Ksmem_read_index = 0;
int Ksmem_write_index = 0;
// We don't need to clear the sK smem tiles since we'll mask out the scores anyway.
Tensor tKrK = make_fragment_like(tKgK);
flash::copy_b64_page_one<Kernel_traits, Is_even_MN, Is_even_K>(gK, tKgK, tKrK, tKVcKV, params.d, n_block,
block_table, params.k_batch_stride, params.k_row_stride, params.page_block_size, binfo.actual_seqlen_k - n_block * kBlockN);
// flash::cp_async_wait<0>();
// __syncthreads();
// if (tidx == 0 && blockIdx.y == 0 && blockIdx.z == 0) { print(tKsK); }
// __syncthreads();
clear(acc_o);
flash::Softmax<size<1>(acc_o)> softmax;
const float alibi_slope = !Has_alibi ? 0.0f : reinterpret_cast<float *>(params.alibi_slopes_ptr)[bidb * params.alibi_slopes_batch_stride + bidh] / params.scale_softmax;
flash::Mask<Is_causal, Is_local, Has_alibi> mask(binfo.actual_seqlen_k, binfo.actual_seqlen_q, params.ngroups, params.window_size_left, params.window_size_right, alibi_slope);
// For performance reason, we separate out two kinds of iterations:
// those that need masking on S, and those that don't.
// We need masking on S for the very last block when K and V has length not multiple of kBlockN.
// We also need masking on S if it's causal, for the last ceil_div(kBlockM, kBlockN) blocks.
// We will have at least 1 "masking" iteration.
// If not even_N, then seqlen_k might end in the middle of a block. In that case we need to
// mask 2 blocks (e.g. when kBlockM == kBlockN), not just 1.
constexpr int n_masking_steps = (!Is_causal && !Is_local)
? 1
: ((Is_even_MN && Is_causal) ? cute::ceil_div(kBlockM, kBlockN) : cute::ceil_div(kBlockM, kBlockN) + 1);
#pragma unroll
for (int masking_step = 0; masking_step < n_masking_steps; ++masking_step, --n_block) {
Tensor acc_s = partition_fragment_C(tiled_mma_s, Shape<Int<kBlockM>, Int<kBlockN>>{}); // (MMA=4, MMA_M, MMA_N)
cute::copy(tKrK, tKsK(_, _, _, Ksmem_write_index));
Ksmem_write_index ^= 1;
clear(acc_s);
flash::sync_threads();
if (n_block > n_block_min) {
flash::copy_b64_page_one<Kernel_traits, /*Is_even_MN=*/true, Is_even_K>(gK, tKgK, tKrK, tKVcKV, params.d, n_block - 1,
block_table, params.k_batch_stride, params.k_row_stride, params.page_block_size);
}
flash::gemm_opt</*A_in_regs=*/Kernel_traits::Is_Q_in_regs>(
acc_s, tSrQ, tSrK, tSsQ, tSsK(_, _, _, Ksmem_read_index), tiled_mma_s, smem_tiled_copy_Q, smem_tiled_copy_K,
smem_thr_copy_Q, smem_thr_copy_K
);
// if (cute::thread0()) { print(acc_s); }
if constexpr (Is_softcap){
flash::apply_softcap(acc_s, params.softcap);
}
mask.template apply_mask<Is_causal, Is_even_MN>(
acc_s, n_block * kBlockN, m_block * kBlockM + (tidx / 64) % kAtomLayoutMS * 16 + (tidx & 0xf), kAtomLayoutMS * 16
);
// We have key_padding_mask so we'll need to Check_inf
masking_step == 0
? softmax.template softmax_rescale_o</*Is_first=*/true, /*Check_inf=*/Is_causal || Is_local || !Is_even_MN, true, true>(acc_s, acc_o, params.scale_softmax_log2)
: softmax.template softmax_rescale_o</*Is_first=*/false, /*Check_inf=*/Is_causal || Is_local || !Is_even_MN, true, true>(acc_s, acc_o, params.scale_softmax_log2);
// if (cute::thread0()) { print(scores_max); print(scores_sum); print(scores); }
// Convert acc_s from fp32 to fp16/bf16
//Tensor rP = flash::convert_type<Element>(acc_s);
CONVERT_TENSOR_TYPE(ElementAccum, Element, acc_s, rP)
// Reshape rP from (MMA=4, MMA_M, MMA_N) to ((4, 2), MMA_M, MMA_N / 2)
// if using m16n8k16 or (4, MMA_M, MMA_N) if using m16n8k8.
//Tensor tOrP = make_tensor(rP.data(), flash::convert_layout_acc_Aregs<Kernel_traits::TiledMma>(rP.layout()));
lds4x4_with_swizzle424(tOsVt(_, _, Ksmem_read_index), tOrVt);
CUTE_STATIC_ASSERT_V(size<2>(tOrVt) == _1{}); // only support MMA_K = 1
Tensor tOrVt_permute_view = make_tensor(tOrVt.data(), make_layout(make_shape(size<0>(tOrVt), size<1, 0>(tOrVt), size<1, 1>(tOrVt))));
permute_4x4_b16(tOrVt_permute_view);
Tensor tOrP = make_tensor(rP.data(), acc_s.layout());
flash::gemm_rr(acc_o, tOrP, tOrVt, tiled_mma_o);
Ksmem_read_index ^= 1;
// This check is at the end of the loop since we always have at least 1 iteration
if (n_masking_steps > 1 && n_block <= n_block_min) {
--n_block;
break;
}
}
// These are the iterations where we don't need masking on S
for (; n_block >= n_block_min; --n_block) {
Tensor acc_s = partition_fragment_C(tiled_mma_s, Shape<Int<kBlockM>, Int<kBlockN>>{}); // (MMA=4, MMA_M, MMA_N)
cute::copy(tKrK, tKsK(_, _, _, Ksmem_write_index));
Ksmem_write_index ^= 1;
clear(acc_s);
flash::sync_threads();
if (n_block > n_block_min) {
// Advance gK
flash::copy_b64_page_one<Kernel_traits, /*Is_even_MN=*/true, Is_even_K>(gK, tKgK, tKrK, tKVcKV, params.d, n_block - 1,
block_table, params.k_batch_stride, params.k_row_stride, params.page_block_size);
}
flash::gemm_opt</*A_in_regs=*/Kernel_traits::Is_Q_in_regs>(
acc_s, tSrQ, tSrK, tSsQ, tSsK(_, _, _, Ksmem_read_index), tiled_mma_s, smem_tiled_copy_Q, smem_tiled_copy_K,
smem_thr_copy_Q, smem_thr_copy_K
);
if constexpr (Is_softcap){
flash::apply_softcap(acc_s, params.softcap);
}
mask.template apply_mask<Is_causal, Is_even_MN>(
acc_s, n_block * kBlockN, m_block * kBlockM + (tidx / 64) % kAtomLayoutMS * 16 + (tidx & 0xf), kAtomLayoutMS * 16
);
softmax.template softmax_rescale_o</*Is_first=*/false, /*Check_inf=*/Is_causal || Is_local || !Is_even_MN, true, true>(acc_s, acc_o, params.scale_softmax_log2);
//Tensor rP = flash::convert_type<Element>(acc_s);
CONVERT_TENSOR_TYPE(ElementAccum, Element, acc_s, rP)
// Reshape rP from (MMA=4, MMA_M, MMA_N) to ((4, 2), MMA_M, MMA_N / 2)
// if using m16n8k16 or (4, MMA_M, MMA_N) if using m16n8k8.
//Tensor tOrP = make_tensor(rP.data(), flash::convert_layout_acc_Aregs<Kernel_traits::TiledMma>(rP.layout()));
lds4x4_with_swizzle424(tOsVt(_, _, Ksmem_read_index), tOrVt);
CUTE_STATIC_ASSERT_V(size<2>(tOrVt) == _1{}); // only support MMA_K = 1
Tensor tOrVt_permute_view = make_tensor(tOrVt.data(), make_layout(make_shape(size<0>(tOrVt), size<1, 0>(tOrVt), size<1, 1>(tOrVt))));
permute_4x4_b16(tOrVt_permute_view);
Tensor tOrP = make_tensor(rP.data(), acc_s.layout());
flash::gemm_rr(acc_o, tOrP, tOrVt, tiled_mma_o);
Ksmem_read_index ^= 1;
}
// Epilogue
Tensor lse = softmax.template normalize_softmax_lse</*Is_dropout=*/false, /*Return_lse*/true, Split>(acc_o, params.scale_softmax);
Tensor acc_o_view = make_tensor(acc_o.data(), make_layout(Shape<_4, Shape<_4, _2>>{},
Stride<_1, Shape<_4, _16>>{}));
@ -361,7 +56,6 @@ __forceinline__ __device__ void compute_attn_1rowblock_splitkv_k64_mla_16x16_4wa
acc_o_copy(row, make_coord(col, k)) = acc_o_view(col, make_coord(row, k));
}
}
// if (cute::thread0()) { print(lse); }
if constexpr (!Split) {
Tensor sOaccum = make_tensor(make_smem_ptr(reinterpret_cast<ElementO *>(smem_)), typename Kernel_traits::SmemLayoutO{}); // (SMEM_M,SMEM_N)
// Partition sO to match the accumulator partitioning
@ -376,8 +70,6 @@ __forceinline__ __device__ void compute_attn_1rowblock_splitkv_k64_mla_16x16_4wa
Stride<_1, _16>{}));
// sOaccum is larger than sQ, so we need to syncthreads here
// TODO: allocate enough smem for sOaccum
if constexpr (Kernel_traits::Share_Q_K_smem) { flash::sync_threads(); }
cute::copy(tOrO, tOsO);
@ -391,7 +83,6 @@ __forceinline__ __device__ void compute_attn_1rowblock_splitkv_k64_mla_16x16_4wa
make_stride(params.o_row_stride, _1{}));
Tensor gLSEaccum = make_tensor(make_gmem_ptr(reinterpret_cast<ElementAccum *>(params.softmax_lse_ptr) + row_offset_lseaccum),
Shape<Int<kBlockM>>{}, Stride<_1>{});
// if (tidx == 0) { printf("row_offset_o = %d, bidh = %d, gOaccum = %p\n", row_offset_o, bidh, gOaccum.data()); }
GmemTiledCopyO gmem_tiled_copy_Oaccum;
auto gmem_thr_copy_Oaccum = gmem_tiled_copy_Oaccum.get_thread_slice(tidx);
@ -426,17 +117,17 @@ __forceinline__ __device__ void compute_attn_1rowblock_splitkv_k64_mla_16x16_4wa
tOrOaccum, tOgOaccum, tOcO, params.d_v, binfo.actual_seqlen_q - m_block * kBlockM
);
} else {
const int split_offset = params.num_splits_ptr[bidb];
Tensor sOaccum = make_tensor(make_smem_ptr(reinterpret_cast<ElementO *>(smem_)), typename Kernel_traits::SmemLayoutO{}); // (SMEM_M,SMEM_N)
const index_t row_offset_oaccum = (((n_split_idx * params.b + bidb) * params.h + bidh) * params.seqlen_q
const index_t row_offset_oaccum = (((split_offset + n_split_idx) * params.h + bidh) * params.seqlen_q
+ m_block * kBlockM) * params.d_v;
const index_t row_offset_lseaccum = ((n_split_idx * params.b + bidb) * params.h + bidh) * params.seqlen_q + m_block * kBlockM;
const index_t row_offset_lseaccum = ((split_offset + n_split_idx) * params.h + bidh) * params.seqlen_q + m_block * kBlockM;
Tensor gOaccum = make_tensor(make_gmem_ptr(reinterpret_cast<ElementO *>(params.oaccum_ptr) + row_offset_oaccum),
Shape<Int<kBlockM>, Int<kHeadDimV>>{},
make_stride(kHeadDimV, _1{}));
Tensor gLSEaccum = make_tensor(make_gmem_ptr(reinterpret_cast<ElementAccum *>(params.softmax_lseaccum_ptr) + row_offset_lseaccum),
Shape<Int<kBlockM>>{}, Stride<_1>{});
// if (tidx == 0) { printf("row_offset_o = %d, bidh = %d, gOaccum = %p\n", row_offset_o, bidh, gOaccum.data()); }
Tensor taccOrOaccum = make_tensor(acc_o_copy.data(), acc_o_copy.layout());
int warp_offset = warp_idx * 16 * 64;
@ -491,5 +182,262 @@ __forceinline__ __device__ void compute_attn_1rowblock_splitkv_k64_mla_16x16_4wa
);
}
}
template<typename Kernel_traits, bool Is_causal, bool Is_even_MN, bool Is_even_K, bool Is_enable_dcp, typename Params>
__forceinline__ __device__ void compute_attn_1rowblock_splitkv_mla_k64_16x16_4waves_xcore1000(
const Params &params, const int bidb, const int bidh, const int m_block, const int n_split_idx, const int n_block_min, const int n_block_max, const bool NoSplit) {
using Element = typename Kernel_traits::Element;
using ElementAccum = typename Kernel_traits::ElementAccum;
using index_t = typename Kernel_traits::index_t;
// Shared memory.
extern __shared__ char smem_[];
// The thread index.
const int tidx = threadIdx.x;
const int warp_idx = tidx / 64;
const int lane_idx = tidx % 64;
constexpr int kBlockM = Kernel_traits::kBlockM;
constexpr int kBlockN = Kernel_traits::kBlockN;
constexpr int kHeadDim = Kernel_traits::kHeadDim;
constexpr int kHeadDimV = Kernel_traits::kHeadDimV;
constexpr int kNWarps = Kernel_traits::kNWarps;
constexpr int kBlockKSmem = Kernel_traits::kBlockKSmem;
constexpr int kAtomLayoutMS = Kernel_traits::kAtomLayoutMS;
constexpr int kAtomLayoutMO = Kernel_traits::kAtomLayoutMO;
constexpr int Num_Stages = Kernel_traits::Num_Stages;
static_assert(kBlockKSmem == 64);
const BlockInfo</*Varlen=*/!Is_even_MN> binfo(params, bidb);
if (m_block * kBlockM >= binfo.actual_seqlen_q) return;
const int actual_seqlen_k = params.cu_seqlens_k[bidb];
// We iterate over the blocks in reverse order. This is because the last block is the only one
// that needs masking when we read K and V from global memory. Moreover, iterating in reverse
// might save us 1 register (we just need n_block instead of both n_block and n_block_max).
const index_t row_offset_q = binfo.q_offset(params.q_batch_stride, params.q_row_stride, bidb)
+ m_block * kBlockM * params.q_row_stride + bidh * params.q_head_stride;
// We move K and V to the last block.
const int bidb_cache = params.cache_batch_idx == nullptr ? bidb : params.cache_batch_idx[bidb];
const int *block_table = params.block_table == nullptr ? nullptr : params.block_table + bidb * params.block_table_batch_stride;
const int block_table_idx = block_table == nullptr ? 0 : (n_block_max - 1) * kBlockN / params.page_block_size;
const int block_table_offset = block_table == nullptr ? 0 : (n_block_max - 1) * kBlockN - block_table_idx * params.page_block_size;
const index_t row_offset_k = block_table == nullptr
? binfo.k_offset(params.k_batch_stride, params.k_row_stride, bidb_cache)
+ (n_block_max - 1) * kBlockN * params.k_row_stride + (bidh / params.h_h_k_ratio) * params.k_head_stride
: (bidh / params.h_h_k_ratio) * params.k_head_stride;
const index_t row_offset_v = block_table == nullptr
? binfo.k_offset(params.v_batch_stride, params.v_row_stride, bidb_cache)
+ (n_block_max - 1) * kBlockN * params.v_row_stride + (bidh / params.h_h_k_ratio) * params.v_head_stride
: (bidh / params.h_h_k_ratio) * params.v_head_stride;
Tensor gQ = make_tensor(make_gmem_ptr(reinterpret_cast<Element *>(params.q_ptr) + row_offset_q),
Shape<Int<kBlockM>, Int<kHeadDim>>{},
make_stride(params.q_row_stride, _1{}));
Tensor gK = make_tensor(make_gmem_ptr(reinterpret_cast<Element *>(params.k_ptr) + row_offset_k),
Shape<Int<kBlockN>, Int<kHeadDim>>{},
make_stride(params.k_row_stride, _1{}));
Tensor gV = make_tensor(make_gmem_ptr(reinterpret_cast<Element *>(params.v_ptr) + row_offset_v),
Shape<Int<kBlockN>, Int<kHeadDimV>>{},
make_stride(params.v_row_stride, _1{}));
Tensor sQ = make_tensor(make_smem_ptr(reinterpret_cast<Element *>(smem_)),
typename Kernel_traits::SmemLayoutQ424{});
Tensor sK = make_tensor(sQ.data() + (Kernel_traits::Share_Q_K_smem ? 0 : size(sQ)),
typename Kernel_traits::SmemLayoutK424{});
Tensor sV = make_tensor(sK.data(), typename Kernel_traits::SmemLayoutVtNoSwizzle{});
Tensor sVt = make_tensor(sV.data(), typename Kernel_traits::SmemLayoutVtransposed424{});
Tensor sVtNoSwizzle = make_tensor(sV.data(), typename Kernel_traits::SmemLayoutVtransposedNoSwizzle{});
typename Kernel_traits::GmemTiledCopyB64 gmem_tiled_copy_Q;
auto gmem_thr_copy_Q = gmem_tiled_copy_Q.get_thread_slice(tidx);
Tensor tQgQ = gmem_thr_copy_Q.partition_S(gQ);
Tensor tQsQ = gmem_thr_copy_Q.partition_D(sQ);
typename Kernel_traits::GmemTiledCopyB64 gmem_tiled_copy_KV;
auto gmem_thr_copy_KV = gmem_tiled_copy_KV.get_thread_slice(tidx);
Tensor tKgK = gmem_thr_copy_KV.partition_S(gK); // (KCPY, KCPY_N, KCPY_K)
Tensor tKsK = gmem_thr_copy_KV.partition_D(sK);
// S is only 16x16 size, so all 4 waves compute the same S
int tidx_mma_s = tidx & 0x3F;
typename Kernel_traits::TiledMmaS tiled_mma_s;
auto thr_mma_s = tiled_mma_s.get_thread_slice(tidx_mma_s);
Tensor tSrQ = thr_mma_s.partition_fragment_A(sQ); // (MMA,MMA_M,MMA_K)
Tensor tSrK = thr_mma_s.partition_fragment_B(sK(_, _, 0)); // (MMA,MMA_N,MMA_K)
typename Kernel_traits::TiledMmaO tiled_mma_o;
auto thr_mma_o = tiled_mma_o.get_thread_slice(tidx);
// Tensor tOrVt = thr_mma_o.partition_fragment_B(sVt); // (MMA, MMA_K,MMA_N)
Tensor tOrVt = make_tensor<Element>(Shape<_4, Shape<_4, _2>, _1>{});
Tensor acc_o = partition_fragment_C(tiled_mma_o, Shape<Int<kBlockM>, Int<kHeadDimV>>{}); // MMA, MMA_M, MMA_K
//
// Copy Atom retiling
//
auto smem_tiled_copy_Q = make_tiled_copy_A(typename Kernel_traits::UniversalCopyAtomB64{}, tiled_mma_s);
auto smem_thr_copy_Q = smem_tiled_copy_Q.get_thread_slice(tidx_mma_s);
Tensor tSsQ = smem_thr_copy_Q.partition_S(sQ);
auto smem_tiled_copy_K = make_tiled_copy_B(typename Kernel_traits::UniversalCopyAtomB64{}, tiled_mma_s);
auto smem_thr_copy_K = smem_tiled_copy_K.get_thread_slice(tidx_mma_s);
Tensor tSsK = smem_thr_copy_K.partition_S(sK);
auto smem_tiled_copy_V = make_tiled_copy_B(typename Kernel_traits::SmemCopyAtomTransposed{}, tiled_mma_o);
auto smem_thr_copy_V = smem_tiled_copy_V.get_thread_slice(tidx);
int warp_offset = warp_idx / kAtomLayoutMO * 16 * 64;
int thread_offset = lane_idx / 16 * 4 * 64;
Element *Vtsmem_ptr_lds = reinterpret_cast<Element *>(sVt.data().get()) + warp_offset + thread_offset;
Tensor tOsVt = make_tensor(make_smem_ptr(Vtsmem_ptr_lds), make_layout(Shape<_4, _2, Int<Num_Stages>>{}, // MMA MMA_N NUM_STAGES
Stride<_1, Int<16*256>, Int<kBlockN*kHeadDim>>{}));
// PREDICATES
// Construct identity layout for sQ and sK
Tensor cQ = make_identity_tensor(make_shape(size<0>(sQ), size<1>(sQ))); // (BLK_M,BLK_K) -> (blk_m,blk_k)
Tensor cKV = make_identity_tensor(make_shape(size<0>(sK), size<1>(sK))); // (BLK_N,BLK_K) -> (blk_n,blk_k)
// Repeat the partitioning with identity layouts
Tensor tQcQ = gmem_thr_copy_Q.partition_S(cQ); // (ACPY,ACPY_M,ACPY_K) -> (blk_m,blk_k)
Tensor tKVcKV = gmem_thr_copy_KV.partition_S(cKV); // (BCPY,BCPY_N,BCPY_K) -> (blk_n,blk_k)
// Prologue
// Read Q from gmem to smem, optionally apply rotary embedding.
Tensor tQrQ = make_fragment_like(tQgQ);
// We don't need to clear the sQ smem tiles since we'll only write out the valid outputs
flash::copy_b64<Is_even_MN, Is_even_K>(tQgQ, tQrQ, tQcQ, params.d, binfo.actual_seqlen_q - m_block * kBlockM);
cute::copy(tQrQ, tQsQ);
if constexpr (Kernel_traits::Is_Q_in_regs) {
flash::sync_threads();
cute::copy(smem_tiled_copy_Q, tSsQ, tSrQ);
flash::sync_threads();
}
int n_block = n_block_max - 1;
int Ksmem_read_index = 0;
int Ksmem_write_index = 0;
// We don't need to clear the sK smem tiles since we'll mask out the scores anyway.
Tensor tKrK = make_fragment_like(tKgK);
flash::copy_b64_page_one<Kernel_traits, Is_even_MN, Is_even_K>(gK, tKgK, tKrK, tKVcKV, params.d, n_block,
block_table, params.k_batch_stride, params.k_row_stride, params.page_block_size, binfo.actual_seqlen_k - n_block * kBlockN);
clear(acc_o);
flash::Softmax<size<1>(acc_o)> softmax;
flash::Mask<Is_causal, Is_enable_dcp> mask(binfo.actual_seqlen_k, binfo.actual_seqlen_q, params.ngroups, binfo.tot_seqlen_k, params.cp_world_size, params.cp_rank);
// For performance reason, we separate out two kinds of iterations:
// those that need masking on S, and those that don't.
// We need masking on S for the very last block when K and V has length not multiple of kBlockN.
// We also need masking on S if it's causal, for the last ceil_div(kBlockM, kBlockN) blocks.
// We will have at least 1 "masking" iteration.
// If not even_N, then seqlen_k might end in the middle of a block. In that case we need to
// mask 2 blocks (e.g. when kBlockM == kBlockN), not just 1.
constexpr int n_masking_steps = (!Is_causal)
? 1
: ((Is_even_MN && Is_causal) ? cute::ceil_div(kBlockM, kBlockN) : cute::ceil_div(kBlockM, kBlockN) + 1);
#pragma unroll
for (int masking_step = 0; masking_step < n_masking_steps; ++masking_step, --n_block) {
Tensor acc_s = partition_fragment_C(tiled_mma_s, Shape<Int<kBlockM>, Int<kBlockN>>{}); // (MMA=4, MMA_M, MMA_N)
cute::copy(tKrK, tKsK(_, _, _, Ksmem_write_index));
Ksmem_write_index ^= 1;
clear(acc_s);
flash::sync_threads();
if (n_block > n_block_min) {
flash::copy_b64_page_one<Kernel_traits, /*Is_even_MN=*/true, Is_even_K>(gK, tKgK, tKrK, tKVcKV, params.d, n_block - 1,
block_table, params.k_batch_stride, params.k_row_stride, params.page_block_size);
}
flash::gemm</*A_in_regs=*/Kernel_traits::Is_Q_in_regs>(
acc_s, tSrQ, tSrK, tSsQ, tSsK(_, _, _, Ksmem_read_index), tiled_mma_s, smem_tiled_copy_Q, smem_tiled_copy_K,
smem_thr_copy_Q, smem_thr_copy_K
);
mask.template apply_mask<Is_causal, Is_even_MN>(
acc_s, n_block * kBlockN, m_block * kBlockM + (tidx / 64) % kAtomLayoutMS * 16 + (tidx & 0xf), kAtomLayoutMS * 16
);
// We have key_padding_mask so we'll need to Check_inf
masking_step == 0
? softmax.template softmax_rescale_o</*Is_first=*/true, /*Check_inf=*/Is_causal || !Is_even_MN, true, true>(acc_s, acc_o, params.scale_softmax_log2)
: softmax.template softmax_rescale_o</*Is_first=*/false, /*Check_inf=*/Is_causal || !Is_even_MN, true, true>(acc_s, acc_o, params.scale_softmax_log2);
// Convert acc_s from fp32 to fp16/bf16
//Tensor rP = flash::convert_type<Element>(acc_s);
CONVERT_TENSOR_TYPE(ElementAccum, Element, acc_s, rP)
// Reshape rP from (MMA=4, MMA_M, MMA_N) to ((4, 2), MMA_M, MMA_N / 2)
// if using m16n8k16 or (4, MMA_M, MMA_N) if using m16n8k8.
//Tensor tOrP = make_tensor(rP.data(), flash::convert_layout_acc_Aregs<Kernel_traits::TiledMma>(rP.layout()));
lds4x4_with_swizzle424(tOsVt(_, _, Ksmem_read_index), tOrVt);
CUTE_STATIC_ASSERT_V(size<2>(tOrVt) == _1{}); // only support MMA_K = 1
Tensor tOrVt_permute_view = make_tensor(tOrVt.data(), make_layout(make_shape(size<0>(tOrVt), size<1, 0>(tOrVt), size<1, 1>(tOrVt))));
permute_4x4_b16(tOrVt_permute_view);
Tensor tOrP = make_tensor(rP.data(), acc_s.layout());
flash::gemm_rr(acc_o, tOrP, tOrVt, tiled_mma_o);
Ksmem_read_index ^= 1;
// This check is at the end of the loop since we always have at least 1 iteration
if (n_masking_steps > 1 && n_block <= n_block_min) {
--n_block;
break;
}
}
// These are the iterations where we don't need masking on S
for (; n_block >= n_block_min; --n_block) {
Tensor acc_s = partition_fragment_C(tiled_mma_s, Shape<Int<kBlockM>, Int<kBlockN>>{}); // (MMA=4, MMA_M, MMA_N)
cute::copy(tKrK, tKsK(_, _, _, Ksmem_write_index));
Ksmem_write_index ^= 1;
clear(acc_s);
flash::sync_threads();
if (n_block > n_block_min) {
// Advance gK
flash::copy_b64_page_one<Kernel_traits, /*Is_even_MN=*/true, Is_even_K>(gK, tKgK, tKrK, tKVcKV, params.d, n_block - 1,
block_table, params.k_batch_stride, params.k_row_stride, params.page_block_size);
}
flash::gemm</*A_in_regs=*/Kernel_traits::Is_Q_in_regs>(
acc_s, tSrQ, tSrK, tSsQ, tSsK(_, _, _, Ksmem_read_index), tiled_mma_s, smem_tiled_copy_Q, smem_tiled_copy_K,
smem_thr_copy_Q, smem_thr_copy_K
);
mask.template apply_mask<Is_causal, Is_even_MN>(
acc_s, n_block * kBlockN, m_block * kBlockM + (tidx / 64) % kAtomLayoutMS * 16 + (tidx & 0xf), kAtomLayoutMS * 16
);
softmax.template softmax_rescale_o</*Is_first=*/false, /*Check_inf=*/Is_causal || !Is_even_MN, true, true>(acc_s, acc_o, params.scale_softmax_log2);
//Tensor rP = flash::convert_type<Element>(acc_s);
CONVERT_TENSOR_TYPE(ElementAccum, Element, acc_s, rP)
// Reshape rP from (MMA=4, MMA_M, MMA_N) to ((4, 2), MMA_M, MMA_N / 2)
// if using m16n8k16 or (4, MMA_M, MMA_N) if using m16n8k8.
//Tensor tOrP = make_tensor(rP.data(), flash::convert_layout_acc_Aregs<Kernel_traits::TiledMma>(rP.layout()));
lds4x4_with_swizzle424(tOsVt(_, _, Ksmem_read_index), tOrVt);
CUTE_STATIC_ASSERT_V(size<2>(tOrVt) == _1{}); // only support MMA_K = 1
Tensor tOrVt_permute_view = make_tensor(tOrVt.data(), make_layout(make_shape(size<0>(tOrVt), size<1, 0>(tOrVt), size<1, 1>(tOrVt))));
permute_4x4_b16(tOrVt_permute_view);
Tensor tOrP = make_tensor(rP.data(), acc_s.layout());
flash::gemm_rr(acc_o, tOrP, tOrVt, tiled_mma_o);
Ksmem_read_index ^= 1;
}
// Epilogue
if (NoSplit) {
store_16x16<Kernel_traits, false, Is_even_MN, Is_even_K>(params, bidb, bidh, m_block, n_split_idx, smem_, acc_o, softmax);
}else{
store_16x16<Kernel_traits, true, Is_even_MN, Is_even_K>(params, bidb, bidh, m_block, n_split_idx, smem_, acc_o, softmax);
}
}
} // namespace flash

View File

@ -13,40 +13,16 @@
#include "utils.h"
#include "softmax.h"
#include "mask.h"
#include "rotary.h"
#include "attn_mask.h"
namespace flash {
using namespace cute;
template<typename Kernel_traits, bool Is_causal, bool Is_local, bool Has_alibi, bool Is_even_MN, bool Is_even_K, bool Is_softcap, bool Split, typename Params>
__forceinline__ __device__ void compute_attn_1rowblock_splitkv_k64_mla_32x16_4waves(const Params &params, const int bidb, const int bidh, const int m_block, const int n_split_idx, const int num_n_splits) {
using Element = typename Kernel_traits::Element;
template<typename Kernel_traits, bool Split, bool Is_even_MN, bool Is_even_K, typename AccO, typename Softmax, typename Params>
__forceinline__ __device__ void store_32x16(const Params &params, const int bidb, const int bidh, const int m_block, const int n_split_idx, char* smem_, AccO acc_o, Softmax softmax ) {
using ElementAccum = typename Kernel_traits::ElementAccum;
using Element = typename Kernel_traits::Element;
using index_t = typename Kernel_traits::index_t;
// Shared memory.
extern __shared__ char smem_[];
// The thread index.
const int tidx = threadIdx.x;
const int warp_idx = tidx / 64;
const int lane_idx = tidx % 64;
constexpr int kBlockM = Kernel_traits::kBlockM;
constexpr int kBlockN = Kernel_traits::kBlockN;
constexpr int kHeadDim = Kernel_traits::kHeadDim;
constexpr int kHeadDimV = Kernel_traits::kHeadDimV;
constexpr int kNWarps = Kernel_traits::kNWarps;
constexpr int kBlockKSmem = Kernel_traits::kBlockKSmem;
constexpr int kAtomLayoutMS = Kernel_traits::kAtomLayoutMS;
constexpr int kAtomLayoutMO = Kernel_traits::kAtomLayoutMO;
constexpr int Num_Stages = Kernel_traits::Num_Stages;
static_assert(kBlockKSmem == 64);
using GmemTiledCopyO = std::conditional_t<
!Split,
typename Kernel_traits::GmemTiledCopyO,
@ -54,300 +30,19 @@ __forceinline__ __device__ void compute_attn_1rowblock_splitkv_k64_mla_32x16_4wa
>;
using ElementO = std::conditional_t<!Split, Element, ElementAccum>;
// The thread index.
const int tidx = threadIdx.x;
const int warp_idx = tidx / 64;
const int lane_idx = tidx % 64;
constexpr int kBlockM = Kernel_traits::kBlockM;
constexpr int kBlockN = Kernel_traits::kBlockN;
constexpr int kNWarps = Kernel_traits::kNWarps;
constexpr int kHeadDimV = Kernel_traits::kHeadDimV;
const BlockInfo</*Varlen=*/!Is_even_MN> binfo(params, bidb);
// if (threadIdx.x == 0 && blockIdx.y == 0 && blockIdx.z == 0) { printf("Is_even_MN = %d, is_cumulativ = %d, seqlen_k_cache = %d, actual_seqlen_k = %d\n", Is_even_MN, params.is_seqlens_k_cumulative, binfo.seqlen_k_cache, binfo.actual_seqlen_k); }
// if (threadIdx.x == 0 && blockIdx.y == 1 && blockIdx.z == 0) { printf("params.knew_ptr = %p, seqlen_k_cache + seqlen_knew = %d\n", params.knew_ptr, binfo.seqlen_k_cache + (params.knew_ptr == nullptr ? 0 : params.seqlen_knew)); }
if (m_block * kBlockM >= binfo.actual_seqlen_q) return;
const int n_blocks_per_split = ((binfo.actual_seqlen_k + kBlockN - 1) / kBlockN + num_n_splits - 1) / num_n_splits;
const int n_block_min = !Is_local
? n_split_idx * n_blocks_per_split
: std::max(n_split_idx * n_blocks_per_split, (m_block * kBlockM + binfo.actual_seqlen_k - binfo.actual_seqlen_q - params.window_size_left) / kBlockN);
int n_block_max = std::min(cute::ceil_div(binfo.actual_seqlen_k, kBlockN), (n_split_idx + 1) * n_blocks_per_split);
if (Is_causal || Is_local) {
n_block_max = std::min(n_block_max,
cute::ceil_div((m_block + 1) * kBlockM + binfo.actual_seqlen_k - binfo.actual_seqlen_q / params.ngroups + params.window_size_right, kBlockN));
}
if (n_block_min >= n_block_max) { // This also covers the case where n_block_max <= 0
// We exit early and write 0 to gOaccum and -inf to gLSEaccum.
// Otherwise we might read OOB elements from gK and gV,
// or get wrong results when we combine gOaccum from different blocks.
const index_t row_offset_o = binfo.q_offset(params.o_batch_stride, params.o_row_stride, bidb)
+ m_block * kBlockM * params.o_row_stride + bidh * params.o_head_stride;
const index_t row_offset_oaccum = (((n_split_idx * params.b + bidb) * params.h + bidh) * params.seqlen_q
+ m_block * kBlockM) * params.d_v;
const index_t row_offset_lseaccum = ((n_split_idx * params.b + bidb) * params.h + bidh) * params.seqlen_q + m_block * kBlockM;
Tensor gOaccum = make_tensor(make_gmem_ptr(reinterpret_cast<ElementO *>(Split ? params.oaccum_ptr : params.o_ptr) + (Split ? row_offset_oaccum : row_offset_o)),
Shape<Int<kBlockM>, Int<kHeadDimV>>{},
make_stride(Split ? kHeadDimV : params.o_row_stride, _1{}));
Tensor gLSEaccum = make_tensor(make_gmem_ptr(reinterpret_cast<ElementAccum *>(Split ? params.softmax_lseaccum_ptr : params.softmax_lse_ptr) + row_offset_lseaccum),
Shape<Int<kBlockM>>{}, Stride<_1>{});
GmemTiledCopyO gmem_tiled_copy_Oaccum;
auto gmem_thr_copy_Oaccum = gmem_tiled_copy_Oaccum.get_thread_slice(tidx);
Tensor tOgOaccum = gmem_thr_copy_Oaccum.partition_D(gOaccum);
Tensor tOrOaccum = make_tensor<ElementO>(shape(tOgOaccum));
clear(tOrOaccum);
// Construct identity layout for sO
Tensor cO = make_identity_tensor(make_shape(size<0>(gOaccum), size<1>(gOaccum))); // (BLK_M,BLK_K) -> (blk_m,blk_k)
// Repeat the partitioning with identity layouts
Tensor tOcO = gmem_thr_copy_Oaccum.partition_D(cO);
Tensor tOpO = make_tensor<bool>(make_shape(size<2>(tOgOaccum)));
if (!Is_even_K) {
#pragma unroll
for (int k = 0; k < size(tOpO); ++k) { tOpO(k) = get<1>(tOcO(0, 0, k)) < params.d_v; }
}
// Clear_OOB_K must be false since we don't want to write zeros to gmem
flash::copy<Is_even_MN, Is_even_K, /*Clear_OOB_MN=*/false, /*Clear_OOB_K=*/false>(
gmem_tiled_copy_Oaccum, tOrOaccum, tOgOaccum, tOcO, tOpO, binfo.actual_seqlen_q - m_block * kBlockM
);
#pragma unroll
for (int m = 0; m < size<1>(tOgOaccum); ++m) {
const int row = get<0>(tOcO(0, m, 0));
if (row < binfo.actual_seqlen_q - m_block * kBlockM && get<1>(tOcO(0, m, 0)) == 0) { gLSEaccum(row) = Split ? -INFINITY : INFINITY; }
}
return;
}
// We iterate over the blocks in reverse order. This is because the last block is the only one
// that needs masking when we read K and V from global memory. Moreover, iterating in reverse
// might save us 1 register (we just need n_block instead of both n_block and n_block_max).
const index_t row_offset_q = binfo.q_offset(params.q_batch_stride, params.q_row_stride, bidb)
+ m_block * kBlockM * params.q_row_stride + bidh * params.q_head_stride;
// We move K and V to the last block.
const int bidb_cache = params.cache_batch_idx == nullptr ? bidb : params.cache_batch_idx[bidb];
const int *block_table = params.block_table == nullptr ? nullptr : params.block_table + bidb * params.block_table_batch_stride;
const int block_table_idx = block_table == nullptr ? 0 : (n_block_max - 1) * kBlockN / params.page_block_size;
const int block_table_offset = block_table == nullptr ? 0 : (n_block_max - 1) * kBlockN - block_table_idx * params.page_block_size;
const index_t row_offset_k = block_table == nullptr
? binfo.k_offset(params.k_batch_stride, params.k_row_stride, bidb_cache)
+ (n_block_max - 1) * kBlockN * params.k_row_stride + (bidh / params.h_h_k_ratio) * params.k_head_stride
: (bidh / params.h_h_k_ratio) * params.k_head_stride;
const index_t row_offset_v = block_table == nullptr
? binfo.k_offset(params.v_batch_stride, params.v_row_stride, bidb_cache)
+ (n_block_max - 1) * kBlockN * params.v_row_stride + (bidh / params.h_h_k_ratio) * params.v_head_stride
: (bidh / params.h_h_k_ratio) * params.v_head_stride;
Tensor gQ = make_tensor(make_gmem_ptr(reinterpret_cast<Element *>(params.q_ptr) + row_offset_q),
Shape<Int<kBlockM>, Int<kHeadDim>>{},
make_stride(params.q_row_stride, _1{}));
Tensor gK = make_tensor(make_gmem_ptr(reinterpret_cast<Element *>(params.k_ptr) + row_offset_k),
Shape<Int<kBlockN>, Int<kHeadDim>>{},
make_stride(params.k_row_stride, _1{}));
// if (threadIdx.x == 0 && blockIdx.y == 0 && blockIdx.z == 0) { printf("k_ptr = %p, row_offset_k = %d, gK_ptr = %p\n", params.k_ptr, row_offset_k, gK.data()); }
Tensor gV = make_tensor(make_gmem_ptr(reinterpret_cast<Element *>(params.v_ptr) + row_offset_v),
Shape<Int<kBlockN>, Int<kHeadDimV>>{},
make_stride(params.v_row_stride, _1{}));
Tensor sQ = make_tensor(make_smem_ptr(reinterpret_cast<Element *>(smem_)),
typename Kernel_traits::SmemLayoutQ{});
Tensor sK = make_tensor(sQ.data() + (Kernel_traits::Share_Q_K_smem ? 0 : size(sQ)),
typename Kernel_traits::SmemLayoutK{});
Tensor sV = make_tensor(sK.data(), typename Kernel_traits::SmemLayoutVtNoSwizzle{});
Tensor sVt = make_tensor(sV.data(), typename Kernel_traits::SmemLayoutVtransposed{});
Tensor sVtNoSwizzle = make_tensor(sV.data(), typename Kernel_traits::SmemLayoutVtransposedNoSwizzle{});
typename Kernel_traits::GmemTiledCopyB128 gmem_tiled_copy_Q;
auto gmem_thr_copy_Q = gmem_tiled_copy_Q.get_thread_slice(tidx);
Tensor tQgQ = gmem_thr_copy_Q.partition_S(gQ);
Tensor tQsQ = gmem_thr_copy_Q.partition_D(sQ);
typename Kernel_traits::GmemTiledCopyB64 gmem_tiled_copy_KV;
auto gmem_thr_copy_KV = gmem_tiled_copy_KV.get_thread_slice(tidx);
Tensor tKgK = gmem_thr_copy_KV.partition_S(gK); // (KCPY, KCPY_N, KCPY_K)
Tensor tKsK = gmem_thr_copy_KV.partition_D(sK);
// wave0 and wave2 compute the same S, wave1 and wave3 compute the same S
int tidx_mma_s = tidx & 0x7F;
typename Kernel_traits::TiledMmaS tiled_mma_s;
auto thr_mma_s = tiled_mma_s.get_thread_slice(tidx_mma_s);
Tensor tSrQ = thr_mma_s.partition_fragment_A(sQ); // (MMA,MMA_M,MMA_K)
Tensor tSrK = thr_mma_s.partition_fragment_B(sK(_, _, 0)); // (MMA,MMA_N,MMA_K)
typename Kernel_traits::TiledMmaO tiled_mma_o;
auto thr_mma_o = tiled_mma_o.get_thread_slice(tidx);
// Tensor tOrVt = thr_mma_o.partition_fragment_B(sVt); // (MMA, MMA_K,MMA_N)
Tensor tOrVt = make_tensor<Element>(Shape<_4, Shape<_4, _4>, _1>{});
Tensor acc_o = partition_fragment_C(tiled_mma_o, Shape<Int<kBlockM>, Int<kHeadDimV>>{}); // MMA, MMA_M, MMA_K
//
// Copy Atom retiling
//
auto smem_tiled_copy_Q = make_tiled_copy_A(typename Kernel_traits::SmemCopyAtomB64{}, tiled_mma_s);
auto smem_thr_copy_Q = smem_tiled_copy_Q.get_thread_slice(tidx_mma_s);
Tensor tSsQ = smem_thr_copy_Q.partition_S(sQ);
auto smem_tiled_copy_K = make_tiled_copy_B(typename Kernel_traits::SmemCopyAtomB64{}, tiled_mma_s);
auto smem_thr_copy_K = smem_tiled_copy_K.get_thread_slice(tidx_mma_s);
Tensor tSsK = smem_thr_copy_K.partition_S(sK);
auto smem_tiled_copy_V = make_tiled_copy_B(typename Kernel_traits::SmemCopyAtomTransposed{}, tiled_mma_o);
auto smem_thr_copy_V = smem_tiled_copy_V.get_thread_slice(tidx);
int warp_offset = warp_idx / kAtomLayoutMO * 16 * 64;
int thread_offset = lane_idx / 16 * 4 * 64;
Element *Vtsmem_ptr_lds = reinterpret_cast<Element *>(sVt.data().get()) + warp_offset + thread_offset;
Tensor tOsVt = make_tensor(make_smem_ptr(Vtsmem_ptr_lds), make_layout(Shape<_4, _4, Int<Num_Stages>>{}, // MMA MMA_N NUM_STAGES
Stride<_1, Int<16*128>, Int<kBlockN*kHeadDim>>{}));
// PREDICATES
// Construct identity layout for sQ and sK
Tensor cQ = make_identity_tensor(make_shape(size<0>(sQ), size<1>(sQ))); // (BLK_M,BLK_K) -> (blk_m,blk_k)
Tensor cKV = make_identity_tensor(make_shape(size<0>(sK), size<1>(sK))); // (BLK_N,BLK_K) -> (blk_n,blk_k)
// Repeat the partitioning with identity layouts
Tensor tQcQ = gmem_thr_copy_Q.partition_S(cQ); // (ACPY,ACPY_M,ACPY_K) -> (blk_m,blk_k)
Tensor tKVcKV = gmem_thr_copy_KV.partition_S(cKV); // (BCPY,BCPY_N,BCPY_K) -> (blk_n,blk_k)
// Prologue
// Read Q from gmem to smem, optionally apply rotary embedding.
Tensor tQrQ = make_fragment_like(tQgQ);
// We don't need to clear the sQ smem tiles since we'll only write out the valid outputs
flash::copy_b128<Is_even_MN, Is_even_K>(tQgQ, tQrQ, tQcQ, params.d, binfo.actual_seqlen_q - m_block * kBlockM);
cute::copy(tQrQ, tQsQ);
if constexpr (Kernel_traits::Is_Q_in_regs) {
flash::sync_threads();
cute::copy(smem_tiled_copy_Q, tSsQ, tSrQ);
flash::sync_threads();
}
int n_block = n_block_max - 1;
int Ksmem_read_index = 0;
int Ksmem_write_index = 0;
// We don't need to clear the sK smem tiles since we'll mask out the scores anyway.
Tensor tKrK = make_fragment_like(tKgK);
flash::copy_b64_page_one<Kernel_traits, Is_even_MN, Is_even_K>(gK, tKgK, tKrK, tKVcKV, params.d, n_block,
block_table, params.k_batch_stride, params.k_row_stride, params.page_block_size, binfo.actual_seqlen_k - n_block * kBlockN);
// flash::cp_async_wait<0>();
// __syncthreads();
// if (tidx == 0 && blockIdx.y == 0 && blockIdx.z == 0) { print(tKsK); }
// __syncthreads();
clear(acc_o);
flash::Softmax<size<1>(acc_o)> softmax;
const float alibi_slope = !Has_alibi ? 0.0f : reinterpret_cast<float *>(params.alibi_slopes_ptr)[bidb * params.alibi_slopes_batch_stride + bidh] / params.scale_softmax;
flash::Mask<Is_causal, Is_local, Has_alibi> mask(binfo.actual_seqlen_k, binfo.actual_seqlen_q, params.ngroups, params.window_size_left, params.window_size_right, alibi_slope);
// For performance reason, we separate out two kinds of iterations:
// those that need masking on S, and those that don't.
// We need masking on S for the very last block when K and V has length not multiple of kBlockN.
// We also need masking on S if it's causal, for the last ceil_div(kBlockM, kBlockN) blocks.
// We will have at least 1 "masking" iteration.
// If not even_N, then seqlen_k might end in the middle of a block. In that case we need to
// mask 2 blocks (e.g. when kBlockM == kBlockN), not just 1.
constexpr int n_masking_steps = (!Is_causal && !Is_local)
? 1
: ((Is_even_MN && Is_causal) ? cute::ceil_div(kBlockM, kBlockN) : cute::ceil_div(kBlockM, kBlockN) + 1);
#pragma unroll
for (int masking_step = 0; masking_step < n_masking_steps; ++masking_step, --n_block) {
Tensor acc_s = partition_fragment_C(tiled_mma_s, Shape<Int<kBlockM>, Int<kBlockN>>{}); // (MMA=4, MMA_M, MMA_N)
cute::copy(tKrK, tKsK(_, _, _, Ksmem_write_index));
Ksmem_write_index ^= 1;
clear(acc_s);
flash::sync_threads();
if (n_block > n_block_min) {
flash::copy_b64_page_one<Kernel_traits, /*Is_even_MN=*/true, Is_even_K>(gK, tKgK, tKrK, tKVcKV, params.d, n_block - 1,
block_table, params.k_batch_stride, params.k_row_stride, params.page_block_size);
}
flash::gemm_opt</*A_in_regs=*/Kernel_traits::Is_Q_in_regs>(
acc_s, tSrQ, tSrK, tSsQ, tSsK(_, _, _, Ksmem_read_index), tiled_mma_s, smem_tiled_copy_Q, smem_tiled_copy_K,
smem_thr_copy_Q, smem_thr_copy_K
);
// if (cute::thread0()) { print(acc_s); }
if constexpr (Is_softcap){
flash::apply_softcap(acc_s, params.softcap);
}
mask.template apply_mask<Is_causal, Is_even_MN>(
acc_s, n_block * kBlockN, m_block * kBlockM + (tidx / 64) % kAtomLayoutMS * 16 + (tidx & 0xf), kAtomLayoutMS * 16
);
// We have key_padding_mask so we'll need to Check_inf
masking_step == 0
? softmax.template softmax_rescale_o</*Is_first=*/true, /*Check_inf=*/Is_causal || Is_local || !Is_even_MN, true, true>(acc_s, acc_o, params.scale_softmax_log2)
: softmax.template softmax_rescale_o</*Is_first=*/false, /*Check_inf=*/Is_causal || Is_local || !Is_even_MN, true, true>(acc_s, acc_o, params.scale_softmax_log2);
// if (cute::thread0()) { print(scores_max); print(scores_sum); print(scores); }
// Convert acc_s from fp32 to fp16/bf16
//Tensor rP = flash::convert_type<Element>(acc_s);
CONVERT_TENSOR_TYPE(ElementAccum, Element, acc_s, rP)
// Reshape rP from (MMA=4, MMA_M, MMA_N) to ((4, 2), MMA_M, MMA_N / 2)
// if using m16n8k16 or (4, MMA_M, MMA_N) if using m16n8k8.
//Tensor tOrP = make_tensor(rP.data(), flash::convert_layout_acc_Aregs<Kernel_traits::TiledMma>(rP.layout()));
lds4x4_with_swizzle424(tOsVt(_, _, Ksmem_read_index), tOrVt);
CUTE_STATIC_ASSERT_V(size<2>(tOrVt) == _1{}); // only support MMA_K = 1
Tensor tOrVt_permute_view = make_tensor(tOrVt.data(), make_layout(make_shape(size<0>(tOrVt), size<1, 0>(tOrVt), size<1, 1>(tOrVt))));
permute_4x4_b16(tOrVt_permute_view);
Tensor tOrP = make_tensor(rP.data(), acc_s.layout());
flash::gemm_rr(acc_o, tOrP, tOrVt, tiled_mma_o);
Ksmem_read_index ^= 1;
// This check is at the end of the loop since we always have at least 1 iteration
if (n_masking_steps > 1 && n_block <= n_block_min) {
--n_block;
break;
}
}
// These are the iterations where we don't need masking on S
for (; n_block >= n_block_min; --n_block) {
Tensor acc_s = partition_fragment_C(tiled_mma_s, Shape<Int<kBlockM>, Int<kBlockN>>{}); // (MMA=4, MMA_M, MMA_N)
cute::copy(tKrK, tKsK(_, _, _, Ksmem_write_index));
Ksmem_write_index ^= 1;
clear(acc_s);
flash::sync_threads();
if (n_block > n_block_min) {
// Advance gK
flash::copy_b64_page_one<Kernel_traits, /*Is_even_MN=*/true, Is_even_K>(gK, tKgK, tKrK, tKVcKV, params.d, n_block - 1,
block_table, params.k_batch_stride, params.k_row_stride, params.page_block_size);
}
flash::gemm_opt</*A_in_regs=*/Kernel_traits::Is_Q_in_regs>(
acc_s, tSrQ, tSrK, tSsQ, tSsK(_, _, _, Ksmem_read_index), tiled_mma_s, smem_tiled_copy_Q, smem_tiled_copy_K,
smem_thr_copy_Q, smem_thr_copy_K
);
if constexpr (Is_softcap){
flash::apply_softcap(acc_s, params.softcap);
}
mask.template apply_mask<Is_causal, Is_even_MN>(
acc_s, n_block * kBlockN, m_block * kBlockM + (tidx / 64) % kAtomLayoutMS * 16 + (tidx & 0xf), kAtomLayoutMS * 16
);
softmax.template softmax_rescale_o</*Is_first=*/false, /*Check_inf=*/Is_causal || Is_local || !Is_even_MN, true, true>(acc_s, acc_o, params.scale_softmax_log2);
//Tensor rP = flash::convert_type<Element>(acc_s);
CONVERT_TENSOR_TYPE(ElementAccum, Element, acc_s, rP)
// Reshape rP from (MMA=4, MMA_M, MMA_N) to ((4, 2), MMA_M, MMA_N / 2)
// if using m16n8k16 or (4, MMA_M, MMA_N) if using m16n8k8.
//Tensor tOrP = make_tensor(rP.data(), flash::convert_layout_acc_Aregs<Kernel_traits::TiledMma>(rP.layout()));
lds4x4_with_swizzle424(tOsVt(_, _, Ksmem_read_index), tOrVt);
CUTE_STATIC_ASSERT_V(size<2>(tOrVt) == _1{}); // only support MMA_K = 1
Tensor tOrVt_permute_view = make_tensor(tOrVt.data(), make_layout(make_shape(size<0>(tOrVt), size<1, 0>(tOrVt), size<1, 1>(tOrVt))));
permute_4x4_b16(tOrVt_permute_view);
Tensor tOrP = make_tensor(rP.data(), acc_s.layout());
flash::gemm_rr(acc_o, tOrP, tOrVt, tiled_mma_o);
Ksmem_read_index ^= 1;
}
// Epilogue
Tensor lse = softmax.template normalize_softmax_lse</*Is_dropout=*/false, /*Return_lse*/true, Split>(acc_o, params.scale_softmax);
Tensor acc_o_view = make_tensor(acc_o.data(), make_layout(Shape<_4, Shape<_4, _4>>{},
Stride<_1, Shape<_4, _16>>{}));
@ -361,7 +56,6 @@ __forceinline__ __device__ void compute_attn_1rowblock_splitkv_k64_mla_32x16_4wa
acc_o_copy(row, make_coord(col, k)) = acc_o_view(col, make_coord(row, k));
}
}
// if (cute::thread0()) { print(lse); }
if constexpr (!Split) {
Tensor sOaccum = make_tensor(make_smem_ptr(reinterpret_cast<ElementO *>(smem_)), typename Kernel_traits::SmemLayoutO{}); // (SMEM_M,SMEM_N)
// Partition sO to match the accumulator partitioning
@ -376,8 +70,6 @@ __forceinline__ __device__ void compute_attn_1rowblock_splitkv_k64_mla_32x16_4wa
Stride<_1, _16>{}));
// sOaccum is larger than sQ, so we need to syncthreads here
// TODO: allocate enough smem for sOaccum
if constexpr (Kernel_traits::Share_Q_K_smem) { flash::sync_threads(); }
cute::copy(tOrO, tOsO);
@ -426,10 +118,11 @@ __forceinline__ __device__ void compute_attn_1rowblock_splitkv_k64_mla_32x16_4wa
tOrOaccum, tOgOaccum, tOcO, params.d_v, binfo.actual_seqlen_q - m_block * kBlockM
);
} else {
const int split_offset = params.num_splits_ptr[bidb];
Tensor sOaccum = make_tensor(make_smem_ptr(reinterpret_cast<ElementO *>(smem_)), typename Kernel_traits::SmemLayoutO{}); // (SMEM_M,SMEM_N)
const index_t row_offset_oaccum = (((n_split_idx * params.b + bidb) * params.h + bidh) * params.seqlen_q
const index_t row_offset_oaccum = (((split_offset + n_split_idx) * params.h + bidh) * params.seqlen_q
+ m_block * kBlockM) * params.d_v;
const index_t row_offset_lseaccum = ((n_split_idx * params.b + bidb) * params.h + bidh) * params.seqlen_q + m_block * kBlockM;
const index_t row_offset_lseaccum = ((split_offset + n_split_idx) * params.h + bidh) * params.seqlen_q + m_block * kBlockM;
Tensor gOaccum = make_tensor(make_gmem_ptr(reinterpret_cast<ElementO *>(params.oaccum_ptr) + row_offset_oaccum),
Shape<Int<kBlockM>, Int<kHeadDimV>>{},
@ -491,5 +184,292 @@ __forceinline__ __device__ void compute_attn_1rowblock_splitkv_k64_mla_32x16_4wa
);
}
}
template<typename Kernel_traits, bool Is_causal, bool Is_even_MN, bool Is_even_K, bool Is_enable_dcp, typename Params>
__forceinline__ __device__ void compute_attn_1rowblock_splitkv_mla_k64_32x16_4waves_xcore1000(
const Params &params, const int bidb, const int bidh, const int m_block, const int n_split_idx, const int n_block_min, const int n_block_max, const bool NoSplit) {
using Element = typename Kernel_traits::Element;
using ElementAccum = typename Kernel_traits::ElementAccum;
using index_t = typename Kernel_traits::index_t;
// Shared memory.
extern __shared__ char smem_[];
// The thread index.
const int tidx = threadIdx.x;
const int warp_idx = tidx / 64;
const int lane_idx = tidx % 64;
constexpr int kBlockM = Kernel_traits::kBlockM;
constexpr int kBlockN = Kernel_traits::kBlockN;
constexpr int kHeadDim = Kernel_traits::kHeadDim;
constexpr int kHeadDimV = Kernel_traits::kHeadDimV;
constexpr int kNWarps = Kernel_traits::kNWarps;
constexpr int kBlockKSmem = Kernel_traits::kBlockKSmem;
constexpr int kAtomLayoutMS = Kernel_traits::kAtomLayoutMS;
constexpr int kAtomLayoutMO = Kernel_traits::kAtomLayoutMO;
constexpr int Num_Stages = Kernel_traits::Num_Stages;
static_assert(kBlockKSmem == 64);
const BlockInfo</*Varlen=*/!Is_even_MN> binfo(params, bidb);
// if (threadIdx.x == 0 && blockIdx.y == 0 && blockIdx.z == 0) { printf("Is_even_MN = %d, is_cumulativ = %d, seqlen_k_cache = %d, actual_seqlen_k = %d\n", Is_even_MN, params.is_seqlens_k_cumulative, binfo.seqlen_k_cache, binfo.actual_seqlen_k); }
// if (threadIdx.x == 0 && blockIdx.y == 1 && blockIdx.z == 0) { printf("params.knew_ptr = %p, seqlen_k_cache + seqlen_knew = %d\n", params.knew_ptr, binfo.seqlen_k_cache + (params.knew_ptr == nullptr ? 0 : params.seqlen_knew)); }
if (m_block * kBlockM >= binfo.actual_seqlen_q) return;
// We iterate over the blocks in reverse order. This is because the last block is the only one
// that needs masking when we read K and V from global memory. Moreover, iterating in reverse
// might save us 1 register (we just need n_block instead of both n_block and n_block_max).
const index_t row_offset_q = binfo.q_offset(params.q_batch_stride, params.q_row_stride, bidb)
+ m_block * kBlockM * params.q_row_stride + bidh * params.q_head_stride;
// We move K and V to the last block.
const int bidb_cache = params.cache_batch_idx == nullptr ? bidb : params.cache_batch_idx[bidb];
const int *block_table = params.block_table == nullptr ? nullptr : params.block_table + bidb * params.block_table_batch_stride;
const int block_table_idx = block_table == nullptr ? 0 : (n_block_max - 1) * kBlockN / params.page_block_size;
const int block_table_offset = block_table == nullptr ? 0 : (n_block_max - 1) * kBlockN - block_table_idx * params.page_block_size;
const index_t row_offset_k = block_table == nullptr
? binfo.k_offset(params.k_batch_stride, params.k_row_stride, bidb_cache)
+ (n_block_max - 1) * kBlockN * params.k_row_stride + (bidh / params.h_h_k_ratio) * params.k_head_stride
: (bidh / params.h_h_k_ratio) * params.k_head_stride;
const index_t row_offset_v = block_table == nullptr
? binfo.k_offset(params.v_batch_stride, params.v_row_stride, bidb_cache)
+ (n_block_max - 1) * kBlockN * params.v_row_stride + (bidh / params.h_h_k_ratio) * params.v_head_stride
: (bidh / params.h_h_k_ratio) * params.v_head_stride;
Tensor gQ = make_tensor(make_gmem_ptr(reinterpret_cast<Element *>(params.q_ptr) + row_offset_q),
Shape<Int<kBlockM>, Int<kHeadDim>>{},
make_stride(params.q_row_stride, _1{}));
Tensor gK = make_tensor(make_gmem_ptr(reinterpret_cast<Element *>(params.k_ptr) + row_offset_k),
Shape<Int<kBlockN>, Int<kHeadDim>>{},
make_stride(params.k_row_stride, _1{}));
// if (threadIdx.x == 0 && blockIdx.y == 0 && blockIdx.z == 0) { printf("k_ptr = %p, row_offset_k = %d, gK_ptr = %p\n", params.k_ptr, row_offset_k, gK.data()); }
Tensor gV = make_tensor(make_gmem_ptr(reinterpret_cast<Element *>(params.v_ptr) + row_offset_v),
Shape<Int<kBlockN>, Int<kHeadDimV>>{},
make_stride(params.v_row_stride, _1{}));
Tensor sQ = make_tensor(make_smem_ptr(reinterpret_cast<Element *>(smem_)),
typename Kernel_traits::SmemLayoutQ{});
Tensor sK = make_tensor(sQ.data() + (Kernel_traits::Share_Q_K_smem ? 0 : size(sQ)),
typename Kernel_traits::SmemLayoutK424{});
Tensor sV = make_tensor(sK.data(), typename Kernel_traits::SmemLayoutVtNoSwizzle{});
Tensor sVt = make_tensor(sV.data(), typename Kernel_traits::SmemLayoutVtransposed424{});
Tensor sVtNoSwizzle = make_tensor(sV.data(), typename Kernel_traits::SmemLayoutVtransposedNoSwizzle{});
typename Kernel_traits::GmemTiledCopyB128 gmem_tiled_copy_Q;
auto gmem_thr_copy_Q = gmem_tiled_copy_Q.get_thread_slice(tidx);
Tensor tQgQ = gmem_thr_copy_Q.partition_S(gQ);
Tensor tQsQ = gmem_thr_copy_Q.partition_D(sQ);
typename Kernel_traits::GmemTiledCopyB64 gmem_tiled_copy_KV;
auto gmem_thr_copy_KV = gmem_tiled_copy_KV.get_thread_slice(tidx);
Tensor tKgK = gmem_thr_copy_KV.partition_S(gK); // (KCPY, KCPY_N, KCPY_K)
Tensor tKsK = gmem_thr_copy_KV.partition_D(sK);
// wave0 and wave2 compute the same S, wave1 and wave3 compute the same S
int tidx_mma_s = tidx & 0x7F;
typename Kernel_traits::TiledMmaS tiled_mma_s;
auto thr_mma_s = tiled_mma_s.get_thread_slice(tidx_mma_s);
Tensor tSrQ = thr_mma_s.partition_fragment_A(sQ); // (MMA,MMA_M,MMA_K)
Tensor tSrK = thr_mma_s.partition_fragment_B(sK(_, _, 0)); // (MMA,MMA_N,MMA_K)
typename Kernel_traits::TiledMmaO tiled_mma_o;
auto thr_mma_o = tiled_mma_o.get_thread_slice(tidx);
// Tensor tOrVt = thr_mma_o.partition_fragment_B(sVt); // (MMA, MMA_K,MMA_N)
Tensor tOrVt = make_tensor<Element>(Shape<_4, Shape<_4, _4>, _1>{});
Tensor acc_o = partition_fragment_C(tiled_mma_o, Shape<Int<kBlockM>, Int<kHeadDimV>>{}); // MMA, MMA_M, MMA_K
//
// Copy Atom retiling
//
auto smem_tiled_copy_Q = make_tiled_copy_A(typename Kernel_traits::UniversalCopyAtomB64{}, tiled_mma_s);
auto smem_thr_copy_Q = smem_tiled_copy_Q.get_thread_slice(tidx_mma_s);
Tensor tSsQ = smem_thr_copy_Q.partition_S(sQ);
auto smem_tiled_copy_K = make_tiled_copy_B(typename Kernel_traits::UniversalCopyAtomB64{}, tiled_mma_s);
auto smem_thr_copy_K = smem_tiled_copy_K.get_thread_slice(tidx_mma_s);
Tensor tSsK = smem_thr_copy_K.partition_S(sK);
auto smem_tiled_copy_V = make_tiled_copy_B(typename Kernel_traits::SmemCopyAtomTransposed{}, tiled_mma_o);
auto smem_thr_copy_V = smem_tiled_copy_V.get_thread_slice(tidx);
int warp_offset = warp_idx / kAtomLayoutMO * 16 * 64;
int thread_offset = lane_idx / 16 * 4 * 64;
Element *Vtsmem_ptr_lds = reinterpret_cast<Element *>(sVt.data().get()) + warp_offset + thread_offset;
Tensor tOsVt = make_tensor(make_smem_ptr(Vtsmem_ptr_lds), make_layout(Shape<_4, _4, Int<Num_Stages>>{}, // MMA MMA_N NUM_STAGES
Stride<_1, Int<16*128>, Int<kBlockN*kHeadDim>>{}));
// PREDICATES
// Construct identity layout for sQ and sK
Tensor cQ = make_identity_tensor(make_shape(size<0>(sQ), size<1>(sQ))); // (BLK_M,BLK_K) -> (blk_m,blk_k)
Tensor cKV = make_identity_tensor(make_shape(size<0>(sK), size<1>(sK))); // (BLK_N,BLK_K) -> (blk_n,blk_k)
// Repeat the partitioning with identity layouts
Tensor tQcQ = gmem_thr_copy_Q.partition_S(cQ); // (ACPY,ACPY_M,ACPY_K) -> (blk_m,blk_k)
Tensor tKVcKV = gmem_thr_copy_KV.partition_S(cKV); // (BCPY,BCPY_N,BCPY_K) -> (blk_n,blk_k)
// Prologue
// Read Q from gmem to smem, optionally apply rotary embedding.
Tensor tQrQ = make_fragment_like(tQgQ);
// We don't need to clear the sQ smem tiles since we'll only write out the valid outputs
flash::copy_b128<Is_even_MN, Is_even_K>(tQgQ, tQrQ, tQcQ, params.d, binfo.actual_seqlen_q - m_block * kBlockM);
cute::copy(tQrQ, tQsQ);
if constexpr (Kernel_traits::Is_Q_in_regs) {
flash::sync_threads();
cute::copy(smem_tiled_copy_Q, tSsQ, tSrQ);
flash::sync_threads();
}
int n_block = n_block_max - 1;
int Ksmem_read_index = 0;
int Ksmem_write_index = 0;
// We don't need to clear the sK smem tiles since we'll mask out the scores anyway.
Tensor tKrK = make_fragment_like(tKgK);
int row_offset = tidx / 16 + n_block * kBlockN;
int virtual_page_idx = row_offset / params.page_block_size;
int page_offset = row_offset - virtual_page_idx * params.page_block_size;
int32_t page_idx = block_table[virtual_page_idx];
flash::copy_b64_page_one<Kernel_traits, Is_even_MN, Is_even_K>(gK, tKgK, tKrK, tKVcKV, params.d, n_block,
block_table, params.k_batch_stride, params.k_row_stride, params.page_block_size, page_idx, page_offset, binfo.actual_seqlen_k - n_block * kBlockN);
// flash::cp_async_wait<0>();
// __syncthreads();
// if (tidx == 0 && blockIdx.y == 0 && blockIdx.z == 0) { print(tKsK); }
// __syncthreads();
clear(acc_o);
flash::Softmax<size<1>(acc_o)> softmax;
flash::Mask<Is_causal, Is_enable_dcp> mask(binfo.actual_seqlen_k, binfo.actual_seqlen_q, params.ngroups, binfo.tot_seqlen_k, params.cp_world_size, params.cp_rank);
// For performance reason, we separate out two kinds of iterations:
// those that need masking on S, and those that don't.
// We need masking on S for the very last block when K and V has length not multiple of kBlockN.
// We also need masking on S if it's causal, for the last ceil_div(kBlockM, kBlockN) blocks.
// We will have at least 1 "masking" iteration.
// If not even_N, then seqlen_k might end in the middle of a block. In that case we need to
// mask 2 blocks (e.g. when kBlockM == kBlockN), not just 1.
constexpr int n_masking_steps = (!Is_causal)
? 1
: ((Is_even_MN && Is_causal) ? cute::ceil_div(kBlockM, kBlockN) : cute::ceil_div(kBlockM, kBlockN) + 1);
#pragma unroll
for (int masking_step = 0; masking_step < n_masking_steps; ++masking_step, --n_block) {
if (n_block > n_block_min) {
// prefetch load page index
row_offset -= kBlockN;
virtual_page_idx = row_offset / params.page_block_size;
page_offset = row_offset - virtual_page_idx * params.page_block_size;
page_idx = block_table[virtual_page_idx];
}
Tensor acc_s = partition_fragment_C(tiled_mma_s, Shape<Int<kBlockM>, Int<kBlockN>>{}); // (MMA=4, MMA_M, MMA_N)
cute::copy(tKrK, tKsK(_, _, _, Ksmem_write_index));
Ksmem_write_index ^= 1;
clear(acc_s);
flash::sync_threads();
if (n_block > n_block_min) {
flash::copy_b64_page_one<Kernel_traits, /*Is_even_MN=*/true, Is_even_K>(gK, tKgK, tKrK, tKVcKV, params.d, n_block - 1,
block_table, params.k_batch_stride, params.k_row_stride, params.page_block_size, page_idx, page_offset);
}
flash::gemm</*A_in_regs=*/Kernel_traits::Is_Q_in_regs>(
acc_s, tSrQ, tSrK, tSsQ, tSsK(_, _, _, Ksmem_read_index), tiled_mma_s, smem_tiled_copy_Q, smem_tiled_copy_K,
smem_thr_copy_Q, smem_thr_copy_K
);
// if (cute::thread0()) { print(acc_s); }
mask.template apply_mask<Is_causal, Is_even_MN>(
acc_s, n_block * kBlockN, m_block * kBlockM + (tidx / 64) % kAtomLayoutMS * 16 + (tidx & 0xf), kAtomLayoutMS * 16
);
// We have key_padding_mask so we'll need to Check_inf
masking_step == 0
? softmax.template softmax_rescale_o</*Is_first=*/true, /*Check_inf=*/Is_causal || !Is_even_MN, true, true>(acc_s, acc_o, params.scale_softmax_log2)
: softmax.template softmax_rescale_o</*Is_first=*/false, /*Check_inf=*/Is_causal || !Is_even_MN, true, true>(acc_s, acc_o, params.scale_softmax_log2);
// if (cute::thread0()) { print(scores_max); print(scores_sum); print(scores); }
// Convert acc_s from fp32 to fp16/bf16
//Tensor rP = flash::convert_type<Element>(acc_s);
CONVERT_TENSOR_TYPE(ElementAccum, Element, acc_s, rP)
// Reshape rP from (MMA=4, MMA_M, MMA_N) to ((4, 2), MMA_M, MMA_N / 2)
// if using m16n8k16 or (4, MMA_M, MMA_N) if using m16n8k8.
//Tensor tOrP = make_tensor(rP.data(), flash::convert_layout_acc_Aregs<Kernel_traits::TiledMma>(rP.layout()));
lds4x4_with_swizzle424(tOsVt(_, _, Ksmem_read_index), tOrVt);
CUTE_STATIC_ASSERT_V(size<2>(tOrVt) == _1{}); // only support MMA_K = 1
Tensor tOrVt_permute_view = make_tensor(tOrVt.data(), make_layout(make_shape(size<0>(tOrVt), size<1, 0>(tOrVt), size<1, 1>(tOrVt))));
permute_4x4_b16(tOrVt_permute_view);
Tensor tOrP = make_tensor(rP.data(), acc_s.layout());
flash::gemm_rr(acc_o, tOrP, tOrVt, tiled_mma_o);
Ksmem_read_index ^= 1;
// This check is at the end of the loop since we always have at least 1 iteration
if (n_masking_steps > 1 && n_block <= n_block_min) {
--n_block;
break;
}
}
// These are the iterations where we don't need masking on S
for (; n_block >= n_block_min; --n_block) {
if (n_block > n_block_min) {
// prefetch load page index
row_offset -= kBlockN;
virtual_page_idx = row_offset / params.page_block_size;
page_offset = row_offset - virtual_page_idx * params.page_block_size;
page_idx = block_table[virtual_page_idx];
}
Tensor acc_s = partition_fragment_C(tiled_mma_s, Shape<Int<kBlockM>, Int<kBlockN>>{}); // (MMA=4, MMA_M, MMA_N)
cute::copy(tKrK, tKsK(_, _, _, Ksmem_write_index));
Ksmem_write_index ^= 1;
clear(acc_s);
flash::sync_threads();
if (n_block > n_block_min) {
// Advance gK
flash::copy_b64_page_one<Kernel_traits, /*Is_even_MN=*/true, Is_even_K>(gK, tKgK, tKrK, tKVcKV, params.d, n_block - 1,
block_table, params.k_batch_stride, params.k_row_stride, params.page_block_size, page_idx, page_offset);
}
flash::gemm</*A_in_regs=*/Kernel_traits::Is_Q_in_regs>(
acc_s, tSrQ, tSrK, tSsQ, tSsK(_, _, _, Ksmem_read_index), tiled_mma_s, smem_tiled_copy_Q, smem_tiled_copy_K,
smem_thr_copy_Q, smem_thr_copy_K
);
mask.template apply_mask<Is_causal, Is_even_MN>(
acc_s, n_block * kBlockN, m_block * kBlockM + (tidx / 64) % kAtomLayoutMS * 16 + (tidx & 0xf), kAtomLayoutMS * 16
);
softmax.template softmax_rescale_o</*Is_first=*/false, /*Check_inf=*/Is_causal || !Is_even_MN, true, true>(acc_s, acc_o, params.scale_softmax_log2);
//Tensor rP = flash::convert_type<Element>(acc_s);
CONVERT_TENSOR_TYPE(ElementAccum, Element, acc_s, rP)
// Reshape rP from (MMA=4, MMA_M, MMA_N) to ((4, 2), MMA_M, MMA_N / 2)
// if using m16n8k16 or (4, MMA_M, MMA_N) if using m16n8k8.
//Tensor tOrP = make_tensor(rP.data(), flash::convert_layout_acc_Aregs<Kernel_traits::TiledMma>(rP.layout()));
lds4x4_with_swizzle424(tOsVt(_, _, Ksmem_read_index), tOrVt);
CUTE_STATIC_ASSERT_V(size<2>(tOrVt) == _1{}); // only support MMA_K = 1
Tensor tOrVt_permute_view = make_tensor(tOrVt.data(), make_layout(make_shape(size<0>(tOrVt), size<1, 0>(tOrVt), size<1, 1>(tOrVt))));
permute_4x4_b16(tOrVt_permute_view);
Tensor tOrP = make_tensor(rP.data(), acc_s.layout());
flash::gemm_rr(acc_o, tOrP, tOrVt, tiled_mma_o);
Ksmem_read_index ^= 1;
}
// Epilogue
// if (cute::thread0()) { print(lse); }
if (NoSplit) {
store_32x16<Kernel_traits, false, Is_even_MN, Is_even_K>(params, bidb, bidh, m_block, n_split_idx, smem_, acc_o, softmax);
}else{
store_32x16<Kernel_traits, true, Is_even_MN, Is_even_K>(params, bidb, bidh, m_block, n_split_idx, smem_, acc_o, softmax);
}
}
} // namespace flash

View File

@ -13,297 +13,35 @@
#include "utils.h"
#include "softmax.h"
#include "mask.h"
#include "rotary.h"
#include "attn_mask.h"
namespace flash {
using namespace cute;
template<typename Kernel_traits, bool Is_causal, bool Is_local, bool Has_alibi, bool Is_even_MN, bool Is_even_K, bool Is_softcap, bool Split, typename Params>
__forceinline__ __device__ void compute_attn_1rowblock_splitkv_k64_mla_64x16_8waves(const Params &params, const int bidb, const int bidh, const int m_block, const int n_split_idx, const int num_n_splits) {
using Element = typename Kernel_traits::Element;
template<typename Kernel_traits, bool Split, bool Is_even_MN, bool Is_even_K, typename AccO, typename Softmax, typename Params>
__forceinline__ __device__ void store_64x16_xcore1000(const Params &params, const int bidb, const int bidh, const int m_block, const int n_split_idx, char* smem_, AccO acc_o, Softmax softmax){
using ElementAccum = typename Kernel_traits::ElementAccum;
using Element = typename Kernel_traits::Element;
using index_t = typename Kernel_traits::index_t;
// Shared memory.
extern __shared__ char smem_[];
// The thread index.
const int tidx = threadIdx.x;
const int warp_idx = tidx / 64;
const int lane_idx = tidx % 64;
constexpr int kBlockM = Kernel_traits::kBlockM;
constexpr int kBlockN = Kernel_traits::kBlockN;
constexpr int kHeadDim = Kernel_traits::kHeadDim;
constexpr int kHeadDimV = Kernel_traits::kHeadDimV;
constexpr int kNWarps = Kernel_traits::kNWarps;
constexpr int kBlockKSmem = Kernel_traits::kBlockKSmem;
constexpr int kAtomLayoutMS = Kernel_traits::kAtomLayoutMS;
constexpr int kAtomLayoutMO = Kernel_traits::kAtomLayoutMO;
constexpr int Num_Stages = Kernel_traits::Num_Stages;
constexpr int kHeadDimNope = kHeadDimV;
constexpr int kHeadDimRope = kHeadDim - kHeadDimV;
static_assert(kBlockKSmem == 64);
static_assert(Kernel_traits::Share_Q_K_smem && Kernel_traits::Is_Q_in_regs);
using GmemTiledCopyO = std::conditional_t<
!Split,
typename Kernel_traits::GmemTiledCopyO,
typename Kernel_traits::GmemTiledCopyOaccum
>;
using ElementO = std::conditional_t<!Split, Element, ElementAccum>;
// The thread index.
const int tidx = threadIdx.x;
const int warp_idx = tidx / 64;
const int lane_idx = tidx % 64;
constexpr int kBlockM = Kernel_traits::kBlockM;
constexpr int kBlockN = Kernel_traits::kBlockN;
constexpr int kNWarps = Kernel_traits::kNWarps;
constexpr int kHeadDimV = Kernel_traits::kHeadDimV;
const BlockInfo</*Varlen=*/!Is_even_MN> binfo(params, bidb);
// if (threadIdx.x == 0 && blockIdx.y == 0 && blockIdx.z == 0) { printf("Is_even_MN = %d, is_cumulativ = %d, seqlen_k_cache = %d, actual_seqlen_k = %d\n", Is_even_MN, params.is_seqlens_k_cumulative, binfo.seqlen_k_cache, binfo.actual_seqlen_k); }
// if (threadIdx.x == 0 && blockIdx.y == 1 && blockIdx.z == 0) { printf("params.knew_ptr = %p, seqlen_k_cache + seqlen_knew = %d\n", params.knew_ptr, binfo.seqlen_k_cache + (params.knew_ptr == nullptr ? 0 : params.seqlen_knew)); }
if (m_block * kBlockM >= binfo.actual_seqlen_q) return;
const int n_blocks_per_split = ((binfo.actual_seqlen_k + kBlockN - 1) / kBlockN + num_n_splits - 1) / num_n_splits;
const int n_block_min = !Is_local
? n_split_idx * n_blocks_per_split
: std::max(n_split_idx * n_blocks_per_split, (m_block * kBlockM + binfo.actual_seqlen_k - binfo.actual_seqlen_q - params.window_size_left) / kBlockN);
int n_block_max = std::min(cute::ceil_div(binfo.actual_seqlen_k, kBlockN), (n_split_idx + 1) * n_blocks_per_split);
if (Is_causal || Is_local) {
n_block_max = std::min(n_block_max,
cute::ceil_div((m_block + 1) * kBlockM + binfo.actual_seqlen_k - binfo.actual_seqlen_q / params.ngroups + params.window_size_right, kBlockN));
}
if (n_block_min >= n_block_max) { // This also covers the case where n_block_max <= 0
// We exit early and write 0 to gOaccum and -inf to gLSEaccum.
// Otherwise we might read OOB elements from gK and gV,
// or get wrong results when we combine gOaccum from different blocks.
const index_t row_offset_o = binfo.q_offset(params.o_batch_stride, params.o_row_stride, bidb)
+ m_block * kBlockM * params.o_row_stride + bidh * params.o_head_stride;
const index_t row_offset_oaccum = (((n_split_idx * params.b + bidb) * params.h + bidh) * params.seqlen_q
+ m_block * kBlockM) * params.d_v;
const index_t row_offset_lseaccum = ((n_split_idx * params.b + bidb) * params.h + bidh) * params.seqlen_q + m_block * kBlockM;
Tensor gOaccum = make_tensor(make_gmem_ptr(reinterpret_cast<ElementO *>(Split ? params.oaccum_ptr : params.o_ptr) + (Split ? row_offset_oaccum : row_offset_o)),
Shape<Int<kBlockM>, Int<kHeadDimV>>{},
make_stride(Split ? kHeadDimV : params.o_row_stride, _1{}));
Tensor gLSEaccum = make_tensor(make_gmem_ptr(reinterpret_cast<ElementAccum *>(Split ? params.softmax_lseaccum_ptr : params.softmax_lse_ptr) + row_offset_lseaccum),
Shape<Int<kBlockM>>{}, Stride<_1>{});
GmemTiledCopyO gmem_tiled_copy_Oaccum;
auto gmem_thr_copy_Oaccum = gmem_tiled_copy_Oaccum.get_thread_slice(tidx);
Tensor tOgOaccum = gmem_thr_copy_Oaccum.partition_D(gOaccum);
Tensor tOrOaccum = make_tensor<ElementO>(shape(tOgOaccum));
clear(tOrOaccum);
// Construct identity layout for sO
Tensor cO = make_identity_tensor(make_shape(size<0>(gOaccum), size<1>(gOaccum))); // (BLK_M,BLK_K) -> (blk_m,blk_k)
// Repeat the partitioning with identity layouts
Tensor tOcO = gmem_thr_copy_Oaccum.partition_D(cO);
Tensor tOpO = make_tensor<bool>(make_shape(size<2>(tOgOaccum)));
if (!Is_even_K) {
#pragma unroll
for (int k = 0; k < size(tOpO); ++k) { tOpO(k) = get<1>(tOcO(0, 0, k)) < params.d_v; }
}
// Clear_OOB_K must be false since we don't want to write zeros to gmem
flash::copy<Is_even_MN, Is_even_K, /*Clear_OOB_MN=*/false, /*Clear_OOB_K=*/false>(
gmem_tiled_copy_Oaccum, tOrOaccum, tOgOaccum, tOcO, tOpO, binfo.actual_seqlen_q - m_block * kBlockM
);
#pragma unroll
for (int m = 0; m < size<1>(tOgOaccum); ++m) {
const int row = get<0>(tOcO(0, m, 0));
if (row < binfo.actual_seqlen_q - m_block * kBlockM && get<1>(tOcO(0, m, 0)) == 0) { gLSEaccum(row) = Split ? -INFINITY : INFINITY; }
}
return;
}
// We iterate over the blocks in reverse order. This is because the last block is the only one
// that needs masking when we read K and V from global memory. Moreover, iterating in reverse
// might save us 1 register (we just need n_block instead of both n_block and n_block_max).
const index_t row_offset_q = binfo.q_offset(params.q_batch_stride, params.q_row_stride, bidb)
+ m_block * kBlockM * params.q_row_stride + bidh * params.q_head_stride;
// We move K and V to the last block.
const int bidb_cache = params.cache_batch_idx == nullptr ? bidb : params.cache_batch_idx[bidb];
const int *block_table = params.block_table == nullptr ? nullptr : params.block_table + bidb * params.block_table_batch_stride;
const int block_table_idx = block_table == nullptr ? 0 : (n_block_max - 1) * kBlockN / params.page_block_size;
const int block_table_offset = block_table == nullptr ? 0 : (n_block_max - 1) * kBlockN - block_table_idx * params.page_block_size;
const index_t row_offset_k = block_table == nullptr
? binfo.k_offset(params.k_batch_stride, params.k_row_stride, bidb_cache)
+ (n_block_max - 1) * kBlockN * params.k_row_stride + (bidh / params.h_h_k_ratio) * params.k_head_stride
: (bidh / params.h_h_k_ratio) * params.k_head_stride;
const index_t row_offset_v = block_table == nullptr
? binfo.k_offset(params.v_batch_stride, params.v_row_stride, bidb_cache)
+ (n_block_max - 1) * kBlockN * params.v_row_stride + (bidh / params.h_h_k_ratio) * params.v_head_stride
: (bidh / params.h_h_k_ratio) * params.v_head_stride;
Tensor gNopeQ = make_tensor(make_gmem_ptr(reinterpret_cast<Element *>(params.q_ptr) + row_offset_q),
Shape<Int<kBlockM>, Int<kHeadDimNope>>{},
make_stride(params.q_row_stride, _1{}));
Tensor gRopeQ = make_tensor(make_gmem_ptr(gNopeQ.data().get() + kHeadDimNope),
Shape<Int<kBlockM>, Int<kHeadDimRope>>{},
make_stride(params.q_row_stride, _1{}));
Tensor gK = make_tensor(make_gmem_ptr(reinterpret_cast<Element *>(params.k_ptr) + row_offset_k),
Shape<Int<kBlockN>, Int<kHeadDim>>{},
make_stride(params.k_row_stride, _1{}));
// if (threadIdx.x == 0 && blockIdx.y == 0 && blockIdx.z == 0) { printf("k_ptr = %p, row_offset_k = %d, gK_ptr = %p\n", params.k_ptr, row_offset_k, gK.data()); }
Tensor gV = make_tensor(make_gmem_ptr(reinterpret_cast<Element *>(params.v_ptr) + row_offset_v),
Shape<Int<kBlockN>, Int<kHeadDimV>>{},
make_stride(params.v_row_stride, _1{}));
Tensor sQ = make_tensor(make_smem_ptr(reinterpret_cast<Element *>(smem_)),
typename Kernel_traits::SmemLayoutQ{});
Tensor sNopeQ = make_tensor(make_smem_ptr(reinterpret_cast<Element *>(smem_)),
typename Kernel_traits::SmemLayoutNopeQ{});
Tensor sRopeQ = make_tensor(make_smem_ptr(reinterpret_cast<Element *>(smem_)),
typename Kernel_traits::SmemLayoutRopeQ{});
Tensor sK = make_tensor(sNopeQ.data(),
typename Kernel_traits::SmemLayoutK{});
Tensor sV = make_tensor(sK.data(), typename Kernel_traits::SmemLayoutVtNoSwizzle{});
Tensor sVt = make_tensor(sV.data(), typename Kernel_traits::SmemLayoutVtransposed{});
Tensor sVtNoSwizzle = make_tensor(sV.data(), typename Kernel_traits::SmemLayoutVtransposedNoSwizzle{});
typename Kernel_traits::GmemTiledCopyB128 gmem_tiled_copy_Q;
auto gmem_thr_copy_Q = gmem_tiled_copy_Q.get_thread_slice(tidx);
Tensor tNopeQgNopeQ = gmem_thr_copy_Q.partition_S(gNopeQ);
Tensor tNopeQsNopeQ = gmem_thr_copy_Q.partition_D(sNopeQ);
Tensor tRopeQgRopeQ = gmem_thr_copy_Q.partition_S(gRopeQ);
Tensor tRopeQsRopeQ = gmem_thr_copy_Q.partition_D(sRopeQ);
typename Kernel_traits::GmemTiledCopyB32 gmem_tiled_copy_KV;
auto gmem_thr_copy_KV = gmem_tiled_copy_KV.get_thread_slice(tidx);
Tensor tKgK = gmem_thr_copy_KV.partition_S(gK); // (KCPY, KCPY_N, KCPY_K)
Tensor tKsK = gmem_thr_copy_KV.partition_D(sK);
// gemm S is 4x1 wave layout, wave(n) compute the same S with wave(n + 4)
int tidx_mma_s = tidx & 0xFF;
typename Kernel_traits::TiledMmaS tiled_mma_s;
auto thr_mma_s = tiled_mma_s.get_thread_slice(tidx_mma_s);
Tensor tSrQ = thr_mma_s.partition_fragment_A(sQ); // (MMA,MMA_M,MMA_K)
Tensor tSrNopeQ = thr_mma_s.partition_fragment_A(sNopeQ); // (MMA,MMA_M,MMA_K)
Tensor tSrRopeQ = thr_mma_s.partition_fragment_A(sRopeQ); // (MMA,MMA_M,MMA_K)
Tensor tSrK = thr_mma_s.partition_fragment_B(sK(_, _, 0)); // (MMA,MMA_N,MMA_K)
typename Kernel_traits::TiledMmaO tiled_mma_o;
auto thr_mma_o = tiled_mma_o.get_thread_slice(tidx);
// Tensor tOrVt = thr_mma_o.partition_fragment_B(sVt); // (MMA, MMA_K,MMA_N)
Tensor tOrVt = make_tensor<Element>(Shape<_4, Shape<_4, _4>, _1>{});
Tensor acc_o = partition_fragment_C(tiled_mma_o, Shape<Int<kBlockM>, Int<kHeadDimV>>{}); // MMA, MMA_M, MMA_K
//
// Copy Atom retiling
//
auto smem_tiled_copy_Q = make_tiled_copy_A(typename Kernel_traits::SmemCopyAtomB64{}, tiled_mma_s);
auto smem_thr_copy_Q = smem_tiled_copy_Q.get_thread_slice(tidx_mma_s);
Tensor tSsNopeQ = smem_thr_copy_Q.partition_S(sNopeQ);
Tensor tSsRopeQ = smem_thr_copy_Q.partition_S(sRopeQ);
auto smem_tiled_copy_K = make_tiled_copy_B(typename Kernel_traits::SmemCopyAtomB64{}, tiled_mma_s);
auto smem_thr_copy_K = smem_tiled_copy_K.get_thread_slice(tidx_mma_s);
Tensor tSsK = smem_thr_copy_K.partition_S(sK);
auto smem_tiled_copy_V = make_tiled_copy_B(typename Kernel_traits::SmemCopyAtomTransposed{}, tiled_mma_o);
auto smem_thr_copy_V = smem_tiled_copy_V.get_thread_slice(tidx);
int warp_offset = warp_idx / kAtomLayoutMO * 16 * 64;
int thread_offset = lane_idx / 16 * 4 * 64;
Element *Vtsmem_ptr_lds = reinterpret_cast<Element *>(sVt.data().get()) + warp_offset + thread_offset;
Tensor tOsVt = make_tensor(make_smem_ptr(Vtsmem_ptr_lds), make_layout(Shape<_4, _4, Int<Num_Stages>>{}, // MMA MMA_N NUM_STAGES
Stride<_1, Int<16*128>, Int<kBlockN*kHeadDim>>{}));
// PREDICATES
// Construct identity layout for sQ and sK
Tensor cQ = make_identity_tensor(make_shape(size<0>(sQ), size<1>(sQ))); // (BLK_M,BLK_K) -> (blk_m,blk_k)
Tensor cKV = make_identity_tensor(make_shape(size<0>(sK), size<1>(sK))); // (BLK_N,BLK_K) -> (blk_n,blk_k)
// Repeat the partitioning with identity layouts
Tensor tQcQ = gmem_thr_copy_Q.partition_S(cQ); // (ACPY,ACPY_M,ACPY_K) -> (blk_m,blk_k)
Tensor tKVcKV = gmem_thr_copy_KV.partition_S(cKV); // (BCPY,BCPY_N,BCPY_K) -> (blk_n,blk_k)
// Prologue
// Read Q from gmem to smem, optionally apply rotary embedding.
Tensor tNopeQrNopeQ = make_fragment_like(tNopeQgNopeQ);
Tensor tRopeQrRopeQ = make_fragment_like(tRopeQgRopeQ);
// We don't need to clear the sQ smem tiles since we'll only write out the valid outputs
flash::copy_b128<Is_even_MN, Is_even_K>(tNopeQgNopeQ, tNopeQrNopeQ, tQcQ, params.d, binfo.actual_seqlen_q - m_block * kBlockM);
flash::copy_b128<Is_even_MN, Is_even_K>(tRopeQgRopeQ, tRopeQrRopeQ, tQcQ, params.d, binfo.actual_seqlen_q - m_block * kBlockM);
cute::copy(tNopeQrNopeQ, tNopeQsNopeQ);
flash::sync_threads();
cute::copy(smem_tiled_copy_Q, tSsNopeQ, tSrNopeQ);
flash::sync_threads();
cute::copy(tRopeQrRopeQ, tRopeQsRopeQ);
flash::sync_threads();
cute::copy(smem_tiled_copy_Q, tSsRopeQ, tSrRopeQ);
flash::sync_threads();
flash::concat(tSrNopeQ, tSrRopeQ, tSrQ);
int n_block = n_block_max - 1;
int Ksmem_read_index = 0;
int Ksmem_write_index = 0;
// We don't need to clear the sK smem tiles since we'll mask out the scores anyway.
Tensor tKrK = make_fragment_like(tKgK);
flash::copy_b32_page_one<Kernel_traits, Is_even_MN, Is_even_K>(gK, tKgK, tKrK, tKVcKV, params.d, n_block,
block_table, params.k_batch_stride, params.k_row_stride, params.page_block_size, binfo.actual_seqlen_k - n_block * kBlockN);
// flash::cp_async_wait<0>();
// __syncthreads();
// if (tidx == 0 && blockIdx.y == 0 && blockIdx.z == 0) { print(tKsK); }
// __syncthreads();
clear(acc_o);
flash::Softmax<size<1>(acc_o)> softmax;
const float alibi_slope = !Has_alibi ? 0.0f : reinterpret_cast<float *>(params.alibi_slopes_ptr)[bidb * params.alibi_slopes_batch_stride + bidh] / params.scale_softmax;
flash::Mask<Is_causal, Is_local, Has_alibi> mask(binfo.actual_seqlen_k, binfo.actual_seqlen_q, params.ngroups, params.window_size_left, params.window_size_right, alibi_slope);
for (; n_block >= n_block_min; --n_block) {
Tensor acc_s = partition_fragment_C(tiled_mma_s, Shape<Int<kBlockM>, Int<kBlockN>>{}); // (MMA=4, MMA_M, MMA_N)
cute::copy(tKrK, tKsK(_, _, _, Ksmem_write_index));
Ksmem_write_index ^= 1;
clear(acc_s);
flash::sync_threads();
if (n_block > n_block_min) {
// Advance gK
flash::copy_b32_page_one<Kernel_traits, /*Is_even_MN=*/true, Is_even_K>(gK, tKgK, tKrK, tKVcKV, params.d, n_block - 1,
block_table, params.k_batch_stride, params.k_row_stride, params.page_block_size);
}
flash::gemm_opt</*A_in_regs=*/Kernel_traits::Is_Q_in_regs>(
acc_s, tSrQ, tSrK, tSsNopeQ, tSsK(_, _, _, Ksmem_read_index), tiled_mma_s, smem_tiled_copy_Q, smem_tiled_copy_K,
smem_thr_copy_Q, smem_thr_copy_K
);
if constexpr (Is_softcap){
flash::apply_softcap(acc_s, params.softcap);
}
mask.template apply_mask<Is_causal, Is_even_MN>(
acc_s, n_block * kBlockN, m_block * kBlockM + (tidx / 64) % kAtomLayoutMS * 16 + (tidx & 0xf), kAtomLayoutMS * 16
);
softmax.template softmax_rescale_o</*Is_first=*/false, /*Check_inf=*/Is_causal || Is_local || !Is_even_MN, true, true>(acc_s, acc_o, params.scale_softmax_log2);
//Tensor rP = flash::convert_type<Element>(acc_s);
CONVERT_TENSOR_TYPE(ElementAccum, Element, acc_s, rP)
// Reshape rP from (MMA=4, MMA_M, MMA_N) to ((4, 2), MMA_M, MMA_N / 2)
// if using m16n8k16 or (4, MMA_M, MMA_N) if using m16n8k8.
//Tensor tOrP = make_tensor(rP.data(), flash::convert_layout_acc_Aregs<Kernel_traits::TiledMma>(rP.layout()));
lds4x4_with_swizzle424(tOsVt(_, _, Ksmem_read_index), tOrVt);
CUTE_STATIC_ASSERT_V(size<2>(tOrVt) == _1{}); // only support MMA_K = 1
Tensor tOrVt_permute_view = make_tensor(tOrVt.data(), make_layout(make_shape(size<0>(tOrVt), size<1, 0>(tOrVt), size<1, 1>(tOrVt))));
permute_4x4_b16(tOrVt_permute_view);
Tensor tOrP = make_tensor(rP.data(), acc_s.layout());
flash::gemm_rr(acc_o, tOrP, tOrVt, tiled_mma_o);
Ksmem_read_index ^= 1;
}
// Epilogue
Tensor lse = softmax.template normalize_softmax_lse</*Is_dropout=*/false, /*Return_lse*/true, Split>(acc_o, params.scale_softmax);
Tensor acc_o_view = make_tensor(acc_o.data(), make_layout(Shape<_4, Shape<_4, _4>>{},
Stride<_1, Shape<_4, _16>>{}));
@ -317,7 +55,6 @@ __forceinline__ __device__ void compute_attn_1rowblock_splitkv_k64_mla_64x16_8wa
acc_o_copy(row, make_coord(col, k)) = acc_o_view(col, make_coord(row, k));
}
}
// if (cute::thread0()) { print(lse); }
if constexpr (!Split) {
Tensor sOaccum = make_tensor(make_smem_ptr(reinterpret_cast<ElementO *>(smem_)), typename Kernel_traits::SmemLayoutO{}); // (SMEM_M,SMEM_N)
// Partition sO to match the accumulator partitioning
@ -332,8 +69,6 @@ __forceinline__ __device__ void compute_attn_1rowblock_splitkv_k64_mla_64x16_8wa
Stride<_1, _16>{}));
// sOaccum is larger than sQ, so we need to syncthreads here
// TODO: allocate enough smem for sOaccum
if constexpr (Kernel_traits::Share_Q_K_smem) { flash::sync_threads(); }
cute::copy(tOrO, tOsO);
@ -382,10 +117,11 @@ __forceinline__ __device__ void compute_attn_1rowblock_splitkv_k64_mla_64x16_8wa
tOrOaccum, tOgOaccum, tOcO, params.d_v, binfo.actual_seqlen_q - m_block * kBlockM
);
} else {
const int split_offset = params.num_splits_ptr[bidb];
Tensor sOaccum = make_tensor(make_smem_ptr(reinterpret_cast<ElementO *>(smem_)), typename Kernel_traits::SmemLayoutO{}); // (SMEM_M,SMEM_N)
const index_t row_offset_oaccum = (((n_split_idx * params.b + bidb) * params.h + bidh) * params.seqlen_q
const index_t row_offset_oaccum = (((split_offset + n_split_idx) * params.h + bidh) * params.seqlen_q
+ m_block * kBlockM) * params.d_v;
const index_t row_offset_lseaccum = ((n_split_idx * params.b + bidb) * params.h + bidh) * params.seqlen_q + m_block * kBlockM;
const index_t row_offset_lseaccum = ((split_offset + n_split_idx) * params.h + bidh) * params.seqlen_q + m_block * kBlockM;
Tensor gOaccum = make_tensor(make_gmem_ptr(reinterpret_cast<ElementO *>(params.oaccum_ptr) + row_offset_oaccum),
Shape<Int<kBlockM>, Int<kHeadDimV/2>>{},
@ -471,4 +207,237 @@ __forceinline__ __device__ void compute_attn_1rowblock_splitkv_k64_mla_64x16_8wa
}
}
template<typename Kernel_traits, bool Is_causal, bool Is_even_MN, bool Is_even_K, bool Is_enable_dcp, typename Params>
__forceinline__ __device__ void compute_attn_1rowblock_splitkv_mla_k64_64x16_8waves_xcore1000(
const Params &params, const int bidb, const int bidh, const int m_block, const int n_split_idx, const int n_block_min, const int n_block_max, const bool NoSplit) {
using Element = typename Kernel_traits::Element;
using ElementAccum = typename Kernel_traits::ElementAccum;
using index_t = typename Kernel_traits::index_t;
// Shared memory.
extern __shared__ char smem_[];
// The thread index.
const int tidx = threadIdx.x;
const int warp_idx = tidx / 64;
const int lane_idx = tidx % 64;
constexpr int kBlockM = Kernel_traits::kBlockM;
constexpr int kBlockN = Kernel_traits::kBlockN;
constexpr int kHeadDim = Kernel_traits::kHeadDim;
constexpr int kHeadDimV = Kernel_traits::kHeadDimV;
constexpr int kNWarps = Kernel_traits::kNWarps;
constexpr int kBlockKSmem = Kernel_traits::kBlockKSmem;
constexpr int kAtomLayoutMS = Kernel_traits::kAtomLayoutMS;
constexpr int kAtomLayoutMO = Kernel_traits::kAtomLayoutMO;
constexpr int Num_Stages = Kernel_traits::Num_Stages;
constexpr int kHeadDimNope = kHeadDimV;
constexpr int kHeadDimRope = kHeadDim - kHeadDimV;
static_assert(kBlockKSmem == 64);
static_assert(Kernel_traits::Share_Q_K_smem && Kernel_traits::Is_Q_in_regs);
const BlockInfo</*Varlen=*/!Is_even_MN> binfo(params, bidb);
if (m_block * kBlockM >= binfo.actual_seqlen_q) return;
// We iterate over the blocks in reverse order. This is because the last block is the only one
// that needs masking when we read K and V from global memory. Moreover, iterating in reverse
// might save us 1 register (we just need n_block instead of both n_block and n_block_max).
const index_t row_offset_q = binfo.q_offset(params.q_batch_stride, params.q_row_stride, bidb)
+ m_block * kBlockM * params.q_row_stride + bidh * params.q_head_stride;
// We move K and V to the last block.
const int bidb_cache = params.cache_batch_idx == nullptr ? bidb : params.cache_batch_idx[bidb];
const int *block_table = params.block_table == nullptr ? nullptr : params.block_table + bidb * params.block_table_batch_stride;
const int block_table_idx = block_table == nullptr ? 0 : (n_block_max - 1) * kBlockN / params.page_block_size;
const int block_table_offset = block_table == nullptr ? 0 : (n_block_max - 1) * kBlockN - block_table_idx * params.page_block_size;
const index_t row_offset_k = block_table == nullptr
? binfo.k_offset(params.k_batch_stride, params.k_row_stride, bidb_cache)
+ (n_block_max - 1) * kBlockN * params.k_row_stride + (bidh / params.h_h_k_ratio) * params.k_head_stride
: (bidh / params.h_h_k_ratio) * params.k_head_stride;
const index_t row_offset_v = block_table == nullptr
? binfo.k_offset(params.v_batch_stride, params.v_row_stride, bidb_cache)
+ (n_block_max - 1) * kBlockN * params.v_row_stride + (bidh / params.h_h_k_ratio) * params.v_head_stride
: (bidh / params.h_h_k_ratio) * params.v_head_stride;
Tensor gNopeQ = make_tensor(make_gmem_ptr(reinterpret_cast<Element *>(params.q_ptr) + row_offset_q),
Shape<Int<kBlockM>, Int<kHeadDimNope>>{},
make_stride(params.q_row_stride, _1{}));
Tensor gRopeQ = make_tensor(make_gmem_ptr(gNopeQ.data().get() + kHeadDimNope),
Shape<Int<kBlockM>, Int<kHeadDimRope>>{},
make_stride(params.q_row_stride, _1{}));
Tensor gK = make_tensor(make_gmem_ptr(reinterpret_cast<Element *>(params.k_ptr) + row_offset_k),
Shape<Int<kBlockN>, Int<kHeadDim>>{},
make_stride(params.k_row_stride, _1{}));
// if (threadIdx.x == 0 && blockIdx.y == 0 && blockIdx.z == 0) { printf("k_ptr = %p, row_offset_k = %d, gK_ptr = %p\n", params.k_ptr, row_offset_k, gK.data()); }
Tensor gV = make_tensor(make_gmem_ptr(reinterpret_cast<Element *>(params.v_ptr) + row_offset_v),
Shape<Int<kBlockN>, Int<kHeadDimV>>{},
make_stride(params.v_row_stride, _1{}));
Tensor sQ = make_tensor(make_smem_ptr(reinterpret_cast<Element *>(smem_)),
typename Kernel_traits::SmemLayoutQ{});
Tensor sNopeQ = make_tensor(make_smem_ptr(reinterpret_cast<Element *>(smem_)),
typename Kernel_traits::SmemLayoutNopeQ{});
Tensor sRopeQ = make_tensor(make_smem_ptr(reinterpret_cast<Element *>(smem_)),
typename Kernel_traits::SmemLayoutRopeQ{});
Tensor sK = make_tensor(sNopeQ.data(), typename Kernel_traits::SmemLayoutK424{});
Tensor sV = make_tensor(sK.data(), typename Kernel_traits::SmemLayoutVtNoSwizzle{});
Tensor sVt = make_tensor(sV.data(), typename Kernel_traits::SmemLayoutVtransposed424{});
Tensor sVtNoSwizzle = make_tensor(sV.data(), typename Kernel_traits::SmemLayoutVtransposedNoSwizzle{});
typename Kernel_traits::GmemTiledCopyB128 gmem_tiled_copy_Q;
auto gmem_thr_copy_Q = gmem_tiled_copy_Q.get_thread_slice(tidx);
Tensor tNopeQgNopeQ = gmem_thr_copy_Q.partition_S(gNopeQ);
Tensor tNopeQsNopeQ = gmem_thr_copy_Q.partition_D(sNopeQ);
Tensor tRopeQgRopeQ = gmem_thr_copy_Q.partition_S(gRopeQ);
Tensor tRopeQsRopeQ = gmem_thr_copy_Q.partition_D(sRopeQ);
typename Kernel_traits::GmemTiledCopyB32 gmem_tiled_copy_KV;
auto gmem_thr_copy_KV = gmem_tiled_copy_KV.get_thread_slice(tidx);
Tensor tKgK = gmem_thr_copy_KV.partition_S(gK); // (KCPY, KCPY_N, KCPY_K)
Tensor tKsK = gmem_thr_copy_KV.partition_D(sK);
// gemm S is 4x1 wave layout, wave(n) compute the same S with wave(n + 4)
int tidx_mma_s = tidx & 0xFF;
typename Kernel_traits::TiledMmaS tiled_mma_s;
auto thr_mma_s = tiled_mma_s.get_thread_slice(tidx_mma_s);
Tensor tSrQ = thr_mma_s.partition_fragment_A(sQ); // (MMA,MMA_M,MMA_K)
Tensor tSrNopeQ = thr_mma_s.partition_fragment_A(sNopeQ); // (MMA,MMA_M,MMA_K)
Tensor tSrRopeQ = thr_mma_s.partition_fragment_A(sRopeQ); // (MMA,MMA_M,MMA_K)
Tensor tSrK = thr_mma_s.partition_fragment_B(sK(_, _, 0)); // (MMA,MMA_N,MMA_K)
typename Kernel_traits::TiledMmaO tiled_mma_o;
auto thr_mma_o = tiled_mma_o.get_thread_slice(tidx);
Tensor tOrVt = make_tensor<Element>(Shape<_4, Shape<_4, _4>, _1>{});
Tensor acc_o = partition_fragment_C(tiled_mma_o, Shape<Int<kBlockM>, Int<kHeadDimV>>{}); // MMA, MMA_M, MMA_K
//
// Copy Atom retiling
//
auto smem_tiled_copy_Q = make_tiled_copy_A(typename Kernel_traits::UniversalCopyAtomB64{}, tiled_mma_s);
auto smem_thr_copy_Q = smem_tiled_copy_Q.get_thread_slice(tidx_mma_s);
Tensor tSsNopeQ = smem_thr_copy_Q.partition_S(sNopeQ);
Tensor tSsRopeQ = smem_thr_copy_Q.partition_S(sRopeQ);
auto smem_tiled_copy_K = make_tiled_copy_B(typename Kernel_traits::UniversalCopyAtomB64{}, tiled_mma_s);
auto smem_thr_copy_K = smem_tiled_copy_K.get_thread_slice(tidx_mma_s);
Tensor tSsK = smem_thr_copy_K.partition_S(sK);
auto smem_tiled_copy_V = make_tiled_copy_B(typename Kernel_traits::SmemCopyAtomTransposed{}, tiled_mma_o);
auto smem_thr_copy_V = smem_tiled_copy_V.get_thread_slice(tidx);
int warp_offset = warp_idx / kAtomLayoutMO * 16 * 64;
int thread_offset = lane_idx / 16 * 4 * 64;
Element *Vtsmem_ptr_lds = reinterpret_cast<Element *>(sVt.data().get()) + warp_offset + thread_offset;
Tensor tOsVt = make_tensor(make_smem_ptr(Vtsmem_ptr_lds), make_layout(Shape<_4, _4, Int<Num_Stages>>{}, // MMA MMA_N NUM_STAGES
Stride<_1, Int<16*128>, Int<kBlockN*kHeadDim>>{}));
// PREDICATES
// Construct identity layout for sQ and sK
Tensor cQ = make_identity_tensor(make_shape(size<0>(sQ), size<1>(sQ))); // (BLK_M,BLK_K) -> (blk_m,blk_k)
Tensor cKV = make_identity_tensor(make_shape(size<0>(sK), size<1>(sK))); // (BLK_N,BLK_K) -> (blk_n,blk_k)
// Repeat the partitioning with identity layouts
Tensor tQcQ = gmem_thr_copy_Q.partition_S(cQ); // (ACPY,ACPY_M,ACPY_K) -> (blk_m,blk_k)
Tensor tKVcKV = gmem_thr_copy_KV.partition_S(cKV); // (BCPY,BCPY_N,BCPY_K) -> (blk_n,blk_k)
// Prologue
// Read Q from gmem to smem, optionally apply rotary embedding.
Tensor tNopeQrNopeQ = make_fragment_like(tNopeQgNopeQ);
Tensor tRopeQrRopeQ = make_fragment_like(tRopeQgRopeQ);
// We don't need to clear the sQ smem tiles since we'll only write out the valid outputs
flash::copy_b128<Is_even_MN, Is_even_K>(tNopeQgNopeQ, tNopeQrNopeQ, tQcQ, params.d, binfo.actual_seqlen_q - m_block * kBlockM);
flash::copy_b128<Is_even_MN, Is_even_K>(tRopeQgRopeQ, tRopeQrRopeQ, tQcQ, params.d, binfo.actual_seqlen_q - m_block * kBlockM);
cute::copy(tNopeQrNopeQ, tNopeQsNopeQ);
flash::sync_threads();
cute::copy(smem_tiled_copy_Q, tSsNopeQ, tSrNopeQ);
flash::sync_threads();
cute::copy(tRopeQrRopeQ, tRopeQsRopeQ);
flash::sync_threads();
cute::copy(smem_tiled_copy_Q, tSsRopeQ, tSrRopeQ);
flash::sync_threads();
flash::concat(tSrNopeQ, tSrRopeQ, tSrQ);
int n_block = n_block_max - 1;
int Ksmem_read_index = 0;
int Ksmem_write_index = 0;
// We don't need to clear the sK smem tiles since we'll mask out the scores anyway.
Tensor tKrK = make_fragment_like(tKgK);
int row_offset = tidx / 32 + n_block * kBlockN;
int virtual_page_idx = row_offset / params.page_block_size;
int page_offset = row_offset - virtual_page_idx * params.page_block_size;
int32_t page_idx = block_table[virtual_page_idx];
flash::copy_b32_page_one<Kernel_traits, Is_even_MN, Is_even_K>(gK, tKgK, tKrK, tKVcKV, params.d, n_block,
block_table, params.k_batch_stride, params.k_row_stride, params.page_block_size, page_idx, page_offset, binfo.actual_seqlen_k - n_block * kBlockN);
// flash::cp_async_wait<0>();
// __syncthreads();
// if (tidx == 0 && blockIdx.y == 0 && blockIdx.z == 0) { print(tKsK); }
// __syncthreads();
clear(acc_o);
flash::Softmax<size<1>(acc_o)> softmax;
flash::Mask<Is_causal, Is_enable_dcp> mask(binfo.actual_seqlen_k, binfo.actual_seqlen_q, params.ngroups, binfo.tot_seqlen_k, params.cp_world_size, params.cp_rank);
for (; n_block >= n_block_min; --n_block) {
if (n_block > n_block_min) {
// prefetch load page index
row_offset -= kBlockN;
virtual_page_idx = row_offset / params.page_block_size;
page_offset = row_offset - virtual_page_idx * params.page_block_size;
page_idx = block_table[virtual_page_idx];
}
Tensor acc_s = partition_fragment_C(tiled_mma_s, Shape<Int<kBlockM>, Int<kBlockN>>{}); // (MMA=4, MMA_M, MMA_N)
cute::copy(tKrK, tKsK(_, _, _, Ksmem_write_index));
Ksmem_write_index ^= 1;
clear(acc_s);
flash::sync_threads();
if (n_block > n_block_min) {
// Advance gK
flash::copy_b32_page_one<Kernel_traits, /*Is_even_MN=*/true, Is_even_K>(gK, tKgK, tKrK, tKVcKV, params.d, n_block - 1,
block_table, params.k_batch_stride, params.k_row_stride, params.page_block_size, page_idx, page_offset);
}
flash::gemm</*A_in_regs=*/Kernel_traits::Is_Q_in_regs>(
acc_s, tSrQ, tSrK, tSsNopeQ, tSsK(_, _, _, Ksmem_read_index), tiled_mma_s, smem_tiled_copy_Q, smem_tiled_copy_K,
smem_thr_copy_Q, smem_thr_copy_K
);
mask.template apply_mask<Is_causal, Is_even_MN>(
acc_s, n_block * kBlockN, m_block * kBlockM + (tidx / 64) % kAtomLayoutMS * 16 + (tidx & 0xf), kAtomLayoutMS * 16
);
softmax.template softmax_rescale_o</*Is_first=*/false, /*Check_inf=*/Is_causal || !Is_even_MN, false, false>(acc_s, acc_o, params.scale_softmax_log2);
lds4x4_with_swizzle424(tOsVt(_, _, Ksmem_read_index), tOrVt);
CONVERT_TENSOR_TYPE(ElementAccum, Element, acc_s, rP)
CUTE_STATIC_ASSERT_V(size<2>(tOrVt) == _1{}); // only support MMA_K = 1
Tensor tOrVt_permute_view = make_tensor(tOrVt.data(), make_layout(make_shape(size<0>(tOrVt), size<1, 0>(tOrVt), size<1, 1>(tOrVt))));
permute_4x4_b16(tOrVt_permute_view);
Tensor tOrP = make_tensor(rP.data(), acc_s.layout());
flash::gemm_rr(acc_o, tOrP, tOrVt, tiled_mma_o);
Ksmem_read_index ^= 1;
}
// Epilogue
if (NoSplit) {
store_64x16_xcore1000<Kernel_traits, false, Is_even_MN, Is_even_K>(params, bidb, bidh, m_block, n_split_idx, smem_, acc_o, softmax);
}else{
store_64x16_xcore1000<Kernel_traits, true, Is_even_MN, Is_even_K>(params, bidb, bidh, m_block, n_split_idx, smem_, acc_o, softmax);
}
}
} // namespace flash

View File

@ -0,0 +1,275 @@
// Adapted from Dao-AILab/flash-attention (https://github.com/Dao-AILab/flash-attention/tree/v2.6.3)
#pragma once
#include <cute/algorithm/copy.hpp>
#include <mctlass/mctlass.h>
#include <mctlass/array.h>
#include <mctlass/numeric_types.h>
#include "block_info.h"
#include "kernel_traits.h"
#include "utils.h"
#include "softmax.h"
#include "mask.h"
#include "flash_fwd_mla_kernel_k64_64x16_8waves_xcore1000.h"
namespace flash {
using namespace cute;
template<typename Kernel_traits, bool Is_causal, bool Is_even_TopK, typename Params>
__forceinline__ __device__ void compute_attn_1rowblock_splitkv_sparse_mla_k64_64x16_8waves_xcore1000(
const Params &params, const int bidb, const int bidh, const int m_block, const int n_split_idx, const int n_block_min, const int n_block_max, const bool NoSplit) {
using Element = typename Kernel_traits::Element;
using ElementAccum = typename Kernel_traits::ElementAccum;
using index_t = typename Kernel_traits::index_t;
// Shared memory.
extern __shared__ char smem_[];
// The thread index.
const int tidx = threadIdx.x;
const int warp_idx = tidx / 64;
const int lane_idx = tidx % 64;
constexpr int kBlockM = Kernel_traits::kBlockM;
constexpr int kBlockN = Kernel_traits::kBlockN;
constexpr int kHeadDim = Kernel_traits::kHeadDim;
constexpr int kHeadDimV = Kernel_traits::kHeadDimV;
constexpr int kNWarps = Kernel_traits::kNWarps;
constexpr int kBlockKSmem = Kernel_traits::kBlockKSmem;
constexpr int kAtomLayoutMS = Kernel_traits::kAtomLayoutMS;
constexpr int kAtomLayoutMO = Kernel_traits::kAtomLayoutMO;
constexpr int Num_Stages = Kernel_traits::Num_Stages;
constexpr int kBlockTopK = Kernel_traits::kBlockTopK;
constexpr int kHeadDimNope = kHeadDimV;
constexpr int kHeadDimRope = kHeadDim - kHeadDimV;
static_assert(kBlockKSmem == 64);
static_assert(Kernel_traits::Share_Q_K_smem && Kernel_traits::Is_Q_in_regs);
const BlockInfo<true> binfo(params, bidb);
if (m_block * kBlockM >= binfo.actual_seqlen_q) return;
// We iterate over the blocks in reverse order. This is because the last block is the only one
// that needs masking when we read K and V from global memory. Moreover, iterating in reverse
// might save us 1 register (we just need n_block instead of both n_block and n_block_max).
const index_t row_offset_q = binfo.q_offset(params.q_batch_stride, params.q_row_stride, bidb)
+ m_block * kBlockM * params.q_row_stride + bidh * params.q_head_stride;
// We move K and V to the last block.
const int bidb_cache = params.cache_batch_idx == nullptr ? bidb : params.cache_batch_idx[bidb];
const int s_q_idx = params.ngroups >= kBlockM ? m_block / (params.ngroups / kBlockM) : 0; //s_q_idx, ngroups = head_q_ori
const index_t offset_indices = bidb * params.indices_batch_stride + s_q_idx * params.indices_row_stride;
const int* gIndices = params.indices_ptr + offset_indices; // top_k values
const bool indices_all_valid_per_q = params.indices_all_valid_per_q_ptr[bidb * params.indices_all_valid_per_q_batch_stride + s_q_idx];
// will calculate row_offset_k later
const index_t row_offset_k = (bidh / params.h_h_k_ratio) * params.k_head_stride;
const index_t row_offset_v = (bidh / params.h_h_k_ratio) * params.v_head_stride;
Tensor gNopeQ = make_tensor(make_gmem_ptr(reinterpret_cast<Element *>(params.q_ptr) + row_offset_q),
Shape<Int<kBlockM>, Int<kHeadDimNope>>{},
make_stride(params.q_row_stride, _1{}));
Tensor gRopeQ = make_tensor(make_gmem_ptr(gNopeQ.data().get() + kHeadDimNope),
Shape<Int<kBlockM>, Int<kHeadDimRope>>{},
make_stride(params.q_row_stride, _1{}));
Tensor gK = make_tensor(make_gmem_ptr(reinterpret_cast<Element *>(params.k_ptr) + row_offset_k),
Shape<Int<kBlockN>, Int<kHeadDim>>{},
make_stride(params.k_row_stride, _1{}));
Tensor gV = make_tensor(make_gmem_ptr(reinterpret_cast<Element *>(params.v_ptr) + row_offset_v),
Shape<Int<kBlockN>, Int<kHeadDimV>>{},
make_stride(params.v_row_stride, _1{}));
Tensor sQ = make_tensor(make_smem_ptr(reinterpret_cast<Element *>(smem_)),
typename Kernel_traits::SmemLayoutQ{});
Tensor sNopeQ = make_tensor(make_smem_ptr(reinterpret_cast<Element *>(smem_)),
typename Kernel_traits::SmemLayoutNopeQ{});
Tensor sRopeQ = make_tensor(make_smem_ptr(reinterpret_cast<Element *>(smem_)),
typename Kernel_traits::SmemLayoutRopeQ{});
Tensor sK = make_tensor(sNopeQ.data(), typename Kernel_traits::SmemLayoutK424{}); // kBlockN * kheadDim * NumStages
Tensor sV = make_tensor(sK.data(), typename Kernel_traits::SmemLayoutVtNoSwizzle{}); // kBlockN * kHeadDimV
Tensor sVt = make_tensor(sV.data(), typename Kernel_traits::SmemLayoutVtransposed424{}); // kBlockN * kheadDimV * NumStages
Tensor sVtNoSwizzle = make_tensor(sV.data(), typename Kernel_traits::SmemLayoutVtransposedNoSwizzle{}); // kBlockN * kheadDimV * NumStages
int32_t* indices_smem_ptr = reinterpret_cast<int32_t *>(reinterpret_cast<char *>(smem_) + (size(sK) * sizeof(Element) + 3) / 4 * 4);
typename Kernel_traits::GmemTiledCopyB128 gmem_tiled_copy_Q;
auto gmem_thr_copy_Q = gmem_tiled_copy_Q.get_thread_slice(tidx);
Tensor tNopeQgNopeQ = gmem_thr_copy_Q.partition_S(gNopeQ);
Tensor tNopeQsNopeQ = gmem_thr_copy_Q.partition_D(sNopeQ);
Tensor tRopeQgRopeQ = gmem_thr_copy_Q.partition_S(gRopeQ);
Tensor tRopeQsRopeQ = gmem_thr_copy_Q.partition_D(sRopeQ);
typename Kernel_traits::GmemTiledCopyB32 gmem_tiled_copy_KV;
auto gmem_thr_copy_KV = gmem_tiled_copy_KV.get_thread_slice(tidx);
Tensor tKgK = gmem_thr_copy_KV.partition_S(gK); // (KCPY, KCPY_N, KCPY_K)
Tensor tKsK = gmem_thr_copy_KV.partition_D(sK);
// gemm S is 4x1 wave layout, wave(n) compute the same S with wave(n + 4)
int tidx_mma_s = tidx & 0xFF;
typename Kernel_traits::TiledMmaS tiled_mma_s;
auto thr_mma_s = tiled_mma_s.get_thread_slice(tidx_mma_s);
Tensor tSrQ = thr_mma_s.partition_fragment_A(sQ); // (MMA,MMA_M,MMA_K)
Tensor tSrNopeQ = thr_mma_s.partition_fragment_A(sNopeQ); // (MMA,MMA_M,MMA_K)
Tensor tSrRopeQ = thr_mma_s.partition_fragment_A(sRopeQ); // (MMA,MMA_M,MMA_K)
Tensor tSrK = thr_mma_s.partition_fragment_B(sK(_, _, 0)); // (MMA,MMA_N,MMA_K)
typename Kernel_traits::TiledMmaO tiled_mma_o;
auto thr_mma_o = tiled_mma_o.get_thread_slice(tidx);
Tensor tOrVt = make_tensor<Element>(Shape<_4, Shape<_4, _4>, _1>{});
Tensor acc_o = partition_fragment_C(tiled_mma_o, Shape<Int<kBlockM>, Int<kHeadDimV>>{}); // MMA, MMA_M, MMA_K
//
// Copy Atom retiling
//
auto smem_tiled_copy_Q = make_tiled_copy_A(typename Kernel_traits::UniversalCopyAtomB64{}, tiled_mma_s);
auto smem_thr_copy_Q = smem_tiled_copy_Q.get_thread_slice(tidx_mma_s);
Tensor tSsNopeQ = smem_thr_copy_Q.partition_S(sNopeQ);
Tensor tSsRopeQ = smem_thr_copy_Q.partition_S(sRopeQ);
auto smem_tiled_copy_K = make_tiled_copy_B(typename Kernel_traits::UniversalCopyAtomB64{}, tiled_mma_s);
auto smem_thr_copy_K = smem_tiled_copy_K.get_thread_slice(tidx_mma_s);
Tensor tSsK = smem_thr_copy_K.partition_S(sK);
auto smem_tiled_copy_V = make_tiled_copy_B(typename Kernel_traits::SmemCopyAtomTransposed{}, tiled_mma_o);
auto smem_thr_copy_V = smem_tiled_copy_V.get_thread_slice(tidx);
int warp_offset = warp_idx / kAtomLayoutMO * 16 * 64;
int thread_offset = lane_idx / 16 * 4 * 64;
Element *Vtsmem_ptr_lds = reinterpret_cast<Element *>(sVt.data().get()) + warp_offset + thread_offset;
Tensor tOsVt = make_tensor(make_smem_ptr(Vtsmem_ptr_lds), make_layout(Shape<_4, _4, Int<Num_Stages>>{}, // MMA MMA_N NUM_STAGES
Stride<_1, Int<16*128>, Int<kBlockN*kHeadDim>>{}));
// PREDICATES
// Construct identity layout for sQ and sK
Tensor cQ = make_identity_tensor(make_shape(size<0>(sQ), size<1>(sQ))); // (BLK_M,BLK_K) -> (blk_m,blk_k)
Tensor cKV = make_identity_tensor(make_shape(size<0>(sK), size<1>(sK))); // (BLK_N,BLK_K) -> (blk_n,blk_k)
// Repeat the partitioning with identity layouts
Tensor tQcQ = gmem_thr_copy_Q.partition_S(cQ); // (ACPY,ACPY_M,ACPY_K) -> (blk_m,blk_k)
Tensor tKVcKV = gmem_thr_copy_KV.partition_S(cKV); // (BCPY,BCPY_N,BCPY_K) -> (blk_n,blk_k)
// Prologue
// Read Q from gmem to smem, optionally apply rotary embedding.
Tensor tNopeQrNopeQ = make_fragment_like(tNopeQgNopeQ);
Tensor tRopeQrRopeQ = make_fragment_like(tRopeQgRopeQ);
// We don't need to clear the sQ smem tiles since we'll only write out the valid outputs
flash::copy_b128</*Is_even_MN=*/false, /*Is_even_K*/true>(tNopeQgNopeQ, tNopeQrNopeQ, tQcQ, params.d, binfo.actual_seqlen_q - m_block * kBlockM);
flash::copy_b128</*Is_even_MN=*/false, /*Is_even_K*/true>(tRopeQgRopeQ, tRopeQrRopeQ, tQcQ, params.d, binfo.actual_seqlen_q - m_block * kBlockM);
cute::copy(tNopeQrNopeQ, tNopeQsNopeQ);
flash::sync_threads();
cute::copy(smem_tiled_copy_Q, tSsNopeQ, tSrNopeQ);
flash::sync_threads();
cute::copy(tRopeQrRopeQ, tRopeQsRopeQ);
flash::sync_threads();
cute::copy(smem_tiled_copy_Q, tSsRopeQ, tSrRopeQ);
flash::sync_threads();
flash::concat(tSrNopeQ, tSrRopeQ, tSrQ);
int n_block = n_block_max - 1;
int Ksmem_read_index = 0;
int Ksmem_write_index = 0;
// We don't need to clear the sK smem tiles since we'll mask out the scores anyway.
Tensor tKrK = make_fragment_like(tKgK);
int indices_start_block = (n_block * kBlockN) / kBlockTopK * kBlockTopK;
int offset_indices_per_q = indices_start_block + (tidx << 2);
int4 indices_vec = {-1, -1, -1, -1};
if (offset_indices_per_q + 4 <= params.topk) {
indices_vec = __ldg(reinterpret_cast<const int4*>(&gIndices[offset_indices_per_q]));
}
else if (offset_indices_per_q < params.topk) {
indices_vec.x = offset_indices_per_q + 0 < params.topk ? gIndices[offset_indices_per_q + 0] : -1;
indices_vec.y = offset_indices_per_q + 1 < params.topk ? gIndices[offset_indices_per_q + 1] : -1;
indices_vec.z = offset_indices_per_q + 2 < params.topk ? gIndices[offset_indices_per_q + 2] : -1;
indices_vec.w = offset_indices_per_q + 3 < params.topk ? gIndices[offset_indices_per_q + 3] : -1;
}
*((int4*)(&indices_smem_ptr[offset_indices_per_q % kBlockTopK])) = indices_vec;
flash::sync_threads();
uint32_t row_offset = tidx / (kNWarps * 64 / kBlockN) + n_block * kBlockN;
int32_t topk_sparse_idx = indices_smem_ptr[row_offset % kBlockTopK];
topk_sparse_idx = topk_sparse_idx < 0 ? 0 : topk_sparse_idx;
flash::copy_b32_sparse<Kernel_traits, /*Is_even_MN=*/false, /*Is_even_K=*/true>(gK, tKgK, tKrK, tKVcKV, params.d, n_block,
params.d, topk_sparse_idx,
params.topk - n_block * kBlockN);
flash::cp_async_wait<0>();
clear(acc_o);
flash::Softmax<size<1>(acc_o)> softmax;
flash::Mask<Is_causal> mask(params.topk, binfo.actual_seqlen_q, params.ngroups);
for (; n_block >= n_block_min; --n_block) {
Tensor acc_s = partition_fragment_C(tiled_mma_s, Shape<Int<kBlockM>, Int<kBlockN>>{}); // (MMA=4, MMA_M, MMA_N)
cute::copy(tKrK, tKsK(_, _, _, Ksmem_write_index));
Ksmem_write_index ^= 1;
clear(acc_s);
flash::sync_threads();
flash::gemm</*A_in_regs=*/Kernel_traits::Is_Q_in_regs>(
acc_s, tSrQ, tSrK, tSsNopeQ, tSsK(_, _, _, Ksmem_read_index), tiled_mma_s, smem_tiled_copy_Q, smem_tiled_copy_K,
smem_thr_copy_Q, smem_thr_copy_K
);
mask.template apply_sparse_attn_mask<kBlockTopK, false>(acc_s, n_block * kBlockN, indices_smem_ptr, indices_all_valid_per_q);
if ((n_block * kBlockN) % kBlockTopK == 0 && n_block > n_block_min) {
flash::barrier();
indices_start_block = ((n_block - 1) * kBlockN) / kBlockTopK * kBlockTopK;
offset_indices_per_q = indices_start_block + (tidx << 2);
indices_vec = __ldg(reinterpret_cast<const int4*>(&gIndices[offset_indices_per_q]));
*((int4*)(&indices_smem_ptr[offset_indices_per_q % kBlockTopK])) = indices_vec;
flash::sync_threads();
}
n_block == n_block_max - 1
? softmax.template softmax_rescale_o</*Is_first=*/true, /*Check_inf=*/Is_causal || true, true, true>(acc_s, acc_o, params.scale_softmax_log2)
: softmax.template softmax_rescale_o</*Is_first=*/false, /*Check_inf=*/Is_causal || true, true, true>(acc_s, acc_o, params.scale_softmax_log2);
if (n_block > n_block_min) {
row_offset -= kBlockN;
topk_sparse_idx = indices_smem_ptr[row_offset % kBlockTopK];
topk_sparse_idx = topk_sparse_idx < 0 ? 0 : topk_sparse_idx;
flash::copy_b32_sparse<Kernel_traits, /*Is_even_MN=*/true, true>(gK, tKgK, tKrK, tKVcKV, params.d, n_block - 1, params.d, topk_sparse_idx);
}
CONVERT_TENSOR_TYPE(ElementAccum, Element, acc_s, rP)
lds4x4_with_swizzle424(tOsVt(_, _, Ksmem_read_index), tOrVt);
CUTE_STATIC_ASSERT_V(size<2>(tOrVt) == _1{}); // only support MMA_K = 1
Tensor tOrVt_permute_view = make_tensor(tOrVt.data(), make_layout(make_shape(size<0>(tOrVt), size<1, 0>(tOrVt), size<1, 1>(tOrVt))));
permute_4x4_b16(tOrVt_permute_view);
Tensor tOrP = make_tensor(rP.data(), acc_s.layout());
flash::gemm_rr(acc_o, tOrP, tOrVt, tiled_mma_o);
Ksmem_read_index ^= 1;
}
// Epilogue
if (NoSplit) {
store_64x16_xcore1000<Kernel_traits, /*Split*/false, /*Is_even_MN=*/false, /*Is_even_K=*/true>(params, bidb, bidh, m_block, n_split_idx, smem_, acc_o, softmax);
}else{
store_64x16_xcore1000<Kernel_traits, /*Split*/true, /*Is_even_MN=*/false, /*Is_even_K=*/true>(params, bidb, bidh, m_block, n_split_idx, smem_, acc_o, softmax);
}
}
} // namespace flash

View File

@ -0,0 +1,347 @@
#pragma once
#include <cute/algorithm/copy.hpp>
#include <mctlass/mctlass.h>
#include <mctlass/array.h>
#include <mctlass/numeric_types.h>
#include "block_info.h"
#include "kernel_traits.h"
#include "utils.h"
#include "softmax.h"
#include "mask.h"
namespace flash {
using namespace cute;
template<typename Kernel_traits, bool Is_causal, bool Is_even_TopK, typename Params>
__forceinline__ __device__ void sparse_attn_fwd_kernel(const Params &params) {
constexpr bool Is_local = false;
using Element = typename Kernel_traits::Element;
using ElementAccum = typename Kernel_traits::ElementAccum;
using index_t = typename Kernel_traits::index_t;
// Shared memory.
extern __shared__ char smem_[];
// The thread index.
const int tidx = threadIdx.x;
const int warp_idx = tidx / 64;
const int lane_idx = tidx % 64;
constexpr int kBlockM = Kernel_traits::kBlockM;
constexpr int kBlockN = Kernel_traits::kBlockN;
constexpr int kHeadDim = Kernel_traits::kHeadDim;
constexpr int kHeadDimV = Kernel_traits::kHeadDimV;
constexpr int kNWarps = Kernel_traits::kNWarps;
constexpr int kBlockKSmem = Kernel_traits::kBlockKSmem;
constexpr int kAtomLayoutMS = Kernel_traits::kAtomLayoutMS;
constexpr int kAtomLayoutMO = Kernel_traits::kAtomLayoutMO;
constexpr int Num_Stages = Kernel_traits::Num_Stages;
constexpr int kBlockTopK = Kernel_traits::kBlockTopK;
constexpr int kHeadDimNope = kHeadDimV;
constexpr int kHeadDimRope = kHeadDim - kHeadDimV;
static_assert(kBlockKSmem == 64);
static_assert(Kernel_traits::Share_Q_K_smem && Kernel_traits::Is_Q_in_regs);
using GmemTiledCopyO = typename Kernel_traits::GmemTiledCopyO;
using ElementO = Element;
const int bidb = 0;
const int h_q_idx = blockIdx.x % (params.h_q / kBlockM); //q_h_idx
const int s_q_idx = blockIdx.x / (params.h_q / kBlockM); //s_q_idx
const int q_block_idx = h_q_idx * kBlockM;
// const BlockInfo</*Varlen=*/!Is_even_MN> binfo(params, bidb);
if (q_block_idx >= params.h_q) return;
const int n_block_min = 0;
int n_block_max = cute::ceil_div(params.topk, kBlockN);
// if (Is_causal || Is_local) {
// n_block_max = std::min(n_block_max,
// cute::ceil_div((m_block + 1) * kBlockM + params.s_kv - params.s_q / params.ngroups + params.window_size_right, kBlockN));
//
// We iterate over the blocks in reverse order. This is because the last block is the only one
// that needs masking when we read K and V from global memory. Moreover, iterating in reverse
// might save us 1 register (we just need n_block instead of both n_block and n_block_max).
const index_t row_offset_q = q_block_idx * params.q_head_stride + s_q_idx * params.q_row_stride;
const index_t row_offset_k = (n_block_max - 1) * kBlockN * params.k_row_stride;
const index_t offset_indices = s_q_idx * params.stride_indices_s_q;
// We move K and V to the last block.
// const int bidb_cache = bidb;
// const int *block_table = nullptr;
const int* gIndices = params.indices_ptr + offset_indices;
const bool indices_all_valid_per_q = params.indices_all_valid_per_q_ptr[s_q_idx];
Tensor gNopeQ = make_tensor(make_gmem_ptr(reinterpret_cast<Element *>(params.q_ptr) + row_offset_q),
Shape<Int<kBlockM>, Int<kHeadDimNope>>{},
make_stride(params.q_head_stride, _1{}));
Tensor gRopeQ = make_tensor(make_gmem_ptr(gNopeQ.data().get() + kHeadDimNope),
Shape<Int<kBlockM>, Int<kHeadDimRope>>{},
make_stride(params.q_head_stride, _1{}));
Tensor gK = make_tensor(make_gmem_ptr(reinterpret_cast<Element *>(params.kv_ptr)),
Shape<Int<kBlockN>, Int<kHeadDim>>{},
make_stride(params.k_row_stride, _1{}));
Tensor sQ = make_tensor(make_smem_ptr(reinterpret_cast<Element *>(smem_)),
typename Kernel_traits::SmemLayoutQ{});
Tensor sNopeQ = make_tensor(make_smem_ptr(reinterpret_cast<Element *>(smem_)),
typename Kernel_traits::SmemLayoutNopeQ{});
Tensor sRopeQ = make_tensor(make_smem_ptr(reinterpret_cast<Element *>(smem_)),
typename Kernel_traits::SmemLayoutRopeQ{});
Tensor sK = make_tensor(sNopeQ.data(),
typename Kernel_traits::SmemLayoutK424{});
Tensor sV = make_tensor(sK.data(), typename Kernel_traits::SmemLayoutVtNoSwizzle{});
Tensor sVt = make_tensor(sV.data(), typename Kernel_traits::SmemLayoutVtransposed424{});
int32_t* indices_smem_ptr = reinterpret_cast<int32_t *>(reinterpret_cast<char *>(smem_) + (size(sK) * sizeof(Element) + 3) / 4 * 4);
typename Kernel_traits::GmemTiledCopyB128 gmem_tiled_copy_Q;
auto gmem_thr_copy_Q = gmem_tiled_copy_Q.get_thread_slice(tidx);
Tensor tNopeQgNopeQ = gmem_thr_copy_Q.partition_S(gNopeQ);
Tensor tNopeQsNopeQ = gmem_thr_copy_Q.partition_D(sNopeQ);
Tensor tRopeQgRopeQ = gmem_thr_copy_Q.partition_S(gRopeQ);
Tensor tRopeQsRopeQ = gmem_thr_copy_Q.partition_D(sRopeQ);
typename Kernel_traits::GmemTiledCopyB32 gmem_tiled_copy_KV;
auto gmem_thr_copy_KV = gmem_tiled_copy_KV.get_thread_slice(tidx);
Tensor tKgK = gmem_thr_copy_KV.partition_S(gK); // (KCPY, KCPY_N, KCPY_K)
Tensor tKsK = gmem_thr_copy_KV.partition_D(sK);
// gemm S is 4x1 wave layout, wave(n) compute the same S with wave(n + 4)
int tidx_mma_s = tidx & 0xFF;
typename Kernel_traits::TiledMmaS tiled_mma_s;
auto thr_mma_s = tiled_mma_s.get_thread_slice(tidx_mma_s);
Tensor tSrQ = thr_mma_s.partition_fragment_A(sQ); // (MMA,MMA_M,MMA_K)
Tensor tSrNopeQ = thr_mma_s.partition_fragment_A(sNopeQ); // (MMA,MMA_M,MMA_K)
Tensor tSrRopeQ = thr_mma_s.partition_fragment_A(sRopeQ); // (MMA,MMA_M,MMA_K)
Tensor tSrK = thr_mma_s.partition_fragment_B(sK(_, _, 0)); // (MMA,MMA_N,MMA_K)
typename Kernel_traits::TiledMmaO tiled_mma_o;
auto thr_mma_o = tiled_mma_o.get_thread_slice(tidx);
// Tensor tOrVt = thr_mma_o.partition_fragment_B(sVt); // (MMA, MMA_K,MMA_N)
Tensor tOrVt = make_tensor<Element>(Shape<_4, Shape<_4, _4>, _1>{});
Tensor acc_o = partition_fragment_C(tiled_mma_o, Shape<Int<kBlockM>, Int<kHeadDimV>>{}); // MMA, MMA_M, MMA_K
//
// Copy Atom retiling
//
auto smem_tiled_copy_Q = make_tiled_copy_A(typename Kernel_traits::UniversalCopyAtomB64{}, tiled_mma_s);
auto smem_thr_copy_Q = smem_tiled_copy_Q.get_thread_slice(tidx_mma_s);
Tensor tSsNopeQ = smem_thr_copy_Q.partition_S(sNopeQ);
Tensor tSsRopeQ = smem_thr_copy_Q.partition_S(sRopeQ);
auto smem_tiled_copy_K = make_tiled_copy_B(typename Kernel_traits::UniversalCopyAtomB64{}, tiled_mma_s);
auto smem_thr_copy_K = smem_tiled_copy_K.get_thread_slice(tidx_mma_s);
Tensor tSsK = smem_thr_copy_K.partition_S(sK);
auto smem_tiled_copy_V = make_tiled_copy_B(typename Kernel_traits::SmemCopyAtomTransposed{}, tiled_mma_o);
auto smem_thr_copy_V = smem_tiled_copy_V.get_thread_slice(tidx);
int warp_offset = warp_idx / kAtomLayoutMO * 16 * 64;
int thread_offset = lane_idx / 16 * 4 * 64;
Element *Vtsmem_ptr_lds = reinterpret_cast<Element *>(sVt.data().get()) + warp_offset + thread_offset;
Tensor tOsVt = make_tensor(make_smem_ptr(Vtsmem_ptr_lds), make_layout(Shape<_4, _4, Int<Num_Stages>>{}, // MMA MMA_N NUM_STAGES
Stride<_1, Int<16*128>, Int<kBlockN*kHeadDim>>{}));
// PREDICATES
// Construct identity layout for sQ and sK
Tensor cQ = make_identity_tensor(make_shape(size<0>(sQ), size<1>(sQ))); // (BLK_M,BLK_K) -> (blk_m,blk_k)
Tensor cKV = make_identity_tensor(make_shape(size<0>(sK), size<1>(sK))); // (BLK_N,BLK_K) -> (blk_n,blk_k)
// Repeat the partitioning with identity layouts
Tensor tQcQ = gmem_thr_copy_Q.partition_S(cQ); // (ACPY,ACPY_M,ACPY_K) -> (blk_m,blk_k)
Tensor tKVcKV = gmem_thr_copy_KV.partition_S(cKV); // (BCPY,BCPY_N,BCPY_K) -> (blk_n,blk_k)
// Prologue
// Read Q from gmem to smem, optionally apply rotary embedding.
Tensor tNopeQrNopeQ = make_fragment_like(tNopeQgNopeQ);
Tensor tRopeQrRopeQ = make_fragment_like(tRopeQgRopeQ);
// We don't need to clear the sQ smem tiles since we'll only write out the valid outputs
flash::copy_b128</*Is_even_MN*/ true, true>(tNopeQgNopeQ, tNopeQrNopeQ, tQcQ, params.d_qk, params.h_q - q_block_idx);
flash::copy_b128</*Is_even_MN*/ true, true>(tRopeQgRopeQ, tRopeQrRopeQ, tQcQ, params.d_qk, params.h_q - q_block_idx);
cute::copy(tNopeQrNopeQ, tNopeQsNopeQ);
flash::sync_threads();
cute::copy(smem_tiled_copy_Q, tSsNopeQ, tSrNopeQ);
flash::sync_threads();
cute::copy(tRopeQrRopeQ, tRopeQsRopeQ);
flash::sync_threads();
cute::copy(smem_tiled_copy_Q, tSsRopeQ, tSrRopeQ);
flash::sync_threads();
flash::concat(tSrNopeQ, tSrRopeQ, tSrQ);
int n_block = n_block_max - 1;
int Ksmem_read_index = 0;
int Ksmem_write_index = 0;
// We don't need to clear the sK smem tiles since we'll mask out the scores anyway.
Tensor tKrK = make_fragment_like(tKgK);
int indices_start_block = (n_block * kBlockN) / kBlockTopK * kBlockTopK;
int offset_indices_per_q = indices_start_block + (tidx << 2);
int4 indices_vec = {-1, -1, -1, -1};
if (offset_indices_per_q + 4 <= params.topk) {
indices_vec = __ldg(reinterpret_cast<const int4*>(&gIndices[offset_indices_per_q]));
}
else if (offset_indices_per_q < params.topk) {
indices_vec.x = offset_indices_per_q + 0 < params.topk ? gIndices[offset_indices_per_q + 0] : -1;
indices_vec.y = offset_indices_per_q + 1 < params.topk ? gIndices[offset_indices_per_q + 1] : -1;
indices_vec.z = offset_indices_per_q + 2 < params.topk ? gIndices[offset_indices_per_q + 2] : -1;
indices_vec.w = offset_indices_per_q + 3 < params.topk ? gIndices[offset_indices_per_q + 3] : -1;
}
*((int4*)(&indices_smem_ptr[offset_indices_per_q % kBlockTopK])) = indices_vec;
flash::sync_threads();
uint32_t row_offset = tidx / 32 + n_block * kBlockN;
int32_t topk_sparse_idx = indices_smem_ptr[row_offset % kBlockTopK];
topk_sparse_idx = topk_sparse_idx < 0 ? 0 : topk_sparse_idx;
flash::copy_b32_sparse<Kernel_traits, /*Is_even_MN=*/false, /*Is_even_K=*/true>(gK, tKgK, tKrK, tKVcKV, params.d_qk, n_block, params.k_row_stride, topk_sparse_idx, params.topk - n_block * kBlockN);
clear(acc_o);
flash::Softmax<size<1>(acc_o)> softmax;
flash::Mask<Is_causal> mask(params.topk, params.s_q, params.h_q);
for (; n_block >= n_block_min; --n_block) {
Tensor acc_s = partition_fragment_C(tiled_mma_s, Shape<Int<kBlockM>, Int<kBlockN>>{}); // (MMA=4, MMA_M, MMA_N)
cute::copy(tKrK, tKsK(_, _, _, Ksmem_write_index));
Ksmem_write_index ^= 1;
clear(acc_s);
flash::sync_threads();
flash::gemm</*A_in_regs=*/Kernel_traits::Is_Q_in_regs>(
acc_s, tSrQ, tSrK, tSsNopeQ, tSsK(_, _, _, Ksmem_read_index), tiled_mma_s, smem_tiled_copy_Q, smem_tiled_copy_K,
smem_thr_copy_Q, smem_thr_copy_K
);
mask.template apply_sparse_attn_mask<kBlockTopK, /*Is_even_MN=*/true>(acc_s, n_block * kBlockN, indices_smem_ptr, indices_all_valid_per_q);
if ((n_block * kBlockN) % kBlockTopK == 0 && n_block > n_block_min) {
flash::barrier();
indices_start_block = ((n_block - 1) * kBlockN) / kBlockTopK * kBlockTopK;
offset_indices_per_q = indices_start_block + (tidx << 2);
int4 indices_vec = __ldg(reinterpret_cast<const int4*>(&gIndices[offset_indices_per_q]));
*((int4*)(&indices_smem_ptr[offset_indices_per_q % kBlockTopK])) = indices_vec;
flash::sync_threads();
}
n_block == n_block_max - 1
? softmax.template softmax_rescale_o</*Is_first=*/true, /*Check_inf=*/true, true, true>(acc_s, acc_o, params.sm_scale_div_log2)
: softmax.template softmax_rescale_o</*Is_first=*/false, /*Check_inf=*/true, true, true>(acc_s, acc_o, params.sm_scale_div_log2);
if (n_block > n_block_min) {
row_offset -= kBlockN;
topk_sparse_idx = indices_smem_ptr[row_offset % kBlockTopK];
topk_sparse_idx = topk_sparse_idx < 0 ? 0 : topk_sparse_idx;
flash::copy_b32_sparse<Kernel_traits, /*Is_even_MN=*/true, true>(gK, tKgK, tKrK, tKVcKV, params.d_qk, n_block - 1, params.k_row_stride, topk_sparse_idx);
}
//Tensor rP = flash::convert_type<Element>(acc_s);
CONVERT_TENSOR_TYPE(ElementAccum, Element, acc_s, rP)
// Reshape rP from (MMA=4, MMA_M, MMA_N) to ((4, 2), MMA_M, MMA_N / 2)
// if using m16n8k16 or (4, MMA_M, MMA_N) if using m16n8k8.
//Tensor tOrP = make_tensor(rP.data(), flash::convert_layout_acc_Aregs<Kernel_traits::TiledMma>(rP.layout()));
lds4x4_with_swizzle424(tOsVt(_, _, Ksmem_read_index), tOrVt);
CUTE_STATIC_ASSERT_V(size<2>(tOrVt) == _1{}); // only support MMA_K = 1
Tensor tOrVt_permute_view = make_tensor(tOrVt.data(), make_layout(make_shape(size<0>(tOrVt), size<1, 0>(tOrVt), size<1, 1>(tOrVt))));
permute_4x4_b16(tOrVt_permute_view);
Tensor tOrP = make_tensor(rP.data(), acc_s.layout());
flash::gemm_rr(acc_o, tOrP, tOrVt, tiled_mma_o);
Ksmem_read_index ^= 1;
}
// if(thread0()){print(acc_o);}
// Epilogue
Tensor lse = softmax.template normalize_softmax_lse</*Is_dropout=*/false, /*Return_lse*/true, false>(acc_o, params.sm_scale);
Tensor acc_o_view = make_tensor(acc_o.data(), make_layout(Shape<_4, Shape<_4, _4>>{},
Stride<_1, Shape<_4, _16>>{}));
Tensor acc_o_copy = make_fragment_like(acc_o_view);
#pragma unroll
for (int k = 0; k < size<1, 1>(acc_o_view); k++) {
#pragma unroll
for (int idx = 0; idx < 16; idx++) {
int row = idx / 4;
int col = idx % 4;
acc_o_copy(row, make_coord(col, k)) = acc_o_view(col, make_coord(row, k));
}
}
// if (cute::thread0()) { print(lse); }
Tensor sOaccum = make_tensor(make_smem_ptr(reinterpret_cast<ElementO *>(smem_)), typename Kernel_traits::SmemLayoutO{}); // (SMEM_M,SMEM_N)
// Partition sO to match the accumulator partitioning
using SmemTiledCopyO = typename Kernel_traits::SmemCopyAtomO;
CONVERT_TENSOR_TYPE(ElementAccum, ElementO, acc_o_copy, rO)
warp_offset = warp_idx * 16 * 64;
thread_offset = lane_idx % 16 * 64 + lane_idx / 16 * 16;
Element *Osmem_ptr_sts = reinterpret_cast<ElementO *>(smem_) + warp_offset + thread_offset;
Tensor tOsO = make_tensor(make_smem_ptr(Osmem_ptr_sts), make_layout(Shape<_16, _4>{},
Stride<_1, Int<16*64*kNWarps>>{}));
Tensor tOrO = make_tensor(rO.data(), make_layout(Shape<_16, _4>{},
Stride<_1, _16>{}));
if constexpr (Kernel_traits::Share_Q_K_smem) { flash::sync_threads(); }
cute::copy(tOrO, tOsO);
const index_t row_offset_o = q_block_idx * params.o_head_stride + s_q_idx * params.o_row_stride;
const index_t row_offset_lseaccum = s_q_idx * params.h_q + q_block_idx;
const index_t row_offset_max_logits = row_offset_lseaccum;
Tensor gOaccum = make_tensor(make_gmem_ptr(reinterpret_cast<ElementO *>(params.out_ptr) + (row_offset_o)),
Shape<Int<kBlockM>, Int<kHeadDimV>>{},
make_stride(params.o_head_stride, _1{}));
Tensor gLSEaccum = make_tensor(make_gmem_ptr(reinterpret_cast<ElementAccum *>(params.lse_ptr) + row_offset_lseaccum),
Shape<Int<kBlockM>>{}, Stride<_1>{});
Tensor gMaxLogits = make_tensor(make_gmem_ptr(reinterpret_cast<ElementAccum *>(params.max_logits) + row_offset_max_logits),
Shape<Int<kBlockM>>{}, Stride<_1>{});
// if (tidx == 0) { printf("row_offset_o = %d, s_q_idx = %d, gOaccum = %p\n", row_offset_o, s_q_idx, gOaccum.data()); }
GmemTiledCopyO gmem_tiled_copy_Oaccum;
auto gmem_thr_copy_Oaccum = gmem_tiled_copy_Oaccum.get_thread_slice(tidx);
Tensor tOsOaccum = gmem_thr_copy_Oaccum.partition_S(sOaccum); // ((Atom,AtomNum),ATOM_M,ATOM_N)
Tensor tOgOaccum = gmem_thr_copy_Oaccum.partition_D(gOaccum);
flash::sync_threads();
Tensor tOrOaccum = make_tensor<ElementO>(shape(tOgOaccum));
cute::copy(gmem_tiled_copy_Oaccum, tOsOaccum, tOrOaccum);
Tensor caccO = make_identity_tensor(Shape<Int<kBlockM>, Int<kHeadDimV>>{}); // (BLK_M,BLK_K) -> (blk_m,blk_k)
Tensor taccOcO = thr_mma_o.partition_C(caccO); // (MMA,MMA_M,MMA_K)
static_assert(decltype(size<0>(taccOcO))::value == 4);
// Convert to ((2, 2), MMA_M, MMA_K) then take only the row indices.
Tensor taccOcO_row = logical_divide(taccOcO, Shape<_4>{})(make_coord(0, _), _, 0);
CUTE_STATIC_ASSERT_V(size(lse) == size(taccOcO_row)); // MMA_M
if (get<1>(taccOcO_row(0)) == 0) {
#pragma unroll
for (int mi = 0; mi < size(lse); ++mi) {
const int row = get<0>(taccOcO_row(mi));
if (row < params.h_q - q_block_idx) {
gMaxLogits(row) = softmax.row_max(mi) * params.sm_scale * M_LOG2E;
gLSEaccum(row) = lse(mi);
}
}
}
// Construct identity layout for sO
Tensor cO = make_identity_tensor(make_shape(size<0>(sOaccum), size<1>(sOaccum))); // (BLK_M,BLK_K) -> (blk_m,blk_k)
// Repeat the partitioning with identity layouts
Tensor tOcO = gmem_thr_copy_Oaccum.partition_D(cO); // (ACPY,ACPY_M,ACPY_K) -> (blk_m,blk_k)
// Clear_OOB_K must be false since we don't want to write zeros to gmem
flash::copy_reg_to_global</*Is_even_MN*/ true, true>(
tOrOaccum, tOgOaccum, tOcO, params.d_v, params.h_q - q_block_idx
);
}
} // namespace flash

View File

@ -0,0 +1,341 @@
// Adapted from Dao-AILab/flash-attention (https://github.com/Dao-AILab/flash-attention/tree/v2.6.3)
#pragma once
#include <cute/algorithm/copy.hpp>
#include <mctlass/mctlass.h>
#include <mctlass/array.h>
#include <mctlass/numeric_types.h>
#include "block_info.h"
#include "kernel_traits.h"
#include "utils.h"
#include "softmax.h"
#include "mask.h"
namespace flash {
using namespace cute;
template<typename Kernel_traits, bool Split, bool Is_even_MN, bool Is_even_K, typename AccO, typename Softmax, typename Params>
__forceinline__ __device__ void store_64x16_xcore1500(const Params &params, const int bidb, const int bidh, const int m_block, const int n_split_idx, char* smem_, AccO acc_o, Softmax softmax){
using ElementAccum = typename Kernel_traits::ElementAccum;
using Element = typename Kernel_traits::Element;
using index_t = typename Kernel_traits::index_t;
using GmemTiledCopyO = std::conditional_t<
!Split,
typename Kernel_traits::GmemTiledCopyO,
typename Kernel_traits::GmemTiledCopyOaccum
>;
using ElementO = std::conditional_t<!Split, Element, ElementAccum>;
// The thread index.
const int tidx = threadIdx.x;
const int warp_idx = tidx / 64;
const int lane_idx = tidx % 64;
constexpr int kBlockM = Kernel_traits::kBlockM;
constexpr int kBlockN = Kernel_traits::kBlockN;
constexpr int kNWarps = Kernel_traits::kNWarps;
constexpr int kHeadDimV = Kernel_traits::kHeadDimV;
const BlockInfo</*Varlen=*/!Is_even_MN> binfo(params, bidb);
typename Kernel_traits::TiledMmaO tiled_mma_o;
auto thr_mma_o = tiled_mma_o.get_thread_slice(tidx);
Tensor lse = softmax.template normalize_softmax_lse</*Is_dropout=*/false, /*Return_lse*/true, Split>(acc_o, params.scale_softmax);
const int split_offset = params.num_splits_ptr[bidb];
Tensor sOaccum = make_tensor(make_smem_ptr(reinterpret_cast<ElementO *>(smem_)), typename Kernel_traits::SmemLayoutO{}); // (SMEM_M,SMEM_N)
// Partition sO to match the accumulator partitioning
using SmemTiledCopyO = std::conditional_t<
!Split,
typename Kernel_traits::SmemCopyAtomO,
typename Kernel_traits::SmemCopyAtomOaccum
>;
CONVERT_TENSOR_TYPE(ElementAccum, ElementO, acc_o, rO)
auto smem_tiled_copy_Oaccum = make_tiled_copy_C(SmemTiledCopyO{}, tiled_mma_o);
auto smem_thr_copy_Oaccum = smem_tiled_copy_Oaccum.get_thread_slice(tidx);
Tensor taccOrOaccum = smem_thr_copy_Oaccum.retile_S(rO); // ((Atom,AtomNum), MMA_M, MMA_N)
Tensor taccOsOaccum = smem_thr_copy_Oaccum.partition_D(sOaccum); // ((Atom,AtomNum),PIPE_M,PIPE_N)
if constexpr (Split || Kernel_traits::Share_Q_K_smem) { flash::sync_threads(); }
cute::copy(smem_tiled_copy_Oaccum, taccOrOaccum, taccOsOaccum);
const index_t row_offset_o = binfo.q_offset(params.o_batch_stride, params.o_row_stride, bidb)
+ m_block * kBlockM * params.o_row_stride + bidh * params.o_head_stride;
const index_t row_offset_oaccum = (((split_offset + n_split_idx) * params.h + bidh) * params.seqlen_q
+ m_block * kBlockM) * params.d_v;
const index_t row_offset_lseaccum = ((split_offset + n_split_idx) * params.h + bidh) * params.seqlen_q + m_block * kBlockM;
Tensor gOaccum = make_tensor(make_gmem_ptr(reinterpret_cast<ElementO *>(Split ? params.oaccum_ptr : params.o_ptr) + (Split ? row_offset_oaccum : row_offset_o)),
Shape<Int<kBlockM>, Int<kHeadDimV>>{},
make_stride(Split ? kHeadDimV : params.o_row_stride, _1{}));
Tensor gLSEaccum = make_tensor(make_gmem_ptr(reinterpret_cast<ElementAccum *>(Split ? params.softmax_lseaccum_ptr : params.softmax_lse_ptr) + row_offset_lseaccum),
Shape<Int<kBlockM>>{}, Stride<_1>{});
// if (tidx == 0) { printf("row_offset_o = %d, bidh = %d, gOaccum = %p\n", row_offset_o, bidh, gOaccum.data()); }
GmemTiledCopyO gmem_tiled_copy_Oaccum;
auto gmem_thr_copy_Oaccum = gmem_tiled_copy_Oaccum.get_thread_slice(tidx);
Tensor tOsOaccum = gmem_thr_copy_Oaccum.partition_S(sOaccum); // ((Atom,AtomNum),ATOM_M,ATOM_N)
Tensor tOgOaccum = gmem_thr_copy_Oaccum.partition_D(gOaccum);
flash::sync_threads();
Tensor tOrOaccum = make_tensor<ElementO>(shape(tOgOaccum));
cute::copy(gmem_tiled_copy_Oaccum, tOsOaccum, tOrOaccum);
Tensor caccO = make_identity_tensor(Shape<Int<kBlockM>, Int<kHeadDimV>>{}); // (BLK_M,BLK_K) -> (blk_m,blk_k)
Tensor taccOcO = thr_mma_o.partition_C(caccO); // (MMA,MMA_M,MMA_K)
static_assert(decltype(size<0>(taccOcO))::value == 4);
// Convert to ((2, 2), MMA_M, MMA_K) then take only the row indices.
Tensor taccOcO_row = logical_divide(taccOcO, Shape<_4>{})(make_coord(0, _), _, 0);
CUTE_STATIC_ASSERT_V(size(lse) == size(taccOcO_row)); // MMA_M
if (get<1>(taccOcO_row(0)) == 0) {
#pragma unroll
for (int mi = 0; mi < size(lse); ++mi) {
const int row = get<0>(taccOcO_row(mi));
if (row < binfo.actual_seqlen_q - m_block * kBlockM) { gLSEaccum(row) = lse(mi); }
}
}
// Construct identity layout for sO
Tensor cO = make_identity_tensor(make_shape(size<0>(sOaccum), size<1>(sOaccum))); // (BLK_M,BLK_K) -> (blk_m,blk_k)
// Repeat the partitioning with identity layouts
Tensor tOcO = gmem_thr_copy_Oaccum.partition_D(cO); // (ACPY,ACPY_M,ACPY_K) -> (blk_m,blk_k)
// Clear_OOB_K must be false since we don't want to write zeros to gmem
flash::copy_reg_to_global<Is_even_MN, Is_even_K>(
tOrOaccum, tOgOaccum, tOcO, params.d_v, binfo.actual_seqlen_q - m_block * kBlockM
);
}
template<typename Kernel_traits, bool Is_causal, bool Is_even_MN, bool Is_even_K, bool Is_enable_dcp, typename Params>
__forceinline__ __device__ void compute_attn_1rowblock_splitkv_mla_k64_64x16_8waves_xcore1500(
const Params &params, const int bidb, const int bidh, const int m_block, const int n_split_idx, const int n_block_min, const int n_block_max, const bool NoSplit) {
using Element = typename Kernel_traits::Element;
using ElementAccum = typename Kernel_traits::ElementAccum;
using index_t = typename Kernel_traits::index_t;
// Shared memory.
extern __shared__ char smem_[];
// The thread index.
const int tidx = threadIdx.x;
const int warp_idx = tidx / 64;
const int lane_idx = tidx % 64;
constexpr int kBlockM = Kernel_traits::kBlockM;
constexpr int kBlockN = Kernel_traits::kBlockN;
constexpr int kHeadDim = Kernel_traits::kHeadDim;
constexpr int kHeadDimV = Kernel_traits::kHeadDimV;
constexpr int kNWarps = Kernel_traits::kNWarps;
constexpr int kBlockKSmem = Kernel_traits::kBlockKSmem;
constexpr int kAtomLayoutMS = Kernel_traits::kAtomLayoutMS;
constexpr int kAtomLayoutMO = Kernel_traits::kAtomLayoutMO;
constexpr int Num_Stages = Kernel_traits::Num_Stages;
constexpr int kHeadDimNope = kHeadDimV;
constexpr int kHeadDimRope = kHeadDim - kHeadDimV;
static_assert(kBlockKSmem == 64);
static_assert(Kernel_traits::Share_Q_K_smem && Kernel_traits::Is_Q_in_regs);
const BlockInfo</*Varlen=*/!Is_even_MN> binfo(params, bidb);
if (m_block * kBlockM >= binfo.actual_seqlen_q) return;
// We iterate over the blocks in reverse order. This is because the last block is the only one
// that needs masking when we read K and V from global memory. Moreover, iterating in reverse
// might save us 1 register (we just need n_block instead of both n_block and n_block_max).
const index_t row_offset_q = binfo.q_offset(params.q_batch_stride, params.q_row_stride, bidb)
+ m_block * kBlockM * params.q_row_stride + bidh * params.q_head_stride;
// We move K and V to the last block.
const int bidb_cache = params.cache_batch_idx == nullptr ? bidb : params.cache_batch_idx[bidb];
const int *block_table = params.block_table == nullptr ? nullptr : params.block_table + bidb * params.block_table_batch_stride;
const int block_table_idx = block_table == nullptr ? 0 : (n_block_max - 1) * kBlockN / params.page_block_size;
const int block_table_offset = block_table == nullptr ? 0 : (n_block_max - 1) * kBlockN - block_table_idx * params.page_block_size;
const index_t row_offset_k = block_table == nullptr
? binfo.k_offset(params.k_batch_stride, params.k_row_stride, bidb_cache)
+ (n_block_max - 1) * kBlockN * params.k_row_stride + (bidh / params.h_h_k_ratio) * params.k_head_stride
: (bidh / params.h_h_k_ratio) * params.k_head_stride;
const index_t row_offset_v = block_table == nullptr
? binfo.k_offset(params.v_batch_stride, params.v_row_stride, bidb_cache)
+ (n_block_max - 1) * kBlockN * params.v_row_stride + (bidh / params.h_h_k_ratio) * params.v_head_stride
: (bidh / params.h_h_k_ratio) * params.v_head_stride;
Tensor gNopeQ = make_tensor(make_gmem_ptr(reinterpret_cast<Element *>(params.q_ptr) + row_offset_q),
Shape<Int<kBlockM>, Int<kHeadDimNope>>{},
make_stride(params.q_row_stride, _1{}));
Tensor gRopeQ = make_tensor(make_gmem_ptr(gNopeQ.data().get() + kHeadDimNope),
Shape<Int<kBlockM>, Int<kHeadDimRope>>{},
make_stride(params.q_row_stride, _1{}));
Tensor gK = make_tensor(make_gmem_ptr(reinterpret_cast<Element *>(params.k_ptr) + row_offset_k),
Shape<Int<kBlockN>, Int<kHeadDim>>{},
make_stride(params.k_row_stride, _1{}));
// if (threadIdx.x == 0 && blockIdx.y == 0 && blockIdx.z == 0) { printf("k_ptr = %p, row_offset_k = %d, gK_ptr = %p\n", params.k_ptr, row_offset_k, gK.data()); }
Tensor gV = make_tensor(make_gmem_ptr(reinterpret_cast<Element *>(params.v_ptr) + row_offset_v),
Shape<Int<kBlockN>, Int<kHeadDimV>>{},
make_stride(params.v_row_stride, _1{}));
Tensor sQ = make_tensor(make_smem_ptr(reinterpret_cast<Element *>(smem_)),
typename Kernel_traits::SmemLayoutQ{});
Tensor sNopeQ = make_tensor(make_smem_ptr(reinterpret_cast<Element *>(smem_)),
typename Kernel_traits::SmemLayoutNopeQ{});
Tensor sRopeQ = make_tensor(make_smem_ptr(reinterpret_cast<Element *>(smem_)),
typename Kernel_traits::SmemLayoutRopeQ{});
Tensor sK = make_tensor(sNopeQ.data(), typename Kernel_traits::SmemLayoutK242{});
Tensor sV = make_tensor(sK.data(), typename Kernel_traits::SmemLayoutVtNoSwizzle{});
Tensor sVt = make_tensor(sV.data(), typename Kernel_traits::SmemLayoutVtransposed242{});
Tensor sVtNoSwizzle = make_tensor(sV.data(), typename Kernel_traits::SmemLayoutVtransposedNoSwizzle{});
typename Kernel_traits::GmemTiledCopyB128 gmem_tiled_copy_Q;
auto gmem_thr_copy_Q = gmem_tiled_copy_Q.get_thread_slice(tidx);
Tensor tNopeQgNopeQ = gmem_thr_copy_Q.partition_S(gNopeQ);
Tensor tNopeQsNopeQ = gmem_thr_copy_Q.partition_D(sNopeQ);
Tensor tRopeQgRopeQ = gmem_thr_copy_Q.partition_S(gRopeQ);
Tensor tRopeQsRopeQ = gmem_thr_copy_Q.partition_D(sRopeQ);
typename Kernel_traits::GmemTiledCopyB32 gmem_tiled_copy_KV;
auto gmem_thr_copy_KV = gmem_tiled_copy_KV.get_thread_slice(tidx);
Tensor tKgK = gmem_thr_copy_KV.partition_S(gK); // (KCPY, KCPY_N, KCPY_K)
Tensor tKsK = gmem_thr_copy_KV.partition_D(sK);
// gemm S is 4x1 wave layout, wave(n) compute the same S with wave(n + 4)
int tidx_mma_s = tidx & 0xFF;
typename Kernel_traits::TiledMmaS tiled_mma_s;
auto thr_mma_s = tiled_mma_s.get_thread_slice(tidx_mma_s);
Tensor tSrQ = thr_mma_s.partition_fragment_A(sQ); // (MMA,MMA_M,MMA_K)
Tensor tSrNopeQ = thr_mma_s.partition_fragment_A(sNopeQ); // (MMA,MMA_M,MMA_K)
Tensor tSrRopeQ = thr_mma_s.partition_fragment_A(sRopeQ); // (MMA,MMA_M,MMA_K)
Tensor tSrK = thr_mma_s.partition_fragment_B(sK(_, _, 0)); // (MMA,MMA_N,MMA_K)
typename Kernel_traits::TiledMmaO tiled_mma_o;
auto thr_mma_o = tiled_mma_o.get_thread_slice(tidx);
Tensor tOrVt = thr_mma_o.partition_fragment_B(sVtNoSwizzle); // (MMA, MMA_K,MMA_N)
Tensor acc_o = partition_fragment_C(tiled_mma_o, Shape<Int<kBlockM>, Int<kHeadDimV>>{}); // MMA, MMA_M, MMA_K
//
// Copy Atom retiling
//
auto smem_tiled_copy_Q = make_tiled_copy_A(typename Kernel_traits::UniversalCopyAtomB64{}, tiled_mma_s);
auto smem_thr_copy_Q = smem_tiled_copy_Q.get_thread_slice(tidx_mma_s);
Tensor tSsNopeQ = smem_thr_copy_Q.partition_S(sNopeQ);
Tensor tSsRopeQ = smem_thr_copy_Q.partition_S(sRopeQ);
auto smem_tiled_copy_K = make_tiled_copy_B(typename Kernel_traits::UniversalCopyAtomB64{}, tiled_mma_s);
auto smem_thr_copy_K = smem_tiled_copy_K.get_thread_slice(tidx_mma_s);
Tensor tSsK = smem_thr_copy_K.partition_S(sK);
auto smem_tiled_copy_V = make_tiled_copy_B(typename Kernel_traits::LDSB64Trans4x16Atom{}, tiled_mma_o);
int warp_offset = warp_idx / 4 * 16;
int thread_offset = lane_idx % 4 * 4
+ lane_idx % 16 / 4 * 64
+ lane_idx / 16 * 4 * 64;
Element *Vtsmem_ptr_lds = reinterpret_cast<Element *>(sVt.data().get()) + warp_offset + thread_offset;
Tensor tOsVt = make_tensor(make_smem_ptr(Vtsmem_ptr_lds), make_layout(Shape<_4, Shape<_2, _8>, Int<kBlockN / 16>, Int<Num_Stages>>{}, // MMA MMA_N NUM_STAGES
Stride<_1, Stride<_32, Int<kBlockN*64>>,Int<16*64>, Int<kBlockN*kHeadDim>>{}));
// PREDICATES
// Construct identity layout for sQ and sK
Tensor cQ = make_identity_tensor(make_shape(size<0>(sQ), size<1>(sQ))); // (BLK_M,BLK_K) -> (blk_m,blk_k)
Tensor cKV = make_identity_tensor(make_shape(size<0>(sK), size<1>(sK))); // (BLK_N,BLK_K) -> (blk_n,blk_k)
// Repeat the partitioning with identity layouts
Tensor tQcQ = gmem_thr_copy_Q.partition_S(cQ); // (ACPY,ACPY_M,ACPY_K) -> (blk_m,blk_k)
Tensor tKVcKV = gmem_thr_copy_KV.partition_S(cKV); // (BCPY,BCPY_N,BCPY_K) -> (blk_n,blk_k)
// Prologue
// Read Q from gmem to smem, optionally apply rotary embedding.
Tensor tNopeQrNopeQ = make_fragment_like(tNopeQgNopeQ);
Tensor tRopeQrRopeQ = make_fragment_like(tRopeQgRopeQ);
// We don't need to clear the sQ smem tiles since we'll only write out the valid outputs
flash::copy_b128<Is_even_MN, Is_even_K>(tNopeQgNopeQ, tNopeQrNopeQ, tQcQ, params.d, binfo.actual_seqlen_q - m_block * kBlockM);
flash::copy_b128<Is_even_MN, Is_even_K>(tRopeQgRopeQ, tRopeQrRopeQ, tQcQ, params.d, binfo.actual_seqlen_q - m_block * kBlockM);
cute::copy(tNopeQrNopeQ, tNopeQsNopeQ);
flash::sync_threads();
cute::copy(smem_tiled_copy_Q, tSsNopeQ, tSrNopeQ);
flash::sync_threads();
cute::copy(tRopeQrRopeQ, tRopeQsRopeQ);
flash::sync_threads();
cute::copy(smem_tiled_copy_Q, tSsRopeQ, tSrRopeQ);
flash::sync_threads();
flash::concat(tSrNopeQ, tSrRopeQ, tSrQ);
int n_block = n_block_max - 1;
int Ksmem_read_index = 0;
int Ksmem_write_index = 0;
// We don't need to clear the sK smem tiles since we'll mask out the scores anyway.
Tensor tKrK = make_fragment_like(tKgK);
int row_offset = tidx / 32 + n_block * kBlockN;
int virtual_page_idx = row_offset / params.page_block_size;
int page_offset = row_offset - virtual_page_idx * params.page_block_size;
int32_t page_idx = block_table[virtual_page_idx];
flash::copy_b32_page_one<Kernel_traits, Is_even_MN, Is_even_K>(gK, tKgK, tKrK, tKVcKV, params.d, n_block,
block_table, params.k_batch_stride, params.k_row_stride, params.page_block_size, page_idx, page_offset, binfo.actual_seqlen_k - n_block * kBlockN);
clear(acc_o);
flash::Softmax<size<1>(acc_o)> softmax;
flash::Mask<Is_causal, Is_enable_dcp> mask(binfo.actual_seqlen_k, binfo.actual_seqlen_q, params.ngroups, binfo.tot_seqlen_k, params.cp_world_size, params.cp_rank);
for (; n_block >= n_block_min; --n_block) {
if (n_block > n_block_min) {
// prefetch load page index
row_offset -= kBlockN;
virtual_page_idx = row_offset / params.page_block_size;
page_offset = row_offset - virtual_page_idx * params.page_block_size;
page_idx = block_table[virtual_page_idx];
}
Tensor acc_s = partition_fragment_C(tiled_mma_s, Shape<Int<kBlockM>, Int<kBlockN>>{}); // (MMA=4, MMA_M, MMA_N)
cute::copy(tKrK, tKsK(_, _, _, Ksmem_write_index));
Ksmem_write_index ^= 1;
clear(acc_s);
flash::sync_threads();
if (n_block > n_block_min) {
// Advance gK
flash::copy_b32_page_one<Kernel_traits, /*Is_even_MN=*/true, Is_even_K>(gK, tKgK, tKrK, tKVcKV, params.d, n_block - 1,
block_table, params.k_batch_stride, params.k_row_stride, params.page_block_size, page_idx, page_offset);
}
flash::gemm</*A_in_regs=*/Kernel_traits::Is_Q_in_regs>(
acc_s, tSrQ, tSrK, tSsNopeQ, tSsK(_, _, _, Ksmem_read_index), tiled_mma_s, smem_tiled_copy_Q, smem_tiled_copy_K,
smem_thr_copy_Q, smem_thr_copy_K
);
mask.template apply_mask<Is_causal, Is_even_MN>(
acc_s, n_block * kBlockN, m_block * kBlockM + (tidx / 64) % kAtomLayoutMS * 16 + (tidx & 0xf), kAtomLayoutMS * 16
);
softmax.template softmax_rescale_o</*Is_first=*/false, /*Check_inf=*/Is_causal || !Is_even_MN, false, false>(acc_s, acc_o, params.scale_softmax_log2);
//Tensor rP = flash::convert_type<Element>(acc_s);
// Reshape rP from (MMA=4, MMA_M, MMA_N) to ((4, 2), MMA_M, MMA_N / 2)
// if using m16n8k16 or (4, MMA_M, MMA_N) if using m16n8k8.
//Tensor tOrP = make_tensor(rP.data(), flash::convert_layout_acc_Aregs<Kernel_traits::TiledMma>(rP.layout()));
cute::copy(smem_tiled_copy_V, tOsVt(_, _, _, Ksmem_read_index), tOrVt);
CONVERT_TENSOR_TYPE(ElementAccum, Element, acc_s, rP)
Tensor tOrP = make_tensor(rP.data(), acc_s.layout());
flash::gemm_rr(acc_o, tOrP, tOrVt, tiled_mma_o);
Ksmem_read_index ^= 1;
}
// Epilogue
if (NoSplit) {
store_64x16_xcore1500<Kernel_traits, false, Is_even_MN, Is_even_K>(params, bidb, bidh, m_block, n_split_idx, smem_, acc_o, softmax);
}else{
store_64x16_xcore1500<Kernel_traits, true, Is_even_MN, Is_even_K>(params, bidb, bidh, m_block, n_split_idx, smem_, acc_o, softmax);
}
}
} // namespace flash

View File

@ -0,0 +1,336 @@
// Adapted from Dao-AILab/flash-attention (https://github.com/Dao-AILab/flash-attention/tree/v2.6.3)
#pragma once
#include <cute/algorithm/copy.hpp>
#include <mctlass/mctlass.h>
#include <mctlass/array.h>
#include <mctlass/numeric_types.h>
#include "block_info.h"
#include "kernel_traits.h"
#include "utils.h"
#include "softmax.h"
#include "mask.h"
namespace flash {
using namespace cute;
template<typename Kernel_traits, bool Split, bool Is_even_MN, bool Is_even_K, typename AccO, typename Softmax, typename Params, typename Tensor0>
__forceinline__ __device__ void store_64x32_xcore1500(const Params &params, const int bidb, const int bidh, const int m_block, const int n_split_idx, char* smem_, AccO acc_o, Softmax softmax, Tensor0& sRowmax) {
using ElementAccum = typename Kernel_traits::ElementAccum;
using Element = typename Kernel_traits::Element;
using index_t = typename Kernel_traits::index_t;
using GmemTiledCopyO = std::conditional_t<
!Split,
typename Kernel_traits::GmemTiledCopyO,
typename Kernel_traits::GmemTiledCopyOaccum
>;
using ElementO = std::conditional_t<!Split, Element, ElementAccum>;
// The thread index.
const int tidx = threadIdx.x;
const int warp_idx = tidx / 64;
const int lane_idx = tidx % 64;
constexpr int kBlockM = Kernel_traits::kBlockM;
constexpr int kHeadDimV = Kernel_traits::kHeadDimV;
const BlockInfo</*Varlen=*/!Is_even_MN> binfo(params, bidb);
typename Kernel_traits::TiledMmaO tiled_mma_o;
auto thr_mma_o = tiled_mma_o.get_thread_slice(tidx);
Tensor lse = softmax.template normalize_softmax_lse</*Is_dropout=*/false, /*Return_lse*/true, Split>(acc_o, sRowmax, params.scale_softmax);
const int split_offset = params.num_splits_ptr[bidb];
Tensor sOaccum = make_tensor(make_smem_ptr(reinterpret_cast<ElementO *>(smem_)), typename Kernel_traits::SmemLayoutO{}); // (SMEM_M,SMEM_N)
// Partition sO to match the accumulator partitioning
using SmemTiledCopyO = std::conditional_t<
!Split,
typename Kernel_traits::SmemCopyAtomO,
typename Kernel_traits::SmemCopyAtomOaccum
>;
CONVERT_TENSOR_TYPE(ElementAccum, ElementO, acc_o, rO)
auto smem_tiled_copy_Oaccum = make_tiled_copy_C(SmemTiledCopyO{}, tiled_mma_o);
auto smem_thr_copy_Oaccum = smem_tiled_copy_Oaccum.get_thread_slice(tidx);
Tensor taccOrOaccum = smem_thr_copy_Oaccum.retile_S(rO); // ((Atom,AtomNum), MMA_M, MMA_N)
Tensor taccOsOaccum = smem_thr_copy_Oaccum.partition_D(sOaccum); // ((Atom,AtomNum),PIPE_M,PIPE_N)
if constexpr (Split || Kernel_traits::Share_Q_K_smem) { flash::sync_threads(); }
cute::copy(smem_tiled_copy_Oaccum, taccOrOaccum, taccOsOaccum);
const index_t row_offset_o = binfo.q_offset(params.o_batch_stride, params.o_row_stride, bidb)
+ m_block * kBlockM * params.o_row_stride + bidh * params.o_head_stride;
const index_t row_offset_oaccum = (((split_offset + n_split_idx) * params.h + bidh) * params.seqlen_q
+ m_block * kBlockM) * params.d_v;
const index_t row_offset_lseaccum = Split ? ((split_offset + n_split_idx) * params.h + bidh) * params.seqlen_q + m_block * kBlockM : ((n_split_idx * params.b + bidb) * params.h + bidh) * params.seqlen_q + m_block * kBlockM;
Tensor gOaccum = make_tensor(make_gmem_ptr(reinterpret_cast<ElementO *>(Split ? params.oaccum_ptr : params.o_ptr) + (Split ? row_offset_oaccum : row_offset_o)),
Shape<Int<kBlockM>, Int<kHeadDimV>>{},
make_stride(Split ? kHeadDimV : params.o_row_stride, _1{}));
Tensor gLSEaccum = make_tensor(make_gmem_ptr(reinterpret_cast<ElementAccum *>(Split ? params.softmax_lseaccum_ptr : params.softmax_lse_ptr) + row_offset_lseaccum),
Shape<Int<kBlockM>>{}, Stride<_1>{});
// if (tidx == 0) { printf("row_offset_o = %d, bidh = %d, gOaccum = %p\n", row_offset_o, bidh, gOaccum.data()); }
GmemTiledCopyO gmem_tiled_copy_Oaccum;
auto gmem_thr_copy_Oaccum = gmem_tiled_copy_Oaccum.get_thread_slice(tidx);
Tensor tOsOaccum = gmem_thr_copy_Oaccum.partition_S(sOaccum); // ((Atom,AtomNum),ATOM_M,ATOM_N)
Tensor tOgOaccum = gmem_thr_copy_Oaccum.partition_D(gOaccum);
flash::sync_threads();
Tensor tOrOaccum = make_tensor<ElementO>(shape(tOgOaccum));
cute::copy(gmem_tiled_copy_Oaccum, tOsOaccum, tOrOaccum);
Tensor caccO = make_identity_tensor(Shape<Int<kBlockM>, Int<kHeadDimV>>{}); // (BLK_M,BLK_K) -> (blk_m,blk_k)
Tensor taccOcO = thr_mma_o.partition_C(caccO); // (MMA,MMA_M,MMA_K)
static_assert(decltype(size<0>(taccOcO))::value == 4);
// Convert to ((2, 2), MMA_M, MMA_K) then take only the row indices.
Tensor taccOcO_row = logical_divide(taccOcO, Shape<_4>{})(make_coord(0, _), _, 0);
CUTE_STATIC_ASSERT_V(size(lse) == size(taccOcO_row)); // MMA_M
if (get<1>(taccOcO_row(0)) == 0) {
#pragma unroll
for (int mi = 0; mi < size(lse); ++mi) {
const int row = get<0>(taccOcO_row(mi));
if (row < binfo.actual_seqlen_q - m_block * kBlockM) { gLSEaccum(row) = lse(mi); }
}
}
// Construct identity layout for sO
Tensor cO = make_identity_tensor(make_shape(size<0>(sOaccum), size<1>(sOaccum))); // (BLK_M,BLK_K) -> (blk_m,blk_k)
// Repeat the partitioning with identity layouts
Tensor tOcO = gmem_thr_copy_Oaccum.partition_D(cO); // (ACPY,ACPY_M,ACPY_K) -> (blk_m,blk_k)
// Clear_OOB_K must be false since we don't want to write zeros to gmem
flash::copy_reg_to_global<Is_even_MN, Is_even_K>(
tOrOaccum, tOgOaccum, tOcO, params.d_v, binfo.actual_seqlen_q - m_block * kBlockM
);
}
template<typename Kernel_traits, bool Is_causal, bool Is_even_MN, bool Is_even_K, bool Is_enable_dcp, typename Params>
__forceinline__ __device__ void compute_attn_1rowblock_splitkv_mla_k64_64x32_8waves_xcore1500(const Params &params, const int bidb, const int bidh, const int m_block, const int n_split_idx, const int n_block_min, const int n_block_max, const bool NoSplit) {
using Element = typename Kernel_traits::Element;
using ElementAccum = typename Kernel_traits::ElementAccum;
using index_t = typename Kernel_traits::index_t;
// Shared memory.
extern __shared__ char smem_[];
// The thread index.
const int tidx = threadIdx.x;
const int warp_idx = tidx / 64;
const int lane_idx = tidx % 64;
constexpr int kBlockM = Kernel_traits::kBlockM;
constexpr int kBlockN = Kernel_traits::kBlockN;
constexpr int kHeadDim = Kernel_traits::kHeadDim;
constexpr int kHeadDimV = Kernel_traits::kHeadDimV;
constexpr int kNWarps = Kernel_traits::kNWarps;
constexpr int kBlockKSmem = Kernel_traits::kBlockKSmem;
constexpr int kAtomLayoutMS = Kernel_traits::kAtomLayoutMS;
constexpr int kAtomLayoutMO = Kernel_traits::kAtomLayoutMO;
constexpr int Num_Stages = Kernel_traits::Num_Stages;
constexpr int kHeadDimNope = kHeadDimV;
constexpr int kHeadDimRope = kHeadDim - kHeadDimV;
static_assert(kBlockKSmem == 64);
static_assert(Kernel_traits::Share_Q_K_smem && Kernel_traits::Is_Q_in_regs);
const BlockInfo</*Varlen=*/!Is_even_MN> binfo(params, bidb);
if (m_block * kBlockM >= binfo.actual_seqlen_q) return;
// We iterate over the blocks in reverse order. This is because the last block is the only one
// that needs masking when we read K and V from global memory. Moreover, iterating in reverse
// might save us 1 register (we just need n_block instead of both n_block and n_block_max).
const index_t row_offset_q = binfo.q_offset(params.q_batch_stride, params.q_row_stride, bidb)
+ m_block * kBlockM * params.q_row_stride + bidh * params.q_head_stride;
// We move K and V to the last block.
const int bidb_cache = params.cache_batch_idx == nullptr ? bidb : params.cache_batch_idx[bidb];
const int *block_table = params.block_table == nullptr ? nullptr : params.block_table + bidb * params.block_table_batch_stride;
const int block_table_idx = block_table == nullptr ? 0 : (n_block_max - 1) * kBlockN / params.page_block_size;
const int block_table_offset = block_table == nullptr ? 0 : (n_block_max - 1) * kBlockN - block_table_idx * params.page_block_size;
const index_t row_offset_k = block_table == nullptr
? binfo.k_offset(params.k_batch_stride, params.k_row_stride, bidb_cache)
+ (n_block_max - 1) * kBlockN * params.k_row_stride + (bidh / params.h_h_k_ratio) * params.k_head_stride
: (bidh / params.h_h_k_ratio) * params.k_head_stride;
const index_t row_offset_v = block_table == nullptr
? binfo.k_offset(params.v_batch_stride, params.v_row_stride, bidb_cache)
+ (n_block_max - 1) * kBlockN * params.v_row_stride + (bidh / params.h_h_k_ratio) * params.v_head_stride
: (bidh / params.h_h_k_ratio) * params.v_head_stride;
Tensor gQ = make_tensor(make_gmem_ptr(reinterpret_cast<Element *>(params.q_ptr) + row_offset_q),
Shape<Int<kBlockM>, Int<kHeadDim>>{},
make_stride(params.q_row_stride, _1{}));
Tensor gK = make_tensor(make_gmem_ptr(reinterpret_cast<Element *>(params.k_ptr) + row_offset_k),
Shape<Int<kBlockN>, Int<kHeadDim>>{},
make_stride(params.k_row_stride, _1{}));
// if (threadIdx.x == 0 && blockIdx.y == 0 && blockIdx.z == 0) { printf("k_ptr = %p, row_offset_k = %d, gK_ptr = %p\n", params.k_ptr, row_offset_k, gK.data()); }
Tensor gV = make_tensor(make_gmem_ptr(reinterpret_cast<Element *>(params.v_ptr) + row_offset_v),
Shape<Int<kBlockN>, Int<kHeadDimV>>{},
make_stride(params.v_row_stride, _1{}));
Tensor sQ = make_tensor(make_smem_ptr(reinterpret_cast<Element *>(smem_)),
typename Kernel_traits::SmemLayoutQ{});
Tensor sQ_NoSwizzle = make_tensor(sQ.data(), typename Kernel_traits::SmemLayoutQNoSwizzle{});
Tensor sK = make_tensor(sQ.data(), typename Kernel_traits::SmemLayoutK242{});
Tensor sK_NoSwizzle = make_tensor(sK.data(), typename Kernel_traits::SmemLayoutKNoswizzle{});
Tensor sVt = make_tensor(sK.data(), typename Kernel_traits::SmemLayoutVtransposed242{});
Tensor sVtNoSwizzle = make_tensor(sVt.data(), typename Kernel_traits::SmemLayoutVtransposedNoSwizzle{});
Tensor sP = make_tensor(sK.data() + size(sK), typename Kernel_traits::SmemLayoutP{}); // sP is also bf16
Tensor sRowMax = make_tensor(make_smem_ptr(reinterpret_cast<ElementAccum *>(smem_) + size(sP) / 2 + size(sK) / 2), //sK is bf16, yet row_max is fp32
typename Kernel_traits::SmemLayoutRowMax{});
const int swz333_offset = cute::get_swizzle_offset<8, 3, 3, 3>(tidx);
typename Kernel_traits::GmemTiledCopyB128 gmem_tiled_copy_Q;
auto gmem_thr_copy_Q = gmem_tiled_copy_Q.get_thread_slice(tidx);
Tensor tQgQ = gmem_thr_copy_Q.partition_S(gQ);
tQgQ = make_tensor(tQgQ.data() + swz333_offset, layout(tQgQ));
Tensor tQsQ = gmem_thr_copy_Q.partition_D(sQ_NoSwizzle);
const int swz242_swap_offset = cute::get_swizzle_offset<8,2,4,2>(tidx) + ((tidx & 63) >= 32 ? ((tidx & 1) == 0 ? 8 : -8): 0);
typename Kernel_traits::GmemTiledCopyB128 gmem_tiled_copy_KV;
int tidx_load_k = tidx & 0xFF;
auto gmem_thr_copy_KV = gmem_tiled_copy_KV.get_thread_slice(tidx_load_k);
Tensor tKgK = gmem_thr_copy_KV.partition_S(gK); // (KCPY, KCPY_N, KCPY_K)
tKgK = make_tensor(tKgK.data() + swz242_swap_offset, layout(tKgK));
Tensor tKsK = gmem_thr_copy_KV.partition_D(sK_NoSwizzle);
typename Kernel_traits::TiledMmaS_16x16x32_4x2 tiled_mma_s;
auto thr_mma_s = tiled_mma_s.get_thread_slice(tidx);
Tensor tSrQ = thr_mma_s.partition_fragment_A(sQ_NoSwizzle); // (MMA,MMA_M,MMA_K)
Tensor tSrK = thr_mma_s.partition_fragment_B(sK_NoSwizzle(_, _, 0)); // (MMA,MMA_N,MMA_K)
typename Kernel_traits::TiledMmaO tiled_mma_o;
auto thr_mma_o = tiled_mma_o.get_thread_slice(tidx);
Tensor tOrVt = thr_mma_o.partition_fragment_B(sVtNoSwizzle(_, _, 0)); // (MMA, MMA_K,MMA_N)
auto smem_tiled_copy_S = make_tiled_copy_C(typename Kernel_traits::UniversalCopyAtomB64{}, tiled_mma_s);
auto smem_thr_copy_S = smem_tiled_copy_S.get_thread_slice(tidx);
Tensor tPsP = smem_thr_copy_S.partition_D(sP); // ((Atom,AtomNum),PIPE_M,PIPE_N)
auto smem_tiled_copy_P = make_tiled_copy_A(typename Kernel_traits::UniversalCopyAtomB64{}, tiled_mma_o);
auto smem_thr_copy_P = smem_tiled_copy_P.get_thread_slice(tidx);
Tensor tOsP = smem_thr_copy_P.partition_S(sP);
Tensor tOrP = thr_mma_o.partition_fragment_A(sP);
Tensor acc_o = partition_fragment_C(tiled_mma_o, Shape<Int<kBlockM>, Int<kHeadDimV>>{}); // MMA, MMA_M, MMA_K
//
// Copy Atom retiling
//
auto smem_tiled_copy_Q = make_tiled_copy_A(typename Kernel_traits::UniversalCopyAtomB128{}, tiled_mma_s);
auto smem_thr_copy_Q = smem_tiled_copy_Q.get_thread_slice(tidx);
Tensor tSsQ = smem_thr_copy_Q.partition_S(sQ);
const int swz242_diff_lds_b128 = (tidx & 7) >= 4 ? ((tidx & 31) < 16 ? 8 : -8) : 0;
auto smem_tiled_copy_K = make_tiled_copy_B(typename Kernel_traits::UniversalCopyAtomB128{}, tiled_mma_s);
auto smem_thr_copy_K = smem_tiled_copy_K.get_thread_slice(tidx);
Tensor tSsK = smem_thr_copy_K.partition_S(sK);
tSsK = make_tensor(tSsK.data() + swz242_diff_lds_b128, layout(tSsK));
const int swz242_diff_lds_trans_b128 = (tidx & 31) >= 16 ? ((tidx & 3) < 2 ? 8 : -8) : 0;
auto smem_tiled_copy_V = make_tiled_copy_B(typename Kernel_traits::LDSB64Trans4x16Atom{}, tiled_mma_o);
auto smem_thr_copy_V = smem_tiled_copy_V.get_thread_slice(tidx);
auto tOsVt = smem_thr_copy_V.partition_S(sVt);
tOsVt = make_tensor(tOsVt.data() + swz242_diff_lds_trans_b128, layout(tOsVt));
// PREDICATES
// Construct identity layout for sQ and sK
Tensor cQ = make_identity_tensor(make_shape(size<0>(sQ), size<1>(sQ))); // (BLK_M,BLK_K) -> (blk_m,blk_k)
Tensor cKV = make_identity_tensor(make_shape(size<0>(sK), size<1>(sK))); // (BLK_N,BLK_K) -> (blk_n,blk_k)
// Repeat the partitioning with identity layouts
Tensor tQcQ_noSwizzle = gmem_thr_copy_Q.partition_S(cQ); // (ACPY,ACPY_M,ACPY_K) -> (blk_m,blk_k)
Tensor tKVcKV_noSwizzle = gmem_thr_copy_KV.partition_S(cKV); // (BCPY,BCPY_N,BCPY_K) -> (blk_n,blk_k)
Tensor tKVcKV = make_tensor(tKVcKV_noSwizzle.data() + make_coord(0, swz242_swap_offset), layout(tKVcKV_noSwizzle));
Tensor tQcQ = make_tensor(tQcQ_noSwizzle.data() + make_coord(0, swz333_offset), layout(tQcQ_noSwizzle));
// Prologue
// We don't need to clear the sQ smem tiles since we'll only write out the valid outputs
flash::copy_b128_bsm_async<Is_even_MN, Is_even_K>(tQgQ, tQsQ, tQcQ, params.d, binfo.actual_seqlen_q - m_block * kBlockM);
flash::barrier_gvm<0>();
cute::copy(smem_tiled_copy_Q, tSsQ, tSrQ);
flash::sync_threads();
int n_block = n_block_max - 1;
uint32_t Ksmem_read_index = 0;
uint32_t Ksmem_write_index = 0;
uint32_t page_idx[size<1>(tKgK)];
uint32_t page_offset[size<1>(tKgK)];
// We don't need to clear the sK smem tiles since we'll mask out the scores anyway.
if (warp_idx < 4) {
flash::copy_page<Kernel_traits>(tKgK, page_idx, page_offset, n_block, block_table, params.page_block_size);
flash::copy_b128_page_bsm_async<Kernel_traits, Is_even_MN, Is_even_K>(gK, tKgK, tKsK(_, _, _, Ksmem_write_index), tKVcKV, params.d, n_block,
block_table, params.k_batch_stride, params.k_row_stride, params.page_block_size, page_idx, page_offset, swz242_swap_offset, binfo.actual_seqlen_k - n_block * kBlockN);
}
Ksmem_write_index ^= 1;
clear(acc_o);
flash::Softmax<size<1>(acc_o)> softmax;
flash::Mask<Is_causal, Is_enable_dcp> mask(binfo.actual_seqlen_k, binfo.actual_seqlen_q, params.ngroups, binfo.tot_seqlen_k, params.cp_world_size, params.cp_rank);
for (; n_block >= n_block_min; --n_block) {
Tensor acc_s = partition_fragment_C(tiled_mma_s, Shape<Int<kBlockM>, Int<kBlockN>>{}); // (MMA=4, MMA_M, MMA_N)
clear(acc_s);
flash::barrier_gvm<0>();
flash::gemm_prefetch_lds</*A_in_regs=*/Kernel_traits::Is_Q_in_regs, /*B_in_regs=*/false, /*prefetch_lds_num*/4>(
acc_s, tSrQ, tSrK, tSsQ, tSsK(_, _, _, Ksmem_read_index), tiled_mma_s, smem_tiled_copy_Q, smem_tiled_copy_K,
smem_thr_copy_Q, smem_thr_copy_K
);
/*
* NOTE:add this schedbound for separate gemmQK and gemmPV into two scheduling blocks
* bring opportunity to lds prefetching
*/
__builtin_mxc_schedbound_begin();
mask.template apply_mask<Is_causal, Is_even_MN>(
acc_s, n_block * kBlockN + (tidx / 64) / 4 * 16, m_block * kBlockM + (tidx / 64) % kAtomLayoutMS * 16 + (tidx & 0xf), kAtomLayoutMS * 16
);
Tensor scores_max_prev = make_fragment_like(softmax.row_max);
softmax.template get_row_max</*Is_first=*/false>(acc_s, scores_max_prev, sRowMax, params.scale_softmax_log2);
if (n_block > n_block_min && warp_idx < 4) {
// Advance gK
flash::copy_page<Kernel_traits>(tKgK, page_idx, page_offset, n_block - 1, block_table, params.page_block_size);
flash::copy_b128_page_bsm_async<Kernel_traits, /*Is_even_MN=*/true, Is_even_K>(gK, tKgK, tKsK(_, _, _, Ksmem_write_index), tKVcKV, params.d, n_block - 1,
block_table, params.k_batch_stride, params.k_row_stride, params.page_block_size, page_idx, page_offset, swz242_swap_offset);
Ksmem_write_index ^= 1;
}
softmax.template softmax_rescale_o_without_row_max</*Is_first=*/false, /*Check_inf=*/Is_causal || !Is_even_MN, true>(acc_s, acc_o, scores_max_prev, params.scale_softmax_log2);
CONVERT_TENSOR_TYPE(ElementAccum, Element, acc_s, rP)
cute::copy(smem_tiled_copy_S, rP, tPsP);
flash::sync_threads();
flash::gemm</*A_in_regs=*/false, /*B_in_regs=*/false>(
acc_o, tOrP, tOrVt, tOsP, tOsVt(_, _, _, Ksmem_read_index), tiled_mma_o, smem_tiled_copy_P, smem_tiled_copy_V,
smem_thr_copy_P, smem_thr_copy_V
);
Ksmem_read_index ^= 1;
}
// Epilogue
if (NoSplit) {
store_64x32_xcore1500<Kernel_traits, false, Is_even_MN, Is_even_K>(params, bidb, bidh, m_block, n_split_idx, smem_, acc_o, softmax, sRowMax);
}else{
store_64x32_xcore1500<Kernel_traits, true, Is_even_MN, Is_even_K>(params, bidb, bidh, m_block, n_split_idx, smem_, acc_o, softmax, sRowMax);
}
}
} // namespace flash

View File

@ -1,12 +0,0 @@
// Adapted from Dao-AILab/flash-attention (https://github.com/Dao-AILab/flash-attention/tree/v2.6.3)
#include <mctlass/numeric_types.h>
#include "run_mha.h"
#include "flash_fwd_dispatch_template.h"
void run_mha_fwd(mcFlashAttn::Flash_fwd_mla_params &params, cudaStream_t stream, bool force_split_kernel) {
constexpr int kHeadDim = 576;
mcFlashAttn::run_mha_fwd_splitkv_dispatch<kHeadDim>(params, stream);
}

View File

@ -4,4 +4,5 @@
#include "flash_mla.h"
void run_mha_fwd(mcFlashAttn::Flash_fwd_mla_params &params, cudaStream_t stream, bool force_split_kernel=false);
void run_mla_fwd(mcFlashAttn::Flash_fwd_mla_params &params, cudaStream_t stream);
void run_mla_fwd(SparsePrefillParams &params);

View File

@ -0,0 +1,24 @@
// Adapted from Dao-AILab/flash-attention (https://github.com/Dao-AILab/flash-attention/tree/v2.6.3)
#include <ATen/cuda/CUDAContext.h>
#include <mctlass/numeric_types.h>
#include "run_mla.h"
#include "static_switch.h"
#include "flash_fwd_dispatch_template.h"
void run_mla_fwd(mcFlashAttn::Flash_fwd_mla_params &params, cudaStream_t stream) {
constexpr int kHeadDim = 576;
ARCH_SWITCH(params.arch, kArch, [&] {
mcFlashAttn::run_mla_fwd_splitkv_dispatch<kHeadDim, kArch>(params, stream);
});
}
void run_mla_fwd(SparsePrefillParams &params) {
constexpr int kHeadDim = 576;
ARCH_SWITCH(params.arch, kArch, [&] {
mcFlashAttn::run_flash_mla_sparse_prefill_dispatch<kHeadDim, kArch>(params, params.stream);
});
}

View File

@ -0,0 +1,97 @@
#pragma once
#include "flash_mla.h"
#include "flash_dense_mla_decode_kernel.h"
#include "static_switch.h"
__global__ void __launch_bounds__(32, 1, 1)
get_mla_metadata_kernel(const GetDecodingMetadataParams params) {
int *seqlens_k_ptr = params.seqlens_k_ptr;
int *tile_scheduler_metadata_ptr = params.tile_scheduler_metadata_ptr;
int *num_splits_ptr = params.num_splits_ptr;
int batch_size = params.batch_size;
int block_size_n = params.block_size_n;
int fixed_overhead_num_blocks = params.fixed_overhead_num_blocks;
int num_sm_parts = params.num_sm_parts;
extern __shared__ int shared_mem[];
int* num_blocks_shared = shared_mem; // [batch_size]
int* num_splits_shared = shared_mem + batch_size; // [batch_size+1]
int* seqlens_k_shared = shared_mem + batch_size*2+1; // [batch_size]
int* first_block_idx_shared = shared_mem + batch_size*3+1; // [batch_size]
int* last_block_idx_shared = shared_mem + batch_size*4+1; // [batch_size]
int total_num_blocks = 0;
for (int i = threadIdx.x; i < batch_size; i += 32) {
int cur_s_k = params.topk == -1 ? __ldg(seqlens_k_ptr + i) : params.topk;
seqlens_k_shared[i] = cur_s_k;
int first_token_idx = 0;
int last_token_idx = max(cur_s_k-1, 0);
int cur_first_block_idx = first_token_idx / block_size_n;
int cur_last_block_idx = last_token_idx / block_size_n;
// NOTE Should attend to tokens [first_token_idx, last_token_idx], i.e. blocks [cur_first_block_idx, cur_last_block_idx]
// NOTE Before clamping, first_token_idx <= last_token_idx always holds, so after clamping, first_token_idx <= last_token_idx still holds.
// NOTE if seqlens_k is 0, then first_token_idx == last_token_idx == cur_first_block_idx == cur_last_block_idx == 0. So the sequence will have 1 block. We will correct this later in this kernel.
int num_blocks = cur_last_block_idx - cur_first_block_idx + 1;
total_num_blocks += num_blocks + fixed_overhead_num_blocks;
num_blocks_shared[i] = num_blocks;
first_block_idx_shared[i] = cur_first_block_idx;
last_block_idx_shared[i] = cur_last_block_idx;
}
for (int offset = 16; offset >= 1; offset /= 2) {
total_num_blocks += __shfl_xor_sync(uint32_t(-1), total_num_blocks, offset);
}
__syncwarp();
if (threadIdx.x == 0) {
int payload = mctlass::ceil_div(total_num_blocks, num_sm_parts) + fixed_overhead_num_blocks;
int now_idx = 0, now_block = 0, now_n_split_idx = 0, cum_num_splits = 0;
num_splits_shared[0] = 0;
for (int i = 0; i < num_sm_parts; ++i) {
int tile_scheduler_metadata0[4], tile_scheduler_metadata1;
tile_scheduler_metadata0[0] = now_idx;
tile_scheduler_metadata0[1] = now_block + first_block_idx_shared[now_idx];
tile_scheduler_metadata1 = now_n_split_idx;
int remain_payload = payload;
while (now_idx < batch_size) {
int num_blocks = num_blocks_shared[now_idx];
int now_remain_blocks = num_blocks - now_block;
if (remain_payload >= now_remain_blocks + fixed_overhead_num_blocks) {
cum_num_splits += now_n_split_idx + 1;
num_splits_shared[now_idx + 1] = cum_num_splits;
remain_payload -= now_remain_blocks + fixed_overhead_num_blocks;
++now_idx;
now_block = 0;
now_n_split_idx = 0;
} else {
if (remain_payload - fixed_overhead_num_blocks > 0) {
now_block += remain_payload - fixed_overhead_num_blocks;
++now_n_split_idx;
remain_payload = 0;
}
break;
}
}
tile_scheduler_metadata0[2] = now_block > 0 ? now_idx : now_idx - 1;
tile_scheduler_metadata0[3] = now_block > 0 ?
now_block + first_block_idx_shared[now_idx] : (seqlens_k_shared[now_idx-1] == 0 ?
0 : last_block_idx_shared[now_idx-1] + 1);
*reinterpret_cast<int4 *>(tile_scheduler_metadata_ptr + i * TileSchedulerMetaDataSize) = *reinterpret_cast<int4 *>(tile_scheduler_metadata0);
tile_scheduler_metadata_ptr[i * TileSchedulerMetaDataSize + 4] = tile_scheduler_metadata1;
}
FLASH_DEVICE_ASSERT(now_idx == batch_size && now_block == 0 && now_n_split_idx == 0);
}
__syncwarp();
for (int i = threadIdx.x; i <= batch_size; i += 32) {
num_splits_ptr[i] = num_splits_shared[i];
}
}
void run_get_mla_metadata_kernel(GetDecodingMetadataParams &params, cudaStream_t stream) {
int smem_size = sizeof(int) * (params.batch_size*5+1);
CUDA_CHECK(cudaFuncSetAttribute(get_mla_metadata_kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size));
get_mla_metadata_kernel<<<1, 32, smem_size, stream>>>(params);
CUDA_KERNEL_LAUNCH_CHECK();
}

View File

@ -0,0 +1,18 @@
// Adapted from Dao-AILab/flash-attention (https://github.com/Dao-AILab/flash-attention/tree/v2.6.3)
#include "flash_mla.h"
#include "flash_fwd_run_template.h"
#include <mctlass/numeric_types.h>
template void run_flash_splitkv_fwd_mla_template<
576,
16,
16,
4,
true,
true,
mctlass::bfloat16_t,
true,
512,
1
>(Flash_fwd_mla_params &params, cudaStream_t stream);

View File

@ -1,17 +1,17 @@
// Adapted from Dao-AILab/flash-attention (https://github.com/Dao-AILab/flash-attention/tree/v2.6.3)
#include "flash_mla.h"
#include "flash_run_fwd_template_impl.h"
#include "flash_fwd_run_template.h"
#include <mctlass/numeric_types.h>
template void run_flash_splitkv_fwd_template<
template void run_flash_splitkv_fwd_mla_template<
576,
16,
16,
4,
true,
true,
cutlass::bfloat16_t,
mctlass::bfloat16_t,
true,
512,
2

View File

@ -1,18 +1,18 @@
// Adapted from Dao-AILab/flash-attention (https://github.com/Dao-AILab/flash-attention/tree/v2.6.3)
#include "flash_mla.h"
#include "flash_run_fwd_template_impl.h"
#include "flash_fwd_run_template.h"
#include <mctlass/numeric_types.h>
template void run_flash_splitkv_fwd_template<
template void run_flash_splitkv_fwd_mla_template<
576,
16,
16,
4,
true,
true,
cutlass::half_t,
false,
mctlass::half_t,
true,
512,
2
1
>(Flash_fwd_mla_params &params, cudaStream_t stream);

View File

@ -1,17 +1,17 @@
// Adapted from Dao-AILab/flash-attention (https://github.com/Dao-AILab/flash-attention/tree/v2.6.3)
#include "flash_mla.h"
#include "flash_run_fwd_template_impl.h"
#include "flash_fwd_run_template.h"
#include <mctlass/numeric_types.h>
template void run_flash_splitkv_fwd_template<
template void run_flash_splitkv_fwd_mla_template<
576,
16,
16,
4,
true,
true,
cutlass::half_t,
mctlass::half_t,
true,
512,
2

View File

@ -1,17 +1,17 @@
// Adapted from Dao-AILab/flash-attention (https://github.com/Dao-AILab/flash-attention/tree/v2.6.3)
#include "flash_mla.h"
#include "flash_run_fwd_template_impl.h"
#include "flash_fwd_run_template.h"
#include <mctlass/numeric_types.h>
template void run_flash_splitkv_fwd_template<
template void run_flash_splitkv_fwd_mla_template<
576,
32,
16,
4,
true,
true,
cutlass::bfloat16_t,
mctlass::bfloat16_t,
true,
512,
2

View File

@ -1,17 +1,17 @@
// Adapted from Dao-AILab/flash-attention (https://github.com/Dao-AILab/flash-attention/tree/v2.6.3)
#include "flash_mla.h"
#include "flash_run_fwd_template_impl.h"
#include "flash_fwd_run_template.h"
#include <mctlass/numeric_types.h>
template void run_flash_splitkv_fwd_template<
template void run_flash_splitkv_fwd_mla_template<
576,
32,
16,
4,
true,
true,
cutlass::half_t,
mctlass::half_t,
true,
512,
2

View File

@ -1,17 +1,17 @@
// Adapted from Dao-AILab/flash-attention (https://github.com/Dao-AILab/flash-attention/tree/v2.6.3)
#include "flash_mla.h"
#include "flash_run_fwd_template_impl.h"
#include "flash_fwd_run_template.h"
#include <mctlass/numeric_types.h>
template void run_flash_splitkv_fwd_template<
template void run_flash_splitkv_fwd_mla_template<
576,
64,
16,
8,
true,
true,
cutlass::bfloat16_t,
mctlass::bfloat16_t,
true,
512,
2

View File

@ -1,17 +1,17 @@
// Adapted from Dao-AILab/flash-attention (https://github.com/Dao-AILab/flash-attention/tree/v2.6.3)
#include "flash_mla.h"
#include "flash_run_fwd_template_impl.h"
#include "flash_fwd_run_template.h"
#include <mctlass/numeric_types.h>
template void run_flash_splitkv_fwd_template<
template void run_flash_splitkv_fwd_mla_template<
576,
64,
16,
8,
true,
true,
cutlass::half_t,
mctlass::half_t,
true,
512,
2

View File

@ -0,0 +1,18 @@
// Adapted from Dao-AILab/flash-attention (https://github.com/Dao-AILab/flash-attention/tree/v2.6.3)
#include "flash_mla.h"
#include "flash_fwd_run_template.h"
#include <mctlass/numeric_types.h>
template void run_flash_mla_sparse_prefill_template<
576,
64,
16,
8,
true,
true,
mctlass::bfloat16_t,
true,
512,
2
>(SparsePrefillParams &params, cudaStream_t stream);

View File

@ -1,18 +1,18 @@
// Adapted from Dao-AILab/flash-attention (https://github.com/Dao-AILab/flash-attention/tree/v2.6.3)
#include "flash_mla.h"
#include "flash_run_fwd_template_impl.h"
#include "flash_fwd_run_template.h"
#include <mctlass/numeric_types.h>
template void run_flash_splitkv_fwd_template<
template void run_flash_splitkv_fwd_sparse_mla_template<
576,
64,
16,
8,
true,
true,
cutlass::half_t,
false,
mctlass::bfloat16_t,
true,
512,
2
>(Flash_fwd_mla_params &params, cudaStream_t stream);
>(Flash_fwd_mla_params &params, cudaStream_t stream);

View File

@ -1,18 +1,18 @@
// Adapted from Dao-AILab/flash-attention (https://github.com/Dao-AILab/flash-attention/tree/v2.6.3)
#include "flash_mla.h"
#include "flash_run_fwd_template_impl.h"
#include "flash_fwd_run_template.h"
#include <mctlass/numeric_types.h>
template void run_flash_splitkv_fwd_template<
template void run_flash_splitkv_fwd_sparse_mla_template<
576,
64,
16,
8,
true,
true,
cutlass::bfloat16_t,
false,
mctlass::half_t,
true,
512,
2
>(Flash_fwd_mla_params &params, cudaStream_t stream);
>(Flash_fwd_mla_params &params, cudaStream_t stream);

View File

@ -0,0 +1,19 @@
// Adapted from Dao-AILab/flash-attention (https://github.com/Dao-AILab/flash-attention/tree/v2.6.3)
#include "flash_mla.h"
#include "flash_fwd_run_template.h"
#include <mctlass/numeric_types.h>
template void run_flash_splitkv_fwd_mla_template<
576,
64,
16,
8,
true,
true,
mctlass::bfloat16_t,
true,
512,
2,
Arch::xcore1500
>(Flash_fwd_mla_params &params, cudaStream_t stream);

View File

@ -0,0 +1,19 @@
// Adapted from Dao-AILab/flash-attention (https://github.com/Dao-AILab/flash-attention/tree/v2.6.3)
#include "flash_mla.h"
#include "flash_fwd_run_template.h"
#include <mctlass/numeric_types.h>
template void run_flash_splitkv_fwd_mla_template<
576,
64,
16,
8,
true,
true,
mctlass::half_t,
true,
512,
2,
Arch::xcore1500
>(Flash_fwd_mla_params &params, cudaStream_t stream);

View File

@ -0,0 +1,19 @@
// Adapted from Dao-AILab/flash-attention (https://github.com/Dao-AILab/flash-attention/tree/v2.6.3)
#include "flash_mla.h"
#include "flash_fwd_run_template.h"
#include <mctlass/numeric_types.h>
template void run_flash_splitkv_fwd_mla_template<
576,
64,
32,
8,
true,
true,
mctlass::bfloat16_t,
true,
512,
2,
Arch::xcore1500
>(Flash_fwd_mla_params &params, cudaStream_t stream);

View File

@ -0,0 +1,19 @@
// Adapted from Dao-AILab/flash-attention (https://github.com/Dao-AILab/flash-attention/tree/v2.6.3)
#include "flash_mla.h"
#include "flash_fwd_run_template.h"
#include <mctlass/numeric_types.h>
template void run_flash_splitkv_fwd_mla_template<
576,
64,
32,
8,
true,
true,
mctlass::half_t,
true,
512,
2,
Arch::xcore1500
>(Flash_fwd_mla_params &params, cudaStream_t stream);

View File

@ -0,0 +1,19 @@
// Adapted from Dao-AILab/flash-attention (https://github.com/Dao-AILab/flash-attention/tree/v2.6.3)
#include "flash_mla.h"
#include "flash_fwd_run_template.h"
#include <mctlass/numeric_types.h>
template void run_flash_mla_sparse_prefill_template<
576,
64,
16,
8,
true,
true,
mctlass::bfloat16_t,
true,
512,
2,
Arch::xcore1500
>(SparsePrefillParams &params, cudaStream_t stream);

View File

@ -0,0 +1,19 @@
// Adapted from Dao-AILab/flash-attention (https://github.com/Dao-AILab/flash-attention/tree/v2.6.3)
#include "flash_mla.h"
#include "flash_fwd_run_template.h"
#include <mctlass/numeric_types.h>
template void run_flash_splitkv_fwd_sparse_mla_template<
576,
64,
16,
8,
true,
true,
mctlass::bfloat16_t,
true,
512,
2,
Arch::xcore1500
>(Flash_fwd_mla_params &params, cudaStream_t stream);

View File

@ -0,0 +1,19 @@
// Adapted from Dao-AILab/flash-attention (https://github.com/Dao-AILab/flash-attention/tree/v2.6.3)
#include "flash_mla.h"
#include "flash_fwd_run_template.h"
#include <mctlass/numeric_types.h>
template void run_flash_splitkv_fwd_sparse_mla_template<
576,
64,
16,
8,
true,
true,
mctlass::half_t,
true,
512,
2,
Arch::xcore1500
>(Flash_fwd_mla_params &params, cudaStream_t stream);

42
csrc/mctlass/.gitignore vendored Normal file
View File

@ -0,0 +1,42 @@
# Compiled Object files
*.slo
*.lo
*.o
*.obj
# Precompiled Headers
*.gch
*.pch
# Compiled Dynamic libraries
*.so
*.dylib
*.dll
# Fortran module files
*.mod
# Compiled Static libraries
*.lai
*.la
*.a
*.lib
# Executables
*.exe
*.out
*.app
# vim tags
tags
.tags
.*.swp
# Visual Studio Code
.vscode
# install.sh build dir
build
# PyCache files
__pycache__

View File

@ -0,0 +1,79 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#pragma once
#include <cute/config.hpp>
#include <cute/tensor.hpp>
namespace cute
{
//
// Accept mutable temporaries
//
template <class Alpha,
class XEngine, class XLayout,
class Beta,
class YEngine, class YLayout>
CUTE_HOST_DEVICE
void
axpby(Alpha const& alpha,
Tensor<XEngine, XLayout> const& x,
Beta const& beta,
Tensor<YEngine, YLayout> && y)
{
return axpby(alpha, x, beta, y);
}
//
// AXPBY
//
template <class Alpha,
class XEngine, class XLayout,
class Beta,
class YEngine, class YLayout>
CUTE_HOST_DEVICE
void
axpby(Alpha const& alpha,
Tensor<XEngine, XLayout> const& x,
Beta const& beta,
Tensor<YEngine, YLayout> & y)
{
auto isBetaZero = (beta == Int<0>{});
CUTE_UNROLL
for (int i = 0; i < size(x); ++i) {
y(i) = (isBetaZero ? alpha * x(i) : alpha * x(i) + beta * y(i));
}
}
} // end namespace cute

View File

@ -0,0 +1,66 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#pragma once
#include <cute/config.hpp>
#include <cute/tensor.hpp>
#include <cute/algorithm/fill.hpp>
namespace cute
{
//
// Accept mutable temporaries
//
template <class Engine, class Layout>
CUTE_HOST_DEVICE
void
clear(Tensor<Engine, Layout>&& tensor)
{
return clear(tensor);
}
//
// Set elements to zero
//
template <class Engine, class Layout>
CUTE_HOST_DEVICE
void
clear(Tensor<Engine, Layout>& tensor)
{
using T = typename Tensor<Engine,Layout>::value_type;
fill(tensor, T{});
}
} // end namespace cute

View File

@ -0,0 +1,523 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#pragma once
#include <cute/config.hpp>
#include <cute/tensor.hpp>
#include <cute/tensor_predicate.hpp>
#include <cute/atom/copy_atom.hpp>
namespace cute
{
//
// Accept mutable temporaries
//
template <class PrdTensor,
class SrcEngine, class SrcLayout,
class DstEngine, class DstLayout>
CUTE_HOST_DEVICE
void
copy_if(PrdTensor const& pred,
Tensor<SrcEngine, SrcLayout> const& src,
Tensor<DstEngine, DstLayout> && dst)
{
return copy_if(pred, src, dst);
}
template <class... CopyArgs,
class PrdTensor,
class SrcEngine, class SrcLayout,
class DstEngine, class DstLayout>
CUTE_HOST_DEVICE
void
copy_if(Copy_Atom<CopyArgs...> const& copy_atom,
PrdTensor const& pred,
Tensor<SrcEngine, SrcLayout> const& src,
Tensor<DstEngine, DstLayout> && dst)
{
return copy_if(copy_atom, pred, src, dst);
}
template <class VecType,
class SrcEngine, class SrcLayout,
class DstEngine, class DstLayout>
CUTE_HOST_DEVICE
void
copy_vec(Tensor<SrcEngine, SrcLayout> const& src,
Tensor<DstEngine, DstLayout> && dst)
{
return copy_vec<VecType>(src, dst);
}
template <class SrcEngine, class SrcLayout,
class DstEngine, class DstLayout>
CUTE_HOST_DEVICE
void
copy(Tensor<SrcEngine, SrcLayout> const& src,
Tensor<DstEngine, DstLayout> && dst)
{
return copy(src, dst);
}
template <class... CopyArgs,
class SrcEngine, class SrcLayout,
class DstEngine, class DstLayout>
CUTE_HOST_DEVICE
void
copy(Copy_Atom<CopyArgs...> const& copy_atom,
Tensor<SrcEngine, SrcLayout> const& src,
Tensor<DstEngine, DstLayout> && dst)
{
return copy(copy_atom, src, dst);
}
//
// copy_if -- Predicated Copy
//
template <class PrdTensor,
class SrcEngine, class SrcLayout,
class DstEngine, class DstLayout>
CUTE_HOST_DEVICE
void
copy_if(PrdTensor const& pred,
Tensor<SrcEngine, SrcLayout> const& src,
Tensor<DstEngine, DstLayout> & dst)
{
auto copy_op = select_elementwise_copy(src, dst);
CUTE_UNROLL
for (int i = 0; i < size(src); ++i) {
if (pred(i)) {
copy_op.copy(src(i), dst(i));
}
}
}
//
// copy_if -- Predicated CopyAtom
//
template <class... CopyArgs,
class PredTensor,
class SrcEngine, class SrcLayout,
class DstEngine, class DstLayout>
CUTE_HOST_DEVICE
void
copy_if(Copy_Atom<CopyArgs...> const& copy_atom,
PredTensor const& pred, // (Rest...)
Tensor<SrcEngine, SrcLayout> const& src, // (V,Rest...)
Tensor<DstEngine, DstLayout> & dst) // (V,Rest...)
{
static_assert(SrcLayout::rank == DstLayout::rank, "CopyAtom rank-mismatch.");
if constexpr (SrcLayout::rank == 1) { // Dispatch the copy
copy_atom.call(src, dst);
} else { // Loop over all but the first mode
constexpr int R = SrcLayout::rank;
auto src_v = group_modes<1,R>(src);
auto dst_v = group_modes<1,R>(dst);
CUTE_UNROLL
for (int i = 0; i < size<1>(src_v); ++i) {
if (pred(i)) {
copy_atom.call(src_v(_,i), dst_v(_,i));
}
}
}
}
//
// copy_vec -- attempt vectorized copy with VecType
//
template <class VecType,
class SrcEngine, class SrcLayout,
class DstEngine, class DstLayout>
CUTE_HOST_DEVICE
void
copy_vec(Tensor<SrcEngine, SrcLayout> const& src,
Tensor<DstEngine, DstLayout> & dst)
{
using SrcType = typename SrcEngine::value_type;
using DstType = typename DstEngine::value_type;
if constexpr (sizeof(SrcType) == sizeof(DstType) && sizeof(VecType) > sizeof(DstType))
{
/* @pre is_aligned<N>(src.data()) &&
* is_aligned<N>(dst.data())
*/
auto src_v = recast<VecType const>(src);
auto dst_v = recast<VecType >(dst);
#if 0
if (thread0()) {
print("copy_vec -- vectorizing copy from %3db to %3db\n", int(8*sizeof(SrcType)), int(8*sizeof(VecType)));
print(" "); print(layout(src)); print(" => "); print(layout(src_v)); print("\n");
print(" "); print(layout(dst)); print(" => "); print(layout(dst_v)); print("\n");
}
#endif
return copy_if(TrivialPredTensor{}, src_v, dst_v);
} else {
#if 0
if (thread0()) {
print("copy_vec -- not vectorizing, copy with %3db and %3db\n", int(8*sizeof(SrcType)), int(8*sizeof(DstType)));
print(" "); print(layout(src)); print("\n");
print(" "); print(layout(dst)); print("\n");
}
#endif
return copy_if(TrivialPredTensor{}, src, dst);
}
}
//
// copy -- auto-vectorizing copy
//
template <class SrcEngine, class SrcLayout,
class DstEngine, class DstLayout>
CUTE_HOST_DEVICE
void
copy(Tensor<SrcEngine, SrcLayout> const& src,
Tensor<DstEngine, DstLayout> & dst)
{
constexpr int N = decltype(max_common_vector(src, dst))::value;
#if 0
if (thread0()) {
print("copy -- found a max_common_vector of %d\n", N);
print(" "); print(src.data()); print(" o "); print(layout(src)); print("\n");
print(" "); print(dst.data()); print(" o "); print(layout(dst)); print("\n");
}
#endif
if constexpr (N <= 1) {
return copy_if(TrivialPredTensor{}, src, dst);
} else {
constexpr int vec_bits = N * sizeof_bits<typename SrcEngine::value_type>::value;
using VecType = uint_bit_t<cute::min(128, vec_bits)>;
return copy_vec<VecType>(src, dst);
}
}
//
// copy -- CopyAtom
//
template <class... CopyArgs,
class SrcEngine, class SrcLayout,
class DstEngine, class DstLayout>
CUTE_HOST_DEVICE
void
copy(Copy_Atom<CopyArgs...> const& copy_atom,
Tensor<SrcEngine, SrcLayout> const& src,
Tensor<DstEngine, DstLayout> & dst)
{
return copy_if(copy_atom, TrivialPredTensor{}, src, dst);
}
template <class... CopyArgs,
class SrcEngine, class SrcLayout,
class DstEngine, class DstLayout>
CUTE_HOST_DEVICE
void
copy(Copy_Atom<DefaultCopy, CopyArgs...> const&,
Tensor<SrcEngine, SrcLayout> const& src,
Tensor<DstEngine, DstLayout> & dst)
{
return copy(src, dst);
}
#if defined(__MERGE_LDS_B32)
CUTE_DEVICE
void reg_trans(uint32_t &a) {
/* ************************************************************
** tmp_0[n]=a[(n&0x3c)+shfl[n%4]]
** 0x0b1 means shfl[0]=1, shfl[1]=0, shfl[2]=3, shfl[3]=2
** and tmp_0[n]/a[n] means the value of tmp_0/a while lane_id=n
* ************************************************************/
auto tmp_0 = __builtin_mxc_mov_raw_shfl(a, 0x0b1, 0xf, 0xf, false);
auto tmp_1 = __builtin_mxc_byte_perm(a, tmp_0, 0x07060302);
a = __builtin_mxc_byte_perm(tmp_0, a, 0x05040100);
if (__lane_id() & 0x1) {
a = tmp_1;
}
}
CUTE_DEVICE
void reg_trans(uint32_t &a, uint32_t &b) {
reg_trans(a);
reg_trans(b);
}
template <
class SrcEngine, class SrcLayout,
class DstEngine, class DstLayout>
CUTE_DEVICE
void copy_trans(Tensor<SrcEngine, SrcLayout> const && src,
Tensor<DstEngine, DstLayout> && dst,
const uint32_t &src_stride,
const uint32_t &dst_stride,
const uint32_t *cpy_offset)
{
auto dst_ptr = reinterpret_cast<uint32_t *>(dst.data());
auto src_addr = reinterpret_cast<uint64_t>(src.data().ptr_);
src_addr = src_addr - cpy_offset[8];
/* *************************************************
** The address attribute of src_addr has benn destoried,
** So we need to use __attribute__((address_space (3)))
* *************************************************/
uint32_t __attribute__((address_space(3))) *src_ptr[8];
CUTE_UNROLL
for (uint32_t i = 0; i < 4; ++i) {
src_ptr[2 * i] = (uint32_t __attribute__((address_space(3))) *)(src_addr) + cpy_offset[2 * i];
src_ptr[2 * i + 1] = (uint32_t __attribute__((address_space(3))) *)(src_addr) + cpy_offset[2 * i + 1];
}
CUTE_UNROLL
for (uint32_t i = 0; i < size(dst) / 4; ++i) {
dst_ptr[i] = src_ptr[i][0];
}
CUTE_UNROLL
for (uint32_t i = 0; i < 8; ++i) {
src_ptr[i] = src_ptr[i] + src_stride;
}
dst_ptr = dst_ptr + dst_stride;
CUTE_UNROLL
for (uint32_t i = 0; i < size(dst) / 2 - size(dst) / 4; ++i) {
dst_ptr[i] = src_ptr[i][0];
}
}
#elif defined(__MERGE_LDS_B64)
CUTE_DEVICE
void reg_trans(uint32_t &a, uint32_t &b) {
const int laneId = __lane_id();
/* ************************************************************
** tmp_0[n]=a[(n&0x3c)+shfl[n%4]]
** 0x0b1 means shfl[0]=1, shfl[1]=0, shfl[2]=3, shfl[3]=2
** and tmp_0[n]/a[n] means the value of tmp_0/a while lane_id=n
* ************************************************************/
auto tmp_0 = __builtin_mxc_mov_raw_shfl(a, 0x0b1, 0xf, 0xf, false);
auto tmp_1 = __builtin_mxc_byte_perm(a, tmp_0, 0x07060302);
a = __builtin_mxc_byte_perm(tmp_0, a, 0x05040100);
auto tmp_2 = __builtin_mxc_mov_raw_shfl(b, 0x0b1, 0xf, 0xf, false);
auto tmp_3 = __builtin_mxc_byte_perm(b, tmp_2, 0x07060302);
b = __builtin_mxc_byte_perm(tmp_2, b, 0x05040100);
if (laneId & 0x1) {
a = tmp_1;
b = tmp_3;
}
/* ************************************************************
** tmp_0[n]=a[(n&0x3c)+shfl[n%4]]
** 0x04e means shfl[0]=2, shfl[1]=3, shfl[2]=0, shfl[3]=1
** and tmp_0[n]/a[n] means the value of tmp_0/a while lane_id=n
* ************************************************************/
tmp_0 = __builtin_mxc_mov_raw_shfl(a, 0x04e, 0xf, 0xf, false);
tmp_1 = __builtin_mxc_mov_raw_shfl(b, 0x04e, 0xf, 0xf, false);
if ((laneId & 0x3) >> 1) {
a = tmp_1;
}
else {
b = tmp_0;
}
}
template <
class SrcEngine, class SrcLayout,
class DstEngine, class DstLayout>
CUTE_DEVICE
void
copy_trans(
Tensor<SrcEngine, SrcLayout> const&& src,
Tensor<DstEngine, DstLayout> && dst,
const int src_stride,
const int dst_stride,
const uint32_t *cpy_offset)
{
auto dst_ptr = reinterpret_cast<uint32_t *>(dst.data());
auto src_addr = reinterpret_cast<uint64_t const>(src.data().ptr_);
src_addr = src_addr - cpy_offset[4];
/* ************************************************
** The address attribute of src_addr has benn destoried
** So we need to use __attribute__((address_space (3)))
* *************************************************/
uint32_t __attribute__((address_space (3))) *src_ptr[4];
CUTE_UNROLL
for (int i = 0; i < 4; ++i) {
src_ptr[i] = (uint32_t __attribute__((address_space(3))) *)(src_addr) + cpy_offset[i];
}
CUTE_UNROLL
for (int i = 0; i < size(dst) / 8; ++i) {
dst_ptr[2 * i] = src_ptr[i][0];
dst_ptr[2 * i + 1] = src_ptr[i][1];
}
CUTE_UNROLL
for (int i = 0; i < 4; ++i) {
src_ptr[i] = src_ptr[i] + src_stride;
}
dst_ptr = dst_ptr + dst_stride;
CUTE_UNROLL
for (int i = 0; i < size(dst) / 4 - size(dst) / 8; ++i) {
dst_ptr[2 * i] = src_ptr[i][0];
dst_ptr[2 * i + 1] = src_ptr[i][1];
}
}
#endif
#if defined(__MERGE_LDS_B32) || defined(__MERGE_LDS_B64)
template<class DstEngine, class DstLayout>
CUTE_DEVICE
void tensor_trans(Tensor<DstEngine, DstLayout> && dst,
const uint32_t stride) {
auto dst_ptr = reinterpret_cast<uint32_t *>(dst.data());
CUTE_UNROLL
for (uint32_t i = 0; i < size(dst) / 8; ++i) {
reg_trans(dst_ptr[2 * i], dst_ptr[2 * i + 1]);
}
dst_ptr = dst_ptr + stride;
CUTE_UNROLL
for (uint32_t i = 0; i < size(dst) / 4 - size(dst) / 8; ++i) {
reg_trans(dst_ptr[2 * i], dst_ptr[2 * i + 1]);
}
}
#endif
template <class SrcEngine, class SrcLayout>
CUTE_HOST_DEVICE
void
copy_global_to_reg(
Tensor<SrcEngine, SrcLayout> const&& src,
uint32_t *dst)
{
typedef __NATIVE_VECTOR__(4, int) VecType;
auto src_ptr = (VecType *)(src.data().ptr_);
auto dst_ptr = (VecType *)(dst);
dst_ptr[0] = __builtin_mxc_load_global_async128(src_ptr);
}
template <class DstEngine, class DstLayout>
CUTE_HOST_DEVICE
void
copy_reg_to_share(
uint32_t *src_ptr,
Tensor<DstEngine, DstLayout> && dst)
{
auto dst_ptr = reinterpret_cast<uint32_t *>(dst.data().ptr_);
dst_ptr[0] = src_ptr[0];
dst_ptr[1] = src_ptr[1];
dst_ptr[2] = src_ptr[2];
dst_ptr[3] = src_ptr[3];
}
//////////////////////////////////////////
// Special Auto-Vectorizing Overloads
//////////////////////////////////////////
#if defined(CUTE_COPY_ATOM_TMA_SM90_ENABLED)
template <class... CT_Args, class... CA_Args,
class SrcEngine, class SrcLayout,
class DstEngine, class DstLayout>
CUTE_HOST_DEVICE
void
copy(Copy_Atom<Copy_Traits<SM90_BULK_COPY_AUTO, CT_Args...>, CA_Args...> const& atom,
Tensor<SrcEngine, SrcLayout> const& src,
Tensor<DstEngine, DstLayout> & dst)
{
using SrcType = typename SrcEngine::value_type;
using DstType = typename DstEngine::value_type;
static_assert(sizeof_bits<SrcType>::value == sizeof_bits<DstType>::value);
static_assert((is_gmem<SrcEngine>::value && is_smem<DstEngine>::value) ||
(is_smem<SrcEngine>::value && is_gmem<DstEngine>::value),
"Bulk Copy only supports gmem -> smem or smem -> gmem movement.");
// Do BulkCopy dispatch
using BULK_COPY_OP = conditional_t<is_gmem<SrcEngine>::value,
SM90_BULK_COPY_G2S,
SM90_BULK_COPY_S2G>;
constexpr int N = decltype(max_common_vector(src, dst))::value;
// Construct a new concrete Atom of the vector size
using N_BITS = Int<N*sizeof_bits<SrcType>::value>;
using COPY_ATOM = Copy_Atom<Copy_Traits<BULK_COPY_OP, N_BITS, CT_Args...>, SrcType>;
auto bulk_atom = apply(atom.opargs_, [&](auto const&... args) { return COPY_ATOM{args...}; });
// Tile the src and dst to the Atom
auto tiler = right_inverse(dst.layout()).compose(Int<N>{});
#if 0
if (thread0()) {
print("copy -- found a max_common_vector of %d\n", N);
print(" "); print(src.data()); print(" o "); print(layout(src)); print("\n");
print(" "); print(dst.data()); print(" o "); print(layout(dst)); print("\n");
}
#endif
return copy(bulk_atom, logical_divide(src, tiler), logical_divide(dst, tiler));
}
#endif // #if defined(CUTE_COPY_ATOM_TMA_SM90_ENABLED)
} // end namespace cute

View File

@ -0,0 +1,87 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#pragma once
#include <cute/config.hpp>
#include <cute/tensor.hpp>
#include <cute/algorithm/prefer.hpp>
namespace cute
{
//
// Accept mutable temporaries
//
template <class Engine, class Layout, class T>
CUTE_HOST_DEVICE
void
fill(Tensor<Engine, Layout>&& tensor, T const& value)
{
return fill(tensor, value);
}
namespace detail
{
// Prefer fill(tensor.data(), value), if possible
template <class Engine, class Layout, class T>
CUTE_HOST_DEVICE
auto
fill(Tensor<Engine, Layout>& tensor, T const& value, prefer<1>)
-> decltype(fill(tensor.data(), value))
{
fill(tensor.data(), value);
}
// Default implementation
template <class Engine, class Layout, class T>
CUTE_HOST_DEVICE
void
fill(Tensor<Engine, Layout>& tensor, T const& value, prefer<0>)
{
CUTE_UNROLL
for (int i = 0; i < size(tensor); ++i) {
tensor(i) = value;
}
}
} // end namespace detail
template <class Engine, class Layout, class T>
CUTE_HOST_DEVICE
void
fill(Tensor<Engine, Layout>& tensor, T const& value)
{
return detail::fill(tensor, value, prefer<1>{});
}
} // end namespace cute

View File

@ -0,0 +1,198 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#pragma once
#include <cute/config.hpp>
#include <cute/util/type_traits.hpp>
/** C++14 <functional> extensions */
namespace cute {
/**************/
/** Identity **/
/**************/
struct identity {
template <class T>
CUTE_HOST_DEVICE constexpr
decltype(auto) operator()(T&& arg) const {
return std::forward<T>(arg);
}
};
template <class R>
struct constant_fn {
template <class... T>
CUTE_HOST_DEVICE constexpr
decltype(auto) operator()(T&&...) const {
return r_;
}
R r_;
};
/***********/
/** Unary **/
/***********/
#define CUTE_LEFT_UNARY_OP(NAME,OP) \
struct NAME { \
template <class T> \
CUTE_HOST_DEVICE constexpr \
decltype(auto) operator()(T&& arg) const { \
return OP std::forward<T>(arg); \
} \
}
#define CUTE_RIGHT_UNARY_OP(NAME,OP) \
struct NAME { \
template <class T> \
CUTE_HOST_DEVICE constexpr \
decltype(auto) operator()(T&& arg) const { \
return std::forward<T>(arg) OP ; \
} \
}
#define CUTE_NAMED_UNARY_OP(NAME,OP) \
struct NAME { \
template <class T> \
CUTE_HOST_DEVICE constexpr \
decltype(auto) operator()(T&& arg) const { \
return OP (std::forward<T>(arg)); \
} \
}
CUTE_LEFT_UNARY_OP(unary_plus, +);
CUTE_LEFT_UNARY_OP(negate, -);
CUTE_LEFT_UNARY_OP(bit_not, ~);
CUTE_LEFT_UNARY_OP(logical_not, !);
CUTE_LEFT_UNARY_OP(dereference, *);
CUTE_LEFT_UNARY_OP(address_of, &);
CUTE_LEFT_UNARY_OP(pre_increment, ++);
CUTE_LEFT_UNARY_OP(pre_decrement, --);
CUTE_RIGHT_UNARY_OP(post_increment, ++);
CUTE_RIGHT_UNARY_OP(post_decrement, --);
CUTE_NAMED_UNARY_OP(abs_fn, abs);
CUTE_NAMED_UNARY_OP(conjugate, cute::conj);
#undef CUTE_LEFT_UNARY_OP
#undef CUTE_RIGHT_UNARY_OP
#undef CUTE_NAMED_UNARY_OP
/************/
/** Binary **/
/************/
#define CUTE_BINARY_OP(NAME,OP) \
struct NAME { \
template <class T, class U> \
CUTE_HOST_DEVICE constexpr \
decltype(auto) operator()(T&& lhs, U&& rhs) const { \
return std::forward<T>(lhs) OP std::forward<U>(rhs); \
} \
}
#define CUTE_NAMED_BINARY_OP(NAME,OP) \
struct NAME { \
template <class T, class U> \
CUTE_HOST_DEVICE constexpr \
decltype(auto) operator()(T&& lhs, U&& rhs) const { \
return OP (std::forward<T>(lhs), std::forward<U>(rhs)); \
} \
}
CUTE_BINARY_OP(plus, +);
CUTE_BINARY_OP(minus, -);
CUTE_BINARY_OP(multiplies, *);
CUTE_BINARY_OP(divides, /);
CUTE_BINARY_OP(modulus, %);
CUTE_BINARY_OP(plus_assign, +=);
CUTE_BINARY_OP(minus_assign, -=);
CUTE_BINARY_OP(multiplies_assign, *=);
CUTE_BINARY_OP(divides_assign, /=);
CUTE_BINARY_OP(modulus_assign, %=);
CUTE_BINARY_OP(bit_and, &);
CUTE_BINARY_OP(bit_or, |);
CUTE_BINARY_OP(bit_xor, ^);
CUTE_BINARY_OP(left_shift, <<);
CUTE_BINARY_OP(right_shift, >>);
CUTE_BINARY_OP(bit_and_assign, &=);
CUTE_BINARY_OP(bit_or_assign, |=);
CUTE_BINARY_OP(bit_xor_assign, ^=);
CUTE_BINARY_OP(left_shift_assign, <<=);
CUTE_BINARY_OP(right_shift_assign, >>=);
CUTE_BINARY_OP(logical_and, &&);
CUTE_BINARY_OP(logical_or, ||);
CUTE_BINARY_OP(equal_to, ==);
CUTE_BINARY_OP(not_equal_to, !=);
CUTE_BINARY_OP(greater, >);
CUTE_BINARY_OP(less, <);
CUTE_BINARY_OP(greater_equal, >=);
CUTE_BINARY_OP(less_equal, <=);
CUTE_NAMED_BINARY_OP(max_fn, cute::max);
CUTE_NAMED_BINARY_OP(min_fn, cute::min);
#undef CUTE_BINARY_OP
#undef CUTE_NAMED_BINARY_OP
/**********/
/** Meta **/
/**********/
template <class Fn, class Arg>
struct bound_fn {
template <class T>
CUTE_HOST_DEVICE constexpr
decltype(auto)
operator()(T&& arg) {
return fn_(arg_, std::forward<T>(arg));
}
Fn fn_;
Arg arg_;
};
template <class Fn, class Arg>
CUTE_HOST_DEVICE constexpr
auto
bind(Fn const& fn, Arg const& arg) {
return bound_fn<Fn,Arg>{fn, arg};
}
} // end namespace cute

View File

@ -0,0 +1,744 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#pragma once
#include <cute/config.hpp>
#include <cute/util/type_traits.hpp>
#include <cute/algorithm/functional.hpp>
#include <cute/tensor.hpp>
#include <cute/atom/mma_atom.hpp>
/** The gemm algorithm takes four (or three) tensors and computes
* D += A * B + C
* It dispatches based on the number of modes each tensor has:
*
* 1. `(V) x (V) => (V)`.
* The element-wise product of vectors. Dispatches to FMA or MMA.
* 2. `(M) x (N) => (M,N)`.
* The outer product of vectors. Dispatches to [3] with new mode K=(1).
* 3. `(M,K) x (N,K) => (M,N)`.
* The product of matrices. Dispatches to [5] with MMA vector-mode V.
* 4. `(V,M) x (V,N) => (V,M,N)`.
* The batched outer product of vectors. Accounts for register reuse and dispatches to [1] for each (m,n).
* 5. `(V,M,K) x (V,N,K) => (V,M,N)`.
* The batched product of matrices. Dispatches to [4] for each (k).
*/
namespace cute
{
//
// Three arguments to four
//
template <class TA, class ALayout,
class TB, class BLayout,
class TC, class CLayout>
CUTE_HOST_DEVICE
void
gemm(Tensor<TA, ALayout> const& A,
Tensor<TB, BLayout> const& B,
Tensor<TC, CLayout> & C)
{
return gemm(C, A, B, C);
}
template <class MMA,
class TA, class ALayout,
class TB, class BLayout,
class TC, class CLayout>
CUTE_HOST_DEVICE
void
gemm(MMA_Atom<MMA> const& mma,
Tensor<TA, ALayout> const& A,
Tensor<TB, BLayout> const& B,
Tensor<TC, CLayout> & C)
{
return gemm(mma, C, A, B, C);
}
//
// Accept mutable temporaries
//
template <class TA, class ALayout,
class TB, class BLayout,
class TC, class CLayout>
CUTE_HOST_DEVICE
void
gemm(Tensor<TA, ALayout> const& A,
Tensor<TB, BLayout> const& B,
Tensor<TC, CLayout> && C)
{
return gemm(C, A, B, C);
}
template <class TD, class DLayout,
class TA, class ALayout,
class TB, class BLayout,
class TC, class CLayout>
CUTE_HOST_DEVICE
void
gemm(Tensor<TD, DLayout> && D,
Tensor<TA, ALayout> const& A,
Tensor<TB, BLayout> const& B,
Tensor<TC, CLayout> const& C)
{
return gemm(D, A, B, C);
}
template <class MMA,
class TA, class ALayout,
class TB, class BLayout,
class TC, class CLayout>
CUTE_HOST_DEVICE
void
gemm(MMA_Atom<MMA> const& mma,
Tensor<TA, ALayout> const& A,
Tensor<TB, BLayout> const& B,
Tensor<TC, CLayout> && C)
{
return gemm(mma, C, A, B, C);
}
template <class MMA,
class TD, class DLayout,
class TA, class ALayout,
class TB, class BLayout,
class TC, class CLayout>
CUTE_HOST_DEVICE
void
gemm(MMA_Atom<MMA> const& mma,
Tensor<TD, DLayout> && D,
Tensor<TA, ALayout> const& A,
Tensor<TB, BLayout> const& B,
Tensor<TC, CLayout> const& C)
{
return gemm(mma, D, A, B, C);
}
//
// Default MMA is UniversalFMA
//
template <class TD, class DLayout,
class TA, class ALayout,
class TB, class BLayout,
class TC, class CLayout>
CUTE_HOST_DEVICE
void
gemm(Tensor<TD, DLayout> & D,
Tensor<TA, ALayout> const& A,
Tensor<TB, BLayout> const& B,
Tensor<TC, CLayout> const& C)
{
using MMA = MMA_Atom<UniversalFMA<typename Tensor<TD,DLayout>::value_type,
typename Tensor<TA,ALayout>::value_type,
typename Tensor<TB,BLayout>::value_type,
typename Tensor<TC,CLayout>::value_type>>;
return gemm(MMA{}, D, A, B, C);
}
//
// Thread-Local Register-Memory GEMMs
//
// Dispatch [1]: (V) x (V) => (V)
template <class MMA,
class TD, class DLayout,
class TA, class ALayout,
class TB, class BLayout,
class TC, class CLayout,
__CUTE_REQUIRES(DLayout::rank == 1 && is_rmem<TD>::value &&
ALayout::rank == 1 && is_rmem<TA>::value &&
BLayout::rank == 1 && is_rmem<TB>::value &&
CLayout::rank == 1 && is_rmem<TC>::value)>
CUTE_HOST_DEVICE
void
gemm(MMA_Atom<MMA> const& mma,
Tensor<TD, DLayout> & D, // (V) Logical data
Tensor<TA, ALayout> const& A, // (V) Logical data
Tensor<TB, BLayout> const& B, // (V) Logical data
Tensor<TC, CLayout> const& C) // (V) Logical data
{
// No static assertions on (V), MMA checks compatibility
mma.call(D, A, B, C);
}
// Dispatch [2]: (M) x (N) => (M,N)
template <class MMA,
class TD, class DLayout,
class TA, class ALayout,
class TB, class BLayout,
class TC, class CLayout,
__CUTE_REQUIRES(DLayout::rank == 2 && is_rmem<TD>::value &&
ALayout::rank == 1 && is_rmem<TA>::value &&
BLayout::rank == 1 && is_rmem<TB>::value &&
CLayout::rank == 2 && is_rmem<TC>::value)>
CUTE_HOST_DEVICE
void
gemm(MMA_Atom<MMA> const& mma,
Tensor<TD, DLayout> & D, // (M,N) Logical data
Tensor<TA, ALayout> const& A, // (M) Logical data
Tensor<TB, BLayout> const& B, // (N) Logical data
Tensor<TC, CLayout> const& C) // (M,N) Logical data
{
CUTE_STATIC_ASSERT_V(size<0>(A) == size<0>(C)); // AM == CM
CUTE_STATIC_ASSERT_V(size<0>(B) == size<1>(C)); // BN == CN
CUTE_STATIC_ASSERT_V(size<0>(C) == size<0>(D) && size<1>(C) == size<1>(D));
gemm(mma,
D, // (M,N)
make_tensor(A.data(), append<2>(A.layout())), // (M,1)
make_tensor(B.data(), append<2>(B.layout())), // (N,1)
C); // (M,N)
}
// Dispatch [3]: (M,K) x (N,K) => (M,N)
template <class MMA,
class TD, class DLayout,
class TA, class ALayout,
class TB, class BLayout,
class TC, class CLayout,
__CUTE_REQUIRES(DLayout::rank == 2 && is_rmem<TD>::value &&
ALayout::rank == 2 && is_rmem<TA>::value &&
BLayout::rank == 2 && is_rmem<TB>::value &&
CLayout::rank == 2 && is_rmem<TC>::value)>
CUTE_HOST_DEVICE
void
gemm(MMA_Atom<MMA> const& mma,
Tensor<TD, DLayout> & D, // (M,N) Logical data
Tensor<TA, ALayout> const& A, // (M,K) Logical data
Tensor<TB, BLayout> const& B, // (N,K) Logical data
Tensor<TC, CLayout> const& C) // (M,N) Logical data
{
CUTE_STATIC_ASSERT_V(size<0>(A) == size<0>(C)); // AM == CM
CUTE_STATIC_ASSERT_V(size<0>(B) == size<1>(C)); // BN == CN
CUTE_STATIC_ASSERT_V(size<1>(A) == size<1>(B)); // AK == BK
CUTE_STATIC_ASSERT_V(size<0>(C) == size<0>(D) && size<1>(C) == size<1>(D));
// Assert this is a 1-value MMA
CUTE_STATIC_ASSERT_V(size<1>(typename MMA_Atom<MMA>::LayoutC_TV{}) == Int<1>{});
CUTE_STATIC_ASSERT_V(size<1>(typename MMA_Atom<MMA>::LayoutA_TV{}) == Int<1>{});
CUTE_STATIC_ASSERT_V(size<1>(typename MMA_Atom<MMA>::LayoutB_TV{}) == Int<1>{});
gemm(mma,
make_tensor(D.data(), prepend<3>(D.layout())), // (1,M,N)
make_tensor(A.data(), prepend<3>(A.layout())), // (1,M,K)
make_tensor(B.data(), prepend<3>(B.layout())), // (1,N,K)
make_tensor(C.data(), prepend<3>(C.layout()))); // (1,M,N)
}
// Dispatch [4]: (V,M) x (V,N) => (V,M,N)
template <class MMA,
class TD, class DLayout,
class TA, class ALayout,
class TB, class BLayout,
class TC, class CLayout,
__CUTE_REQUIRES(DLayout::rank == 3 && is_rmem<TD>::value &&
ALayout::rank == 2 && is_rmem<TA>::value &&
BLayout::rank == 2 && is_rmem<TB>::value &&
CLayout::rank == 3 && is_rmem<TC>::value)>
CUTE_HOST_DEVICE
void
gemm(MMA_Atom<MMA> const& mma,
Tensor<TD, DLayout> & D, // (V,M,N) Logical data
Tensor<TA, ALayout> const& A, // (V,M) Logical data
Tensor<TB, BLayout> const& B, // (V,N) Logical data
Tensor<TC, CLayout> const& C) // (V,M,N) Logical data
{
CUTE_STATIC_ASSERT_V(size<1>(A) == size<1>(C)); // AM == CM
CUTE_STATIC_ASSERT_V(size<1>(B) == size<2>(C)); // BN == CN
CUTE_STATIC_ASSERT_V(size<0>(C) == size<0>(D) && size<1>(C) == size<1>(D) && size<2>(C) == size<2>(D));
auto M = size<1>(A);
auto N = size<1>(B);
// REGISTER .reuse OPTIMIZATIONS
// 64-bit traversal specialization -- serpentine path
if constexpr (decltype(size<0>(A))::value * sizeof(typename TA::value_type) == 8 &&
decltype(size<0>(B))::value * sizeof(typename TB::value_type) == 8)
{
#if 1 // NOTE: Row- vs Col- major could depend on the C-matrix order... (which we can test)
// Row-major serpentine iteration
CUTE_UNROLL
for (int m = 0; m < M; ++m) {
CUTE_UNROLL
for (int n = 0; n < N; ++n) {
int ns = (m & 1) ? N-1-n : n; // Serpentine coordinate
gemm(mma, D(_,m,ns), A(_,m), B(_,ns), C(_,m,ns));
}
}
#else
// Col-major serpentine iteration
CUTE_UNROLL
for (int n = 0; n < N; ++n) {
CUTE_UNROLL
for (int m = 0; m < M; ++m) {
int ms = (n & 1) ? M-1-m : m; // Serpentine coordinate
gemm(mma, D(_,ms,n), A(_,ms), B(_,n), C(_,ms,n));
}
}
#endif
} else
// 32-bit traversal specialization -- kinked serpentine path
if constexpr (decltype(size<0>(A))::value * sizeof(typename TA::value_type) == 4 &&
decltype(size<0>(B))::value * sizeof(typename TB::value_type) == 4)
{
#if 1 // NOTE: Row- vs Col- major could depend on the C-matrix order... (which we can test)
// Row-major kinked serpentine iteration
CUTE_UNROLL
for (int m = 0; m < M; m += 2) {
CUTE_UNROLL
for (int n = 0; n < N; ++n) {
int ns = (m & 2) ? N-1-n : n;
gemm(mma, D(_,m+0,ns), A(_,m+0), B(_,ns), C(_,m+0,ns));
if (m+1 < M) {
gemm(mma, D(_,m+1,ns), A(_,m+1), B(_,ns), C(_,m+1,ns));
}
}
}
#else
// Col-major kinked serpentine iteration
CUTE_UNROLL
for (int n = 0; n < N; n += 2) {
CUTE_UNROLL
for (int m = 0; m < M; ++m) {
// Kinked serpentine traversal for maximum register reuse
int ms = (n & 2) ? M-1-m : m;
gemm(mma, D(_,ms,n+0), A(_,ms), B(_,n+0), C(_,ms,n+0));
if (n+1 < N) {
gemm(mma, D(_,ms,n+1), A(_,ms), B(_,n+1), C(_,ms,n+1));
}
}
}
#endif
} else
// 64-bit + 32-bit traversal order -- keep A (64-bit) in the outer loop and serpentine B
if constexpr (decltype(size<0>(A))::value * sizeof(typename TA::value_type) == 8 &&
decltype(size<0>(B))::value * sizeof(typename TB::value_type) == 4) {
// Row-major serpentine iteration
CUTE_UNROLL
for (int m = 0; m < M; ++m) {
CUTE_UNROLL
for (int n = 0; n < N; ++n) {
int ns = (m & 1) ? N-1-n : n; // Serpentine coordinate
gemm(mma, D(_,m,ns), A(_,m), B(_,ns), C(_,m,ns));
}
}
} else
// 32-bit + 64-bit traversal order -- keep B (64-bit) in the outer loop and serpentine A
if constexpr (decltype(size<0>(A))::value * sizeof(typename TA::value_type) == 4 &&
decltype(size<0>(B))::value * sizeof(typename TB::value_type) == 8) {
// Col-major serpentine iteration
CUTE_UNROLL
for (int n = 0; n < N; ++n) {
CUTE_UNROLL
for (int m = 0; m < M; ++m) {
int ms = (n & 1) ? M-1-m : m; // Serpentine coordinate
gemm(mma, D(_,ms,n), A(_,ms), B(_,n), C(_,ms,n));
}
}
} else
// Fallback to serpentine loop
{
// Col-major serpentine iteration
CUTE_UNROLL
for (int n = 0; n < N; ++n) {
CUTE_UNROLL
for (int m = 0; m < M; ++m) {
int ms = (n & 1) ? M-1-m : m; // Serpentine coordinate
gemm(mma, D(_,ms,n), A(_,ms), B(_,n), C(_,ms,n));
}
}
}
}
// Dispatch [5]: (V,M,K) x (V,N,K) => (V,M,N)
template <class MMA,
class TD, class DLayout,
class TA, class ALayout,
class TB, class BLayout,
class TC, class CLayout,
__CUTE_REQUIRES(DLayout::rank == 3 && is_rmem<TD>::value &&
ALayout::rank == 3 && is_rmem<TA>::value &&
BLayout::rank == 3 && is_rmem<TB>::value &&
CLayout::rank == 3 && is_rmem<TC>::value)>
CUTE_HOST_DEVICE
void
gemm(MMA_Atom<MMA> const& mma,
Tensor<TD, DLayout> & D, // (V,M,N) Logical data
Tensor<TA, ALayout> const& A, // (V,M,K) Logical data
Tensor<TB, BLayout> const& B, // (V,N,K) Logical data
Tensor<TC, CLayout> const& C) // (V,M,N) Logical data
{
CUTE_STATIC_ASSERT_V(size<1>(A) == size<1>(C)); // AM == CM
CUTE_STATIC_ASSERT_V(size<1>(B) == size<2>(C)); // BN == CN
CUTE_STATIC_ASSERT_V(size<2>(A) == size<2>(B)); // AK == BK
CUTE_STATIC_ASSERT_V(size<0>(C) == size<0>(D) && size<1>(C) == size<1>(D) && size<2>(C) == size<2>(D));
auto K = size<2>(A);
CUTE_UNROLL
for (int k = 0; k < K; ++k) {
gemm(mma, D, A(_,_,k), B(_,_,k), C);
}
}
//
// Thread-Local Shared-Memory GEMMs
//
// Dispatch [1]: (V) x (V) => (V)
// Dispatch [2]: (M) x (N) => (M,N)
// Dispatch [3]: (M,K) x (N,K) => (M,N)
// Dispatch [4]: (V,M) x (V,N) => (V,M,N)
// Dispatch [5]: (V,M,K) x (V,N,K) => (V,M,N)
// Dispatch [3]: (M,K) x (N,K) => (M,N)
template <class MMA,
class TD, class DLayout,
class TA, class ALayout,
class TB, class BLayout,
class TC, class CLayout,
__CUTE_REQUIRES(DLayout::rank == 2 && is_rmem<TD>::value &&
ALayout::rank == 2 && is_smem<TA>::value &&
BLayout::rank == 2 && is_smem<TB>::value &&
CLayout::rank == 2 && is_rmem<TC>::value)>
CUTE_HOST_DEVICE
void
gemm(MMA_Atom<MMA> const& mma,
Tensor<TD, DLayout> & D, // (M,N) Logical data
Tensor<TA, ALayout> const& A, // (M,K) Logical data
Tensor<TB, BLayout> const& B, // (N,K) Logical data
Tensor<TC, CLayout> const& C) // (M,N) Logical data
{
CUTE_STATIC_ASSERT_V(size<0>(A) == size<0>(C)); // AM == CM
CUTE_STATIC_ASSERT_V(size<0>(B) == size<1>(C)); // BN == CN
CUTE_STATIC_ASSERT_V(size<1>(A) == size<1>(B)); // AK == BK
CUTE_STATIC_ASSERT_V(size<0>(C) == size<0>(D) && size<1>(C) == size<1>(D));
// Assert this is a 1-value MMA
CUTE_STATIC_ASSERT_V(size<1>(typename MMA_Atom<MMA>::LayoutC_TV{}) == Int<1>{});
CUTE_STATIC_ASSERT_V(size<1>(typename MMA_Atom<MMA>::LayoutA_TV{}) == Int<1>{});
CUTE_STATIC_ASSERT_V(size<1>(typename MMA_Atom<MMA>::LayoutB_TV{}) == Int<1>{});
gemm(mma,
make_tensor(D.data(), prepend<3>(D.layout())), // (1,M,N)
make_tensor(A.data(), prepend<3>(A.layout())), // (1,M,K)
make_tensor(B.data(), prepend<3>(B.layout())), // (1,N,K)
make_tensor(C.data(), prepend<3>(C.layout()))); // (1,M,N)
}
// Dispatch [5]: (V,M,K) x (V,N,K) => (V,M,N)
template <class MMA,
class TD, class DLayout,
class TA, class ALayout,
class TB, class BLayout,
class TC, class CLayout,
__CUTE_REQUIRES(DLayout::rank == 3 && is_rmem<TD>::value &&
ALayout::rank == 3 && is_smem<TA>::value &&
BLayout::rank == 3 && is_smem<TB>::value &&
CLayout::rank == 3 && is_rmem<TC>::value)>
CUTE_HOST_DEVICE
void
gemm(MMA_Atom<MMA> const& mma,
Tensor<TD, DLayout> & D, // (V,M,N) Logical data
Tensor<TA, ALayout> const& A, // (V,M,K) Logical data
Tensor<TB, BLayout> const& B, // (V,N,K) Logical data
Tensor<TC, CLayout> const& C) // (V,M,N) Logical data
{
CUTE_STATIC_ASSERT_V(size<1>(A) == size<1>(C)); // AM == CM
CUTE_STATIC_ASSERT_V(size<1>(B) == size<2>(C)); // BN == CN
CUTE_STATIC_ASSERT_V(size<2>(A) == size<2>(B)); // AK == BK
CUTE_STATIC_ASSERT_V(size<0>(C) == size<0>(D) && size<1>(C) == size<1>(D) && size<2>(C) == size<2>(D));
auto rA = MMA_Atom<MMA>::make_fragment_A(A);
auto rB = MMA_Atom<MMA>::make_fragment_B(B);
auto K = size<2>(A);
CUTE_UNROLL
for (int k = 0; k < K; ++k)
{
copy(A(_,_,k), rA(_,_,k));
copy(B(_,_,k), rB(_,_,k));
// Thread-level register gemm for k
gemm(mma, D, rA(_,_,k), rB(_,_,k), C);
}
}
//
// Collective Shared-Memory GEMMs
//
template <class... Args,
class Alpha, class TA, class ALayout, class TB, class BLayout,
class Beta, class TC, class CLayout,
class ALoadTransformOp, class BLoadTransformOp,
__CUTE_REQUIRES(ALayout::rank == 2 && is_smem<TA>::value &&
BLayout::rank == 2 && is_smem<TB>::value &&
CLayout::rank == 2 && is_smem<TC>::value)>
CUTE_HOST_DEVICE
void
gemm(ThrMMA<Args...> const& thr_mma,
Alpha const& alpha,
Tensor<TA, ALayout> sA,
Tensor<TB, BLayout> sB,
Beta const& beta,
Tensor<TC, CLayout> sC,
ALoadTransformOp const& sA_load_op /* transforms A values before used in GEMM */,
BLoadTransformOp const& sB_load_op /* transforms B values before used in GEMM */)
{
CUTE_STATIC_ASSERT_V(size<0>(sA) == size<0>(sC)); // AM == CM
CUTE_STATIC_ASSERT_V(size<0>(sB) == size<1>(sC)); // BN == CN
CUTE_STATIC_ASSERT_V(size<1>(sA) == size<1>(sB)); // AK == BK
using TypeA = typename TA::value_type;
using TypeB = typename TB::value_type;
using TypeC = typename TC::value_type;
static_assert(is_same_v<decay_t<invoke_result_t<ALoadTransformOp, TypeA>>, TypeA>,
"ALoadTransformOp functor must accept and return value of type TA::value_type");
static_assert(is_same_v<decay_t<invoke_result_t<BLoadTransformOp, TypeB>>, TypeB>,
"BLoadTransformOp functor must accept and return value of type TB::value_type");
// Original, static size of the problem
auto M = size<0>(sC);
auto N = size<1>(sC);
auto K = size<1>(sA);
// Block size of the compute tile
auto BLK_M = tile_size<0>(thr_mma);
auto BLK_N = tile_size<1>(thr_mma);
auto BLK_K = tile_size<2>(thr_mma);
// Compute the "residues"
auto m_residue = M - BLK_M * (ceil_div(M, BLK_M) - Int<1>{}); // (0,BLK_M]
auto n_residue = N - BLK_N * (ceil_div(N, BLK_N) - Int<1>{}); // (0,BLK_N]
auto k_residue = K - BLK_K * (ceil_div(K, BLK_K) ); // (-BLK_K,0]
// Shift the origin so k_residue is zeroth tile
sA.data() = &sA(0,k_residue);
sB.data() = &sB(0,k_residue);
#if 0
if (thread0()) {
printf("%d in BLK_M (%d)\n", int(m_residue), int(BLK_M));
printf("%d in BLK_N (%d)\n", int(n_residue), int(BLK_N));
printf("%d in BLK_K (%d)\n", int(k_residue), int(BLK_K));
}
#endif
//
// MMA Partitioning
//
// Round the layout extents up to BLK_X
Tensor rounded_sA = sA.compose(make_shape(ceil_div(M, BLK_M) * BLK_M, ceil_div(K, BLK_K) * BLK_K));
Tensor rounded_sB = sB.compose(make_shape(ceil_div(N, BLK_N) * BLK_N, ceil_div(K, BLK_K) * BLK_K));
Tensor rounded_sC = sC.compose(make_shape(ceil_div(M, BLK_M) * BLK_M, ceil_div(N, BLK_N) * BLK_N));
#if 0
if (thread0()) {
print(rounded_sA.layout()); print("\n");
print(rounded_sB.layout()); print("\n");
print(rounded_sC.layout()); print("\n");
}
#endif
// Partition the sA and sB tiles across the threads for the MMA
Tensor tCsA = thr_mma.partition_A(rounded_sA); // (MMA,MMA_M,MMA_K)
Tensor tCsB = thr_mma.partition_B(rounded_sB); // (MMA,MMA_N,MMA_K)
Tensor tCsC = thr_mma.partition_C(rounded_sC); // (MMA,MMA_M,MMA_N)
// Create register tensors for the MMA to operate on
Tensor tCrA = thr_mma.make_fragment_A(tCsA); // (MMA,MMA_M,MMA_K)
Tensor tCrB = thr_mma.make_fragment_B(tCsB); // (MMA,MMA_N,MMA_K)
Tensor tCrC = thr_mma.make_fragment_C(tCsC); // (MMA,MMA_M,MMA_N)
#if 0
if (thread0()) {
print(tCsA.layout()); print("\n");
print(tCsB.layout()); print("\n");
print(tCsC.layout()); print("\n");
print(tCrA.layout()); print("\n");
print(tCrB.layout()); print("\n");
print(tCrC.layout()); print("\n");
}
#endif
//
// PREDICATION
//
// Allocate the preds for only the MMA-mode of tCsA and tCsB
Tensor tCpA = make_tensor<bool>(size<0>(tCsA));
Tensor tCpB = make_tensor<bool>(size<0>(tCsB));
// Create coordinate tensors on a single compute block for predication
Tensor cA = make_identity_tensor(make_shape(BLK_M, BLK_K)); // (BLK_M,BLK_K) -> (blk_m,blk_k)
Tensor cB = make_identity_tensor(make_shape(BLK_N, BLK_K)); // (BLK_M,BLK_K) -> (blk_n,blk_k)
// Repeat partitioning with thr_mma
Tensor tCcA = thr_mma.partition_A(cA); // (MMA,1,1) -> (blk_m,blk_k)
Tensor tCcB = thr_mma.partition_B(cB); // (MMA,1,1) -> (blk_n,blk_k)
// Populate the m and n predicates
CUTE_UNROLL
for (int i = 0; i < size(tCpA); ++i) {
tCpA(i) = elem_less(get<0>(tCcA(i)), m_residue);
}
CUTE_UNROLL
for (int i = 0; i < size(tCpB); ++i) {
tCpB(i) = elem_less(get<0>(tCcB(i)), n_residue);
}
#if 0
printf("Thr %d: A(%d,%d):%d B(%d,%d):%d\n",
threadIdx.x,
int(get<0>(tCcA(0))), int(get<1>(tCcA(0))), int(tCpA(0)),
int(get<0>(tCcB(0))), int(get<1>(tCcB(0))), int(tCpB(0)));
#endif
//
// PREFETCH k_block = 0 (with k-predication)
//
CUTE_UNROLL
for (int i = 0; i < size<0>(tCsA); ++i) { // Copy MMA_I
if (k_residue == 0 || get<1>(tCcA(i)) >= -k_residue) { // k_block = 0, predicated on k
CUTE_UNROLL
for (int m = 0; m < size<1>(tCsA); ++m) { // Copy MMA_M, predicated on m
tCrA(i,m,0) = (m_residue == BLK_M || m < size<1>(tCsA)-1 || tCpA(i)) ? sA_load_op(tCsA(i,m,0)) : TypeA{};
}
}
}
CUTE_UNROLL
for (int i = 0; i < size<0>(tCsB); ++i) { // Copy MMA_I
if (k_residue == 0 || get<1>(tCcB(i)) >= -k_residue) { // k_block = 0, predicated on k
CUTE_UNROLL
for (int n = 0; n < size<1>(tCsB); ++n) { // Copy MMA_N, predicated on n
tCrB(i,n,0) = (n_residue == BLK_N || n < size<1>(tCsB)-1 || tCpB(i)) ? sB_load_op(tCsB(i,n,0)) : TypeB{};
}
}
}
//
// MAINLOOP
//
// Clear accumulators
clear(tCrC);
constexpr int K_BLOCK_MAX = size<2>(tCrA);
CUTE_UNROLL
for (int k_block = 0; k_block < K_BLOCK_MAX; ++k_block)
{
// static-if load the next k_block. No k-predication required on these loads.
if (k_block < K_BLOCK_MAX-1)
{
// Load the next k_block
int k_next = k_block + 1;
CUTE_UNROLL
for (int m = 0; m < size<1>(tCsA); ++m) { // Copy MMA_M
CUTE_UNROLL
for (int i = 0; i < size<0>(tCsA); ++i) { // Copy_if MMA_I predicated on m
tCrA(i,m,k_next) = (m_residue == BLK_M || m < size<1>(tCsA)-1 || tCpA(i)) ? sA_load_op(tCsA(i,m,k_next)) : TypeA{};
}
}
CUTE_UNROLL
for (int n = 0; n < size<1>(tCsB); ++n) { // Copy MMA_N
CUTE_UNROLL
for (int i = 0; i < size<0>(tCsB); ++i) { // Copy MMA_I predicated on n
tCrB(i,n,k_next) = (n_residue == BLK_N || n < size<1>(tCsB)-1 || tCpB(i)) ? sB_load_op(tCsB(i,n,k_next)) : TypeB{};
}
}
}
// GEMM on k_block in registers
gemm(thr_mma, tCrA(_,_,k_block), tCrB(_,_,k_block), tCrC);
}
//
// Epilogue
//
Tensor cC = make_identity_tensor(make_shape(BLK_M, BLK_N)); // (BLK_M,BLK_N) -> (blk_m,blk_n)
Tensor tCcC = thr_mma.partition_C(cC); // (MMA, 1, 1) -> (blk_m,blk_n)
const bool isBetaZero = (beta == Beta{});
// Custom axpby_if for now
CUTE_UNROLL
for (int m = 0; m < size<1>(tCsC); ++m)
{
CUTE_UNROLL
for (int n = 0; n < size<2>(tCsC); ++n)
{
CUTE_UNROLL
for (int i = 0; i < size<0>(tCsC); ++i)
{
if ((m_residue == BLK_M || m < size<1>(tCrC)-1 || get<0>(tCcC(i)) < m_residue) &&
(n_residue == BLK_N || n < size<2>(tCrC)-1 || get<1>(tCcC(i)) < n_residue))
{
tCsC(i,m,n) = isBetaZero ? alpha * tCrC(i,m,n) : alpha * tCrC(i,m,n) + beta * tCsC(i,m,n);
}
}
}
}
}
template <class... Args,
class Alpha, class TA, class ALayout, class TB, class BLayout,
class Beta, class TC, class CLayout,
__CUTE_REQUIRES(ALayout::rank == 2 && is_smem<TA>::value &&
BLayout::rank == 2 && is_smem<TB>::value &&
CLayout::rank == 2 && is_smem<TC>::value)>
CUTE_HOST_DEVICE
void
gemm(ThrMMA<Args...> const& thr_mma,
Alpha const& alpha,
Tensor<TA, ALayout> sA,
Tensor<TB, BLayout> sB,
Beta const& beta,
Tensor<TC, CLayout> sC)
{
gemm(thr_mma, alpha, sA, sB, beta, sC, identity() /* sA_load_op */, identity() /* sB_load_op */);
}
} // end namespace cute

View File

@ -0,0 +1,46 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#pragma once
namespace cute
{
// Infinite types that inherit from each other
template <size_t N>
struct prefer : prefer<N-1> {};
template <>
struct prefer<0> {};
// Can be used to preferencially overload implementations
// Higher N in prefer<N> have higher priority.
} // end namespace cute

View File

@ -0,0 +1,123 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
/** Common algorithms on (hierarchical) tensors */
#pragma once
#include <cute/config.hpp>
#include <cute/tensor.hpp>
namespace cute
{
//
// for_each
//
template <class Engine, class Layout, class UnaryOp>
CUTE_HOST_DEVICE constexpr
void
for_each(Tensor<Engine,Layout> const& tensor, UnaryOp&& op)
{
CUTE_UNROLL
for (int i = 0; i < size(tensor); ++i) {
static_cast<UnaryOp&&>(op)(tensor(i));
}
}
template <class Engine, class Layout, class UnaryOp>
CUTE_HOST_DEVICE constexpr
void
for_each(Tensor<Engine,Layout>& tensor, UnaryOp&& op)
{
CUTE_UNROLL
for (int i = 0; i < size(tensor); ++i) {
static_cast<UnaryOp&&>(op)(tensor(i));
}
}
// Accept mutable temporaries
template <class Engine, class Layout, class UnaryOp>
CUTE_HOST_DEVICE constexpr
void
for_each(Tensor<Engine,Layout>&& tensor, UnaryOp&& op)
{
return for_each(tensor, static_cast<UnaryOp&&>(op));
}
//
// transform
//
// Similar to std::transform but does not return number of elements affected
template <class Engine, class Layout, class UnaryOp>
CUTE_HOST_DEVICE constexpr
void
transform(Tensor<Engine,Layout>& tensor, UnaryOp&& op)
{
CUTE_UNROLL
for (int i = 0; i < size(tensor); ++i) {
tensor(i) = static_cast<UnaryOp&&>(op)(tensor(i));
}
}
// Accept mutable temporaries
template <class Engine, class Layout, class UnaryOp>
CUTE_HOST_DEVICE constexpr
void
transform(Tensor<Engine,Layout>&& tensor, UnaryOp&& op)
{
return transform(tensor, std::forward<UnaryOp>(op));
}
// Similar to std::transform transforms one tensors and assigns it to another
template <class EngineIn, class LayoutIn, class EngineOut, class LayoutOut, class UnaryOp>
CUTE_HOST_DEVICE constexpr
void
transform(Tensor<EngineIn,LayoutIn>& tensor_in, Tensor<EngineOut,LayoutOut>& tensor_out, UnaryOp&& op)
{
CUTE_UNROLL
for (int i = 0; i < size(tensor_in); ++i) {
tensor_out(i) = static_cast<UnaryOp&&>(op)(tensor_in(i));
}
}
// Accept mutable temporaries
template <class EngineIn, class LayoutIn, class EngineOut, class LayoutOut, class UnaryOp>
CUTE_HOST_DEVICE constexpr
void
transform(Tensor<EngineIn,LayoutIn>&& tensor_in, Tensor<EngineOut,LayoutOut>&& tensor_out, UnaryOp&& op)
{
return transform(tensor_in, tensor_out, std::forward<UnaryOp>(op));
}
} // end namespace cute

View File

@ -0,0 +1,875 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#pragma once
#include <cute/config.hpp>
#include <cute/util/type_traits.hpp>
#include <cute/container/tuple.hpp>
#include <cute/algorithm/functional.hpp>
#include <cute/numeric/integer_sequence.hpp>
#include <cute/numeric/integral_constant.hpp>
/** Common algorithms on (hierarchical) tuples */
/** Style choice:
* Forward params [using static_cast<T&&>(.)] for const/non-const/ref/non-ref args
* but don't bother forwarding functions as ref-qualified member fns are extremely rare
*/
namespace cute
{
//
// Apply (Unpack)
// (t, f) => f(t_0,t_1,...,t_n)
//
namespace detail {
template <class T, class F, int... I>
CUTE_HOST_DEVICE constexpr
auto
apply(T&& t, F&& f, seq<I...>)
{
return f(get<I>(static_cast<T&&>(t))...);
}
} // end namespace detail
template <class T, class F>
CUTE_HOST_DEVICE constexpr
auto
apply(T&& t, F&& f)
{
return detail::apply(static_cast<T&&>(t), f, tuple_seq<T>{});
}
//
// Transform Apply
// (t, f, g) => g(f(t_0),f(t_1),...)
//
namespace detail {
template <class T, class F, class G, int... I>
CUTE_HOST_DEVICE constexpr
auto
tapply(T&& t, F&& f, G&& g, seq<I...>)
{
return g(f(get<I>(static_cast<T&&>(t)))...);
}
template <class T0, class T1, class F, class G, int... I>
CUTE_HOST_DEVICE constexpr
auto
tapply(T0&& t0, T1&& t1, F&& f, G&& g, seq<I...>)
{
return g(f(get<I>(static_cast<T0&&>(t0)),
get<I>(static_cast<T1&&>(t1)))...);
}
template <class T0, class T1, class T2, class F, class G, int... I>
CUTE_HOST_DEVICE constexpr
auto
tapply(T0&& t0, T1&& t1, T2&& t2, F&& f, G&& g, seq<I...>)
{
return g(f(get<I>(static_cast<T0&&>(t0)),
get<I>(static_cast<T1&&>(t1)),
get<I>(static_cast<T2&&>(t2)))...);
}
} // end namespace detail
template <class T, class F, class G>
CUTE_HOST_DEVICE constexpr
auto
transform_apply(T&& t, F&& f, G&& g)
{
return detail::tapply(static_cast<T&&>(t), f, g, tuple_seq<T>{});
}
template <class T0, class T1, class F, class G>
CUTE_HOST_DEVICE constexpr
auto
transform_apply(T0&& t0, T1&& t1, F&& f, G&& g)
{
return detail::tapply(static_cast<T0&&>(t0), static_cast<T1&&>(t1), f, g, tuple_seq<T0>{});
}
template <class T0, class T1, class T2, class F, class G>
CUTE_HOST_DEVICE constexpr
auto
transform_apply(T0&& t0, T1&& t1, T2&& t2, F&& f, G&& g)
{
return detail::tapply(static_cast<T0&&>(t0), static_cast<T1&&>(t1), static_cast<T2&&>(t2), f, g, tuple_seq<T0>{});
}
//
// For Each
// (t, f) => f(t_0),f(t_1),...,f(t_n)
//
template <class T, class F>
CUTE_HOST_DEVICE constexpr
void
for_each(T&& t, F&& f)
{
detail::apply(t, [&](auto&&... a) { (f(static_cast<decltype(a)&&>(a)), ...); }, tuple_seq<T>{});
}
template <class T, class F>
CUTE_HOST_DEVICE constexpr
auto
for_each_leaf(T&& t, F&& f)
{
if constexpr (is_tuple<remove_cvref_t<T>>::value) {
return detail::apply(static_cast<T&&>(t), [&](auto&&... a){ return (for_each_leaf(static_cast<decltype(a)&&>(a), f), ...); }, tuple_seq<T>{});
} else {
return f(static_cast<T&&>(t));
}
CUTE_GCC_UNREACHABLE;
}
//
// Transform
// (t, f) => (f(t_0),f(t_1),...,f(t_n))
//
template <class T, class F>
CUTE_HOST_DEVICE constexpr
auto
transform(T const& t, F&& f)
{
return detail::tapply(t, f, [](auto const&... a){ return cute::make_tuple(a...); }, tuple_seq<T>{});
}
template <class T0, class T1, class F>
CUTE_HOST_DEVICE constexpr
auto
transform(T0 const& t0, T1 const& t1, F&& f)
{
static_assert(tuple_size<T0>::value == tuple_size<T1>::value, "Mismatched tuple_size");
return detail::tapply(t0, t1, f, [](auto const&... a){ return cute::make_tuple(a...); }, tuple_seq<T0>{});
}
template <class T0, class T1, class T2, class F>
CUTE_HOST_DEVICE constexpr
auto
transform(T0 const& t0, T1 const& t1, T2 const& t2, F&& f)
{
static_assert(tuple_size<T0>::value == tuple_size<T1>::value, "Mismatched tuple_size");
static_assert(tuple_size<T0>::value == tuple_size<T2>::value, "Mismatched tuple_size");
return detail::tapply(t0, t1, t2, f, [](auto const&... a){ return cute::make_tuple(a...); }, tuple_seq<T0>{});
}
template <class T, class F>
CUTE_HOST_DEVICE constexpr
auto
transform_leaf(T const& t, F&& f)
{
if constexpr (is_tuple<T>::value) {
return transform(t, [&](auto const& a) { return transform_leaf(a, f); });
} else {
return f(t);
}
CUTE_GCC_UNREACHABLE;
}
template <class T0, class T1, class F>
CUTE_HOST_DEVICE constexpr
auto
transform_leaf(T0 const& t0, T1 const& t1, F&& f)
{
if constexpr (is_tuple<T0>::value) {
return transform(t0, t1, [&](auto const& a, auto const& b) { return transform_leaf(a, b, f); });
} else {
return f(t0, t1);
}
CUTE_GCC_UNREACHABLE;
}
//
// find and find_if
//
namespace detail {
template <class T, class F>
CUTE_HOST_DEVICE constexpr
auto
find_if(T const& t, F&& f, seq<>)
{
return cute::integral_constant<int, tuple_size<T>::value>{};
}
template <class T, class F, int I, int... Is>
CUTE_HOST_DEVICE constexpr
auto
find_if(T const& t, F&& f, seq<I,Is...>)
{
if constexpr (decltype(f(get<I>(t)))::value) {
return cute::integral_constant<int, I>{};
} else {
return find_if(t, f, seq<Is...>{});
}
CUTE_GCC_UNREACHABLE;
}
} // end namespace detail
template <class T, class F>
CUTE_HOST_DEVICE constexpr
auto
find_if(T const& t, F&& f)
{
if constexpr (is_tuple<T>::value) {
return detail::find_if(t, f, tuple_seq<T>{});
} else {
return cute::integral_constant<int, decltype(f(t))::value ? 0 : 1>{};
}
CUTE_GCC_UNREACHABLE;
}
template <class T, class X>
CUTE_HOST_DEVICE constexpr
auto
find(T const& t, X const& x)
{
return find_if(t, [&](auto const& v) { return v == x; }); // This should always return a static true/false
}
template <class T, class F>
CUTE_HOST_DEVICE constexpr
auto
none_of(T const& t, F&& f)
{
if constexpr (is_tuple<T>::value) {
return cute::integral_constant<bool, decltype(find_if(t, f))::value == tuple_size<T>::value>{};
} else {
return not f(t);
}
CUTE_GCC_UNREACHABLE;
}
template <class T, class F>
CUTE_HOST_DEVICE constexpr
auto
all_of(T const& t, F&& f)
{
if constexpr (is_tuple<T>::value) {
auto not_f = [&](auto const& a) { return not f(a); };
return cute::integral_constant<bool, decltype(find_if(t, not_f))::value == tuple_size<T>::value>{};
} else {
return f(t);
}
CUTE_GCC_UNREACHABLE;
}
template <class T, class F>
CUTE_HOST_DEVICE constexpr
auto
any_of(T const& t, F&& f)
{
return not none_of(t, f);
}
//
// Filter
// (t, f) => <f(t_0),f(t_1),...,f(t_n)>
//
template <class T, class F>
CUTE_HOST_DEVICE constexpr
auto
filter_tuple(T const& t, F&& f)
{
return transform_apply(t, f, [](auto const&... a) { return cute::tuple_cat(a...); });
}
template <class T0, class T1, class F>
CUTE_HOST_DEVICE constexpr
auto
filter_tuple(T0 const& t0, T1 const& t1, F&& f)
{
return transform_apply(t0, t1, f, [](auto const&... a) { return cute::tuple_cat(a...); });
}
//
// Fold (Reduce, Accumulate)
// (t, v, f) => f(...f(f(v,t_0),t_1),...,t_n)
//
namespace detail {
// This impl compiles much faster than cute::apply and variadic args
template <class T, class V, class F>
CUTE_HOST_DEVICE constexpr
decltype(auto)
fold(T&& t, V&& v, F&& f, seq<>)
{
return static_cast<V&&>(v);
}
template <class T, class V, class F, int I, int... Is>
CUTE_HOST_DEVICE constexpr
decltype(auto)
fold(T&& t, V&& v, F&& f, seq<I,Is...>)
{
if constexpr (sizeof...(Is) == 0) {
return f(static_cast<V&&>(v), get<I>(static_cast<T&&>(t)));
} else {
return fold(static_cast<T&&>(t),
f(static_cast<V&&>(v), get<I>(static_cast<T&&>(t))),
f,
seq<Is...>{});
}
CUTE_GCC_UNREACHABLE;
}
} // end namespace detail
template <class T, class V, class F>
CUTE_HOST_DEVICE constexpr
auto
fold(T&& t, V&& v, F&& f)
{
if constexpr (is_tuple<remove_cvref_t<T>>::value) {
return detail::fold(static_cast<T&&>(t),
static_cast<V&&>(v),
f,
tuple_seq<T>{});
} else {
return f(static_cast<V&&>(v), static_cast<T&&>(t));
}
CUTE_GCC_UNREACHABLE;
}
template <class T, class F>
CUTE_HOST_DEVICE constexpr
decltype(auto)
fold_first(T&& t, F&& f)
{
if constexpr (is_tuple<remove_cvref_t<T>>::value) {
return detail::fold(static_cast<T&&>(t),
get<0>(static_cast<T&&>(t)),
f,
make_range<1,tuple_size<remove_cvref_t<T>>::value>{});
} else {
return static_cast<T&&>(t);
}
CUTE_GCC_UNREACHABLE;
}
//
// front, back, take, unwrap
//
// Get the first non-tuple element in a hierarchical tuple
template <class T>
CUTE_HOST_DEVICE constexpr
decltype(auto)
front(T&& t)
{
if constexpr (is_tuple<remove_cvref_t<T>>::value) {
return front(get<0>(static_cast<T&&>(t)));
} else {
return static_cast<T&&>(t);
}
CUTE_GCC_UNREACHABLE;
}
// Get the last non-tuple element in a hierarchical tuple
template <class T>
CUTE_HOST_DEVICE constexpr
decltype(auto)
back(T&& t)
{
if constexpr (is_tuple<remove_cvref_t<T>>::value) {
constexpr int N = tuple_size<remove_cvref_t<T>>::value;
return back(get<N-1>(static_cast<T&&>(t)));
} else {
return static_cast<T&&>(t);
}
CUTE_GCC_UNREACHABLE;
}
// Takes the elements in the range [B,E)
template <int B, int E, class T>
CUTE_HOST_DEVICE constexpr
auto
take(T const& t)
{
return detail::apply(t, [](auto const&... a) { return cute::make_tuple(a...); }, make_range<B,E>{});
}
// Unwrap rank-1 tuples until we're left with a rank>1 tuple or a non-tuple
template <class T>
CUTE_HOST_DEVICE constexpr
auto
unwrap(T const& t)
{
if constexpr (is_tuple<T>::value) {
if constexpr (tuple_size<T>::value == 1) {
return unwrap(get<0>(t));
} else {
return t;
}
} else {
return t;
}
CUTE_GCC_UNREACHABLE;
}
//
// Flatten a hierarchical tuple to a tuple of depth one.
//
template <class T>
CUTE_HOST_DEVICE constexpr
auto
flatten_to_tuple(T const& t)
{
if constexpr (is_tuple<T>::value) {
return filter_tuple(t, [](auto const& a) { return flatten_to_tuple(a); });
} else {
return cute::make_tuple(t);
}
CUTE_GCC_UNREACHABLE;
}
template <class T>
CUTE_HOST_DEVICE constexpr
auto
flatten(T const& t)
{
if constexpr (is_tuple<T>::value) {
return filter_tuple(t, [](auto const& a) { return flatten_to_tuple(a); });
} else {
return t;
}
CUTE_GCC_UNREACHABLE;
}
//
// insert and remove and replace
//
namespace detail {
// Shortcut around cute::tuple_cat for common insert/remove/repeat cases
template <class T, class X, int... I, int... J, int... K>
CUTE_HOST_DEVICE constexpr
auto
construct(T const& t, X const& x, seq<I...>, seq<J...>, seq<K...>)
{
return cute::make_tuple(get<I>(t)..., (void(J),x)..., get<K>(t)...);
}
} // end namespace detail
// Insert x into the Nth position of the tuple
template <int N, class T, class X>
CUTE_HOST_DEVICE constexpr
auto
insert(T const& t, X const& x)
{
return detail::construct(t, x, make_seq<N>{}, seq<0>{}, make_range<N,tuple_size<T>::value>{});
}
// Remove the Nth element of the tuple
template <int N, class T>
CUTE_HOST_DEVICE constexpr
auto
remove(T const& t)
{
return detail::construct(t, 0, make_seq<N>{}, seq<>{}, make_range<N+1,tuple_size<T>::value>{});
}
// Replace the Nth element of the tuple with x
template <int N, class T, class X>
CUTE_HOST_DEVICE constexpr
auto
replace(T const& t, X const& x)
{
return detail::construct(t, x, make_seq<N>{}, seq<0>{}, make_range<N+1,tuple_size<T>::value>{});
}
// Replace the first element of the tuple with x
template <class T, class X>
CUTE_HOST_DEVICE constexpr
auto
replace_front(T const& t, X const& x)
{
if constexpr (is_tuple<T>::value) {
return detail::construct(t, x, seq<>{}, seq<0>{}, make_range<1,tuple_size<T>::value>{});
} else {
return x;
}
CUTE_GCC_UNREACHABLE;
}
// Replace the last element of the tuple with x
template <class T, class X>
CUTE_HOST_DEVICE constexpr
auto
replace_back(T const& t, X const& x)
{
if constexpr (is_tuple<T>::value) {
return detail::construct(t, x, make_seq<tuple_size<T>::value-1>{}, seq<0>{}, seq<>{});
} else {
return x;
}
CUTE_GCC_UNREACHABLE;
}
//
// Make a tuple of Xs of tuple_size N
//
template <int N, class X>
CUTE_HOST_DEVICE constexpr
auto
repeat(X const& x)
{
return detail::construct(0, x, seq<>{}, make_seq<N>{}, seq<>{});
}
//
// Make a tuple of Xs the same profile as tuple
//
template <class T, class X>
CUTE_HOST_DEVICE constexpr
auto
repeat_like(T const& t, X const& x)
{
if constexpr (is_tuple<T>::value) {
return transform(t, [&](auto const& a) { return repeat_like(a,x); });
} else {
return x;
}
CUTE_GCC_UNREACHABLE;
}
// Group the elements [B,E) of a T into a single element
// e.g. group<2,4>(T<_1,_2,_3,_4,_5,_6>{})
// => T<_1,_2,T<_3,_4>,_5,_6>{}
template <int B, int E, class T>
CUTE_HOST_DEVICE constexpr
auto
group(T const& t)
{
return detail::construct(t, take<B,E>(t), make_seq<B>{}, seq<0>{}, make_range<E,tuple_size<T>::value>{});
}
//
// Extend a T to rank N by appending/prepending an element
//
template <int N, class T, class X>
CUTE_HOST_DEVICE constexpr
auto
append(T const& a, X const& x)
{
if constexpr (is_tuple<T>::value) {
if constexpr (N == tuple_size<T>::value) {
return a;
} else {
static_assert(N > tuple_size<T>::value);
return detail::construct(a, x, make_seq<tuple_size<T>::value>{}, make_seq<N-tuple_size<T>::value>{}, seq<>{});
}
} else {
if constexpr (N == 1) {
return a;
} else {
return detail::construct(cute::make_tuple(a), x, seq<0>{}, make_seq<N-1>{}, seq<>{});
}
}
CUTE_GCC_UNREACHABLE;
}
template <class T, class X>
CUTE_HOST_DEVICE constexpr
auto
append(T const& a, X const& x)
{
if constexpr (is_tuple<T>::value) {
return detail::construct(a, x, make_seq<tuple_size<T>::value>{}, seq<0>{}, seq<>{});
} else {
return cute::make_tuple(a, x);
}
CUTE_GCC_UNREACHABLE;
}
template <int N, class T, class X>
CUTE_HOST_DEVICE constexpr
auto
prepend(T const& a, X const& x)
{
if constexpr (is_tuple<T>::value) {
if constexpr (N == tuple_size<T>::value) {
return a;
} else {
static_assert(N > tuple_size<T>::value);
return detail::construct(a, x, seq<>{}, make_seq<N-tuple_size<T>::value>{}, make_seq<tuple_size<T>::value>{});
}
} else {
if constexpr (N == 1) {
return a;
} else {
static_assert(N > 1);
return detail::construct(cute::make_tuple(a), x, seq<>{}, make_seq<N-1>{}, seq<0>{});
}
}
CUTE_GCC_UNREACHABLE;
}
template <class T, class X>
CUTE_HOST_DEVICE constexpr
auto
prepend(T const& a, X const& x)
{
if constexpr (is_tuple<T>::value) {
return detail::construct(a, x, seq<>{}, seq<0>{}, make_seq<tuple_size<T>::value>{});
} else {
return cute::make_tuple(x, a);
}
CUTE_GCC_UNREACHABLE;
}
//
// Inclusive scan (prefix sum)
//
namespace detail {
template <class T, class V, class F, int I, int... Is>
CUTE_HOST_DEVICE constexpr
auto
iscan(T const& t, V const& v, F&& f, seq<I,Is...>)
{
// Apply the function to v and the element at I
auto v_next = f(v, get<I>(t));
// Replace I with v_next
auto t_next = replace<I>(t, v_next);
#if 0
std::cout << "ISCAN i" << I << std::endl;
std::cout << " t " << t << std::endl;
std::cout << " i " << v << std::endl;
std::cout << " f(i,t) " << v_next << std::endl;
std::cout << " t_n " << t_next << std::endl;
#endif
if constexpr (sizeof...(Is) == 0) {
return t_next;
} else {
return iscan(t_next, v_next, f, seq<Is...>{});
}
CUTE_GCC_UNREACHABLE;
}
} // end namespace detail
template <class T, class V, class F>
CUTE_HOST_DEVICE constexpr
auto
iscan(T const& t, V const& v, F&& f)
{
return detail::iscan(t, v, f, tuple_seq<T>{});
}
//
// Exclusive scan (prefix sum)
//
namespace detail {
template <class T, class V, class F, int I, int... Is>
CUTE_HOST_DEVICE constexpr
auto
escan(T const& t, V const& v, F&& f, seq<I,Is...>)
{
if constexpr (sizeof...(Is) == 0) {
// Replace I with v
return replace<I>(t, v);
} else {
// Apply the function to v and the element at I
auto v_next = f(v, get<I>(t));
// Replace I with v
auto t_next = replace<I>(t, v);
#if 0
std::cout << "ESCAN i" << I << std::endl;
std::cout << " t " << t << std::endl;
std::cout << " i " << v << std::endl;
std::cout << " f(i,t) " << v_next << std::endl;
std::cout << " t_n " << t_next << std::endl;
#endif
// Recurse
return escan(t_next, v_next, f, seq<Is...>{});
}
CUTE_GCC_UNREACHABLE;
}
} // end namespace detail
template <class T, class V, class F>
CUTE_HOST_DEVICE constexpr
auto
escan(T const& t, V const& v, F&& f)
{
return detail::escan(t, v, f, tuple_seq<T>{});
}
//
// Zip (Transpose)
//
// Take ((a,b,c,...),(x,y,z,...),...) rank-R0 x rank-R1 input
// to produce ((a,x,...),(b,y,...),(c,z,...),...) rank-R1 x rank-R0 output
namespace detail {
template <int J, class... Ts>
CUTE_HOST_DEVICE constexpr
auto
zip_(Ts const&... ts)
{
return cute::make_tuple(get<J>(ts)...);
}
template <class T, int... Is, int... Js>
CUTE_HOST_DEVICE constexpr
auto
zip(T const& t, seq<Is...>, seq<Js...>)
{
static_assert(conjunction<bool_constant<tuple_size<tuple_element_t<0,T>>::value == tuple_size<tuple_element_t<Is,T>>::value>...>::value, "Mismatched Ranks");
return cute::make_tuple(zip_<Js>(get<Is>(t)...)...);
}
} // end namespace detail
template <class T>
CUTE_HOST_DEVICE constexpr
auto
zip(T const& t)
{
if constexpr (is_tuple<T>::value) {
if constexpr (is_tuple<tuple_element_t<0,T>>::value) {
return detail::zip(t, tuple_seq<T>{}, tuple_seq<tuple_element_t<0,T>>{});
} else {
return cute::make_tuple(t);
}
} else {
return t;
}
CUTE_GCC_UNREACHABLE;
}
// Convenient to pass them in separately
template <class T0, class T1, class... Ts>
CUTE_HOST_DEVICE constexpr
auto
zip(T0 const& t0, T1 const& t1, Ts const&... ts)
{
return zip(cute::make_tuple(t0, t1, ts...));
}
//
// zip2_by -- A guided zip for rank-2 tuples
// Take a tuple like ((A,a),((B,b),(C,c)),d)
// and produce a tuple ((A,(B,C)),(a,(b,c),d))
// where the rank-2 modes are selected by the terminals of the guide (X,(X,X))
//
namespace detail {
template <class T, class TG, int... Is, int... Js>
CUTE_HOST_DEVICE constexpr
auto
zip2_by(T const& t, TG const& guide, seq<Is...>, seq<Js...>)
{
// zip2_by produces the modes like ((A,a),(B,b),...)
auto split = cute::make_tuple(zip2_by(get<Is>(t), get<Is>(guide))...);
// Rearrange and append missing modes from t to make ((A,B,...),(a,b,...,x,y))
return cute::make_tuple(cute::make_tuple(get<0>(get<Is>(split))...),
cute::make_tuple(get<1>(get<Is>(split))..., get<Js>(t)...));
}
} // end namespace detail
template <class T, class TG>
CUTE_HOST_DEVICE constexpr
auto
zip2_by(T const& t, TG const& guide)
{
if constexpr (is_tuple<TG>::value) {
constexpr int TR = tuple_size<T>::value;
constexpr int GR = tuple_size<TG>::value;
static_assert(TR >= GR, "Mismatched ranks");
return detail::zip2_by(t, guide,
make_range< 0, GR>{},
make_range<GR, TR>{});
} else {
static_assert(tuple_size<T>::value == 2, "Mismatched ranks");
return t;
}
CUTE_GCC_UNREACHABLE;
}
} // end namespace cute

View File

@ -0,0 +1,243 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#pragma once
#include <cute/config.hpp>
// Config
// #if (defined(__MACA_ARCH__) && (__MACA_ARCH__ >= 900) && \
// ((__CUDACC_VER_MAJOR__ >= 12) || ((__CUDACC_VER_MAJOR__ == 11) && (__CUDACC_VER_MINOR__ >= 8))))
// # define CUTE_ARCH_CLUSTER_SM90_ENABLED
// #endif
// #if (defined(__MACA_ARCH__) && (__MACA_ARCH__ >= 900) && (__CUDACC_VER_MAJOR__ >= 12))
// # define CUTE_ARCH_ELECT_ONE_SM90_ENABLED
// #endif
namespace cute {
CUTE_DEVICE void cluster_arrive_relaxed()
{
#if defined(CUTE_ARCH_CLUSTER_SM90_ENABLED)
asm volatile("barrier.cluster.arrive.relaxed.aligned;\n" : : );
#else
CUTE_RUNTIME_ASSERT("CUTE_ARCH_CLUSTER_SM90_ENABLED is not defined");
#endif
}
CUTE_DEVICE void cluster_arrive()
{
#if defined(CUTE_ARCH_CLUSTER_SM90_ENABLED)
asm volatile("barrier.cluster.arrive.aligned;\n" : : );
#else
CUTE_RUNTIME_ASSERT("CUTE_ARCH_CLUSTER_SM90_ENABLED is not defined");
#endif
}
CUTE_DEVICE void cluster_wait()
{
#if defined(CUTE_ARCH_CLUSTER_SM90_ENABLED)
asm volatile("barrier.cluster.wait.aligned;\n" : : );
#else
CUTE_RUNTIME_ASSERT("CUTE_ARCH_CLUSTER_SM90_ENABLED is not defined");
#endif
}
CUTE_DEVICE void cluster_sync()
{
#if defined(CUTE_ARCH_CLUSTER_SM90_ENABLED)
cluster_arrive();
cluster_wait();
#else
CUTE_RUNTIME_ASSERT("CUTE_ARCH_CLUSTER_SM90_ENABLED is not defined");
#endif
}
// Returns the dim3 grid size in terms of number of clusters.
CUTE_DEVICE dim3 cluster_grid_dims()
{
#if defined(CUTE_ARCH_CLUSTER_SM90_ENABLED)
uint32_t x, y, z;
asm volatile("mov.u32 %0, %nclusterid.x;\n" : "=r"(x) : );
asm volatile("mov.u32 %0, %nclusterid.y;\n" : "=r"(y) : );
asm volatile("mov.u32 %0, %nclusterid.z;\n" : "=r"(z) : );
return {x, y, z};
#elif defined(__MACA_ARCH__)
// MSVC requires protecting use of gridDim with __MACA_ARCH__.
return gridDim;
#elif defined(_MSC_VER)
CUTE_RUNTIME_ASSERT("cluster_grid_dims() can only be called on device");
#else
return {0, 0, 0};
#endif
}
// Returns the dim3 cluster rank in the grid.
CUTE_DEVICE dim3 cluster_id_in_grid()
{
#if defined(CUTE_ARCH_CLUSTER_SM90_ENABLED)
uint32_t x, y, z;
asm volatile("mov.u32 %0, %clusterid.x;\n" : "=r"(x) : );
asm volatile("mov.u32 %0, %clusterid.y;\n" : "=r"(y) : );
asm volatile("mov.u32 %0, %clusterid.z;\n" : "=r"(z) : );
return {x, y, z};
#elif defined(__MACA_ARCH__)
// MSVC requires protecting use of blockIdx with __MACA_ARCH__.
return blockIdx;
#elif defined(_MSC_VER)
CUTE_RUNTIME_ASSERT("cluster_id_in_grid() can only be called on device");
#else
return {0, 0, 0};
#endif
}
// Returns the relative dim3 block rank local to the cluster.
CUTE_DEVICE dim3 block_id_in_cluster()
{
#if defined(CUTE_ARCH_CLUSTER_SM90_ENABLED)
uint32_t x, y, z;
asm volatile("mov.u32 %0, %cluster_ctaid.x;\n" : "=r"(x) : );
asm volatile("mov.u32 %0, %cluster_ctaid.y;\n" : "=r"(y) : );
asm volatile("mov.u32 %0, %cluster_ctaid.z;\n" : "=r"(z) : );
return {x, y, z};
#else
return {0,0,0};
#endif
}
// Returns the dim3 cluster shape.
CUTE_DEVICE dim3 cluster_shape()
{
#if defined(CUTE_ARCH_CLUSTER_SM90_ENABLED)
uint32_t x, y, z;
asm volatile("mov.u32 %0, %cluster_nctaid.x;\n" : "=r"(x) : );
asm volatile("mov.u32 %0, %cluster_nctaid.y;\n" : "=r"(y) : );
asm volatile("mov.u32 %0, %cluster_nctaid.z;\n" : "=r"(z) : );
return {x, y, z};
#else
return {1,1,1};
#endif
}
// Get 1D ctaid in a cluster.
MCTLASS_DEVICE uint32_t block_rank_in_cluster()
{
#if defined(CUTE_ARCH_CLUSTER_SM90_ENABLED)
uint32_t rank;
asm volatile("mov.u32 %0, %cluster_ctarank;\n" : "=r"(rank) :);
return rank;
#else
return 0;
#endif
}
// Set the destination block-ID in cluster for a given SMEM Address
MCTLASS_DEVICE uint32_t set_block_rank(uint32_t smemAddr, uint32_t rank)
{
#if defined(CUTE_ARCH_CLUSTER_SM90_ENABLED)
uint32_t result;
asm volatile("mapa.shared::cluster.u32 %0, %1, %2;\n"
: "=r"(result)
: "r"(smemAddr), "r"(rank));
return result;
#else
return smemAddr;
#endif
}
// Elect one thread in the warp. The elected thread gets its predicate set to true, all others obtain false.
CUTE_HOST_DEVICE uint32_t elect_one_sync()
{
#if defined(CUTE_ARCH_ELECT_ONE_SM90_ENABLED)
uint32_t pred = 0;
uint32_t laneid = 0;
asm volatile(
"{\n"
".reg .b32 %rx;\n"
".reg .pred %px;\n"
" elect.sync %rx|%px, %2;\n"
"@%px mov.s32 %1, 1;\n"
" mov.s32 %0, %rx;\n"
"}\n"
: "+r"(laneid), "+r"(pred)
: "r"(0xFFFFFFFF));
return pred;
#elif defined(__MACA_ARCH__)
return (threadIdx.x % 64) == 0;
#else
return true;
#endif
}
struct ElectOneLaneIdReturnType {
uint32_t is_leader;
uint32_t leader_lane_id;
};
CUTE_HOST_DEVICE
ElectOneLaneIdReturnType
elect_one_leader_sync()
{
#if defined(CUTE_ARCH_ELECT_ONE_SM90_ENABLED)
uint32_t pred = 0;
uint32_t laneid = 0;
asm volatile(
"{\n"
".reg .b32 %rx;\n"
".reg .pred %px;\n"
" elect.sync %rx|%px, %2;\n"
"@%px mov.s32 %1, 1;\n"
" mov.s32 %0, %rx;\n"
"}\n"
: "+r"(laneid), "+r"(pred)
: "r"(0xFFFFFFFF));
return {pred, laneid};
#elif defined(__MACA_ARCH__)
return {(threadIdx.x % 64) == 0, 0};
#else
return {true, 0};
#endif
}
// Store value to remote shared memory in the cluster
CUTE_DEVICE
void
store_shared_remote(uint32_t value, uint32_t smem_addr, uint32_t mbarrier_addr, uint32_t dst_cta_rank)
{
#if defined(CUTE_ARCH_CLUSTER_SM90_ENABLED)
uint32_t dsmem_addr = set_block_rank(smem_addr, dst_cta_rank);
uint32_t remote_barrier_addr = set_block_rank(mbarrier_addr, dst_cta_rank);
asm volatile("st.async.shared::cluster.mbarrier::complete_tx::bytes.u32 [%0], %1, [%2];"
: : "r"(dsmem_addr), "r"(value), "r"(remote_barrier_addr));
#endif
}
} // end namespace cute

View File

@ -0,0 +1,71 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#pragma once
#include <cute/config.hpp>
#include <cute/arch/util.hpp>
#include <cute/numeric/uint128.hpp>
namespace cute
{
//
// Direct Copy for any type
//
template <class S, class D = S>
struct UniversalCopy
{
using SRegisters = S[1];
using DRegisters = D[1];
CUTE_HOST_DEVICE static constexpr void
copy(S const& src,
D & dst)
{
dst = src;
}
};
//
// Placeholder for the copy algorithm's default, auto-vectorizing behavior
//
struct DefaultCopy
{
using SRegisters = uint128_t[1];
using DRegisters = uint128_t[1];
};
using AutoVectorizingCopy = DefaultCopy;
} // end namespace cute

View File

@ -0,0 +1,322 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#pragma once
#include <cute/config.hpp>
#include <cute/arch/copy.hpp>
// Config
#if defined(__clang__) && defined(__MACA__)
// ldmatrix PTX instructions added in Clang 14: https://reviews.llvm.org/D107046
// ... but will not work until Clang 15:
// * https://reviews.llvm.org/D121666
// * https://reviews.llvm.org/D126846
#define CUTE_ARCH_CLANG_SUPPORTS_LDSM_SM75 (__clang_major__ >= 15)
#endif
#if defined(__MXCC__) || defined(__MACACC_RTC__)
// ldmatrix PTX instruction added in CUDA 10.2+
#define CUTE_ARCH_NVCC_SUPPORTS_LDSM_SM75 ((__CUDACC_VER_MAJOR__ == 10 && __CUDACC_VER_MINOR__ >= 2) || __CUDACC_VER_MAJOR__ >= 11)
#endif
#if ! defined(CUTE_ARCH_LDSM_SM75_SUPPORTED)
#define CUTE_ARCH_LDSM_SM75_SUPPORTED (CUTE_ARCH_NVCC_SUPPORTS_LDSM_SM75 || CUTE_ARCH_CLANG_SUPPORTS_LDSM_SM75)
#endif
#if ! defined(CUTE_ARCH_LDSM_SM75_ENABLED)
#define CUTE_ARCH_LDSM_SM75_ENABLED (CUTE_ARCH_LDSM_SM75_SUPPORTED)
#endif
#if 0
#define CUTE_ARCH_LDSM_SM75_ACTIVATED 1
#endif
namespace cute
{
struct SM75_U32x1_LDSM_N
{
using SRegisters = uint128_t[1];
using DRegisters = uint32_t[1];
CUTE_HOST_DEVICE static void
copy(uint128_t const& smem_src,
uint32_t& dst)
{
#if defined(CUTE_ARCH_LDSM_SM75_ACTIVATED)
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(&smem_src);
asm volatile ("ldmatrix.sync.aligned.x1.m8n8.shared.b16 {%0}, [%1];\n"
: "=r"(dst)
: "r"(smem_int_ptr));
#else
CUTE_RUNTIME_ASSERT("Trying to use ldmatrix without CUTE_ARCH_LDSM_SM75_ACTIVATED.");
#endif
}
};
struct SM75_U32x2_LDSM_N
{
using SRegisters = uint128_t[1];
using DRegisters = uint32_t[2];
CUTE_HOST_DEVICE static void
copy(uint128_t const& smem_src,
uint32_t& dst0, uint32_t& dst1)
{
#if defined(CUTE_ARCH_LDSM_SM75_ACTIVATED)
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(&smem_src);
asm volatile ("ldmatrix.sync.aligned.x2.m8n8.shared.b16 {%0, %1}, [%2];\n"
: "=r"(dst0), "=r"(dst1)
: "r"(smem_int_ptr));
#else
CUTE_RUNTIME_ASSERT("Trying to use ldmatrix without CUTE_ARCH_LDSM_SM75_ACTIVATED.");
#endif
}
};
struct SM75_U32x4_LDSM_N
{
using SRegisters = uint128_t[1];
using DRegisters = uint32_t[4];
CUTE_HOST_DEVICE static void
copy(uint128_t const& smem_src,
uint32_t& dst0, uint32_t& dst1, uint32_t& dst2, uint32_t& dst3)
{
#if defined(CUTE_ARCH_LDSM_SM75_ACTIVATED)
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(&smem_src);
asm volatile ("ldmatrix.sync.aligned.x4.m8n8.shared.b16 {%0, %1, %2, %3}, [%4];\n"
: "=r"(dst0), "=r"(dst1), "=r"(dst2), "=r"(dst3)
: "r"(smem_int_ptr));
#elif defined(__MACA_ARCH__)
const int lane_id = __lane_id();
if (lane_id >= 32) return;
uint64_t sm_ptr = reinterpret_cast<uint64_t>(&smem_src);
uint64_t row_ptr[32];
for (int i = 0; i < 32; ++i) {
row_ptr[i] = __shfl_down_sync(ULONG_MAX, sm_ptr, i - lane_id);
}
const int row_id = lane_id / 4;
const int col_offset = lane_id % 4;
dst0 = *(reinterpret_cast<uint32_t *>(row_ptr[0 + row_id]) + col_offset);
dst1 = *(reinterpret_cast<uint32_t *>(row_ptr[8 + row_id]) + col_offset);
dst2 = *(reinterpret_cast<uint32_t *>(row_ptr[16 + row_id]) + col_offset);
dst3 = *(reinterpret_cast<uint32_t *>(row_ptr[24 + row_id]) + col_offset);
#else
CUTE_RUNTIME_ASSERT("Trying to use ldmatrix without CUTE_ARCH_LDSM_SM75_ACTIVATED.");
#endif
}
};
struct SM75_U32x4_LDSM_N_B
{
using SRegisters = uint128_t[1];
using DRegisters = uint32_t[4];
CUTE_HOST_DEVICE static void
copy(uint128_t const& smem_src,
uint32_t& dst0, uint32_t& dst1, uint32_t& dst2, uint32_t& dst3)
{
#if defined(CUTE_ARCH_LDSM_SM75_ACTIVATED)
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(&smem_src);
asm volatile ("ldmatrix.sync.aligned.x4.m8n8.shared.b16 {%0, %1, %2, %3}, [%4];\n"
: "=r"(dst0), "=r"(dst1), "=r"(dst2), "=r"(dst3)
: "r"(smem_int_ptr));
#elif defined(__MACA_ARCH__)
const int lane_id = __lane_id();
if (lane_id >= 32) return;
uint64_t sm_ptr = reinterpret_cast<uint64_t>(&smem_src);
uint64_t row_ptr[32];
for (int i = 0; i < 32; ++i) {
row_ptr[i] = __shfl_down_sync(ULONG_MAX, sm_ptr, i - lane_id);
}
const int row_id = lane_id / 4;
const int col_offset = lane_id % 4;
dst0 = *(reinterpret_cast<uint32_t *>(row_ptr[0 + row_id]) + col_offset);
dst1 = *(reinterpret_cast<uint32_t *>(row_ptr[8 + row_id]) + col_offset);
dst2 = *(reinterpret_cast<uint32_t *>(row_ptr[16 + row_id]) + col_offset);
dst3 = *(reinterpret_cast<uint32_t *>(row_ptr[24 + row_id]) + col_offset);
#else
CUTE_RUNTIME_ASSERT("Trying to use ldmatrix without CUTE_ARCH_LDSM_SM75_ACTIVATED.");
#endif
}
};
struct SM75_U16x2_LDSM_T
{
using SRegisters = uint128_t[1];
using DRegisters = uint32_t[1];
CUTE_HOST_DEVICE static void
copy(uint128_t const& smem_src,
uint32_t& dst)
{
#if defined(CUTE_ARCH_LDSM_SM75_ACTIVATED)
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(&smem_src);
asm volatile ("ldmatrix.sync.aligned.x1.trans.m8n8.shared.b16 {%0}, [%1];\n"
: "=r"(dst)
: "r"(smem_int_ptr));
#else
CUTE_RUNTIME_ASSERT("Trying to use ldmatrix without CUTE_ARCH_LDSM_SM75_ACTIVATED.");
#endif
}
};
struct SM75_U16x4_LDSM_T
{
using SRegisters = uint128_t[1];
using DRegisters = uint32_t[2];
CUTE_HOST_DEVICE static void
copy(uint128_t const& smem_src,
uint32_t& dst0, uint32_t& dst1)
{
#if defined(CUTE_ARCH_LDSM_SM75_ACTIVATED)
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(&smem_src);
asm volatile ("ldmatrix.sync.aligned.x2.trans.m8n8.shared.b16 {%0, %1}, [%2];\n"
: "=r"(dst0), "=r"(dst1)
: "r"(smem_int_ptr));
#else
CUTE_RUNTIME_ASSERT("Trying to use ldmatrix without CUTE_ARCH_LDSM_SM75_ACTIVATED.");
#endif
}
};
struct SM75_U16x8_LDSM_T
{
using SRegisters = uint128_t[1];
using DRegisters = uint32_t[4];
CUTE_HOST_DEVICE static void
copy(uint128_t const& smem_src,
uint32_t& dst0, uint32_t& dst1, uint32_t& dst2, uint32_t& dst3)
{
#if defined(CUTE_ARCH_LDSM_SM75_ACTIVATED)
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(&smem_src);
asm volatile ("ldmatrix.sync.aligned.x4.trans.m8n8.shared.b16 {%0, %1, %2, %3}, [%4];\n"
: "=r"(dst0), "=r"(dst1), "=r"(dst2), "=r"(dst3)
: "r"(smem_int_ptr));
#elif defined(__MACA_ARCH__)
const int lane_id = __lane_id();
if (lane_id >= 32) return;
uint64_t sm_ptr = reinterpret_cast<uint64_t>(&smem_src);
uint64_t row_ptr[32];
for (int i = 0; i < 32; ++i) {
row_ptr[i] = __shfl_down_sync(ULONG_MAX, sm_ptr, i - lane_id);
}
const int row_offset = lane_id % 4 * 2;
const int col_offset = lane_id / 4;
auto low_b16_addr = reinterpret_cast<uint16_t *>(row_ptr[0 + row_offset]) + col_offset;
auto high_b16_addr = reinterpret_cast<uint16_t *>(row_ptr[0 + row_offset + 1]) + col_offset;
auto dst_b16 = reinterpret_cast<uint16_t *>(&dst0);
*dst_b16 = *low_b16_addr;
*(dst_b16 + 1) = *high_b16_addr;
low_b16_addr = reinterpret_cast<uint16_t *>(row_ptr[8 + row_offset]) + col_offset;
high_b16_addr = reinterpret_cast<uint16_t *>(row_ptr[8 + row_offset + 1]) + col_offset;
dst_b16 = reinterpret_cast<uint16_t *>(&dst1);
*dst_b16 = *low_b16_addr;
*(dst_b16 + 1) = *high_b16_addr;
low_b16_addr = reinterpret_cast<uint16_t *>(row_ptr[16 + row_offset]) + col_offset;
high_b16_addr = reinterpret_cast<uint16_t *>(row_ptr[16 + row_offset + 1]) + col_offset;
dst_b16 = reinterpret_cast<uint16_t *>(&dst2);
*dst_b16 = *low_b16_addr;
*(dst_b16 + 1) = *high_b16_addr;
low_b16_addr = reinterpret_cast<uint16_t *>(row_ptr[24 + row_offset]) + col_offset;
high_b16_addr = reinterpret_cast<uint16_t *>(row_ptr[24 + row_offset + 1]) + col_offset;
dst_b16 = reinterpret_cast<uint16_t *>(&dst3);
*dst_b16 = *low_b16_addr;
*(dst_b16 + 1) = *high_b16_addr;
#else
CUTE_RUNTIME_ASSERT("Trying to use ldmatrix without CUTE_ARCH_LDSM_SM75_ACTIVATED.");
#endif
}
};
//
// Legacy LDSM interfaces that aren't very useful
//
template <class T>
CUTE_HOST_DEVICE
void
copy_ldsm(uint128_t const* const smem_ptr,
T* rmem_ptr)
{
uint32_t* reg_ptr = reinterpret_cast<uint32_t*>(rmem_ptr);
// if constexpr
if (sizeof(T) == 4) {
SM75_U32x1_LDSM_N::copy(smem_ptr[0], reg_ptr[0]);
}
else if (sizeof(T) == 8) {
SM75_U32x2_LDSM_N::copy(smem_ptr[0], reg_ptr[0], reg_ptr[1]);
}
else if (sizeof(T) == 16) {
SM75_U32x4_LDSM_N::copy(smem_ptr[0], reg_ptr[0], reg_ptr[1], reg_ptr[2], reg_ptr[3]);
}
else {
static_assert(sizeof(T) == 4 || sizeof(T) == 8 || sizeof(T) == 16, "sizeof(T) is not supported");
}
}
template <class T>
CUTE_HOST_DEVICE
void
copy_ldsm_trans(uint128_t const* const smem_ptr,
T* rmem_ptr)
{
uint32_t* reg_ptr = reinterpret_cast<uint32_t*>(rmem_ptr);
// if constexpr
if (sizeof(T) == 4) {
SM75_U16x2_LDSM_T::copy(smem_ptr[0], reg_ptr[0]);
}
else if (sizeof(T) == 8) {
SM75_U16x4_LDSM_T::copy(smem_ptr[0], reg_ptr[0], reg_ptr[1]);
}
else if (sizeof(T) == 16) {
SM75_U16x8_LDSM_T::copy(smem_ptr[0], reg_ptr[0], reg_ptr[1], reg_ptr[2], reg_ptr[3]);
}
else {
static_assert(sizeof(T) == 4 || sizeof(T) == 8 || sizeof(T) == 16, "sizeof(T) is not supported");
}
}
} // end namespace cute

View File

@ -0,0 +1,201 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#pragma once
#include <cute/config.hpp>
#include <cute/arch/copy.hpp>
#if defined(__MACA_ARCH__) && 0
# define CUTE_ARCH_CP_ASYNC_SM80_ENABLED
#endif
#if defined(__MACA_ARCH__) && (__MACA_ARCH__ == 1500)
# define MACA_ARCH_LDS_TRANS_ENABLED
#endif
namespace cute
{
extern __device__ void
_vmov_cp_async(__attribute__((address_space (3)))int8_t *, __attribute__((address_space (1)))int8_t *, int32_t, int32_t)
__asm("llvm.mxc.cp.async.global.to.shared");
/// Copy via cp.async with caching at all levels
template <class TS, class TD = TS>
struct SM80_CP_ASYNC_CACHEALWAYS
{
using SRegisters = TS[1];
using DRegisters = TD[1];
static_assert(sizeof(TS) == sizeof(TD), "cp.async requires sizeof(src_value_type) == sizeof(dst_value_type)");
static_assert(sizeof(TS) == 4 || sizeof(TS) == 8 || sizeof(TS) == 16, "cp.async sizeof(TS) is not supported");
CUTE_HOST_DEVICE static void
copy(TS const& gmem_src,
TD & smem_dst)
{
#if defined(CUTE_ARCH_CP_ASYNC_SM80_ENABLED)
TS const* gmem_ptr = &gmem_src;
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(&smem_dst);
asm volatile("cp.async.ca.shared.global [%0], [%1], %2;\n"
:: "r"(smem_int_ptr),
"l"(gmem_ptr),
"n"(sizeof(TS)));
#elif defined(__MACA_ARCH__)
TS const *gmem_ptr = &gmem_src;
TD *smem_ptr = &smem_dst;
*static_cast<TD *>(smem_ptr) = *static_cast<TD const *>(gmem_ptr);
#else
CUTE_RUNTIME_ASSERT("Support for cp.async instructions has not been enabled");
#endif
}
};
/// Copy via cp.async with caching at global level
template <class TS, class TD = TS>
struct SM80_CP_ASYNC_CACHEGLOBAL
{
using SRegisters = TS[1];
using DRegisters = TD[1];
static_assert(sizeof(TS) == sizeof(TD), "cp.async requires sizeof(src_value_type) == sizeof(dst_value_type)");
static_assert(sizeof(TS) == 4 || sizeof(TS) == 8 || sizeof(TS) == 16, "cp.async sizeof(TS) is not supported");
CUTE_HOST_DEVICE static void
copy(TS const& gmem_src,
TD & smem_dst)
{
#if defined(CUTE_ARCH_CP_ASYNC_SM80_ENABLED)
TS const* gmem_ptr = &gmem_src;
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(&smem_dst);
asm volatile("cp.async.cg.shared.global [%0], [%1], %2;\n"
:: "r"(smem_int_ptr),
"l"(gmem_ptr),
"n"(sizeof(TS)));
#elif defined(__MACA_ARCH__)
TS const *gmem_ptr = &gmem_src;
TD *smem_ptr = &smem_dst;
*static_cast<TD *>(smem_ptr) = *static_cast<TD const *>(gmem_ptr);
#else
CUTE_RUNTIME_ASSERT("Support for cp.async instructions has not been enabled");
#endif
}
};
////////////////////////////////////////////////////////////////////////////////////////////////////
/// Copy via cp.async with caching at global level
template <class TS, class TD = TS>
struct MACA_CP_ASYNC_CACHEGLOBAL
{
using SRegisters = TS[1];
using DRegisters = TD[1];
static_assert(sizeof(TS) == sizeof(TD), "cp.async requires sizeof(src_value_type) == sizeof(dst_value_type)");
static_assert(sizeof(TS) == 4 || sizeof(TS) == 8 || sizeof(TS) == 16, "cp.async sizeof(TS) is not supported");
CUTE_HOST_DEVICE static void
copy(TS const& gmem_src,
TD & smem_dst)
{
#if defined(__MACA_ARCH__)
TS const *gmem_ptr = &gmem_src;
TD *smem_ptr = &smem_dst;
_vmov_cp_async((__attribute__((address_space (3)))int8_t *)(smem_ptr),
(__attribute__((address_space (1)))int8_t *)(gmem_ptr), 0, sizeof(TS));
//__builtin_mxc_arrive(64);
//__builtin_mxc_barrier();
#else
CUTE_RUNTIME_ASSERT("Support for cp.async instructions has not been enabled");
#endif
}
};
////////////////////////////////////////////////////////////////////////////////////////////////////
/// Establishes an ordering w.r.t previously issued cp.async instructions. Does not block.
CUTE_HOST_DEVICE
void
cp_async_fence()
{
#if defined(CUTE_ARCH_CP_ASYNC_SM80_ENABLED)
asm volatile("cp.async.commit_group;\n" ::);
#endif
}
////////////////////////////////////////////////////////////////////////////////////////////////////
/// Blocks until all but N previous cp.async.commit_group operations have committed.
template <int N>
CUTE_HOST_DEVICE
void
cp_async_wait()
{
#if defined(CUTE_ARCH_CP_ASYNC_SM80_ENABLED)
if constexpr (N == 0) {
asm volatile("cp.async.wait_all;\n" ::);
} else {
asm volatile("cp.async.wait_group %0;\n" :: "n"(N));
}
#endif
}
template <int N>
CUTE_HOST_DEVICE
void
cp_async_wait(Int<N>)
{
return cp_async_wait<N>();
}
/////////////////////////////////////////////////////////////////////////////////////////////////
struct MACA_LDS_TRANS_4X16
{
using SRegisters = uint64_t[1];
using DRegisters = uint64_t[1];
CUTE_HOST_DEVICE static void
copy(uint64_t& smem_src,
uint64_t& dst)
{
#if defined(MACA_ARCH_LDS_TRANS_ENABLED)
int64_t *smem_src_ptr = reinterpret_cast<int64_t*>(&smem_src);
dst = __builtin_mxc_load_shared_trans_4x16_i64(smem_src_ptr);
#else
CUTE_RUNTIME_ASSERT("Trying to use lds_b64_trans_4x16 without MACA_ARCH_LDS_TRANS_ENABLED.");
#endif
}
};
} // end namespace cute

View File

@ -0,0 +1,225 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#pragma once
#include <cute/config.hpp>
#include <cute/arch/copy.hpp>
// Config
// #if (defined(__MACA_ARCH__) && (__MACA_ARCH__ >= 900) && (__CUDACC_VER_MAJOR__ >= 12))
// # define CUTE_ARCH_STSM_SM90_ENABLED
// # define CUTE_ARCH_TMA_SM90_ENABLED
// #endif
namespace cute
{
struct SM90_U32x1_STSM_N
{
using SRegisters = uint32_t[1];
using DRegisters = uint128_t[1];
CUTE_HOST_DEVICE static void
copy(uint32_t const& src,
uint128_t & smem_dst)
{
#if defined(CUTE_ARCH_STSM_SM90_ENABLED)
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(&smem_dst);
asm volatile ("stmatrix.sync.aligned.x1.m8n8.shared.b16 [%0], {%1};\n"
:: "r"(smem_int_ptr),
"r"(src));
#else
CUTE_RUNTIME_ASSERT("Trying to use stmatrix without CUTE_ARCH_STSM_SM90_ENABLED.");
#endif
}
};
struct SM90_U32x2_STSM_N
{
using SRegisters = uint32_t[2];
using DRegisters = uint128_t[1];
CUTE_HOST_DEVICE static void
copy(uint32_t const& src0, uint32_t const& src1,
uint128_t& smem_dst)
{
#if defined(CUTE_ARCH_STSM_SM90_ENABLED)
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(&smem_dst);
asm volatile ("stmatrix.sync.aligned.x2.m8n8.shared.b16 [%0], {%1, %2};\n"
:: "r"(smem_int_ptr),
"r"(src0), "r"(src1));
#else
CUTE_RUNTIME_ASSERT("Trying to use stmatrix without CUTE_ARCH_STSM_SM90_ENABLED.");
#endif
}
};
struct SM90_U32x4_STSM_N
{
using SRegisters = uint32_t[4];
using DRegisters = uint128_t[1];
CUTE_HOST_DEVICE static void
copy(uint32_t const& src0, uint32_t const& src1, uint32_t const& src2, uint32_t const& src3,
uint128_t& smem_dst)
{
#if defined(CUTE_ARCH_STSM_SM90_ENABLED)
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(&smem_dst);
asm volatile ("stmatrix.sync.aligned.x4.m8n8.shared.b16 [%0], {%1, %2, %3, %4};\n"
:: "r"(smem_int_ptr),
"r"(src0), "r"(src1), "r"(src2), "r"(src3));
#else
CUTE_RUNTIME_ASSERT("Trying to use stmatrix without CUTE_ARCH_STSM_SM90_ENABLED.");
#endif
}
};
struct SM90_U16x2_STSM_T
{
using SRegisters = uint32_t[1];
using DRegisters = uint128_t[1];
CUTE_HOST_DEVICE static void
copy(uint32_t const& src,
uint128_t& smem_dst)
{
#if defined(CUTE_ARCH_STSM_SM90_ENABLED)
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(&smem_dst);
asm volatile ("stmatrix.sync.aligned.x1.trans.m8n8.shared.b16 [%0], {%1};\n"
:: "r"(smem_int_ptr),
"r"(src));
#else
CUTE_RUNTIME_ASSERT("Trying to use stmatrix without CUTE_ARCH_STSM_SM90_ENABLED.");
#endif
}
};
struct SM90_U16x4_STSM_T
{
using SRegisters = uint32_t[2];
using DRegisters = uint128_t[1];
CUTE_HOST_DEVICE static void
copy(uint32_t const& src0, uint32_t const& src1,
uint128_t& smem_dst)
{
#if defined(CUTE_ARCH_STSM_SM90_ENABLED)
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(&smem_dst);
asm volatile ("stmatrix.sync.aligned.x2.trans.m8n8.shared.b16 [%0], {%1, %2};\n"
:: "r"(smem_int_ptr),
"r"(src0), "r"(src1));
#else
CUTE_RUNTIME_ASSERT("Trying to use stmatrix without CUTE_ARCH_STSM_SM90_ENABLED.");
#endif
}
};
struct SM90_U16x8_STSM_T
{
using SRegisters = uint32_t[4];
using DRegisters = uint128_t[1];
CUTE_HOST_DEVICE static void
copy(uint32_t const& src0, uint32_t const& src1, uint32_t const& src2, uint32_t const& src3,
uint128_t& smem_dst)
{
#if defined(CUTE_ARCH_STSM_SM90_ENABLED)
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(&smem_dst);
asm volatile ("stmatrix.sync.aligned.x4.trans.m8n8.shared.b16 [%0], {%1, %2, %3, %4};\n"
:: "r"(smem_int_ptr),
"r"(src0), "r"(src1), "r"(src2), "r"(src3));
#else
CUTE_RUNTIME_ASSERT("Trying to use stmatrix without CUTE_ARCH_STSM_SM90_ENABLED.");
#endif
}
};
//
// Legacy STSM interfaces that aren't very useful
//
template <class T>
CUTE_HOST_DEVICE
void
copy_stsm(T const* const rmem_ptr,
uint128_t* const smem_ptr)
{
uint32_t const* reg_ptr = reinterpret_cast<uint32_t const*>(rmem_ptr);
// if constexpr
if (sizeof(T) == 4) {
SM90_U32x1_STSM_N::copy(reg_ptr[0], smem_ptr[0]);
}
else if (sizeof(T) == 8) {
SM90_U32x2_STSM_N::copy(reg_ptr[0], reg_ptr[1], smem_ptr[0]);
}
else if (sizeof(T) == 16) {
SM90_U32x4_STSM_N::copy(reg_ptr[0], reg_ptr[1], reg_ptr[2], reg_ptr[3], smem_ptr[0]);
}
else {
static_assert(sizeof(T) == 4 || sizeof(T) == 8 || sizeof(T) == 16, "sizeof(T) is not supported");
}
}
template <class T>
CUTE_HOST_DEVICE
void
copy_stsm_trans(T const* const rmem_ptr,
uint128_t* const smem_ptr)
{
uint32_t const* reg_ptr = reinterpret_cast<uint32_t const*>(rmem_ptr);
// if constexpr
if (sizeof(T) == 4) {
SM90_U16x2_STSM_T::copy(reg_ptr[0], smem_ptr[0]);
}
else if (sizeof(T) == 8) {
SM90_U16x4_STSM_T::copy(reg_ptr[0], reg_ptr[1], smem_ptr[0]);
}
else if (sizeof(T) == 16) {
SM90_U16x8_STSM_T::copy(reg_ptr[0], reg_ptr[1], reg_ptr[2], reg_ptr[3], smem_ptr[0]);
}
else {
static_assert(sizeof(T) == 4 || sizeof(T) == 8 || sizeof(T) == 16, "sizeof(T) is not supported");
}
}
////////////////////////////////////////////////////////////////////////////////////////////////////
} // end namespace cute
////////////////////////////////////////////////////////////////////////////////////////////////////
#include <cute/arch/copy_sm90_desc.hpp>
#include <cute/arch/copy_sm90_tma.hpp>
////////////////////////////////////////////////////////////////////////////////////////////////////

View File

@ -0,0 +1,201 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#pragma once
#if !defined(__MACACC_RTC__)
#include <maca.h>
#include <cinttypes>
#endif
#include <cute/config.hpp>
#include <cute/arch/copy.hpp>
#include <cute/arch/copy_sm90.hpp>
#include <cute/container/alignment.hpp>
#include <cute/container/bit_field.hpp>
#include <cute/numeric/int.hpp> // to_Format<[u]intX>
#include <cute/numeric/half.hpp> // to_Format<half_t>
namespace cute
{
//////////////////////////////////////////////////////////////////////////////////////////////////////
/// Barriers are 64-bit of user-managed information used in broadly two types syncronization patterns
/// 1) arrive/wait on threads (usage: cp.async and warp-specialized kernels)
/// 2) transaction-based (usage: TMA transaction where a CTA issues one transaction)
//////////////////////////////////////////////////////////////////////////////////////////////////////
// Initialize barrier present in shared memory
CUTE_HOST_DEVICE
void
initialize_barrier(uint64_t& smem_barrier, // 64 bits user-manged barrier in smem
int thread_count = 1) // Thread count expected to arrive/wait on this barrier
{
#if defined(CUTE_ARCH_TMA_SM90_ENABLED)
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(&smem_barrier);
asm volatile ("mbarrier.init.shared.b64 [%0], %1;\n"
:: "r"(smem_int_ptr),
"r"(thread_count));
#endif
}
// Set the number of bytes transfered per transaction
CUTE_HOST_DEVICE
void
set_barrier_transaction_bytes(uint64_t& smem_barrier, // 64 bits user-manged barrier in smem
uint32_t bytes) // Number of bytes transfered by per TMA transaction
{
#if defined(CUTE_ARCH_TMA_SM90_ENABLED)
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(&smem_barrier);
asm volatile ("mbarrier.arrive.expect_tx.shared.b64 _, [%0], %1;\n"
:: "r"(smem_int_ptr),
"r"(bytes));
#endif
}
// Barrier wait
CUTE_HOST_DEVICE
void
wait_barrier(uint64_t& smem_barrier, // 64 bits user-manged barrier in smem
int phase_bit) // Current phase bit the barrier waiting to flip
{
#if defined(CUTE_ARCH_TMA_SM90_ENABLED)
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(&smem_barrier);
asm volatile(
"{\n"
".reg .pred P1;\n"
"LAB_WAIT:\n"
"mbarrier.try_wait.parity.shared.b64 P1, [%0], %1;\n"
"@P1 bra.uni DONE;\n"
"bra.uni LAB_WAIT;\n"
"DONE:\n"
"}\n"
:: "r"(smem_int_ptr),
"r"(phase_bit));
#endif
}
// Barrier arrive
CUTE_HOST_DEVICE
void
arrive_barrier(uint64_t& smem_barrier) // 64 bits user-manged barrier in smem
{
#if defined(CUTE_ARCH_TMA_SM90_ENABLED)
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(&smem_barrier);
asm volatile(
"{\n"
".reg .b64 state; \n"
"mbarrier.arrive.shared.b64 state, [%0];\n"
"}\n"
:: "r"(smem_int_ptr));
#endif
}
////////////////////////////////////////////////////////////////////////////////////////////////////
// TMA Descriptor and utilities
////////////////////////////////////////////////////////////////////////////////////////////////////
namespace TMA {
enum class SmemSwizzleBits : uint8_t {
DISABLE = 0,
B32 = 1,
B64 = 2,
B128 = 3,
};
#if !defined(__MACACC_RTC__)
// #if (__CUDACC_VER_MAJOR__ >= 12)
#if 0
template <class T>
inline CUtensorMapDataType to_CUtensorMapDataType() {
if constexpr (is_same<T, int8_t>::value) { return CU_TENSOR_MAP_DATA_TYPE_UINT8; } else
if constexpr (is_same<T, uint8_t>::value) { return CU_TENSOR_MAP_DATA_TYPE_UINT8; } else
if constexpr (is_same<T, uint16_t>::value) { return CU_TENSOR_MAP_DATA_TYPE_UINT16; } else
if constexpr (is_same<T, uint32_t>::value) { return CU_TENSOR_MAP_DATA_TYPE_UINT32; } else
if constexpr (is_same<T, uint64_t>::value) { return CU_TENSOR_MAP_DATA_TYPE_UINT64; } else
if constexpr (is_same<T, int32_t>::value) { return CU_TENSOR_MAP_DATA_TYPE_INT32; } else
if constexpr (is_same<T, int64_t>::value) { return CU_TENSOR_MAP_DATA_TYPE_INT64; } else
if constexpr (is_same<T, half_t>::value) { return CU_TENSOR_MAP_DATA_TYPE_FLOAT16; } else
if constexpr (is_same<T, float>::value) { return CU_TENSOR_MAP_DATA_TYPE_FLOAT32; } else
if constexpr (is_same<T, double>::value) { return CU_TENSOR_MAP_DATA_TYPE_FLOAT64; } else
if constexpr (is_same<T, bfloat16_t>::value) { return CU_TENSOR_MAP_DATA_TYPE_BFLOAT16; } else
if constexpr (is_same<T, tfloat32_t>::value) { return CU_TENSOR_MAP_DATA_TYPE_TFLOAT32; } else
{ static_assert(sizeof(T) < 0, "Unknown TMA Format!"); }
}
inline CUtensorMapSwizzle to_CUtensorMapSwizzle(SmemSwizzleBits const& t) {
switch (t) {
default: assert(false && "Unknown SmemSwizzleBits!");
case SmemSwizzleBits::DISABLE: return CU_TENSOR_MAP_SWIZZLE_NONE;
case SmemSwizzleBits::B32: return CU_TENSOR_MAP_SWIZZLE_32B;
case SmemSwizzleBits::B64: return CU_TENSOR_MAP_SWIZZLE_64B;
case SmemSwizzleBits::B128: return CU_TENSOR_MAP_SWIZZLE_128B;
}
}
#endif // (__CUDACC_VER_MAJOR__ >= 12)
#endif // !defined(__MACACC_RTC__)
} // end namespace TMA
// #if (__CUDACC_VER_MAJOR__ >= 12) && !defined(__MACACC_RTC__)
#if 0
using TmaDescriptor = CUtensorMap;
#else
using TmaDescriptor = struct { char bytes[128]; };
#endif
////////////////////////////////////////////////////////////////////////////////////////////////////
/// Initiates a TensorMap Prefetch
////////////////////////////////////////////////////////////////////////////////////////////////////
CUTE_HOST_DEVICE
void
prefetch_tma_descriptor(TmaDescriptor const* desc_ptr)
{
#if defined(CUTE_ARCH_TMA_SM90_ENABLED)
uint64_t gmem_int_desc = reinterpret_cast<uint64_t>(desc_ptr);
// Prefetch TMA Descriptor using generic addressing (i.e. no specific state space: const or param)
asm volatile (
"prefetch.tensormap [%0];"
:
: "l"(gmem_int_desc)
: "memory");
#else
CUTE_RUNTIME_ASSERT("Trying to use TMA Descriptor Prefetch without CUTE_ARCH_TMA_SM90_ENABLED.");
#endif
}
///////////////////////////////////////////////////////////////////////////////
} // end namespace cute

View File

@ -0,0 +1,861 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#pragma once
#include <cute/config.hpp>
#include <cute/arch/copy.hpp>
#include <cute/arch/copy_sm90.hpp>
namespace cute
{
////////////////////////////////////////////////////////////////////////////////////////////////////
/// TMA_LOAD : Initiates a TMA copy from global memory to shared memory
////////////////////////////////////////////////////////////////////////////////////////////////////
struct SM90_TMA_LOAD_1D
{
CUTE_HOST_DEVICE static void
copy(void const* const desc_ptr, uint64_t& smem_mbar,
void const* const smem_ptr,
int32_t const& crd0)
{
#if defined(CUTE_ARCH_TMA_SM90_ENABLED)
uint64_t gmem_int_desc = reinterpret_cast<uint64_t>(desc_ptr);
uint32_t smem_int_mbar = cast_smem_ptr_to_uint(&smem_mbar);
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(smem_ptr);
asm volatile (
"cp.async.bulk.tensor.1d.shared::cluster.global.mbarrier::complete_tx::bytes"
" [%0], [%1, {%3}], [%2];"
:
: "r"(smem_int_ptr), "l"(gmem_int_desc), "r"(smem_int_mbar),
"r"(crd0)
: "memory");
#else
CUTE_RUNTIME_ASSERT("Trying to use tma without CUTE_ARCH_TMA_SM90_ENABLED.");
#endif
}
};
struct SM90_TMA_LOAD_2D
{
CUTE_HOST_DEVICE static void
copy(void const* const desc_ptr, uint64_t& smem_mbar,
void const* const smem_ptr,
int32_t const& crd0, int32_t const& crd1)
{
#if defined(CUTE_ARCH_TMA_SM90_ENABLED)
uint64_t gmem_int_desc = reinterpret_cast<uint64_t>(desc_ptr);
uint32_t smem_int_mbar = cast_smem_ptr_to_uint(&smem_mbar);
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(smem_ptr);
asm volatile (
"cp.async.bulk.tensor.2d.shared::cluster.global.mbarrier::complete_tx::bytes"
" [%0], [%1, {%3, %4}], [%2];"
:
: "r"(smem_int_ptr), "l"(gmem_int_desc), "r"(smem_int_mbar),
"r"(crd0), "r"(crd1)
: "memory");
#else
CUTE_RUNTIME_ASSERT("Trying to use tma without CUTE_ARCH_TMA_SM90_ENABLED.");
#endif
}
};
struct SM90_TMA_LOAD_3D
{
CUTE_HOST_DEVICE static void
copy(void const* const desc_ptr, uint64_t& smem_mbar,
void const* const smem_ptr,
int32_t const& crd0, int32_t const& crd1, int32_t const& crd2)
{
#if defined(CUTE_ARCH_TMA_SM90_ENABLED)
uint64_t gmem_int_desc = reinterpret_cast<uint64_t>(desc_ptr);
uint32_t smem_int_mbar = cast_smem_ptr_to_uint(&smem_mbar);
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(smem_ptr);
asm volatile (
"cp.async.bulk.tensor.3d.shared::cluster.global.mbarrier::complete_tx::bytes"
" [%0], [%1, {%3, %4, %5}], [%2];"
:
: "r"(smem_int_ptr), "l"(gmem_int_desc), "r"(smem_int_mbar),
"r"(crd0), "r"(crd1), "r"(crd2)
: "memory");
#else
CUTE_RUNTIME_ASSERT("Trying to use tma without CUTE_ARCH_TMA_SM90_ENABLED.");
#endif
}
};
struct SM90_TMA_LOAD_4D
{
CUTE_HOST_DEVICE static void
copy(void const* const desc_ptr, uint64_t& smem_mbar,
void const* const smem_ptr,
int32_t const& crd0, int32_t const& crd1, int32_t const& crd2, int32_t const& crd3)
{
#if defined(CUTE_ARCH_TMA_SM90_ENABLED)
uint64_t gmem_int_desc = reinterpret_cast<uint64_t>(desc_ptr);
uint32_t smem_int_mbar = cast_smem_ptr_to_uint(&smem_mbar);
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(smem_ptr);
asm volatile (
"cp.async.bulk.tensor.4d.shared::cluster.global.mbarrier::complete_tx::bytes"
" [%0], [%1, {%3, %4, %5, %6}], [%2];"
:
: "r"(smem_int_ptr), "l"(gmem_int_desc), "r"(smem_int_mbar),
"r"(crd0), "r"(crd1), "r"(crd2), "r"(crd3)
: "memory");
#else
CUTE_RUNTIME_ASSERT("Trying to use tma without CUTE_ARCH_TMA_SM90_ENABLED.");
#endif
}
};
struct SM90_TMA_LOAD_5D
{
CUTE_HOST_DEVICE static void
copy(void const* const desc_ptr, uint64_t& smem_mbar,
void const* const smem_ptr,
int32_t const& crd0, int32_t const& crd1, int32_t const& crd2, int32_t const& crd3, int32_t const& crd4)
{
#if defined(CUTE_ARCH_TMA_SM90_ENABLED)
uint64_t gmem_int_desc = reinterpret_cast<uint64_t>(desc_ptr);
uint32_t smem_int_mbar = cast_smem_ptr_to_uint(&smem_mbar);
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(smem_ptr);
asm volatile (
"cp.async.bulk.tensor.5d.shared::cluster.global.mbarrier::complete_tx::bytes"
" [%0], [%1, {%3, %4, %5, %6, %7}], [%2];"
:
: "r"(smem_int_ptr), "l"(gmem_int_desc), "r"(smem_int_mbar),
"r"(crd0), "r"(crd1), "r"(crd2), "r"(crd3), "r"(crd4)
: "memory");
#else
CUTE_RUNTIME_ASSERT("Trying to use tma without CUTE_ARCH_TMA_SM90_ENABLED.");
#endif
}
};
struct SM90_TMA_LOAD
{
CUTE_HOST_DEVICE static void
copy(void const* const desc_ptr, uint64_t& smem_mbar,
void const* const smem_ptr,
int32_t const& crd0)
{
return SM90_TMA_LOAD_1D::copy(desc_ptr, smem_mbar, smem_ptr, crd0);
}
CUTE_HOST_DEVICE static void
copy(void const* const desc_ptr, uint64_t& smem_mbar,
void const* const smem_ptr,
int32_t const& crd0, int32_t const& crd1)
{
return SM90_TMA_LOAD_2D::copy(desc_ptr, smem_mbar, smem_ptr, crd0, crd1);
}
CUTE_HOST_DEVICE static void
copy(void const* const desc_ptr, uint64_t& smem_mbar,
void const* const smem_ptr,
int32_t const& crd0, int32_t const& crd1, int32_t const& crd2)
{
return SM90_TMA_LOAD_3D::copy(desc_ptr, smem_mbar, smem_ptr, crd0, crd1, crd2);
}
CUTE_HOST_DEVICE static void
copy(void const* const desc_ptr, uint64_t& smem_mbar,
void const* const smem_ptr,
int32_t const& crd0, int32_t const& crd1, int32_t const& crd2, int32_t const& crd3)
{
return SM90_TMA_LOAD_4D::copy(desc_ptr, smem_mbar, smem_ptr, crd0, crd1, crd2, crd3);
}
CUTE_HOST_DEVICE static void
copy(void const* const desc_ptr, uint64_t& smem_mbar,
void const* const smem_ptr,
int32_t const& crd0, int32_t const& crd1, int32_t const& crd2, int32_t const& crd3, int32_t const& crd4)
{
return SM90_TMA_LOAD_5D::copy(desc_ptr, smem_mbar, smem_ptr, crd0, crd1, crd2, crd3, crd4);
}
};
////////////////////////////////////////////////////////////////////////////////////////////////////
/// TMA_LOAD im2col: Initiates a TMA copy, in im2col mode, from global memory to shared memory
////////////////////////////////////////////////////////////////////////////////////////////////////
struct SM90_TMA_LOAD_IM2COL_3D
{
CUTE_HOST_DEVICE static void
copy(void const* const desc_ptr, uint64_t& smem_mbar,
void const* const smem_ptr,
int32_t const& coord_c, int32_t const& coord_w, int32_t const& coord_n,
uint16_t const& offset_w)
{
#if defined(CUTE_ARCH_TMA_SM90_ENABLED)
uint64_t gmem_int_desc = reinterpret_cast<uint64_t>(desc_ptr);
uint32_t smem_int_mbar = cast_smem_ptr_to_uint(&smem_mbar);
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(smem_ptr);
// Copy from global to shared::cluster.
asm volatile (
"cp.async.bulk.tensor.3d.shared::cluster.global.im2col.mbarrier::complete_tx::bytes"
" [%0], [%1, {%3, %4, %5}], [%2], {%6};"
:
: "r"(smem_int_ptr), "l"(gmem_int_desc), "r"(smem_int_mbar),
"r"(coord_c), "r"(coord_w), "r"(coord_n),
"h"(offset_w)
: "memory");
#else
CUTE_RUNTIME_ASSERT("Trying to use tma without CUTE_ARCH_TMA_SM90_ENABLED.");
#endif
}
};
struct SM90_TMA_LOAD_IM2COL_4D
{
CUTE_HOST_DEVICE static void
copy(void const* const desc_ptr, uint64_t& smem_mbar,
void const* const smem_ptr,
int32_t const& coord_c, int32_t const& coord_w, int32_t const& coord_h, int32_t const& coord_n,
uint16_t const& offset_w,
uint16_t const& offset_h)
{
#if defined(CUTE_ARCH_TMA_SM90_ENABLED)
uint64_t gmem_int_desc = reinterpret_cast<uint64_t>(desc_ptr);
uint32_t smem_int_mbar = cast_smem_ptr_to_uint(&smem_mbar);
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(smem_ptr);
// Copy from global to shared::cluster.
asm volatile (
"cp.async.bulk.tensor.4d.shared::cluster.global.im2col.mbarrier::complete_tx::bytes"
" [%0], [%1, {%3, %4, %5, %6}], [%2], {%7, %8};"
:
: "r"(smem_int_ptr), "l"(gmem_int_desc), "r"(smem_int_mbar),
"r"(coord_c), "r"(coord_w), "r"(coord_h), "r"(coord_n),
"h"(offset_w), "h"(offset_h)
: "memory");
#else
CUTE_RUNTIME_ASSERT("Trying to use tma without CUTE_ARCH_TMA_SM90_ENABLED.");
#endif
}
};
struct SM90_TMA_LOAD_IM2COL_5D
{
CUTE_HOST_DEVICE static void
copy(void const* const desc_ptr, uint64_t& smem_mbar,
void const* const smem_ptr,
int32_t const& coord_c, int32_t const& coord_w, int32_t const& coord_h, int32_t const& coord_d, int32_t const& coord_n,
uint16_t const& offset_w,
uint16_t const& offset_h,
uint16_t const& offset_d)
{
#if defined(CUTE_ARCH_TMA_SM90_ENABLED)
uint64_t gmem_int_desc = reinterpret_cast<uint64_t>(desc_ptr);
uint32_t smem_int_mbar = cast_smem_ptr_to_uint(&smem_mbar);
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(smem_ptr);
// Copy from global to shared::cluster.
asm volatile (
"cp.async.bulk.tensor.5d.shared::cluster.global.im2col.mbarrier::complete_tx::bytes"
" [%0], [%1, {%3, %4, %5, %6, %7}], [%2], {%8, %9, %10};"
:
: "r"(smem_int_ptr), "l"(gmem_int_desc), "r"(smem_int_mbar),
"r"(coord_c), "r"(coord_w), "r"(coord_h), "r"(coord_d), "r"(coord_n),
"h"(offset_w), "h"(offset_h), "h"(offset_d)
: "memory");
#else
CUTE_RUNTIME_ASSERT("Trying to use tma without CUTE_ARCH_TMA_SM90_ENABLED.");
#endif
}
};
struct SM90_TMA_LOAD_IM2COL
{
CUTE_HOST_DEVICE static void
copy(void const* const desc_ptr, uint64_t& smem_mbar,
void const* const smem_ptr,
int32_t const& coord_c, int32_t const& coord_w, int32_t const& coord_n,
uint16_t const& offset_w)
{
return SM90_TMA_LOAD_IM2COL_3D::copy(desc_ptr, smem_mbar, smem_ptr,
coord_c, coord_w, coord_n,
offset_w);
}
CUTE_HOST_DEVICE static void
copy(void const* const desc_ptr, uint64_t& smem_mbar,
void const* const smem_ptr,
int32_t const& coord_c, int32_t const& coord_w, int32_t const& coord_h, int32_t const& coord_n,
uint16_t const& offset_w,
uint16_t const& offset_h)
{
return SM90_TMA_LOAD_IM2COL_4D::copy(desc_ptr, smem_mbar, smem_ptr,
coord_c, coord_w, coord_h, coord_n,
offset_w, offset_h);
}
CUTE_HOST_DEVICE static void
copy(void const* const desc_ptr, uint64_t& smem_mbar,
void const* const smem_ptr,
int32_t const& coord_c, int32_t const& coord_w, int32_t const& coord_h, int32_t const& coord_d, int32_t const& coord_n,
uint16_t const& offset_w,
uint16_t const& offset_h,
uint16_t const& offset_d)
{
return SM90_TMA_LOAD_IM2COL_5D::copy(desc_ptr, smem_mbar, smem_ptr,
coord_c, coord_w, coord_h, coord_d, coord_n,
offset_w, offset_h, offset_d);
}
};
////////////////////////////////////////////////////////////////////////////////////////////////////
/// TMA_LOAD_MULTICAST: Initiates a TMA copy from global memory to shared memory
////////////////////////////////////////////////////////////////////////////////////////////////////
struct SM90_TMA_LOAD_MULTICAST_1D
{
CUTE_HOST_DEVICE static void
copy(void const* const desc_ptr, uint64_t& smem_mbar, uint16_t multicast_mask,
void const* const smem_ptr,
int32_t const& crd0)
{
#if defined(CUTE_ARCH_TMA_SM90_ENABLED)
uint64_t gmem_int_desc = reinterpret_cast<uint64_t>(desc_ptr);
uint32_t smem_int_mbar = cast_smem_ptr_to_uint(&smem_mbar);
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(smem_ptr);
asm volatile (
"cp.async.bulk.tensor.1d.shared::cluster.global.mbarrier::complete_tx::bytes.multicast::cluster"
" [%0], [%1, {%4}], [%2], %3;"
:
: "r"(smem_int_ptr), "l"(gmem_int_desc), "r"(smem_int_mbar),
"h"(multicast_mask),
"r"(crd0)
: "memory");
#else
CUTE_RUNTIME_ASSERT("Trying to use tma without CUTE_ARCH_TMA_SM90_ENABLED.");
#endif
}
};
struct SM90_TMA_LOAD_MULTICAST_2D
{
CUTE_HOST_DEVICE static void
copy(void const* const desc_ptr, uint64_t& smem_mbar, uint16_t multicast_mask,
void const* const smem_ptr,
int32_t const& crd0, int32_t const& crd1)
{
#if defined(CUTE_ARCH_TMA_SM90_ENABLED)
uint64_t gmem_int_desc = reinterpret_cast<uint64_t>(desc_ptr);
uint32_t smem_int_mbar = cast_smem_ptr_to_uint(&smem_mbar);
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(smem_ptr);
asm volatile (
"cp.async.bulk.tensor.2d.shared::cluster.global.mbarrier::complete_tx::bytes.multicast::cluster"
" [%0], [%1, {%4, %5}], [%2], %3;"
:
: "r"(smem_int_ptr), "l"(gmem_int_desc), "r"(smem_int_mbar),
"h"(multicast_mask),
"r"(crd0), "r"(crd1)
: "memory");
#else
CUTE_RUNTIME_ASSERT("Trying to use tma without CUTE_ARCH_TMA_SM90_ENABLED.");
#endif
}
};
struct SM90_TMA_LOAD_MULTICAST_3D
{
CUTE_HOST_DEVICE static void
copy(void const* const desc_ptr, uint64_t& smem_mbar, uint16_t multicast_mask,
void const* const smem_ptr,
int32_t const& crd0, int32_t const& crd1, int32_t const& crd2)
{
#if defined(CUTE_ARCH_TMA_SM90_ENABLED)
uint64_t gmem_int_desc = reinterpret_cast<uint64_t>(desc_ptr);
uint32_t smem_int_mbar = cast_smem_ptr_to_uint(&smem_mbar);
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(smem_ptr);
asm volatile (
"cp.async.bulk.tensor.3d.shared::cluster.global.mbarrier::complete_tx::bytes.multicast::cluster"
" [%0], [%1, {%4, %5, %6}], [%2], %3;"
:
: "r"(smem_int_ptr), "l"(gmem_int_desc), "r"(smem_int_mbar),
"h"(multicast_mask),
"r"(crd0), "r"(crd1), "r"(crd2)
: "memory");
#else
CUTE_RUNTIME_ASSERT("Trying to use tma without CUTE_ARCH_TMA_SM90_ENABLED.");
#endif
}
};
struct SM90_TMA_LOAD_MULTICAST_4D
{
CUTE_HOST_DEVICE static void
copy(void const* const desc_ptr, uint64_t& smem_mbar, uint16_t multicast_mask,
void const* const smem_ptr,
int32_t const& crd0, int32_t const& crd1, int32_t const& crd2, int32_t const& crd3)
{
#if defined(CUTE_ARCH_TMA_SM90_ENABLED)
uint64_t gmem_int_desc = reinterpret_cast<uint64_t>(desc_ptr);
uint32_t smem_int_mbar = cast_smem_ptr_to_uint(&smem_mbar);
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(smem_ptr);
asm volatile (
"cp.async.bulk.tensor.4d.shared::cluster.global.mbarrier::complete_tx::bytes.multicast::cluster"
" [%0], [%1, {%4, %5, %6, %7}], [%2], %3;"
:
: "r"(smem_int_ptr), "l"(gmem_int_desc), "r"(smem_int_mbar),
"h"(multicast_mask),
"r"(crd0), "r"(crd1), "r"(crd2), "r"(crd3)
: "memory");
#else
CUTE_RUNTIME_ASSERT("Trying to use tma without CUTE_ARCH_TMA_SM90_ENABLED.");
#endif
}
};
struct SM90_TMA_LOAD_MULTICAST_5D
{
CUTE_HOST_DEVICE static void
copy(void const* const desc_ptr, uint64_t& smem_mbar, uint16_t multicast_mask,
void const* const smem_ptr,
int32_t const& crd0, int32_t const& crd1, int32_t const& crd2, int32_t const& crd3, int32_t const& crd4)
{
#if defined(CUTE_ARCH_TMA_SM90_ENABLED)
uint64_t gmem_int_desc = reinterpret_cast<uint64_t>(desc_ptr);
uint32_t smem_int_mbar = cast_smem_ptr_to_uint(&smem_mbar);
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(smem_ptr);
asm volatile (
"cp.async.bulk.tensor.5d.shared::cluster.global.mbarrier::complete_tx::bytes.multicast::cluster"
" [%0], [%1, {%4, %5, %6, %7, %8}], [%2], %3;"
:
: "r"(smem_int_ptr), "l"(gmem_int_desc), "r"(smem_int_mbar),
"h"(multicast_mask),
"r"(crd0), "r"(crd1), "r"(crd2), "r"(crd3), "r"(crd4)
: "memory");
#else
CUTE_RUNTIME_ASSERT("Trying to use tma without CUTE_ARCH_TMA_SM90_ENABLED.");
#endif
}
};
struct SM90_TMA_LOAD_MULTICAST
{
CUTE_HOST_DEVICE static void
copy(void const* const desc_ptr, uint64_t& smem_mbar, uint16_t multicast_mask,
void const* const smem_ptr,
int32_t const& crd0)
{
return SM90_TMA_LOAD_MULTICAST_1D::copy(desc_ptr, smem_mbar, multicast_mask, smem_ptr, crd0);
}
CUTE_HOST_DEVICE static void
copy(void const* const desc_ptr, uint64_t& smem_mbar, uint16_t multicast_mask,
void const* const smem_ptr,
int32_t const& crd0, int32_t const& crd1)
{
return SM90_TMA_LOAD_MULTICAST_2D::copy(desc_ptr, smem_mbar, multicast_mask, smem_ptr, crd0, crd1);
}
CUTE_HOST_DEVICE static void
copy(void const* const desc_ptr, uint64_t& smem_mbar, uint16_t multicast_mask,
void const* const smem_ptr,
int32_t const& crd0, int32_t const& crd1, int32_t const& crd2)
{
return SM90_TMA_LOAD_MULTICAST_3D::copy(desc_ptr, smem_mbar, multicast_mask, smem_ptr, crd0, crd1, crd2);
}
CUTE_HOST_DEVICE static void
copy(void const* const desc_ptr, uint64_t& smem_mbar, uint16_t multicast_mask,
void const* const smem_ptr,
int32_t const& crd0, int32_t const& crd1, int32_t const& crd2, int32_t const& crd3)
{
return SM90_TMA_LOAD_MULTICAST_4D::copy(desc_ptr, smem_mbar, multicast_mask, smem_ptr, crd0, crd1, crd2, crd3);
}
CUTE_HOST_DEVICE static void
copy(void const* const desc_ptr, uint64_t& smem_mbar, uint16_t multicast_mask,
void const* const smem_ptr,
int32_t const& crd0, int32_t const& crd1, int32_t const& crd2, int32_t const& crd3, int32_t const& crd4)
{
return SM90_TMA_LOAD_MULTICAST_5D::copy(desc_ptr, smem_mbar, multicast_mask, smem_ptr, crd0, crd1, crd2, crd3, crd4);
}
};
////////////////////////////////////////////////////////////////////////////////////////////////////
/// TMA_LOAD_MULTICAST im2col: Initiates a TMA copy, in im2col mode, from global memory to shared memory
////////////////////////////////////////////////////////////////////////////////////////////////////
struct SM90_TMA_LOAD_IM2COL_MULTICAST_3D
{
CUTE_HOST_DEVICE static void
copy(void const* const desc_ptr, uint64_t& smem_mbar,
uint16_t const& multicast_mask,
void const* const smem_ptr,
int32_t const& coord_c, int32_t const& coord_w, int32_t const& coord_n,
uint16_t const& offset_w)
{
#if defined(CUTE_ARCH_TMA_SM90_ENABLED)
uint64_t gmem_int_desc = reinterpret_cast<uint64_t>(desc_ptr);
uint32_t smem_int_mbar = cast_smem_ptr_to_uint(&smem_mbar);
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(smem_ptr);
// Copy from global to shared::cluster.
asm volatile (
"cp.async.bulk.tensor.3d.shared::cluster.global.im2col.mbarrier::complete_tx::bytes.multicast::cluster"
" [%0], [%1, {%3, %4, %5}], [%2], {%6}, %7;"
:
: "r"(smem_int_ptr), "l"(gmem_int_desc), "r"(smem_int_mbar),
"r"(coord_c), "r"(coord_w), "r"(coord_n),
"h"(offset_w),
"h"(multicast_mask)
: "memory");
#else
CUTE_RUNTIME_ASSERT("Trying to use tma without CUTE_ARCH_TMA_SM90_ENABLED.");
#endif
}
};
struct SM90_TMA_LOAD_IM2COL_MULTICAST_4D
{
CUTE_HOST_DEVICE static void
copy(void const* const desc_ptr, uint64_t& smem_mbar,
uint16_t const& multicast_mask,
void const* const smem_ptr,
int32_t const& coord_c, int32_t const& coord_w, int32_t const& coord_h, int32_t const& coord_n,
uint16_t const& offset_w,
uint16_t const& offset_h)
{
#if defined(CUTE_ARCH_TMA_SM90_ENABLED)
uint64_t gmem_int_desc = reinterpret_cast<uint64_t>(desc_ptr);
uint32_t smem_int_mbar = cast_smem_ptr_to_uint(&smem_mbar);
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(smem_ptr);
// Copy from global to shared::cluster.
asm volatile (
"cp.async.bulk.tensor.4d.shared::cluster.global.im2col.mbarrier::complete_tx::bytes"
" [%0], [%1, {%3, %4, %5, %6}], [%2], {%7, %8}, %9;"
:
: "r"(smem_int_ptr), "l"(gmem_int_desc), "r"(smem_int_mbar),
"r"(coord_c), "r"(coord_w), "r"(coord_h), "r"(coord_n),
"h"(offset_w), "h"(offset_h),
"h"(multicast_mask)
: "memory");
#else
CUTE_RUNTIME_ASSERT("Trying to use tma without CUTE_ARCH_TMA_SM90_ENABLED.");
#endif
}
};
struct SM90_TMA_LOAD_IM2COL_MULTICAST_5D
{
CUTE_HOST_DEVICE static void
copy(void const* const desc_ptr, uint64_t& smem_mbar,
uint16_t const& multicast_mask,
void const* const smem_ptr,
int32_t const& coord_c, int32_t const& coord_w, int32_t const& coord_h, int32_t const& coord_d, int32_t const& coord_n,
uint16_t const& offset_w,
uint16_t const& offset_h,
uint16_t const& offset_d)
{
#if defined(CUTE_ARCH_TMA_SM90_ENABLED)
uint64_t gmem_int_desc = reinterpret_cast<uint64_t>(desc_ptr);
uint32_t smem_int_mbar = cast_smem_ptr_to_uint(&smem_mbar);
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(smem_ptr);
// Copy from global to shared::cluster.
asm volatile (
"cp.async.bulk.tensor.5d.shared::cluster.global.im2col.mbarrier::complete_tx::bytes"
" [%0], [%1, {%3, %4, %5, %6, %7}], [%2], {%8, %9, %10}, %11;"
:
: "r"(smem_int_ptr), "l"(gmem_int_desc), "r"(smem_int_mbar),
"r"(coord_c), "r"(coord_w), "r"(coord_h), "r"(coord_d), "r"(coord_n),
"h"(offset_w), "h"(offset_h), "h"(offset_d),
"h"(multicast_mask)
: "memory");
#else
CUTE_RUNTIME_ASSERT("Trying to use tma without CUTE_ARCH_TMA_SM90_ENABLED.");
#endif
}
};
struct SM90_TMA_LOAD_IM2COL_MULTICAST
{
CUTE_HOST_DEVICE static void
copy(void const* const desc_ptr, uint64_t& smem_mbar,
uint16_t const& multicast_mask,
void const* const smem_ptr,
int32_t const& coord_c, int32_t const& coord_w, int32_t const& coord_n,
uint16_t const& offset_w)
{
return SM90_TMA_LOAD_IM2COL_MULTICAST_3D::copy(desc_ptr, smem_mbar,
multicast_mask, smem_ptr,
coord_c, coord_w, coord_n,
offset_w);
}
CUTE_HOST_DEVICE static void
copy(void const* const desc_ptr, uint64_t& smem_mbar,
uint16_t const& multicast_mask,
void const* const smem_ptr,
int32_t const& coord_c, int32_t const& coord_w, int32_t const& coord_h, int32_t const& coord_n,
uint16_t const& offset_w,
uint16_t const& offset_h)
{
return SM90_TMA_LOAD_IM2COL_MULTICAST_4D::copy(desc_ptr, smem_mbar,
multicast_mask, smem_ptr,
coord_c, coord_w, coord_h, coord_n,
offset_w, offset_h);
}
CUTE_HOST_DEVICE static void
copy(void const* const desc_ptr, uint64_t& smem_mbar,
uint16_t const& multicast_mask,
void const* const smem_ptr,
int32_t const& coord_c, int32_t const& coord_w, int32_t const& coord_h, int32_t const& coord_d, int32_t const& coord_n,
uint16_t const& offset_w,
uint16_t const& offset_h,
uint16_t const& offset_d)
{
return SM90_TMA_LOAD_IM2COL_MULTICAST_5D::copy(desc_ptr, smem_mbar,
multicast_mask, smem_ptr,
coord_c, coord_w, coord_h, coord_d, coord_n,
offset_w, offset_h, offset_d);
}
};
////////////////////////////////////////////////////////////////////////////////////////////////////
/// TMA_STORE : Initiates a TMA copy from shared memory to global memory
////////////////////////////////////////////////////////////////////////////////////////////////////
struct SM90_TMA_STORE_1D
{
CUTE_HOST_DEVICE static void
copy(void const* const desc_ptr,
void const* const smem_ptr,
int32_t const& crd0)
{
#if defined(CUTE_ARCH_TMA_SM90_ENABLED)
uint64_t gmem_int_desc = reinterpret_cast<uint64_t>(desc_ptr);
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(smem_ptr);
asm volatile (
"cp.async.bulk.tensor.1d.global.shared::cta.bulk_group [%0, {%2}], [%1];"
:
: "l"(gmem_int_desc), "r"(smem_int_ptr),
"r"(crd0)
: "memory");
#else
CUTE_RUNTIME_ASSERT("Trying to use tma without CUTE_ARCH_TMA_SM90_ENABLED.");
#endif
}
};
struct SM90_TMA_STORE_2D
{
CUTE_HOST_DEVICE static void
copy(void const* const desc_ptr,
void const* const smem_ptr,
int32_t const& crd0, int32_t const& crd1)
{
#if defined(CUTE_ARCH_TMA_SM90_ENABLED)
uint64_t gmem_int_desc = reinterpret_cast<uint64_t>(desc_ptr);
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(smem_ptr);
asm volatile (
"cp.async.bulk.tensor.2d.global.shared::cta.bulk_group [%0, {%2, %3}], [%1];"
:
: "l"(gmem_int_desc), "r"(smem_int_ptr),
"r"(crd0), "r"(crd1)
: "memory");
#else
CUTE_RUNTIME_ASSERT("Trying to use tma without CUTE_ARCH_TMA_SM90_ENABLED.");
#endif
}
};
struct SM90_TMA_STORE_3D
{
CUTE_HOST_DEVICE static void
copy(void const* const desc_ptr,
void const* const smem_ptr,
int32_t const& crd0, int32_t const& crd1, int32_t const& crd2)
{
#if defined(CUTE_ARCH_TMA_SM90_ENABLED)
uint64_t gmem_int_desc = reinterpret_cast<uint64_t>(desc_ptr);
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(smem_ptr);
asm volatile (
"cp.async.bulk.tensor.3d.global.shared::cta.bulk_group [%0, {%2, %3, %4}], [%1];"
:
: "l"(gmem_int_desc), "r"(smem_int_ptr),
"r"(crd0), "r"(crd1), "r"(crd2)
: "memory");
#else
CUTE_RUNTIME_ASSERT("Trying to use tma without CUTE_ARCH_TMA_SM90_ENABLED.");
#endif
}
};
struct SM90_TMA_STORE_4D
{
CUTE_HOST_DEVICE static void
copy(void const* const desc_ptr,
void const* const smem_ptr,
int32_t const& crd0, int32_t const& crd1, int32_t const& crd2, int32_t const& crd3)
{
#if defined(CUTE_ARCH_TMA_SM90_ENABLED)
uint64_t gmem_int_desc = reinterpret_cast<uint64_t>(desc_ptr);
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(smem_ptr);
asm volatile (
"cp.async.bulk.tensor.4d.global.shared::cta.bulk_group [%0, {%2, %3, %4, %5}], [%1];"
:
: "l"(gmem_int_desc), "r"(smem_int_ptr),
"r"(crd0), "r"(crd1), "r"(crd2), "r"(crd3)
: "memory");
#else
CUTE_RUNTIME_ASSERT("Trying to use tma without CUTE_ARCH_TMA_SM90_ENABLED.");
#endif
}
};
struct SM90_TMA_STORE_5D
{
CUTE_HOST_DEVICE static void
copy(void const* const desc_ptr,
void const* const smem_ptr,
int32_t const& crd0, int32_t const& crd1, int32_t const& crd2, int32_t const& crd3, int32_t const& crd4)
{
#if defined(CUTE_ARCH_TMA_SM90_ENABLED)
uint64_t gmem_int_desc = reinterpret_cast<uint64_t>(desc_ptr);
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(smem_ptr);
asm volatile (
"cp.async.bulk.tensor.5d.global.shared::cta.bulk_group [%0, {%2, %3, %4, %5, %6}], [%1];"
:
: "l"(gmem_int_desc), "r"(smem_int_ptr),
"r"(crd0), "r"(crd1), "r"(crd2), "r"(crd3), "r"(crd4)
: "memory");
#else
CUTE_RUNTIME_ASSERT("Trying to use tma without CUTE_ARCH_TMA_SM90_ENABLED.");
#endif
}
};
struct SM90_TMA_STORE
{
CUTE_HOST_DEVICE static void
copy(void const* const desc_ptr,
void const* const smem_ptr,
int32_t const& crd0)
{
return SM90_TMA_STORE_1D::copy(desc_ptr, smem_ptr, crd0);
}
CUTE_HOST_DEVICE static void
copy(void const* const desc_ptr,
void const* const smem_ptr,
int32_t const& crd0, int32_t const& crd1)
{
return SM90_TMA_STORE_2D::copy(desc_ptr, smem_ptr, crd0, crd1);
}
CUTE_HOST_DEVICE static void
copy(void const* const desc_ptr,
void const* const smem_ptr,
int32_t const& crd0, int32_t const& crd1, int32_t const& crd2)
{
return SM90_TMA_STORE_3D::copy(desc_ptr, smem_ptr, crd0, crd1, crd2);
}
CUTE_HOST_DEVICE static void
copy(void const* const desc_ptr,
void const* const smem_ptr,
int32_t const& crd0, int32_t const& crd1, int32_t const& crd2, int32_t const& crd3)
{
return SM90_TMA_STORE_4D::copy(desc_ptr, smem_ptr, crd0, crd1, crd2, crd3);
}
CUTE_HOST_DEVICE static void
copy(void const* const desc_ptr,
void const* const smem_ptr,
int32_t const& crd0, int32_t const& crd1, int32_t const& crd2, int32_t const& crd3, int32_t const& crd4)
{
return SM90_TMA_STORE_5D::copy(desc_ptr, smem_ptr, crd0, crd1, crd2, crd3, crd4);
}
};
// Indicate arrival of warp issuing TMA_STORE
CUTE_HOST_DEVICE static void
tma_store_arrive() {
#if defined(CUTE_ARCH_TMA_SM90_ENABLED)
asm volatile("cp.async.bulk.commit_group;");
#else
CUTE_RUNTIME_ASSERT("Trying to use tma without CUTE_ARCH_TMA_SM90_ENABLED.");
#endif
}
// Wait on prior N (Count) TMA_STORE instructions to complete
template <int Count>
CUTE_HOST_DEVICE static void
tma_store_wait() {
#if defined(CUTE_ARCH_TMA_SM90_ENABLED)
asm volatile(
"cp.async.bulk.wait_group.read %0;"
:
: "n"(Count)
: "memory");
#else
CUTE_RUNTIME_ASSERT("Trying to use tma without CUTE_ARCH_TMA_SM90_ENABLED.");
#endif
}
////////////////////////////////////////////////////////////////////////////////////////////////////
/// BULK_COPY : Copy a bulk of memory between shared memory and global memory
////////////////////////////////////////////////////////////////////////////////////////////////////
struct SM90_BULK_COPY_G2S
{
CUTE_HOST_DEVICE static void
copy(void const* const gmem_ptr, uint64_t& smem_mbar,
void const* const smem_ptr, int32_t load_bytes)
{
#if defined(CUTE_ARCH_TMA_SM90_ENABLED)
uint32_t smem_int_mbar = cast_smem_ptr_to_uint(&smem_mbar);
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(smem_ptr);
asm volatile("cp.async.bulk.shared::cluster.global.mbarrier::complete_tx::bytes [%0], [%1], %2, [%3];\n"
:
: "r"(smem_int_ptr), "l"(gmem_ptr), "r"(load_bytes), "r"(smem_int_mbar)
: "memory");
#else
CUTE_RUNTIME_ASSERT("Trying to use BULK_COPY without CUTE_ARCH_TMA_SM90_ENABLED.");
#endif
}
};
struct SM90_BULK_COPY_S2G
{
CUTE_HOST_DEVICE static void
copy(void const* const smem_ptr,
void const* const gmem_ptr, int32_t store_bytes)
{
#if defined(CUTE_ARCH_TMA_SM90_ENABLED)
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(smem_ptr);
asm volatile("cp.async.bulk.global.shared::cta.bulk_group [%0], [%1], %2;\n"
:
: "l"(gmem_ptr), "r"(smem_int_ptr), "r"(store_bytes)
: "memory");
#else
CUTE_RUNTIME_ASSERT("Trying to use BULK_COPY without CUTE_ARCH_TMA_SM90_ENABLED.");
#endif
}
};
struct SM90_BULK_COPY_AUTO {};
////////////////////////////////////////////////////////////////////////////////////////////////////
} // end namespace cute

View File

@ -0,0 +1,64 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#pragma once
#include <cute/config.hpp>
#include <cute/arch/util.hpp>
namespace cute
{
//
// Direct FMA for any type
//
template <class D, class A = D, class B = A, class C = D>
struct UniversalFMA
{
using DRegisters = D[1];
using ARegisters = A[1];
using BRegisters = B[1];
using CRegisters = C[1];
CUTE_HOST_DEVICE static constexpr void
fma(D & d,
A const& a,
B const& b,
C const& c)
{
// Forward to an ADL/cute free function for these types
using cute::fma;
fma(d, a, b, c);
}
};
} // end namespace cute

View File

@ -0,0 +1,120 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#pragma once
#include <cute/config.hpp>
#include <cute/arch/mma.hpp>
// Config
// #if ((__CUDACC_VER_MAJOR__ > 10) || (__CUDACC_VER_MAJOR__ == 10 && __CUDACC_VER_MINOR__ >= 2))
// # define CUTE_ARCH_MMA_SM75_SUPPORTED
// # if (defined(__MACA_ARCH__) && (__MACA_ARCH__ >= 750))
// # define CUTE_ARCH_MMA_SM75_ENABLED
// # endif
// #endif
namespace cute
{
//
// SM75 MMA 1688 F16F16F32
//
struct SM75_16x8x8_F32F16F16F32_TN
{
using DRegisters = float[4];
using ARegisters = uint32_t[2];
using BRegisters = uint32_t[1];
using CRegisters = float[4];
// Register asm fma
CUTE_HOST_DEVICE static void
fma(float & d0, float & d1, float & d2, float & d3,
uint32_t const& a0, uint32_t const& a1,
uint32_t const& b0,
float const& c0, float const& c1, float const& c2, float const& c3)
{
#if defined(CUTE_ARCH_MMA_SM75_ENABLED)
asm volatile("mma.sync.aligned.m16n8k8.row.col.f32.f16.f16.f32"
"{%0, %1, %2, %3},"
"{%4, %5},"
"{%6},"
"{%7, %8, %9, %10};\n"
: "=f"(d0), "=f"(d1), "=f"(d2), "=f"(d3)
: "r"(a0), "r"(a1),
"r"(b0),
"f"(c0), "f"(c1), "f"(c2), "f"(c3));
#else
CUTE_RUNTIME_ASSERT("Attempting to use SM75_16x8x8_F32F16F16F32_TN without CUTE_ARCH_MMA_SM75_ENABLED");
#endif
}
};
////////////////////////////////////////////////////////////////////////////////////////////////////
//
// SM75 MMA 8816 S8S8S32
//
struct SM75_8x8x16_S32S8S8S32_TN
{
using DRegisters = uint32_t[2];
using ARegisters = uint32_t[1];
using BRegisters = uint32_t[1];
using CRegisters = uint32_t[2];
// Register asm fma
CUTE_HOST_DEVICE static void
fma(uint32_t & d0, uint32_t & d1,
uint32_t const& a0,
uint32_t const& b0,
uint32_t const& c0, uint32_t const& c1)
{
#if defined(CUTE_ARCH_MMA_SM75_ENABLED)
asm volatile("mma.sync.aligned.m8n8k16.row.col.s32.s8.s8.s32"
"{%0, %1},"
"{%2},"
"{%3},"
"{%4, %5};\n"
: "=r"(d0), "=r"(d1)
: "r"(a0),
"r"(b0),
"r"(c0), "r"(c1));
#else
CUTE_RUNTIME_ASSERT("Attempting to use SM75_8x8x16_S32S8S8S32_TN without CUTE_ARCH_MMA_SM75_ENABLED");
#endif
}
};
////////////////////////////////////////////////////////////////////////////////////////////////////
} // end namespace cute

File diff suppressed because it is too large Load Diff

View File

@ -0,0 +1,961 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#pragma once
#include <cute/config.hpp>
#include <cute/arch/mma.hpp>
// Config
// #if defined(__MACA_ARCH__) && (__MACA_ARCH__ >= 900)
// # define CUTE_ARCH_MMA_SM90_ENABLED
// #endif
////////////////////////////////////////////////////////////////////////////////////////////////////
namespace cute {
////////////////////////////////////////////////////////////////////////////////////////////////////
// MMA 16x8x4 TN
struct SM90_16x8x4_F64F64F64F64_TN
{
using DRegisters = double[4];
using ARegisters = double[2];
using BRegisters = double[1];
using CRegisters = double[4];
CUTE_HOST_DEVICE static void
fma(double & d0, double & d1, double & d2, double & d3,
double const& a0, double const& a1,
double const& b0,
double const& c0, double const& c1, double const& c2, double const& c3)
{
#if defined(CUTE_ARCH_MMA_SM90_ENABLED)
asm volatile(
"mma.sync.aligned.m16n8k4.row.col.f64.f64.f64.f64"
"{%0, %1, %2, %3},"
"{%4, %5},"
"{%6},"
"{%7, %8, %9, %10};\n"
: "=d"(d0), "=d"(d1), "=d"(d2), "=d"(d3)
: "d"(a0), "d"(a1),
"d"(b0),
"d"(c0), "d"(c1), "d"(c2), "d"(c3));
#else
CUTE_RUNTIME_ASSERT("Attempting to use SM90_16x8x4_F64F64F64F64_TN without CUTE_ARCH_MMA_SM90_ENABLED");
#endif
}
};
////////////////////////////////////////////////////////////////////////////////////////////////////
// MMA 16x8x8 TN
struct SM90_16x8x8_F64F64F64F64_TN
{
using DRegisters = double[4];
using ARegisters = double[4];
using BRegisters = double[2];
using CRegisters = double[4];
CUTE_HOST_DEVICE static void
fma(double & d0, double & d1, double & d2, double & d3,
double const& a0, double const& a1, double const& a2, double const& a3,
double const& b0, double const& b1,
double const& c0, double const& c1, double const& c2, double const& c3)
{
#if defined(CUTE_ARCH_MMA_SM90_ENABLED)
asm volatile(
"mma.sync.aligned.m16n8k8.row.col.f64.f64.f64.f64"
"{%0, %1, %2, %3},"
"{%4, %5, %6, %7},"
"{%8, %9},"
"{%10, %11, %12, %13};\n"
: "=d"(d0), "=d"(d1), "=d"(d2), "=d"(d3)
: "d"(a0), "d"(a1), "d"(a2), "d"(a3),
"d"(b0), "d"(b1),
"d"(c0), "d"(c1), "d"(c2), "d"(c3));
#else
CUTE_RUNTIME_ASSERT("Attempting to use SM90_16x8x8_F64F64F64F64_TN without CUTE_ARCH_MMA_SM90_ENABLED");
#endif
}
};
////////////////////////////////////////////////////////////////////////////////////////////////////
// MMA 16x8x16 TN
struct SM90_16x8x16_F64F64F64F64_TN
{
using DRegisters = double[4];
using ARegisters = double[8];
using BRegisters = double[4];
using CRegisters = double[4];
CUTE_HOST_DEVICE static void
fma(double & d0, double & d1, double & d2, double & d3,
double const& a0, double const& a1, double const& a2, double const& a3,
double const& a4, double const& a5, double const& a6, double const& a7,
double const& b0, double const& b1, double const& b2, double const& b3,
double const& c0, double const& c1, double const& c2, double const& c3)
{
#if defined(CUTE_ARCH_MMA_SM90_ENABLED)
asm volatile(
"mma.sync.aligned.m16n8k16.row.col.f64.f64.f64.f64"
"{%0, %1, %2, %3},"
"{%4, %5, %6, %7, %8, %9, %10, %11},"
"{%12, %13, %14, %15},"
"{%16, %17, %18, %19};\n"
: "=d"(d0), "=d"(d1), "=d"(d2), "=d"(d3)
: "d"(a0), "d"(a1), "d"(a2), "d"(a3),
"d"(a4), "d"(a5), "d"(a6), "d"(a7),
"d"(b0), "d"(b1), "d"(b2), "d"(b3),
"d"(c0), "d"(c1), "d"(c2), "d"(c3));
#else
CUTE_RUNTIME_ASSERT("Attempting to use SM90_16x8x16_F64F64F64F64_TN without CUTE_ARCH_MMA_SM90_ENABLED");
#endif
}
};
////////////////////////////////////////////////////////////////////////////////////////////////////
// MMA 16x8x4 TN
struct SM90_16x8x4_C64C64C64C64_TN
{
using DRegisters = complex<double>[4];
using ARegisters = complex<double>[2];
using BRegisters = complex<double>[1];
using CRegisters = complex<double>[4];
CUTE_HOST_DEVICE static void
fma(complex<double> & d0, complex<double> & d1,
complex<double> & d2, complex<double> & d3,
complex<double> const& a0, complex<double> const& a1,
complex<double> const& b0,
complex<double> const& c0, complex<double> const& c1,
complex<double> const& c2, complex<double> const& c3)
{
// Because thrust::complex does not provide a mutable ref
double& rd0 = reinterpret_cast<double(&)[2]>(d0)[0];
double& id0 = reinterpret_cast<double(&)[2]>(d0)[1];
double& rd1 = reinterpret_cast<double(&)[2]>(d1)[0];
double& id1 = reinterpret_cast<double(&)[2]>(d1)[1];
double& rd2 = reinterpret_cast<double(&)[2]>(d2)[0];
double& id2 = reinterpret_cast<double(&)[2]>(d2)[1];
double& rd3 = reinterpret_cast<double(&)[2]>(d3)[0];
double& id3 = reinterpret_cast<double(&)[2]>(d3)[1];
// d.real() = a.real() * b.real() + c.real();
SM90_16x8x4_F64F64F64F64_TN::fma(
rd0, rd1, rd2, rd3,
a0.real(), a1.real(),
b0.real(),
c0.real(), c1.real(), c2.real(), c3.real());
// d.imag() = a.imag() * b.real() + c.imag();
SM90_16x8x4_F64F64F64F64_TN::fma(
id0, id1, id2, id3,
a0.imag(), a1.imag(),
b0.real(),
c0.imag(), c1.imag(), c2.imag(), c3.imag());
// d.real() = -a.imag() * b.imag() + d.real();
SM90_16x8x4_F64F64F64F64_TN::fma(
rd0, rd1, rd2, rd3,
-a0.imag(), -a1.imag(),
b0.imag(),
d0.real(), d1.real(), d2.real(), d3.real());
// d.imag() = a.real() * b.imag() + d.imag();
SM90_16x8x4_F64F64F64F64_TN::fma(
id0, id1, id2, id3,
a0.real(), a1.real(),
b0.imag(),
d0.imag(), d1.imag(), d2.imag(), d3.imag());
}
};
////////////////////////////////////////////////////////////////////////////////////////////////////
// MMA 16x8x8 TN
struct SM90_16x8x8_C64C64C64C64_TN
{
using DRegisters = complex<double>[4];
using ARegisters = complex<double>[4];
using BRegisters = complex<double>[2];
using CRegisters = complex<double>[4];
CUTE_HOST_DEVICE static void
fma(complex<double> & d0, complex<double> & d1,
complex<double> & d2, complex<double> & d3,
complex<double> const& a0, complex<double> const& a1,
complex<double> const& a2, complex<double> const& a3,
complex<double> const& b0, complex<double> const& b1,
complex<double> const& c0, complex<double> const& c1,
complex<double> const& c2, complex<double> const& c3)
{
// Because thrust::complex does not provide a mutable ref
double& rd0 = reinterpret_cast<double(&)[2]>(d0)[0];
double& id0 = reinterpret_cast<double(&)[2]>(d0)[1];
double& rd1 = reinterpret_cast<double(&)[2]>(d1)[0];
double& id1 = reinterpret_cast<double(&)[2]>(d1)[1];
double& rd2 = reinterpret_cast<double(&)[2]>(d2)[0];
double& id2 = reinterpret_cast<double(&)[2]>(d2)[1];
double& rd3 = reinterpret_cast<double(&)[2]>(d3)[0];
double& id3 = reinterpret_cast<double(&)[2]>(d3)[1];
// d.real() = a.real() * b.real() + c.real();
SM90_16x8x8_F64F64F64F64_TN::fma(
rd0, rd1, rd2, rd3,
a0.real(), a1.real(), a2.real(), a3.real(),
b0.real(), b1.real(),
c0.real(), c1.real(), c2.real(), c3.real());
// d.imag() = a.imag() * b.real() + c.imag();
SM90_16x8x8_F64F64F64F64_TN::fma(
id0, id1, id2, id3,
a0.imag(), a1.imag(), a2.imag(), a3.imag(),
b0.real(), b1.real(),
c0.imag(), c1.imag(), c2.imag(), c3.imag());
// d.real() = -a.imag() * b.imag() + d.real();
SM90_16x8x8_F64F64F64F64_TN::fma(
rd0, rd1, rd2, rd3,
-a0.imag(), -a1.imag(), -a2.imag(), -a3.imag(),
b0.imag(), b1.imag(),
d0.real(), d1.real(), d2.real(), d3.real());
// d.imag() = a.real() * b.imag() + d.imag();
SM90_16x8x8_F64F64F64F64_TN::fma(
id0, id1, id2, id3,
a0.real(), a1.real(), a2.real(), a3.real(),
b0.imag(), b1.imag(),
d0.imag(), d1.imag(), d2.imag(), d3.imag());
}
};
////////////////////////////////////////////////////////////////////////////////////////////////////
// MMA 16x8x16 TN
struct SM90_16x8x16_C64C64C64C64_TN
{
using DRegisters = complex<double>[4];
using ARegisters = complex<double>[8];
using BRegisters = complex<double>[4];
using CRegisters = complex<double>[4];
CUTE_HOST_DEVICE static void
fma(complex<double> & d0, complex<double> & d1,
complex<double> & d2, complex<double> & d3,
complex<double> const& a0, complex<double> const& a1,
complex<double> const& a2, complex<double> const& a3,
complex<double> const& a4, complex<double> const& a5,
complex<double> const& a6, complex<double> const& a7,
complex<double> const& b0, complex<double> const& b1,
complex<double> const& b2, complex<double> const& b3,
complex<double> const& c0, complex<double> const& c1,
complex<double> const& c2, complex<double> const& c3)
{
// Because thrust::complex does not provide a mutable ref
double& rd0 = reinterpret_cast<double(&)[2]>(d0)[0];
double& id0 = reinterpret_cast<double(&)[2]>(d0)[1];
double& rd1 = reinterpret_cast<double(&)[2]>(d1)[0];
double& id1 = reinterpret_cast<double(&)[2]>(d1)[1];
double& rd2 = reinterpret_cast<double(&)[2]>(d2)[0];
double& id2 = reinterpret_cast<double(&)[2]>(d2)[1];
double& rd3 = reinterpret_cast<double(&)[2]>(d3)[0];
double& id3 = reinterpret_cast<double(&)[2]>(d3)[1];
// d.real() = a.real() * b.real() + c.real();
SM90_16x8x16_F64F64F64F64_TN::fma(
rd0, rd1, rd2, rd3,
a0.real(), a1.real(), a2.real(), a3.real(),
a4.real(), a5.real(), a6.real(), a7.real(),
b0.real(), b1.real(), b2.real(), b3.real(),
c0.real(), c1.real(), c2.real(), c3.real());
// d.imag() = a.imag() * b.real() + c.imag();
SM90_16x8x16_F64F64F64F64_TN::fma(
id0, id1, id2, id3,
a0.imag(), a1.imag(), a2.imag(), a3.imag(),
a4.imag(), a5.imag(), a6.imag(), a7.imag(),
b0.real(), b1.real(), b2.real(), b3.real(),
c0.imag(), c1.imag(), c2.imag(), c3.imag());
// d.real() = -a.imag() * b.imag() + d.real();
SM90_16x8x16_F64F64F64F64_TN::fma(
rd0, rd1, rd2, rd3,
-a0.imag(), -a1.imag(), -a2.imag(), -a3.imag(),
-a4.imag(), -a5.imag(), -a6.imag(), -a7.imag(),
b0.imag(), b1.imag(), b2.imag(), b3.imag(),
d0.real(), d1.real(), d2.real(), d3.real());
// d.imag() = a.real() * b.imag() + d.imag();
SM90_16x8x16_F64F64F64F64_TN::fma(
id0, id1, id2, id3,
a0.real(), a1.real(), a2.real(), a3.real(),
a4.real(), a5.real(), a6.real(), a7.real(),
b0.imag(), b1.imag(), b2.imag(), b3.imag(),
d0.imag(), d1.imag(), d2.imag(), d3.imag());
}
};
////////////////////////////////////////////////////////////////////////////////////////////////////
} // namespace cute
////////////////////////////////////////////////////////////////////////////////////////////////////
#include <cute/arch/mma_sm90_desc.hpp>
#include <cute/arch/mma_sm90_gmma.hpp>
////////////////////////////////////////////////////////////////////////////////////////////////////
namespace cute {
namespace GMMA {
template <
class ElementA,
class ElementB,
class ElementC,
class TileShape_MNK,
GMMA::Major MajorA = GMMA::Major::K,
GMMA::Major MajorB = GMMA::Major::K,
auto... Args // e.g. GMMA::ScaleOut::One, [GMMA::ScaleIn::One, GMMA::ScaleIn::One]
// But most commonly leave empty for defaults
>
CUTE_HOST_DEVICE constexpr
auto
ss_op_selector()
{
static_assert(is_static<TileShape_MNK>::value, "TileShape_MNK must be static.");
static_assert(rank(TileShape_MNK{}) == 3, "TileShape_MNK must be rank 3.");
static_assert(size<0>(TileShape_MNK{}) % 64 == 0, "Tile_M must be a multiple of 64.");
auto Tile_N = size<1>(TileShape_MNK{});
// FP16 accumulator
if constexpr (is_same_v<ElementC, half_t>) {
static_assert(is_same_v<ElementA, half_t>, "Element types for AB must be half if ElementC is half.");
static_assert(is_same_v<ElementB, half_t>, "Element types for AB must be half if ElementC is half.");
static_assert(size<2>(TileShape_MNK{}) % 16 == 0, "Tile_K must be a multiple of 16.");
// Dispatch against the Tile N mode size
if constexpr (Tile_N % 256 == 0) {
return SM90_64x256x16_F16F16F16_SS<MajorA, MajorB, Args...>{};
}
else if constexpr (Tile_N % 192 == 0) {
return SM90_64x192x16_F16F16F16_SS<MajorA, MajorB, Args...>{};
}
else if constexpr (Tile_N % 128 == 0) {
return SM90_64x128x16_F16F16F16_SS<MajorA, MajorB, Args...>{};
}
else if constexpr (Tile_N % 96 == 0) {
return SM90_64x96x16_F16F16F16_SS<MajorA, MajorB, Args...>{};
}
else if constexpr (Tile_N % 64 == 0) {
return SM90_64x64x16_F16F16F16_SS<MajorA, MajorB, Args...>{};
}
else if constexpr (Tile_N % 32 == 0) {
return SM90_64x32x16_F16F16F16_SS<MajorA, MajorB, Args...>{};
}
else if constexpr (Tile_N % 16 == 0) {
return SM90_64x16x16_F16F16F16_SS<MajorA, MajorB, Args...>{};
}
else if constexpr (Tile_N % 8 == 0) {
return SM90_64x8x16_F16F16F16_SS<MajorA, MajorB, Args...>{};
}
else {
static_assert(Tile_N % 8 == 0, "Tile_N must be a multiple of 8.");
}
}
// FP32 accumulator
else if constexpr (is_same_v<ElementC, float>) {
// FP16 inputs
if constexpr (is_same_v<ElementA, half_t>) {
static_assert(is_same_v<ElementA, ElementB>, "ElementA and ElementB must be the same type for this config.");
static_assert(size<2>(TileShape_MNK{}) % 16 == 0, "Tile_K must be a multiple of 16.");
if constexpr (Tile_N % 256 == 0) {
return SM90_64x256x16_F32F16F16_SS<MajorA, MajorB, Args...>{};
}
else if constexpr (Tile_N % 192 == 0) {
return SM90_64x192x16_F32F16F16_SS<MajorA, MajorB, Args...>{};
}
else if constexpr (Tile_N % 128 == 0) {
return SM90_64x128x16_F32F16F16_SS<MajorA, MajorB, Args...>{};
}
else if constexpr (Tile_N % 96 == 0) {
return SM90_64x96x16_F32F16F16_SS<MajorA, MajorB, Args...>{};
}
else if constexpr (Tile_N % 64 == 0) {
return SM90_64x64x16_F32F16F16_SS<MajorA, MajorB, Args...>{};
}
else if constexpr (Tile_N % 32 == 0) {
return SM90_64x32x16_F32F16F16_SS<MajorA, MajorB, Args...>{};
}
else if constexpr (Tile_N % 16 == 0) {
return SM90_64x16x16_F32F16F16_SS<MajorA, MajorB, Args...>{};
}
else if constexpr (Tile_N % 8 == 0) {
return SM90_64x8x16_F32F16F16_SS<MajorA, MajorB, Args...>{};
}
else {
static_assert(Tile_N % 8 == 0, "Tile_N must be a multiple of 8.");
}
}
// BF16 inputs
else if constexpr (is_same_v<ElementA, bfloat16_t>) {
static_assert(is_same_v<ElementA, ElementB>, "ElementA and ElementB must be the same type for this config.");
static_assert(size<2>(TileShape_MNK{}) % 16 == 0, "Tile_K must be a multiple of 16.");
if constexpr (Tile_N % 256 == 0) {
return SM90_64x256x16_F32BF16BF16_SS<MajorA, MajorB, Args...>{};
}
else if constexpr (Tile_N % 192 == 0) {
return SM90_64x192x16_F32BF16BF16_SS<MajorA, MajorB, Args...>{};
}
else if constexpr (Tile_N % 128 == 0) {
return SM90_64x128x16_F32BF16BF16_SS<MajorA, MajorB, Args...>{};
}
else if constexpr (Tile_N % 96 == 0) {
return SM90_64x96x16_F32BF16BF16_SS<MajorA, MajorB, Args...>{};
}
else if constexpr (Tile_N % 64 == 0) {
return SM90_64x64x16_F32BF16BF16_SS<MajorA, MajorB, Args...>{};
}
else if constexpr (Tile_N % 32 == 0) {
return SM90_64x32x16_F32BF16BF16_SS<MajorA, MajorB, Args...>{};
}
else if constexpr (Tile_N % 16 == 0) {
return SM90_64x16x16_F32BF16BF16_SS<MajorA, MajorB, Args...>{};
}
else if constexpr (Tile_N % 8 == 0) {
return SM90_64x8x16_F32BF16BF16_SS<MajorA, MajorB, Args...>{};
}
else {
static_assert(Tile_N % 8 == 0, "Tile_N must be a multiple of 8.");
}
}
// TF32 inputs
else if constexpr (is_same_v<ElementA, tfloat32_t>) {
static_assert(is_same_v<ElementA, ElementB>, "ElementA and ElementB must be the same type for this config.");
static_assert(MajorA == GMMA::Major::K, "MajorA must be GMMA::Major::K for this config.");
static_assert(MajorB == GMMA::Major::K, "MajorB must be GMMA::Major::K for this config.");
static_assert(size<2>(TileShape_MNK{}) % 8 == 0, "Tile_K must be a multiple of 8.");
if constexpr (Tile_N % 256 == 0) {
return SM90_64x256x8_F32TF32TF32_SS_TN<Args...>{};
}
else if constexpr (Tile_N % 192 == 0) {
return SM90_64x192x8_F32TF32TF32_SS_TN<Args...>{};
}
else if constexpr (Tile_N % 128 == 0) {
return SM90_64x128x8_F32TF32TF32_SS_TN<Args...>{};
}
else if constexpr (Tile_N % 96 == 0) {
return SM90_64x96x8_F32TF32TF32_SS_TN<Args...>{};
}
else if constexpr (Tile_N % 64 == 0) {
return SM90_64x64x8_F32TF32TF32_SS_TN<Args...>{};
}
else if constexpr (Tile_N % 32 == 0) {
return SM90_64x32x8_F32TF32TF32_SS_TN<Args...>{};
}
else if constexpr (Tile_N % 16 == 0) {
return SM90_64x16x8_F32TF32TF32_SS_TN<Args...>{};
}
else if constexpr (Tile_N % 8 == 0) {
return SM90_64x8x8_F32TF32TF32_SS_TN<Args...>{};
}
else {
static_assert(Tile_N % 8 == 0, "Tile_N must be a multiple of 8.");
}
}
else {
static_assert(sizeof(ElementA) == 0, "No eligible GMMA operator for request configuration.");
}
}
// S32 accumulator
else if constexpr (is_same_v<ElementC, int32_t>) {
static_assert(MajorA == GMMA::Major::K, "MajorA must be GMMA::Major::K for this config.");
static_assert(MajorB == GMMA::Major::K, "MajorB must be GMMA::Major::K for this config.");
static_assert(size<2>(TileShape_MNK{}) % 32 == 0, "Tile_K must be a multiple of 32.");
// ElementA == int8_t && ElementB == int8_t
if constexpr (is_same_v<ElementA, int8_t> && is_same_v<ElementB, int8_t>) {
if constexpr (Tile_N % 256 == 0) {
return SM90_64x256x32_S32S8S8_SS_TN{};
}
else if constexpr (Tile_N % 192 == 0) {
return SM90_64x192x32_S32S8S8_SS_TN{};
}
else if constexpr (Tile_N % 128 == 0) {
return SM90_64x128x32_S32S8S8_SS_TN{};
}
else if constexpr (Tile_N % 96 == 0) {
return SM90_64x96x32_S32S8S8_SS_TN{};
}
else if constexpr (Tile_N % 64 == 0) {
return SM90_64x64x32_S32S8S8_SS_TN{};
}
else if constexpr (Tile_N % 32 == 0) {
return SM90_64x32x32_S32S8S8_SS_TN{};
}
else if constexpr (Tile_N % 16 == 0) {
return SM90_64x16x32_S32S8S8_SS_TN{};
}
else if constexpr (Tile_N % 8 == 0) {
return SM90_64x8x32_S32S8S8_SS_TN{};
}
else {
static_assert(Tile_N % 8 == 0, "Tile_N must be a multiple of 8.");
}
}
// ElementA == int8_t && ElementB == uint8_t
else if constexpr (is_same_v<ElementA, int8_t> && is_same_v<ElementB, uint8_t>) {
static_assert(size<2>(TileShape_MNK{}) % 32 == 0, "Tile_K must be a multiple of 32.");
if constexpr (Tile_N % 256 == 0) {
return SM90_64x256x32_S32S8U8_SS_TN{};
}
else if constexpr (Tile_N % 192 == 0) {
return SM90_64x192x32_S32S8U8_SS_TN{};
}
else if constexpr (Tile_N % 128 == 0) {
return SM90_64x128x32_S32S8U8_SS_TN{};
}
else if constexpr (Tile_N % 96 == 0) {
return SM90_64x96x32_S32S8U8_SS_TN{};
}
else if constexpr (Tile_N % 64 == 0) {
return SM90_64x64x32_S32S8U8_SS_TN{};
}
else if constexpr (Tile_N % 32 == 0) {
return SM90_64x32x32_S32S8U8_SS_TN{};
}
else if constexpr (Tile_N % 16 == 0) {
return SM90_64x16x32_S32S8U8_SS_TN{};
}
else if constexpr (Tile_N % 8 == 0) {
return SM90_64x8x32_S32S8U8_SS_TN{};
}
else {
static_assert(Tile_N % 8 == 0, "Tile_N must be a multiple of 8.");
}
}
// ElementA == uint8_t && ElementB == int8_t
else if constexpr (is_same_v<ElementA, uint8_t> && is_same_v<ElementB, int8_t>) {
static_assert(size<2>(TileShape_MNK{}) % 32 == 0, "Tile_K must be a multiple of 32.");
if constexpr (Tile_N % 256 == 0) {
return SM90_64x256x32_S32U8S8_SS_TN{};
}
else if constexpr (Tile_N % 192 == 0) {
return SM90_64x192x32_S32U8S8_SS_TN{};
}
else if constexpr (Tile_N % 128 == 0) {
return SM90_64x128x32_S32U8S8_SS_TN{};
}
else if constexpr (Tile_N % 96 == 0) {
return SM90_64x96x32_S32U8S8_SS_TN{};
}
else if constexpr (Tile_N % 64 == 0) {
return SM90_64x64x32_S32U8S8_SS_TN{};
}
else if constexpr (Tile_N % 32 == 0) {
return SM90_64x32x32_S32U8S8_SS_TN{};
}
else if constexpr (Tile_N % 16 == 0) {
return SM90_64x16x32_S32U8S8_SS_TN{};
}
else if constexpr (Tile_N % 8 == 0) {
return SM90_64x8x32_S32U8S8_SS_TN{};
}
else {
static_assert(Tile_N % 8 == 0, "Tile_N must be a multiple of 8.");
}
}
// ElementA == uint8_t && ElementB == uint8_t
else if constexpr (is_same_v<ElementA, uint8_t> && is_same_v<ElementB, uint8_t>) {
static_assert(size<2>(TileShape_MNK{}) % 32 == 0, "Tile_K must be a multiple of 32.");
if constexpr (Tile_N % 256 == 0) {
return SM90_64x256x32_S32U8U8_SS_TN{};
}
else if constexpr (Tile_N % 192 == 0) {
return SM90_64x192x32_S32U8U8_SS_TN{};
}
else if constexpr (Tile_N % 128 == 0) {
return SM90_64x128x32_S32U8U8_SS_TN{};
}
else if constexpr (Tile_N % 96 == 0) {
return SM90_64x96x32_S32U8U8_SS_TN{};
}
else if constexpr (Tile_N % 64 == 0) {
return SM90_64x64x32_S32U8U8_SS_TN{};
}
else if constexpr (Tile_N % 32 == 0) {
return SM90_64x32x32_S32U8U8_SS_TN{};
}
else if constexpr (Tile_N % 16 == 0) {
return SM90_64x16x32_S32U8U8_SS_TN{};
}
else if constexpr (Tile_N % 8 == 0) {
return SM90_64x8x32_S32U8U8_SS_TN{};
}
else {
static_assert(Tile_N % 8 == 0, "Tile_N must be a multiple of 8.");
}
}
}
// Unknown accumulator type
else {
static_assert(sizeof(ElementC) == 0, "Unknown ElementC accumulator type.");
}
}
template <
class ElementA,
class ElementB,
class ElementC,
class TileShape_MNK,
GMMA::Major MajorA = GMMA::Major::K,
GMMA::Major MajorB = GMMA::Major::K,
auto... Args // e.g. GMMA::ScaleOut::One, [GMMA::ScaleIn::One, GMMA::ScaleIn::One]
// But most commonly leave empty for defaults
>
CUTE_HOST_DEVICE constexpr
auto
rs_op_selector()
{
static_assert(is_static<TileShape_MNK>::value, "TileShape_MNK must be static.");
static_assert(rank(TileShape_MNK{}) == 3, "TileShape_MNK must be rank 3.");
static_assert(size<0>(TileShape_MNK{}) % 64 == 0, "Tile_M must be a multiple of 64.");
static_assert(MajorA == GMMA::Major::K, "Register source A operand GMMAs must have K-major A layout.");
auto Tile_N = size<1>(TileShape_MNK{});
// FP16 accumulator
if constexpr (is_same_v<ElementC, half_t>) {
static_assert(is_same_v<ElementA, half_t>, "Element types for AB must be half if ElementC is half.");
static_assert(is_same_v<ElementB, half_t>, "Element types for AB must be half if ElementC is half.");
static_assert(size<2>(TileShape_MNK{}) % 16 == 0, "Tile_K must be a multiple of 16.");
// Dispatch against the Tile N mode size
if constexpr (Tile_N % 256 == 0) {
return SM90_64x256x16_F16F16F16_RS<MajorA, MajorB, Args...>{};
}
else if constexpr (Tile_N % 192 == 0) {
return SM90_64x192x16_F16F16F16_RS<MajorA, MajorB, Args...>{};
}
else if constexpr (Tile_N % 128 == 0) {
return SM90_64x128x16_F16F16F16_RS<MajorA, MajorB, Args...>{};
}
else if constexpr (Tile_N % 96 == 0) {
return SM90_64x96x16_F16F16F16_RS<MajorA, MajorB, Args...>{};
}
else if constexpr (Tile_N % 64 == 0) {
return SM90_64x64x16_F16F16F16_RS<MajorA, MajorB, Args...>{};
}
else if constexpr (Tile_N % 32 == 0) {
return SM90_64x32x16_F16F16F16_RS<MajorA, MajorB, Args...>{};
}
else if constexpr (Tile_N % 16 == 0) {
return SM90_64x16x16_F16F16F16_RS<MajorA, MajorB, Args...>{};
}
else if constexpr (Tile_N % 8 == 0) {
return SM90_64x8x16_F16F16F16_RS<MajorA, MajorB, Args...>{};
}
else {
static_assert(Tile_N % 8 == 0, "Tile_N must be a multiple of 8.");
}
}
// FP32 accumulator
else if constexpr (is_same_v<ElementC, float>) {
static_assert(is_same_v<ElementA, ElementB>, "ElementA and ElementB must be the same type for this config.");
static_assert(size<2>(TileShape_MNK{}) % 16 == 0, "Tile_K must be a multiple of 16.");
// FP16 inputs
if constexpr (is_same_v<ElementA, half_t>) {
if constexpr (Tile_N % 256 == 0) {
return SM90_64x256x16_F32F16F16_RS<MajorA, MajorB, Args...>{};
}
else if constexpr (Tile_N % 192 == 0) {
return SM90_64x192x16_F32F16F16_RS<MajorA, MajorB, Args...>{};
}
else if constexpr (Tile_N % 128 == 0) {
return SM90_64x128x16_F32F16F16_RS<MajorA, MajorB, Args...>{};
}
else if constexpr (Tile_N % 96 == 0) {
return SM90_64x96x16_F32F16F16_RS<MajorA, MajorB, Args...>{};
}
else if constexpr (Tile_N % 64 == 0) {
return SM90_64x64x16_F32F16F16_RS<MajorA, MajorB, Args...>{};
}
else if constexpr (Tile_N % 32 == 0) {
return SM90_64x32x16_F32F16F16_RS<MajorA, MajorB, Args...>{};
}
else if constexpr (Tile_N % 16 == 0) {
return SM90_64x16x16_F32F16F16_RS<MajorA, MajorB, Args...>{};
}
else if constexpr (Tile_N % 8 == 0) {
return SM90_64x8x16_F32F16F16_RS<MajorA, MajorB, Args...>{};
}
else {
static_assert(Tile_N % 8 == 0, "Tile_N must be a multiple of 8.");
}
}
// BF16 inputs
else if constexpr (is_same_v<ElementA, bfloat16_t>) {
static_assert(size<2>(TileShape_MNK{}) % 16 == 0, "Tile_K must be a multiple of 16.");
if constexpr (Tile_N % 256 == 0) {
return SM90_64x256x16_F32BF16BF16_RS<MajorA, MajorB, Args...>{};
}
else if constexpr (Tile_N % 192 == 0) {
return SM90_64x192x16_F32BF16BF16_RS<MajorA, MajorB, Args...>{};
}
else if constexpr (Tile_N % 128 == 0) {
return SM90_64x128x16_F32BF16BF16_RS<MajorA, MajorB, Args...>{};
}
else if constexpr (Tile_N % 96 == 0) {
return SM90_64x96x16_F32BF16BF16_RS<MajorA, MajorB, Args...>{};
}
else if constexpr (Tile_N % 64 == 0) {
return SM90_64x64x16_F32BF16BF16_RS<MajorA, MajorB, Args...>{};
}
else if constexpr (Tile_N % 32 == 0) {
return SM90_64x32x16_F32BF16BF16_RS<MajorA, MajorB, Args...>{};
}
else if constexpr (Tile_N % 16 == 0) {
return SM90_64x16x16_F32BF16BF16_RS<MajorA, MajorB, Args...>{};
}
else if constexpr (Tile_N % 8 == 0) {
return SM90_64x8x16_F32BF16BF16_RS<MajorA, MajorB, Args...>{};
}
else {
static_assert(Tile_N % 8 == 0, "Tile_N must be a multiple of 8.");
}
}
// TF32 inputs
else if constexpr (is_same_v<ElementA, tfloat32_t>) {
static_assert(MajorB == GMMA::Major::K, "MajorB must be GMMA::Major::K for this config.");
static_assert(size<2>(TileShape_MNK{}) % 8 == 0, "Tile_K must be a multiple of 8.");
if constexpr (Tile_N % 256 == 0) {
return SM90_64x256x8_F32TF32TF32_RS_TN<Args...>{};
}
else if constexpr (Tile_N % 192 == 0) {
return SM90_64x192x8_F32TF32TF32_RS_TN<Args...>{};
}
else if constexpr (Tile_N % 128 == 0) {
return SM90_64x128x8_F32TF32TF32_RS_TN<Args...>{};
}
else if constexpr (Tile_N % 96 == 0) {
return SM90_64x96x8_F32TF32TF32_RS_TN<Args...>{};
}
else if constexpr (Tile_N % 64 == 0) {
return SM90_64x64x8_F32TF32TF32_RS_TN<Args...>{};
}
else if constexpr (Tile_N % 32 == 0) {
return SM90_64x32x8_F32TF32TF32_RS_TN<Args...>{};
}
else if constexpr (Tile_N % 16 == 0) {
return SM90_64x16x8_F32TF32TF32_RS_TN<Args...>{};
}
else if constexpr (Tile_N % 8 == 0) {
return SM90_64x8x8_F32TF32TF32_RS_TN<Args...>{};
}
else {
static_assert(Tile_N % 8 == 0, "Tile_N must be a multiple of 8.");
}
}
else {
static_assert(sizeof(ElementA) == 0, "No eligible GMMA operator for request configuration.");
}
}
// S32 accumulator
else if constexpr (is_same_v<ElementC, int32_t>) {
static_assert(MajorB == GMMA::Major::K, "MajorB must be GMMA::Major::K for this config.");
static_assert(size<2>(TileShape_MNK{}) % 32 == 0, "Tile_K must be a multiple of 32.");
// ElementA == int8_t && ElementB == int8_t
if constexpr (is_same_v<ElementA, int8_t> && is_same_v<ElementB, int8_t>) {
if constexpr (Tile_N % 256 == 0) {
return SM90_64x256x32_S32S8S8_RS_TN{};
}
else if constexpr (Tile_N % 192 == 0) {
return SM90_64x192x32_S32S8S8_RS_TN{};
}
else if constexpr (Tile_N % 128 == 0) {
return SM90_64x128x32_S32S8S8_RS_TN{};
}
else if constexpr (Tile_N % 96 == 0) {
return SM90_64x96x32_S32S8S8_RS_TN{};
}
else if constexpr (Tile_N % 64 == 0) {
return SM90_64x64x32_S32S8S8_RS_TN{};
}
else if constexpr (Tile_N % 32 == 0) {
return SM90_64x32x32_S32S8S8_RS_TN{};
}
else if constexpr (Tile_N % 16 == 0) {
return SM90_64x16x32_S32S8S8_RS_TN{};
}
else if constexpr (Tile_N % 8 == 0) {
return SM90_64x8x32_S32S8S8_RS_TN{};
}
else {
static_assert(Tile_N % 8 == 0, "Tile_N must be a multiple of 8.");
}
}
// ElementA == int8_t && ElementB == uint8_t
else if constexpr (is_same_v<ElementA, int8_t> && is_same_v<ElementB, uint8_t>) {
static_assert(size<2>(TileShape_MNK{}) % 32 == 0, "Tile_K must be a multiple of 32.");
if constexpr (Tile_N % 256 == 0) {
return SM90_64x256x32_S32S8U8_RS_TN{};
}
else if constexpr (Tile_N % 192 == 0) {
return SM90_64x192x32_S32S8U8_RS_TN{};
}
else if constexpr (Tile_N % 128 == 0) {
return SM90_64x128x32_S32S8U8_RS_TN{};
}
else if constexpr (Tile_N % 96 == 0) {
return SM90_64x96x32_S32S8U8_RS_TN{};
}
else if constexpr (Tile_N % 64 == 0) {
return SM90_64x64x32_S32S8U8_RS_TN{};
}
else if constexpr (Tile_N % 32 == 0) {
return SM90_64x32x32_S32S8U8_RS_TN{};
}
else if constexpr (Tile_N % 16 == 0) {
return SM90_64x16x32_S32S8U8_RS_TN{};
}
else if constexpr (Tile_N % 8 == 0) {
return SM90_64x8x32_S32S8U8_RS_TN{};
}
else {
static_assert(Tile_N % 8 == 0, "Tile_N must be a multiple of 8.");
}
}
// ElementA == uint8_t && ElementB == int8_t
else if constexpr (is_same_v<ElementA, uint8_t> && is_same_v<ElementB, int8_t>) {
static_assert(size<2>(TileShape_MNK{}) % 32 == 0, "Tile_K must be a multiple of 32.");
if constexpr (Tile_N % 256 == 0) {
return SM90_64x256x32_S32U8S8_RS_TN{};
}
else if constexpr (Tile_N % 192 == 0) {
return SM90_64x192x32_S32U8S8_RS_TN{};
}
else if constexpr (Tile_N % 128 == 0) {
return SM90_64x128x32_S32U8S8_RS_TN{};
}
else if constexpr (Tile_N % 96 == 0) {
return SM90_64x96x32_S32U8S8_RS_TN{};
}
else if constexpr (Tile_N % 64 == 0) {
return SM90_64x64x32_S32U8S8_RS_TN{};
}
else if constexpr (Tile_N % 32 == 0) {
return SM90_64x32x32_S32U8S8_RS_TN{};
}
else if constexpr (Tile_N % 16 == 0) {
return SM90_64x16x32_S32U8S8_RS_TN{};
}
else if constexpr (Tile_N % 8 == 0) {
return SM90_64x8x32_S32U8S8_RS_TN{};
}
else {
static_assert(Tile_N % 8 == 0, "Tile_N must be a multiple of 8.");
}
}
// ElementA == uint8_t && ElementB == uint8_t
else if constexpr (is_same_v<ElementA, uint8_t> && is_same_v<ElementB, uint8_t>) {
static_assert(size<2>(TileShape_MNK{}) % 32 == 0, "Tile_K must be a multiple of 32.");
if constexpr (Tile_N % 256 == 0) {
return SM90_64x256x32_S32U8U8_RS_TN{};
}
else if constexpr (Tile_N % 192 == 0) {
return SM90_64x192x32_S32U8U8_RS_TN{};
}
else if constexpr (Tile_N % 128 == 0) {
return SM90_64x128x32_S32U8U8_RS_TN{};
}
else if constexpr (Tile_N % 96 == 0) {
return SM90_64x96x32_S32U8U8_RS_TN{};
}
else if constexpr (Tile_N % 64 == 0) {
return SM90_64x64x32_S32U8U8_RS_TN{};
}
else if constexpr (Tile_N % 32 == 0) {
return SM90_64x32x32_S32U8U8_RS_TN{};
}
else if constexpr (Tile_N % 16 == 0) {
return SM90_64x16x32_S32U8U8_RS_TN{};
}
else if constexpr (Tile_N % 8 == 0) {
return SM90_64x8x32_S32U8U8_RS_TN{};
}
else {
static_assert(Tile_N % 8 == 0, "Tile_N must be a multiple of 8.");
}
}
}
// Unknown accumulator type
else {
static_assert(sizeof(ElementC) == 0, "Unknown ElementC accumulator type.");
}
}
} // end namespace GMMA
} // end namespace cute
////////////////////////////////////////////////////////////////////////////////////////////////////

View File

@ -0,0 +1,135 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#pragma once
#if !defined(__MACACC_RTC__)
#include <cinttypes>
#endif
#include <cute/config.hpp>
#include <cute/arch/mma.hpp>
////////////////////////////////////////////////////////////////////////////////////////////////////
namespace cute {
////////////////////////////////////////////////////////////////////////////////////////////////////
// GMMA Descriptor and utilities
// GMMA enums and utilities
namespace GMMA
{
enum class LayoutType : uint8_t {
INTERLEAVE = 0,
B128 = 1,
B64 = 2,
B32 = 3,
};
CUTE_HOST_DEVICE char const* to_string(LayoutType const& t) {
switch (t) {
case LayoutType::INTERLEAVE: return "INTERLEAVE";
case LayoutType::B128: return "B128";
case LayoutType::B64: return "B64";
case LayoutType::B32: return "B32";
}
return nullptr;
}
#if !defined(__MACACC_RTC__)
// Output operator for all enums in this namespace
CUTE_HOST std::ostream& operator<<(std::ostream& os, LayoutType const& t) {
char const* s = to_string(t);
if (s) {
std::operator<<(os, s); // Explicit call to avoid ambiguity
} else {
os.setstate(std::ios_base::failbit);
}
return os;
}
#endif // !defined(__MACACC_RTC__)
} // end namespace GMMA
union GmmaDescriptor
{
uint64_t desc_;
uint32_t reg32_[2];
uint16_t reg16_[4];
// Bitfield implementation avoids the need for shifts in assignment
struct {
// start_address, bit [0,14), 4LSB not included
uint16_t start_address_ : 14, : 2; // 14 bits [0,14), 2 bits unused
// leading dimension byte offset, bit [16,30), 4LSB not included
// For N: This is the stride from the first col to the second col of the 8x2 brick in INTERLEAVED
// Unused for all SWIZZLE_* layouts (and assumed to be 1)
// For T: This is the stride from the first 8 rows to the next 8 rows.
uint16_t leading_byte_offset_ : 14, : 2; // 14 bits [0,14), 2 bits unused
// stride dimension byte offset, bit [32,46), 4LSB not included
// For N: This is the stride from the first 8 rows to the next 8 rows.
// For T: This is the stride fro mthe first 8 cols to the next 8 cols.
uint16_t stride_byte_offset_ : 14, : 2; // 14 bits [0,14), 2 bits unused
// base_offset, bit [49,52)
// Valid only for SWIZZLE_128B and SWIZZLE_64B
uint8_t : 1, base_offset_ : 3, : 4; // 1 bit unused, 3 bits [1,4), 4 bits unused
// layout type, bit [62,64)
// SWIZZLE_NONE = 0, SWIZZLE_32B = 3, SWIZZLE_64B = 2, SWIZZLE_128B = 1
uint8_t : 6, layout_type_ : 2; // 6 bits unused, 2 bits [6,8)
};
// Decay to a uint64_t
CUTE_HOST_DEVICE constexpr
operator uint64_t() const noexcept { return desc_; }
// Printer
CUTE_HOST_DEVICE friend void print(GmmaDescriptor const& t)
{
#if !defined(__MACACC_RTC__)
printf("GmmaDescriptor: 0x%016" PRIx64 "\n", t.desc_);
printf(" start_addr : 0x%04x\n", t.start_address_);
printf(" leading_off: 0x%04x (%d)\n", t.leading_byte_offset_, t.leading_byte_offset_);
printf(" stride_off : 0x%04x (%d)\n", t.stride_byte_offset_, t.stride_byte_offset_);
printf(" base_offset: 0x%01x\n", t.base_offset_);
printf(" layout_type: 0x%01x (%s)\n", t.layout_type_, to_string(static_cast<GMMA::LayoutType>(t.layout_type_)));
#endif
}
};
////////////////////////////////////////////////////////////////////////////////////////////////////
} // namespace cute
////////////////////////////////////////////////////////////////////////////////////////////////////

File diff suppressed because it is too large Load Diff

View File

@ -0,0 +1,249 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#pragma once
#include <cute/config.hpp>
#include <cute/numeric/integer_sequence.hpp>
#if defined(__clang__) && defined(__MACA__)
// __cvta_generic_to_shared was added in Clang 14: https://reviews.llvm.org/D111665
#if __clang_major__ >= 14
#define CUTE_CLANG_SUPPORTS_CVTA_GENERIC_TO_SHARED 1
#endif
// __nvvm_get_smem_pointer added in Clang 14: https://reviews.llvm.org/D111665
// ... but will not work on Windows until Clang 15: https://reviews.llvm.org/D122897
#if (!defined(_WIN32) && __clang_major__ >= 14) || __clang_major__ >= 15
#define CUTE_CLANG_SUPPORTS_NVVM_GET_SMEM_POINTER 1
#endif
#endif
#if defined(__MXCC__) || defined(__MACACC_RTC__)
#if 1
#define CUTE_MACA_SUPPORTS_CVTA_GENERIC_TO_SHARED 1
#endif
#if 1
#define CUTE_MACA_SUPPORTS_NVVM_GET_SMEM_POINTER 1
#endif
#endif
#if CUTE_MACA_SUPPORTS_CVTA_GENERIC_TO_SHARED || CUTE_CLANG_SUPPORTS_CVTA_GENERIC_TO_SHARED
#define CUTE_CVTA_GENERIC_TO_SHARED_SUPPORTED 1
#endif
#if !defined(CUTE_CVTA_GENERIC_TO_SHARED_ACTIVATED) && CUTE_CVTA_GENERIC_TO_SHARED_SUPPORTED && defined(__MACA_ARCH__)
#define CUTE_CVTA_GENERIC_TO_SHARED_ACTIVATED 1
#endif
#if CUTE_MACA_SUPPORTS_NVVM_GET_SMEM_POINTER || CUTE_CLANG_SUPPORTS_NVVM_GET_SMEM_POINTER
#define CUTE_NVVM_GET_SMEM_POINTER_SUPPORTED 1
#endif
#if !defined(CUTE_NVVM_GET_SMEM_POINTER_ACTIVATED) && CUTE_NVVM_GET_SMEM_POINTER_SUPPORTED && defined(__MACA_ARCH__)
#define CUTE_NVVM_GET_SMEM_POINTER_ACTIVATED 1
#endif
// Clang 14+ provides a declaration of __nvvm_get_smem_pointer, so we only need
// to provide one for NVCC
#if CUTE_MACA_SUPPORTS_NVVM_GET_SMEM_POINTER
extern "C" {
// This NVVM intrinsic is subject to change in future versions of CUDA.
// Clients should not call it directly.
CUTE_DEVICE uint32_t __nvvm_get_smem_pointer(void*);
}
#endif
namespace cute
{
/// CUTE helper to cast SMEM pointer to unsigned
CUTE_DEVICE
uint32_t
cast_smem_ptr_to_uint(void const* const ptr)
{
// We prefer to use the new CVTA intrinsics if they are available, otherwise we will fall back to
// the previous internal intrinsics if they are available.
#if CUTE_CVTA_GENERIC_TO_SHARED_ACTIVATED
//
// This NVVM intrinsic converts an address in shared memory to a plain
// unsigned integer. This is necessary to pass to shared memory instructions
// in inline PTX.
//
// In CUDA 11 and beyond, this replaces __nvvm_get_smem_pointer() [only available in 10.2].
//
//__device__ size_t __cvta_generic_to_shared(void* ptr);
/// CUTE helper to get SMEM pointer
return static_cast<uint32_t>(__cvta_generic_to_shared(ptr));
#elif CUTE_NVVM_GET_SMEM_POINTER_ACTIVATED
return __nvvm_get_smem_pointer(ptr);
#elif defined(__MACA_ARCH__)
uint32_t smem_ptr;
return smem_ptr;
#else
(void) ptr;
printf("ERROR: cast_smem_ptr_to_uint not supported but used.\n");
return 0;
#endif
}
//
// Utility for pointer interfaces
//
namespace detail {
template <class Fn,
class PtrS, int... Is,
class PtrD, int... Id>
CUTE_HOST_DEVICE constexpr
void
explode(Fn fn,
PtrS&& s, int_sequence<Is...>,
PtrD&& d, int_sequence<Id...>)
{
return fn(s[Is]..., d[Id]...);
}
template <class Fn,
class PtrA, int... Ia,
class PtrB, int... Ib,
class PtrC, int... Ic>
CUTE_HOST_DEVICE constexpr
void
explode(Fn fn,
PtrA&& a, int_sequence<Ia...>,
PtrB&& b, int_sequence<Ib...>,
PtrC&& c, int_sequence<Ic...>)
{
return fn(a[Ia]..., b[Ib]..., c[Ic]...);
}
template <class Fn,
class PtrD, int... Id,
class PtrA, int... Ia,
class PtrB, int... Ib,
class PtrC, int... Ic>
CUTE_HOST_DEVICE constexpr
void
explode(Fn fn,
PtrD&& d, int_sequence<Id...>,
PtrA&& a, int_sequence<Ia...>,
PtrB&& b, int_sequence<Ib...>,
PtrC&& c, int_sequence<Ic...>)
{
return fn(d[Id]..., a[Ia]..., b[Ib]..., c[Ic]...);
}
template <class Fn,
class PtrA, int... Ia,
class PtrB, int... Ib,
class PtrC, int... Ic,
class ParamType>
CUTE_HOST_DEVICE constexpr
void
explode_with_d_scaling(Fn fn,
PtrA&& a, int_sequence<Ia...>,
PtrB&& b, int_sequence<Ib...>,
PtrC&& c, int_sequence<Ic...>,
ParamType&& p0)
{
return fn(a[Ia]..., b[Ib]..., c[Ic]..., p0);
}
template <class Fn,
class PtrD, int... Id,
class PtrA, int... Ia,
class PtrB, int... Ib,
class PtrC, int... Ic,
class ParamType>
CUTE_HOST_DEVICE constexpr
void
explode_with_d_scaling(Fn fn,
PtrD&& d, int_sequence<Id...>,
PtrA&& a, int_sequence<Ia...>,
PtrB&& b, int_sequence<Ib...>,
PtrC&& c, int_sequence<Ic...>,
ParamType&& p0)
{
return fn(d[Id]..., a[Ia]..., b[Ib]..., c[Ic]..., p0);
}
} // end namespace detail
template <int SRegCount, int DRegCount,
class Fn, class PtrS, class PtrD>
CUTE_HOST_DEVICE constexpr
void
explode(Fn fn, PtrS&& s, PtrD&& d)
{
return detail::explode(fn,
s, make_int_sequence<SRegCount>{},
d, make_int_sequence<DRegCount>{});
}
template <int ARegCount, int BRegCount, int CRegCount,
class Fn, class PtrA, class PtrB, class PtrC>
CUTE_HOST_DEVICE constexpr
void
explode(Fn fn, PtrA&& a, PtrB&& b, PtrC&& c)
{
return detail::explode(fn,
a, make_int_sequence<ARegCount>{},
b, make_int_sequence<BRegCount>{},
c, make_int_sequence<CRegCount>{});
}
template <int DRegCount, int ARegCount, int BRegCount, int CRegCount,
class Fn, class PtrD, class PtrA, class PtrB, class PtrC>
CUTE_HOST_DEVICE constexpr
void
explode(Fn fn, PtrD&& d, PtrA&& a, PtrB&& b, PtrC&& c)
{
return detail::explode(fn,
d, make_int_sequence<DRegCount>{},
a, make_int_sequence<ARegCount>{},
b, make_int_sequence<BRegCount>{},
c, make_int_sequence<CRegCount>{});
}
} // end namespace cute

View File

@ -0,0 +1,707 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#pragma once
#include <cute/config.hpp>
#include <cute/arch/copy.hpp>
#include <cute/tensor.hpp>
#include <cute/atom/copy_traits.hpp>
#include <cute/util/type_traits.hpp>
namespace cute
{
template <class... Args>
struct Copy_Atom;
template <class CopyOperation, class T>
struct Copy_Atom<CopyOperation, T> : Copy_Atom<Copy_Traits<CopyOperation>, T>
{};
template <class... Args, class T>
struct Copy_Atom<Copy_Traits<Args...>, T>
: Copy_Traits<Args...>
{
using Traits = Copy_Traits<Args...>;
// Bit and Thr layouts from the Copy_Traits
using ThrID = typename Traits::ThrID;
using BitLayoutSrc = typename Traits::SrcLayout;
using BitLayoutDst = typename Traits::DstLayout;
using BitLayoutRef = typename Traits::RefLayout;
using ValType = T;
using ValLayoutSrc = decltype(upcast<sizeof_bits<ValType>::value>(BitLayoutSrc{}));
using ValLayoutDst = decltype(upcast<sizeof_bits<ValType>::value>(BitLayoutDst{}));
using ValLayoutRef = decltype(upcast<sizeof_bits<ValType>::value>(BitLayoutRef{}));
CUTE_STATIC_ASSERT_V(size<0>(ValLayoutSrc{}) == size(ThrID{}), "CopyOperation is not valid for Src of ValType.");
CUTE_STATIC_ASSERT_V(size<0>(ValLayoutDst{}) == size(ThrID{}), "CopyOperation is not valid for Dst of ValType.");
CUTE_STATIC_ASSERT_V(size<0>(ValLayoutRef{}) == size(ThrID{}), "CopyOperation is not valid for Ref of ValType.");
static constexpr int NumValSrc = size<1>(ValLayoutSrc{});
static constexpr int NumValDst = size<1>(ValLayoutDst{});
// Additional Trait parameters/transformations
template <class... TraitsArgs>
CUTE_HOST_DEVICE
auto
with(TraitsArgs&&... args) const {
auto traits = Traits::with(std::forward<TraitsArgs>(args)...);
return Copy_Atom<decltype(traits), T>{traits};
}
//
// Tensor call interfaces
//
// Check and call instruction, or recurse
template <class TS, class SLayout,
class TD, class DLayout>
CUTE_HOST_DEVICE
void
call(Tensor<TS,SLayout> const& src,
Tensor<TD,DLayout> & dst) const
{
static_assert(SLayout::rank == 1, "Expected rank-1 src tensor");
static_assert(DLayout::rank == 1, "Expected rank-1 dst tensor");
if constexpr (is_constant<NumValSrc, decltype(size(src))>::value ||
is_constant<NumValDst, decltype(size(dst))>::value) {
// Dispatch to unpack for instruction
return copy_unpack(*this, src, dst);
} else
if constexpr (is_tuple<decltype(shape(src))>::value &&
is_tuple<decltype(shape(dst))>::value) {
// If the size of the src/dst doesn't match the instruction,
// recurse this rank-1 layout by peeling off the mode
// ((A,B,C,...)) -> (A,B,C,...)
return copy(*this, tensor<0>(src), tensor<0>(dst));
} else {
static_assert(sizeof(TS) < 0, "No instruction match and no recursion possible.");
}
}
// Accept mutable temporaries
template <class SEngine, class SLayout,
class DEngine, class DLayout>
CUTE_HOST_DEVICE
void
call(Tensor<SEngine,SLayout> const& src,
Tensor<DEngine,DLayout> && dst) const
{
return call(src, dst);
}
};
//
// A tiling of copy atoms
//
template <class TiledCopy, class ThrIdx>
struct ThrCopy;
template <class Copy_Atom,
class LayoutCopy_TV, // (tid,vid) -> coord [Need not be 2D...]
class ShapeTile_MN> // coord space
struct TiledCopy : Copy_Atom
{
// Layout information from the CopyAtom
using AtomThrID = typename Copy_Atom::ThrID; // thrid -> thr_idx
using AtomLayoutSrc = typename Copy_Atom::ValLayoutSrc; // (thr,val) -> offset
using AtomLayoutDst = typename Copy_Atom::ValLayoutDst; // (thr,val) -> offset
using AtomLayoutRef = typename Copy_Atom::ValLayoutRef; // (thr,val) -> offset
using AtomNumThr = decltype(size<0>(AtomLayoutRef{}));
using AtomNumVal = decltype(size<1>(AtomLayoutRef{}));
// Layout information for the TiledCopy
using Tiler_MN = ShapeTile_MN;
using TiledShape_MN = decltype(shape(ShapeTile_MN{}));
using TiledLayout_TV = LayoutCopy_TV;
using TiledNumThr = decltype(size<0>(TiledLayout_TV{}));
using TiledNumVal = decltype(size<1>(TiledLayout_TV{}));
CUTE_STATIC_ASSERT_V(TiledNumThr{} % AtomNumThr{} == Int<0>{}, "TiledCopy uses too few thrs for selected CopyAtom");
CUTE_STATIC_ASSERT_V(TiledNumVal{} % AtomNumVal{} == Int<0>{}, "TiledCopy uses too few vals for selected CopyAtom");
// Tile a tensor or a layout from shape
// (M,N,...)
// to shape
// ((ThrV,ThrX),FrgV,(RestM,RestN,...))
// where
// ThrV: The threads local to a COPY_ATOM Src.
// ThrX: The threads tiled across COPY_ATOMs Src.
// FrgV: The values local to a COPY_ATOM Src.
// RestM: The values tiled in M.
// RestN: The values tiled in N.
template <class STensor>
CUTE_HOST_DEVICE constexpr static
auto
tidfrg_S(STensor&& stensor)
{
constexpr int R = remove_cvref_t<STensor>::rank;
static_assert(R >= rank_v<TiledShape_MN>, "Rank of tensor to be partitioned too small.");
// Generalize the dimension checks for arbitrary rank
//CUTE_STATIC_ASSERT_V(size<0>(stensor) % size<0>(TiledShape_MNK{}) == Int<0>{});
//CUTE_STATIC_ASSERT_V(size<1>(stensor) % size<1>(TiledShape_MNK{}) == Int<0>{});
return tile2thrfrg(zipped_divide(stensor,Tiler_MN{}), right_inverse(AtomLayoutRef{}).compose(AtomLayoutSrc{}));
}
// Tile a tensor or a layout from shape
// (M,N,...)
// to shape
// ((ThrV,ThrX),FrgV,(RestM,RestN,...))
// where
// ThrV: The threads local to a COPY_ATOM Dst.
// ThrX: The threads tiled across COPY_ATOMs Dst.
// FrgV: The values local to a COPY_ATOM Dst.
// RestM: The values tiled in M.
// RestN: The values tiled in N.
template <class DTensor>
CUTE_HOST_DEVICE constexpr static
auto
tidfrg_D(DTensor&& dtensor)
{
constexpr int R = remove_cvref_t<DTensor>::rank;
static_assert(R >= rank_v<TiledShape_MN>, "Rank of tensor to be partitioned too small.");
// Generalize the dimension checks for arbitrary rank
//CUTE_STATIC_ASSERT_V(size<0>(stensor) % size<0>(TiledShape_MNK{}) == Int<0>{});
//CUTE_STATIC_ASSERT_V(size<1>(stensor) % size<1>(TiledShape_MNK{}) == Int<0>{});
return tile2thrfrg(zipped_divide(dtensor,Tiler_MN{}), right_inverse(AtomLayoutRef{}).compose(AtomLayoutDst{}));
}
// Tile a tensor or a layout from shape
// (Tile,(RestM,RestN,...))
// to shape
// ((ThrV,ThrX),FrgV,(RestM,RestN,...))
template <class Tensor, class Ref2TrgLayout>
CUTE_HOST_DEVICE constexpr static
auto
tile2thrfrg(Tensor&& tensor, Ref2TrgLayout const& ref2trg)
{
// Take the thrs/vals that the atom is interested in
// NOTE: Assumes the AtomNumThr are contiguous and identity within TiledThrID
auto atom_layout_TV = zipped_divide(TiledLayout_TV{}, make_shape(AtomNumThr{}, AtomNumVal{}));
// ((atom_tid,atom_val),(rest_tid,rest_val)) -> (m,n)
// Transform to the trg layout
auto trg_layout_TV = atom_layout_TV.compose(ref2trg, _);
// ((trg_tid,trg_val),(rest_tid,rest_val)) -> (m,n)
// Transform the thrs mode from thrid to thr_idx
// NOTE: Assumes the AtomNumThr are contiguous and identity within TiledThrID
auto thrval2mn = coalesce(zip(trg_layout_TV), Shape<_1,Shape<_1,_1>>{});
// ((trg_tid,rest_tid),(trg_val,rest_val)) -> (m,n)
/// ==================
// Transform the tile mode
auto tv_tensor = tensor.compose(thrval2mn, _);
// ((thrid,val),(RM,RN,...))
// Unfold and return
return tv_tensor(make_coord(_,_), _);
}
// retile_S and retile_D assume they are working with the reference layout -- they are the same
template <class Tensor>
CUTE_HOST_DEVICE constexpr static
auto
retile(Tensor&& tensor)
{
constexpr int R = remove_cvref_t<Tensor>::rank;
// Assert that AtomLayoutSrc|Dst is identity so we can skip the Ref transformation
// Assume the first size<0>(tensor) elements are the first val_ids in TiledLayout_TV.
// Then, we only need the shape+layout of those size<0>(tensor) elements in TiledLayout_TV
// and that shape is what we gather from the other modes of tensor
auto V = size<0>(tensor);
auto frg_layout_mn = upcast<TiledNumThr{} * V>(right_inverse(TiledLayout_TV{}).with_shape(TiledShape_MN{}));
// (m,n) -> v_idx -- The shape and order of the V inside of TiledLayout_TV
auto frg_layout_v = zipped_divide(logical_product(make_layout(V), right_inverse(frg_layout_mn)), make_layout(AtomNumVal{}));
// (atom_vals,rest_vals) -> (v,m,n)
/// =======
// Tile the tensor for TileFrg
auto t_tensor = zipped_divide(tensor, prepend(product_each(shape(frg_layout_mn)), V));
// ((TileV,TileM,TileN,...),(1,RestM,RestN,...))
// Transform the tile mode
auto v_tensor = t_tensor.compose(frg_layout_v, _);
// ((atom_vals,rest_vals),(1,RM,RN,...))
// Unfold and return
return v_tensor(_, append<R>(Int<0>{},_));
}
CUTE_HOST_DEVICE constexpr static
auto
get_layoutS_TV()
{
// (M,N) -> (M,N)
auto ref_S = make_layout(make_shape(TiledShape_MN{}, Int<1>{}));
// (thr_idx,val_idx) -> (M,N)
return tile2thrfrg(ref_S, right_inverse(AtomLayoutRef{}).compose(AtomLayoutSrc{}))(_,_,Int<0>{});
}
CUTE_HOST_DEVICE constexpr static
auto
get_layoutS_MN()
{
// (thr_idx,val_idx) -> (M,N)
auto layoutS_TV = get_layoutS_TV();
// (M,K) -> (thr_idx,val_idx)
auto layoutS_MK = right_inverse(layoutS_TV).with_shape(TiledShape_MN{});
// athrid = (v,m,k) -> thr_idx
auto thrID_S = make_layout(size<0>(TiledLayout_TV{}));
return cute::make_tuple(layoutS_MK, thrID_S);
}
CUTE_HOST_DEVICE constexpr static
auto
get_layoutD_TV()
{
// (M,N) -> (M,N)
auto ref_D = make_layout(make_shape(TiledShape_MN{}, Int<1>{}));
// (thr_idx,val_idx) -> (M,N)
return tile2thrfrg(ref_D, right_inverse(AtomLayoutRef{}).compose(AtomLayoutDst{}))(_,_,Int<0>{});
}
CUTE_HOST_DEVICE constexpr static
auto
get_layoutD_MN()
{
// (thr_idx,val_idx) -> (M,N)
auto layoutD_TV = get_layoutD_TV();
// (M,K) -> (thr_idx,val_idx)
auto layoutD_MK = right_inverse(layoutD_TV).with_shape(TiledShape_MN{});
// athrid = (v,m,k) -> thr_idx
auto thrID_D = make_layout(size<0>(TiledLayout_TV{}));
return cute::make_tuple(layoutD_MK, thrID_D);
}
template <class ThrIdx,
__CUTE_REQUIRES(is_integral<ThrIdx>::value)>
CUTE_HOST_DEVICE static
auto
get_slice(ThrIdx const& thr_idx)
{
return ThrCopy<TiledCopy, ThrIdx>(thr_idx);
}
template <class ThrIdx,
__CUTE_REQUIRES(is_integral<ThrIdx>::value)>
CUTE_HOST_DEVICE static
auto
get_thread_slice(ThrIdx const& thr_idx)
{
return get_slice(thr_idx);
}
};
template <class TiledCopy, class ThrIdx>
struct ThrCopy
{
ThrIdx thr_idx_;
CUTE_HOST_DEVICE
ThrCopy(ThrIdx const& thr_idx) : thr_idx_(thr_idx) {}
template <class STensor>
CUTE_HOST_DEVICE
auto
partition_S(STensor&& stensor) {
//static_assert(sizeof(typename remove_cvref_t<STensor>::value_type) == sizeof(typename TiledCopy::ValType),
// "Expected ValType for tiling SrcTensor.");
auto thr_tensor = make_tensor(std::forward<STensor>(stensor).data(), TiledCopy::tidfrg_S(stensor.layout()));
return thr_tensor(thr_idx_, _, repeat<rank_v<STensor>>(_));
}
template <class DTensor>
CUTE_HOST_DEVICE
auto
partition_D(DTensor&& dtensor) {
//static_assert(sizeof(typename remove_cvref_t<DTensor>::value_type) == sizeof(typename TiledCopy::ValType),
// "Expected ValType for tiling DstTensor.");
auto thr_tensor = make_tensor(std::forward<DTensor>(dtensor).data(), TiledCopy::tidfrg_D(dtensor.layout()));
return thr_tensor(thr_idx_, _, repeat<rank_v<DTensor>>(_));
}
template <class STensor>
CUTE_HOST_DEVICE static
auto
retile_S(STensor&& stensor) {
// static_assert(sizeof(typename remove_cvref_t<STensor>::value_type) == sizeof(typename TiledCopy::ValType),
// "Expected ValType for tiling SrcTensor.");
return make_tensor(std::forward<STensor>(stensor).data(), TiledCopy::retile(stensor.layout()));
}
template <class DTensor>
CUTE_HOST_DEVICE static
auto
retile_D(DTensor&& dtensor) {
// static_assert(sizeof(typename remove_cvref_t<DTensor>::value_type) == sizeof(typename TiledCopy::ValType),
// "Expected ValType for tiling DstTensor.");
return make_tensor(std::forward<DTensor>(dtensor).data(), TiledCopy::retile(dtensor.layout()));
}
};
template <class... Args,
class LayoutCopy_TV,
class Tiler>
CUTE_HOST_DEVICE
auto
make_tiled_copy_impl(Copy_Atom<Args...> const& atom,
LayoutCopy_TV const&,
Tiler const&)
{
return TiledCopy<Copy_Atom<Args...>, LayoutCopy_TV, Tiler>{atom};
}
//
// These tile the Copy_Atom as a whole
//
template <class... Args,
class TiledMMA>
CUTE_HOST_DEVICE
auto
make_tiled_copy_A(Copy_Atom<Args...> const& copy_atom,
TiledMMA const& tiled_mma)
{
using MNK = typename TiledMMA::TiledShape_MNK;
return make_tiled_copy_impl(copy_atom, tiled_mma.get_layoutA_TV(), make_shape(size<0>(MNK{}),size<2>(MNK{})));
}
template <class... Args,
class TiledMMA>
CUTE_HOST_DEVICE
auto
make_tiled_copy_B(Copy_Atom<Args...> const& copy_atom,
TiledMMA const& tiled_mma)
{
using MNK = typename TiledMMA::TiledShape_MNK;
return make_tiled_copy_impl(copy_atom, tiled_mma.get_layoutB_TV(), make_shape(size<1>(MNK{}),size<2>(MNK{})));
}
template <class... Args,
class TiledMMA>
CUTE_HOST_DEVICE
auto
make_tiled_copy_C(Copy_Atom<Args...> const& copy_atom,
TiledMMA const& tiled_mma)
{
using MNK = typename TiledMMA::TiledShape_MNK;
return make_tiled_copy_impl(copy_atom, tiled_mma.get_layoutC_TV(), make_shape(size<0>(MNK{}),size<1>(MNK{})));
}
// returns the smallest tiled copy that can retile LayoutC_TV
// for use with pipelined epilogues with subtiled stores
template <class... Args,
class TiledMMA>
CUTE_HOST_DEVICE
auto
make_tiled_copy_C_atom(Copy_Atom<Args...> const& copy_atom,
TiledMMA const& tiled_mma)
{
// Truncate the V-layout to just the Copy_Atom, keep the V-order
auto layoutC_TV = tiled_mma.get_layoutC_TV();
auto copy_V = Int<Copy_Atom<Args...>::NumValSrc>{};
CUTE_STATIC_ASSERT_V(copy_V <= size<1>(layoutC_TV));
auto layout_TV = composition(layoutC_TV, make_layout(make_shape(size<0>(layoutC_TV), copy_V)));
// Recompute tiler and restride the TV layout for the new tiler
// Tiler -- Find the active elements in the MMA tensor and generate a tiler to extract them
// Convert to the awkward by-mode tiler to preserve the modes of the tiled MMA
using MNK = typename TiledMMA::TiledShape_MNK;
auto mma_tiler = make_shape(size<0>(MNK{}),size<1>(MNK{}));
auto mma_zeros = repeat_like(mma_tiler, Int<0>{});
auto tiler = transform(make_seq<rank(mma_tiler)>{}, [&](auto i) {
return filter(composition(make_layout(mma_tiler, replace<i>(mma_zeros, Int<1>{})), layout_TV));
});
// Layout_TV -- Find the (tid,vid) -> tile coord transformation
// Apply the tiler to a reference and transform the codomain
// tile_coord -> mma_coord
auto tile2mma = composition(make_layout(mma_tiler), tiler);
// (tid,vid) -> tile_coord
auto layout_tv = composition(left_inverse(tile2mma), layout_TV);
using MNK = typename TiledMMA::TiledShape_MNK;
return make_tiled_copy_impl(copy_atom, layout_tv, tiler);
}
template <class... Args,
class ThrLayout,
class ValLayout = Layout<_1>>
CUTE_HOST_DEVICE
auto
make_tiled_copy(Copy_Atom<Args...> const& copy_atom,
ThrLayout const& thr_layout = {}, // (m,n) -> thr_idx
ValLayout const& val_layout = {})
{
constexpr int R = cute::max(rank_v<ThrLayout>, rank_v<ValLayout>);
auto thr_layout_mn = append<R>(thr_layout, Layout<_1>{});
auto val_layout_mn = append<R>(val_layout, Layout<_1>{});
// Take the raked_products to compute the Layout_MN
auto layout_mn = raked_product(thr_layout_mn, val_layout_mn);
auto layout_tv = right_inverse(layout_mn).with_shape(make_shape(size(thr_layout), size(val_layout)));
// print("thr_layout: "); print(thr_layout_mn); print("\n");
// print("val_layout: "); print(val_layout_mn); print("\n");
// print("layout_mn : "); print(layout_mn); print("\n");
// print("layout_tv : "); print(layout_tv); print("\n");
return make_tiled_copy_impl(copy_atom, layout_tv, product_each(shape(layout_mn)));
}
// Make a TiledCopy out of the copy_atom that matches the Src-Layout of tiled_copy
template <class... Args,
class TiledCopy>
CUTE_HOST_DEVICE
auto
make_tiled_copy_S(Copy_Atom<Args...> const& copy_atom,
TiledCopy const& tiled_copy)
{
return make_tiled_copy_impl(copy_atom, tiled_copy.get_layoutS_TV(), typename TiledCopy::Tiler_MN{});
}
// Make a TiledCopy out of the copy_atom that matches the Dst-Layout of tiled_copy
template <class... Args,
class TiledCopy>
CUTE_HOST_DEVICE
auto
make_tiled_copy_D(Copy_Atom<Args...> const& copy_atom,
TiledCopy const& tiled_copy)
{
return make_tiled_copy_impl(copy_atom, tiled_copy.get_layoutD_TV(), typename TiledCopy::Tiler_MN{});
}
//
// Size
//
// The logical size of a TileCopy
template <int... I, class... Args>
CUTE_HOST_DEVICE constexpr
auto
tile_size(TiledCopy<Args...> const&)
{
return size<I...>(typename TiledCopy<Args...>::TiledShape_MN{});
}
// The number of threads involved in a TiledCopy
template <class... Args>
CUTE_HOST_DEVICE constexpr
auto
size(TiledCopy<Args...> const&)
{
return typename TiledCopy<Args...>::TiledNumThr{};
}
//
// Display utilities
//
template <class... Args, class T>
CUTE_HOST_DEVICE
void
print(Copy_Atom<Copy_Traits<Args...>, T> const&)
{
using Atom = Copy_Atom<Copy_Traits<Args...>, T>;
print("Copy_Atom\n");
print(" ThrID: "); print(typename Atom::ThrID{}); print("\n");
print(" ValLayoutSrc: "); print(typename Atom::ValLayoutSrc{}); print("\n");
print(" ValLayoutDst: "); print(typename Atom::ValLayoutDst{}); print("\n");
print(" ValLayoutRef: "); print(typename Atom::ValLayoutRef{}); print("\n");
print(" ValueType: %db\n", int(sizeof_bits<typename Atom::ValType>::value));
}
template <class Atom, class... Args>
CUTE_HOST_DEVICE
void
print(TiledCopy<Atom, Args...> const& copy, char const* pad = "")
{
using Copy = TiledCopy<Atom, Args...>;
print("TiledCopy\n");
print(" Tiler_MN: "); print(typename Copy::Tiler_MN{}); print("\n");
print(" TiledLayout_TV: "); print(typename Copy::TiledLayout_TV{}); print("\n");
print(static_cast<Atom const&>(copy));
}
template <class TiledCopy, class ThrIdx>
CUTE_HOST_DEVICE
void
print(ThrCopy<TiledCopy, ThrIdx> const&)
{
print(TiledCopy{});
}
template <class... Args>
CUTE_HOST_DEVICE
auto
print_latex(TiledCopy<Args...> const& copy)
{
auto [layoutS_MN, thrID_S] = copy.get_layoutS_MN();
auto [layoutD_MN, thrID_D] = copy.get_layoutD_MN();
print_latex_copy(layoutS_MN, thrID_S,
layoutD_MN, thrID_D);
}
// MNK Copy Layout to Latex TIKZ -- 8-value color coded by thread
template <class LayoutS, class ThrIDS,
class LayoutD, class ThrIDD>
CUTE_HOST_DEVICE
void
print_latex_copy(LayoutS const& S, ThrIDS const& TS, // (m,n) -> (tid,vid) and tid -> thr_idx
LayoutD const& D, ThrIDD const& TD) // (m,n) -> (tid,vid) and tid -> thr_idx
{
CUTE_STATIC_ASSERT_V(rank(S) == Int<2>{});
CUTE_STATIC_ASSERT_V(rank(D) == Int<2>{});
assert(size<0>(S) == size<0>(D));
assert(size<1>(S) == size<1>(D));
char const* latex_header =
"\\documentclass{standalone}\n"
"\\usepackage{tikz}\n"
"\\usetikzlibrary{external}\n"
"\\tikzexternalize\n"
"\\begin{document}\n"
"\\begin{tikzpicture}[x={(0cm,-1cm)},y={(1cm,0cm)},box/.style={rectangle,draw=black,thick,minimum size=1cm,anchor=center}]\n\n";
char const* latex_footer =
"\\end{tikzpicture}\n"
"\\end{document}\n";
char const* color_map[8] = {"{rgb,255:red,175;green,175;blue,255}",
"{rgb,255:red,175;green,255;blue,175}",
"{rgb,255:red,255;green,255;blue,175}",
"{rgb,255:red,255;green,175;blue,175}",
"{rgb,255:red,210;green,210;blue,255}",
"{rgb,255:red,210;green,255;blue,210}",
"{rgb,255:red,255;green,255;blue,210}",
"{rgb,255:red,255;green,210;blue,210}",};
// Header
printf("%% LayoutS: "); print(S); printf("\n");
printf("%% ThrIDS : "); print(TS); printf("\n");
printf("%% LayoutD: "); print(D); printf("\n");
printf("%% ThrIDD : "); print(TD); printf("\n\n");
printf(latex_header);
// S starting at 0,0
for (int i = 0; i < size<0>(S); ++i) {
for (int j = 0; j < size<1>(S); ++j) {
int thrid = S(i,j) % size(TS);
int val_idx = S(i,j) / size(TS);
int thr_idx = TS(thrid);
printf("\\node[box,fill=%s] at (%d,%d) {\\shortstack{T%d \\\\ V%d}};\n",
color_map[thr_idx % 8],
i, j,
thr_idx, val_idx);
}
}
// D starting at 0,size<1>(S)+3
for (int i = 0; i < size<0>(D); ++i) {
for (int j = 0; j < size<1>(D); ++j) {
int thrid = D(i,j) % size(TD);
int val_idx = D(i,j) / size(TD);
int thr_idx = TD(thrid);
printf("\\node[box,fill=%s] at (%d,%d) {\\shortstack{T%d \\\\ V%d}};\n",
color_map[thr_idx % 8],
i, j + size<1>(S) + 3,
thr_idx, val_idx);
}
}
// S Labels
for (int i = 0, j = -1; i < size<0>(S); ++i) {
printf("\\node at (%d,%d) {\\Large{\\texttt{%d}}};\n", i, j, i);
}
for (int j = 0, i = -1; j < size<1>(S); ++j) {
printf("\\node at (%d,%d) {\\Large{\\texttt{%d}}};\n", i, j, j);
}
// D Labels
for (int i = 0, j = size<1>(D); i < size<0>(S); ++i) {
printf("\\node at (%d,%d) {\\Large{\\texttt{%d}}};\n", i, j + size<1>(S) + 3, i);
}
for (int j = 0, i = -1; j < size<1>(D); ++j) {
printf("\\node at (%d,%d) {\\Large{\\texttt{%d}}};\n", i, j + size<1>(S) + 3, j);
}
// Footer
printf(latex_footer);
}
} // end namespace cute
////////////////////////////////////////////////////////////////////////////////////////////////////
#include <cute/atom/copy_traits_sm75.hpp>
#include <cute/atom/copy_traits_sm80.hpp>
// #include <cute/atom/copy_traits_sm90.hpp>
// Config
// #if (__CUDACC_VER_MAJOR__ >= 12)
#if 0
# define CUTE_COPY_ATOM_TMA_SM90_ENABLED
#endif
#if defined(CUTE_COPY_ATOM_TMA_SM90_ENABLED)
#include <cute/atom/copy_traits_sm90_tma.hpp>
#endif
////////////////////////////////////////////////////////////////////////////////////////////////////

View File

@ -0,0 +1,131 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#pragma once
#include <cute/arch/copy.hpp>
#include <cute/tensor.hpp>
namespace cute
{
/**
* concept Copy_Traits
* {
* using ThrID = // Logical thread id (tid) -> tidx
*
* using SrcLayout = // (Logical src thread id (tid), Logical src value id (vid)) -> bit
* using DstLayout = // (Logical dst thread id (tid), Logical dst value id (vid)) -> bit
* using RefLayout = // (Logical ref thread id (tid), Logical ref value id (vid)) -> bit
* };
*
* The abstract bit ordering of the Copy_Traits (the codomain of SrcLayout, DstLayout, and RefLayout)
* is arbitrary and only used to construct maps
* (ref-tid,ref-vid) -> (src-tid,src-vid)
* (ref-tid,ref-vid) -> (dst-tid,dst-vid)
* in TiledCopy. The Layout_TV in TiledCopy is in accordance with the RefLayout of a Traits, then mapped to
* the Src or Dst (tid,vid) representation on demand.
*
*/
template <class CopyOperation, class... CopyOpArgs>
struct Copy_Traits
{
static_assert(sizeof(CopyOperation) == 0, "Copy_Traits not implemented for this Copy_Operation.");
};
template <class S, class D>
struct Copy_Traits<UniversalCopy<S,D>>
{
// Logical thread id to thread idx (one-thread)
using ThrID = Layout<_1>;
// Map from (src-thr,src-val) to bit
using SrcLayout = Layout<Shape<_1,Int<sizeof_bits<S>::value>>>;
// Map from (dst-thr,dst-val) to bit
using DstLayout = Layout<Shape<_1,Int<sizeof_bits<D>::value>>>;
// Reference map from (thr,val) to bit
using RefLayout = SrcLayout;
};
template <>
struct Copy_Traits<DefaultCopy>
{
// Logical thread id to thread idx (one-thread)
using ThrID = Layout<_1>;
// Map from (src-thr,src-val) to bit
using SrcLayout = Layout<Shape<_1,_1>, Stride<_0,_0>>;
// Map from (dst-thr,dst-val) to bit
using DstLayout = Layout<Shape<_1,_1>, Stride<_0,_0>>;
// Reference map from (thr,val) to bit
using RefLayout = SrcLayout;
};
//
// Generic copy_unpack for any Copy_Traits
//
template <class Operation, class... Args,
class TS, class SLayout,
class TD, class DLayout>
CUTE_HOST_DEVICE constexpr
void
copy_unpack(Copy_Traits<Operation, Args...> const&,
Tensor<TS,SLayout> const& src,
Tensor<TD,DLayout> & dst)
{
// Specializations can generalize on these checks
//static_assert(is_smem<TS>::value, "Expected smem for this Copy_Traits<Operation>");
//static_assert(is_rmem<TD>::value, "Expected rmem for this Copy_Traits<Operation>");
using RegistersSrc = typename Operation::SRegisters;
using RegistersDst = typename Operation::DRegisters;
using RegTypeSrc = typename remove_extent<RegistersSrc>::type;
using RegTypeDst = typename remove_extent<RegistersDst>::type;
constexpr int RegNumSrc = extent<RegistersSrc>::value;
constexpr int RegNumDst = extent<RegistersDst>::value;
Tensor rS = recast<RegTypeSrc>(src);
Tensor rD = recast<RegTypeDst>(dst);
CUTE_STATIC_ASSERT_V(size(rS) == Int<RegNumSrc>{},
"In CopyAtom, src layout doesn't vectorize into registers. This src layout is incompatible with this tiled copy.");
CUTE_STATIC_ASSERT_V(size(rD) == Int<RegNumDst>{},
"In CopyAtom, dst layout doesn't vectorize into registers. This dst layout is incompatible with this tiled copy.");
detail::explode(Operation::copy,
rS, make_int_sequence<RegNumSrc>{},
rD, make_int_sequence<RegNumDst>{});
}
} // end namespace cute

View File

@ -0,0 +1,160 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#pragma once
#include <cute/arch/copy_sm75.hpp>
#include <cute/atom/copy_traits.hpp>
#include <cute/layout.hpp>
namespace cute
{
template <>
struct Copy_Traits<SM75_U32x1_LDSM_N>
{
// Logical thread id to thread idx (warp)
using ThrID = Layout<_32>;
// Map from (src-thr,src-val) to bit
using SrcLayout = Layout<Shape <Shape < _8,_4>,_128>,
Stride<Stride<_128,_0>, _1>>;
// Map from (dst-thr,dst-val) to bit
using DstLayout = Layout<Shape <_32,_32>,
Stride<_32, _1>>;
// Reference map from (thr,val) to bit
using RefLayout = DstLayout;
};
template <>
struct Copy_Traits<SM75_U32x2_LDSM_N>
{
// Logical thread id to thread idx (warp)
using ThrID = Layout<_32>;
// Map from (src-thr,src-val) to bit
using SrcLayout = Layout<Shape <Shape < _16,_2>,_128>,
Stride<Stride<_128,_0>, _1>>;
// Map from (dst-thr,dst-val) to bit
using DstLayout = Layout<Shape <_32,Shape <_32, _2>>,
Stride<_32,Stride< _1,_1024>>>;
// Reference map from (thr,val) to bit
using RefLayout = DstLayout;
};
template <>
struct Copy_Traits<SM75_U32x4_LDSM_N>
{
// Logical thread id to thread idx (warp)
using ThrID = Layout<_32>;
// Map from (src-thr,src-val) to bit
using SrcLayout = Layout<Shape < _32,_128>,
Stride<_128, _1>>;
// Map from (dst-thr,dst-val) to bit
using DstLayout = Layout<Shape <_32,Shape <_32, _4>>,
Stride<_32,Stride< _1,_1024>>>;
// Reference map from (thr,val) to bit
using RefLayout = DstLayout;
};
template <>
struct Copy_Traits<SM75_U32x4_LDSM_N_B>
{
// Logical thread id to thread idx (warp)
using ThrID = Layout<_32>;
// Map from (src-thr,src-val) to bit
using SrcLayout = Layout<Shape < _32,_128>,
Stride<_128, _1>>;
// Map from (dst-thr,dst-val) to bit
using DstLayout = Layout<Shape <_32,Shape <_32, _4>>,
Stride<_32,Stride< _1,_1024>>>;
// Reference map from (thr,val) to bit
using RefLayout = DstLayout;
};
template <>
struct Copy_Traits<SM75_U16x2_LDSM_T>
{
// Logical thread id to thread idx (warp)
using ThrID = Layout<_32>;
// Map from (src-thr,src-val) to bit
using SrcLayout = Layout<Shape <Shape < _8,_4>,_128>,
Stride<Stride<_128,_0>, _1>>;
// Map from (dst-thr,dst-val) to bit
using DstLayout = Layout<Shape <Shape < _4, _8>,Shape <_16, _2>>,
Stride<Stride<_256,_16>,Stride< _1,_128>>>;
// Reference map from (thr,val) to bit
using RefLayout = DstLayout;
};
template <>
struct Copy_Traits<SM75_U16x4_LDSM_T>
{
// Logical thread id to thread idx (warp)
using ThrID = Layout<_32>;
// Map from (src-thr,src-val) to bit
using SrcLayout = Layout<Shape <Shape < _16,_2>,_128>,
Stride<Stride<_128,_0>, _1>>;
// Map from (dst-thr,dst-val) to bit
using DstLayout = Layout<Shape <Shape < _4, _8>,Shape <_16, _2, _2>>,
Stride<Stride<_256,_16>,Stride< _1,_128,_1024>>>;
// Reference map from (thr,val) to bit
using RefLayout = DstLayout;
};
template <>
struct Copy_Traits<SM75_U16x8_LDSM_T>
{
// Logical thread id to thread idx (warp)
using ThrID = Layout<_32>;
// Map from (src-thr,src-val) to bit
using SrcLayout = Layout<Shape < _32,_128>,
Stride<_128, _1>>;
// Map from (dst-thr,dst-val) to bit
using DstLayout = Layout<Shape <Shape < _4, _8>,Shape <_16, _2, _4>>,
Stride<Stride<_256,_16>,Stride< _1,_128,_1024>>>;
// Reference map from (thr,val) to bit
using RefLayout = DstLayout;
};
} // end namespace cute

View File

@ -0,0 +1,130 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#pragma once
#include <cute/arch/copy_sm80.hpp>
#include <cute/atom/copy_traits.hpp>
#include <cute/layout.hpp>
namespace cute
{
template <class S, class D>
struct Copy_Traits<SM80_CP_ASYNC_CACHEALWAYS<S,D>>
{
// Logical thread id to thread idx (one-thread)
using ThrID = Layout<_1>;
// Map from (src-thr,src-val) to bit
using SrcLayout = Layout<Shape<_1,Int<sizeof_bits<S>::value>>>;
// Map from (dst-thr,dst-val) to bit
using DstLayout = Layout<Shape<_1,Int<sizeof_bits<D>::value>>>;
// Reference map from (thr,val) to bit
using RefLayout = SrcLayout;
};
template <class S, class D>
struct Copy_Traits<SM80_CP_ASYNC_CACHEGLOBAL<S,D>>
{
// Logical thread id to thread idx (one-thread)
using ThrID = Layout<_1>;
// Map from (src-thr,src-val) to bit
using SrcLayout = Layout<Shape<_1,Int<sizeof_bits<S>::value>>>;
// Map from (dst-thr,dst-val) to bit
using DstLayout = Layout<Shape<_1,Int<sizeof_bits<D>::value>>>;
// Reference map from (thr,val) to bit
using RefLayout = SrcLayout;
};
template <class S, class D>
struct Copy_Traits<MACA_CP_ASYNC_CACHEGLOBAL<S,D>>
{
// Logical thread id to thread idx (one-thread)
using ThrID = Layout<_1>;
// Map from (src-thr,src-val) to bit
using SrcLayout = Layout<Shape<_1,Int<sizeof_bits<S>::value>>>;
// Map from (dst-thr,dst-val) to bit
using DstLayout = Layout<Shape<_1,Int<sizeof_bits<D>::value>>>;
// Reference map from (thr,val) to bit
using RefLayout = SrcLayout;
};
////////////////////////////////////////////////////////////////////////////////////////////////////
template <>
struct Copy_Traits<MACA_LDS_TRANS_4X16>
{
// Logical thread id to thread idx (16 threads)
using ThrID = Layout<_16>;
// Map from (src-thr,src-val) to bit
using SrcLayout = Layout<Shape <Shape<_4, _4>, _64>,
Stride<Stride<_64, _256>, _1>>;
// Map from (dst-thr,dst-val) to bit
using DstLayout = Layout<Shape <_16, Shape<_16, _4>>,
Stride<_16, Stride<_1, _256>>>;
// Reference map from (thr,val) to bit
using RefLayout = DstLayout;
};
// Element copy selector
template <class SrcTensor, class DstTensor>
CUTE_HOST_DEVICE constexpr
auto
select_elementwise_copy(SrcTensor const&, DstTensor const&)
{
using SrcType = typename SrcTensor::value_type;
using DstType = typename DstTensor::value_type;
#if defined(CUTE_ARCH_CP_ASYNC_SM80_ENABLED)
if constexpr (is_gmem<SrcTensor>::value && is_smem<DstTensor>::value &&
sizeof(SrcType) == sizeof(DstType) &&
(sizeof(SrcType) == 4 || sizeof(SrcType) == 8 || sizeof(SrcType) == 16))
{
return SM80_CP_ASYNC_CACHEALWAYS<SrcType,DstType>{};
} else {
return UniversalCopy<SrcType,DstType>{};
}
CUTE_GCC_UNREACHABLE;
#else
return UniversalCopy<SrcType,DstType>{};
#endif
}
}

View File

@ -0,0 +1,132 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#pragma once
#include <cute/arch/copy_sm90.hpp>
#include <cute/atom/copy_traits.hpp>
#include <cute/atom/copy_traits_sm75.hpp>
#include <cute/layout.hpp>
namespace cute
{
template <>
struct Copy_Traits<SM90_U32x1_STSM_N>
{
// Logical thread id to thread idx (warp)
using ThrID = Layout<_32>;
// Map from (src-thr,src-val) to bit
using SrcLayout = typename Copy_Traits<SM75_U32x1_LDSM_N>::DstLayout;
// Map from (dst-thr,dst-val) to bit
using DstLayout = typename Copy_Traits<SM75_U32x1_LDSM_N>::SrcLayout;
// Reference map from (thr,val) to bit
using RefLayout = SrcLayout;
};
template <>
struct Copy_Traits<SM90_U32x2_STSM_N>
{
// Logical thread id to thread idx (warp)
using ThrID = Layout<_32>;
// Map from (src-thr,src-val) to bit
using SrcLayout = typename Copy_Traits<SM75_U32x2_LDSM_N>::DstLayout;
// Map from (dst-thr,dst-val) to bit
using DstLayout = typename Copy_Traits<SM75_U32x2_LDSM_N>::SrcLayout;
// Reference map from (thr,val) to bit
using RefLayout = SrcLayout;
};
template <>
struct Copy_Traits<SM90_U32x4_STSM_N>
{
// Logical thread id to thread idx (warp)
using ThrID = Layout<_32>;
// Map from (src-thr,src-val) to bit
using SrcLayout = typename Copy_Traits<SM75_U32x4_LDSM_N>::DstLayout;
// Map from (dst-thr,dst-val) to bit
using DstLayout = typename Copy_Traits<SM75_U32x4_LDSM_N>::SrcLayout;
// Reference map from (thr,val) to bit
using RefLayout = SrcLayout;
};
template <>
struct Copy_Traits<SM90_U16x2_STSM_T>
{
// Logical thread id to thread idx (warp)
using ThrID = Layout<_32>;
// Map from (src-thr,src-val) to bit
using SrcLayout = typename Copy_Traits<SM75_U16x2_LDSM_T>::DstLayout;
// Map from (dst-thr,dst-val) to bit
using DstLayout = typename Copy_Traits<SM75_U16x2_LDSM_T>::SrcLayout;
// Reference map from (thr,val) to bit
using RefLayout = SrcLayout;
};
template <>
struct Copy_Traits<SM90_U16x4_STSM_T>
{
// Logical thread id to thread idx (warp)
using ThrID = Layout<_32>;
// Map from (src-thr,src-val) to bit
using SrcLayout = typename Copy_Traits<SM75_U16x4_LDSM_T>::DstLayout;
// Map from (dst-thr,dst-val) to bit
using DstLayout = typename Copy_Traits<SM75_U16x4_LDSM_T>::SrcLayout;
// Reference map from (thr,val) to bit
using RefLayout = SrcLayout;
};
template <>
struct Copy_Traits<SM90_U16x8_STSM_T>
{
// Logical thread id to thread idx (warp)
using ThrID = Layout<_32>;
// Map from (src-thr,src-val) to bit
using SrcLayout = typename Copy_Traits<SM75_U16x8_LDSM_T>::DstLayout;
// Map from (dst-thr,dst-val) to bit
using DstLayout = typename Copy_Traits<SM75_U16x8_LDSM_T>::SrcLayout;
// Reference map from (thr,val) to bit
using RefLayout = SrcLayout;
};
} // end namespace cute

View File

@ -0,0 +1,973 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#pragma once
#if !defined(__MACACC_RTC__)
#include <cuda.h>
#endif
#include <cute/arch/copy_sm90_desc.hpp>
#include <cute/arch/copy_sm90_tma.hpp>
#include <cute/tensor.hpp>
#include <cute/atom/copy_traits.hpp>
#include <cute/atom/copy_atom.hpp>
namespace cute
{
//////////////////////////////////////////////////////////////////////////////
///////////////////////////// TMA_LOAD ///////////////////////////////////////
//////////////////////////////////////////////////////////////////////////////
struct SM90_TMA_LOAD_OP : SM90_TMA_LOAD {};
// The executable SM90_TMA_LOAD with tma_desc and tma_mbar
template <class NumBits>
struct Copy_Traits<SM90_TMA_LOAD_OP, NumBits>
{
using ThrID = Layout<_1>;
// Map from (src-thr,src-val) to bit
using SrcLayout = Layout<Shape<_1,NumBits>>;
// Map from (dst-thr,dst-val) to bit
using DstLayout = Layout<Shape<_1,NumBits>>;
// Reference map from (thr,val) to bit
using RefLayout = SrcLayout;
// SM90_TMA_LOAD arguments
TmaDescriptor const& tma_desc_;
uint64_t& tma_load_mbar_;
template <class Coord, int... Is>
CUTE_HOST_DEVICE constexpr
void
copy_unpack_(void const* const dst_ptr,
Coord const& src_coord, seq<Is...>) const
{
#if 0
print("THR (%d,%d,%d) BLK (%d,%d,%d)\n",
threadIdx.x, threadIdx.y, threadIdx.z,
blockIdx.x, blockIdx.y, blockIdx.z);
print(" TMA Coord "); print(src_coord); print("\n");
print(" TMA Shape "); print(make_tuple(uint64_t(tma_desc_.size0_),
uint64_t(tma_desc_.size1_),
uint64_t(tma_desc_.size2_),
uint64_t(tma_desc_.size3_))); print("\n");
#endif
SM90_TMA_LOAD::copy(&tma_desc_,
tma_load_mbar_,
dst_ptr,
get<Is>(src_coord)...);
}
// This is the copy_unpack dispatch for this Copy_Traits
// Src needs to be a gmem tensor with TmaCoordIterator .data()
// Dst needs to be a smem tensor
template <class TS, class SLayout,
class TD, class DLayout>
CUTE_HOST_DEVICE friend constexpr
void
copy_unpack(Copy_Traits const& traits,
Tensor<TS,SLayout> const& src,
Tensor<TD,DLayout> & dst)
{
//static_assert(is_gmem<TS>::value, "Expected gmem src for SM90_TMA_LOAD"); // TMA spoofed src tensor
static_assert(is_smem<TD>::value, "Expected smem dst for SM90_TMA_LOAD");
traits.copy_unpack_(dst.data().get(), src.data().coord_, tuple_seq<decltype(src.data().coord_)>{});
}
};
// The non-executable SM90_TMA_LOAD with tma_desc and no tma_mbar
// Use .with(tma_mbar) to construct an executable version
template <class NumBits, class GmemStrides>
struct Copy_Traits<SM90_TMA_LOAD, NumBits, GmemStrides>
{
using ThrID = Layout<_1>;
// Map from (src-thr,src-val) to bit
using SrcLayout = Layout<Shape<_1,NumBits>>;
// Map from (dst-thr,dst-val) to bit
using DstLayout = Layout<Shape<_1,NumBits>>;
// Reference map from (thr,val) to bit
using RefLayout = SrcLayout;
// SM90_TMA_LOAD arguments
TmaDescriptor tma_desc_;
GmemStrides g_stride_;
// Return TmaDescriptor/TensorMap
CUTE_HOST_DEVICE constexpr
TmaDescriptor const*
get_tma_descriptor() const {
return &tma_desc_;
}
// Construct an executable SM90_TMA_LOAD with tma_mbar
CUTE_HOST_DEVICE constexpr
Copy_Traits<SM90_TMA_LOAD_OP, NumBits>
with(uint64_t& tma_mbar, uint16_t const& multicast_mask = 0) const {
// We accept multicast_mask here to keep the API for both atoms consistent
// assert(multicast_mask == 0);
(void) multicast_mask;
return {tma_desc_, tma_mbar};
}
// Generate the TMA coord tensor
template <class GShape>
CUTE_HOST_DEVICE constexpr
auto
get_tma_tensor(GShape const& g_shape) const {
static_assert(is_congruent<decltype(g_shape), decltype(g_stride_)>::value);
constexpr int tma_rank = decltype(cute::min(rank(flatten(g_stride_)), Int<5>{}))::value;
return make_tensor(ArithmeticTupleIterator(as_arithmetic_tuple(repeat<tma_rank>(Int<0>{}))),
g_shape,
g_stride_);
}
// Don't try to execute a copy with SM90_TMA_LOAD before calling .with()
template <class TS, class SLayout,
class TD, class DLayout>
CUTE_HOST_DEVICE friend constexpr void
copy_unpack(Copy_Traits const& traits,
Tensor<TS,SLayout> const& src,
Tensor<TD,DLayout> & dst) = delete;
};
//////////////////////////////////////////////////////////////////////////////
///////////////////////////// TMA_LOAD_MULTICAST /////////////////////////////
//////////////////////////////////////////////////////////////////////////////
struct SM90_TMA_LOAD_MULTICAST_OP : SM90_TMA_LOAD_MULTICAST {};
template <class NumBits>
struct Copy_Traits<SM90_TMA_LOAD_MULTICAST_OP, NumBits>
{
using ThrID = Layout<_1>;
// Map from (src-thr,src-val) to bit
using SrcLayout = Layout<Shape<_1,NumBits>>;
// Map from (dst-thr,dst-val) to bit
using DstLayout = Layout<Shape<_1,NumBits>>;
// Reference map from (thr,val) to bit
using RefLayout = SrcLayout;
// SM90_TMA_LOAD_MULTICAST arguments
TmaDescriptor const& tma_desc_;
uint64_t& tma_load_mbar_;
uint16_t const& multicast_mask_;
template <class Coord, int... Is>
CUTE_HOST_DEVICE constexpr
void
copy_unpack_(void const* const dst_ptr,
Coord const& src_coord, seq<Is...>) const
{
#if 0
print("THR (%d,%d,%d) BLK (%d,%d,%d)\n",
threadIdx.x, threadIdx.y, threadIdx.z,
blockIdx.x, blockIdx.y, blockIdx.z);
print(" TMA Coord "); print(src_coord); print("\n");
print(" TMA Shape "); print(make_tuple(uint64_t(tma_desc_.size0_),
uint64_t(tma_desc_.size1_),
uint64_t(tma_desc_.size2_),
uint64_t(tma_desc_.size3_))); print("\n");
#endif
SM90_TMA_LOAD_MULTICAST::copy(&tma_desc_,
tma_load_mbar_,
multicast_mask_,
dst_ptr,
get<Is>(src_coord)...);
}
template <class TS, class SLayout,
class TD, class DLayout>
CUTE_HOST_DEVICE friend constexpr
void
copy_unpack(Copy_Traits const& traits,
Tensor<TS,SLayout> const& src,
Tensor<TD,DLayout> & dst)
{
//static_assert(is_gmem<TS>::value, "Expected gmem src for SM90_TMA_LOAD"); // TMA spoofed src tensor
static_assert(is_smem<TD>::value, "Expected smem dst for SM90_TMA_LOAD_MULTICAST");
traits.copy_unpack_(dst.data().get(), src.data().coord_, tuple_seq<decltype(src.data().coord_)>{});
}
};
template <class NumBits, class GmemStrides>
struct Copy_Traits<SM90_TMA_LOAD_MULTICAST, NumBits, GmemStrides>
{
using ThrID = Layout<_1>;
// Map from (src-thr,src-val) to bit
using SrcLayout = Layout<Shape<_1,NumBits>>;
// Map from (dst-thr,dst-val) to bit
using DstLayout = Layout<Shape<_1,NumBits>>;
// Reference map from (thr,val) to bit
using RefLayout = SrcLayout;
// SM90_TMA_LOAD_MULTICAST arguments
TmaDescriptor tma_desc_;
GmemStrides g_stride_;
// Return TmaDescriptor/TensorMap
CUTE_HOST_DEVICE constexpr
TmaDescriptor const*
get_tma_descriptor() const {
return &tma_desc_;
}
// Construct an executable SM90_TMA_LOAD_MULTICAST with tma_mbar
CUTE_HOST_DEVICE constexpr
Copy_Traits<SM90_TMA_LOAD_MULTICAST_OP, NumBits>
with(uint64_t& tma_load_mbar, uint16_t const& multicast_mask) const {
return {tma_desc_, tma_load_mbar, multicast_mask};
}
// Generate the TMA coord tensor
template <class GShape>
CUTE_HOST_DEVICE constexpr
auto
get_tma_tensor(GShape const& g_shape) const {
static_assert(is_congruent<decltype(g_shape), decltype(g_stride_)>::value);
constexpr int tma_rank = decltype(cute::min(rank(flatten(g_stride_)), Int<5>{}))::value;
return make_tensor(ArithmeticTupleIterator(as_arithmetic_tuple(repeat<tma_rank>(Int<0>{}))),
g_shape,
g_stride_);
}
// Don't try to execute a copy with SM90_TMA_LOAD_MULTICAST before calling .with()
template <class TS, class SLayout,
class TD, class DLayout>
CUTE_HOST_DEVICE friend constexpr void
copy_unpack(Copy_Traits const& traits,
Tensor<TS,SLayout> const& src,
Tensor<TD,DLayout> & dst) = delete;
};
//////////////////////////////////////////////////////////////////////////////
///////////////////////////// TMA_STORE //////////////////////////////////////
//////////////////////////////////////////////////////////////////////////////
// The executable SM90_TMA_STORE with tma_desc
template <class NumBits, class GmemStrides>
struct Copy_Traits<SM90_TMA_STORE, NumBits, GmemStrides>
{
using ThrID = Layout<_1>;
// Map from (src-thr,src-val) to bit
using SrcLayout = Layout<Shape<_1,NumBits>>;
// Map from (dst-thr,dst-val) to bit
using DstLayout = Layout<Shape<_1,NumBits>>;
// Reference map from (thr,val) to bit
using RefLayout = SrcLayout;
// SM90_TMA_STORE arguments
TmaDescriptor tma_desc_;
GmemStrides g_stride_;
// Return TmaDescriptor/TensorMap
CUTE_HOST_DEVICE constexpr
TmaDescriptor const*
get_tma_descriptor() const {
return &tma_desc_;
}
// Generate the TMA coord tensor
template <class GShape>
CUTE_HOST_DEVICE constexpr
auto
get_tma_tensor(GShape const& g_shape) const {
static_assert(is_congruent<decltype(g_shape), decltype(g_stride_)>::value);
constexpr int tma_rank = decltype(cute::min(rank(flatten(g_stride_)), Int<5>{}))::value;
return make_tensor(ArithmeticTupleIterator(as_arithmetic_tuple(repeat<tma_rank>(Int<0>{}))),
g_shape,
g_stride_);
}
template <class Coord, int... Is>
CUTE_HOST_DEVICE constexpr
void
copy_unpack_(void const* const src_ptr,
Coord const& dst_coord, seq<Is...>) const
{
#if 0
print("THR (%d,%d,%d) BLK (%d,%d,%d)\n",
threadIdx.x, threadIdx.y, threadIdx.z,
blockIdx.x, blockIdx.y, blockIdx.z);
print(" TMA Coord "); print(dst_coord); print("\n");
print(" TMA Shape "); print(make_tuple(uint64_t(tma_desc_.size0_),
uint64_t(tma_desc_.size1_),
uint64_t(tma_desc_.size2_),
uint64_t(tma_desc_.size3_))); print("\n");
#endif
SM90_TMA_STORE::copy(&tma_desc_,
src_ptr,
get<Is>(dst_coord)...);
}
// This is the copy_unpack dispatch for this Copy_Traits
// Src needs to be a smem tensor
// Dst needs to be a gmem tensor with TmaCoordIterator .data()
template <class TS, class SLayout,
class TD, class DLayout>
CUTE_HOST_DEVICE friend constexpr
void
copy_unpack(Copy_Traits const& traits,
Tensor<TS,SLayout> const& src,
Tensor<TD,DLayout> & dst)
{
static_assert(is_smem<TS>::value, "Expected smem src for SM90_TMA_STORE");
//static_assert(is_gmem<TD>::value, "Expected gmem dst for SM90_TMA_STORE"); // TMA spoofed src tensor
traits.copy_unpack_(src.data().get(), dst.data().coord_, tuple_seq<decltype(dst.data().coord_)>{});
}
};
//////////////////////////////////////////////////////////////////////////////
///////////////////////////// BULK COPY //////////////////////////////////////
//////////////////////////////////////////////////////////////////////////////
template <class NumBits, class... OpArgs>
struct Copy_Traits<SM90_BULK_COPY_G2S, NumBits, OpArgs...>
{
static_assert(int32_t(NumBits::value / 8) % 16 == 0,
"Bulk Copy requires copy vector size align to 16B.");
using ThrID = Layout<_1>;
// Map from (src-thr,src-val) to bit
using SrcLayout = Layout<Shape<_1,NumBits>>;
// Map from (dst-thr,dst-val) to bit
using DstLayout = Layout<Shape<_1,NumBits>>;
// Reference map from (thr,val) to bit
using RefLayout = SrcLayout;
// SM90_BULK_COPY_G2S arguments
// 0: uint64_t* bulk_load_memory_barrier
cute::tuple<OpArgs...> bulk_load_mbar_;
template <class TS, class SLayout,
class TD, class DLayout>
CUTE_HOST_DEVICE friend constexpr
void
copy_unpack(Copy_Traits const& traits,
Tensor<TS,SLayout> const& src,
Tensor<TD,DLayout> & dst)
{
static_assert(is_same<cute::tuple<OpArgs...>, cute::tuple<uint64_t*>>::value,
"Extra arguments not set. Set .with() before use.");
static_assert(is_gmem<TS>::value, "Expected gmem src for SM90_BULK_COPY_G2S");
static_assert(is_smem<TD>::value, "Expected smem dst for SM90_BULK_COPY_G2S");
SM90_BULK_COPY_G2S::copy(src.data().get(), *get<0>(traits.bulk_load_mbar_),
dst.data().get(), int32_t(NumBits::value / 8));
}
// Record the memory barrier for the instruction
CUTE_HOST_DEVICE constexpr
Copy_Traits<SM90_BULK_COPY_G2S, NumBits, uint64_t*>
with(uint64_t& bulk_mbar) const {
return {{&bulk_mbar}};
}
};
template <class NumBits>
struct Copy_Traits<SM90_BULK_COPY_S2G, NumBits>
{
static_assert(int32_t(NumBits::value / 8) % 16 == 0,
"Bulk Copy requires copy vector size align to 16B.");
using ThrID = Layout<_1>;
// Map from (src-thr,src-val) to bit
using SrcLayout = Layout<Shape<_1,NumBits>>;
// Map from (dst-thr,dst-val) to bit
using DstLayout = Layout<Shape<_1,NumBits>>;
// Reference map from (thr,val) to bit
using RefLayout = SrcLayout;
template <class TS, class SLayout,
class TD, class DLayout>
CUTE_HOST_DEVICE friend constexpr
void
copy_unpack(Copy_Traits const& traits,
Tensor<TS,SLayout> const& src,
Tensor<TD,DLayout> & dst)
{
static_assert(is_smem<TS>::value, "Expected smem src for SM90_BULK_COPY_S2G");
static_assert(is_gmem<TD>::value, "Expected gmem dst for SM90_BULK_COPY_S2G");
SM90_BULK_COPY_S2G::copy(src.data().get(), dst.data().get(), int32_t(NumBits::value / 8));
}
};
//
// Placeholder for the bulk copy algorithm's default, auto-vectorizing behavior
//
template <class... OpArgs>
struct Copy_Traits<SM90_BULK_COPY_AUTO, OpArgs...>
{
// Logical thread id to thread idx (one-thread)
using ThrID = Layout<_1>;
// Map from (src-thr,src-val) to bit
using SrcLayout = Layout<Shape<_1,_1>, Stride<_0,_0>>;
// Map from (dst-thr,dst-val) to bit
using DstLayout = Layout<Shape<_1,_1>, Stride<_0,_0>>;
// Reference map from (thr,val) to bit
using RefLayout = SrcLayout;
// SM90_UBULK_COPY arguments
// 0: uint64_t* bulk_load_memory_barrier [if this is a BULK_LOAD_G2S]
cute::tuple<OpArgs...> opargs_;
// Record the memory barrier for the instruction
CUTE_HOST_DEVICE constexpr
Copy_Traits<SM90_BULK_COPY_AUTO, uint64_t*>
with(uint64_t& bulk_mbar) const {
return {{&bulk_mbar}};
}
};
//
// MAKE_TMA_COPY and related
//
namespace detail
{
template <int B, int M, int S, class Offset, class SLayout>
auto
get_swizzle_portion(ComposedLayout<Swizzle<B,M,S>,Offset,SLayout>)
{
return Swizzle<B,M,S>{};
}
template <class Shape, class Stride>
auto
get_swizzle_portion(Layout<Shape,Stride>)
{
return Swizzle<0,4,3>{};
}
template <int B, int M, int S, class Offset, class SLayout>
auto
get_nonswizzle_portion(ComposedLayout<Swizzle<B,M,S>,Offset,SLayout> const& slayout)
{
return slayout.layout_fn();
}
template <class Shape, class Stride>
auto
get_nonswizzle_portion(Layout<Shape,Stride> const& slayout)
{
return slayout;
}
template <int B, int M, int S>
TMA::SmemSwizzleBits
get_tma_swizzle_bits(Swizzle<B,M,S>)
{
if constexpr (M == 4) {
switch (B) {
default: static_assert(0 <= B && B <= 3, "Expected B = 0,1,2, or 3 when M == 4. Unsupported layout swizzle.");
case 3: return TMA::SmemSwizzleBits::B128;
case 2: return TMA::SmemSwizzleBits::B64;
case 1: return TMA::SmemSwizzleBits::B32;
case 0: return TMA::SmemSwizzleBits::DISABLE;
}
} else
{
static_assert(M < 0, "Unsupported layout swizzle.");
}
}
template <class Layout>
TMA::SmemSwizzleBits
get_tma_swizzle_bits(Layout const& layout)
{
return get_tma_swizzle_bits(get_swizzle_portion(layout));
}
#if !defined(__MACACC_RTC__)
// Use a smem2gmode map to read through the GMEM tensor
// and construct a TMA Descriptor for the resulting instruction
template <class GEngine, class GLayout,
class SShape, class SStride,
int B, int M, int S>
CUTE_HOST
auto
make_tma_copy_desc(Tensor<GEngine,GLayout> const& gtensor, // The original GMEM Tensor
Layout<SShape,SStride> const& smem_inv, // smem_idx to flat gmode
Swizzle<B,M,S> const& swizzle) // Swizzle fn on smem_idx
{
using T = typename GEngine::value_type;
auto flat_glayout = flatten(gtensor.layout());
CUTE_STATIC_ASSERT_V(rank(flat_glayout) == rank(smem_inv));
constexpr int rank_smem_inv = decltype(rank(smem_inv))::value;
auto tma_multimode = rank(flat_glayout) > Int<5>{};
constexpr uint32_t tma_dim = cute::min(rank(flat_glayout), 5);;
//
// TMA gmem desc info
//
void* gmem_address = (void*) gtensor.data();
cute::array<uint64_t, 5> gmem_prob_shape = {1,1,1,1,1};
cute::array<uint64_t, 5> gmem_prob_stride = {0,0,0,0,0};
for_each(make_seq<rank_smem_inv>{}, [&](auto i) {
auto e = stride<i>(smem_inv); // For g++-7.5, let it deduce e rather than fuse with below
constexpr int j = decltype(e.mode())::value;
constexpr int tma_i = i < 5 ? i : 4;
// Problem stride
uint64_t stride_j = stride<j>(flat_glayout) * sizeof(T);
uint64_t old_stride = gmem_prob_stride[tma_i];
gmem_prob_stride[tma_i] = gcd(gmem_prob_stride[tma_i], stride_j);
// Problem shape
uint64_t shape_j = shape<j>(flat_glayout);
if (gmem_prob_stride[tma_i] != 0) {
// We're "resetting" this TMA mode and using it as a "multimode"
// Recurrence: g_shape = (s_i - 1) * (d_i / gcd_j d_j) + 1
gmem_prob_shape[tma_i] = (gmem_prob_shape[tma_i]-1) * (old_stride / gmem_prob_stride[tma_i])
+ (shape_j-1) * (stride_j / gmem_prob_stride[tma_i])
+ 1;
} else {
gmem_prob_shape[tma_i] = shape_j;
}
});
assert((reinterpret_cast<uint64_t>(gmem_address) & 0b1111) == 0); // Address must be 16B-aligned
assert(gmem_prob_shape[0] >= (uint64_t(1))); // Size must be min 1
assert(gmem_prob_shape[0] <= (uint64_t(1) << 32)); // Size must be max 2^32
assert(gmem_prob_shape[1] >= (uint64_t(1))); // Size must be min 1
assert(gmem_prob_shape[1] <= (uint64_t(1) << 32)); // Size must be max 2^32
assert(gmem_prob_shape[2] >= (uint64_t(1))); // Size must be min 1
assert(gmem_prob_shape[2] <= (uint64_t(1) << 32)); // Size must be max 2^32
assert(gmem_prob_shape[3] >= (uint64_t(1))); // Size must be min 1
assert(gmem_prob_shape[3] <= (uint64_t(1) << 32)); // Size must be max 2^32
assert(gmem_prob_shape[4] >= (uint64_t(1))); // Size must be min 1
assert(gmem_prob_shape[4] <= (uint64_t(1) << 32)); // Size must be max 2^32
assert((gmem_prob_stride[0]) == sizeof(T)); // First stride is implicitly 1
assert((gmem_prob_stride[1]) < (uint64_t(1) << 40)); // Stride must be max 2^40
assert((gmem_prob_stride[1] & 0b1111) == 0); // Stride must be multiple of 16B (128b)
assert((gmem_prob_stride[2]) < (uint64_t(1) << 40)); // Stride must be max 2^40
assert((gmem_prob_stride[2] & 0b1111) == 0); // Stride must be multiple of 16B (128b)
assert((gmem_prob_stride[3]) < (uint64_t(1) << 40)); // Stride must be max 2^40
assert((gmem_prob_stride[3] & 0b1111) == 0); // Stride must be multiple of 16B (128b)
assert((gmem_prob_stride[4]) < (uint64_t(1) << 40)); // Stride must be max 2^40
assert((gmem_prob_stride[4] & 0b1111) == 0); // Stride must be multiple of 16B (128b)
//
// TMA smem desc info
//
cute::array<uint32_t, 5> smem_box_shape = {1,1,1,1,1};
cute::array<uint32_t, 5> smem_box_stride = {1,1,1,1,1};
for_each(make_seq<rank_smem_inv>{}, [&](auto i) {
uint32_t shape_i = shape<i>(smem_inv);
constexpr int tma_i = i < 5 ? i : 4;
if (tma_multimode && tma_i == 4) {
// We're "reusing" this TMA mode and using it as a "multimode"
smem_box_shape[tma_i] = 1;
} else {
smem_box_shape[tma_i] = shape_i;
}
});
assert(smem_box_shape[0] >= (uint64_t(1))); // Size must be min 1
assert(smem_box_shape[0] <= (uint64_t(1) << 8)); // Size must be max 2^8
assert(smem_box_shape[0] >= (uint64_t(1))); // Size must be min 1
assert(smem_box_shape[0] <= (uint64_t(1) << 8)); // Size must be max 2^8
assert(smem_box_shape[0] >= (uint64_t(1))); // Size must be min 1
assert(smem_box_shape[0] <= (uint64_t(1) << 8)); // Size must be max 2^8
assert(smem_box_shape[0] >= (uint64_t(1))); // Size must be min 1
assert(smem_box_shape[0] <= (uint64_t(1) << 8)); // Size must be max 2^8
assert(smem_box_stride[0] >= (uint32_t(1))); // Stride must be min 1
assert(smem_box_stride[0] <= (uint32_t(8))); // Stride must be max 2^3
assert(smem_box_stride[1] >= (uint32_t(1))); // Stride must be min 1
assert(smem_box_stride[1] <= (uint32_t(8))); // Stride must be max 2^3
assert(smem_box_stride[2] >= (uint32_t(1))); // Stride must be min 1
assert(smem_box_stride[2] <= (uint32_t(8))); // Stride must be max 2^3
assert(smem_box_stride[3] >= (uint32_t(1))); // Stride must be min 1
assert(smem_box_stride[3] <= (uint32_t(8))); // Stride must be max 2^3
assert(smem_box_stride[4] >= (uint32_t(1))); // Stride must be min 1
assert(smem_box_stride[4] <= (uint32_t(8))); // Stride must be max 2^3
//
// Construct the descriptor
//
TmaDescriptor tma_desc = {0};
//
// TMA general info
//
// #if (__CUDACC_VER_MAJOR__ >= 12)
#if 0
CUtensorMapDataType tma_format = TMA::to_CUtensorMapDataType<T>();
CUtensorMapInterleave tma_interleave = CU_TENSOR_MAP_INTERLEAVE_NONE;
CUtensorMapL2promotion tma_l2Promotion = CU_TENSOR_MAP_L2_PROMOTION_NONE;
CUtensorMapFloatOOBfill tma_oobFill = CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE;
// TMA smem swizzle type
CUtensorMapSwizzle smem_swizzle = TMA::to_CUtensorMapSwizzle(get_tma_swizzle_bits(swizzle));
CUresult result = cuTensorMapEncodeTiled(
&tma_desc,
tma_format,
tma_dim,
gmem_address,
gmem_prob_shape.data(),
gmem_prob_stride.data() + 1, // gmem_prob_stride[0] implicitly 1
smem_box_shape.data(),
smem_box_stride.data(),
tma_interleave,
smem_swizzle,
tma_l2Promotion,
tma_oobFill);
if (result != CUDA_SUCCESS) {
std::cerr << "TMA Desc Addr: " << &tma_desc
<< "\nformat " << tma_format
<< "\ndim " << tma_dim
<< "\ngmem_address " << gmem_address
<< "\nglobalDim " << gmem_prob_shape
<< "\nglobalStrides " << gmem_prob_stride
<< "\nboxDim " << smem_box_shape
<< "\nelementStrides " << smem_box_stride
<< "\ninterleave " << tma_interleave
<< "\nswizzle " << smem_swizzle
<< "\nl2Promotion " << tma_l2Promotion
<< "\noobFill " << tma_oobFill << std::endl;
std::cerr << "Error: Failed to initialize the TMA descriptor " << result << std::endl;
assert(false);
}
#endif // (__CUDACC_VER_MAJOR__ >= 12)
// Finally, get the inverse permutation of the E<i> bases for the mocked gmem stride
auto gmem_stride_bases_flat = transform(make_seq<rank_smem_inv>{}, [&](auto i) {
auto k = find(stride(smem_inv), E<i>{});
// For gcc 7.5 -- avoid 'if constexpr'
int32_t tma_coord_stride = int32_t(stride<i>(flat_glayout) * sizeof(T) / (gmem_prob_stride[4] != 0 ? gmem_prob_stride[4] : 16));
return conditional_return(tma_multimode && (k >= Int<4>{}),
E<4>{} * tma_coord_stride, // The 4th TMA mode is the multimode, use int32_t coord stride
E<k>{});
});
// Give that the profile of gtensor and fold it
// NOTE: This is the only reason we want the original gtensor shape rather than the more intuitive flattened shape
auto gmem_stride_bases = stride(composition(make_layout(repeat_like(shape(flat_glayout), Int<2>{}), gmem_stride_bases_flat),
make_layout(repeat_like(shape(gtensor), Int<2>{}))));
return make_tuple(tma_desc, gmem_stride_bases);
}
template <class CopyOp,
class GEngine, class GLayout,
class SLayout,
class TShape, class TStride,
class VShape, class VStride>
CUTE_HOST
auto
make_tma_copy_tiled(CopyOp,
Tensor<GEngine,GLayout> const& gtensor, // Full GMEM Tensor
SLayout const& slayout, // CTA Tile of SMEM
Layout<TShape,TStride> const& cta_t_map, // T: CTA thr idx -> logical TMA tid
Layout<VShape,VStride> const& cta_v_map) // V: CTA val idx -> gmem coord
{
//
// TMA parameter checking
//
CUTE_STATIC_ASSERT_V(product_each(shape(slayout)) == product_each(shape(cta_v_map)),
"TMA requires CTA_Tile and SLayout top-level shape equivalence.");
CUTE_STATIC_ASSERT_V(size(slayout) % cosize(cta_t_map) == Int<0>{},
"Number of active CTAs in TMA must divide domain size of slayout.");
//
// TMA slayout manipulation
//
auto flat_glayout = flatten(gtensor.layout());
// Invert the smem to get the largest contiguous vector in the smem layout
auto inv_smem_layout = right_inverse(get_nonswizzle_portion(slayout));
// trunc_smem_idx -> trunc_smem_coord
// Map from smem idx to a gmem mode
auto sidx_to_gmode = coalesce(composition(cta_v_map, inv_smem_layout));
// Truncate any incompatibilities
auto smem_rank = find_if(stride(sidx_to_gmode), [](auto e) {
auto v = basis_value(e);
return not is_constant<1,decltype(v)>{};
});
static_assert(smem_rank > 0, "Could not find a common smem-gmem vectorization for TMA. Do they have a common majorness?");
// TMA uses a maximum of 5 modes
// If the gtensor has more than 5 modes, we need to reserve the last TMA-mode as a "multimode"
constexpr int smem_tma_rank = cute::min(int(smem_rank), (rank(flat_glayout) > Int<5>{} ? 4 : 5));
// Keep only the static-1 basis modes into gmem
auto sidx_to_gmode_trunc = take<0,smem_tma_rank>(sidx_to_gmode);
// Split according to the portion each multicast CTA will be responsible for
auto sidx_to_gmode_vt = logical_divide(sidx_to_gmode_trunc, shape_div(size(sidx_to_gmode_trunc), cosize(cta_t_map)));
#if 0
print("g_layout : "); print(gtensor.layout()); print("\n");
print("s_layout : "); print(slayout); print("\n");
print("cta_t_map : "); print(cta_t_map); print("\n");
print("cta_v_map : "); print(cta_v_map); print("\n");
print("inv_smem : "); print(inv_smem_layout); print("\n");
print("sidx_to_gmode : "); print(sidx_to_gmode); print("\n");
print("sidx_to_gmode_trunc : "); print(sidx_to_gmode_trunc); print("\n");
print("sidx_to_gmode_vt : "); print(sidx_to_gmode_vt); print("\n");
#endif
//
// TMA gtensor manipulation
//
// Generate a TupleBasis for the gtensor
auto flat_gbasis = make_basis_like(shape(flat_glayout));
// Fold the flat_gbasis into the glayout
auto glayout_basis = make_layout(shape(gtensor),
stride(composition(make_layout(repeat_like(shape(flat_glayout), Int<2>{}), flat_gbasis),
make_layout(repeat_like(shape(gtensor), Int<2>{})))));
// Tile the modes of gtensor with the truncated cta_v_map o inv_smem_layout_trunc
auto tma_layout_v_trunc = flatten(composition(glayout_basis, layout<0>(sidx_to_gmode_vt)));
// Append any missing basis on the end as size-1 modes b/c they got truncated
// NOTE This is essentially ArithmeticTuple complement...
auto missing_basis = fold(stride(tma_layout_v_trunc), flat_gbasis, [](auto init, auto e) {
auto k = find(init, e);
return remove<k>(init);
});
// The appended map from truncated smem codomain to gmem mode: trunc_smem_idx -> gmem_mode
auto tma_layout_v = make_layout(flatten(cute::make_tuple(tma_layout_v_trunc.shape(), repeat<rank(missing_basis)>(Int<1>{}))),
flatten(cute::make_tuple(tma_layout_v_trunc.stride(), missing_basis)));
#if 0
print("flat_gbasis : "); print(flat_gbasis); print("\n");
print("missing_b : "); print(missing_basis); print("\n");
print("tma_layout_v : "); print(tma_layout_v); print("\n");
#endif
//
// Construct the TMA Desc and GMEM mode ordering
//
auto [tma_desc, gmem_stride_bases] = detail::make_tma_copy_desc(gtensor, tma_layout_v, get_swizzle_portion(slayout));
//
// Construct the Copy_Traits
//
using T = typename GEngine::value_type;
constexpr int num_bits = decltype(size<0>(sidx_to_gmode_vt))::value * sizeof(T) * 8;
using Traits = Copy_Traits<CopyOp, Int<num_bits>, decltype(gmem_stride_bases)>;
#if 0
print("num_bits : "); print(num_bits); print("\n");
print("g_stride_bases: "); print(gmem_stride_bases); print("\n");
#endif
Traits tma_traits{tma_desc, gmem_stride_bases};
//
// Construct the TiledCopy
//
auto cta_tiler = product_each(shape(cta_v_map));
// (CTA V, CTA T) -> smem_coord
auto layout_vt = composition(inv_smem_layout, make_layout(shape(sidx_to_gmode_vt)));
// Scale that up to cover all of the smem_coords
auto layout_VT = tile_to_shape(layout_vt, make_shape(size(cta_v_map)/size<1>(layout_vt), size<1>(layout_vt)));
// Flip it and change the domain of the T from logical thr to thr_idx
auto layout_TV = make_layout(composition(layout<1>(layout_VT), cta_t_map), layout<0>(layout_VT));
#if 0
print("cta_tiler : "); print(cta_tiler); print("\n");
print("layout_VT : "); print(layout_VT); print("\n");
print("layout_TV : "); print(layout_TV); print("\n");
#endif
using T = typename GEngine::value_type;
return TiledCopy<Copy_Atom<Traits,T>, decltype(layout_TV), decltype(cta_tiler)>{tma_traits};
}
#endif // !defined(__MACACC_RTC__)
} // end namespace detail
/** Make a CuTe CTA-collective TiledCopy for a TMA operation.
*
* @param CopyOp The target copy operation: SM90_TMA_LOAD, SM90_TMA_LOAD_MULTICAST, SM90_TMA_STORE
* @param gtensor The GMEM Tensor to be involved in the TMA.
* @param slayout The SMEM Layout to be involved in the TMA.
* @param cta_tile The CTA-local tile that each CTA will be tiling GMEM with.
* This is often the blk_shape that is used to tile the GMEM for CTAs:
* local_tile(gtensor, blk_shape, blk_coord) -> CTA-local tile of gtensor
* @param cluster_size When using SM90_TMA_LOAD_MULTICAST, this can be a (static) power-of-2 <= 16
* defining the multicast size (used to further partition the SMEM)
* Else, static-1
*
* This code attempts to maximize the TMA box size. It does this by tracing
* the SMEM "vector" -- the inverse of the smem layout -- to find the largest
* contiguous array of smem that can be written to/from global memory given
* the constraints that the TMA instruction imposes.
*
* This is accomplished by assigning "basis" strides to the GMEM to track which
* modes of SMEM map to which modes of GMEM, then reorder the modes of GMEM according
* to the SMEM vector, and then using those GMEM/SMEM modes to fill in the desc.
*
* Examples:
using T = float;
T* gptr = nullptr;
{
// Simple 2D
Tensor gtensor = make_tensor(gptr, make_shape(1024, 256), GenRowMajor{}); // K-Major GMEM
auto slayout = make_layout(make_shape(_64{}, _32{}), GenRowMajor{}); // K-Major SMEM
auto tma = make_tma_copy(SM90_TMA_LOAD{}, gtensor, slayout);
}
{
// GMMA 2D
Tensor gtensor = make_tensor(gptr, make_shape(1024, 256)); // MN-Major GMEM
auto slayout = tile_to_shape(GMMA::Layout_MN_SW128_Atom<T>{}, make_shape(_128{},_64{})); // MN-Major Swizzled+Tiled 128x64 SMEM
auto tma = make_tma_copy(SM90_TMA_LOAD{}, gtensor, slayout);
}
{
// 3D
Tensor gtensor = make_tensor(gptr, make_shape(1024, 32, 512), make_stride(64, Int<1>{}, 65536)); // GMEM
auto slayout = make_layout(make_shape(_16{}, _8{}, _2{}), make_stride(_16{}, _1{}, _8{})); // SMEM w/ same major-mode
auto tma = make_tma_copy(SM90_TMA_LOAD{}, gtensor, slayout);
}
{
// cuTENSOR 4D
auto layout = make_shape(make_shape(32,40),make_shape(make_shape(8,8),656)); // GMEM
auto cta_tile = make_shape(_128{},make_shape(_32{},_2{})); // GMEM Tiling:
// Take 128-elem from m: m0 must divide 128,
// m-last may be predicated
// Take 32-elem from k0, 2-elem from k1
auto slayout = make_layout(cta_tile); // Col-Major SMEM
auto tma = make_tma_copy(SM90_TMA_LOAD{}, gtensor, slayout, cta_tile, Int<1>{});
}
*
* Check the TMA box size and desc:
print("TMA Box size: "); print(typename decltype(tma)::Tiler_MN{}); print("\n");
print("TMA desc : "); print(tma.tma_desc_); print("\n");
*
* Usage:
Tensor mA = tma_a.get_tma_tensor(make_shape(M,N)); // (M,N) TMA coord tensor
Tensor gA = local_tile(mA, cta_tile, cta_coord); // (BLK_M,BLK_N) TMA coord tensor for this CTA
Tensor sA = make_tensor(make_smem_ptr<T>(sptr), slayout); // (BLK_M,BLK_N) SMEM tensor
auto cta_tma = tma.get_slice(cta_idx_in_cluster); // Slice for multicast partitioning
Tensor tAgA = cta_tma.partition_S(gA); // Partition for src
Tensor tAsA = cta_tma.partition_D(sA); // Partition for dst
copy(tma.with(barrier, mcast_mask), tAgA, tAsA); // copy with supporting TMA params
*/
#if !defined(__MACACC_RTC__)
template <class CopyOp,
class GEngine, class GLayout,
class SLayout,
class CTA_Tile,
class Cluster_Size>
CUTE_HOST
auto
make_tma_copy(CopyOp const& copy_op,
Tensor<GEngine,GLayout> const& gtensor,
SLayout const& slayout,
CTA_Tile const& cta_tile,
Cluster_Size const& cluster_size)
{
return detail::make_tma_copy_tiled(copy_op,
gtensor,
slayout,
make_layout(cluster_size),
make_identity_layout(cta_tile));
}
// Explicit defaulting
template <class CopyOp,
class GEngine, class GLayout,
class SLayout>
CUTE_HOST
auto
make_tma_copy(CopyOp const& copy_op,
Tensor<GEngine,GLayout> const& gtensor,
SLayout const& slayout)
{
return make_tma_copy(copy_op, gtensor, slayout, product_each(shape(slayout)), Int<1>{});
}
template <class CopyOp,
class GEngine, class GLayout,
class SLayout,
class Cluster_Size>
CUTE_HOST
auto
make_tma_copy(CopyOp const& copy_op,
Tensor<GEngine,GLayout> const& gtensor,
SLayout const& slayout,
Cluster_Size const& cluster_size)
{
return make_tma_copy(copy_op, gtensor, slayout, product_each(shape(slayout)), cluster_size);
}
#endif // !defined(__MACACC_RTC__)
} // end namespace cute

File diff suppressed because it is too large Load Diff

View File

@ -0,0 +1,208 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#pragma once
#include <cute/arch/mma.hpp>
#include <cute/tensor.hpp>
namespace cute
{
namespace detail {
template <class X, class = void>
struct supports_output_scaling { static constexpr bool value = false; };
template <class X>
struct supports_output_scaling<X, void_t<decltype(declval<X>().accumulate_)>> { static constexpr bool value = true; };
} // end namespace detail
/**
* concept MMA_Traits
* {
* using ElementDVal = // Logical A-value type
* using ElementAVal = // Logical B-value type
* using ElementBVal = // Logical C-value type
* using ElementCVal = // Logical D-value type (NOTE: Not used? Assumed == ElementDVal)
*
* using ElementAFrg = // A-type consumed by MMA (if ommitted, same as ElementAVal)
* using ElementBFrg = // B_type consumed by MMA (if ommitted, same as ElementBVal)
* using ElementCFrg = // C_type consumed by MMA (if ommitted, same as ElementCVal)
*
* using Shape_MNK = // Logical MxNxK shape of the MMA
*
* using ThrID = // Logical thread id (tid) -> tidx
*
* using ALayout = // (Logical thread id (tid), Logical value id (vid)) -> Flat MK-coord
* using BLayout = // (Logical thread id (tid), Logical value id (vid)) -> Flat NK-coord
* using CLayout = // (Logical thread id (tid), Logical value id (vid)) -> Flat MN-coord
* };
*/
template <class MMAOperation, class... MMAOpArgs>
struct MMA_Traits
{
static_assert(sizeof(MMAOperation) == 0, "MMA_Traits not implemented for this MMA_Operation.");
};
template <class D, class A, class B, class C>
struct MMA_Traits<UniversalFMA<D,A,B,C>>
{
using ElementDVal = D;
using ElementAVal = A;
using ElementBVal = B;
using ElementCVal = C;
// Logical shape of the MMA
using Shape_MNK = Shape<_1,_1,_1>;
// Logical thread id (tid) -> tidx
using ThrID = Layout<_1>;
// (Logical thread id (tid), Logical value id (vid)) -> coord
// (tid,vid) -> (m,k)
using ALayout = Layout<Shape<_1,_1>>;
// (tid,vid) -> (n,k)
using BLayout = Layout<Shape<_1,_1>>;
// (tid,vid) -> (m,n)
using CLayout = Layout<Shape<_1,_1>>;
};
//
// Generic mma_unpack for any MMA_Traits
//
template <class MMA_Op, class... MMA_Args,
class TD, class DLayout,
class TA, class ALayout,
class TB, class BLayout,
class TC, class CLayout>
CUTE_HOST_DEVICE constexpr
void
mma_unpack(MMA_Traits<MMA_Op, MMA_Args...> const& traits,
Tensor<TD, DLayout> & D,
Tensor<TA, ALayout> const& A,
Tensor<TB, BLayout> const& B,
Tensor<TC, CLayout> const& C)
{
static_assert(is_rmem<TD>::value, "Expected registers in MMA_Atom::call");
static_assert(is_rmem<TA>::value, "Expected registers in MMA_Atom::call");
static_assert(is_rmem<TB>::value, "Expected registers in MMA_Atom::call");
static_assert(is_rmem<TC>::value, "Expected registers in MMA_Atom::call");
// Register value types from the MMA_Operation register arrays
using RegTypeD = typename remove_extent<typename MMA_Op::DRegisters>::type;
using RegTypeA = typename remove_extent<typename MMA_Op::ARegisters>::type;
using RegTypeB = typename remove_extent<typename MMA_Op::BRegisters>::type;
using RegTypeC = typename remove_extent<typename MMA_Op::CRegisters>::type;
using MMATraits = MMA_Traits<MMA_Op, MMA_Args...>;
constexpr int RegNumD = extent<typename MMA_Op::DRegisters>::value;
constexpr int RegNumA = extent<typename MMA_Op::ARegisters>::value;
constexpr int RegNumB = extent<typename MMA_Op::BRegisters>::value;
constexpr int RegNumC = extent<typename MMA_Op::CRegisters>::value;
Tensor rA = recast<RegTypeA>(A);
Tensor rB = recast<RegTypeB>(B);
CUTE_STATIC_ASSERT_V(size(rA) == Int<RegNumA>{});
CUTE_STATIC_ASSERT_V(size(rB) == Int<RegNumB>{});
if constexpr (is_same<RegTypeD, void>::value)
{
static_assert(is_same<typename TD::value_type, typename TC::value_type>::value, "GMMA C and D value_type must match.");
static_assert(is_same<DLayout, CLayout>::value, "GMMA C and D layouts must match.");
// assert((void*)&C == (void*)&D);
Tensor rC = recast<RegTypeC>(D); // NOTE: D and C are same, so use mutable D
//CUTE_STATIC_ASSERT_V(size(rC) == Int<RegNumC>{});
if constexpr (detail::supports_output_scaling<MMATraits>::value) {
detail::explode_with_d_scaling(MMA_Op::fma,
rA, make_int_sequence<RegNumA>{},
rB, make_int_sequence<RegNumB>{},
rC, make_int_sequence<RegNumC>{},
traits.accumulate_);
}
else {
detail::explode(MMA_Op::fma,
rA, make_int_sequence<RegNumA>{},
rB, make_int_sequence<RegNumB>{},
rC, make_int_sequence<RegNumC>{});
}
}
else {
Tensor rD = recast<RegTypeD>(D);
Tensor rC = recast<RegTypeC>(C);
CUTE_STATIC_ASSERT_V(size(rD) == Int<RegNumD>{});
CUTE_STATIC_ASSERT_V(size(rC) == Int<RegNumC>{});
if constexpr (detail::supports_output_scaling<MMATraits>::value) {
detail::explode_with_d_scaling(MMA_Op::fma,
rD, make_int_sequence<RegNumD>{},
rA, make_int_sequence<RegNumA>{},
rB, make_int_sequence<RegNumB>{},
rC, make_int_sequence<RegNumC>{},
traits.accumulate_);
}
else {
detail::explode(MMA_Op::fma,
rD, make_int_sequence<RegNumD>{},
rA, make_int_sequence<RegNumA>{},
rB, make_int_sequence<RegNumB>{},
rC, make_int_sequence<RegNumC>{});
}
}
}
namespace detail {
template <class X, class = void>
struct FrgTypeA_or_Default { using type = typename X::ElementAVal; };
template <class X>
struct FrgTypeA_or_Default<X,void_t<typename X::ElementAFrg>> { using type = typename X::ElementAFrg; };
template <class X, class = void>
struct FrgTypeB_or_Default { using type = typename X::ElementBVal; };
template <class X>
struct FrgTypeB_or_Default<X,void_t<typename X::ElementBFrg>> { using type = typename X::ElementBFrg; };
template <class X, class = void>
struct FrgTypeC_or_Default { using type = typename X::ElementCVal; };
template <class X>
struct FrgTypeC_or_Default<X,void_t<typename X::ElementCFrg>> { using type = typename X::ElementCFrg; };
} // end namespace detail
} // namespace cute

View File

@ -0,0 +1,81 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#pragma once
#include <cute/arch/mma_sm75.hpp>
#include <cute/atom/mma_traits.hpp>
#include <cute/layout.hpp>
namespace cute
{
template <>
struct MMA_Traits<SM75_16x8x8_F32F16F16F32_TN>
{
using ElementDVal = float;
using ElementAVal = half_t;
using ElementBVal = half_t;
using ElementCVal = float;
using Shape_MNK = Shape<_16,_8,_8>;
using ThrID = Layout<_32>;
using ALayout = Layout<Shape <Shape < _4,_8>,Shape < _2,_2>>,
Stride<Stride<_32,_2>,Stride<_16,_1>>>;
using BLayout = Layout<Shape <Shape < _4,_8>,_2>,
Stride<Stride<_16,_1>,_8>>;
using CLayout = Layout<Shape <Shape < _4,_8>,Shape < _2,_2>>,
Stride<Stride<_32,_2>,Stride<_16,_1>>>;
};
///////////////////////////////////////////////////////////////////////////////
template <>
struct MMA_Traits<SM75_8x8x16_S32S8S8S32_TN>
{
using ElementDVal = int32_t;
using ElementAVal = int8_t;
using ElementBVal = int8_t;
using ElementCVal = int32_t;
using Shape_MNK = Shape<_8,_8,_16>;
using ThrID = Layout<_32>;
using ALayout = Layout<Shape <Shape < _4,_8>,_4>,
Stride<Stride<_32,_1>,_8>>;
using BLayout = Layout<Shape <Shape < _4,_8>,_4>,
Stride<Stride<_32,_1>,_8>>;
using CLayout = Layout<Shape <Shape < _4,_8>,_2>,
Stride<Stride<_16,_1>,_8>>;
};
///////////////////////////////////////////////////////////////////////////////
} // namespace cute

View File

@ -0,0 +1,604 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#pragma once
#include <cute/arch/mma_sm80.hpp>
#include <cute/atom/mma_traits.hpp>
#include <cute/layout.hpp>
#include <cute/numeric/integer_subbyte.hpp>
#include <mctlass/numeric_types.h>
namespace cute
{
namespace {
// (T32,V1) -> (M8,N8)
using SM80_8x4 = Layout<Shape <Shape < _4,_8>,_1>,
Stride<Stride< _8,_1>,_0>>;
// (T32,V2) -> (M8,N8)
using SM80_8x8_Row = Layout<Shape <Shape < _4,_8>,_2>,
Stride<Stride<_16,_1>,_8>>;
// (T32,V4) -> (M8,N16)
using SM80_8x16_Row = Layout<Shape <Shape < _4,_8>,_4>,
Stride<Stride<_32,_1>,_8>>;
// (T32,V4) -> (M16,N8)
using SM80_16x8_Row = Layout<Shape <Shape < _4,_8>,Shape < _2,_2>>,
Stride<Stride<_32,_1>,Stride<_16,_8>>>;
}
///////////////////////////////////////////////////////////////////////////////
//////////////////////// fp16 = fp16 * fp16 + fp16 ////////////////////////////
///////////////////////////////////////////////////////////////////////////////
template <>
struct MMA_Traits<SM80_16x8x8_F16F16F16F16_TN>
{
using ElementDVal = half_t;
using ElementAVal = half_t;
using ElementBVal = half_t;
using ElementCVal = half_t;
using Shape_MNK = Shape<_16,_8,_8>;
using ThrID = Layout<_32>;
using ALayout = SM80_16x8_Row;
using BLayout = SM80_8x8_Row;
using CLayout = SM80_16x8_Row;
};
template <>
struct MMA_Traits<SM80_16x8x16_F16F16F16F16_TN>
{
using ElementDVal = half_t;
using ElementAVal = half_t;
using ElementBVal = half_t;
using ElementCVal = half_t;
using Shape_MNK = Shape<_16,_8,_16>;
using ThrID = Layout<_32>;
using ALayout = Layout<Shape <Shape < _4,_8>,Shape < _2,_2, _2>>,
Stride<Stride<_32,_1>,Stride<_16,_8,_128>>>;
using BLayout = Layout<Shape <Shape < _4,_8>,Shape <_2, _2>>,
Stride<Stride<_16,_1>,Stride<_8,_64>>>;
using CLayout = SM80_16x8_Row;
};
///////////////////////////////////////////////////////////////////////////////
//////////////////////// fp32 = fp16 * fp16 + fp32 ////////////////////////////
///////////////////////////////////////////////////////////////////////////////
template <>
struct MMA_Traits<SM80_16x8x8_F32F16F16F32_TN>
: MMA_Traits<SM80_16x8x8_F16F16F16F16_TN>
{
using ElementDVal = float;
using ElementAVal = half_t;
using ElementBVal = half_t;
using ElementCVal = float;
};
template <>
struct MMA_Traits<SM80_16x8x16_F32F16F16F32_TN>
: MMA_Traits<SM80_16x8x16_F16F16F16F16_TN>
{
using ElementDVal = float;
using ElementAVal = half_t;
using ElementBVal = half_t;
using ElementCVal = float;
};
template <>
struct MMA_Traits<MACA_16x16x16_F32F16F16F32>
{
using ElementDVal = float;
using ElementAVal = half_t;
using ElementBVal = half_t;
using ElementCVal = float;
using Shape_MNK = Shape<_16, _16, _16>;
using ThrID = Layout<_64>;
using ALayout = Layout<Shape<Shape<_16, _4>, _4>,
Stride<Stride<_1, _64>, _16>>;
using BLayout = Layout<Shape<Shape<_16, _4>, _4>,
Stride<Stride<_1, _64>, _16>>;
using CLayout = Layout<Shape<Shape<_16, _4>, _4>,
Stride<Stride<_1, _64>, _16>>;
};
template <>
struct MMA_Traits<MACA_16x16x32_F32F16F16F32>
{
using ElementDVal = float;
using ElementAVal = half_t;
using ElementBVal = half_t;
using ElementCVal = float;
using Shape_MNK = Shape<_16, _16, _32>;
using ThrID = Layout<_64>;
using ALayout = Layout<Shape<Shape<_16, _4>, _8>,
Stride<Stride<_1, _128>, _16>>;
using BLayout = Layout<Shape<Shape<_16, _4>, _8>,
Stride<Stride<_1, _128>, _16>>;
using CLayout = Layout<Shape<Shape<_16, _4>, _4>,
Stride<Stride<_1, _64>, _16>>;
};
// use for lds4x4 + perm4x4
template <>
struct MMA_Traits<MACA_16x64x16_F32F16F16F32>
{
using ElementDVal = float;
using ElementAVal = half_t;
using ElementBVal = half_t;
using ElementCVal = float;
using Shape_MNK = Shape<_16, _64, _16>;
using ThrID = Layout<_64>;
using ALayout = Layout<Shape<Shape<_16, _4>, _4>,
Stride<Stride<_1, _64>, _16>>;
using BLayout = Layout<Shape<Shape<_16, _4>, Shape<_4, _4>>,
Stride<Stride<_4, _256>, Stride<_1, _64>>>;
using CLayout = Layout<Shape<Shape<_16, _4>, _16>,
Stride<Stride<_1, _256>, _16>>;
};
///////////////////////////////////////////////////////////////////////////////
//////////////////////// int32 = int8 * int8 + int32 //////////////////////////
///////////////////////////////////////////////////////////////////////////////
template <>
struct MMA_Traits<MACA_16x16x16_I32I8I8I32>
{
using ElementDVal = int32_t;
using ElementAVal = int8_t;
using ElementBVal = int8_t;
using ElementCVal = int32_t;
using Shape_MNK = Shape<_16, _16, _16>;
using ThrID = Layout<_64>;
using ALayout = Layout<Shape<Shape<_16, _4>, _4>,
Stride<Stride<_1, _64>, _16>>;
using BLayout = Layout<Shape<Shape<_16, _4>, _4>,
Stride<Stride<_1, _64>, _16>>;
using CLayout = Layout<Shape<Shape<_16, _4>, _4>,
Stride<Stride<_1, _64>, _16>>;
};
template <>
struct MMA_Traits<MACA_16x16x32_I32I8I8I32>
{
using ElementDVal = int32_t;
using ElementAVal = int8_t;
using ElementBVal = int8_t;
using ElementCVal = int32_t;
using Shape_MNK = Shape<_16, _16, _32>;
using ThrID = Layout<_64>;
using ALayout = Layout<Shape<Shape<_16, _4>, Shape<_8>>,
Stride<Stride<_1, _128>, Stride<_16>>>;
using BLayout = Layout<Shape<Shape<_16, _4>, Shape<_8>>,
Stride<Stride<_1, _128>, Stride<_16>>>;
using CLayout = Layout<Shape<Shape<_16, _4>, _4>,
Stride<Stride<_1, _64>, _16>>;
};
///////////////////////////////////////////////////////////////////////////////
//////////////////////// fp32 = bf16 * bf16 + fp32 ////////////////////////////
///////////////////////////////////////////////////////////////////////////////
template <>
struct MMA_Traits<SM80_16x8x8_F32BF16BF16F32_TN>
: MMA_Traits<SM80_16x8x8_F16F16F16F16_TN>
{
using ElementDVal = float;
using ElementAVal = bfloat16_t;
using ElementBVal = bfloat16_t;
using ElementCVal = float;
};
template <>
struct MMA_Traits<SM80_16x8x16_F32BF16BF16F32_TN>
: MMA_Traits<SM80_16x8x16_F16F16F16F16_TN>
{
using ElementDVal = float;
using ElementAVal = bfloat16_t;
using ElementBVal = bfloat16_t;
using ElementCVal = float;
};
template <>
struct MMA_Traits<MACA_16x16x16_F32BF16BF16F32>
{
using ElementDVal = float;
using ElementAVal = bfloat16_t;
using ElementBVal = bfloat16_t;
using ElementCVal = float;
using Shape_MNK = Shape<_16, _16, _16>;
using ThrID = Layout<_64>;
using ALayout = Layout<Shape<Shape<_16, _4>, _4>,
Stride<Stride<_1, _64>, _16>>;
using BLayout = Layout<Shape<Shape<_16, _4>, _4>,
Stride<Stride<_1, _64>, _16>>;
using CLayout = Layout<Shape<Shape<_16, _4>, _4>,
Stride<Stride<_1, _64>, _16>>;
};
template <>
struct MMA_Traits<MACA_16x16x32_F32BF16BF16F32>
{
using ElementDVal = float;
using ElementAVal = bfloat16_t;
using ElementBVal = bfloat16_t;
using ElementCVal = float;
using Shape_MNK = Shape<_16, _16, _32>;
using ThrID = Layout<_64>;
using ALayout = Layout<Shape<Shape<_16, _4>, _8>,
Stride<Stride<_1, _128>, _16>>;
using BLayout = Layout<Shape<Shape<_16, _4>, _8>,
Stride<Stride<_1, _128>, _16>>;
using CLayout = Layout<Shape<Shape<_16, _4>, _4>,
Stride<Stride<_1, _64>, _16>>;
};
// use for lds4x4 + perm4x4
template <>
struct MMA_Traits<MACA_16x64x16_F32BF16BF16F32>
{
using ElementDVal = float;
using ElementAVal = bfloat16_t;
using ElementBVal = bfloat16_t;
using ElementCVal = float;
using Shape_MNK = Shape<_16, _64, _16>;
using ThrID = Layout<_64>;
using ALayout = Layout<Shape<Shape<_16, _4>, _4>,
Stride<Stride<_1, _64>, _16>>;
using BLayout = Layout<Shape<Shape<_16, _4>, Shape<_4, _4>>,
Stride<Stride<_4, _256>, Stride<_1, _64>>>;
using CLayout = Layout<Shape<Shape<_16, _4>, _16>,
Stride<Stride<_1, _256>, _16>>;
};
///////////////////////////////////////////////////////////////////////////////
//////////////////////// fp32 = tf32 * tf32 + fp32 ////////////////////////////
///////////////////////////////////////////////////////////////////////////////
template <>
struct MMA_Traits<SM80_16x8x4_F32TF32TF32F32_TN>
{
using ElementDVal = float;
using ElementAVal = mctlass::tfloat32_t;
using ElementBVal = mctlass::tfloat32_t;
using ElementCVal = float;
using Shape_MNK = Shape<_16,_8,_4>;
using ThrID = Layout<_32>;
using ALayout = Layout<Shape <Shape < _4,_8>,_2>,
Stride<Stride<_16,_1>,_8>>;
using BLayout = SM80_8x4;
using CLayout = SM80_16x8_Row;
};
template <>
struct MMA_Traits<SM80_16x8x8_F32TF32TF32F32_TN>
{
using ElementDVal = float;
using ElementAVal = mctlass::tfloat32_t;
using ElementBVal = mctlass::tfloat32_t;
using ElementCVal = float;
using Shape_MNK = Shape<_16,_8,_8>;
using ThrID = Layout<_32>;
using ALayout = Layout<Shape <Shape < _4,_8>,Shape <_2, _2>>,
Stride<Stride<_16,_1>,Stride<_8,_64>>>;
using BLayout = Layout<Shape <Shape <_4,_8>, _2>,
Stride<Stride<_8,_1>,_32>>;
using CLayout = SM80_16x8_Row;
};
///////////////////////////////////////////////////////////////////////////////
//////////////////////// fp64 = fp64 * fp64 + fp64 ////////////////////////////
///////////////////////////////////////////////////////////////////////////////
template <>
struct MMA_Traits<SM80_8x8x4_F64F64F64F64_TN>
{
using ElementDVal = double;
using ElementAVal = double;
using ElementBVal = double;
using ElementCVal = double;
using Shape_MNK = Shape<_8,_8,_4>;
using ThrID = Layout<_32>;
using ALayout = SM80_8x4;
using BLayout = SM80_8x4;
using CLayout = SM80_8x8_Row;
};
// Custom complex fp64 MMA composed of 4 fp64 MMAs -- same layouts
template <>
struct MMA_Traits<SM80_8x8x4_C64C64C64C64_TN>
: MMA_Traits<SM80_8x8x4_F64F64F64F64_TN>
{
using ElementDVal = complex<double>;
using ElementAVal = complex<double>;
using ElementBVal = complex<double>;
using ElementCVal = complex<double>;
};
// Custom complex fp64 MMA composed of 3 fp64 MMAs -- same layouts
template <>
struct MMA_Traits<SM80_8x8x4_GC64C64C64GC64_TN>
: MMA_Traits<SM80_8x8x4_F64F64F64F64_TN>
{
using ElementDVal = typename SM80_8x8x4_GC64C64C64GC64_TN::GaussComplex;
using ElementAVal = complex<double>;
using ElementBVal = complex<double>;
using ElementCVal = typename SM80_8x8x4_GC64C64C64GC64_TN::GaussComplex;
};
///////////////////////////////////////////////////////////////////////////////
/////////////////////////// s32 = s8 * s8 + s32 ///////////////////////////////
///////////////////////////////////////////////////////////////////////////////
template <>
struct MMA_Traits<SM80_8x8x16_S32S8S8S32_TN>
{
using ElementDVal = int32_t;
using ElementAVal = int8_t;
using ElementBVal = int8_t;
using ElementCVal = int32_t;
using Shape_MNK = Shape<_8,_8,_16>;
using ThrID = Layout<_32>;
using ALayout = SM80_8x16_Row;
using BLayout = SM80_8x16_Row;
using CLayout = SM80_8x8_Row;
};
template <>
struct MMA_Traits<SM80_8x8x16_S32S8S8S32_TN_SATURATE>
: MMA_Traits<SM80_8x8x16_S32S8S8S32_TN> {};
template <>
struct MMA_Traits<SM80_16x8x16_S32S8S8S32_TN>
{
using ElementDVal = int32_t;
using ElementAVal = int8_t;
using ElementBVal = int8_t;
using ElementCVal = int32_t;
using Shape_MNK = Shape<_16,_8,_16>;
using ThrID = Layout<_32>;
using ALayout = Layout<Shape <Shape < _4,_8>,Shape < _4,_2>>,
Stride<Stride<_64,_1>,Stride<_16,_8>>>;
using BLayout = SM80_8x16_Row;
using CLayout = SM80_16x8_Row;
};
template <>
struct MMA_Traits<SM80_16x8x16_S32S8S8S32_TN_SATURATE>
: MMA_Traits<SM80_16x8x16_S32S8S8S32_TN> {};
template <>
struct MMA_Traits<SM80_16x8x32_S32S8S8S32_TN>
{
using ElementDVal = int32_t;
using ElementAVal = int8_t;
using ElementBVal = int8_t;
using ElementCVal = int32_t;
using Shape_MNK = Shape<_16,_8,_32>;
using ThrID = Layout<_32>;
using ALayout = Layout<Shape <Shape < _4,_8>,Shape < _4,_2, _2>>,
Stride<Stride<_64,_1>,Stride<_16,_8,_256>>>;
using BLayout = Layout<Shape <Shape < _4,_8>, Shape <_4, _2>>,
Stride<Stride<_32,_1>, Stride<_8,_128>>>;
using CLayout = SM80_16x8_Row;
};
template <>
struct MMA_Traits<SM80_16x8x32_S32S8S8S32_TN_SATURATE>
: MMA_Traits<SM80_16x8x32_S32S8S8S32_TN> {};
///////////////////////////////////////////////////////////////////////////////
/////////////////////////// s32 = s8 * u8 + s32 ///////////////////////////////
///////////////////////////////////////////////////////////////////////////////
template <>
struct MMA_Traits<SM80_8x8x16_S32S8U8S32_TN>
: MMA_Traits<SM80_8x8x16_S32S8S8S32_TN>
{
using ElementDVal = int32_t;
using ElementAVal = int8_t;
using ElementBVal = uint8_t;
using ElementCVal = int32_t;
};
template <>
struct MMA_Traits<SM80_8x8x16_S32S8U8S32_TN_SATURATE>
: MMA_Traits<SM80_8x8x16_S32S8U8S32_TN> {};
template <>
struct MMA_Traits<SM80_16x8x16_S32S8U8S32_TN>
: MMA_Traits<SM80_16x8x16_S32S8S8S32_TN>
{
using ElementDVal = int32_t;
using ElementAVal = int8_t;
using ElementBVal = uint8_t;
using ElementCVal = int32_t;
};
template <>
struct MMA_Traits<SM80_16x8x16_S32S8U8S32_TN_SATURATE>
: MMA_Traits<SM80_16x8x16_S32S8U8S32_TN> {};
template <>
struct MMA_Traits<SM80_16x8x32_S32S8U8S32_TN>
: MMA_Traits<SM80_16x8x32_S32S8S8S32_TN>
{
using ElementDVal = int32_t;
using ElementAVal = int8_t;
using ElementBVal = uint8_t;
using ElementCVal = int32_t;
};
template <>
struct MMA_Traits<SM80_16x8x32_S32S8U8S32_TN_SATURATE>
: MMA_Traits<SM80_16x8x32_S32S8U8S32_TN> {};
///////////////////////////////////////////////////////////////////////////////
/////////////////////////// s32 = u8 * s8 + s32 ///////////////////////////////
///////////////////////////////////////////////////////////////////////////////
template <>
struct MMA_Traits<SM80_8x8x16_S32U8S8S32_TN>
: MMA_Traits<SM80_8x8x16_S32S8S8S32_TN>
{
using ElementDVal = int32_t;
using ElementAVal = uint8_t;
using ElementBVal = int8_t;
using ElementCVal = int32_t;
};
template <>
struct MMA_Traits<SM80_8x8x16_S32U8S8S32_TN_SATURATE>
: MMA_Traits<SM80_8x8x16_S32U8S8S32_TN> {};
template <>
struct MMA_Traits<SM80_16x8x16_S32U8S8S32_TN>
: MMA_Traits<SM80_16x8x16_S32S8S8S32_TN>
{
using ElementDVal = int32_t;
using ElementAVal = uint8_t;
using ElementBVal = int8_t;
using ElementCVal = int32_t;
};
template <>
struct MMA_Traits<SM80_16x8x16_S32U8S8S32_TN_SATURATE>
: MMA_Traits<SM80_16x8x16_S32U8S8S32_TN> {};
template <>
struct MMA_Traits<SM80_16x8x32_S32U8S8S32_TN>
: MMA_Traits<SM80_16x8x32_S32S8S8S32_TN>
{
using ElementDVal = int32_t;
using ElementAVal = uint8_t;
using ElementBVal = int8_t;
using ElementCVal = int32_t;
};
template <>
struct MMA_Traits<SM80_16x8x32_S32U8S8S32_TN_SATURATE>
: MMA_Traits<SM80_16x8x32_S32U8S8S32_TN> {};
///////////////////////////////////////////////////////////////////////////////
/////////////////////////// s32 = u8 * u8 + s32 ///////////////////////////////
///////////////////////////////////////////////////////////////////////////////
template <>
struct MMA_Traits<SM80_8x8x16_S32U8U8S32_TN>
: MMA_Traits<SM80_8x8x16_S32S8S8S32_TN>
{
using ElementDVal = int32_t;
using ElementAVal = uint8_t;
using ElementBVal = uint8_t;
using ElementCVal = int32_t;
};
template <>
struct MMA_Traits<SM80_8x8x16_S32U8U8S32_TN_SATURATE>
: MMA_Traits<SM80_8x8x16_S32U8U8S32_TN> {};
template <>
struct MMA_Traits<SM80_16x8x16_S32U8U8S32_TN>
: MMA_Traits<SM80_16x8x16_S32S8S8S32_TN>
{
using ElementDVal = int32_t;
using ElementAVal = uint8_t;
using ElementBVal = uint8_t;
using ElementCVal = int32_t;
};
template <>
struct MMA_Traits<SM80_16x8x16_S32U8U8S32_TN_SATURATE>
: MMA_Traits<SM80_16x8x16_S32U8U8S32_TN> {};
template <>
struct MMA_Traits<SM80_16x8x32_S32U8U8S32_TN>
: MMA_Traits<SM80_16x8x32_S32S8S8S32_TN>
{
using ElementDVal = int32_t;
using ElementAVal = uint8_t;
using ElementBVal = uint8_t;
using ElementCVal = int32_t;
};
template <>
struct MMA_Traits<SM80_16x8x32_S32U8U8S32_TN_SATURATE>
: MMA_Traits<SM80_16x8x32_S32U8U8S32_TN> {};
///////////////////////////////////////////////////////////////////////////////
/////////////////////////// s32 = b1 ^ b1 + s32 ///////////////////////////////
///////////////////////////////////////////////////////////////////////////////
template <>
struct MMA_Traits<SM80_16x8x256_S32U1U1S32_TN_XORPOPC>
{
using ElementDVal = int32_t;
using ElementAVal = cute::uint1b_t;
using ElementBVal = cute::uint1b_t;
using ElementCVal = int32_t;
using Shape_MNK = Shape<_16,_8,_256>;
using ThrID = Layout<_32>;
using ALayout = Layout<Shape <_32,Shape < _8, _4,_2, _2>>,
Stride<_64,Stride<_64,_16,_8,_2048>>>;
using BLayout = Layout<Shape <_32,Shape <_32, _2>>,
Stride<_32,Stride< _1,_1024>>>;
using CLayout = SM80_16x8_Row;
};
} // end namespace cute

View File

@ -0,0 +1,132 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#pragma once
#include <cute/arch/mma_sm90.hpp>
#include <cute/atom/mma_traits.hpp>
#include <cute/layout.hpp>
namespace cute {
///////////////////////////////////////////////////////////////////////////////
//////////////////////// fp64 = fp64 * fp64 + fp64 ////////////////////////////
///////////////////////////////////////////////////////////////////////////////
template <>
struct MMA_Traits<SM90_16x8x4_F64F64F64F64_TN>
{
using ElementDVal = double;
using ElementAVal = double;
using ElementBVal = double;
using ElementCVal = double;
using Shape_MNK = Shape<_16,_8,_4>;
using ThrID = Layout<_32>;
using ALayout = Layout<Shape <Shape < _4,_8>,_2>,
Stride<Stride<_16,_1>,_8>>;
using BLayout = Layout<Shape <Shape < _4,_8>,_1>,
Stride<Stride< _8,_1>,_0>>;
using CLayout = Layout<Shape <Shape < _4,_8>,Shape < _2,_2>>,
Stride<Stride<_32,_1>,Stride<_16,_8>>>;
};
template <>
struct MMA_Traits<SM90_16x8x8_F64F64F64F64_TN>
{
using ElementDVal = double;
using ElementAVal = double;
using ElementBVal = double;
using ElementCVal = double;
using Shape_MNK = Shape<_16,_8,_8>;
using ThrID = Layout<_32>;
using ALayout = Layout<Shape <Shape < _4,_8>,Shape <_2, _2>>,
Stride<Stride<_16,_1>,Stride<_8,_64>>>;
using BLayout = Layout<Shape <Shape < _4,_8>, _2>,
Stride<Stride< _8,_1>,_32>>;
using CLayout = Layout<Shape <Shape < _4,_8>,Shape < _2,_2>>,
Stride<Stride<_32,_1>,Stride<_16,_8>>>;
};
template <>
struct MMA_Traits<SM90_16x8x16_F64F64F64F64_TN>
{
using ElementDVal = double;
using ElementAVal = double;
using ElementBVal = double;
using ElementCVal = double;
using Shape_MNK = Shape<_16,_8,_16>;
using ThrID = Layout<_32>;
using ALayout = Layout<Shape <Shape < _4,_8>,Shape <_2, _4>>,
Stride<Stride<_16,_1>,Stride<_8,_64>>>;
using BLayout = Layout<Shape <Shape < _4,_8>, _4>,
Stride<Stride< _8,_1>,_32>>;
using CLayout = Layout<Shape <Shape < _4,_8>,Shape < _2,_2>>,
Stride<Stride<_32,_1>,Stride<_16,_8>>>;
};
///////////////////////////////////////////////////////////////////////////////////
//////////////////////// cfp64 = cfp64 * cfp64 + cfp64 ////////////////////////////
///////////////////////////////////////////////////////////////////////////////////
template <>
struct MMA_Traits<SM90_16x8x4_C64C64C64C64_TN>
: MMA_Traits<SM90_16x8x4_F64F64F64F64_TN>
{
using ElementDVal = complex<double>;
using ElementAVal = complex<double>;
using ElementBVal = complex<double>;
using ElementCVal = complex<double>;
};
template <>
struct MMA_Traits<SM90_16x8x8_C64C64C64C64_TN>
: MMA_Traits<SM90_16x8x8_F64F64F64F64_TN>
{
using ElementDVal = complex<double>;
using ElementAVal = complex<double>;
using ElementBVal = complex<double>;
using ElementCVal = complex<double>;
};
template <>
struct MMA_Traits<SM90_16x8x16_C64C64C64C64_TN>
: MMA_Traits<SM90_16x8x16_F64F64F64F64_TN>
{
using ElementDVal = complex<double>;
using ElementAVal = complex<double>;
using ElementBVal = complex<double>;
using ElementCVal = complex<double>;
};
} // end namespace cute

File diff suppressed because it is too large Load Diff

View File

@ -0,0 +1,169 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#pragma once
#if defined(__MACA_ARCH__) || defined(__clang__)
# define CUTE_HOST_DEVICE __forceinline__ __host__ __device__
# define CUTE_DEVICE __forceinline__ __device__
# define CUTE_HOST __forceinline__ __host__
#else
# define CUTE_HOST_DEVICE inline
# define CUTE_DEVICE inline
# define CUTE_HOST inline
#endif // CUTE_HOST_DEVICE, CUTE_DEVICE
#if !defined(__MACACC_RTC__) && (defined(__MACA_ARCH__))
# define CUTE_UNROLL _Pragma("unroll")
# define CUTE_NO_UNROLL _Pragma("unroll 1")
#elif defined(__MACACC_RTC__)
# define CUTE_UNROLL _Pragma("unroll")
# define CUTE_NO_UNROLL _Pragma("unroll 1")
#else
# define CUTE_UNROLL
# define CUTE_NO_UNROLL
#endif // CUTE_UNROLL
#if defined(__MACA_ARCH__)
# define CUTE_INLINE_CONSTANT static const __device__
#else
# define CUTE_INLINE_CONSTANT static constexpr
#endif
#if defined(__MACA_ARCH__) && (__MACA_ARCH__ == 1000)
# define CUTE_MACA_XCORE1000_ENABLED
#endif
#if defined(__MACA_ARCH__) && (__MACA_ARCH__ == 1500)
# define CUTE_MACA_XCORE1500_ENABLED
#endif
// __grid_constant__ was introduced in CUDA 11.7.
// #if ((__CUDACC_VER_MAJOR__ >= 12) || ((__CUDACC_VER_MAJOR__ == 11) && (__CUDACC_VER_MINOR__ >= 7)))
#if 0
# define CUTE_GRID_CONSTANT_SUPPORTED
#endif
// __grid_constant__ can be enabled only on SM70+.
#if defined(__MACA_ARCH__)
# define CUTE_GRID_CONSTANT_ENABLED
#endif
#if ! defined(CUTE_GRID_CONSTANT)
# if defined(CUTE_GRID_CONSTANT_SUPPORTED) && defined(CUTE_GRID_CONSTANT_ENABLED)
# define CUTE_GRID_CONSTANT __grid_constant__
# else
# define CUTE_GRID_CONSTANT
# endif
#endif
// Some versions of GCC < 11 have trouble deducing that a
// function with "auto" return type and all of its returns in an "if
// constexpr ... else" statement must actually return. Thus, GCC
// emits spurious "missing return statement" build warnings.
// Developers can suppress these warnings by using the
// CUTE_GCC_UNREACHABLE macro, which must be followed by a semicolon.
// It's harmless to use the macro for other GCC versions or other
// compilers, but it has no effect.
#if ! defined(CUTE_GCC_UNREACHABLE)
# if defined(__GNUC__) && __GNUC__ < 11
// GCC 10, but not 7.5, 9.4.0, or 11, issues "missing return
// statement" warnings without this little bit of help.
# define CUTE_GCC_UNREACHABLE __builtin_unreachable()
# else
# define CUTE_GCC_UNREACHABLE
# endif
#endif
#ifdef _MSC_VER
// Provides support for alternative operators 'and', 'or', and 'not'
#include <iso646.h>
#endif // _MSC_VER
#if defined(__MACACC_RTC__)
#define CUTE_STL_NAMESPACE cuda::std
#define CUTE_STL_NAMESPACE_IS_CUDA_STD
#else
#define CUTE_STL_NAMESPACE std
#endif
//
// Assertion helpers
//
#if defined(__MACACC_RTC__)
#include <cuda/std/cassert>
#else
#include <cassert>
#endif
#define CUTE_STATIC_ASSERT static_assert
#define CUTE_STATIC_ASSERT_V(x,...) static_assert(decltype(x)::value, ##__VA_ARGS__)
#if defined(__MACA_ARCH__)
# define CUTE_RUNTIME_ASSERT(x) assert(0 && x);__brkpt()
#else
# define CUTE_RUNTIME_ASSERT(x) assert(0 && x)
#endif
//
// IO
//
#if !defined(__MACACC_RTC__)
#include <cstdio>
#include <iostream>
#include <iomanip>
#endif
//
// Support
//
#include <cute/util/type_traits.hpp>
//
// Basic types
//
#include <cute/numeric/int.hpp>
#include <cute/numeric/real.hpp>
#include <cute/numeric/half.hpp>
#include <cute/numeric/float8.hpp>
#include <cute/numeric/bfloat.hpp>
#include <cute/numeric/tfloat.hpp>
#include <cute/numeric/complex.hpp>
//
// Debugging utilities
//
#include <cute/util/print.hpp>
#include <cute/util/debug.hpp>

View File

@ -0,0 +1,70 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#pragma once
#include <cute/config.hpp>
#include <cute/numeric/int.hpp>
#include <cute/numeric/math.hpp>
namespace cute
{
// Test if a pointer is aligned to N bytes
template <int N>
CUTE_HOST_DEVICE constexpr
bool
is_byte_aligned(void const* const ptr)
{
static_assert(N > 0 && (N & (N - 1)) == 0, "N must be a power of 2 in alignment check");
return (reinterpret_cast<uintptr_t>(ptr) & (N-1)) == 0;
}
#if defined(__MACACC__)
# define CUTE_ALIGNAS(n) __align__(n)
#else
# define CUTE_ALIGNAS(n) alignas(n)
#endif
template <size_t Alignment>
struct aligned_struct {};
template <> struct CUTE_ALIGNAS( 1) aligned_struct< 1> {};
template <> struct CUTE_ALIGNAS( 2) aligned_struct< 2> {};
template <> struct CUTE_ALIGNAS( 4) aligned_struct< 4> {};
template <> struct CUTE_ALIGNAS( 8) aligned_struct< 8> {};
template <> struct CUTE_ALIGNAS( 16) aligned_struct< 16> {};
template <> struct CUTE_ALIGNAS( 32) aligned_struct< 32> {};
template <> struct CUTE_ALIGNAS( 64) aligned_struct< 64> {};
template <> struct CUTE_ALIGNAS(128) aligned_struct<128> {};
template <> struct CUTE_ALIGNAS(256) aligned_struct<256> {};
} // end namespace cute

View File

@ -0,0 +1,334 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#pragma once
#include <cute/config.hpp>
#include <cute/numeric/integral_constant.hpp>
#include <cute/util/type_traits.hpp>
namespace cute
{
template <class T, size_t N>
struct array
{
using value_type = T;
using size_type = size_t;
using difference_type = ptrdiff_t;
using reference = value_type&;
using const_reference = const value_type&;
using pointer = value_type*;
using const_pointer = const value_type*;
using iterator = pointer;
using const_iterator = const_pointer;
CUTE_HOST_DEVICE constexpr
reference operator[](size_type pos)
{
return begin()[pos];
}
CUTE_HOST_DEVICE constexpr
const_reference operator[](size_type pos) const
{
return begin()[pos];
}
CUTE_HOST_DEVICE constexpr
reference front()
{
return *begin();
}
CUTE_HOST_DEVICE constexpr
const_reference front() const
{
return *begin();
}
CUTE_HOST_DEVICE constexpr
reference back()
{
// return *rbegin();
return operator[](N-1);
}
CUTE_HOST_DEVICE constexpr
const_reference back() const
{
// return *rbegin();
return operator[](N-1);
}
CUTE_HOST_DEVICE constexpr
T* data()
{
return __elems_;
}
CUTE_HOST_DEVICE constexpr
T const* data() const
{
return __elems_;
}
CUTE_HOST_DEVICE constexpr
iterator begin()
{
return data();
}
CUTE_HOST_DEVICE constexpr
const_iterator begin() const
{
return data();
}
CUTE_HOST_DEVICE constexpr
const_iterator cbegin()
{
return begin();
}
CUTE_HOST_DEVICE constexpr
const_iterator cbegin() const
{
return begin();
}
CUTE_HOST_DEVICE constexpr
iterator end()
{
return data() + size();
}
CUTE_HOST_DEVICE constexpr
const_iterator end() const
{
return data() + size();
}
CUTE_HOST_DEVICE constexpr
const_iterator cend()
{
return end();
}
CUTE_HOST_DEVICE constexpr
const_iterator cend() const
{
return end();
}
CUTE_HOST_DEVICE constexpr
bool empty() const
{
return size() == 0;
}
CUTE_HOST_DEVICE constexpr
size_type size() const
{
return N;
}
CUTE_HOST_DEVICE constexpr
size_type max_size() const
{
return size();
}
CUTE_HOST_DEVICE constexpr
void fill(const T& value)
{
for (auto& e : *this) {
e = value;
}
}
CUTE_HOST_DEVICE constexpr
void clear()
{
fill(T(0));
}
CUTE_HOST_DEVICE constexpr
void swap(array& other)
{
using CUTE_STL_NAMESPACE::swap;
for (size_type i = 0; i < size(); ++i) {
swap((*this)[i], other[i]);
}
}
value_type __elems_[N > 0 ? N : 1];
};
template <class T, size_t N>
CUTE_HOST_DEVICE constexpr
bool operator==(array<T,N> const& lhs, array<T,N> const& rhs)
{
for (size_t i = 0; i < N; ++i) {
if (lhs[i] != rhs[i]) {
return false;
}
}
return true;
}
template <class T, size_t N>
CUTE_HOST_DEVICE constexpr
void clear(array<T,N>& a)
{
a.fill(T(0));
}
template <typename T, size_t N>
CUTE_HOST_DEVICE constexpr
void fill(array<T,N>& a, T const& value)
{
a.fill(value);
}
template <class T, size_t N>
CUTE_HOST_DEVICE constexpr
void swap(array<T,N>& a, array<T,N>& b)
{
a.swap(b);
}
} // end cute
//
// Specialize tuple-related functionality for cute::array
//
#if defined(__MACACC_RTC__)
#include <cuda/std/tuple>
#else
#include <tuple>
#endif
namespace cute
{
template <size_t I, class T, size_t N>
CUTE_HOST_DEVICE constexpr
T& get(array<T,N>& a)
{
static_assert(I < N, "Index out of range");
return a[I];
}
template <size_t I, class T, size_t N>
CUTE_HOST_DEVICE constexpr
T const& get(array<T,N> const& a)
{
static_assert(I < N, "Index out of range");
return a[I];
}
template <size_t I, class T, size_t N>
CUTE_HOST_DEVICE constexpr
T&& get(array<T,N>&& a)
{
static_assert(I < N, "Index out of range");
return std::move(a[I]);
}
} // end namespace cute
namespace CUTE_STL_NAMESPACE
{
template <class T, size_t N>
struct tuple_size<cute::array<T,N>>
: cute::integral_constant<size_t, N>
{};
template <size_t I, class T, size_t N>
struct tuple_element<I, cute::array<T,N>>
{
using type = T;
};
template <class T, size_t N>
struct tuple_size<const cute::array<T,N>>
: cute::integral_constant<size_t, N>
{};
template <size_t I, class T, size_t N>
struct tuple_element<I, const cute::array<T,N>>
{
using type = T;
};
} // end namespace CUTE_STL_NAMESPACE
#ifdef CUTE_STL_NAMESPACE_IS_CUDA_STD
namespace std
{
#if defined(__MACACC_RTC__)
template <class... _Tp>
struct tuple_size;
template<size_t _Ip, class... _Tp>
struct tuple_element;
#endif
template <class T, size_t N>
struct tuple_size<cute::array<T,N>>
: cute::integral_constant<size_t, N>
{};
template <size_t I, class T, size_t N>
struct tuple_element<I, cute::array<T,N>>
{
using type = T;
};
template <class T, size_t N>
struct tuple_size<const cute::array<T,N>>
: cute::integral_constant<size_t, N>
{};
template <size_t I, class T, size_t N>
struct tuple_element<I, const cute::array<T,N>>
{
using type = T;
};
} // end namepsace std
#endif // CUTE_STL_NAMESPACE_IS_CUDA_STD

View File

@ -0,0 +1,42 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#pragma once
#include <cute/container/array.hpp>
#include <cute/container/alignment.hpp>
namespace cute
{
template <class T, size_t N, size_t Alignment = 16>
struct CUTE_ALIGNAS(Alignment) array_aligned : cute::array<T,N> {};
} // end namespace cute

View File

@ -0,0 +1,633 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
/*! \file
\brief Statically sized array of elements that accommodates subbyte trivial types
in a packed storage.
*/
#pragma once
#include <cute/config.hpp>
#include <cute/numeric/int.hpp> // sizeof_bits
#include <cute/numeric/integral_constant.hpp>
namespace cute
{
////////////////////////////////////////////////////////////////////////////////////////////////////
/// Statically sized array for any data type
template <class T, size_t N>
class array_subbyte
{
public:
/// Number of total bits in the array
static constexpr int kSizeBits = sizeof_bits<T>::value * N;
/// Storage type
using Storage = conditional_t<(kSizeBits % 32) == 0, uint32_t,
conditional_t<(kSizeBits % 16) == 0, uint16_t,
uint8_t>>;
/// Number of logical elements per stored object
static constexpr int kElementsPerStoredItem = sizeof_bits<Storage>::value / sizeof_bits<T>::value;
/// Number of storage elements
static constexpr size_t kStorageElements = (N + kElementsPerStoredItem - 1) / kElementsPerStoredItem;
/// Bitmask for covering one item
static constexpr Storage bit_mask_ = ((Storage(1) << sizeof_bits<T>::value) - 1);
//
// C++ standard members with reference and iterator types omitted
//
using value_type = T;
using pointer = value_type*;
using const_pointer = value_type const*;
using size_type = size_t;
using difference_type = ptrdiff_t;
//
// References
//
/// Reference object inserts or extracts sub-byte items
class reference {
/// Pointer to storage element
Storage* ptr_;
/// Index into elements packed into Storage object
int idx_;
public:
/// Default ctor
CUTE_HOST_DEVICE constexpr
reference() : ptr_(nullptr), idx_(0) {}
/// Ctor
CUTE_HOST_DEVICE constexpr
reference(Storage* ptr, int idx = 0) : ptr_(ptr), idx_(idx) {}
/// Assignment
CUTE_HOST_DEVICE constexpr
reference& operator=(T x) {
Storage item = (x & bit_mask_);
Storage kUpdateMask = Storage(~(bit_mask_ << (idx_ * sizeof_bits<T>::value)));
*ptr_ = Storage((*ptr_ & kUpdateMask) | (item << (idx_ * sizeof_bits<T>::value)));
return *this;
}
CUTE_HOST_DEVICE constexpr
T get() const {
if constexpr (is_same<bool, T>::value) {
// Extract to bool -- potentially faster impl
return bool((*ptr_) & (bit_mask_ << (idx_ * sizeof_bits<T>::value)));
} else {
// Extract to T
Storage item = Storage((*ptr_ >> (idx_ * sizeof_bits<T>::value)) & bit_mask_);
return reinterpret_cast<T const&>(item);
}
}
/// Extract to type T
CUTE_HOST_DEVICE constexpr
operator T() const {
return get();
}
};
/// Reference object extracts sub-byte items
class const_reference {
/// Pointer to storage element
Storage const* ptr_;
/// Index into elements packed into Storage object
int idx_;
public:
/// Default ctor
CUTE_HOST_DEVICE constexpr
const_reference(): ptr_(nullptr), idx_(0) { }
/// Ctor
CUTE_HOST_DEVICE constexpr
const_reference(Storage const* ptr, int idx = 0): ptr_(ptr), idx_(idx) { }
CUTE_HOST_DEVICE constexpr
const T get() const {
if constexpr (is_same<bool, T>::value) {
// Extract to bool -- potentially faster impl
return bool((*ptr_) & (bit_mask_ << (idx_ * sizeof_bits<T>::value)));
} else {
// Extract to T
Storage item = Storage((*ptr_ >> (idx_ * sizeof_bits<T>::value)) & bit_mask_);
return reinterpret_cast<T const&>(item);
}
}
/// Extract to type T
CUTE_HOST_DEVICE constexpr
operator T() const {
return get();
}
};
//
// Iterators
//
/// Bidirectional iterator over elements
class iterator {
/// Pointer to storage element
Storage* ptr_;
/// Index into elements packed into Storage object
int idx_;
public:
CUTE_HOST_DEVICE constexpr
iterator(): ptr_(nullptr), idx_(0) { }
CUTE_HOST_DEVICE constexpr
iterator(Storage* ptr, int idx = 0): ptr_(ptr), idx_(idx) { }
CUTE_HOST_DEVICE constexpr
iterator& operator++() {
++idx_;
if (idx_ == kElementsPerStoredItem) {
++ptr_;
idx_ = 0;
}
return *this;
}
CUTE_HOST_DEVICE constexpr
iterator& operator--() {
if (idx_) {
--idx_;
} else {
--ptr_;
idx_ = kElementsPerStoredItem - 1;
}
return *this;
}
CUTE_HOST_DEVICE constexpr
iterator operator++(int) {
iterator ret(*this);
++(*this);
return ret;
}
CUTE_HOST_DEVICE constexpr
iterator operator--(int) {
iterator ret(*this);
--(*this);
return ret;
}
CUTE_HOST_DEVICE constexpr
iterator& operator+=(int k) {
idx_ += k;
ptr_ += idx_ / kElementsPerStoredItem;
idx_ = idx_ % kElementsPerStoredItem;
return *this;
}
CUTE_HOST_DEVICE constexpr
iterator operator+(int k) const {
return iterator(ptr_,idx_) += k;
}
CUTE_HOST_DEVICE constexpr
reference operator*() const {
return reference(ptr_, idx_);
}
CUTE_HOST_DEVICE constexpr
reference operator[](int k) const {
return *(*this + k);
}
CUTE_HOST_DEVICE constexpr
bool operator==(iterator const& other) const {
return ptr_ == other.ptr_ && idx_ == other.idx_;
}
CUTE_HOST_DEVICE constexpr
bool operator!=(iterator const& other) const {
return !(*this == other);
}
};
/// Bidirectional constant iterator over elements
class const_iterator {
/// Pointer to storage element
Storage const* ptr_;
/// Index into elements packed into Storage object
int idx_;
public:
CUTE_HOST_DEVICE constexpr
const_iterator(): ptr_(nullptr), idx_(0) { }
CUTE_HOST_DEVICE constexpr
const_iterator(Storage const* ptr, int idx = 0): ptr_(ptr), idx_(idx) { }
CUTE_HOST_DEVICE constexpr
const_iterator& operator++() {
++idx_;
if (idx_ == kElementsPerStoredItem) {
++ptr_;
idx_ = 0;
}
return *this;
}
CUTE_HOST_DEVICE constexpr
const_iterator& operator--() {
if (idx_) {
--idx_;
} else {
--ptr_;
idx_ = kElementsPerStoredItem - 1;
}
return *this;
}
CUTE_HOST_DEVICE constexpr
const_iterator operator++(int) {
iterator ret(*this);
++idx_;
if (idx_ == kElementsPerStoredItem) {
++ptr_;
idx_ = 0;
}
return ret;
}
CUTE_HOST_DEVICE constexpr
const_iterator operator--(int) {
iterator ret(*this);
if (idx_) {
--idx_;
} else {
--ptr_;
idx_ = kElementsPerStoredItem - 1;
}
return ret;
}
CUTE_HOST_DEVICE constexpr
const_iterator& operator+=(int k) {
idx_ += k;
ptr_ += idx_ / kElementsPerStoredItem;
idx_ = idx_ % kElementsPerStoredItem;
return *this;
}
CUTE_HOST_DEVICE constexpr
const_iterator operator+(int k) const {
return const_iterator(ptr_,idx_) += k;
}
CUTE_HOST_DEVICE constexpr
const_reference operator*() const {
return const_reference(ptr_, idx_);
}
CUTE_HOST_DEVICE constexpr
const_reference operator[](int k) const {
return *(*this + k);
}
CUTE_HOST_DEVICE constexpr
bool operator==(iterator const& other) const {
return ptr_ == other.ptr_ && idx_ == other.idx_;
}
CUTE_HOST_DEVICE constexpr
bool operator!=(iterator const& other) const {
return !(*this == other);
}
};
private:
/// Internal storage
Storage storage[kStorageElements];
public:
CUTE_HOST_DEVICE constexpr
array_subbyte() { }
CUTE_HOST_DEVICE constexpr
array_subbyte(array_subbyte const& x) {
CUTE_UNROLL
for (unsigned i = 0; i < kStorageElements; ++i) {
storage[i] = x.storage[i];
}
}
CUTE_HOST_DEVICE constexpr
size_type size() const {
return N;
}
CUTE_HOST_DEVICE constexpr
size_type max_size() const {
return N;
}
CUTE_HOST_DEVICE constexpr
bool empty() const {
return !N;
}
/// Efficient clear method
CUTE_HOST_DEVICE constexpr
void clear() {
CUTE_UNROLL
for (unsigned i = 0; i < kStorageElements; ++i) {
storage[i] = Storage(0);
}
}
// Efficient fill method
CUTE_HOST_DEVICE constexpr
void fill(T const& value) {
Storage item = (reinterpret_cast<Storage const&>(value) & bit_mask_);
// Reproduce the value over the bits of the storage item
CUTE_UNROLL
for (unsigned s = sizeof_bits<T>::value; s < sizeof_bits<Storage>::value; s *= 2) {
item |= item << s;
}
CUTE_UNROLL
for (unsigned i = 0; i < kStorageElements; ++i) {
storage[i] = item;
}
}
CUTE_HOST_DEVICE constexpr
reference at(size_type pos) {
return reference(storage + pos / kElementsPerStoredItem, pos % kElementsPerStoredItem);
}
CUTE_HOST_DEVICE constexpr
const_reference at(size_type pos) const {
return const_reference(storage + pos / kElementsPerStoredItem, pos % kElementsPerStoredItem);
}
CUTE_HOST_DEVICE constexpr
reference operator[](size_type pos) {
return at(pos);
}
CUTE_HOST_DEVICE constexpr
const_reference operator[](size_type pos) const {
return at(pos);
}
CUTE_HOST_DEVICE constexpr
reference front() {
return at(0);
}
CUTE_HOST_DEVICE constexpr
const_reference front() const {
return at(0);
}
CUTE_HOST_DEVICE constexpr
reference back() {
return reference(storage + kStorageElements - 1, kElementsPerStoredItem - 1);
}
CUTE_HOST_DEVICE constexpr
const_reference back() const {
return const_reference(storage + kStorageElements - 1, kElementsPerStoredItem - 1);
}
CUTE_HOST_DEVICE constexpr
pointer data() {
return reinterpret_cast<pointer>(storage);
}
CUTE_HOST_DEVICE constexpr
const_pointer data() const {
return reinterpret_cast<const_pointer>(storage);
}
CUTE_HOST_DEVICE constexpr
Storage* raw_data() {
return storage;
}
CUTE_HOST_DEVICE constexpr
Storage const* raw_data() const {
return storage;
}
CUTE_HOST_DEVICE constexpr
iterator begin() {
return iterator(storage);
}
CUTE_HOST_DEVICE constexpr
const_iterator begin() const {
return const_iterator(storage);
}
CUTE_HOST_DEVICE constexpr
const_iterator cbegin() const {
return begin();
}
CUTE_HOST_DEVICE constexpr
iterator end() {
return iterator(storage + N / kElementsPerStoredItem, N % kElementsPerStoredItem);
}
CUTE_HOST_DEVICE constexpr
const_iterator end() const {
return const_iterator(storage + N / kElementsPerStoredItem, N % kElementsPerStoredItem);
}
CUTE_HOST_DEVICE constexpr
const_iterator cend() const {
return end();
}
//
// Comparison operators
//
};
//
// Operators
//
template <class T, size_t N>
CUTE_HOST_DEVICE constexpr
void clear(array_subbyte<T,N>& a)
{
a.clear();
}
template <class T, size_t N>
CUTE_HOST_DEVICE constexpr
void fill(array_subbyte<T,N>& a, T const& value)
{
a.fill(value);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
} // namespace cute
//
// Specialize tuple-related functionality for cute::array_subbyte
//
#if defined(__MACACC_RTC__)
#include <cuda/std/tuple>
#else
#include <tuple>
#endif
namespace cute
{
template <size_t I, class T, size_t N>
CUTE_HOST_DEVICE constexpr
T& get(array_subbyte<T,N>& a)
{
static_assert(I < N, "Index out of range");
return a[I];
}
template <size_t I, class T, size_t N>
CUTE_HOST_DEVICE constexpr
T const& get(array_subbyte<T,N> const& a)
{
static_assert(I < N, "Index out of range");
return a[I];
}
template <size_t I, class T, size_t N>
CUTE_HOST_DEVICE constexpr
T&& get(array_subbyte<T,N>&& a)
{
static_assert(I < N, "Index out of range");
return std::move(a[I]);
}
} // end namespace cute
namespace CUTE_STL_NAMESPACE
{
template <class T, size_t N>
struct tuple_size<cute::array_subbyte<T,N>>
: cute::integral_constant<size_t, N>
{};
template <size_t I, class T, size_t N>
struct tuple_element<I, cute::array_subbyte<T,N>>
{
using type = T;
};
template <class T, size_t N>
struct tuple_size<const cute::array_subbyte<T,N>>
: cute::integral_constant<size_t, N>
{};
template <size_t I, class T, size_t N>
struct tuple_element<I, const cute::array_subbyte<T,N>>
{
using type = T;
};
} // end namespace CUTE_STL_NAMESPACE
#ifdef CUTE_STL_NAMESPACE_IS_CUDA_STD
namespace std
{
#if defined(__MACACC_RTC__)
template <class... _Tp>
struct tuple_size;
template<size_t _Ip, class... _Tp>
struct tuple_element;
#endif
template <class T, size_t N>
struct tuple_size<cute::array_subbyte<T,N>>
: cute::integral_constant<size_t, N>
{};
template <size_t I, class T, size_t N>
struct tuple_element<I, cute::array_subbyte<T,N>>
{
using type = T;
};
template <class T, size_t N>
struct tuple_size<const cute::array_subbyte<T,N>>
: cute::integral_constant<size_t, N>
{};
template <size_t I, class T, size_t N>
struct tuple_element<I, const cute::array_subbyte<T,N>>
{
using type = T;
};
} // end namespace std
#endif // CUTE_STL_NAMESPACE_IS_CUDA_STD

View File

@ -0,0 +1,131 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
/*! \file
\brief Portable bit field that supports byte and word straddling that can
be used in unions to bit-wise define parameters.
*/
#pragma once
#include <cute/config.hpp>
#include <cute/numeric/int.hpp> // uint_bit_t
namespace cute
{
class dummy_type {};
template <uint32_t BitStart, uint32_t NumBits, class OtherValueType = dummy_type>
struct bit_field
{
static_assert(0 < NumBits && NumBits <= 64, "bit_fields with more than 64 bits are not supported.");
// value_type: Use the smallest value type that fits NumBits
static constexpr uint32_t value_type_bits = (NumBits <= 8) ? 8 :
(NumBits <= 16) ? 16 :
(NumBits <= 32) ? 32 : 64;
using value_type = cute::uint_bit_t<value_type_bits>;
// storage_type: Use the smallest storage_type that avoids boundary crossing
static constexpr uint32_t storage_type_bits = (BitStart / 8 == (BitStart + NumBits - 1) / 8) ? 8 :
(BitStart / 16 == (BitStart + NumBits - 1) / 16) ? 16 :
(BitStart / 32 == (BitStart + NumBits - 1) / 32) ? 32 : 64;
using storage_type = cute::uint_bit_t<storage_type_bits>;
static_assert(sizeof(OtherValueType) == sizeof(value_type) || is_same<OtherValueType,dummy_type>::value,
"sizeof(OtherValueType) must be same as sizeof(value_type).");
// Number of storage values needed: ceil_div(BitStart + NumBits, storage_type_bits)
static constexpr uint32_t N = (BitStart + NumBits + storage_type_bits - 1) / storage_type_bits;
// Index of storage value for BitStart
static constexpr uint32_t idx = BitStart / storage_type_bits;
// Bit of data_[idx] for BitStart
static constexpr uint32_t bit_lo = BitStart % storage_type_bits;
// Number of bits in data_[idx] used for NumBits if straddling, else 0
static constexpr uint32_t bit_hi = (idx + 1 < N) ? (storage_type_bits - bit_lo) : 0;
// NumBits mask
static constexpr value_type mask = (NumBits < 64) ? ((uint64_t(1) << NumBits) - 1) : uint64_t(-1);
// NumBits mask for BitStart
static constexpr storage_type mask_lo = storage_type(mask) << bit_lo;
// NumBits mask for leftover bits in data_[idx+1] if straddling, else 0
static constexpr storage_type mask_hi = (idx + 1 < N) ? (storage_type(mask) >> bit_hi) : 0;
storage_type data_[N];
// Get value
CUTE_HOST_DEVICE constexpr
value_type get() const {
storage_type result = (data_[idx] & mask_lo) >> bit_lo;
if constexpr (bit_hi) {
result |= (data_[idx+1] & mask_hi) << bit_hi;
}
return static_cast<value_type>(result);
}
// Set value
CUTE_HOST_DEVICE constexpr
void set(value_type x) {
storage_type item = static_cast<storage_type>(x & mask);
data_[idx] = static_cast<storage_type>((data_[idx] & ~mask_lo) | (item << bit_lo));
if constexpr (bit_hi) {
data_[idx+1] = static_cast<storage_type>((data_[idx+1] & ~mask_hi) | (item >> bit_hi));
}
}
// Assign value
CUTE_HOST_DEVICE constexpr
bit_field& operator=(value_type x) {
set(x);
return *this;
}
// Cast to value
CUTE_HOST_DEVICE constexpr
operator value_type () const {
return get();
}
// Assign OtherValueType
CUTE_HOST_DEVICE constexpr
bit_field& operator=(OtherValueType x) {
return *this = *reinterpret_cast<value_type*>(&x);
}
// Cast to OtherValueType
CUTE_HOST_DEVICE constexpr
operator OtherValueType () const {
value_type x = get();
return *reinterpret_cast<OtherValueType*>(&x);
}
};
} // end namespace cute

View File

@ -0,0 +1,186 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#pragma once
#include <cute/config.hpp>
#include <cute/util/type_traits.hpp>
#include <cute/numeric/integral_constant.hpp>
namespace cute
{
//
// dim3
//
using dim3 = ::dim3;
// MSVC doesn't define its C++ version macro to match
// its C++ language version. This means that when
// building with MSVC, dim3 isn't constexpr-friendly.
template <size_t I>
CUTE_HOST_DEVICE
#if ! defined(_MSC_VER)
constexpr
#endif
uint32_t& get(dim3& a)
{
static_assert(I < 3, "Index out of range");
if constexpr (I == 0) {
return a.x;
} else if constexpr (I == 1) {
return a.y;
} else if constexpr (I == 2) {
return a.z;
}
CUTE_GCC_UNREACHABLE;
}
template <size_t I>
CUTE_HOST_DEVICE
#if ! defined(_MSC_VER)
constexpr
#endif
uint32_t const& get(dim3 const& a)
{
static_assert(I < 3, "Index out of range");
if constexpr (I == 0) {
return a.x;
} else if constexpr (I == 1) {
return a.y;
} else if constexpr (I == 2) {
return a.z;
}
CUTE_GCC_UNREACHABLE;
}
template <size_t I>
CUTE_HOST_DEVICE
#if ! defined(_MSC_VER)
constexpr
#endif
uint32_t&& get(dim3&& a)
{
static_assert(I < 3, "Index out of range");
if constexpr (I == 0) {
return std::move(a.x);
} else if constexpr (I == 1) {
return std::move(a.y);
} else if constexpr (I == 2) {
return std::move(a.z);
}
CUTE_GCC_UNREACHABLE;
}
// Specialize cute::tuple-traits for external types
template <>
struct tuple_size<dim3>
: integral_constant<size_t, 3>
{};
template <size_t I>
struct tuple_element<I, dim3>
{
using type = uint32_t;
};
//
// uint3
//
using uint3 = ::uint3;
template <size_t I>
CUTE_HOST_DEVICE constexpr
uint32_t& get(uint3& a)
{
static_assert(I < 3, "Index out of range");
if constexpr (I == 0) {
return a.x;
} else if constexpr (I == 1) {
return a.y;
} else if constexpr (I == 2) {
return a.z;
}
CUTE_GCC_UNREACHABLE;
}
template <size_t I>
CUTE_HOST_DEVICE constexpr
uint32_t const& get(uint3 const& a)
{
static_assert(I < 3, "Index out of range");
if constexpr (I == 0) {
return a.x;
} else if constexpr (I == 1) {
return a.y;
} else if constexpr (I == 2) {
return a.z;
}
CUTE_GCC_UNREACHABLE;
}
template <size_t I>
CUTE_HOST_DEVICE constexpr
uint32_t&& get(uint3&& a)
{
static_assert(I < 3, "Index out of range");
if constexpr (I == 0) {
return std::move(a.x);
} else if constexpr (I == 1) {
return std::move(a.y);
} else if constexpr (I == 2) {
return std::move(a.z);
}
CUTE_GCC_UNREACHABLE;
}
// Specialize cute::tuple-traits for external types
template <>
struct tuple_size<uint3>
: integral_constant<size_t, 3>
{};
template <size_t I>
struct tuple_element<I, uint3>
{
using type = uint32_t;
};
} // end namespace cute

View File

@ -0,0 +1,702 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#pragma once
#include <cute/config.hpp>
#include <cute/util/type_traits.hpp>
#include <cute/numeric/integral_constant.hpp> // cute::true_type, cute::false_type
#include <cute/numeric/integer_sequence.hpp>
#include <cute/container/cuda_types.hpp>
//#include <cute/container/array.hpp> // Advanced optimizations
//
// cute::tuple is like std::tuple, with two differences.
//
// 1. It works on both host and device.
// 2. Its template arguments must be semiregular types.
//
// Semiregular types are default constructible and copyable.
// They include "value types" like int or float,
// but do _not_ include references like int& or float&.
// (See std::tie for an example of a tuple of references.)
//
// This is simplified over the implementations in std::, cuda::std::, and thrust:: by ignoring much of
// the conversion SFINAE, special overloading, and avoiding cvref template types.
// Furthermore, the empty base optimization (EBO) is MORE aggressive by avoiding
// construction calls, and ignoring any need for unique element addresses.
//
// Over standard-conforming tuple implementations, this appears to accelerate compilation times by over 3x.
namespace cute
{
namespace detail
{
// EBO stands for "empty base optimization."
// We use this technique to ensure that cute::tuple
// doesn't need to waste space storing any template arguments
// of cute::tuple that have no data (like integral_constant).
// Otherwise, cute::tuple would need to spend at least 1 byte
// for each of its template arguments.
//
// EBO always "holds" a single value of type T.
// N is like an array index that TupleBase uses
// to access the desired tuple element.
template <size_t N, class T, bool IsEmpty = is_empty<T>::value>
struct EBO;
// Specialization for types T that have no data;
// the "static tuple leaf." Valid T here include
// integral_constant<U, Value>, Int<Value>,
// and any other semiregular type
// for which std::is_empty_v<T> is true.
template <size_t N, class T>
struct EBO<N, T, true>
{
CUTE_HOST_DEVICE constexpr
EBO() {}
CUTE_HOST_DEVICE constexpr
EBO(T const&) {}
};
template <size_t N, class T>
CUTE_HOST_DEVICE constexpr T getv(EBO<N, T, true> const&)
{ return {}; }
// Specialization for types T that are not empty;
// the "dynamic tuple leaf." Valid T here include int,
// any other integral or floating-point type,
// or any semiregular type for which std::is_empty_v<T> is false.
template <size_t N, class T>
struct EBO<N, T, false>
{
CUTE_HOST_DEVICE constexpr
EBO() : t_{} {}
template <class U>
CUTE_HOST_DEVICE constexpr
EBO(U const& u) : t_{u} {}
T t_;
};
template <size_t N, class T>
CUTE_HOST_DEVICE constexpr T const& getv(EBO<N, T, false> const& x)
{ return x.t_; }
template <size_t N, class T>
CUTE_HOST_DEVICE constexpr T& getv(EBO<N, T, false>& x)
{ return x.t_; }
template <size_t N, class T>
CUTE_HOST_DEVICE constexpr T&& getv(EBO<N, T, false>&& x)
{ return static_cast<T&&>(x.t_); }
template <class IdxSeq, class... T>
struct TupleBase;
// Base class of cute::tuple.
// It inherits from EBO<i, t> for each (i, t) in (I..., T...).
// The actual storage (for nonempty t) lives in the base classes.
// index_sequence is a way to wrap up a sequence of zero or more
// compile-time integer values in a single type.
// We only ever use index_sequence<0, 1, ..., sizeof...(T)> in practice,
// as the type alias TupleBase below indicates.
template <size_t... I, class... T>
struct TupleBase<index_sequence<I...>, T...>
: EBO<I,T>...
{
CUTE_HOST_DEVICE constexpr
TupleBase() {}
template <class... U>
CUTE_HOST_DEVICE constexpr explicit
TupleBase(U const&... u)
: EBO<I,T>(u)... {}
template <class... U>
CUTE_HOST_DEVICE constexpr
TupleBase(TupleBase<index_sequence<I...>, U...> const& u)
: EBO<I,T>(getv(static_cast<EBO<I,U> const&>(u)))... {}
};
} // end namespace detail
// Attempting to use the following commented-out alias
// in the declaration of `struct tuple` causes MSVC 2022 build errors.
//
//template <class... T>
//using TupleBase = detail::TupleBase<make_index_sequence<sizeof...(T)>, T...>;
// This is the actual cute::tuple class.
// The storage (if any) lives in TupleBase's EBO base classes.
//
// Inheriting from the above alias TupleBase
// causes MSVC 2022 build errors when assigning one tuple to another:
//
// illegal member initialization:
// 'TupleBase< /* template arguments */ >' is not a base or member
//
// Not using the alias or any kind of alias fixed the errors.
// In summary: this is verbose as a work-around for MSVC build errors.
template <class... T>
struct tuple : detail::TupleBase<make_index_sequence<sizeof...(T)>, T...>
{
CUTE_HOST_DEVICE constexpr
tuple() {}
template <class... U>
CUTE_HOST_DEVICE constexpr
tuple(U const&... u) : detail::TupleBase<make_index_sequence<sizeof...(T)>, T...>(u...) {}
template <class... U>
CUTE_HOST_DEVICE constexpr
tuple(tuple<U...> const& u)
: detail::TupleBase<make_index_sequence<sizeof...(T)>, T...>(static_cast<detail::TupleBase<make_index_sequence<sizeof...(U)>, U...> const&>(u)) {}
};
//
// get for cute::tuple (just like std::get for std::tuple)
//
template <size_t I, class... T>
CUTE_HOST_DEVICE constexpr
decltype(auto)
get(tuple<T...> const& t) noexcept
{
static_assert(I < sizeof...(T), "Index out of range");
return detail::getv<I>(t);
}
template <size_t I, class... T>
CUTE_HOST_DEVICE constexpr
decltype(auto)
get(tuple<T...>& t) noexcept
{
static_assert(I < sizeof...(T), "Index out of range");
return detail::getv<I>(t);
}
template <size_t I, class... T>
CUTE_HOST_DEVICE constexpr
decltype(auto)
get(tuple<T...>&& t) noexcept
{
static_assert(I < sizeof...(T), "Index out of range");
return detail::getv<I>(static_cast<tuple<T...>&&>(t));
}
//
// Custom is_tuple trait simply checks the existence of tuple_size
// and assumes std::get<I>(.), std::tuple_element<I,.>
//
namespace detail {
template <class T>
auto has_tuple_size( T*) -> integral_constant<bool, 0 <= tuple_size<T>::value>;
auto has_tuple_size(...) -> false_type;
} // end namespace detail
template <class T>
struct is_tuple : decltype(detail::has_tuple_size((T*)0)) {};
//
// make_tuple (value-based implementation)
//
template <class... T>
CUTE_HOST_DEVICE constexpr
tuple<T...>
make_tuple(T const&... t)
{
return {t...};
}
//
// tuple_cat concatenates multiple cute::tuple into a single cute::tuple,
// just like std::tuple_cat for std::tuple.
//
#if 0
// Original implementation
namespace detail {
template <class T0, class T1,
size_t... I0, size_t... I1>
CUTE_HOST_DEVICE constexpr
auto
tuple_cat(T0 const& t0, T1 const& t1,
index_sequence<I0...>, index_sequence<I1...>)
{
return cute::make_tuple(get<I0>(t0)..., get<I1>(t1)...);
}
} // end namespace detail
CUTE_HOST_DEVICE constexpr
tuple<>
tuple_cat()
{
return {};
}
template <class Tuple,
__CUTE_REQUIRES(is_tuple<Tuple>::value)>
CUTE_HOST_DEVICE constexpr
Tuple const&
tuple_cat(Tuple const& t)
{
return t;
}
template <class T0, class T1>
CUTE_HOST_DEVICE constexpr
auto
tuple_cat(T0 const& t0, T1 const& t1)
{
return detail::tuple_cat(t0, t1,
make_index_sequence<tuple_size<T0>::value>{},
make_index_sequence<tuple_size<T1>::value>{});
}
template <class T0, class T1, class T2, class... Ts>
CUTE_HOST_DEVICE constexpr
auto
tuple_cat(T0 const& t0, T1 const& t1, T2 const& t2, Ts const&... ts)
{
return cute::tuple_cat(cute::tuple_cat(t0,t1),t2,ts...);
}
#endif
#if 1
// Extended implementation
namespace detail {
template <class T0, class T1,
size_t... I0, size_t... I1>
CUTE_HOST_DEVICE constexpr
auto
tuple_cat(T0 const& t0, T1 const& t1,
index_sequence<I0...>, index_sequence<I1...>)
{
return cute::make_tuple(get<I0>(t0)..., get<I1>(t1)...);
}
template <class T0, class T1, class T2,
size_t... I0, size_t... I1, size_t... I2>
CUTE_HOST_DEVICE constexpr
auto
tuple_cat(T0 const& t0, T1 const& t1, T2 const& t2,
index_sequence<I0...>, index_sequence<I1...>, index_sequence<I2...>)
{
return cute::make_tuple(get<I0>(t0)..., get<I1>(t1)..., get<I2>(t2)...);
}
template <class T0, class T1, class T2, class T3,
size_t... I0, size_t... I1, size_t... I2, size_t... I3>
CUTE_HOST_DEVICE constexpr
auto
tuple_cat(T0 const& t0, T1 const& t1, T2 const& t2, T3 const& t3,
index_sequence<I0...>, index_sequence<I1...>, index_sequence<I2...>, index_sequence<I3...>)
{
return cute::make_tuple(get<I0>(t0)..., get<I1>(t1)..., get<I2>(t2)..., get<I3>(t3)...);
}
template <class T0, class T1, class T2, class T3, class T4,
size_t... I0, size_t... I1, size_t... I2, size_t... I3, size_t... I4>
CUTE_HOST_DEVICE constexpr
auto
tuple_cat(T0 const& t0, T1 const& t1, T2 const& t2, T3 const& t3, T4 const& t4,
index_sequence<I0...>, index_sequence<I1...>, index_sequence<I2...>, index_sequence<I3...>, index_sequence<I4...>)
{
return cute::make_tuple(get<I0>(t0)..., get<I1>(t1)..., get<I2>(t2)..., get<I3>(t3)..., get<I4>(t4)...);
}
} // end namespace detail
CUTE_HOST_DEVICE constexpr
tuple<>
tuple_cat()
{
return {};
}
template <class Tuple,
__CUTE_REQUIRES(is_tuple<Tuple>::value)>
CUTE_HOST_DEVICE constexpr
Tuple const&
tuple_cat(Tuple const& t)
{
return t;
}
template <class T0, class T1>
CUTE_HOST_DEVICE constexpr
auto
tuple_cat(T0 const& t0, T1 const& t1)
{
return detail::tuple_cat(t0, t1,
make_index_sequence<tuple_size<T0>::value>{},
make_index_sequence<tuple_size<T1>::value>{});
}
template <class T0, class T1, class T2>
CUTE_HOST_DEVICE constexpr
auto
tuple_cat(T0 const& t0, T1 const& t1, T2 const& t2)
{
return detail::tuple_cat(t0, t1, t2,
make_index_sequence<tuple_size<T0>::value>{},
make_index_sequence<tuple_size<T1>::value>{},
make_index_sequence<tuple_size<T2>::value>{});
}
template <class T0, class T1, class T2, class T3>
CUTE_HOST_DEVICE constexpr
auto
tuple_cat(T0 const& t0, T1 const& t1, T2 const& t2, T3 const& t3)
{
return detail::tuple_cat(t0, t1, t2, t3,
make_index_sequence<tuple_size<T0>::value>{},
make_index_sequence<tuple_size<T1>::value>{},
make_index_sequence<tuple_size<T2>::value>{},
make_index_sequence<tuple_size<T3>::value>{});
}
template <class T0, class T1, class T2, class T3, class T4>
CUTE_HOST_DEVICE constexpr
auto
tuple_cat(T0 const& t0, T1 const& t1, T2 const& t2, T3 const& t3, T4 const& t4)
{
return detail::tuple_cat(t0, t1, t2, t3, t4,
make_index_sequence<tuple_size<T0>::value>{},
make_index_sequence<tuple_size<T1>::value>{},
make_index_sequence<tuple_size<T2>::value>{},
make_index_sequence<tuple_size<T3>::value>{},
make_index_sequence<tuple_size<T4>::value>{});
}
template <class T0, class T1, class T2, class T3, class T4, class T5, class... Ts>
CUTE_HOST_DEVICE constexpr
auto
tuple_cat(T0 const& t0, T1 const& t1, T2 const& t2, T3 const& t3, T4 const& t4, T5 const& t5, Ts const&... ts)
{
return cute::tuple_cat(cute::tuple_cat(t0,t1,t2,t3,t4), t5, ts...);
}
#endif
#if 0
// Outer-Inner indexing trick to concat all tuples at once
namespace detail {
template <size_t... Ns>
struct tuple_cat_helper
{
static constexpr cute::array<size_t,sizeof...(Ns)> ns = {Ns...};
static constexpr size_t total_size() {
size_t sum = 0;
for (size_t n : ns) sum += n;
return sum;
}
static constexpr size_t total_size_ = total_size();
static constexpr auto values() {
cute::array<size_t[2],total_size_> outer_inner = {};
size_t idx = 0;
for (size_t i = 0; i < ns.size(); ++i) {
for (size_t j = 0; j < ns[i]; ++j, ++idx) {
outer_inner[idx][0] = i;
outer_inner[idx][1] = j;
}
}
return outer_inner;
}
static constexpr auto outer_inner_ = values();
using total_sequence = make_index_sequence<total_size_>;
};
template <class Helper, class Tuple, size_t... I>
CUTE_HOST_DEVICE constexpr
auto
tuple_cat(Tuple const& t, index_sequence<I...>)
{
return cute::make_tuple(get<Helper::outer_inner_[I][1]>(get<Helper::outer_inner_[I][0]>(t))...);
}
template <class T0, class T1,
size_t... I0, size_t... I1>
CUTE_HOST_DEVICE constexpr
auto
tuple_cat(T0 const& t0, T1 const& t1,
index_sequence<I0...>, index_sequence<I1...>)
{
return cute::make_tuple(get<I0>(t0)..., get<I1>(t1)...);
}
} // end namespace detail
CUTE_HOST_DEVICE constexpr
tuple<>
tuple_cat()
{
return {};
}
template <class Tuple,
__CUTE_REQUIRES(is_tuple<Tuple>::value)>
CUTE_HOST_DEVICE constexpr
Tuple const&
tuple_cat(Tuple const& t)
{
return t;
}
template <class T0, class T1>
CUTE_HOST_DEVICE constexpr
auto
tuple_cat(T0 const& t0, T1 const& t1)
{
return detail::tuple_cat(t0, t1,
make_index_sequence<tuple_size<T0>::value>{},
make_index_sequence<tuple_size<T1>::value>{});
}
template <class... Tuples>
CUTE_HOST_DEVICE constexpr
auto
tuple_cat(Tuples const&... ts)
{
using Helper = detail::tuple_cat_helper<tuple_size<Tuples>::value...>;
return detail::tuple_cat<Helper>(cute::make_tuple(ts...), typename Helper::total_sequence{});
}
#endif
//
// Equality operators
//
namespace detail {
template <size_t I, class TupleA, class TupleB>
CUTE_HOST_DEVICE constexpr
auto
equal_impl(TupleA const& a, TupleB const& b)
{
if constexpr (I == tuple_size<TupleA>::value) {
return cute::true_type{}; // Terminal: TupleA is exhausted
} else if constexpr (I == tuple_size<TupleB>::value) {
return cute::false_type{}; // Terminal: TupleA is not exhausted, TupleB is exhausted
} else {
return (get<I>(a) == get<I>(b)) && equal_impl<I+1>(a,b);
}
CUTE_GCC_UNREACHABLE;
}
} // end namespace detail
template <class TupleT, class TupleU,
__CUTE_REQUIRES(is_tuple<TupleT>::value && is_tuple<TupleU>::value)>
CUTE_HOST_DEVICE constexpr
auto
operator==(TupleT const& t, TupleU const& u)
{
return detail::equal_impl<0>(t, u);
}
template <class TupleT, class TupleU,
__CUTE_REQUIRES(is_tuple<TupleT>::value ^ is_tuple<TupleU>::value)>
CUTE_HOST_DEVICE constexpr
auto
operator==(TupleT const& t, TupleU const& u)
{
return cute::false_type{};
}
template <class TupleT, class TupleU,
__CUTE_REQUIRES(is_tuple<TupleT>::value && is_tuple<TupleU>::value)>
CUTE_HOST_DEVICE constexpr
auto
operator!=(TupleT const& t, TupleU const& u)
{
return !(t == u);
}
template <class TupleT, class TupleU,
__CUTE_REQUIRES(is_tuple<TupleT>::value ^ is_tuple<TupleU>::value)>
CUTE_HOST_DEVICE constexpr
auto
operator!=(TupleT const& t, TupleU const& u)
{
return cute::true_type{};
}
//
// Comparison operators
//
//
// There are many ways to compare tuple of elements and because CuTe is built
// on parameterizing layouts of coordinates, some comparisons are appropriate
// only in certain cases.
// -- lexicographical comparison [reverse, reflected, revref]
// -- colexicographical comparison [reverse, reflected, revref]
// -- element-wise comparison [any,all]
// This can be very confusing. To avoid errors in selecting the appropriate
// comparison, op<|op<=|op>|op>= are *not* implemented for cute::tuple.
//
// That said, see int_tuple for more explicitly named common comparison ops.
//
//
// Display utilities
//
namespace detail {
template <class Tuple, size_t... Is>
CUTE_HOST_DEVICE void print_tuple(Tuple const& t,
index_sequence<Is...>, char s = '(', char e = ')')
{
using eat = int[];
using cute::print;
(void) eat {(print(s), 0),
(print(Is == 0 ? "" : ","), print(get<Is>(t)), 0)...,
(print(e), 0)};
}
#if !defined(__MACACC_RTC__)
template <class Tuple, std::size_t... Is>
CUTE_HOST std::ostream& print_tuple_os(std::ostream& os, Tuple const& t,
index_sequence<Is...>, char s = '(', char e = ')')
{
using eat = int[];
(void) eat {(void(os << s), 0),
(void(os << (Is == 0 ? "" : ",") << get<Is>(t)), 0)...,
(void(os << e), 0)};
return os;
}
#endif // !defined(__MACACC_RTC__)
} // end namespace detail
template <class Tuple,
__CUTE_REQUIRES(is_tuple<Tuple>::value)>
CUTE_HOST_DEVICE void print(Tuple const& t)
{
return detail::print_tuple(t, make_index_sequence<tuple_size<Tuple>::value>{});
}
#if !defined(__MACACC_RTC__)
template <class Tuple,
__CUTE_REQUIRES(is_tuple<Tuple>::value)>
CUTE_HOST std::ostream& operator<<(std::ostream& os, Tuple const& t)
{
return detail::print_tuple_os(os, t, make_index_sequence<tuple_size<Tuple>::value>{});
}
#endif // !defined(__MACACC_RTC__)
} // end namespace cute
namespace CUTE_STL_NAMESPACE
{
template <class... T>
struct tuple_size<cute::tuple<T...>>
: cute::integral_constant<size_t, sizeof...(T)>
{};
template <size_t I, class... T>
struct tuple_element<I, cute::tuple<T...>>
: CUTE_STL_NAMESPACE::tuple_element<I, CUTE_STL_NAMESPACE::tuple<T...>>
{};
template <class... T>
struct tuple_size<const cute::tuple<T...>>
: cute::integral_constant<size_t, sizeof...(T)>
{};
template <size_t I, class... T>
struct tuple_element<I, const cute::tuple<T...>>
: CUTE_STL_NAMESPACE::tuple_element<I, const CUTE_STL_NAMESPACE::tuple<T...>>
{};
} // end namespace CUTE_STL_NAMESPACE
//
// std compatibility
//
#ifdef CUTE_STL_NAMESPACE_IS_CUDA_STD
namespace std
{
#if defined(__MACACC_RTC__)
template <class... _Tp>
struct tuple_size;
template<size_t _Ip, class... _Tp>
struct tuple_element;
#endif
template <class... T>
struct tuple_size<cute::tuple<T...>>
: cute::integral_constant<size_t, sizeof...(T)>
{};
template <size_t I, class... T>
struct tuple_element<I, cute::tuple<T...>>
: CUTE_STL_NAMESPACE::tuple_element<I, CUTE_STL_NAMESPACE::tuple<T...>>
{};
template <class... T>
struct tuple_size<const cute::tuple<T...>>
: cute::integral_constant<size_t, sizeof...(T)>
{};
template <size_t I, class... T>
struct tuple_element<I, const cute::tuple<T...>>
: CUTE_STL_NAMESPACE::tuple_element<I, const CUTE_STL_NAMESPACE::tuple<T...>>
{};
} // end namepsace std
#endif // CUTE_STL_NAMESPACE_IS_CUDA_STD

View File

@ -0,0 +1,136 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#pragma once
#include <cute/numeric/integral_constant.hpp>
namespace cute
{
template <class T>
struct type_c {
using type = T;
};
template <class... T>
struct type_list {};
} // end namespace cute
//
// Specialize tuple-related functionality for cute::type_list
//
#if defined(__MACACC_RTC__)
#include <cuda/std/tuple>
#else
#include <tuple>
#endif
#include <cute/container/tuple.hpp>
namespace cute
{
template <int I, class... T>
CUTE_HOST_DEVICE constexpr
CUTE_STL_NAMESPACE::tuple_element_t<I, type_list<T...>>
get(type_list<T...>&) noexcept {
return {};
}
template <int I, class... T>
CUTE_HOST_DEVICE constexpr
CUTE_STL_NAMESPACE::tuple_element_t<I, type_list<T...>>
get(type_list<T...> const& t) noexcept {
return {};
}
} // end namespace cute
namespace CUTE_STL_NAMESPACE
{
template <class... T>
struct tuple_size<cute::type_list<T...>>
: cute::integral_constant<size_t, sizeof...(T)>
{};
template <size_t I, class... T>
struct tuple_element<I, cute::type_list<T...>>
: cute::type_c<typename CUTE_STL_NAMESPACE::tuple_element<I, CUTE_STL_NAMESPACE::tuple<T...>>::type>
{};
template <class... T>
struct tuple_size<const cute::type_list<T...>>
: cute::integral_constant<size_t, sizeof...(T)>
{};
template <size_t I, class... T>
struct tuple_element<I, const cute::type_list<T...>>
: cute::type_c<typename CUTE_STL_NAMESPACE::tuple_element<I, CUTE_STL_NAMESPACE::tuple<T...>>::type>
{};
} // end namespace std
#ifdef CUTE_STL_NAMESPACE_IS_CUDA_STD
namespace std
{
#if defined(__MACACC_RTC__)
template <class... _Tp>
struct tuple_size;
template<size_t _Ip, class... _Tp>
struct tuple_element;
#endif
template <class... T>
struct tuple_size<cute::type_list<T...>>
: cute::integral_constant<size_t, sizeof...(T)>
{};
template <size_t I, class... T>
struct tuple_element<I, cute::type_list<T...>>
: cute::type_c<typename CUTE_STL_NAMESPACE::tuple_element<I, CUTE_STL_NAMESPACE::tuple<T...>>::type>
{};
template <class... T>
struct tuple_size<const cute::type_list<T...>>
: cute::integral_constant<size_t, sizeof...(T)>
{};
template <size_t I, class... T>
struct tuple_element<I, const cute::type_list<T...>>
: cute::type_c<typename CUTE_STL_NAMESPACE::tuple_element<I, CUTE_STL_NAMESPACE::tuple<T...>>::type>
{};
} // end namespace std
#endif // CUTE_STL_NAMESPACE_IS_CUDA_STD

View File

@ -0,0 +1,875 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#pragma once
#include <cute/config.hpp>
#include <cute/container/tuple.hpp>
#include <cute/container/array.hpp>
#include <cute/algorithm/tuple_algorithms.hpp>
#include <cute/numeric/integral_constant.hpp>
namespace cute
{
template <class... Ts>
using IntTuple = cute::tuple<Ts...>;
// Construct an IntTuple with all value-elements
template <class... Ts>
CUTE_HOST_DEVICE constexpr
IntTuple<Ts...>
make_int_tuple(Ts const&... t)
{
return {t...};
}
/** if rank(int) == 1, then get<0>(int) should work too
*/
template <size_t I, class T, __CUTE_REQUIRES(is_integral<remove_cvref_t<T>>::value)>
CUTE_HOST_DEVICE constexpr
decltype(auto)
get(T&& t) noexcept
{
static_assert(I == 0, "Index out of range");
return static_cast<T&&>(t);
}
/** Custom recursive get for anything that implements get<I>(.)
*/
template <size_t I0, size_t I1, size_t... Is, class Tuple>
CUTE_HOST_DEVICE constexpr
decltype(auto)
get(Tuple&& t) noexcept
{
return get<I1,Is...>(get<I0>(static_cast<Tuple&&>(t)));
}
//
// rank
//
template <int... Is, class IntTuple>
CUTE_HOST_DEVICE constexpr
auto
rank(IntTuple const& t)
{
if constexpr (sizeof...(Is) == 0) {
if constexpr (is_tuple<IntTuple>::value) {
return Int<tuple_size<IntTuple>::value>{};
} else {
return Int<1>{};
}
} else {
return rank(get<Is...>(t));
}
CUTE_GCC_UNREACHABLE;
}
template <class IntTuple>
using rank_t = decltype(rank(declval<IntTuple>()));
template <class IntTuple>
static constexpr int rank_v = rank_t<IntTuple>::value;
//
// shape
//
template <class IntTuple>
CUTE_HOST_DEVICE constexpr
auto
shape(IntTuple const& s)
{
if constexpr (is_tuple<IntTuple>::value) {
return transform(s, [](auto const& a) { return shape(a); });
} else {
return s;
}
CUTE_GCC_UNREACHABLE;
}
template <int I, int... Is, class IntTuple>
CUTE_HOST_DEVICE constexpr
auto
shape(IntTuple const& s)
{
if constexpr (is_tuple<IntTuple>::value) {
return shape<Is...>(get<I>(s));
} else {
return get<I,Is...>(shape(s));
}
CUTE_GCC_UNREACHABLE;
}
//
// max
//
template <class T0, class... Ts>
CUTE_HOST_DEVICE constexpr
auto
max(T0 const& t0, Ts const&... ts)
{
if constexpr (is_tuple<T0>::value) {
return cute::max(cute::apply(t0, [](auto const&... a){ return cute::max(a...); }), ts...);
} else if constexpr (sizeof...(Ts) == 0) {
return t0;
} else {
return cute::max(t0, cute::max(ts...));
}
CUTE_GCC_UNREACHABLE;
}
//
// min
//
template <class T0, class... Ts>
CUTE_HOST_DEVICE constexpr
auto
min(T0 const& t0, Ts const&... ts)
{
if constexpr (is_tuple<T0>::value) {
return cute::min(cute::apply(t0, [](auto const&... a){ return cute::min(a...); }), ts...);
} else if constexpr (sizeof...(Ts) == 0) {
return t0;
} else {
return cute::min(t0, cute::min(ts...));
}
CUTE_GCC_UNREACHABLE;
}
//
// depth
//
template <int... Is, class IntTuple>
CUTE_HOST_DEVICE constexpr
auto
depth(IntTuple const& t)
{
if constexpr (sizeof...(Is) == 0) {
if constexpr (is_tuple<IntTuple>::value) {
return Int<1>{} + cute::apply(t, [](auto const&... v){ return cute::max(depth(v)...); });
} else {
return Int<0>{};
}
} else {
return depth(get<Is...>(t));
}
CUTE_GCC_UNREACHABLE;
}
template <class Tuple>
using depth_t = decltype(depth(declval<Tuple>()));
template <class Tuple>
static constexpr int depth_v = depth_t<Tuple>::value;
//
// product
//
template <class IntTuple>
CUTE_HOST_DEVICE constexpr
auto
product(IntTuple const& a)
{
if constexpr (is_tuple<IntTuple>::value) {
return cute::apply(a, [](auto const&... v){ return (Int<1>{} * ... * product(v)); });
} else {
return a;
}
CUTE_GCC_UNREACHABLE;
}
template <class Tuple>
CUTE_HOST_DEVICE constexpr
auto
product_each(Tuple const& t)
{
return transform(t, [](auto const& x) { return product(x); });
}
// Return the product of elements in a mode
template <int... Is, class IntTuple>
CUTE_HOST_DEVICE constexpr
auto
size(IntTuple const& a)
{
if constexpr (sizeof...(Is) == 0) {
return product(a);
} else {
return product(get<Is...>(a));
}
CUTE_GCC_UNREACHABLE;
}
template <class IntTuple>
static constexpr int size_v = decltype(size(declval<IntTuple>()))::value;
//
// sum
//
template <class IntTuple>
CUTE_HOST_DEVICE constexpr
auto
sum(IntTuple const& a)
{
if constexpr (is_tuple<IntTuple>::value) {
return cute::apply(a, [](auto const&... v){ return (Int<0>{} + ... + sum(v)); });
} else {
return a;
}
CUTE_GCC_UNREACHABLE;
}
//
// inner_product
//
template <class IntTupleA, class IntTupleB>
CUTE_HOST_DEVICE constexpr
auto
inner_product(IntTupleA const& a, IntTupleB const& b)
{
if constexpr (is_tuple<IntTupleA>::value && is_tuple<IntTupleB>::value) {
static_assert(tuple_size<IntTupleA>::value == tuple_size<IntTupleB>::value, "Mismatched ranks");
return transform_apply(a, b, [](auto const& x, auto const& y) { return inner_product(x,y); },
[](auto const&... v) { return (Int<0>{} + ... + v); });
} else {
return a * b;
}
CUTE_GCC_UNREACHABLE;
}
//
// ceil_div
//
template <class IntTupleA, class IntTupleB>
CUTE_HOST_DEVICE constexpr
auto
ceil_div(IntTupleA const& a, IntTupleB const& b)
{
if constexpr (is_tuple<IntTupleA>::value && is_tuple<IntTupleB>::value) {
static_assert(tuple_size<IntTupleA>::value >= tuple_size<IntTupleB>::value, "Mismatched ranks");
constexpr int R = tuple_size<IntTupleA>::value; // Missing ranks in TupleB are implictly 1
return transform(a, append<R>(b,Int<1>{}), [](auto const& x, auto const& y) { return ceil_div(x,y); });
} else {
return (a + b - Int<1>{}) / b;
}
CUTE_GCC_UNREACHABLE;
}
/** Division for Shapes
*/
template <class IntTupleA, class IntTupleB>
CUTE_HOST_DEVICE constexpr
auto
shape_div(IntTupleA const& a, IntTupleB const& b)
{
if constexpr (is_tuple<IntTupleA>::value) {
if constexpr (is_tuple<IntTupleB>::value) { // tuple tuple
static_assert(tuple_size<IntTupleA>::value == tuple_size<IntTupleB>::value, "Mismatched ranks");
return transform(a, b, [](auto const& x, auto const& y) { return shape_div(x,y); });
} else { // tuple int
auto const [result, rest] = fold(a, cute::make_tuple(cute::make_tuple(), b),
[] (auto const& init, auto const& ai) {
return cute::make_tuple(append(get<0>(init), shape_div(ai, get<1>(init))), shape_div(get<1>(init), ai));
});
return result;
}
} else {
if constexpr (is_tuple<IntTupleB>::value) { // int tuple
return shape_div(a, product(b));
} else { // int int
//assert(a % b == 0 || b % a == 0);
return a / b != 0 ? a / b : signum(a) * signum(b); // divide with rounding away from zero
}
}
CUTE_GCC_UNREACHABLE;
}
/** Division for Shapes that are static constants
* @pre t % u == 0 || u % t == 0
* @result if t % u == 0, then t / u
* if u % t == 0, then signum(t) * signum(u)
*/
template <class T, T t, class U, U u>
CUTE_HOST_DEVICE constexpr
constant<decltype(shape_div(t,u)), shape_div(t,u)>
shape_div(constant<T, t> const&, constant<U, u> const&)
{
static_assert(t % u == 0 || u % t == 0, "Static shape_div failure");
return {};
}
/** Return a tuple the same profile as A scaled by corresponding elements in B
*/
template <class A, class B>
CUTE_HOST_DEVICE constexpr
auto
elem_scale(A const& a, B const& b)
{
if constexpr (is_tuple<A>::value) {
return transform(a, b, [](auto const& x, auto const& y) { return elem_scale(x,y); });
} else {
return a * product(b);
}
CUTE_GCC_UNREACHABLE;
}
/** Test if two IntTuple have the same profile (hierarchical rank division)
*/
template <class IntTupleA, class IntTupleB>
CUTE_HOST_DEVICE constexpr
auto
congruent(IntTupleA const& a, IntTupleB const& b)
{
return bool_constant<is_same<decltype(repeat_like(shape(a),_0{})),
decltype(repeat_like(shape(b),_0{}))>::value>{};
}
template <class A, class B>
using is_congruent = decltype(congruent(declval<A>(), declval<B>()));
/** Test if two IntTuple have the similar profiles up to Shape A (hierarchical rank division)
*/
template <class IntTupleA, class IntTupleB>
CUTE_HOST_DEVICE constexpr
auto
weakly_congruent(IntTupleA const& a, IntTupleB const& b)
{
if constexpr (is_tuple<IntTupleA>::value && is_tuple<IntTupleB>::value) {
if constexpr (tuple_size<IntTupleA>::value != tuple_size<IntTupleB>::value) {
return false_type{};
} else {
return transform_apply(a, b, [](auto const& x, auto const& y) { return weakly_congruent(x,y); },
[](auto const&... z) { return (true_type{} && ... && z); });
}
} else if constexpr (is_integral<IntTupleA>::value) {
return true_type{};
} else if constexpr (is_integral<IntTupleB>::value) {
return false_type{};
} else {
return weakly_congruent(shape(a), shape(b));
}
CUTE_GCC_UNREACHABLE;
}
template <class A, class B>
using is_weakly_congruent = decltype(weakly_congruent(declval<A>(), declval<B>()));
/** Test if Shape B is compatible with Shape A:
* Any coordinate into A can also be used as a coordinate into B
* A <= B is a partially ordered set of factored shapes
*/
template <class IntTupleA, class IntTupleB>
CUTE_HOST_DEVICE constexpr
auto
compatible(IntTupleA const& a, IntTupleB const& b)
{
if constexpr (is_tuple<IntTupleA>::value && is_tuple<IntTupleB>::value) {
if constexpr (tuple_size<IntTupleA>::value != tuple_size<IntTupleB>::value) {
return false_type{};
} else {
return transform_apply(a, b, [](auto const& x, auto const& y) { return compatible(x,y); },
[](auto const&... z) { return (true_type{} && ... && z); });
}
} else if constexpr (is_integral<IntTupleA>::value) {
return a == size(b);
} else if constexpr (is_integral<IntTupleB>::value) {
return false_type{};
} else {
return compatible(shape(a), shape(b));
}
CUTE_GCC_UNREACHABLE;
}
template <class A, class B>
using is_compatible = decltype(compatible(declval<A>(), declval<B>()));
/** Test if Shape B is weakly compatible with Shape A:
* Shape B divides Shape A at some level of refinement
*/
template <class IntTupleA, class IntTupleB>
CUTE_HOST_DEVICE constexpr
auto
weakly_compatible(IntTupleA const& a, IntTupleB const& b)
{
if constexpr (is_tuple<IntTupleA>::value && is_tuple<IntTupleB>::value) {
if constexpr (tuple_size<IntTupleA>::value != tuple_size<IntTupleB>::value) {
return false_type{};
} else {
return transform_apply(a, b, [](auto const& x, auto const& y) { return weakly_compatible(x,y); },
[](auto const&... z) { return (true_type{} && ... && z); });
}
} else if constexpr (is_integral<IntTupleA>::value) {
return a % size(b) == Int<0>{};
} else if constexpr (is_integral<IntTupleB>::value) {
return false_type{};
} else {
return weakly_compatible(shape(a), shape(b));
}
CUTE_GCC_UNREACHABLE;
}
template <class A, class B>
using is_weakly_compatible = decltype(weakly_compatible(declval<A>(), declval<B>()));
/** Replace the elements of Tuple B that are paired with an Int<0> with an Int<1>
*/
template <class IntTupleA, class IntTupleB>
CUTE_HOST_DEVICE constexpr
auto
filter_zeros(IntTupleA const& a, IntTupleB const& b)
{
if constexpr (is_tuple<IntTupleA>::value) {
return transform(a, b, [](auto const& x, auto const& y) { return filter_zeros(x,y); });
} else if constexpr (is_constant<0, IntTupleA>::value) {
return Int<1>{};
} else {
return b;
}
CUTE_GCC_UNREACHABLE;
}
template <class Tuple>
CUTE_HOST_DEVICE constexpr
auto
filter_zeros(Tuple const& t)
{
return filter_zeros(t, t);
}
//
// Converters and constructors with arrays and params
//
/** Make an IntTuple of rank N from an Indexable array.
* Access elements up to a dynamic index n, then use init (requires compatible types)
* Consider cute::take<B,E> if all indexing is known to be valid
* \code
* std::vector<int> a = {6,3,4};
* auto tup = make_int_tuple<5>(a, a.size(), 0) // (6,3,4,0,0)
* \endcode
*/
template <int N, class Indexable, class T>
CUTE_HOST_DEVICE constexpr
auto
make_int_tuple(Indexable const& t, int n, T const& init)
{
static_assert(N > 0);
if constexpr (N == 1) {
return 0 < n ? t[0] : init;
} else {
return transform(make_seq<N>{}, [&](auto i) { return i < n ? t[i] : init; });
}
CUTE_GCC_UNREACHABLE;
}
/** Fill the dynamic values of a Tuple with values from another Tuple
* \code
* auto params = make_int_tuple(6,3,4);
* cute::tuple<Int<1>, cute::tuple<int, int, Int<3>>, int, Int<2>> result;
* fill_int_tuple_from(result, params); // (_1,(6,3,_3),4,_2)
* \endcode
*/
template <class Tuple, class TupleV>
CUTE_HOST_DEVICE constexpr
auto
fill_int_tuple_from(Tuple& result, TupleV const& vals)
{
return fold(result, vals, [](auto const& init, auto&& r) {
if constexpr (is_static<remove_cvref_t<decltype(r)>>::value) { // Skip static elements of result
return init;
} else if constexpr (is_tuple<remove_cvref_t<decltype(r)>>::value) { // Recurse into tuples
return fill_int_tuple_from(r, init);
} else { // Assign and consume arg
static_assert(tuple_size<remove_cvref_t<decltype(init)>>::value > 0, "Not enough values to fill with!");
r = get<0>(init);
return remove<0>(init);
}
CUTE_GCC_UNREACHABLE;
});
}
/** Make a "Tuple" by filling in the dynamic values in order from the arguments
* \code
* using result_t = cute::tuple<Int<1>, cute::tuple<int, int, Int<3>>, int, Int<2>>;
* auto result = make_int_tuple_from<result_t>(6,3,4); // (_1,(6,3,_3),4,_2)
* \endcode
*/
template <class Tuple, class... Ts>
CUTE_HOST_DEVICE constexpr
Tuple
make_int_tuple_from(Ts const&... ts)
{
Tuple result = Tuple{};
fill_int_tuple_from(result, cute::make_tuple(ts...));
return result;
}
/** Convert a tuple to a flat homogeneous array of type T
* \code
* auto tup = cute::make_tuple(Int<1>{}, cute::make_tuple(6,3,Int<3>{}),4,Int<2>{});
* cute::array<uint64_t,6> result = to_array<uint64_t>(tup); // [1,6,3,3,4,2]
* \endcode
*/
template <class T = int64_t, class IntTuple>
CUTE_HOST_DEVICE constexpr
auto
to_array(IntTuple const& t)
{
auto flat_t = flatten_to_tuple(t);
constexpr int N = tuple_size<decltype(flat_t)>::value;
cute::array<T,N> result;
for_each(make_seq<N>{}, [&] (auto i) { result[i] = get<i>(flat_t); });
return result;
}
//
// Comparison operators
//
//
// There are many ways to compare tuple of elements and because CuTe is built
// on parameterizing layouts of coordinates, some comparisons are appropriate
// only in certain cases.
// -- lexicographical comparison [reverse, reflected, revref] : Correct for coords in RowMajor Layout
// -- colexicographical comparison [reverse, reflected, revref] : Correct for coords in ColMajor Layout
// -- element-wise comparison [any,all] :
// This can be very confusing. To avoid errors in selecting the appropriate
// comparison, op<|op<=|op>|op>= are *not* implemented for cute::tuple.
//
// When actually desiring to order coordinates, the user should map them to
// their indices within the Layout they came from:
// e.g. layoutX(coordA) < layoutX(coordB)
// That said, we implement the three most common ways to compare tuples below.
// These are implemented with slighly more explicit names than op<.
//
template <class IntTupleA, class IntTupleB>
CUTE_HOST_DEVICE constexpr
auto
lex_less(IntTupleA const& a, IntTupleB const& b);
template <class IntTupleA, class IntTupleB>
CUTE_HOST_DEVICE constexpr
auto
colex_less(IntTupleA const& a, IntTupleB const& b);
template <class IntTupleA, class IntTupleB>
CUTE_HOST_DEVICE constexpr
auto
elem_less(IntTupleA const& a, IntTupleB const& b);
namespace detail {
template <size_t I, class TupleA, class TupleB>
CUTE_HOST_DEVICE constexpr
auto
lex_less_impl(TupleA const& a, TupleB const& b)
{
if constexpr (I == tuple_size<TupleB>::value) {
return cute::false_type{}; // Terminal: TupleB is exhausted
} else if constexpr (I == tuple_size<TupleA>::value) {
return cute::true_type{}; // Terminal: TupleA is exhausted, TupleB is not exhausted
} else {
return lex_less(get<I>(a), get<I>(b)) || (get<I>(a) == get<I>(b) && lex_less_impl<I+1>(a,b));
}
CUTE_GCC_UNREACHABLE;
}
template <size_t I, class TupleA, class TupleB>
CUTE_HOST_DEVICE constexpr
auto
colex_less_impl(TupleA const& a, TupleB const& b)
{
if constexpr (I == tuple_size<TupleB>::value) {
return cute::false_type{}; // Terminal: TupleB is exhausted
} else if constexpr (I == tuple_size<TupleA>::value) {
return cute::true_type{}; // Terminal: TupleA is exhausted, TupleB is not exhausted
} else {
constexpr size_t A = tuple_size<TupleA>::value - 1 - I;
constexpr size_t B = tuple_size<TupleB>::value - 1 - I;
return colex_less(get<A>(a), get<B>(b)) || (get<A>(a) == get<B>(b) && colex_less_impl<I+1>(a,b));
}
CUTE_GCC_UNREACHABLE;
}
template <size_t I, class TupleA, class TupleB>
CUTE_HOST_DEVICE constexpr
auto
elem_less_impl(TupleA const& a, TupleB const& b)
{
if constexpr (I == tuple_size<TupleA>::value) {
return cute::true_type{}; // Terminal: TupleA is exhausted
} else if constexpr (I == tuple_size<TupleB>::value) {
return cute::false_type{}; // Terminal: TupleA is not exhausted, TupleB is exhausted
} else {
return elem_less(get<I>(a), get<I>(b)) && elem_less_impl<I+1>(a,b);
}
CUTE_GCC_UNREACHABLE;
}
} // end namespace detail
// Lexicographical comparison
template <class IntTupleA, class IntTupleB>
CUTE_HOST_DEVICE constexpr
auto
lex_less(IntTupleA const& a, IntTupleB const& b)
{
if constexpr (is_tuple<IntTupleA>::value && is_tuple<IntTupleB>::value) {
return detail::lex_less_impl<0>(a, b);
} else {
return a < b;
}
CUTE_GCC_UNREACHABLE;
}
template <class T, class U>
CUTE_HOST_DEVICE constexpr
auto
lex_leq(T const& t, U const& u) {
return !lex_less(u, t);
}
template <class T, class U>
CUTE_HOST_DEVICE constexpr
auto
lex_gtr(T const& t, U const& u) {
return lex_less(u, t);
}
template <class T, class U>
CUTE_HOST_DEVICE constexpr
auto
lex_geq(T const& t, U const& u) {
return !lex_less(t, u);
}
// Colexicographical comparison
template <class IntTupleA, class IntTupleB>
CUTE_HOST_DEVICE constexpr
auto
colex_less(IntTupleA const& a, IntTupleB const& b)
{
if constexpr (is_tuple<IntTupleA>::value && is_tuple<IntTupleB>::value) {
return detail::colex_less_impl<0>(a, b);
} else {
return a < b;
}
CUTE_GCC_UNREACHABLE;
}
template <class T, class U>
CUTE_HOST_DEVICE constexpr
auto
colex_leq(T const& t, U const& u) {
return !colex_less(u, t);
}
template <class T, class U>
CUTE_HOST_DEVICE constexpr
auto
colex_gtr(T const& t, U const& u) {
return colex_less(u, t);
}
template <class T, class U>
CUTE_HOST_DEVICE constexpr
auto
colex_geq(T const& t, U const& u) {
return !colex_less(t, u);
}
// Elementwise [all] comparison
template <class IntTupleA, class IntTupleB>
CUTE_HOST_DEVICE constexpr
auto
elem_less(IntTupleA const& a, IntTupleB const& b)
{
if constexpr (is_tuple<IntTupleA>::value && is_tuple<IntTupleB>::value) {
return detail::elem_less_impl<0>(a, b);
} else {
return a < b;
}
CUTE_GCC_UNREACHABLE;
}
template <class T, class U>
CUTE_HOST_DEVICE constexpr
auto
elem_leq(T const& t, U const& u) {
return !elem_less(u, t);
}
template <class T, class U>
CUTE_HOST_DEVICE constexpr
auto
elem_gtr(T const& t, U const& u) {
return elem_less(u, t);
}
template <class T, class U>
CUTE_HOST_DEVICE constexpr
auto
elem_geq(T const& t, U const& u) {
return !elem_less(t, u);
}
/** Increment a (dynamic) coord lexicographically within a shape
* \code
* auto shape = make_shape(1,2,make_shape(2,3),3);
*
* int i = 0;
* for (auto coord = repeat_like(shape, 0); back(coord) != back(shape); increment(coord, shape)) {
* std::cout << i++ << ": " << coord << std::endl;
* }
* assert(i == size(shape));
* \endcode
*/
template <class Coord, class Shape>
CUTE_HOST_DEVICE constexpr
void
increment(Coord& coord, Shape const& shape);
namespace detail {
template <class Coord, class Shape, int I0, int... Is>
CUTE_HOST_DEVICE constexpr
void
increment(Coord& coord, Shape const& shape, seq<I0,Is...>)
{
cute::increment(get<I0>(coord), get<I0>(shape));
if constexpr (sizeof...(Is) != 0) {
if (back(get<I0>(coord)) == back(get<I0>(shape))) {
back(get<I0>(coord)) = 0;
increment(coord, shape, seq<Is...>{});
}
}
}
} // end namespace detail
template <class Coord, class Shape>
CUTE_HOST_DEVICE constexpr
void
increment(Coord& coord, Shape const& shape)
{
if constexpr (is_integral<Coord>::value && is_integral<Shape>::value) {
++coord;
} else if constexpr (is_tuple<Coord>::value && is_tuple<Shape>::value) {
static_assert(tuple_size<Coord>::value == tuple_size<Shape>::value, "Mismatched ranks");
detail::increment(coord, shape, tuple_seq<Coord>{});
} else {
static_assert(sizeof(Coord) == 0, "Invalid parameters");
}
}
struct ForwardCoordIteratorSentinal
{};
// A forward iterator for a coordinate that starts from zero and goes to shape
template <class Coord, class Shape>
struct ForwardCoordIterator
{
static_assert(is_congruent<Coord, Shape>::value);
CUTE_HOST_DEVICE constexpr
Coord const& operator*() const { return coord; }
CUTE_HOST_DEVICE constexpr
ForwardCoordIterator& operator++() { increment(coord, shape); return *this; }
// Sentinal for the end of the implied range
CUTE_HOST_DEVICE constexpr
bool operator< (ForwardCoordIteratorSentinal const&) const { return back(coord) < back(shape); }
CUTE_HOST_DEVICE constexpr
bool operator==(ForwardCoordIteratorSentinal const&) const { return back(coord) == back(shape); }
CUTE_HOST_DEVICE constexpr
bool operator!=(ForwardCoordIteratorSentinal const&) const { return back(coord) != back(shape); }
// NOTE: These are expensive, avoid use
CUTE_HOST_DEVICE constexpr
bool operator< (ForwardCoordIterator const& other) const { return colex_less(coord, other.coord); }
CUTE_HOST_DEVICE constexpr
bool operator==(ForwardCoordIterator const& other) const { return coord == other.coord; }
CUTE_HOST_DEVICE constexpr
bool operator!=(ForwardCoordIterator const& other) const { return coord != other.coord; }
Coord coord;
Shape const& shape;
};
// A forward iterator for a coordinate that starts from zero
template <class Shape>
CUTE_HOST_DEVICE constexpr
auto
make_coord_iterator(Shape const& shape)
{
auto coord = repeat_like(shape, int(0));
return ForwardCoordIterator<decltype(coord),Shape>{coord,shape};
}
} // end namespace cute

File diff suppressed because it is too large Load Diff

Some files were not shown because too many files have changed in this diff Show More