[MoE] Deprecate act_and_mul_triton; fold filter_expert into JIT silu/gelu_and_mul (#23707)
Co-authored-by: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 4.7
parent
d49a0377de
commit
c7878dbb6d
@@ -44,9 +44,14 @@ struct ActivationParams {
|
||||
void* __restrict__ out;
|
||||
uint32_t hidden_dim;
|
||||
uint32_t num_tokens;
|
||||
// Optional MoE expert filtering: when expert_ids != nullptr, a token is
|
||||
// skipped if expert_ids[token_id / expert_step] == -1. expert_step is 1
|
||||
// for per-token routing and BLOCK_SIZE_M for sorted/TMA routing.
|
||||
const int32_t* __restrict__ expert_ids;
|
||||
uint32_t expert_step;
|
||||
};
|
||||
|
||||
template <typename T, ActivationKind kAct, bool kUsePDL>
|
||||
template <typename T, ActivationKind kAct, bool kUsePDL, bool kFilterExpert>
|
||||
__global__ void act_and_mul_kernel(const __grid_constant__ ActivationParams params) {
|
||||
using namespace device;
|
||||
constexpr auto kVecSize = kMaxVecBytes / sizeof(T);
|
||||
@@ -56,6 +61,9 @@ __global__ void act_and_mul_kernel(const __grid_constant__ ActivationParams para
|
||||
const auto token_id = tid / num_vecs;
|
||||
|
||||
if (token_id >= params.num_tokens) return;
|
||||
if constexpr (kFilterExpert) {
|
||||
if (params.expert_ids[token_id / params.expert_step] == -1) return;
|
||||
}
|
||||
const auto offset = tid % num_vecs;
|
||||
const auto input_offset = token_id * (num_vecs * 2) + offset;
|
||||
const auto output_offset = tid;
|
||||
@@ -78,11 +86,33 @@ struct ActivationKernel {
|
||||
static constexpr auto kVecSize = device::kMaxVecBytes / sizeof(T);
|
||||
static constexpr auto kBlockSize = 256u;
|
||||
|
||||
template <ActivationKind kAct>
|
||||
static constexpr auto activation_kernel = act_and_mul_kernel<T, kAct, kUsePDL>;
|
||||
template <ActivationKind kAct, bool kFilterExpert>
|
||||
static constexpr auto activation_kernel = act_and_mul_kernel<T, kAct, kUsePDL, kFilterExpert>;
|
||||
|
||||
static_assert(device::kMaxVecBytes % sizeof(T) == 0, "unsupported data type");
|
||||
static void run_activation(const tvm::ffi::TensorView input, const tvm::ffi::TensorView out, std::string type) {
|
||||
|
||||
template <bool kFilterExpert>
|
||||
static auto select_kernel(const std::string& type)
|
||||
-> decltype(activation_kernel<ActivationKind::kSiLU, kFilterExpert>) {
|
||||
using namespace host;
|
||||
if (type == "silu") {
|
||||
return activation_kernel<ActivationKind::kSiLU, kFilterExpert>;
|
||||
} else if (type == "gelu") {
|
||||
return activation_kernel<ActivationKind::kGELU, kFilterExpert>;
|
||||
} else if (type == "gelu_tanh") {
|
||||
return activation_kernel<ActivationKind::kGELUTanh, kFilterExpert>;
|
||||
} else {
|
||||
Panic("unsupported activation type: ", type);
|
||||
}
|
||||
return nullptr;
|
||||
}
|
||||
|
||||
static void launch(
|
||||
const tvm::ffi::TensorView& input,
|
||||
const tvm::ffi::TensorView& out,
|
||||
const std::string& type,
|
||||
const int32_t* expert_ids,
|
||||
uint32_t expert_step) {
|
||||
using namespace host;
|
||||
|
||||
auto N = SymbolicSize{"num_tokens"};
|
||||
@@ -106,18 +136,6 @@ struct ActivationKernel {
|
||||
if (num_tokens == 0) return;
|
||||
RuntimeCheck(hidden_size * 2 == D_in.unwrap(), "invalid activation dimension");
|
||||
RuntimeCheck(hidden_size % kVecSize == 0, "hidden size must be divisible by vector size");
|
||||
const auto kernel = [&]() -> decltype(activation_kernel<ActivationKind::kSiLU>) {
|
||||
if (type == "silu") {
|
||||
return activation_kernel<ActivationKind::kSiLU>;
|
||||
} else if (type == "gelu") {
|
||||
return activation_kernel<ActivationKind::kGELU>;
|
||||
} else if (type == "gelu_tanh") {
|
||||
return activation_kernel<ActivationKind::kGELUTanh>;
|
||||
} else {
|
||||
Panic("unsupported activation type: ", type);
|
||||
}
|
||||
return nullptr;
|
||||
}();
|
||||
// only get once to avoid overhead
|
||||
const auto num_total_items = num_tokens * (hidden_size / kVecSize);
|
||||
RuntimeCheck(num_total_items <= std::numeric_limits<uint32_t>::max(), "too many items for 32-bit indexing");
|
||||
@@ -127,8 +145,33 @@ struct ActivationKernel {
|
||||
.out = out.data_ptr(),
|
||||
.hidden_dim = hidden_size,
|
||||
.num_tokens = num_tokens,
|
||||
.expert_ids = expert_ids,
|
||||
.expert_step = expert_step,
|
||||
};
|
||||
LaunchKernel(num_blocks, kBlockSize, device).enable_pdl(kUsePDL)(kernel, params);
|
||||
if (expert_ids != nullptr) {
|
||||
RuntimeCheck(expert_step > 0, "expert_step must be positive");
|
||||
const auto kernel = select_kernel<true>(type);
|
||||
LaunchKernel(num_blocks, kBlockSize, device).enable_pdl(kUsePDL)(kernel, params);
|
||||
} else {
|
||||
const auto kernel = select_kernel<false>(type);
|
||||
LaunchKernel(num_blocks, kBlockSize, device).enable_pdl(kUsePDL)(kernel, params);
|
||||
}
|
||||
}
|
||||
|
||||
static void run_activation(const tvm::ffi::TensorView input, const tvm::ffi::TensorView out, std::string type) {
|
||||
launch(input, out, type, /*expert_ids=*/nullptr, /*expert_step=*/1);
|
||||
}
|
||||
|
||||
static void run_activation_filtered(
|
||||
const tvm::ffi::TensorView input,
|
||||
const tvm::ffi::TensorView out,
|
||||
const tvm::ffi::TensorView expert_ids,
|
||||
int64_t expert_step,
|
||||
std::string type) {
|
||||
using namespace host;
|
||||
RuntimeCheck(is_type<int32_t>(expert_ids.dtype()), "expert_ids must have dtype int32");
|
||||
RuntimeCheck(expert_step >= 1, "expert_step must be positive");
|
||||
launch(input, out, type, static_cast<const int32_t*>(expert_ids.data_ptr()), static_cast<uint32_t>(expert_step));
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
Reference in New Issue
Block a user