Revert "feat: batched MoE expert GEMM — replaces Python for-loop"
This reverts commit 06d7713db6.
This commit is contained in:
@@ -1,114 +0,0 @@
|
|||||||
// moe_batched_gemm.cu — Fused MoE expert GEMM for BI-V100
|
|
||||||
//
|
|
||||||
// Replaces the Python for-loop over 256 experts with:
|
|
||||||
// 1. Gather tokens by expert (sorted by moe_compute_index)
|
|
||||||
// 2. Per-expert GEMM via cublas (torch::mm)
|
|
||||||
// 3. Fused silu activation
|
|
||||||
// 4. Per-expert down GEMM
|
|
||||||
// 5. Weighted scatter-add back to output
|
|
||||||
//
|
|
||||||
// This eliminates Python loop overhead (~256 iterations) and reduces
|
|
||||||
// kernel launch overhead by batching small GEMMs.
|
|
||||||
//
|
|
||||||
// Dimensions (Qwen3.5-35B-A3B, TP=4):
|
|
||||||
// w13: (256, 256, 2048) -> E=256, 2*I=256, H=2048
|
|
||||||
// w2: (256, 2048, 128) -> E=256, H=2048, I=128
|
|
||||||
|
|
||||||
#include <torch/extension.h>
|
|
||||||
#include <c10/cuda/CUDAGuard.h>
|
|
||||||
#include <ATen/cuda/CUDAContext.h>
|
|
||||||
|
|
||||||
namespace xllm::kernel::cuda {
|
|
||||||
|
|
||||||
// Fused SiLU-and-mul kernel (gate_up → act)
|
|
||||||
// gate_up: (N, 2*I), output: (N, I)
|
|
||||||
__global__ void silu_and_mul_inplace_kernel(
|
|
||||||
const __half* __restrict__ gate_up,
|
|
||||||
__half* __restrict__ output,
|
|
||||||
int64_t N, int64_t I) {
|
|
||||||
int64_t idx = blockIdx.x * blockDim.x + threadIdx.x;
|
|
||||||
if (idx >= N * I) return;
|
|
||||||
|
|
||||||
int64_t row = idx / I;
|
|
||||||
int64_t col = idx % I;
|
|
||||||
|
|
||||||
float gate = __half2float(gate_up[row * 2 * I + col]);
|
|
||||||
float up = __half2float(gate_up[row * 2 * I + I + col]);
|
|
||||||
float silu_gate = gate / (1.0f + expf(-gate));
|
|
||||||
output[idx] = __float2half(silu_gate * up);
|
|
||||||
}
|
|
||||||
|
|
||||||
// Main function: batched expert forward
|
|
||||||
// Called from Python with pre-sorted token indices
|
|
||||||
torch::Tensor moe_experts_forward(
|
|
||||||
torch::Tensor hidden_states, // (T, H) — all tokens
|
|
||||||
torch::Tensor w13, // (E, 2*I, H) — gate+up weights
|
|
||||||
torch::Tensor w2, // (E, H, I) — down weights
|
|
||||||
torch::Tensor sorted_tok_ids, // (T*topk,) — which token for each slot
|
|
||||||
torch::Tensor sorted_weights, // (T*topk,) — routing weight for each slot
|
|
||||||
torch::Tensor expert_offsets, // (E+1,) — cumsum of expert_sizes, expert_offsets[0]=0
|
|
||||||
int64_t topk) {
|
|
||||||
|
|
||||||
auto stream = at::cuda::getCurrentCUDAStream();
|
|
||||||
int64_t T = hidden_states.size(0);
|
|
||||||
int64_t H = hidden_states.size(1);
|
|
||||||
int64_t E = w13.size(0);
|
|
||||||
int64_t two_I = w13.size(1); // 2*I
|
|
||||||
int64_t I = two_I / 2;
|
|
||||||
|
|
||||||
auto out = torch::zeros({T, H}, hidden_states.options());
|
|
||||||
|
|
||||||
// Get expert offsets on CPU for loop control
|
|
||||||
auto offsets_cpu = expert_offsets.to(torch::kCPU, torch::kInt64);
|
|
||||||
auto offsets_ptr = offsets_cpu.data_ptr<int64_t>();
|
|
||||||
|
|
||||||
for (int64_t eid = 0; eid < E; ++eid) {
|
|
||||||
int64_t start = offsets_ptr[eid];
|
|
||||||
int64_t end = offsets_ptr[eid + 1];
|
|
||||||
int64_t count = end - start;
|
|
||||||
if (count == 0) continue;
|
|
||||||
|
|
||||||
// Gather tokens for this expert
|
|
||||||
auto tok_ids = sorted_tok_ids.slice(0, start, end); // (count,)
|
|
||||||
auto tokens = hidden_states.index_select(0, tok_ids); // (count, H)
|
|
||||||
|
|
||||||
// GEMM 1: gate+up projection
|
|
||||||
// tokens (count, H) × w13[eid].T (H, 2*I) → (count, 2*I)
|
|
||||||
auto gate_up = torch::mm(tokens, w13[eid].t()); // (count, 2*I)
|
|
||||||
|
|
||||||
// Fused SiLU activation
|
|
||||||
auto act = torch::empty({count, I}, hidden_states.options());
|
|
||||||
if (hidden_states.dtype() == torch::kFloat16) {
|
|
||||||
int64_t total = count * I;
|
|
||||||
int block = 256;
|
|
||||||
int grid = (total + block - 1) / block;
|
|
||||||
silu_and_mul_inplace_kernel<<<grid, block, 0, stream>>>(
|
|
||||||
gate_up.data_ptr<at::Half>(),
|
|
||||||
act.data_ptr<at::Half>(),
|
|
||||||
count, I);
|
|
||||||
} else {
|
|
||||||
auto chunks = gate_up.chunk(2, /*dim=*/1);
|
|
||||||
act = torch::silu(chunks[0]) * chunks[1];
|
|
||||||
}
|
|
||||||
|
|
||||||
// GEMM 2: down projection
|
|
||||||
// act (count, I) × w2[eid].T (I, H) → (count, H)
|
|
||||||
auto expert_out = torch::mm(act, w2[eid].t()); // (count, H)
|
|
||||||
|
|
||||||
// Weighted scatter-add
|
|
||||||
auto weights = sorted_weights.slice(0, start, end).unsqueeze(1); // (count, 1)
|
|
||||||
out.index_add_(0, tok_ids, (expert_out * weights).to(out.dtype()));
|
|
||||||
}
|
|
||||||
|
|
||||||
return out;
|
|
||||||
}
|
|
||||||
|
|
||||||
} // namespace xllm::kernel::cuda
|
|
||||||
|
|
||||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
|
||||||
m.def("moe_experts_forward", &xllm::kernel::cuda::moe_experts_forward,
|
|
||||||
"Batched MoE expert forward (gather → GEMM → silu → GEMM → scatter)",
|
|
||||||
py::arg("hidden_states"), py::arg("w13"), py::arg("w2"),
|
|
||||||
py::arg("sorted_tok_ids"), py::arg("sorted_weights"),
|
|
||||||
py::arg("expert_offsets"), py::arg("topk"));
|
|
||||||
}
|
|
||||||
@@ -67,8 +67,5 @@ build_kernel "xllm_cache" \
|
|||||||
build_kernel "xllm_moe" \
|
build_kernel "xllm_moe" \
|
||||||
"${CUDA_DIR}/moe/moe_fused_topk.cu" "${CUDA_DIR}/moe/moe_compute_index.cu" "${CUDA_DIR}/moe/moe_combine.cu" "${BIND_DIR}/xllm_moe_bind.cpp"
|
"${CUDA_DIR}/moe/moe_fused_topk.cu" "${CUDA_DIR}/moe/moe_compute_index.cu" "${CUDA_DIR}/moe/moe_combine.cu" "${BIND_DIR}/xllm_moe_bind.cpp"
|
||||||
|
|
||||||
build_kernel "xllm_moe_gemm" \
|
|
||||||
"${CUDA_DIR}/moe_batched_gemm.cu"
|
|
||||||
|
|
||||||
echo "=== All kernels built ==="
|
echo "=== All kernels built ==="
|
||||||
ls -lh "${PREBUILT_DIR}"/xllm_*.so 2>/dev/null || echo "No .so files found"
|
ls -lh "${PREBUILT_DIR}"/xllm_*.so 2>/dev/null || echo "No .so files found"
|
||||||
|
|||||||
@@ -156,14 +156,6 @@ except ImportError:
|
|||||||
except ImportError:
|
except ImportError:
|
||||||
_xllm_moe = None
|
_xllm_moe = None
|
||||||
|
|
||||||
try:
|
|
||||||
from vllm import xllm_moe_gemm as _xllm_moe_gemm
|
|
||||||
except ImportError:
|
|
||||||
try:
|
|
||||||
import xllm_moe_gemm as _xllm_moe_gemm
|
|
||||||
except ImportError:
|
|
||||||
_xllm_moe_gemm = None
|
|
||||||
|
|
||||||
try:
|
try:
|
||||||
from vllm import corex_gdn_chunk_recurrent as _corex_gdn_chunk_recurrent
|
from vllm import corex_gdn_chunk_recurrent as _corex_gdn_chunk_recurrent
|
||||||
except ImportError:
|
except ImportError:
|
||||||
@@ -221,13 +213,8 @@ _USE_COREX_MOE_INDEX_COMBINE = (
|
|||||||
_USE_XLLM_MOE = (
|
_USE_XLLM_MOE = (
|
||||||
_xllm_moe is not None
|
_xllm_moe is not None
|
||||||
and env_bool("BI100_MOE_XLLM", True))
|
and env_bool("BI100_MOE_XLLM", True))
|
||||||
_USE_XLLM_MOE_GEMM = (
|
|
||||||
_xllm_moe_gemm is not None
|
|
||||||
and env_bool("BI100_MOE_XLLM_GEMM", True))
|
|
||||||
if _USE_XLLM_MOE:
|
if _USE_XLLM_MOE:
|
||||||
logger.info("xllm_moe ENABLED — fused_topk + compute_index + combine_result")
|
logger.info("xllm_moe ENABLED — fused_topk + compute_index + combine_result")
|
||||||
if _USE_XLLM_MOE_GEMM:
|
|
||||||
logger.info("xllm_moe_gemm ENABLED — batched expert GEMM (replaces Python loop)")
|
|
||||||
_USE_FUSED_MOE_ACTIVATION = env_bool("BI100_MOE_FUSED_ACTIVATION", True)
|
_USE_FUSED_MOE_ACTIVATION = env_bool("BI100_MOE_FUSED_ACTIVATION", True)
|
||||||
|
|
||||||
# ix_fused_moe: full 7-step fused MoE pipeline via ixformer C++ API
|
# ix_fused_moe: full 7-step fused MoE pipeline via ixformer C++ API
|
||||||
@@ -1821,37 +1808,21 @@ class Qwen3_5MoeSparseBlock(nn.Module):
|
|||||||
expert_counts = torch.bincount(
|
expert_counts = torch.bincount(
|
||||||
flat_eids, minlength=w13.shape[0]).tolist()
|
flat_eids, minlength=w13.shape[0]).tolist()
|
||||||
|
|
||||||
if _USE_XLLM_MOE_GEMM:
|
start = 0
|
||||||
# Batched expert GEMM via CUDA — eliminates Python for-loop
|
for eid, count in enumerate(expert_counts):
|
||||||
# Build expert_offsets from expert_counts (cumsum with leading 0)
|
end = start + count
|
||||||
expert_counts_t = torch.tensor(
|
if count == 0:
|
||||||
expert_counts if isinstance(expert_counts, list)
|
|
||||||
else expert_counts.tolist(),
|
|
||||||
dtype=torch.int64, device=hidden_states.device)
|
|
||||||
expert_offsets = torch.zeros(
|
|
||||||
len(expert_counts) + 1, dtype=torch.int64,
|
|
||||||
device=hidden_states.device)
|
|
||||||
torch.cumsum(expert_counts_t, dim=0, out=expert_offsets[1:])
|
|
||||||
out = _xllm_moe_gemm.moe_experts_forward(
|
|
||||||
hidden_states, w13, w2,
|
|
||||||
sorted_tok_ids, sorted_weights.float(),
|
|
||||||
expert_offsets, self.top_k)
|
|
||||||
else:
|
|
||||||
start = 0
|
|
||||||
for eid, count in enumerate(expert_counts):
|
|
||||||
end = start + count
|
|
||||||
if count == 0:
|
|
||||||
start = end
|
|
||||||
continue
|
|
||||||
tok_ids = sorted_tok_ids[start:end]
|
|
||||||
tokens = hidden_states[tok_ids] # (n, H)
|
|
||||||
gate_up = F.linear(tokens, w13[eid]) # (n, 2*I)
|
|
||||||
gate, up = gate_up.chunk(2, dim=-1)
|
|
||||||
act = F.silu(gate) * up # (n, I)
|
|
||||||
expert_out = F.linear(act, w2[eid]) # (n, H)
|
|
||||||
weights = sorted_weights[start:end].unsqueeze(-1)
|
|
||||||
out.index_add_(0, tok_ids, (expert_out * weights).to(out.dtype))
|
|
||||||
start = end
|
start = end
|
||||||
|
continue
|
||||||
|
tok_ids = sorted_tok_ids[start:end]
|
||||||
|
tokens = hidden_states[tok_ids] # (n, H)
|
||||||
|
gate_up = F.linear(tokens, w13[eid]) # (n, 2*I)
|
||||||
|
gate, up = gate_up.chunk(2, dim=-1)
|
||||||
|
act = F.silu(gate) * up # (n, I)
|
||||||
|
expert_out = F.linear(act, w2[eid]) # (n, H)
|
||||||
|
weights = sorted_weights[start:end].unsqueeze(-1)
|
||||||
|
out.index_add_(0, tok_ids, (expert_out * weights).to(out.dtype))
|
||||||
|
start = end
|
||||||
|
|
||||||
return out # partial, all-reduce done in forward()
|
return out # partial, all-reduce done in forward()
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user