[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:
Chenggang Zhao
2026-04-17 09:45:14 +08:00
committed by GitHub
parent d30fc36c8f
commit 7f2a703ed5
109 changed files with 12101 additions and 3219 deletions

View File

@@ -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):