feat: add sm90 megamoe phase1 interfaces
This commit is contained in:
+116
-31
@@ -13,6 +13,14 @@ except Exception as exception:
|
||||
from .. import _C
|
||||
|
||||
|
||||
def _is_sm90() -> bool:
|
||||
return torch.cuda.get_device_capability(torch.cuda.current_device())[0] == 9
|
||||
|
||||
|
||||
def _is_sm100() -> bool:
|
||||
return torch.cuda.get_device_capability(torch.cuda.current_device())[0] == 10
|
||||
|
||||
|
||||
class SymmBuffer:
|
||||
def __init__(self, group: dist.ProcessGroup,
|
||||
# MoE arguments
|
||||
@@ -28,13 +36,23 @@ class SymmBuffer:
|
||||
self.hidden = hidden
|
||||
self.intermediate_hidden = intermediate_hidden
|
||||
|
||||
# Allocate a symmetric buffer
|
||||
num_bytes, slice_input_buffers = _C.get_symm_buffer_size_for_mega_moe(
|
||||
group.size(), num_experts,
|
||||
num_max_tokens_per_rank, num_topk,
|
||||
hidden, intermediate_hidden,
|
||||
use_fp8_dispatch, activation
|
||||
)
|
||||
# Allocate a symmetric buffer (route by architecture)
|
||||
if _is_sm90():
|
||||
num_bytes, slice_input_buffers = _C.get_symm_buffer_size_for_sm90_mega_moe(
|
||||
group.size(), num_experts,
|
||||
num_max_tokens_per_rank, num_topk,
|
||||
hidden, intermediate_hidden,
|
||||
use_fp8_dispatch, activation
|
||||
)
|
||||
elif _is_sm100():
|
||||
num_bytes, slice_input_buffers = _C.get_symm_buffer_size_for_mega_moe(
|
||||
group.size(), num_experts,
|
||||
num_max_tokens_per_rank, num_topk,
|
||||
hidden, intermediate_hidden,
|
||||
use_fp8_dispatch, activation
|
||||
)
|
||||
else:
|
||||
raise RuntimeError('Unsupported architecture for MegaMoE')
|
||||
self.buffer = symm_mem.empty(num_bytes, dtype=torch.int8, device='cuda')
|
||||
self.handle = symm_mem.rendezvous(self.buffer, group=group)
|
||||
self.buffer.zero_()
|
||||
@@ -46,6 +64,10 @@ class SymmBuffer:
|
||||
self.topk_idx, self.topk_weights,
|
||||
self.l1_acts, self.l1_acts_sf,
|
||||
self.l2_acts, self.l2_acts_sf) = slice_input_buffers(self.buffer)
|
||||
self.l1_topk_weights = None
|
||||
self.expert_recv_count_sum = None
|
||||
self.l1_arrival_count = None
|
||||
self.token_src_metadata = None
|
||||
|
||||
def destroy(self):
|
||||
self.handle = None
|
||||
@@ -53,6 +75,16 @@ class SymmBuffer:
|
||||
self.group = None
|
||||
self.x = None
|
||||
self.x_sf = None
|
||||
self.topk_idx = None
|
||||
self.topk_weights = None
|
||||
self.l1_acts = None
|
||||
self.l1_acts_sf = None
|
||||
self.l1_topk_weights = None
|
||||
self.l2_acts = None
|
||||
self.l2_acts_sf = None
|
||||
self.expert_recv_count_sum = None
|
||||
self.l1_arrival_count = None
|
||||
self.token_src_metadata = None
|
||||
|
||||
|
||||
def get_symm_buffer_for_mega_moe(group: dist.ProcessGroup,
|
||||
@@ -62,7 +94,13 @@ def get_symm_buffer_for_mega_moe(group: dist.ProcessGroup,
|
||||
use_fp8_dispatch: bool = True,
|
||||
activation: str = 'swiglu') -> SymmBuffer:
|
||||
# Token count must be aligned to block sizes
|
||||
num_max_tokens_per_rank = align(num_max_tokens_per_rank, _C.get_token_alignment_for_mega_moe())
|
||||
if _is_sm90():
|
||||
alignment = _C.get_token_alignment_for_sm90_mega_moe()
|
||||
elif _is_sm100():
|
||||
alignment = _C.get_token_alignment_for_mega_moe()
|
||||
else:
|
||||
raise RuntimeError('Unsupported architecture for MegaMoE')
|
||||
num_max_tokens_per_rank = align(num_max_tokens_per_rank, alignment)
|
||||
|
||||
return SymmBuffer(
|
||||
group, num_experts,
|
||||
@@ -72,16 +110,17 @@ def get_symm_buffer_for_mega_moe(group: dist.ProcessGroup,
|
||||
)
|
||||
|
||||
|
||||
def _interleave_l1_weights(l1_weights: Tuple[torch.Tensor, torch.Tensor]) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
def _interleave_l1_weight_tensor(t: torch.Tensor, gran: int = 8) -> torch.Tensor:
|
||||
# [gate: 0..7, up: 0..7, gate: 8..15, up: 8..15, ...] instead of [gate | up]
|
||||
def interleave(t, gran: int = 8) -> torch.Tensor:
|
||||
g, n, *rest = t.shape
|
||||
half = n // 2
|
||||
gate = t[:, :half].reshape(g, half // gran, gran, *rest)
|
||||
up = t[:, half:].reshape(g, half // gran, gran, *rest)
|
||||
return torch.empty_like(t).copy_(torch.stack([gate, up], dim=2).reshape(g, n, *rest))
|
||||
g, n, *rest = t.shape
|
||||
half = n // 2
|
||||
gate = t[:, :half].reshape(g, half // gran, gran, *rest)
|
||||
up = t[:, half:].reshape(g, half // gran, gran, *rest)
|
||||
return torch.empty_like(t).copy_(torch.stack([gate, up], dim=2).reshape(g, n, *rest))
|
||||
|
||||
return interleave(l1_weights[0]), interleave(l1_weights[1])
|
||||
|
||||
def _interleave_l1_weights(l1_weights: Tuple[torch.Tensor, torch.Tensor]) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
return _interleave_l1_weight_tensor(l1_weights[0]), _interleave_l1_weight_tensor(l1_weights[1])
|
||||
|
||||
|
||||
def _transpose_sf_for_utccp(sf: torch.Tensor) -> torch.Tensor:
|
||||
@@ -93,36 +132,82 @@ def _transpose_sf_for_utccp(sf: torch.Tensor) -> torch.Tensor:
|
||||
return torch.empty_like(sf).copy_(result)
|
||||
|
||||
|
||||
def transform_weights_for_mega_moe_sm90(
|
||||
l1_weights: Tuple[torch.Tensor, torch.Tensor],
|
||||
l2_weights: Tuple[torch.Tensor, torch.Tensor]
|
||||
) -> Tuple[Tuple[torch.Tensor, torch.Tensor], Tuple[torch.Tensor, torch.Tensor]]:
|
||||
# L1: interleave FP8 gate/up weights only; SM90 float weight SF stays natural MN-major.
|
||||
l1_weights = (_interleave_l1_weight_tensor(l1_weights[0]), l1_weights[1])
|
||||
# L2: no transform
|
||||
return l1_weights, l2_weights
|
||||
|
||||
|
||||
def transform_weights_for_mega_moe(
|
||||
l1_weights: Tuple[torch.Tensor, torch.Tensor],
|
||||
l2_weights: Tuple[torch.Tensor, torch.Tensor]
|
||||
) -> Tuple[Tuple[torch.Tensor, torch.Tensor], Tuple[torch.Tensor, torch.Tensor]]:
|
||||
# L1: interleave gate/up, then transpose SF for UTCCP
|
||||
if _is_sm90():
|
||||
return transform_weights_for_mega_moe_sm90(l1_weights, l2_weights)
|
||||
# SM100: L1 interleave gate/up + UTCCP SF transpose, L2 UTCCP SF transpose
|
||||
l1_interleaved = _interleave_l1_weights(l1_weights)
|
||||
l1_weights = (l1_interleaved[0], _transpose_sf_for_utccp(l1_interleaved[1]))
|
||||
# L2: only transpose SF for UTCCP
|
||||
l2_weights = (l2_weights[0], _transpose_sf_for_utccp(l2_weights[1]))
|
||||
return l1_weights, l2_weights
|
||||
|
||||
|
||||
def fp8_mega_moe(y: torch.Tensor,
|
||||
l1_weights: Tuple[torch.Tensor, torch.Tensor],
|
||||
l2_weights: Tuple[torch.Tensor, torch.Tensor],
|
||||
sym_buffer: SymmBuffer,
|
||||
cumulative_local_expert_recv_stats: Optional[torch.Tensor] = None,
|
||||
recipe: Optional[Tuple[int, int, int]] = None,
|
||||
activation: str = 'swiglu',
|
||||
activation_clamp: Optional[float] = None,
|
||||
fast_math: bool = True):
|
||||
if _is_sm90():
|
||||
if recipe is None:
|
||||
recipe = (1, 128, 128)
|
||||
_C.fp8_mega_moe(
|
||||
y,
|
||||
l1_weights, l2_weights,
|
||||
cumulative_local_expert_recv_stats,
|
||||
sym_buffer.buffer,
|
||||
sym_buffer.handle.buffer_ptrs, sym_buffer.group.rank(),
|
||||
sym_buffer.num_max_tokens_per_rank,
|
||||
sym_buffer.num_experts, sym_buffer.num_topk,
|
||||
recipe,
|
||||
activation, activation_clamp,
|
||||
fast_math
|
||||
)
|
||||
elif _is_sm100():
|
||||
if recipe is None:
|
||||
recipe = (1, 1, 32)
|
||||
_C.fp8_fp4_mega_moe(
|
||||
y,
|
||||
l1_weights, l2_weights,
|
||||
cumulative_local_expert_recv_stats,
|
||||
sym_buffer.buffer,
|
||||
sym_buffer.handle.buffer_ptrs, sym_buffer.group.rank(),
|
||||
sym_buffer.num_max_tokens_per_rank,
|
||||
sym_buffer.num_experts, sym_buffer.num_topk,
|
||||
recipe,
|
||||
activation, activation_clamp,
|
||||
fast_math
|
||||
)
|
||||
else:
|
||||
raise RuntimeError('Unsupported architecture for MegaMoE')
|
||||
|
||||
|
||||
# Backward-compatible alias
|
||||
def fp8_fp4_mega_moe(y: torch.Tensor,
|
||||
l1_weights: Tuple[torch.Tensor, torch.Tensor],
|
||||
l2_weights: Tuple[torch.Tensor, torch.Tensor],
|
||||
sym_buffer: SymmBuffer,
|
||||
cumulative_local_expert_recv_stats: Optional[torch.Tensor] = None,
|
||||
recipe: Tuple[int, int, int] = (1, 1, 32),
|
||||
recipe: Optional[Tuple[int, int, int]] = None,
|
||||
activation: str = 'swiglu',
|
||||
activation_clamp: Optional[float] = None,
|
||||
fast_math: bool = True):
|
||||
_C.fp8_fp4_mega_moe(
|
||||
y,
|
||||
l1_weights, l2_weights,
|
||||
cumulative_local_expert_recv_stats,
|
||||
sym_buffer.buffer,
|
||||
sym_buffer.handle.buffer_ptrs, sym_buffer.group.rank(),
|
||||
sym_buffer.num_max_tokens_per_rank,
|
||||
sym_buffer.num_experts, sym_buffer.num_topk,
|
||||
recipe,
|
||||
activation, activation_clamp,
|
||||
fast_math
|
||||
)
|
||||
fp8_mega_moe(y, l1_weights, l2_weights, sym_buffer,
|
||||
cumulative_local_expert_recv_stats, recipe,
|
||||
activation, activation_clamp, fast_math)
|
||||
|
||||
Reference in New Issue
Block a user