初始化项目,由ModelHub XC社区提供模型
Model: ayh015/myLightningOPD Source: Original Platform
This commit is contained in:
384
slime/backends/fsdp_utils/kernels/fused_experts.py
Normal file
384
slime/backends/fsdp_utils/kernels/fused_experts.py
Normal file
@@ -0,0 +1,384 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
import triton.language as tl
|
||||
from sglang.srt.layers.moe.fused_moe_triton.fused_moe import (
|
||||
invoke_fused_moe_kernel,
|
||||
moe_align_block_size,
|
||||
moe_sum_reduce,
|
||||
silu_and_mul,
|
||||
)
|
||||
|
||||
from .fused_moe_triton_backward_kernels import invoke_fused_moe_backward_kernel
|
||||
|
||||
|
||||
class GateUpProjFunction(torch.autograd.Function):
|
||||
@staticmethod
|
||||
def forward(
|
||||
ctx,
|
||||
hidden_states: torch.Tensor,
|
||||
w1: torch.Tensor,
|
||||
topk_weights: torch.Tensor,
|
||||
topk_ids: torch.Tensor,
|
||||
):
|
||||
num_tokens, _ = hidden_states.shape
|
||||
E, N, _ = w1.shape
|
||||
# We execute the fused_moe kernel in chunks to circumvent this issue:
|
||||
# https://github.com/vllm-project/vllm/issues/5938
|
||||
CHUNK_SIZE = 64 * 1024
|
||||
|
||||
# default deterministic config
|
||||
config = {
|
||||
"BLOCK_SIZE_M": 64,
|
||||
"BLOCK_SIZE_N": 64,
|
||||
"BLOCK_SIZE_K": 32,
|
||||
"GROUP_SIZE_M": 8,
|
||||
}
|
||||
|
||||
topk = topk_ids.shape[1]
|
||||
|
||||
intermediate_cache1 = torch.empty(
|
||||
(num_tokens * topk, N),
|
||||
device=hidden_states.device,
|
||||
dtype=hidden_states.dtype,
|
||||
)
|
||||
|
||||
for chunk in range((num_tokens // CHUNK_SIZE) + 1):
|
||||
begin_chunk_idx, end_chunk_idx = (
|
||||
chunk * CHUNK_SIZE,
|
||||
min((chunk + 1) * CHUNK_SIZE, num_tokens),
|
||||
)
|
||||
curr_hidden_states = hidden_states[begin_chunk_idx:end_chunk_idx]
|
||||
cur_intermediate_cache1 = intermediate_cache1[begin_chunk_idx * topk : end_chunk_idx * topk]
|
||||
|
||||
curr_topk_ids = topk_ids[begin_chunk_idx:end_chunk_idx]
|
||||
curr_topk_weights = topk_weights[begin_chunk_idx:end_chunk_idx]
|
||||
|
||||
sorted_token_ids, expert_ids, num_tokens_post_padded = moe_align_block_size(
|
||||
curr_topk_ids, config["BLOCK_SIZE_M"], E
|
||||
)
|
||||
|
||||
invoke_fused_moe_kernel(
|
||||
curr_hidden_states,
|
||||
w1,
|
||||
None,
|
||||
cur_intermediate_cache1,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
curr_topk_weights,
|
||||
curr_topk_ids,
|
||||
sorted_token_ids,
|
||||
expert_ids,
|
||||
num_tokens_post_padded,
|
||||
False,
|
||||
topk_ids.shape[1],
|
||||
config,
|
||||
compute_type=tl.bfloat16,
|
||||
use_fp8_w8a8=False,
|
||||
use_int8_w8a8=False,
|
||||
use_int8_w8a16=False,
|
||||
use_int4_w4a16=False,
|
||||
per_channel_quant=False,
|
||||
block_shape=None,
|
||||
c_sorted=False,
|
||||
filter_expert=True,
|
||||
)
|
||||
|
||||
ctx.save_for_backward(hidden_states, w1, topk_weights, topk_ids)
|
||||
ctx.config = config
|
||||
ctx.num_tokens = num_tokens
|
||||
ctx.topk = topk
|
||||
|
||||
return intermediate_cache1
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
"""
|
||||
Backward pass for GateUpProjFunction using Triton kernels.
|
||||
|
||||
Args:
|
||||
grad_output: shape (num_tokens * topk, N)
|
||||
|
||||
Returns:
|
||||
(grad_hidden_states, grad_w1, grad_topk_weights, None)
|
||||
"""
|
||||
|
||||
hidden_states, w1, topk_weights, topk_ids = ctx.saved_tensors
|
||||
config = ctx.config
|
||||
num_tokens = ctx.num_tokens
|
||||
topk = ctx.topk
|
||||
|
||||
E, N, D_in = w1.shape
|
||||
CHUNK_SIZE = 64 * 1024
|
||||
|
||||
# Initialize gradient tensors
|
||||
grad_hidden_states = torch.zeros_like(hidden_states)
|
||||
grad_w1 = torch.zeros_like(w1)
|
||||
# GateUpProj stage doesn't need topk_weights gradient
|
||||
grad_topk_weights = torch.zeros_like(topk_weights)
|
||||
|
||||
# Process in chunks to match forward pass
|
||||
for chunk in range((num_tokens // CHUNK_SIZE) + 1):
|
||||
begin_chunk_idx, end_chunk_idx = (
|
||||
chunk * CHUNK_SIZE,
|
||||
min((chunk + 1) * CHUNK_SIZE, num_tokens),
|
||||
)
|
||||
|
||||
curr_num_tokens = end_chunk_idx - begin_chunk_idx
|
||||
if curr_num_tokens == 0:
|
||||
continue
|
||||
|
||||
curr_hidden_states = hidden_states[begin_chunk_idx:end_chunk_idx]
|
||||
curr_topk_ids = topk_ids[begin_chunk_idx:end_chunk_idx]
|
||||
curr_topk_weights = topk_weights[begin_chunk_idx:end_chunk_idx]
|
||||
curr_grad_output = grad_output[begin_chunk_idx * topk : end_chunk_idx * topk]
|
||||
|
||||
# Get aligned metadata
|
||||
sorted_token_ids, expert_ids, num_tokens_post_padded = moe_align_block_size(
|
||||
curr_topk_ids, config["BLOCK_SIZE_M"], E
|
||||
)
|
||||
|
||||
# Prepare gradient buffer for this chunk
|
||||
curr_grad_hidden_states = torch.zeros_like(curr_hidden_states)
|
||||
curr_grad_w1 = torch.zeros_like(w1)
|
||||
|
||||
# Call Triton backward kernel with MUL_ROUTED_WEIGHT=False
|
||||
# Use chunk of hidden_states to match sorted_token_ids indices
|
||||
invoke_fused_moe_backward_kernel(
|
||||
grad_output=curr_grad_output,
|
||||
input=curr_hidden_states, # Use chunk of hidden_states to match sorted_token_ids
|
||||
weight=w1,
|
||||
grad_input=curr_grad_hidden_states,
|
||||
grad_weight=curr_grad_w1,
|
||||
grad_topk_weights=None, # Not needed for GateUpProj
|
||||
topk_weights=curr_topk_weights,
|
||||
topk_ids=curr_topk_ids,
|
||||
sorted_token_ids=sorted_token_ids,
|
||||
expert_ids=expert_ids,
|
||||
num_tokens_post_padded=num_tokens_post_padded,
|
||||
mul_routed_weight=False,
|
||||
top_k=topk,
|
||||
config=config,
|
||||
compute_type=tl.bfloat16,
|
||||
)
|
||||
|
||||
# Accumulate gradients
|
||||
grad_hidden_states[begin_chunk_idx:end_chunk_idx] += curr_grad_hidden_states
|
||||
grad_w1 += curr_grad_w1
|
||||
|
||||
return grad_hidden_states, grad_w1, grad_topk_weights, None
|
||||
|
||||
|
||||
class SiluAndMulFunction(torch.autograd.Function):
|
||||
@staticmethod
|
||||
def forward(ctx, intermediate_cache1: torch.Tensor):
|
||||
num_tokens, N = intermediate_cache1.shape
|
||||
intermediate_cache2 = torch.empty(
|
||||
(num_tokens, N // 2),
|
||||
device=intermediate_cache1.device,
|
||||
dtype=intermediate_cache1.dtype,
|
||||
)
|
||||
silu_and_mul(intermediate_cache1.view(-1, N), intermediate_cache2)
|
||||
|
||||
ctx.save_for_backward(intermediate_cache1)
|
||||
return intermediate_cache2
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
(intermediate_cache1,) = ctx.saved_tensors
|
||||
N = intermediate_cache1.shape[-1]
|
||||
x1, x2 = intermediate_cache1.view(-1, N).chunk(2, dim=-1)
|
||||
silu_x1 = torch.nn.functional.silu(x1)
|
||||
|
||||
sig = torch.sigmoid(x1)
|
||||
dsilu_dx1 = sig + x1 * sig * (1 - sig)
|
||||
grad_x1 = grad_output * x2 * dsilu_dx1
|
||||
grad_x2 = grad_output * silu_x1
|
||||
grad_input = torch.cat([grad_x1, grad_x2], dim=-1)
|
||||
|
||||
return grad_input.view_as(intermediate_cache1)
|
||||
|
||||
|
||||
class DownProjFunction(torch.autograd.Function):
|
||||
@staticmethod
|
||||
def forward(
|
||||
ctx,
|
||||
intermediate_cache2: torch.Tensor,
|
||||
w2: torch.Tensor,
|
||||
topk_weights: torch.Tensor,
|
||||
topk_ids: torch.Tensor,
|
||||
):
|
||||
num_tokens, _ = intermediate_cache2.shape
|
||||
topk = topk_ids.shape[1]
|
||||
num_tokens //= topk
|
||||
E, _, _ = w2.shape
|
||||
# We execute the fused_moe kernel in chunks to circumvent this issue:
|
||||
# https://github.com/vllm-project/vllm/issues/5938
|
||||
CHUNK_SIZE = 64 * 1024
|
||||
|
||||
# default deterministic config
|
||||
config = {
|
||||
"BLOCK_SIZE_M": 64,
|
||||
"BLOCK_SIZE_N": 64,
|
||||
"BLOCK_SIZE_K": 32,
|
||||
"GROUP_SIZE_M": 8,
|
||||
}
|
||||
|
||||
intermediate_cache3 = torch.empty(
|
||||
(num_tokens, topk, w2.shape[1]),
|
||||
device=intermediate_cache2.device,
|
||||
dtype=intermediate_cache2.dtype,
|
||||
)
|
||||
|
||||
for chunk in range((num_tokens // CHUNK_SIZE) + 1):
|
||||
begin_chunk_idx, end_chunk_idx = (
|
||||
chunk * CHUNK_SIZE,
|
||||
min((chunk + 1) * CHUNK_SIZE, num_tokens),
|
||||
)
|
||||
cur_intermediate_cache2 = intermediate_cache2[begin_chunk_idx * topk : end_chunk_idx * topk]
|
||||
cur_intermediate_cache3 = intermediate_cache3[begin_chunk_idx:end_chunk_idx]
|
||||
|
||||
curr_topk_ids = topk_ids[begin_chunk_idx:end_chunk_idx]
|
||||
curr_topk_weights = topk_weights[begin_chunk_idx:end_chunk_idx]
|
||||
|
||||
sorted_token_ids, expert_ids, num_tokens_post_padded = moe_align_block_size(
|
||||
curr_topk_ids, config["BLOCK_SIZE_M"], E
|
||||
)
|
||||
invoke_fused_moe_kernel(
|
||||
cur_intermediate_cache2,
|
||||
w2,
|
||||
None,
|
||||
cur_intermediate_cache3,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
curr_topk_weights,
|
||||
curr_topk_ids,
|
||||
sorted_token_ids,
|
||||
expert_ids,
|
||||
num_tokens_post_padded,
|
||||
True,
|
||||
1,
|
||||
config,
|
||||
compute_type=tl.bfloat16,
|
||||
use_fp8_w8a8=False,
|
||||
use_int8_w8a8=False,
|
||||
use_int8_w8a16=False,
|
||||
use_int4_w4a16=False,
|
||||
per_channel_quant=False,
|
||||
block_shape=None,
|
||||
a_use_tma=False,
|
||||
b_use_tma=False,
|
||||
)
|
||||
|
||||
ctx.save_for_backward(intermediate_cache2, w2, topk_weights, topk_ids)
|
||||
ctx.config = config
|
||||
ctx.num_tokens = num_tokens
|
||||
ctx.topk = topk
|
||||
|
||||
return intermediate_cache3
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
"""
|
||||
Backward pass for DownProjFunction using Triton kernels.
|
||||
|
||||
Args:
|
||||
grad_output: shape (num_tokens, topk, hidden_size)
|
||||
|
||||
Returns:
|
||||
(grad_intermediate_cache2, grad_w2, grad_topk_weights, None)
|
||||
"""
|
||||
intermediate_cache2, w2, topk_weights, topk_ids = ctx.saved_tensors
|
||||
config = ctx.config
|
||||
num_tokens = ctx.num_tokens
|
||||
topk = ctx.topk
|
||||
|
||||
E, hidden_size, intermediate_size = w2.shape
|
||||
CHUNK_SIZE = 64 * 1024
|
||||
|
||||
# Initialize gradient tensors
|
||||
grad_intermediate_cache2 = torch.zeros_like(intermediate_cache2)
|
||||
grad_w2 = torch.zeros_like(w2)
|
||||
grad_topk_weights = torch.zeros_like(topk_weights)
|
||||
|
||||
# Process in chunks to match forward pass
|
||||
for chunk in range((num_tokens // CHUNK_SIZE) + 1):
|
||||
begin_chunk_idx, end_chunk_idx = (
|
||||
chunk * CHUNK_SIZE,
|
||||
min((chunk + 1) * CHUNK_SIZE, num_tokens),
|
||||
)
|
||||
|
||||
curr_num_tokens = end_chunk_idx - begin_chunk_idx
|
||||
if curr_num_tokens == 0:
|
||||
continue
|
||||
|
||||
curr_intermediate_cache2 = intermediate_cache2[begin_chunk_idx * topk : end_chunk_idx * topk]
|
||||
curr_topk_ids = topk_ids[begin_chunk_idx:end_chunk_idx]
|
||||
curr_topk_weights = topk_weights[begin_chunk_idx:end_chunk_idx]
|
||||
curr_grad_output = grad_output[begin_chunk_idx:end_chunk_idx]
|
||||
|
||||
# Get aligned metadata
|
||||
sorted_token_ids, expert_ids, num_tokens_post_padded = moe_align_block_size(
|
||||
curr_topk_ids, config["BLOCK_SIZE_M"], E
|
||||
)
|
||||
|
||||
# Prepare gradient buffers for this chunk
|
||||
curr_grad_intermediate_cache2 = torch.zeros_like(curr_intermediate_cache2)
|
||||
curr_grad_w2 = torch.zeros_like(w2)
|
||||
curr_grad_topk_weights = torch.zeros_like(curr_topk_weights)
|
||||
|
||||
# Call Triton backward kernel with MUL_ROUTED_WEIGHT=True
|
||||
# Note: Use top_k=1 to match forward pass indexing
|
||||
invoke_fused_moe_backward_kernel(
|
||||
grad_output=curr_grad_output,
|
||||
input=curr_intermediate_cache2,
|
||||
weight=w2,
|
||||
grad_input=curr_grad_intermediate_cache2,
|
||||
grad_weight=curr_grad_w2,
|
||||
grad_topk_weights=curr_grad_topk_weights,
|
||||
topk_weights=curr_topk_weights,
|
||||
topk_ids=curr_topk_ids,
|
||||
sorted_token_ids=sorted_token_ids,
|
||||
expert_ids=expert_ids,
|
||||
num_tokens_post_padded=num_tokens_post_padded,
|
||||
mul_routed_weight=True,
|
||||
top_k=1,
|
||||
config=config,
|
||||
compute_type=tl.bfloat16,
|
||||
)
|
||||
|
||||
# Accumulate gradients
|
||||
grad_intermediate_cache2[begin_chunk_idx * topk : end_chunk_idx * topk] = curr_grad_intermediate_cache2
|
||||
grad_w2 += curr_grad_w2
|
||||
grad_topk_weights[begin_chunk_idx:end_chunk_idx] = curr_grad_topk_weights
|
||||
|
||||
return grad_intermediate_cache2, grad_w2, grad_topk_weights, None
|
||||
|
||||
|
||||
class MoeSumReduceFunction(torch.autograd.Function):
|
||||
@staticmethod
|
||||
def forward(
|
||||
ctx,
|
||||
intermediate_cache3: torch.Tensor,
|
||||
hidden_states_shape,
|
||||
):
|
||||
out_hidden_states = torch.empty(
|
||||
hidden_states_shape, device=intermediate_cache3.device, dtype=intermediate_cache3.dtype
|
||||
)
|
||||
moe_sum_reduce(
|
||||
intermediate_cache3,
|
||||
out_hidden_states,
|
||||
1.0,
|
||||
)
|
||||
ctx.save_for_backward(intermediate_cache3)
|
||||
return out_hidden_states
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
(intermediate_cache3,) = ctx.saved_tensors
|
||||
return grad_output.unsqueeze(1).expand_as(intermediate_cache3), None
|
||||
Reference in New Issue
Block a user