feat: implement sm90 megamoe phase5 l2 scatter

This commit is contained in:
Xinyi Liu
2026-06-18 15:17:20 +08:00
parent fc8218750c
commit 9bd0519605
5 changed files with 379 additions and 72 deletions

View File

@@ -85,6 +85,8 @@ def main() -> None:
assert buffer.l2_acts.shape[1] == args.intermediate_hidden
assert buffer.l2_acts_sf.shape[1] == args.intermediate_hidden // 64
assert buffer.l2_acts_sf.dtype == torch.float32
assert buffer.combine_acts.shape == (args.num_topk, buffer.num_max_tokens_per_rank, args.hidden)
assert buffer.combine_acts.dtype == torch.bfloat16
num_tokens = args.num_tokens
buffer.x[:num_tokens].copy_(torch.randn((num_tokens, args.hidden), device='cuda').to(torch.float8_e4m3fn))

View File

@@ -0,0 +1,201 @@
import argparse
import inspect
import os
import pathlib
import sys
from typing import Tuple
import torch
import torch.distributed as dist
REPO_ROOT = pathlib.Path(__file__).resolve().parents[2]
if str(REPO_ROOT) not in sys.path:
sys.path.insert(0, str(REPO_ROOT))
import deep_gemm
from deep_gemm.utils.math import ceil_div
def init_test_dist(local_rank_arg: int = None) -> Tuple[int, int, dist.ProcessGroup]:
local_rank = local_rank_arg if local_rank_arg is not None else int(os.environ.get('LOCAL_RANK', '0'))
rank = int(os.environ.get('RANK', '0'))
world_size = int(os.environ.get('WORLD_SIZE', '1'))
master_addr = os.environ.get('MASTER_ADDR', '127.0.0.1')
master_port = int(os.environ.get('MASTER_PORT', '8364'))
torch.cuda.set_device(local_rank)
sig = inspect.signature(dist.init_process_group)
params = {
'backend': 'nccl',
'init_method': f'tcp://{master_addr}:{master_port}',
'world_size': world_size,
'rank': rank,
}
if 'device_id' in sig.parameters:
params['device_id'] = torch.device(f'cuda:{local_rank}')
dist.init_process_group(**params)
torch.set_default_device('cuda')
return rank, world_size, dist.new_group(list(range(world_size)))
def make_weights(num_experts_per_rank: int, hidden: int, intermediate_hidden: int):
torch.manual_seed(2345 + dist.get_rank())
l1_weights = (torch.randn(
(num_experts_per_rank, intermediate_hidden * 2, hidden),
dtype=torch.float32, device='cuda') * 0.25).to(torch.float8_e4m3fn)
l2_weights = (torch.randn(
(num_experts_per_rank, hidden, intermediate_hidden),
dtype=torch.float32, device='cuda') * 0.25).to(torch.float8_e4m3fn)
l1_weights_sf = torch.empty(
(num_experts_per_rank, ceil_div(intermediate_hidden * 2, 128), hidden // 128),
dtype=torch.float32, device='cuda')
l2_weights_sf = torch.empty(
(num_experts_per_rank, ceil_div(hidden, 128), intermediate_hidden // 128),
dtype=torch.float32, device='cuda')
for expert in range(num_experts_per_rank):
for n_group in range(l1_weights_sf.shape[1]):
for k_group in range(l1_weights_sf.shape[2]):
l1_weights_sf[expert, n_group, k_group] = 0.5 + 0.125 * expert + 0.0625 * n_group + 0.03125 * k_group
for n_group in range(l2_weights_sf.shape[1]):
for k_group in range(l2_weights_sf.shape[2]):
l2_weights_sf[expert, n_group, k_group] = 0.75 + 0.125 * expert + 0.0625 * n_group + 0.03125 * k_group
return deep_gemm.transform_weights_for_mega_moe(
(l1_weights, l1_weights_sf), (l2_weights, l2_weights_sf))
def dequant_l2_acts(buffer: deep_gemm.SymmBuffer,
start: int,
count: int,
intermediate_hidden: int) -> torch.Tensor:
x = buffer.l2_acts[start:start + count, :intermediate_hidden].to(torch.float32)
out = torch.empty_like(x)
for sf_group in range(intermediate_hidden // 64):
col_start = sf_group * 64
col_end = col_start + 64
sf = buffer.l2_acts_sf[start:start + count, sf_group].to(torch.float32)
out[:, col_start:col_end] = x[:, col_start:col_end] * sf[:, None]
return out
def verify_l2_scatter(buffer: deep_gemm.SymmBuffer,
l2_weights: torch.Tensor,
l2_weights_sf: torch.Tensor,
hidden: int,
intermediate_hidden: int,
num_tokens: int,
atol: float,
rtol: float) -> None:
block_m = 128
ref = torch.zeros((num_tokens, hidden), dtype=torch.bfloat16, device='cuda')
counts = [int(v) & 0xffffffff for v in buffer.expert_recv_count_sum.to(torch.int64).cpu().tolist()]
pool_block_offset = 0
expected_mask = (1 << (intermediate_hidden // 64)) - 1
for expert, count in enumerate(counts):
if count == 0:
continue
pool_start = pool_block_offset * block_m
x = dequant_l2_acts(buffer, pool_start, count, intermediate_hidden)
expected = torch.zeros((count, hidden), dtype=torch.float32, device='cuda')
for k_group in range(intermediate_hidden // 128):
k_start = k_group * 128
k_end = k_start + 128
w = l2_weights[expert, :, k_start:k_end].to(torch.float32)
partial = x[:, k_start:k_end] @ w.t()
for n_group in range(hidden // 128):
n_start = n_group * 128
n_end = n_start + 128
sfb = l2_weights_sf[expert, n_group, k_group].to(torch.float32)
expected[:, n_start:n_end] += partial[:, n_start:n_end] * sfb
metadata = buffer.token_src_metadata[pool_start:pool_start + count].to(torch.int64)
for row in range(count):
rank_idx, token_idx, topk_idx = [int(v) for v in metadata[row].tolist()]
assert rank_idx == dist.get_rank()
assert topk_idx == 0
ref[token_idx] = expected[row].to(torch.bfloat16)
for block in range(ceil_div(count, block_m)):
mask = int(buffer.l2_arrival_mask[pool_block_offset + block].item())
assert mask == expected_mask, (expert, block, mask, expected_mask)
pool_block_offset += ceil_div(count, block_m)
actual = buffer.combine_acts[0, :num_tokens, :hidden]
torch.testing.assert_close(actual.cpu(), ref.cpu(), rtol=rtol, atol=atol)
def run_case(args: argparse.Namespace, group: dist.ProcessGroup, rank_idx: int, num_ranks: int) -> None:
assert num_ranks == 1, 'Phase 5 milestone verifies multi-expert single-rank scatter first'
hidden = args.hidden
intermediate_hidden = args.intermediate_hidden
num_tokens = args.num_tokens
num_topk = 1
num_experts = args.num_experts
num_experts_per_rank = num_experts // num_ranks
assert num_experts % num_ranks == 0
buffer = deep_gemm.get_symm_buffer_for_mega_moe(
group, num_experts, args.num_max_tokens_per_rank, num_topk,
hidden, intermediate_hidden)
weights = make_weights(num_experts_per_rank, hidden, intermediate_hidden)
torch.manual_seed(6789 + rank_idx)
x = (torch.randn((num_tokens, hidden), dtype=torch.float32, device='cuda') * 0.25).to(torch.float8_e4m3fn)
x_sf = torch.rand((num_tokens, hidden // 128), dtype=torch.float32, device='cuda') * 0.25 + 0.875
topk_idx = (torch.arange(num_tokens, dtype=torch.long, device='cuda') // 128).reshape(num_tokens, 1)
topk_idx = torch.clamp(topk_idx, max=num_experts - 1)
topk_weights = torch.linspace(0.75, 1.25, num_tokens, dtype=torch.float32, device='cuda').reshape(num_tokens, 1)
buffer.x[:num_tokens].copy_(x)
buffer.x_sf[:num_tokens].copy_(x_sf)
buffer.topk_idx[:num_tokens].copy_(topk_idx)
buffer.topk_weights[:num_tokens].copy_(topk_weights)
torch.cuda.synchronize()
dist.barrier(group=group)
y = torch.empty((num_tokens, hidden), dtype=torch.bfloat16, device='cuda')
deep_gemm.fp8_mega_moe(y, weights[0], weights[1], buffer,
activation_clamp=args.activation_clamp,
fast_math=False)
torch.cuda.synchronize()
verify_l2_scatter(buffer, weights[1][0], weights[1][1], hidden, intermediate_hidden,
num_tokens, args.atol, args.rtol)
dist.barrier(group=group)
if rank_idx == 0:
print('[PASSED] Phase 5 L2 GEMM scatter correctness', flush=True)
buffer.destroy()
def main() -> None:
parser = argparse.ArgumentParser(description='SM90 MegaMoE Phase 5 L2 GEMM scatter correctness')
parser.add_argument('--num-tokens', type=int, default=256)
parser.add_argument('--num-max-tokens-per-rank', type=int, default=384)
parser.add_argument('--hidden', type=int, default=256)
parser.add_argument('--intermediate-hidden', type=int, default=128)
parser.add_argument('--num-experts', type=int, default=2)
parser.add_argument('--activation-clamp', type=float, default=None)
parser.add_argument('--local-rank', type=int, default=None)
parser.add_argument('--atol', type=float, default=3e-2)
parser.add_argument('--rtol', type=float, default=5e-2)
args = parser.parse_args()
rank_idx, num_ranks, group = init_test_dist(args.local_rank)
assert torch.cuda.get_device_capability(torch.cuda.current_device())[0] == 9
assert args.num_tokens == 256
assert args.hidden % 128 == 0 and args.intermediate_hidden % 128 == 0
run_case(args, group, rank_idx, num_ranks)
dist.destroy_process_group()
if __name__ == '__main__':
main()