[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):
|
||||
|
||||
Reference in New Issue
Block a user