feat: implement sm90 megamoe phase5 l2 scatter
This commit is contained in:
@@ -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))
|
||||
|
||||
201
megamoe_dev_test_scripts/phase5/l2_gemm_scatter_correctness.py
Normal file
201
megamoe_dev_test_scripts/phase5/l2_gemm_scatter_correctness.py
Normal 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()
|
||||
Reference in New Issue
Block a user