[perf] MoE intermediate cache pre-allocation: eliminate 189 CUDA mallocs per decode step
fused_experts() is called 64 times per decode step (once per MoE layer). Each call allocated 3 intermediate tensors via torch.empty = 192 mallocs. For decode (M=1, topk=8), all 64 calls use identical shapes. Fix: module-level _moe_intermediate_cache dict that reuses tensors when shapes match. First layer call allocates, subsequent 63 calls reuse. Saves 189 CUDA mallocs per decode step = 74,655 mallocs/second at 395 TPS. Design follows CCCL's dispatch_reduce.cuh pattern: pre-allocate temp_storage once via alias_temporaries, reuse across kernel invocations. No functional change — tensors are .empty() (uninitialized), overwritten before use by ixformer kernels.
This commit is contained in:
@@ -15,6 +15,25 @@ from vllm.platforms import current_platform
|
|||||||
|
|
||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
|
# Module-level cache for MoE intermediate tensors.
|
||||||
|
# Eliminates 192 torch.empty (CUDA malloc) calls per decode step by reusing
|
||||||
|
# buffers across the 64 MoE layer invocations within a single forward pass.
|
||||||
|
# Key insight from CCCL dispatch_reduce.cuh: NVIDIA pre-allocates temp_storage
|
||||||
|
# once and reuses it across kernel invocations rather than re-allocating.
|
||||||
|
_moe_intermediate_cache = {}
|
||||||
|
|
||||||
|
|
||||||
|
def _get_or_alloc(name, shape, dtype, device):
|
||||||
|
"""Get a cached tensor or allocate a new one. Reuses if shape fits."""
|
||||||
|
key = (name, dtype, device)
|
||||||
|
cached = _moe_intermediate_cache.get(key)
|
||||||
|
if cached is not None and cached.shape == shape:
|
||||||
|
return cached
|
||||||
|
# Shape changed (different M, different topk) — reallocate
|
||||||
|
t = torch.empty(shape, dtype=dtype, device=device)
|
||||||
|
_moe_intermediate_cache[key] = t
|
||||||
|
return t
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
@triton.jit
|
||||||
def fused_moe_kernel(
|
def fused_moe_kernel(
|
||||||
@@ -534,15 +553,15 @@ def fused_experts(hidden_states: torch.Tensor,
|
|||||||
|
|
||||||
config = get_config_func(M)
|
config = get_config_func(M)
|
||||||
|
|
||||||
intermediate_cache1 = torch.empty((M, topk_ids.shape[1], N),
|
intermediate_cache1 = _get_or_alloc(
|
||||||
device=hidden_states.device,
|
'moe_c1', (M, topk_ids.shape[1], N),
|
||||||
dtype=hidden_states.dtype)
|
hidden_states.dtype, hidden_states.device)
|
||||||
intermediate_cache2 = torch.empty((M * topk_ids.shape[1], N // 2),
|
intermediate_cache2 = _get_or_alloc(
|
||||||
device=hidden_states.device,
|
'moe_c2', (M * topk_ids.shape[1], N // 2),
|
||||||
dtype=hidden_states.dtype)
|
hidden_states.dtype, hidden_states.device)
|
||||||
intermediate_cache3 = torch.empty((M, topk_ids.shape[1], w2.shape[1]),
|
intermediate_cache3 = _get_or_alloc(
|
||||||
device=hidden_states.device,
|
'moe_c3', (M, topk_ids.shape[1], w2.shape[1]),
|
||||||
dtype=hidden_states.dtype)
|
hidden_states.dtype, hidden_states.device)
|
||||||
|
|
||||||
compute_type = (tl.bfloat16
|
compute_type = (tl.bfloat16
|
||||||
if hidden_states.dtype == torch.bfloat16 else tl.float16)
|
if hidden_states.dtype == torch.bfloat16 else tl.float16)
|
||||||
|
|||||||
Reference in New Issue
Block a user