[Public release 26/04] Introducing Mega MoE, FP4 Indexer and other features/fixes (#304)
* Merge with private repo * Update README * Update README * Update README * Add PyTorch requirements * Fix sync scopes for MQA logits (#256) * Update README
This commit is contained in:
@@ -8,7 +8,8 @@ from deep_gemm.utils import (
|
||||
align, ceil_div,
|
||||
per_token_cast_to_fp8, per_channel_cast_to_fp8, per_block_cast_to_fp8,
|
||||
per_token_cast_to_fp4, transpose_packed_fp4,
|
||||
get_mk_alignment_for_contiguous_layout
|
||||
get_mk_alignment_for_contiguous_layout,
|
||||
set_mk_alignment_for_contiguous_layout
|
||||
)
|
||||
|
||||
|
||||
@@ -107,7 +108,7 @@ def get_major_ab(allow_a_mn_major: bool, allow_b_mn_major: bool) -> Generator:
|
||||
|
||||
|
||||
def get_psum_layout_usage() -> tuple:
|
||||
return (False, True) if get_arch_major() == 10 else (False, )
|
||||
return True, False
|
||||
|
||||
|
||||
def enumerate_normal(dtype: torch.dtype) -> Generator:
|
||||
@@ -168,7 +169,7 @@ def enumerate_m_grouped_contiguous(dtype: torch.dtype) -> Generator:
|
||||
def enumerate_m_grouped_masked(dtype: torch.dtype) -> Generator:
|
||||
quant_config_list = QuantConfig.get_list_from_dtype(dtype)
|
||||
max_m = 4096
|
||||
m_group_list = [(6, 1024), (32, 192), (32, 50)]
|
||||
m_group_list = [(32, 192), (6, 1024), (32, 20), (6, 20)]
|
||||
n_k_list = [(6144, 7168), (7168, 3072), (4096, 4096), (4096, 2048)]
|
||||
for kernel_type in get_kernel_types(dtype):
|
||||
for quant_config in quant_config_list:
|
||||
@@ -182,6 +183,7 @@ def enumerate_m_grouped_masked(dtype: torch.dtype) -> Generator:
|
||||
|
||||
|
||||
def enumerate_k_grouped_contiguous(dtype: torch.dtype):
|
||||
gran_k_list = (128, ) if get_arch_major() == 9 else (32, 128)
|
||||
# Only K-major is supported for SM90 FP8
|
||||
major_a, major_b = (MajorTypeAB.KMajor, MajorTypeAB.KMajor) if get_arch_major() == 9 and dtype == torch.float8_e4m3fn \
|
||||
else (MajorTypeAB.MNMajor, MajorTypeAB.MNMajor)
|
||||
@@ -189,26 +191,36 @@ def enumerate_k_grouped_contiguous(dtype: torch.dtype):
|
||||
for num_groups, m, n, expected_k_per_group in (( 4, 4096, 7168, 8192), ( 4, 7168, 2048, 8192), # EP64
|
||||
( 8, 4096, 7168, 4096), ( 8, 7168, 2048, 4096), # EP32
|
||||
(16, 4096, 7168, 2048), (16, 7168, 2048, 2048)): # EP16
|
||||
ks = [align(int(expected_k_per_group * random.uniform(0.7, 1.3)), get_mk_alignment_for_contiguous_layout()) for _ in range(num_groups)]
|
||||
yield num_groups, m, n, major_a, major_b, ks, expected_k_per_group
|
||||
if dtype == torch.bfloat16:
|
||||
ks = [align(int(expected_k_per_group * random.uniform(0.7, 1.3)), get_mk_alignment_for_contiguous_layout()) for _ in range(num_groups)]
|
||||
yield num_groups, m, n, major_a, major_b, ks, expected_k_per_group
|
||||
else:
|
||||
for gran_k in gran_k_list:
|
||||
set_mk_alignment_for_contiguous_layout(gran_k)
|
||||
ks = [align(int(expected_k_per_group * random.uniform(0.7, 1.3)), gran_k) for _ in range(num_groups)]
|
||||
yield num_groups, m, n, major_a, major_b, ks, expected_k_per_group, gran_k
|
||||
|
||||
|
||||
def enumerate_sf_layout():
|
||||
gran_k_list = (128, ) if get_arch_major() == 9 else (32, 128)
|
||||
for use_ue8m0 in (False, True):
|
||||
for with_transpose in (True, False):
|
||||
for mn in (4096, 4097, 8192):
|
||||
for k in (128, 7168, 7296):
|
||||
for num_groups in (1, 2, 4):
|
||||
yield mn, k, with_transpose, use_ue8m0, num_groups
|
||||
for gran_k in gran_k_list:
|
||||
set_mk_alignment_for_contiguous_layout(gran_k)
|
||||
yield mn, k, with_transpose, use_ue8m0, num_groups, gran_k
|
||||
|
||||
|
||||
def enumerate_k_grouped_sf_layout():
|
||||
alignment = get_mk_alignment_for_contiguous_layout()
|
||||
assert alignment % 128 == 0
|
||||
gran_k_list = (128, ) if get_arch_major() == 9 else (32, 128)
|
||||
for mn in (4096, 7168):
|
||||
for num_groups, avg_k in ((16, 2048), (8, 4096), (72, 384), (128, 256)):
|
||||
ks = [align(int(random.uniform(0.7, 1.3) * avg_k), alignment) for _ in range(num_groups)]
|
||||
yield mn, ks, num_groups
|
||||
for gran_k in gran_k_list:
|
||||
set_mk_alignment_for_contiguous_layout(gran_k)
|
||||
ks = [align(int(random.uniform(0.7, 1.3) * avg_k), gran_k) for _ in range(num_groups)]
|
||||
yield mn, ks, num_groups, gran_k
|
||||
|
||||
|
||||
def enumerate_transpose():
|
||||
@@ -222,25 +234,24 @@ def cast_fp8_fp4_with_major(x: torch.Tensor, major: MajorTypeAB, gran_k: int, is
|
||||
use_ue8m0: bool, use_block_cast_for_fp8: bool = False):
|
||||
if is_fp4:
|
||||
x_fp4 = per_token_cast_to_fp4(x, use_ue8m0=use_ue8m0, gran_k=gran_k)
|
||||
x = x_fp4 if major.is_k_major() else (transpose_packed_fp4(x_fp4[0]).T, x_fp4[1])
|
||||
return x_fp4 if major.is_k_major() else (transpose_packed_fp4(x_fp4[0]).T, x_fp4[1])
|
||||
else:
|
||||
x_fp8 = per_block_cast_to_fp8(x, use_ue8m0=use_ue8m0, gran_k=gran_k) if use_block_cast_for_fp8 \
|
||||
else per_token_cast_to_fp8(x, use_ue8m0=use_ue8m0, gran_k=gran_k)
|
||||
x = x_fp8 if major.is_k_major() else (x_fp8[0].T.contiguous().T, x_fp8[1])
|
||||
return x
|
||||
return x_fp8 if major.is_k_major() else (x_fp8[0].T.contiguous().T, x_fp8[1])
|
||||
|
||||
|
||||
def grouped_cast_fp8_fp4_with_major(x: torch.Tensor, major: MajorTypeAB, gran_k: int, is_fp4: bool,
|
||||
use_ue8m0: bool, use_block_cast_for_fp8: bool = False):
|
||||
num_groups, mn, k = x.size()
|
||||
if is_fp4:
|
||||
x_fp4 = (torch.empty((num_groups, mn, k // 2), device='cuda', dtype=torch.uint8) if major.is_k_major() else \
|
||||
torch.empty((num_groups, k, mn // 2), device='cuda', dtype=torch.uint8),
|
||||
x_fp4 = (torch.empty((num_groups, mn, k // 2), device='cuda', dtype=torch.int8) if major.is_k_major() else \
|
||||
torch.empty((num_groups, k, mn // 2), device='cuda', dtype=torch.int8),
|
||||
torch.empty((num_groups, mn, ceil_div(k, gran_k)), device='cuda', dtype=torch.float))
|
||||
for i in range(num_groups):
|
||||
x_i_fp4 = per_token_cast_to_fp4(x[i], use_ue8m0=use_ue8m0, gran_k=gran_k)
|
||||
x_fp4[0][i], x_fp4[1][i] = x_i_fp4 if major.is_k_major() else (transpose_packed_fp4(x_i_fp4[0]), x_i_fp4[1])
|
||||
x = x_fp4 if major.is_k_major() else (x_fp4[0].mT, x_fp4[1])
|
||||
return x_fp4 if major.is_k_major() else (x_fp4[0].mT, x_fp4[1])
|
||||
else:
|
||||
x_fp8 = (torch.empty_like(x, dtype=torch.float8_e4m3fn),
|
||||
torch.empty((num_groups, ceil_div(mn, gran_k), ceil_div(k, gran_k)), device='cuda', dtype=torch.float) if use_block_cast_for_fp8 \
|
||||
@@ -248,8 +259,7 @@ def grouped_cast_fp8_fp4_with_major(x: torch.Tensor, major: MajorTypeAB, gran_k:
|
||||
for i in range(num_groups):
|
||||
x_fp8[0][i], x_fp8[1][i] = per_block_cast_to_fp8(x[i], use_ue8m0=use_ue8m0, gran_k=gran_k) if use_block_cast_for_fp8 \
|
||||
else per_token_cast_to_fp8(x[i], use_ue8m0=use_ue8m0, gran_k=gran_k)
|
||||
x = x_fp8 if major.is_k_major() else (x_fp8[0].mT.contiguous().mT, x_fp8[1])
|
||||
return x
|
||||
return x_fp8 if major.is_k_major() else (x_fp8[0].mT.contiguous().mT, x_fp8[1])
|
||||
|
||||
|
||||
def generate_normal(m: int, n: int, k: int,
|
||||
@@ -325,7 +335,7 @@ def layout_masked_to_psum(x: torch.Tensor, psum_m: torch.Tensor):
|
||||
last_psum_m = 0
|
||||
for i in range(num_groups):
|
||||
x_psum[last_psum_m: psum_m[i]] = x[i, :psum_m[i] - last_psum_m]
|
||||
last_psum_m = align(psum_m[i], 128)
|
||||
last_psum_m = align(psum_m[i], get_mk_alignment_for_contiguous_layout())
|
||||
return x_psum
|
||||
|
||||
|
||||
@@ -342,7 +352,7 @@ def generate_m_grouped_masked(num_groups: int, max_m: int, expected_m_per_group:
|
||||
psum_m = torch.empty((num_groups, ), device='cuda', dtype=torch.int)
|
||||
for j in range(num_groups):
|
||||
masked_m[j] = int(expected_m_per_group * random.uniform(0.7, 1.3))
|
||||
psum_m[j] = (0 if j == 0 else align(psum_m[j - 1], 128)) + masked_m[j]
|
||||
psum_m[j] = (0 if j == 0 else align(psum_m[j - 1], get_mk_alignment_for_contiguous_layout())) + masked_m[j]
|
||||
assert masked_m.amax().item() <= max_m
|
||||
|
||||
if use_bf16:
|
||||
@@ -356,8 +366,8 @@ def generate_m_grouped_masked(num_groups: int, max_m: int, expected_m_per_group:
|
||||
|
||||
|
||||
def generate_k_grouped_contiguous(num_groups: int, m: int, n: int, major_a: MajorTypeAB, major_b: MajorTypeAB, ks: List[int],
|
||||
use_ue8m0: bool = False, use_bf16: bool = False):
|
||||
assert get_mk_alignment_for_contiguous_layout() % 128 == 0
|
||||
use_ue8m0: bool = False, use_bf16: bool = False, gran_k = 128):
|
||||
assert get_mk_alignment_for_contiguous_layout() % gran_k == 0
|
||||
k = sum(ks)
|
||||
|
||||
a = torch.randn((k, m), device='cuda', dtype=torch.bfloat16)
|
||||
@@ -376,8 +386,8 @@ def generate_k_grouped_contiguous(num_groups: int, m: int, n: int, major_a: Majo
|
||||
assert (major_a, major_b) == (MajorTypeAB.MNMajor, MajorTypeAB.MNMajor)
|
||||
return k, a, b, c, d, ref_d
|
||||
|
||||
a_fp8 = per_channel_cast_to_fp8(a, use_ue8m0=use_ue8m0)
|
||||
b_fp8 = per_channel_cast_to_fp8(b, use_ue8m0=use_ue8m0)
|
||||
a_fp8 = per_channel_cast_to_fp8(a, use_ue8m0=use_ue8m0, gran_k=gran_k)
|
||||
b_fp8 = per_channel_cast_to_fp8(b, use_ue8m0=use_ue8m0, gran_k=gran_k)
|
||||
|
||||
# Transpose for K Major A/B
|
||||
if (major_a, major_b) == (MajorTypeAB.KMajor, MajorTypeAB.KMajor):
|
||||
|
||||
@@ -10,9 +10,9 @@ from deep_gemm.testing import (
|
||||
ignore_env, get_arch_major,
|
||||
test_filter
|
||||
)
|
||||
from deep_gemm.utils import ceil_div, per_custom_dims_cast_to_fp8
|
||||
from deep_gemm.utils import ceil_div, per_custom_dims_cast_to_fp8, per_token_cast_to_fp4, cast_back_from_fp4
|
||||
|
||||
from generators import generate_normal, get_ue8m0_usage, get_kernel_types, MajorTypeAB
|
||||
from generators import get_arch_major, generate_normal, get_ue8m0_usage, get_kernel_types, reset_seed, MajorTypeAB
|
||||
|
||||
|
||||
def apply_skip_head_mid(d: torch.Tensor, head_splits: Tuple[int, int, int]):
|
||||
@@ -53,40 +53,14 @@ def test_gemm_skip_head_mid() -> None:
|
||||
assert diff < 0.001, f'{m=}, {n=}, {k=}, {kernel_opt}, {diff:.5f}'
|
||||
|
||||
t = bench_kineto(lambda: deep_gemm.fp8_gemm_nt_skip_head_mid(a, b, d, head_splits, disable_ue8m0_cast=disable_ue8m0_cast),
|
||||
'fp8_gemm', suppress_kineto_output=True)
|
||||
'gemm_', suppress_kineto_output=True)
|
||||
print(f' > Perf (m={m:5}, n={n:5}, k={k:5}, {kernel_opt}): '
|
||||
f'{t * 1e6:4.0f} us | '
|
||||
f'{2 * m * n * k / t / 1e12:4.0f} TFLOPS | '
|
||||
f'{(count_bytes(a, b, d)) / 1e9 / t:4.0f} GB/s')
|
||||
f'{t * 1e6:4.0f} us | '
|
||||
f'{2 * m * n * k / t / 1e12:4.0f} TFLOPS | '
|
||||
f'{(count_bytes(a, b, d)) / 1e9 / t:4.0f} GB/s')
|
||||
print()
|
||||
|
||||
|
||||
def kv_cache_cast_to_fp8(x: torch.Tensor) -> torch.Tensor:
|
||||
num_blocks, block_size, num_heads, head_dim = x.shape
|
||||
assert num_heads == 1
|
||||
x_amax = x.abs().float().amax(dim=3, keepdim=True).clamp(1e-4)
|
||||
sf = x_amax / 448.0
|
||||
x_scaled = (x * (1.0 / sf)).to(torch.float8_e4m3fn)
|
||||
x_fp8 = torch.empty((num_blocks, block_size * (head_dim + 4)), device=x.device, dtype=torch.uint8)
|
||||
x_fp8[ :, : block_size * head_dim] = x_scaled.view(num_blocks, block_size * head_dim).view(dtype=torch.uint8)
|
||||
x_fp8[ :, block_size * head_dim :] = sf.view(num_blocks, block_size).view(dtype=torch.uint8)
|
||||
return x_fp8.view(num_blocks, block_size, num_heads, head_dim + 4)
|
||||
|
||||
|
||||
def generate_cp_test_data(seq_len, seq_len_kv):
|
||||
assert seq_len_kv % seq_len == 0 and seq_len % 2 == 0
|
||||
chunk_size = seq_len // 2
|
||||
cp_size = seq_len_kv // seq_len
|
||||
# Select an arbitrary CP rank
|
||||
cp_id = cp_size // 3
|
||||
ks = torch.zeros(seq_len, dtype=torch.int, device='cuda')
|
||||
ke = torch.zeros(seq_len, dtype=torch.int, device='cuda')
|
||||
for i in range(chunk_size):
|
||||
ke[i] = cp_id * chunk_size + i
|
||||
ke[i + chunk_size] = (cp_size * 2 - 1 - cp_id) * chunk_size + i
|
||||
return ks, ke
|
||||
|
||||
|
||||
def ref_fp8_mqa_logits(q: torch.Tensor, kv: torch.Tensor, weights: torch.Tensor,
|
||||
cu_seqlen_ks: torch.Tensor, cu_seqlen_ke: torch.Tensor, cost_only: bool = False):
|
||||
seq_len_kv = kv.shape[0]
|
||||
@@ -113,92 +87,137 @@ def ref_fp8_mqa_logits(q: torch.Tensor, kv: torch.Tensor, weights: torch.Tensor,
|
||||
return logits, cost
|
||||
|
||||
|
||||
@ignore_env('DG_JIT_PTXAS_CHECK', lambda: get_arch_major() == 10)
|
||||
def test_mqa_logits():
|
||||
|
||||
# Helper functions
|
||||
def generate_ks_ke_tests(seq_len: int, seq_len_kv: int, disable_cp: bool):
|
||||
if disable_cp:
|
||||
ks = torch.zeros(seq_len, dtype=torch.int, device='cuda')
|
||||
ke = torch.arange(seq_len, dtype=torch.int, device='cuda') + (seq_len_kv - seq_len)
|
||||
return ks, ke
|
||||
assert seq_len_kv % seq_len == 0 and seq_len % 2 == 0
|
||||
chunk_size = seq_len // 2
|
||||
cp_size = seq_len_kv // seq_len
|
||||
# Select an arbitrary CP rank
|
||||
cp_id = cp_size // 3
|
||||
ks = torch.zeros(seq_len, dtype=torch.int, device='cuda')
|
||||
ke = torch.zeros(seq_len, dtype=torch.int, device='cuda')
|
||||
for i in range(chunk_size):
|
||||
ke[i] = cp_id * chunk_size + i
|
||||
ke[i + chunk_size] = (cp_size * 2 - 1 - cp_id) * chunk_size + i
|
||||
return ks, ke
|
||||
|
||||
def enumerate_mqa_logits():
|
||||
for is_fp4 in ((True, False) if get_arch_major() == 10 else (False, )):
|
||||
for logits_dtype in (torch.float, torch.bfloat16):
|
||||
for compressed_logits, clean_logits in [(False, True), (True, False)]:
|
||||
for seq_len in (2048, 4096):
|
||||
for seq_len_kv in (4096, 8192):
|
||||
for num_heads, head_dim in [(64, 128)]:
|
||||
for disable_cp in (False, True):
|
||||
yield is_fp4, logits_dtype, compressed_logits, clean_logits, seq_len, seq_len_kv, num_heads, head_dim, disable_cp
|
||||
|
||||
print('Testing FP8 MQA Logits:')
|
||||
num_heads, head_dim = 64, 128
|
||||
for seq_len in (2048, 4096):
|
||||
for compressed_logits in (False, True):
|
||||
for seq_len_kv in (4096, 8192):
|
||||
for disable_cp in (False, True):
|
||||
q = torch.randn(seq_len, num_heads, head_dim, device='cuda', dtype=torch.bfloat16)
|
||||
kv = torch.randn(seq_len_kv, head_dim, device='cuda', dtype=torch.bfloat16)
|
||||
weights = torch.randn(seq_len, num_heads, device='cuda', dtype=torch.float32)
|
||||
for is_fp4, logits_dtype, compressed_logits, clean_logits, seq_len, seq_len_kv, num_heads, head_dim, disable_cp in enumerate_mqa_logits():
|
||||
# Generate random inputs
|
||||
q = torch.randn(seq_len, num_heads, head_dim, device='cuda', dtype=torch.bfloat16)
|
||||
kv = torch.randn(seq_len_kv, head_dim, device='cuda', dtype=torch.bfloat16)
|
||||
weights = torch.randn(seq_len, num_heads, device='cuda', dtype=torch.float32)
|
||||
ks, ke = generate_ks_ke_tests(seq_len, seq_len_kv, disable_cp)
|
||||
|
||||
if disable_cp:
|
||||
ks = torch.zeros(seq_len, dtype=torch.int, device='cuda')
|
||||
ke = torch.arange(seq_len, dtype=torch.int, device='cuda') + (seq_len_kv - seq_len)
|
||||
else:
|
||||
ks, ke = generate_cp_test_data(seq_len, seq_len_kv)
|
||||
# Calculate reference logits
|
||||
ref_logits, ref_cost = ref_fp8_mqa_logits(q, kv, weights, ks, ke)
|
||||
|
||||
q_fp8 = q.to(torch.float8_e4m3fn)
|
||||
kv_fp8 = per_custom_dims_cast_to_fp8(kv, (0, ), False)
|
||||
# Quantize Q and KV to FP4 / FP8
|
||||
if is_fp4:
|
||||
q_fp4 = per_token_cast_to_fp4(q.view(-1, head_dim), use_ue8m0=True, gran_k=32, use_packed_ue8m0=True)
|
||||
q_in = (q_fp4[0].view(seq_len, num_heads, head_dim // 2), q_fp4[1].view(seq_len, num_heads))
|
||||
q_simulated = cast_back_from_fp4(q_fp4[0], q_fp4[1], gran_k=32, use_packed_ue8m0=True).view(seq_len, num_heads, head_dim).to(torch.bfloat16)
|
||||
|
||||
if compressed_logits:
|
||||
max_seqlen_k = (ke - ks).max().item()
|
||||
logits = deep_gemm.fp8_mqa_logits(q_fp8, kv_fp8, weights, ks, ke, max_seqlen_k=max_seqlen_k, clean_logits=False)
|
||||
assert logits.size() == (seq_len, max_seqlen_k)
|
||||
tmp = torch.full((seq_len, seq_len_kv), float('-inf'), device='cuda')
|
||||
for i in range(seq_len):
|
||||
tmp[i, ks[i] : ke[i]] = logits[i, : ke[i] - ks[i]]
|
||||
logits = tmp
|
||||
else:
|
||||
logits = deep_gemm.fp8_mqa_logits(q_fp8, kv_fp8, weights, ks, ke)
|
||||
kv_fp4 = per_token_cast_to_fp4(kv.view(-1, head_dim), use_ue8m0=True, gran_k=32, use_packed_ue8m0=True)
|
||||
kv_in = (kv_fp4[0].view(seq_len_kv, head_dim // 2), kv_fp4[1].view(seq_len_kv))
|
||||
kv_simulated = cast_back_from_fp4(kv_fp4[0], kv_fp4[1], gran_k=32, use_packed_ue8m0=True).view(seq_len_kv, head_dim).to(torch.bfloat16)
|
||||
else:
|
||||
q_in = q.to(torch.float8_e4m3fn), None
|
||||
q_simulated = q_in[0].to(torch.bfloat16)
|
||||
kv_in = per_custom_dims_cast_to_fp8(kv, (0, ), False)
|
||||
kv_simulated = (kv_in[0].float() * kv_in[1].unsqueeze(1)).to(torch.bfloat16)
|
||||
|
||||
do_check = (seq_len_kv < 32768)
|
||||
if do_check:
|
||||
ref_logits, ref_cost = ref_fp8_mqa_logits(q=q, kv=kv, weights=weights, cu_seqlen_ks=ks, cu_seqlen_ke=ke)
|
||||
# Calculate reference logits
|
||||
simulated_logits, _ = ref_fp8_mqa_logits(q_simulated, kv_simulated, weights, ks, ke)
|
||||
|
||||
ref_neginf_mask = (ref_logits == float('-inf'))
|
||||
neginf_mask = (logits == float('-inf'))
|
||||
assert torch.equal(neginf_mask, ref_neginf_mask)
|
||||
# Prepare kwargs
|
||||
kernel_kwargs = dict(
|
||||
q=q_in, kv=kv_in, weights=weights,
|
||||
cu_seq_len_k_start=ks, cu_seq_len_k_end=ke,
|
||||
clean_logits=clean_logits, max_seqlen_k=0,
|
||||
logits_dtype=logits_dtype
|
||||
)
|
||||
if compressed_logits:
|
||||
max_seqlen_k = (ke - ks).max().item()
|
||||
kernel_kwargs['max_seqlen_k'] = max_seqlen_k
|
||||
|
||||
ref_logits = ref_logits.masked_fill(ref_neginf_mask, 0)
|
||||
logits = logits.masked_fill(neginf_mask, 0)
|
||||
diff = calc_diff(logits, ref_logits)
|
||||
assert diff < 1e-3, f'{diff=}'
|
||||
else:
|
||||
ref_cost = ref_fp8_mqa_logits(q=q, kv=kv, weights=weights, cu_seqlen_ks=ks, cu_seqlen_ke=ke, cost_only=True)
|
||||
# Run kernel
|
||||
logits = deep_gemm.fp8_fp4_mqa_logits(**kernel_kwargs)
|
||||
|
||||
tflops = 2 * ref_cost * num_heads * head_dim / 1e12
|
||||
if compressed_logits:
|
||||
t = bench_kineto(lambda: deep_gemm.fp8_mqa_logits(q_fp8, kv_fp8, weights, ks, ke, max_seqlen_k=max_seqlen_k, clean_logits=False), 'fp8_mqa_logits')
|
||||
else:
|
||||
t, clean_t = bench_kineto(lambda: deep_gemm.fp8_mqa_logits(q_fp8, kv_fp8, weights, ks, ke), ('fp8_mqa_logits', 'clean_logits'))
|
||||
clean_bytes = (seq_len * seq_len_kv - ref_cost) * 4 + count_bytes(ks, ke)
|
||||
print(f' > S={seq_len:4}, SKV={seq_len_kv:6}, H={num_heads:3}, D={head_dim:3}, CP={0 if disable_cp else 1}: '
|
||||
f'{tflops / t:4.0f} TFLOPS, {t * 1e6:4.0f} us, '
|
||||
f'{(count_bytes(q_fp8, kv_fp8, weights, ks, ke) + ref_cost * 4) / t / 1e9:4.0f} GB/s', end='')
|
||||
# noinspection PyUnboundLocalVariable
|
||||
print(f' | clean: {clean_t * 1e6:3.0f} us, {clean_bytes / clean_t / 1e9:4.0f} GB/s' if not compressed_logits else '')
|
||||
# Post process for compressed logits
|
||||
if compressed_logits:
|
||||
assert logits.size() == (seq_len, max_seqlen_k)
|
||||
tmp = torch.full((seq_len, seq_len_kv), float('-inf'), device='cuda')
|
||||
for i in range(seq_len):
|
||||
tmp[i, ks[i] : ke[i]] = logits[i, : ke[i] - ks[i]]
|
||||
logits = tmp
|
||||
|
||||
# Validation
|
||||
ref_neginf_mask = (ref_logits == float('-inf'))
|
||||
neginf_mask = (logits == float('-inf'))
|
||||
assert torch.equal(neginf_mask, ref_neginf_mask)
|
||||
|
||||
ref_logits = ref_logits.masked_fill(ref_neginf_mask, 0)
|
||||
simulated_logits = simulated_logits.masked_fill(ref_neginf_mask, 0)
|
||||
logits = logits.masked_fill(ref_neginf_mask, 0)
|
||||
diff = calc_diff(logits, ref_logits)
|
||||
simulated_diff = calc_diff(logits, simulated_logits)
|
||||
assert diff < 0.02 if is_fp4 else 1e-3, f"Diff: {diff}"
|
||||
assert simulated_diff < 5e-6, f"Simulated Diff: {simulated_diff}"
|
||||
|
||||
# Profiling
|
||||
tflops = 2 * ref_cost * num_heads * head_dim / 1e12
|
||||
t, clean_t = bench_kineto(lambda: deep_gemm.fp8_fp4_mqa_logits(**kernel_kwargs), ('mqa_logits', 'clean_logits'))
|
||||
clean_bytes = (seq_len * seq_len_kv - ref_cost) * 4 + count_bytes(ks, ke)
|
||||
|
||||
print(f' > FP4={is_fp4}, BF16={logits_dtype == torch.bfloat16}, S={seq_len:4}, SKV={seq_len_kv:6}, H={num_heads:3}, D={head_dim:3}, CP={0 if disable_cp else 1}: '
|
||||
f'{tflops / t:4.0f} TFLOPS, {t * 1e6:4.0f} us, '
|
||||
f'{(count_bytes(q_in, kv_in, weights, ks, ke) + ref_cost * 4) / t / 1e9:4.0f} GB/s', end='')
|
||||
print(f' | clean: {clean_t * 1e6:3.0f} us, {clean_bytes / clean_t / 1e9:4.0f} GB/s' if clean_logits else '')
|
||||
print()
|
||||
|
||||
|
||||
def ref_fp8_paged_mqa_logits(q: torch.Tensor, kv_cache: torch.Tensor,
|
||||
weights: torch.Tensor, context_lens: torch.Tensor, block_tables: torch.Tensor,
|
||||
max_model_len: int, is_context_lens_2d: bool):
|
||||
batch_size, next_n, heads, dim = q.size()
|
||||
def ref_paged_mqa_logits(q: torch.Tensor, kv_cache: torch.Tensor,
|
||||
weights: torch.Tensor, context_lens: torch.Tensor, block_tables: torch.Tensor,
|
||||
max_model_len: int, use_2d_context_lens: bool):
|
||||
batch_size, next_n, num_heads, dim = q.size()
|
||||
num_block, block_size, _, dim = kv_cache.size()
|
||||
logits = torch.full([batch_size * next_n, max_model_len], float('-inf'), device=q.device, dtype=torch.float32)
|
||||
context_lens = context_lens.tolist()
|
||||
for i in range(batch_size):
|
||||
context_len = context_lens[i]
|
||||
q_offsets = torch.full((next_n, ), context_len, device='cuda', dtype=torch.int32) if is_context_lens_2d \
|
||||
else torch.arange(context_len - next_n, context_len, device='cuda')
|
||||
q_offsets = torch.full((next_n, ), context_len, device='cuda', dtype=torch.int32) if use_2d_context_lens \
|
||||
else torch.arange(context_len - next_n, context_len, device='cuda')
|
||||
weight_slice = weights[i * next_n:(i + 1) * next_n, :].transpose(0, 1).contiguous()
|
||||
|
||||
num_blocks = (context_len + block_size - 1) // block_size
|
||||
block_idxs = block_tables[i][:num_blocks]
|
||||
kv_slice = kv_cache[block_idxs] # [num_blocks, block_size, kv_heads, dim]
|
||||
kx = kv_slice.permute(2, 3, 0, 1).reshape(kv_slice.size(2), dim, -1) # [kv_heads, dim, total_tokens]
|
||||
qx = q[i].transpose(0, 1) # q[i]: [next_n, heads, dim] -> [heads, next_n, dim]
|
||||
s = torch.matmul(qx, kx).to(logits.dtype) # [heads, next_n, dim] @ [1, dim, total_tokens] -> [heads, next_n, total_tokens]
|
||||
qx = q[i].transpose(0, 1) # q[i]: [next_n, num_heads, dim] -> [num_heads, next_n, dim]
|
||||
s = torch.matmul(qx, kx).to(logits.dtype) # [num_heads, next_n, dim] @ [1, dim, total_tokens] -> [num_heads, next_n, total_tokens]
|
||||
|
||||
total_len = num_blocks * block_size
|
||||
k_offsets = torch.arange(0, total_len, device=q.device)
|
||||
mask = (k_offsets[None, :] < context_len) & (k_offsets[None, :] <= q_offsets[:, None])
|
||||
s = torch.where(mask[None, :, :], s, float('-inf')) # mask shape: [1, next_n, total_tokens]
|
||||
s = torch.relu(s) * weight_slice[..., None] # weight_slice: [heads, next_n] -> [heads, next_n, 1]
|
||||
s = torch.relu(s) * weight_slice[..., None] # weight_slice: [num_heads, next_n] -> [num_heads, next_n, 1]
|
||||
s = s.sum(dim=0) # [next_n, total_tokens]
|
||||
logits[i * next_n:(i + 1) * next_n, :total_len] = torch.where(k_offsets[None, :] <= q_offsets[:, None], s, float('-inf'))
|
||||
|
||||
@@ -206,70 +225,129 @@ def ref_fp8_paged_mqa_logits(q: torch.Tensor, kv_cache: torch.Tensor,
|
||||
|
||||
|
||||
def test_paged_mqa_logits():
|
||||
print('Testing FP8 Paged MQA Logits:')
|
||||
max_model_len = 111 * 1000
|
||||
for is_context_lens_2d in (False, True):
|
||||
for batch_size, next_n in [(64, 1), (64, 2), (128, 1)]:
|
||||
for heads, index_dim in [(64, 128)]:
|
||||
for avg_kv in (8192, 32768):
|
||||
num_blocks, blocksize = max_model_len * 3, 64
|
||||
|
||||
q = torch.randn((batch_size, next_n, heads, index_dim), device='cuda', dtype=torch.bfloat16)
|
||||
kv_cache = torch.randn((num_blocks, blocksize, 1, index_dim), device='cuda', dtype=torch.bfloat16)
|
||||
weights = torch.randn((batch_size * next_n, heads), device='cuda', dtype=torch.float32)
|
||||
q_fp8 = q.to(torch.float8_e4m3fn)
|
||||
kv_cache_fp8 = kv_cache_cast_to_fp8(kv_cache)
|
||||
# Helper functions
|
||||
def kv_cache_cast_to_fp8(x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
num_blocks, block_size, num_heads, head_dim = x.shape
|
||||
assert num_heads == 1
|
||||
x_amax = x.abs().float().amax(dim=3, keepdim=True).clamp(1e-4)
|
||||
sf = x_amax / 448.0
|
||||
x_scaled = (x * (1.0 / sf)).to(torch.float8_e4m3fn)
|
||||
x_cast_back = x_scaled.float() * sf
|
||||
|
||||
context_lens = torch.randint(int(0.7 * avg_kv), int(1.3 * avg_kv), (batch_size, )).cuda().to(torch.int32)
|
||||
context_lens_list = context_lens.tolist()
|
||||
max_block_len = (max(context_lens_list) + blocksize - 1) // blocksize * blocksize
|
||||
block_tables = torch.zeros((batch_size, max_block_len), device='cuda', dtype=torch.int32)
|
||||
x_fp8 = torch.empty((num_blocks, block_size * (head_dim + 4)), device=x.device, dtype=torch.uint8)
|
||||
x_fp8[ :, : block_size * head_dim] = x_scaled.view(num_blocks, block_size * head_dim).view(torch.uint8)
|
||||
x_fp8[ :, block_size * head_dim :] = sf.view(num_blocks, block_size).view(torch.uint8)
|
||||
return x_fp8.view(num_blocks, block_size, num_heads, head_dim + 4), x_cast_back.to(x.dtype)
|
||||
|
||||
counter, block_idx_pool = 0, torch.randperm(num_blocks, device='cuda', dtype=torch.int32)
|
||||
for i in range(batch_size):
|
||||
num_blocks = ceil_div(context_lens_list[i], blocksize)
|
||||
block_tables[i][:num_blocks] = block_idx_pool[counter: counter+num_blocks]
|
||||
counter += num_blocks
|
||||
def kv_cache_cast_to_fp4(x: torch.Tensor) -> torch.Tensor:
|
||||
num_blocks, block_size, num_heads, head_dim = x.shape
|
||||
assert num_heads == 1 and head_dim == 128
|
||||
x_scaled, sf = per_token_cast_to_fp4(x.view(-1, head_dim), use_ue8m0=True, gran_k=32, use_packed_ue8m0=True)
|
||||
x_cast_back = cast_back_from_fp4(x_scaled, sf, gran_k=32, use_packed_ue8m0=True).view(num_blocks, block_size, 1, head_dim)
|
||||
|
||||
ref_logits = ref_fp8_paged_mqa_logits(q, kv_cache, weights, context_lens, block_tables, max_model_len, is_context_lens_2d)
|
||||
positions = torch.arange(max_model_len, device='cuda').unsqueeze(0).expand(batch_size * next_n, -1)
|
||||
x_fp4 = torch.empty((num_blocks, block_size * (head_dim // 2 + 4)), device=x.device, dtype=torch.uint8)
|
||||
x_fp4[ :, : block_size * head_dim // 2] = x_scaled.view(num_blocks, block_size * head_dim // 2).view(torch.uint8)
|
||||
x_fp4[ :, block_size * head_dim // 2 :] = sf.view(num_blocks, block_size).view(torch.uint8)
|
||||
return x_fp4.view(num_blocks, block_size, num_heads, head_dim // 2 + 4), x_cast_back.to(x.dtype)
|
||||
|
||||
if is_context_lens_2d:
|
||||
context_lens_2d = ((context_lens.unsqueeze(1) + 1) * torch.rand(batch_size, next_n, device='cuda')).int()
|
||||
context_lens_2d[:, next_n-1] = context_lens
|
||||
schedule_metadata = deep_gemm.get_paged_mqa_logits_metadata(context_lens_2d, blocksize, deep_gemm.get_num_sms())
|
||||
logits = deep_gemm.fp8_paged_mqa_logits(q_fp8, kv_cache_fp8, weights, context_lens_2d, block_tables, schedule_metadata, max_model_len, clean_logits=False)
|
||||
ref_neginf_mask = ~(positions < context_lens_2d.view(-1).unsqueeze(1))
|
||||
else:
|
||||
schedule_metadata = deep_gemm.get_paged_mqa_logits_metadata(context_lens, blocksize, deep_gemm.get_num_sms())
|
||||
logits = deep_gemm.fp8_paged_mqa_logits(q_fp8, kv_cache_fp8, weights, context_lens, block_tables, schedule_metadata, max_model_len, clean_logits=True)
|
||||
row_indices = torch.arange(batch_size * next_n, device='cuda') // next_n
|
||||
next_n_offset = torch.arange(batch_size * next_n, device='cuda') % next_n
|
||||
ref_neginf_mask = ~(positions <= (context_lens[row_indices] - next_n + next_n_offset).unsqueeze(1))
|
||||
neginf_mask = (logits == float('-inf'))
|
||||
assert torch.equal(neginf_mask, ref_neginf_mask)
|
||||
def enumerate_paged_mqa_logits():
|
||||
arch_major = get_arch_major()
|
||||
for is_fp4 in ((True, False) if arch_major == 10 else (False, )):
|
||||
for logits_dtype in (torch.float, torch.bfloat16):
|
||||
for block_kv in ((32, 64) if arch_major == 10 else (64, )):
|
||||
for use_2d_context_lens, clean_logits in [(True, False)]:
|
||||
for batch_size in (256, ):
|
||||
for next_n in (1, 2, 4, 5, 6) if arch_major == 10 else (1, 2):
|
||||
for num_heads, head_dim in [(64, 128)]:
|
||||
for avg_kv in (8192, 32768):
|
||||
yield is_fp4, logits_dtype, block_kv, use_2d_context_lens, clean_logits, batch_size, next_n, num_heads, head_dim, avg_kv
|
||||
|
||||
logits = logits.masked_fill(ref_neginf_mask, 0)
|
||||
ref_logits = ref_logits.masked_fill(ref_neginf_mask, 0)
|
||||
diff = calc_diff(logits, ref_logits)
|
||||
assert diff < 1e-3, f"{diff=}"
|
||||
|
||||
sum_lens = sum(context_lens.to(torch.int64))
|
||||
tflops = 2 * sum_lens * next_n * heads * index_dim / 1e12
|
||||
input_bytes = count_bytes(q_fp8, weights, context_lens) + sum_lens * (index_dim + 4) + (sum_lens / blocksize) * 4
|
||||
output_bytes = sum_lens * next_n * 4
|
||||
if is_context_lens_2d:
|
||||
t = bench_kineto(lambda: deep_gemm.fp8_paged_mqa_logits(q_fp8, kv_cache_fp8, weights, context_lens_2d, block_tables, schedule_metadata, max_model_len, clean_logits=False),
|
||||
'fp8_paged_mqa_logits')
|
||||
else:
|
||||
t, clean_t = bench_kineto(lambda: deep_gemm.fp8_paged_mqa_logits(q_fp8, kv_cache_fp8, weights, context_lens, block_tables, schedule_metadata, max_model_len, clean_logits=True),
|
||||
('fp8_paged_mqa_logits', 'clean_logits'))
|
||||
clean_bytes = (batch_size * next_n * max_model_len - neginf_mask.sum().item()) * 4 + count_bytes(context_lens)
|
||||
print(f' > BSZ={batch_size:3}, NextN={next_n:1}, H={heads:2}, D={index_dim:2}, L={avg_kv:6}: '
|
||||
f'{tflops / t:4.0f} TFLOPS, {t * 1e6:3.0f} us, '
|
||||
f'{(input_bytes + output_bytes) / t / 1e9:4.0f} GB/s', end='')
|
||||
# noinspection PyUnboundLocalVariable
|
||||
print(f' | clean: {clean_t * 1e6:3.0f} us, {clean_bytes / clean_t / 1e9:4.0f} GB/s' if not is_context_lens_2d else '')
|
||||
print('Testing FP8/FP4 Paged MQA Logits:')
|
||||
max_model_len = 111 * 1024
|
||||
num_total_blocks = max_model_len * 5
|
||||
|
||||
for is_fp4, logits_dtype, block_kv, use_2d_context_lens, clean_logits, batch_size, next_n, num_heads, head_dim, avg_kv in enumerate_paged_mqa_logits():
|
||||
# Generate random inputs
|
||||
q = torch.randn((batch_size, next_n, num_heads, head_dim), device='cuda', dtype=torch.bfloat16)
|
||||
kv_cache = torch.randn((num_total_blocks, block_kv, 1, head_dim), device='cuda', dtype=torch.bfloat16)
|
||||
weights = torch.randn((batch_size * next_n, num_heads), device='cuda', dtype=torch.float)
|
||||
context_lens = torch.randint(int(0.7 * avg_kv), int(1.3 * avg_kv), (batch_size,), device='cuda', dtype=torch.int)
|
||||
|
||||
# Assign block tables
|
||||
num_blocks_per_query = ceil_div(context_lens, block_kv)
|
||||
block_table = torch.empty((batch_size, num_blocks_per_query.max().item()), device='cuda', dtype=torch.int)
|
||||
block_idx_pool = torch.randperm(num_total_blocks, device='cuda', dtype=torch.int)
|
||||
offset = 0
|
||||
for i, num_blocks in enumerate(num_blocks_per_query.tolist()):
|
||||
block_table[i, :num_blocks] = block_idx_pool[offset : offset + num_blocks]
|
||||
offset += num_blocks
|
||||
|
||||
# Calculate reference logits
|
||||
ref_logits = ref_paged_mqa_logits(q, kv_cache, weights, context_lens, block_table, max_model_len, use_2d_context_lens)
|
||||
|
||||
# Quantize Q and KV cache to FP4 / FP8
|
||||
if is_fp4:
|
||||
q_fp4 = per_token_cast_to_fp4(q.view(-1, head_dim), use_ue8m0=True, gran_k=32, use_packed_ue8m0=True)
|
||||
q_in = (q_fp4[0].view(batch_size, next_n, num_heads, head_dim // 2), q_fp4[1].view(batch_size, next_n, num_heads))
|
||||
q_simulated = cast_back_from_fp4(q_fp4[0], q_fp4[1], gran_k=32, use_packed_ue8m0=True).view(batch_size, next_n, num_heads, head_dim).to(torch.bfloat16)
|
||||
kv_in, kv_simulated = kv_cache_cast_to_fp4(kv_cache)
|
||||
else:
|
||||
q_in = q.to(torch.float8_e4m3fn), None
|
||||
q_simulated = q_in[0].to(torch.bfloat16)
|
||||
kv_in, kv_simulated = kv_cache_cast_to_fp8(kv_cache)
|
||||
|
||||
# Calculate simulated reference logits
|
||||
simulated_logits = ref_paged_mqa_logits(q_simulated, kv_simulated, weights, context_lens, block_table, max_model_len, use_2d_context_lens)
|
||||
|
||||
# Prepare masks and context lengths with NextN
|
||||
positions = torch.arange(max_model_len, device='cuda').unsqueeze(0).expand(batch_size * next_n, -1)
|
||||
if use_2d_context_lens:
|
||||
context_lens_nextn = ((context_lens.unsqueeze(1) + 1) * torch.rand(batch_size, next_n, device='cuda')).int()
|
||||
# Ensure last token matches actual length
|
||||
context_lens_nextn[:, -1] = context_lens
|
||||
ref_neginf_mask = ~(positions < context_lens_nextn.view(-1, 1))
|
||||
else:
|
||||
context_lens_nextn = context_lens
|
||||
offsets = torch.arange(batch_size * next_n, device='cuda')
|
||||
limits = (context_lens[offsets // next_n] - next_n + offsets % next_n).unsqueeze(1)
|
||||
ref_neginf_mask = ~(positions <= limits)
|
||||
|
||||
# Run Kernel
|
||||
kernel_kwargs = dict(
|
||||
q=q_in, kv_cache=kv_in, weights=weights,
|
||||
context_lens=context_lens_nextn, block_table=block_table,
|
||||
schedule_meta=deep_gemm.get_paged_mqa_logits_metadata(context_lens_nextn, block_kv, deep_gemm.get_num_sms()),
|
||||
max_context_len=max_model_len, clean_logits=clean_logits, logits_dtype=logits_dtype
|
||||
)
|
||||
logits = deep_gemm.fp8_fp4_paged_mqa_logits(**kernel_kwargs)
|
||||
|
||||
# Validation
|
||||
assert logits.dtype == logits_dtype
|
||||
logits = logits.to(torch.float)
|
||||
|
||||
if clean_logits:
|
||||
assert torch.equal(logits == float('-inf'), ref_neginf_mask), "Mask mismatch"
|
||||
|
||||
logits_masked = logits.masked_fill(ref_neginf_mask, 0)
|
||||
ref_masked = ref_logits.masked_fill(ref_neginf_mask, 0)
|
||||
simulated_masked = simulated_logits.masked_fill(ref_neginf_mask, 0)
|
||||
diff = calc_diff(logits_masked, ref_masked)
|
||||
simulated_diff = calc_diff(logits_masked, simulated_masked)
|
||||
assert diff < 0.02 if is_fp4 else 1e-3, f"Diff: {diff}"
|
||||
assert simulated_diff < 5e-6, f"Simulated Diff: {simulated_diff}"
|
||||
|
||||
# Profiling
|
||||
sum_lens = context_lens.sum().item()
|
||||
tflops_calc = 2 * sum_lens * next_n * num_heads * head_dim / 1e12
|
||||
kv_bytes_per_token = head_dim / (2 if is_fp4 else 1) + 4
|
||||
total_bytes = count_bytes(q, weights) + sum_lens * kv_bytes_per_token + (sum_lens * next_n * logits_dtype.itemsize)
|
||||
|
||||
t, clean_t = bench_kineto(lambda: deep_gemm.fp8_fp4_paged_mqa_logits(**kernel_kwargs), ('paged_mqa_logits', 'clean_logits'))
|
||||
print(f' > FP4={is_fp4}, BF16={logits_dtype == torch.bfloat16}, BLOCK_KV={block_kv}, BSZ={batch_size:3}, NextN={next_n:1}, H={num_heads:2}, D={head_dim:2}, L={avg_kv:6}: '
|
||||
f'{tflops_calc / t:4.0f} TFLOPS, {t * 1e6:3.0f} us, {total_bytes / t / 1e9:4.0f} GB/s', end='')
|
||||
print(f' | clean: {clean_t*1e6:3.0f} us' if clean_logits else '')
|
||||
print()
|
||||
|
||||
|
||||
@@ -280,6 +358,5 @@ if __name__ == '__main__':
|
||||
random.seed(0)
|
||||
|
||||
test_gemm_skip_head_mid()
|
||||
|
||||
test_mqa_logits()
|
||||
test_paged_mqa_logits()
|
||||
|
||||
@@ -11,7 +11,8 @@ from deep_gemm.testing import (
|
||||
from generators import (
|
||||
get_arch_major, layout_masked_to_psum, align,
|
||||
enumerate_normal, enumerate_m_grouped_contiguous, enumerate_m_grouped_masked, enumerate_k_grouped_contiguous,
|
||||
generate_normal, generate_m_grouped_contiguous, generate_m_grouped_masked, generate_k_grouped_contiguous
|
||||
generate_normal, generate_m_grouped_contiguous, generate_m_grouped_masked, generate_k_grouped_contiguous,
|
||||
get_mk_alignment_for_contiguous_layout
|
||||
)
|
||||
|
||||
|
||||
@@ -56,6 +57,10 @@ def test_m_grouped_gemm_contiguous() -> None:
|
||||
major_opt = 'N' if major_a.is_k_major() else 'T'
|
||||
major_opt += 'T' if major_b.is_k_major() else 'N'
|
||||
|
||||
# Select best alignment
|
||||
alignment = deep_gemm.get_theoretical_mk_alignment_for_contiguous_layout()
|
||||
deep_gemm.set_mk_alignment_for_contiguous_layout(alignment)
|
||||
|
||||
for test_alias in (False, True):
|
||||
m, a, b, grouped_layout, d, ref_d = generate_m_grouped_contiguous(num_groups, expected_m_per_group, n, k, major_a, major_b,
|
||||
use_bf16=True, use_psum_layout=use_psum_layout)
|
||||
@@ -65,8 +70,15 @@ def test_m_grouped_gemm_contiguous() -> None:
|
||||
b = b if major_b.is_k_major() else b.mT
|
||||
assert a[0].is_contiguous() and b[0].is_contiguous()
|
||||
getattr(deep_gemm, func_name)(a, b, d, grouped_layout, use_psum_layout=use_psum_layout)
|
||||
diff = calc_diff(d, ref_d)
|
||||
assert diff < 1e-5, f'{m=}, {n=}, {k=}, {major_opt}, {diff:.5f}, alias={test_alias}'
|
||||
if use_psum_layout:
|
||||
for j in range(num_groups):
|
||||
start = 0 if j == 0 else align(grouped_layout[j - 1], get_mk_alignment_for_contiguous_layout())
|
||||
end = grouped_layout[j]
|
||||
diff = calc_diff(d[start : end], ref_d[start : end])
|
||||
assert diff < 1e-5, f'{m=}, {n=}, {k=}, {major_opt}, {diff:.5f}, alias={test_alias}'
|
||||
else:
|
||||
diff = calc_diff(d, ref_d)
|
||||
assert diff < 1e-5, f'{m=}, {n=}, {k=}, {major_opt}, {diff:.5f}, alias={test_alias}'
|
||||
m, a, b, grouped_layout, d, ref_d = generate_m_grouped_contiguous(num_groups, expected_m_per_group, n, k, major_a, major_b,
|
||||
use_bf16=True, use_psum_layout=use_psum_layout)
|
||||
|
||||
@@ -91,6 +103,10 @@ def test_m_grouped_gemm_masked() -> None:
|
||||
sum_t, max_t = 0, 0
|
||||
sum_ops, sum_bytes = 0, 0
|
||||
|
||||
# Select best alignment
|
||||
alignment = deep_gemm.get_theoretical_mk_alignment_for_contiguous_layout(int(expected_m_per_group * 1.2))
|
||||
deep_gemm.set_mk_alignment_for_contiguous_layout(alignment)
|
||||
|
||||
for i in range(num_tests):
|
||||
a, b, masked_m, psum_m, d, ref_d = generate_m_grouped_masked(num_groups, max_m, expected_m_per_group, n, k,
|
||||
use_bf16=True, use_psum_layout=use_psum_layout)
|
||||
@@ -111,7 +127,7 @@ def test_m_grouped_gemm_masked() -> None:
|
||||
if masked_m[j].item() == 0:
|
||||
continue
|
||||
if use_psum_layout:
|
||||
d_slice = d_psum[: psum_m[j]] if j == 0 else d_psum[align(psum_m[j - 1], 128): psum_m[j]]
|
||||
d_slice = d_psum[: psum_m[j]] if j == 0 else d_psum[align(psum_m[j - 1], get_mk_alignment_for_contiguous_layout()): psum_m[j]]
|
||||
else:
|
||||
d_slice = d[j, :masked_m[j].item()]
|
||||
diff = calc_diff(d_slice, ref_d[j, :masked_m[j].item()])
|
||||
@@ -137,6 +153,9 @@ def test_m_grouped_gemm_masked() -> None:
|
||||
|
||||
def test_k_grouped_gemm_contiguous() -> None:
|
||||
print('Testing k-grouped contiguous GEMM:')
|
||||
|
||||
# TODO: Support arbitrary alignment
|
||||
deep_gemm.set_mk_alignment_for_contiguous_layout(128)
|
||||
|
||||
for num_groups, m, n, major_a, major_b, ks, expected_k_per_group in enumerate_k_grouped_contiguous(torch.bfloat16):
|
||||
for test_empty_groups in (False, True):
|
||||
|
||||
@@ -99,7 +99,7 @@ def test_fp8_bhr_hdr_bhd(use_ue8m0: bool = True):
|
||||
deep_gemm.fp8_einsum('bhr,hdr->bhd', x_fp8, y_fp8, z)
|
||||
assert calc_diff(z, ref_z) < 1e-3
|
||||
|
||||
t = bench_kineto(lambda: deep_gemm.fp8_einsum('bhr,hdr->bhd', x_fp8, y_fp8, z), 'fp8_gemm', suppress_kineto_output=True)
|
||||
t = bench_kineto(lambda: deep_gemm.fp8_einsum('bhr,hdr->bhd', x_fp8, y_fp8, z), 'gemm_', suppress_kineto_output=True)
|
||||
t_cublaslt = bench_kineto(lambda: deep_gemm.einsum('bhr,hdr->bhd', x, y, z, use_cublaslt=True), 'nvjet', suppress_kineto_output=True)
|
||||
print(f' > Perf ({b=:4.0f}, {h=}, {r=}, {d=}): ',
|
||||
f'{t * 1e6:4.0f} us | '
|
||||
@@ -129,7 +129,7 @@ def test_fp8_bhd_hdr_bhr(use_ue8m0: bool = True):
|
||||
deep_gemm.fp8_einsum('bhd,hdr->bhr', x_fp8, y_fp8, z)
|
||||
assert calc_diff(z, ref_z) < 1e-3
|
||||
|
||||
t = bench_kineto(lambda: deep_gemm.fp8_einsum('bhd,hdr->bhr', x_fp8, y_fp8, z), 'fp8_gemm', suppress_kineto_output=True)
|
||||
t = bench_kineto(lambda: deep_gemm.fp8_einsum('bhd,hdr->bhr', x_fp8, y_fp8, z), 'gemm_', suppress_kineto_output=True)
|
||||
t_cublaslt = bench_kineto(lambda: deep_gemm.einsum('bhd,hdr->bhr', x, y, z, use_cublaslt=True), 'nvjet', suppress_kineto_output=True)
|
||||
print(f' > Perf ({b=:4.0f}, {h=}, {r=}, {d=}): ',
|
||||
f'{t * 1e6:4.0f} us | '
|
||||
@@ -157,7 +157,7 @@ def test_fp8_bhd_bhr_hdr(use_ue8m0: bool = True):
|
||||
deep_gemm.fp8_einsum('bhd,bhr->hdr', x_fp8, y_fp8, z, z, recipe=(1, 1, 128))
|
||||
assert calc_diff(z, ref_z) < 1e-3
|
||||
|
||||
t = bench_kineto(lambda: deep_gemm.fp8_einsum('bhd,bhr->hdr', x_fp8, y_fp8, z, z, recipe=(1, 1, 128)), 'fp8_gemm', suppress_kineto_output=True)
|
||||
t = bench_kineto(lambda: deep_gemm.fp8_einsum('bhd,bhr->hdr', x_fp8, y_fp8, z, z, recipe=(1, 1, 128)), 'gemm_', suppress_kineto_output=True)
|
||||
print(f' > Perf ({b=:4.0f}, {h=}, {r=}, {d=}): ',
|
||||
f'{t * 1e6:4.0f} us | '
|
||||
f'{2 * b * h * r * d / t / 1e12:4.0f} TFLOPS | '
|
||||
|
||||
@@ -13,11 +13,11 @@ from deep_gemm.testing import (
|
||||
from generators import (
|
||||
KernelType, get_ue8m0_usage, layout_masked_to_psum, align,
|
||||
enumerate_normal, enumerate_m_grouped_contiguous, enumerate_m_grouped_masked, enumerate_k_grouped_contiguous,
|
||||
generate_normal, generate_m_grouped_contiguous, generate_m_grouped_masked, generate_k_grouped_contiguous
|
||||
generate_normal, generate_m_grouped_contiguous, generate_m_grouped_masked, generate_k_grouped_contiguous,
|
||||
get_mk_alignment_for_contiguous_layout
|
||||
)
|
||||
|
||||
|
||||
@ignore_env('DG_JIT_PTXAS_CHECK', lambda: get_arch_major() == 9)
|
||||
def test_gemm() -> None:
|
||||
print('Testing GEMM:')
|
||||
scores = []
|
||||
@@ -45,7 +45,7 @@ def test_gemm() -> None:
|
||||
|
||||
a, b, c, d, ref_d = generate_normal(m, n, k, major_a, major_b, accumulate, out_dtype, kernel_type, use_ue8m0=use_ue8m0, quant_config=quant_config)
|
||||
t = bench_kineto(lambda: deep_gemm.fp8_fp4_gemm_nt(a, b, d, c=c, disable_ue8m0_cast=disable_ue8m0_cast, recipe=recipe, recipe_a=recipe_a, recipe_b=recipe_b),
|
||||
'fp8_gemm', suppress_kineto_output=True)
|
||||
'gemm_', suppress_kineto_output=True)
|
||||
cublas_t, split_k_t = bench_kineto(lambda: deep_gemm.cublaslt_gemm_nt(a[0], b[0], d, c=c), ('nvjet', 'reduce'), suppress_kineto_output=True) \
|
||||
if not quant_config.is_fp4_a and not quant_config.is_fp4_b else (0, 0)
|
||||
print(f' > Perf (m={m:6}, n={n:6}, k={k:6}, {kernel_opt}, layout={major_opt}, {out_opt}, {acc_opt}): '
|
||||
@@ -68,6 +68,10 @@ def test_m_grouped_gemm_contiguous() -> None:
|
||||
disable_ue8m0_cast = not use_ue8m0
|
||||
recipe, recipe_a, recipe_b = quant_config.get_recipes()
|
||||
|
||||
# Select best alignment
|
||||
alignment = deep_gemm.get_theoretical_mk_alignment_for_contiguous_layout()
|
||||
deep_gemm.set_mk_alignment_for_contiguous_layout(alignment)
|
||||
|
||||
for test_alias in (False, True):
|
||||
m, a, b, grouped_layout, d, ref_d = generate_m_grouped_contiguous(num_groups, expected_m_per_group, n, k, major_a, major_b,
|
||||
use_ue8m0=use_ue8m0, use_psum_layout=use_psum_layout,
|
||||
@@ -79,8 +83,15 @@ def test_m_grouped_gemm_contiguous() -> None:
|
||||
assert a[0].is_contiguous() and b[0].is_contiguous()
|
||||
getattr(deep_gemm, func_name)(a, b, d, grouped_layout, disable_ue8m0_cast=disable_ue8m0_cast, use_psum_layout=use_psum_layout,
|
||||
recipe=recipe, recipe_a=recipe_a, recipe_b=recipe_b)
|
||||
diff = calc_diff(d, ref_d)
|
||||
assert diff < quant_config.max_diff(), f'{m=}, {n=}, {k=}, {major_opt}, {kernel_opt}, {diff:.5f}, alias={test_alias}'
|
||||
if use_psum_layout:
|
||||
for j in range(num_groups):
|
||||
start = 0 if j == 0 else align(grouped_layout[j - 1], get_mk_alignment_for_contiguous_layout())
|
||||
end = grouped_layout[j]
|
||||
diff = calc_diff(d[start : end], ref_d[start : end])
|
||||
assert diff < quant_config.max_diff(), f'{m=}, {n=}, {k=}, {major_opt}, {kernel_opt}, {diff:.5f}, alias={test_alias}'
|
||||
else:
|
||||
diff = calc_diff(d, ref_d)
|
||||
assert diff < quant_config.max_diff(), f'{m=}, {n=}, {k=}, {major_opt}, {kernel_opt}, {diff:.5f}, alias={test_alias}'
|
||||
m, a, b, grouped_layout, d, ref_d = generate_m_grouped_contiguous(num_groups, expected_m_per_group, n, k, major_a, major_b,
|
||||
use_ue8m0=use_ue8m0, use_psum_layout=use_psum_layout,
|
||||
quant_config=quant_config)
|
||||
@@ -90,7 +101,7 @@ def test_m_grouped_gemm_contiguous() -> None:
|
||||
deep_gemm.m_grouped_fp8_fp4_gemm_nt_contiguous(a, b, d, grouped_layout, disable_ue8m0_cast=disable_ue8m0_cast, use_psum_layout=use_psum_layout,
|
||||
recipe=recipe, recipe_a=recipe_a, recipe_b=recipe_b)
|
||||
|
||||
t = bench_kineto(test_func, 'fp8_gemm', suppress_kineto_output=True)
|
||||
t = bench_kineto(test_func, 'gemm_', suppress_kineto_output=True)
|
||||
print(f' > Perf ({num_groups=}, m={m:5}, n={n:6}, k={k:5}, {kernel_opt}, layout={major_opt}, psum={use_psum_layout}): '
|
||||
f'{t * 1e6:4.0f} us | '
|
||||
f'{2 * m * n * k / t / 1e12:4.0f} TFLOPS | '
|
||||
@@ -112,6 +123,10 @@ def test_m_grouped_gemm_masked() -> None:
|
||||
sum_t, max_t = 0, 0
|
||||
sum_ops, sum_bytes = 0, 0
|
||||
|
||||
# Select best alignment
|
||||
alignment = deep_gemm.get_theoretical_mk_alignment_for_contiguous_layout(int(expected_m_per_group * 1.2))
|
||||
deep_gemm.set_mk_alignment_for_contiguous_layout(alignment)
|
||||
|
||||
for i in range(num_tests):
|
||||
a, b, masked_m, psum_m, d, ref_d = generate_m_grouped_masked(num_groups, max_m, expected_m_per_group, n, k,
|
||||
use_ue8m0=use_ue8m0, use_psum_layout=use_psum_layout,
|
||||
@@ -124,10 +139,10 @@ def test_m_grouped_gemm_masked() -> None:
|
||||
def test_func():
|
||||
if use_psum_layout:
|
||||
deep_gemm.m_grouped_fp8_fp4_gemm_nt_contiguous(a_psum, b, d_psum, psum_m, disable_ue8m0_cast=disable_ue8m0_cast,
|
||||
use_psum_layout=True, expected_m_for_psum_layout=expected_m_per_group,
|
||||
use_psum_layout=True, expected_m_for_psum_layout=int(expected_m_per_group * 1.2),
|
||||
recipe=recipe, recipe_a=recipe_a, recipe_b=recipe_b)
|
||||
else:
|
||||
deep_gemm.m_grouped_fp8_fp4_gemm_nt_masked(a, b, d, masked_m, expected_m_per_group, disable_ue8m0_cast=disable_ue8m0_cast,
|
||||
deep_gemm.m_grouped_fp8_fp4_gemm_nt_masked(a, b, d, masked_m, int(expected_m_per_group * 1.2), disable_ue8m0_cast=disable_ue8m0_cast,
|
||||
recipe=recipe, recipe_a=recipe_a, recipe_b=recipe_b)
|
||||
|
||||
test_func()
|
||||
@@ -135,7 +150,7 @@ def test_m_grouped_gemm_masked() -> None:
|
||||
if masked_m[j].item() == 0:
|
||||
continue
|
||||
if use_psum_layout:
|
||||
d_slice = d_psum[: psum_m[j]] if j == 0 else d_psum[align(psum_m[j - 1], 128): psum_m[j]]
|
||||
d_slice = d_psum[: psum_m[j]] if j == 0 else d_psum[align(psum_m[j - 1], get_mk_alignment_for_contiguous_layout()): psum_m[j]]
|
||||
else:
|
||||
d_slice = d[j, :masked_m[j].item()]
|
||||
diff = calc_diff(d_slice, ref_d[j, :masked_m[j].item()])
|
||||
@@ -143,7 +158,7 @@ def test_m_grouped_gemm_masked() -> None:
|
||||
|
||||
# Test performance with fixed shapes
|
||||
valid_m = masked_m.sum().item()
|
||||
t = bench_kineto(test_func, 'fp8_gemm', suppress_kineto_output=True)
|
||||
t = bench_kineto(test_func, 'gemm_', suppress_kineto_output=True)
|
||||
|
||||
sum_t += t
|
||||
max_t = max(max_t, t)
|
||||
@@ -158,36 +173,36 @@ def test_m_grouped_gemm_masked() -> None:
|
||||
print()
|
||||
|
||||
|
||||
@ignore_env('DG_JIT_PTXAS_CHECK', lambda: get_arch_major() == 9)
|
||||
def test_k_grouped_gemm_contiguous() -> None:
|
||||
print('Testing k-grouped contiguous GEMM:')
|
||||
|
||||
k_grouped_fp8_gemm_contiguous = deep_gemm.k_grouped_fp8_gemm_nt_contiguous if get_arch_major() == 9 \
|
||||
else deep_gemm.k_grouped_fp8_gemm_tn_contiguous
|
||||
for num_groups, m, n, major_a, major_b, ks, expected_k_per_group in enumerate_k_grouped_contiguous(torch.float8_e4m3fn):
|
||||
for num_groups, m, n, major_a, major_b, ks, expected_k_per_group, gran_k in enumerate_k_grouped_contiguous(torch.float8_e4m3fn):
|
||||
recipe = (1, 1, gran_k)
|
||||
use_ue8m0 = get_ue8m0_usage(KernelType.Kernel1D1D)
|
||||
|
||||
for test_empty_groups in (False, True):
|
||||
new_ks = copy.deepcopy(ks)
|
||||
if test_empty_groups and len(ks) > 1:
|
||||
new_ks[random.randint(0, num_groups - 1)] = 0
|
||||
k, a, b, c, d, ref_d = generate_k_grouped_contiguous(num_groups, m, n, major_a, major_b, new_ks, use_ue8m0=use_ue8m0)
|
||||
k, a, b, c, d, ref_d = generate_k_grouped_contiguous(num_groups, m, n, major_a, major_b, new_ks, use_ue8m0=use_ue8m0, gran_k=gran_k)
|
||||
new_ks_tensor = torch.tensor(new_ks, dtype=torch.int, device='cuda')
|
||||
k_grouped_fp8_gemm_contiguous(a, b, d, new_ks, new_ks_tensor, c)
|
||||
k_grouped_fp8_gemm_contiguous(a, b, d, new_ks, new_ks_tensor, c, recipe=recipe)
|
||||
|
||||
diff = calc_diff(d, ref_d)
|
||||
assert diff < 0.001, f'{m=}, {n=}, {k=}, {ks=}, {diff:.5f}'
|
||||
|
||||
# Test performance
|
||||
k, a, b, c, d, ref_d = generate_k_grouped_contiguous(num_groups, m, n, major_a, major_b, ks, use_ue8m0=use_ue8m0)
|
||||
k, a, b, c, d, ref_d = generate_k_grouped_contiguous(num_groups, m, n, major_a, major_b, ks, use_ue8m0=use_ue8m0, gran_k=gran_k)
|
||||
ks_tensor = torch.tensor(ks, dtype=torch.int, device='cuda')
|
||||
|
||||
# noinspection PyShadowingNames
|
||||
def test_func():
|
||||
k_grouped_fp8_gemm_contiguous(a, b, d, ks, ks_tensor, c)
|
||||
k_grouped_fp8_gemm_contiguous(a, b, d, ks, ks_tensor, c, recipe=recipe)
|
||||
|
||||
t = bench_kineto(test_func, 'fp8_gemm', suppress_kineto_output=True)
|
||||
print(f' > Perf ({num_groups=:2}, m={m:5}, n={n:5}, k={k:5}): '
|
||||
t = bench_kineto(test_func, 'gemm_', suppress_kineto_output=True)
|
||||
print(f' > Perf ({num_groups=:2}, m={m:5}, n={n:5}, k={k:5}, gran_k={gran_k:3}): '
|
||||
f'{t * 1e6:4.0f} us | '
|
||||
f'{2 * m * n * k / t / 1e12:4.0f} TFLOPS | '
|
||||
f'{count_bytes(a, b, c, d) / 1e9 / t:4.0f} GB/s')
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
import torch
|
||||
import random
|
||||
from deep_gemm.testing import bench_kineto, count_bytes
|
||||
from deep_gemm.testing import bench_kineto, count_bytes, get_arch_major
|
||||
from deep_gemm.utils import (
|
||||
align, ceil_div,
|
||||
per_token_cast_to_fp8, per_channel_cast_to_fp8,
|
||||
@@ -43,9 +43,9 @@ def get_mn_major_tma_aligned_packed_ue8m0_tensor_torch_impl(x: torch.Tensor) ->
|
||||
|
||||
def test_sf_layout_kernels() -> None:
|
||||
print('Testing SF layout kernels:')
|
||||
for mn, k, with_transpose, use_ue8m0, num_groups in enumerate_sf_layout():
|
||||
for mn, k, with_transpose, use_ue8m0, num_groups, gran_k in enumerate_sf_layout():
|
||||
x = torch.randn((num_groups * mn, k), dtype=torch.bfloat16, device='cuda')
|
||||
x, fp32_sf = per_token_cast_to_fp8(x, use_ue8m0=use_ue8m0)
|
||||
x, fp32_sf = per_token_cast_to_fp8(x, use_ue8m0=use_ue8m0, gran_k=gran_k)
|
||||
fp32_sf = fp32_sf if num_groups == 1 else fp32_sf.view(num_groups, mn, -1)
|
||||
fp32_sf = fp32_sf if with_transpose else fp32_sf.transpose(-1, -2).contiguous().transpose(-1, -2)
|
||||
|
||||
@@ -60,7 +60,7 @@ def test_sf_layout_kernels() -> None:
|
||||
else:
|
||||
impl, name = get_mn_major_tma_aligned_tensor, 'transpose'
|
||||
transposed_sf = get_mn_major_tma_aligned_tensor(fp32_sf)
|
||||
tma_aligned_mn, sf_k = get_tma_aligned_size(mn, fp32_sf.element_size()), ceil_div(k, 128)
|
||||
tma_aligned_mn, sf_k = get_tma_aligned_size(mn, fp32_sf.element_size()), ceil_div(k, gran_k)
|
||||
if num_groups > 1:
|
||||
assert transposed_sf.size(0) == num_groups
|
||||
assert transposed_sf.stride(0) == tma_aligned_mn * sf_k
|
||||
@@ -74,22 +74,22 @@ def test_sf_layout_kernels() -> None:
|
||||
except AssertionError as e:
|
||||
# Some cases may fallback to PyTorch impl
|
||||
t = 0
|
||||
print(f' > Perf ({num_groups=:2}, {mn=:5}, {k=:5}, transpose={int(with_transpose)}, use_ue8m0={int(use_ue8m0)}): '
|
||||
print(f' > Perf ({num_groups=:2}, {mn=:5}, {k=:5}, transpose={int(with_transpose)}, use_ue8m0={int(use_ue8m0)}, gran_k={gran_k:3}): '
|
||||
f'{t * 1e6:4.0f} us | {count_bytes(fp32_sf, impl(fp32_sf)) / 1e9 / t if t else 0:4.0f} GB/s')
|
||||
print()
|
||||
|
||||
|
||||
def test_k_grouped_sf_layout_kernels() -> None:
|
||||
print('Testing k-grouped SF layout kernels:')
|
||||
for mn, ks, num_groups in enumerate_k_grouped_sf_layout():
|
||||
sf_ks = [k // 128 for k in ks]
|
||||
packed_sf_ks = [ceil_div(k, 512) for k in ks]
|
||||
for mn, ks, num_groups, gran_k in enumerate_k_grouped_sf_layout():
|
||||
sf_ks = [k // gran_k for k in ks]
|
||||
packed_sf_ks = [ceil_div(k, gran_k * 4) for k in ks]
|
||||
ks_tensor = torch.tensor(ks, dtype=torch.int, device='cuda')
|
||||
x = torch.randn((sum(ks), mn), dtype=torch.bfloat16, device='cuda')
|
||||
x, fp32_sf = per_channel_cast_to_fp8(x, use_ue8m0=True)
|
||||
x, fp32_sf = per_channel_cast_to_fp8(x, use_ue8m0=True, gran_k=gran_k)
|
||||
|
||||
# Correctness
|
||||
packed_sf = get_k_grouped_mn_major_tma_aligned_packed_ue8m0_tensor(fp32_sf, ks_tensor, ks)
|
||||
packed_sf = get_k_grouped_mn_major_tma_aligned_packed_ue8m0_tensor(fp32_sf, ks_tensor, ks, gran_k)
|
||||
split_packed_sf = packed_sf.split(packed_sf_ks)
|
||||
split_fp32_sf = fp32_sf.split(sf_ks)
|
||||
for i in range(num_groups):
|
||||
@@ -97,8 +97,8 @@ def test_k_grouped_sf_layout_kernels() -> None:
|
||||
assert torch.equal(split_packed_sf[i], ref_packed_sf), f'{i=}'
|
||||
|
||||
# Performance
|
||||
t = bench_kineto(lambda: get_k_grouped_mn_major_tma_aligned_packed_ue8m0_tensor(fp32_sf, ks_tensor, ks), 'pack_fp32_into_ue8m0')
|
||||
print(f' > Perf ({num_groups=:3}, {mn=:5}, sum_k={sum(ks):5}):'
|
||||
t = bench_kineto(lambda: get_k_grouped_mn_major_tma_aligned_packed_ue8m0_tensor(fp32_sf, ks_tensor, ks, gran_k), 'pack_fp32_into_ue8m0')
|
||||
print(f' > Perf ({num_groups=:3}, {mn=:5}, sum_k={sum(ks):5}, gran_k={gran_k:3}):'
|
||||
f'{t * 1e6:4.0f} us | '
|
||||
f'{count_bytes(fp32_sf, packed_sf, ks_tensor) / 1e9 / t:4.0f} GB/s')
|
||||
print()
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import argparse
|
||||
import torch
|
||||
import torch.multiprocessing as mp
|
||||
import deep_gemm
|
||||
@@ -8,7 +9,11 @@ def main(local_rank: int):
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
procs = [mp.Process(target=main, args=(i, ), ) for i in range(8)]
|
||||
parser = argparse.ArgumentParser(description='Test lazy initialization')
|
||||
parser.add_argument('--num-processes', type=int, default=8, help='Number of processes to spawn (default: 8)')
|
||||
args = parser.parse_args()
|
||||
|
||||
procs = [mp.Process(target=main, args=(i, ), ) for i in range(args.num_processes)]
|
||||
for p in procs:
|
||||
p.start()
|
||||
for p in procs:
|
||||
|
||||
264
tests/test_mega_moe.py
Normal file
264
tests/test_mega_moe.py
Normal file
@@ -0,0 +1,264 @@
|
||||
import argparse
|
||||
import os
|
||||
import random
|
||||
import sys
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from typing import Tuple
|
||||
|
||||
import deep_gemm
|
||||
from deep_gemm.utils import per_token_cast_to_fp4, per_token_cast_to_fp8
|
||||
from deep_gemm.utils.dist import dist_print, init_dist, uneven_all_gather
|
||||
from deep_gemm.testing import bench, bench_kineto, calc_diff
|
||||
|
||||
# Load legacy implements from third-party
|
||||
# noinspection PyBroadException
|
||||
try:
|
||||
import deep_ep
|
||||
import importlib.util
|
||||
from tilelang.profiler.bench import do_bench
|
||||
spec = importlib.util.spec_from_file_location(
|
||||
'tilelang_ops',
|
||||
os.path.join(os.path.dirname(os.path.realpath(__file__)), '..', 'third-party', 'tilelang_ops', '__init__.py'))
|
||||
tilelang_ops = importlib.util.module_from_spec(spec)
|
||||
sys.modules['tilelang_ops'] = tilelang_ops
|
||||
spec.loader.exec_module(tilelang_ops)
|
||||
is_legacy_loaded = True
|
||||
except Exception as ex:
|
||||
print(f'Failed to load legacy code: {ex}, skip baseline benchmarking')
|
||||
is_legacy_loaded = False
|
||||
|
||||
|
||||
# TODO: skip the test for SM90
|
||||
# noinspection PyUnboundLocalVariable,PyShadowingNames
|
||||
def test(local_rank: int, num_local_ranks: int, args: argparse.Namespace):
|
||||
rank_idx, num_ranks, group = init_dist(local_rank, num_local_ranks)
|
||||
torch.manual_seed(rank_idx)
|
||||
random.seed(rank_idx)
|
||||
|
||||
# Settings
|
||||
num_max_tokens_per_rank = args.num_max_tokens_per_rank
|
||||
num_tokens = max(0, args.num_max_tokens_per_rank - random.randint(0, args.num_max_removed_tokens)) \
|
||||
if args.num_tokens == 0 else args.num_tokens
|
||||
hidden, intermediate_hidden = args.hidden, args.intermediate_hidden
|
||||
num_experts, num_topk = args.num_experts, args.num_topk
|
||||
num_experts_per_rank = num_experts // num_ranks
|
||||
assert num_tokens <= num_max_tokens_per_rank
|
||||
|
||||
# Allocate symmetric memory
|
||||
buffer = deep_gemm.get_symm_buffer_for_mega_moe(
|
||||
group, num_experts,
|
||||
num_max_tokens_per_rank, num_topk,
|
||||
hidden, intermediate_hidden
|
||||
)
|
||||
dist_print('Config:', once_in_node=True)
|
||||
dist_print(f' > Tokens: {num_tokens}/{num_max_tokens_per_rank}', once_in_node=True)
|
||||
dist_print(f' > Hidden: {hidden}', once_in_node=True)
|
||||
dist_print(f' > Intermediate: {intermediate_hidden}', once_in_node=True)
|
||||
dist_print(f' > Experts: {num_topk}/{num_experts}', once_in_node=True)
|
||||
dist_print(f' > Buffer: {buffer.buffer.nbytes / 2 ** 30:.3f} GiB', once_in_node=True)
|
||||
dist_print(once_in_node=True)
|
||||
|
||||
# Non-overlapped baseline: EP dispatch + GEMM + EP combine
|
||||
alignment = deep_gemm.get_theoretical_mk_alignment_for_contiguous_layout()
|
||||
deep_gemm.set_mk_alignment_for_contiguous_layout(alignment)
|
||||
ep_buffer = deep_ep.ElasticBuffer(
|
||||
group,
|
||||
num_max_tokens_per_rank=num_max_tokens_per_rank, hidden=hidden,
|
||||
num_topk=num_topk, use_fp8_dispatch=True,
|
||||
explicitly_destroy=True,
|
||||
allow_multiple_reduction=False,
|
||||
gpu_timeout_secs=10, cpu_timeout_secs=30
|
||||
) if is_legacy_loaded else None
|
||||
|
||||
# Create inputs
|
||||
def create_inputs():
|
||||
global x, topk_idx, topk_weights, l1_weights, l2_weights, transformed_l1_weights, transformed_l2_weights
|
||||
x = torch.randn((num_tokens, hidden), dtype=torch.bfloat16, device='cuda')
|
||||
l1_weights = torch.randn(
|
||||
(num_experts_per_rank, intermediate_hidden * 2, hidden), dtype=torch.bfloat16, device='cuda')
|
||||
l2_weights = torch.randn(
|
||||
(num_experts_per_rank, hidden, intermediate_hidden), dtype=torch.bfloat16, device='cuda')
|
||||
scores = torch.randn((num_tokens, num_experts), dtype=torch.float, device='cuda')
|
||||
topk_weights, topk_idx = torch.topk(scores, num_topk, dim=-1, largest=True, sorted=False)
|
||||
if args.masked_ratio > 0:
|
||||
rand_mask = torch.rand_like(topk_idx, dtype=torch.float)
|
||||
topk_idx.masked_fill_(rand_mask < args.masked_ratio, -1)
|
||||
topk_weights.masked_fill_(topk_idx < 0, 0)
|
||||
|
||||
# Check SF requirements
|
||||
assert hidden % 128 == 0
|
||||
assert intermediate_hidden % 128 == 0
|
||||
assert l1_weights.shape[2] % 128 == 0 and l2_weights.shape[2] % 128 == 0
|
||||
|
||||
# Cast inputs to FP8 with per-32 UE8M0 SF
|
||||
x = per_token_cast_to_fp8(x, use_ue8m0=True, gran_k=32, use_packed_ue8m0=True)
|
||||
|
||||
# Cast grouped BF16 weights to FP4 with MN-major SF
|
||||
# TODO: merge with `cast_fp8_fp4_with_major`
|
||||
def cast_grouped_weights_to_fp4(bf16_weights: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
num_groups, n, k = bf16_weights.shape
|
||||
w = torch.empty((num_groups, n, k // 2), device='cuda', dtype=torch.int8)
|
||||
w_sf = torch.empty((num_groups, n, k // 32), device='cuda', dtype=torch.float)
|
||||
for i in range(num_groups):
|
||||
w[i], w_sf[i] = per_token_cast_to_fp4(bf16_weights[i], use_ue8m0=True, gran_k=32)
|
||||
w_sf = deep_gemm.transform_sf_into_required_layout(w_sf, n, k, (1, 32), num_groups)
|
||||
return w, w_sf
|
||||
|
||||
l1_weights = cast_grouped_weights_to_fp4(l1_weights)
|
||||
l2_weights = cast_grouped_weights_to_fp4(l2_weights)
|
||||
transformed_l1_weights, transformed_l2_weights = deep_gemm.transform_weights_for_mega_moe(l1_weights, l2_weights)
|
||||
|
||||
def run_baseline():
|
||||
recv_x, _, recv_topk_weights, handle, _ = ep_buffer.dispatch(
|
||||
x, topk_idx=topk_idx, topk_weights=topk_weights,
|
||||
num_experts=num_experts, expert_alignment=alignment,
|
||||
do_cpu_sync=False, do_handle_copy=False,
|
||||
do_expand=True, use_tma_aligned_col_major_sf=True
|
||||
)
|
||||
n = recv_x[0].size(0)
|
||||
l1_y = torch.empty((n, intermediate_hidden * 2), dtype=torch.bfloat16, device='cuda')
|
||||
deep_gemm.m_grouped_fp8_fp4_gemm_nt_contiguous(
|
||||
recv_x, l1_weights, l1_y, handle.psum_num_recv_tokens_per_expert,
|
||||
use_psum_layout=True, recipe=(1, 1, 32))
|
||||
# noinspection PyCallingNonCallable
|
||||
l1_y = tilelang_ops.swiglu_apply_weight_to_fp8(
|
||||
x=l1_y,
|
||||
topk_weights=recv_topk_weights,
|
||||
avail_tokens=handle.psum_num_recv_tokens_per_expert[-1],
|
||||
num_per_channels=32,
|
||||
use_col_major_scales=True,
|
||||
round_scale=True,
|
||||
ue8m0_scale=True,
|
||||
output_bf16=False,
|
||||
clamp_value=args.activation_clamp,
|
||||
fast_math=bool(args.fast_math)
|
||||
)
|
||||
l2_y = torch.empty((n, hidden), dtype=torch.bfloat16, device='cuda')
|
||||
deep_gemm.m_grouped_fp8_fp4_gemm_nt_contiguous(
|
||||
l1_y, l2_weights, l2_y, handle.psum_num_recv_tokens_per_expert,
|
||||
use_psum_layout=True, recipe=(1, 1, 32))
|
||||
return ep_buffer.combine(l2_y, handle=handle)[0]
|
||||
|
||||
# Run fused mega MoE
|
||||
# NOTES: copy x into buffer before each call because debug mode zeros the entire buffer
|
||||
def run_fused():
|
||||
buffer.x[:num_tokens].copy_(x[0])
|
||||
buffer.x_sf[:num_tokens].copy_(x[1])
|
||||
buffer.topk_idx[:num_tokens].copy_(topk_idx)
|
||||
buffer.topk_weights[:num_tokens].copy_(topk_weights)
|
||||
|
||||
y = torch.empty((num_tokens, hidden), dtype=torch.bfloat16, device='cuda')
|
||||
# noinspection PyTypeChecker
|
||||
deep_gemm.fp8_fp4_mega_moe(
|
||||
y,
|
||||
transformed_l1_weights, transformed_l2_weights,
|
||||
buffer,
|
||||
activation_clamp=args.activation_clamp,
|
||||
fast_math=bool(args.fast_math)
|
||||
)
|
||||
return y
|
||||
|
||||
# Check correctness (must be bitwise identical)
|
||||
num_correctness_tests = 1 if args.num_correctness_tests is None else args.num_correctness_tests
|
||||
# noinspection PyBroadException
|
||||
if is_legacy_loaded and num_correctness_tests > 0:
|
||||
dist_print('Running correctness tests:', once_in_node=True)
|
||||
for i in range(num_correctness_tests):
|
||||
create_inputs()
|
||||
assert torch.equal(run_fused(), run_baseline())
|
||||
if (i + 1) % 100 == 0 or i == num_correctness_tests - 1:
|
||||
dist_print(f' > Correctness test #{i + 1}/{args.num_correctness_tests} passed', once_in_node=True)
|
||||
dist_print(once_in_node=True)
|
||||
else:
|
||||
create_inputs()
|
||||
|
||||
# Count local received tokens
|
||||
gathered_topk_idx = uneven_all_gather(topk_idx, group=group)
|
||||
num_recv_tokens = (rank_idx * num_experts_per_rank <= gathered_topk_idx) & \
|
||||
(gathered_topk_idx < (rank_idx + 1) * num_experts_per_rank)
|
||||
num_recv_tokens = num_recv_tokens.sum().item()
|
||||
|
||||
# Benchmark
|
||||
t_fused = bench_kineto(
|
||||
run_fused, 'mega_moe',
|
||||
barrier=lambda: ep_buffer.barrier(use_comm_stream=False) if ep_buffer else dist.barrier(),
|
||||
trace_path=None if not args.dump_profile_traces else f'{args.dump_profile_traces}/mega_moe_rank{rank_idx}.json')
|
||||
t_baseline = do_bench(run_baseline, _n_warmup=5, _n_repeat=1, backend='cudagraph', return_mode='median') / 1e3 if is_legacy_loaded else 0
|
||||
|
||||
# TFLOPS: 3 matmuls (L1 left, L1 right, L2), each 2 * M * N * K
|
||||
safe_div = lambda a, b: float('nan') if b == 0 else a / b
|
||||
tflops = safe_div(2 * num_recv_tokens * (hidden * intermediate_hidden * 3) / 1e12, t_fused)
|
||||
|
||||
# HBM bytes: weights (FP4 packed = 0.5 bytes) + activations (FP8 = 1 byte) + output (BF16 = 2 bytes)
|
||||
num_hbm_bytes = (
|
||||
num_experts_per_rank * intermediate_hidden * 2 * hidden // 2 + # L1 weights (FP4)
|
||||
num_experts_per_rank * hidden * intermediate_hidden // 2 + # L2 weights (FP4)
|
||||
num_recv_tokens * hidden + # L1 acts read (FP8)
|
||||
num_recv_tokens * intermediate_hidden + # L1 output write (FP8)
|
||||
num_recv_tokens * intermediate_hidden + # L2 acts read (FP8)
|
||||
num_recv_tokens * hidden * 2 # L2 output write (BF16)
|
||||
)
|
||||
hbm_gbs = safe_div(num_hbm_bytes / 1e9, t_fused)
|
||||
|
||||
# NVLink bytes: dispatch pull + combine write-back
|
||||
num_nvlink_bytes = num_recv_tokens * hidden * 3
|
||||
nvlink_gbs = safe_div(num_nvlink_bytes / 1e9, t_fused)
|
||||
|
||||
# Combine reduction (serial) time approximation
|
||||
t_reduction = num_tokens * hidden * 2 * (1 + num_topk) / 6.5e12
|
||||
|
||||
# Summary
|
||||
approx_factor = t_fused / (t_fused - t_reduction)
|
||||
dist_print('Performance:', once_in_node=True)
|
||||
dist_print(f' > EP: {rank_idx:2}/{num_ranks} | '
|
||||
f'{tflops:4.0f} TFLOPS | '
|
||||
f'overlap: '
|
||||
f'{tflops * approx_factor:4.0f} TFLOPS, '
|
||||
f'HBM {hbm_gbs * approx_factor:4.0f} GB/s, '
|
||||
f'NVL {nvlink_gbs * approx_factor:3.0f} GB/s | '
|
||||
f'{t_fused * 1e6:4.0f} us, '
|
||||
f'reduction: {t_reduction * 1e6:4.1f} us | '
|
||||
f'{safe_div(t_baseline, t_fused):.2f}x legacy')
|
||||
|
||||
# Exit
|
||||
dist.barrier()
|
||||
buffer.destroy()
|
||||
ep_buffer.destroy() if is_legacy_loaded else None
|
||||
dist.destroy_process_group()
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
parser = argparse.ArgumentParser(description='Test PyTorch symmetric memory')
|
||||
# Resource settings
|
||||
parser.add_argument('--num-processes', type=int, default=8, help='Number of processes to spawn (default: 8)')
|
||||
|
||||
# Model settings
|
||||
parser.add_argument('--num-max-tokens-per-rank', type=int, default=8192, help='Number of maximum tokens per rank')
|
||||
parser.add_argument('--num-tokens', type=int, default=0, help='Number of tokens per rank (follow max minus removed if 0)')
|
||||
parser.add_argument('--num-max-removed-tokens', type=int, default=0, help='Maximum number of tokens to remove')
|
||||
parser.add_argument('--hidden', type=int, default=7168, help='Hidden size')
|
||||
parser.add_argument('--intermediate-hidden', type=int, default=3072, help='Intermediate hidden size')
|
||||
parser.add_argument('--activation-clamp', type=float, default=10, help='Clamp value for activation')
|
||||
parser.add_argument('--num-experts', type=int, default=384, help='Number of experts')
|
||||
parser.add_argument('--num-topk', type=int, default=6, help='Number of expert selections')
|
||||
parser.add_argument('--masked-ratio', type=float, default=0.0, help='Mask some expert selections')
|
||||
parser.add_argument('--fast-math', type=int, default=1, help='Enable fast math (0 or 1, default: 1)')
|
||||
|
||||
# Test settings
|
||||
parser.add_argument('--num-correctness-tests', type=int, default=None, help='Pressure test')
|
||||
parser.add_argument('--dump-profile-traces', type=str, default='', help='Dump profiling trace JSONs')
|
||||
parser.add_argument('--local-rank-idx', type=int, default=None, help='Run as single process with this local rank (e.g. for NCU prof)')
|
||||
args = parser.parse_args()
|
||||
|
||||
# Create dump trace directories
|
||||
if args.dump_profile_traces:
|
||||
os.makedirs(args.dump_profile_traces, exist_ok=True)
|
||||
|
||||
if args.local_rank_idx is not None:
|
||||
# Single-process mode: each process is launched separately (e.g. by NCU)
|
||||
test(args.local_rank_idx, args.num_processes, args)
|
||||
else:
|
||||
# Launch tests
|
||||
num_processes = args.num_processes
|
||||
torch.multiprocessing.spawn(test, args=(num_processes, args), nprocs=num_processes)
|
||||
@@ -21,7 +21,7 @@ sys.path.append('{script_dir}')
|
||||
torch.manual_seed(0)
|
||||
random.seed(0)
|
||||
|
||||
from tests.{module_name} import {func_name}
|
||||
from {module_name} import {func_name}
|
||||
{func_name}()
|
||||
"""
|
||||
|
||||
@@ -40,7 +40,7 @@ if __name__ == '__main__':
|
||||
else:
|
||||
# Get all test functions except those related to cuBLAS
|
||||
files = [f for f in os.listdir(script_dir) if f.endswith('.py')]
|
||||
exclude_files = ['test_sanitizer.py', 'generators.py']
|
||||
exclude_files = ['test_sanitizer.py', 'generators.py', 'test_mega_moe.py']
|
||||
funcs = [
|
||||
(module_name, name)
|
||||
for module_name in [os.path.splitext(f)[0] for f in files if f not in exclude_files]
|
||||
@@ -53,6 +53,7 @@ if __name__ == '__main__':
|
||||
env['CUDA_LAUNCH_BLOCKING'] = '1'
|
||||
env['DG_JIT_PTXAS_CHECK'] = '1'
|
||||
env['DG_USE_NVIDIA_TOOLS'] = '1'
|
||||
env['DG_USE_TEMP_CUBLASLT_WORKSPACE'] = '1' # Avoid holding CUDA tensor that crashes during shutdown
|
||||
env['PYTORCH_NO_CUDA_MEMORY_CACHING'] = '1'
|
||||
env['TORCH_SHOW_CPP_STACKTRACES'] = '1'
|
||||
|
||||
|
||||
Reference in New Issue
Block a user