Skip to content

Commit af9e1d1

Browse files
authored
[rocm-libraries] ROCm/rocm-libraries#9396 (commit e1aa8ce)
feat(ck): Added Gelu with Tanh approx to XDL 2-stage MoE epilogue ## Motivation Enable the tanh-approximation GELU activation `(gelu_tanh, 0.5 * x * (1 + tanh(sqrt(2/pi) * (x + 0.044715 * x^3))))` in the Composable Kernel XDL 2-stage MoE path. The MoE gridwise kernel epilogue currently supports only `silu/gelu/swiglustep/swiglu_oai`; this adds `gelu_tanh_and_mul` so models whose MoE experts use the GELU tanh approximation (e.g. [Gemma-family MoE](https://huggingface.co/google/gemma-4-26B-A4B/blob/main/config.json)) can use this path. JIRA ID : ROCM-27619 ## Technical Details - `gridwise_gemm_xdl_cshuffle_common.hpp`: add `Activation::gelu_tanh_and_mul = 4` to the activation enum. - `gridwise_moe_gemm.hpp`, `gridwise_moe_gemm_blockscale.hpp`, `gridwise_moe_mx_gemm_<>.hpp`: wire `gelu_tanh_and_mul` into epilogue paths, delegating to the existing `ck::tensor_operation::element_wise::FastGelu` helper (the single source of truth for the tanh-GELU math, `FastGelu(gate) * up`). Also added `static_assert` for validation of supported activations - The activation is applied in fp32 in the epilogue and is orthogonal to the GEMM compute (MFMA/tile/pipeline untouched) and to quantization (existing per-token dequant reused). Only the non-blockscale gridwise kernel is changed. - Then I plan to port these changes to AITER after ROCm/aiter#3886 to avoid merge conflicts ## Test Plan Use `ActOP = 4` in the example `moe_gemm1_xdl_fp8`, rebuild example and launch ctest ## Test Result ``` ctest -R "^example_moe_gemm1_xdl_fp8$" -V' Constructing a list of tests Done constructing a list of tests Updating test list for fixtures Added 0 tests to meet fixture requirements Checking test dependency graph... Checking test dependency graph end test 257 Start 257: example_moe_gemm1_xdl_fp8 257: Test command: example_moe_gemm1_xdl_fp8 257: Working Directory: example/65_gemm_multiply_multiply 257: Test timeout computed to be: 1500 257: a0_t_k: dim 2, lengths {16384, 6144}, strides {6144, 1} 257: b0_e_n_k: dim 3, lengths {8, 6144, 8192}, strides {50331648, 1, 6144} 257: d1_e_n: dim 2, lengths {8, 8192}, strides {8192, 1} 257: d2_e_n: dim 2, lengths {32768, 4096}, strides {1, 0} 257: d0_t_n: dim 2, lengths {16384, 4096}, strides {1, 16384} 257: d2_e_n: dim 2, lengths {32768, 4096}, strides {1, 0} 257: e_t_n: dim 3, lengths {16384, 2, 4096}, strides {8192, 4096, 1} 1/1 Test #257: example_moe_gemm1_xdl_fp8 ........ Passed 83.51 sec The following tests passed: example_moe_gemm1_xdl_fp8 100% tests passed, 0 tests failed out of 1 Label Time Summary: SMOKE_TEST = 83.51 sec*proc (1 test) Total Test time (real) = 83.59 sec ``` ## Submission Checklist - [ ] Look over the contributing guidelines at https://github.com/ROCm/ROCm/blob/develop/CONTRIBUTING.md#pull-requests.
1 parent b974e7e commit af9e1d1

16 files changed

Lines changed: 310 additions & 30 deletions

‎example/65_gemm_multiply_multiply/CMakeLists.txt‎

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -20,6 +20,11 @@ add_example_executable(example_moe_gemm2_xdl_fp8_blockscale moe_gemm2_xdl_fp8_bl
2020
add_example_executable(example_moe_gemm1_xdl_fp8_blockscale moe_gemm1_xdl_fp8_blockscale.cpp)
2121
add_example_executable(example_moe_gemm1_xdl_fp8_blockscale_splitk moe_gemm1_xdl_fp8_blockscale_splitk.cpp)
2222

23+
add_example_executable(example_moe_gemm1_xdl_fp8_gelu_tanh moe_gemm1_xdl_fp8.cpp)
24+
if(TARGET example_moe_gemm1_xdl_fp8_gelu_tanh)
25+
target_compile_definitions(example_moe_gemm1_xdl_fp8_gelu_tanh PRIVATE MOE_ACTOP=4)
26+
endif()
27+
2328
list(APPEND gpu_list gfx942 gfx950 gfx1100 gfx1101 gfx1102 gfx1103 gfx1150 gfx1151 gfx1152 gfx1153 gfx1200 gfx1201 gfx11-generic gfx12-generic gfx1250)
2429

2530
set(target 0)
@@ -69,6 +74,7 @@ if(HAS_MAX_OCCUPANCY_EXPERIMENTAL)
6974
endif()
7075
example_compile_options(example_gemm_multiply_multiply_xdl_fp8_bpreshuffle PRIVATE ${GEMM_OPTIONS})
7176
example_compile_options(example_moe_gemm1_xdl_fp8 PRIVATE ${GEMM_OPTIONS})
77+
example_compile_options(example_moe_gemm1_xdl_fp8_gelu_tanh PRIVATE ${GEMM_OPTIONS})
7278
example_compile_options(example_moe_gemm2_xdl_fp8 PRIVATE ${GEMM_OPTIONS})
7379
example_compile_options(example_gemm_multiply_multiply_xdl_fp8_ab_scale PRIVATE ${BLOCKSCALE_GEMM_OPTIONS})
7480
example_compile_options(example_gemm_multiply_multiply_xdl_fp8_blockscale_bpreshuffle PRIVATE ${BLOCKSCALE_GEMM_OPTIONS})

‎example/65_gemm_multiply_multiply/moe_gemm1_xdl_fp8.cpp‎

Lines changed: 10 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -176,9 +176,16 @@ static constexpr ck::index_t BK1 = 16 / sizeof(B0DataType);
176176
static constexpr ck::index_t EVec = 8 / sizeof(EDataType);
177177
static constexpr ck::index_t D0Vec = 1;
178178
static constexpr ck::index_t D1Vec = 1;
179-
static constexpr ck::index_t ActOP = 1; // 0: gelu_and_mul, 1: silu_and_mul
180-
static constexpr bool MulRoutedWeight = false;
181-
using DeviceOpInstance = ck::tensor_operation::device::DeviceMoeGemm
179+
// Activation (ck::Activation): 0: gelu_and_mul, 1: silu_and_mul, 2: swiglustep_and_mul,
180+
// 3: swiglu_oai_and_mul, 4: gelu_tanh_and_mul
181+
// MOE_ACTOP may be overridden at compile time (e.g. -DMOE_ACTOP=4) so the same example
182+
// can be built as separate ctest variants exercising different activations.
183+
#ifndef MOE_ACTOP
184+
#define MOE_ACTOP 1
185+
#endif
186+
static constexpr ck::index_t ActOP = MOE_ACTOP;
187+
static constexpr bool MulRoutedWeight = false;
188+
using DeviceOpInstance = ck::tensor_operation::device::DeviceMoeGemm
182189
// clang-format off
183190
< Row, Col, DsLayout, ELayout, A0DataType, B0DataType, DsDataType, EDataType, AccDataType, CShuffleDataType,
184191
AElementOp, BElementOp, CDEElementOp, GemmSpec,

‎example/65_gemm_multiply_multiply/moe_gemm1_xdl_fp8_blockscale.cpp‎

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -130,7 +130,9 @@ static constexpr ck::index_t Scale_Block_N = 128;
130130
static constexpr ck::index_t Scale_Block_K = 128;
131131

132132
static constexpr ck::index_t Nswizzle = false;
133-
static constexpr ck::index_t ActOP = 0; // 0: gelu_and_mul, 1: silu_and_mul
133+
// Activation (ck::Activation): 0: gelu_and_mul, 1: silu_and_mul, 2: swiglustep_and_mul,
134+
// 4: gelu_tanh_and_mul
135+
static constexpr ck::index_t ActOP = 0;
134136
static constexpr bool MulRoutedWeight = true;
135137

136138
#if 0

‎example/65_gemm_multiply_multiply/moe_gemm1_xdl_fp8_blockscale_splitk.cpp‎

Lines changed: 6 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -113,10 +113,12 @@ static constexpr ck::index_t Scale_Block_N = 128;
113113
static constexpr ck::index_t Scale_Block_K = 128;
114114

115115
static constexpr ck::index_t Nswizzle = false;
116-
static constexpr ck::index_t IsInputGemm = true; // splitk gemm1 goes to gemm2 pipeline.
117-
static constexpr ck::index_t IsSplitK = true; // splitk gemm1
118-
static constexpr ck::index_t ActOP = 0; // 0: gelu_and_mul, 1: silu_and_mul
119-
static constexpr bool MulRoutedWeight = false; // splitk gemm1 does not do routedWeight.
116+
static constexpr ck::index_t IsInputGemm = true; // splitk gemm1 goes to gemm2 pipeline.
117+
static constexpr ck::index_t IsSplitK = true; // splitk gemm1
118+
// NOTE: ActOP is unused in this split-K path. The fused epilogue activation is only applied
119+
// The fused epilogue activation is only applied when (IsInputGemm && !IsSplitK)
120+
static constexpr ck::index_t ActOP = 0;
121+
static constexpr bool MulRoutedWeight = false; // splitk gemm1 does not do routedWeight.
120122

121123
#if 1
122124
static constexpr ck::index_t MPerBlock = 64;

‎example/67_gemm_microscaling/moe_gemm1_xdl_mx_fp4.cpp‎

Lines changed: 5 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -77,10 +77,11 @@ using CDEElementOp = MulABScaleExpertWeight;
7777

7878
static constexpr auto GemmSpec = ck::tensor_operation::device::GemmSpecialization::Default;
7979

80-
constexpr ck::index_t ScaleBlockSize = 32; // scaling block size
81-
constexpr ck::index_t KPerBlock = 128;
82-
static constexpr ck::index_t Nswizzle = false;
83-
static constexpr ck::index_t ActOP = 0; // 0: gelu_and_mul, 1: silu_and_mul
80+
constexpr ck::index_t ScaleBlockSize = 32; // scaling block size
81+
constexpr ck::index_t KPerBlock = 128;
82+
static constexpr ck::index_t Nswizzle = false;
83+
// Activation (ck::Activation): 0: gelu_and_mul, 1: silu_and_mul, 4: gelu_tanh_and_mul
84+
static constexpr ck::index_t ActOP = 0;
8485
static constexpr ck::index_t MPerBlock = 128;
8586
static constexpr ck::index_t NPerBlock = 64;
8687
static constexpr ck::index_t BlockSize = 256;

‎example/67_gemm_microscaling/moe_gemm1_xdl_mx_fp4_bns.cpp‎

Lines changed: 5 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -77,10 +77,11 @@ using CDEElementOp = MulABScaleExpertWeight;
7777

7878
static constexpr auto GemmSpec = ck::tensor_operation::device::GemmSpecialization::Default;
7979

80-
constexpr ck::index_t ScaleBlockSize = 32; // scaling block size
81-
constexpr ck::index_t KPerBlock = 128; // 128 fp4x2 or 128 fp8
82-
static constexpr ck::index_t Nswizzle = false;
83-
static constexpr ck::index_t ActOP = 0; // 0: gelu_and_mul, 1: silu_and_mul
80+
constexpr ck::index_t ScaleBlockSize = 32; // scaling block size
81+
constexpr ck::index_t KPerBlock = 128; // 128 fp4x2 or 128 fp8
82+
static constexpr ck::index_t Nswizzle = false;
83+
// Activation (ck::Activation): 0: gelu_and_mul, 1: silu_and_mul, 4: gelu_tanh_and_mul
84+
static constexpr ck::index_t ActOP = 0;
8485
static constexpr ck::index_t MPerBlock = 128;
8586
static constexpr ck::index_t NPerBlock = 64;
8687
static constexpr ck::index_t BlockSize = 256;

‎example/67_gemm_microscaling/moe_gemm1_xdl_mx_fp4_bpreshuffle.cpp‎

Lines changed: 5 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -79,10 +79,11 @@ using CDEElementOp = MulABScaleExpertWeight;
7979

8080
static constexpr auto GemmSpec = ck::tensor_operation::device::GemmSpecialization::Default;
8181

82-
constexpr ck::index_t ScaleBlockSize = 32; // scaling block size
83-
constexpr ck::index_t KPerBlock = 128;
84-
static constexpr ck::index_t Nswizzle = false;
85-
static constexpr ck::index_t ActOP = 0; // 0: gelu_and_mul, 1: silu_and_mul
82+
constexpr ck::index_t ScaleBlockSize = 32; // scaling block size
83+
constexpr ck::index_t KPerBlock = 128;
84+
static constexpr ck::index_t Nswizzle = false;
85+
// Activation (ck::Activation): 0: gelu_and_mul, 1: silu_and_mul, 4: gelu_tanh_and_mul
86+
static constexpr ck::index_t ActOP = 0;
8687
static constexpr ck::index_t MPerBlock = 32;
8788
static constexpr bool MulRoutedWeight = true;
8889

‎include/ck/tensor_operation/gpu/grid/gridwise_gemm_xdl_cshuffle_common.hpp‎

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -31,7 +31,8 @@ enum Activation
3131
gelu_and_mul = 0,
3232
silu_and_mul = 1,
3333
swiglustep_and_mul = 2,
34-
swiglu_oai_and_mul = 3
34+
swiglu_oai_and_mul = 3,
35+
gelu_tanh_and_mul = 4
3536
};
3637

3738
// OAI / gpt-oss SwiGLU activation: gate * sigmoid(alpha * gate) * (up + 1), with a

‎include/ck/tensor_operation/gpu/grid/gridwise_moe_gemm.hpp‎

Lines changed: 78 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1156,6 +1156,13 @@ struct GridwiseMoeGemm : public GridwiseGemm_xdl_cshuffle_base<
11561156
BElementwiseOperation b_element_op,
11571157
CElementwiseOperation c_element_op)
11581158
{
1159+
static_assert(ActivationOperation == Activation::gelu_and_mul ||
1160+
ActivationOperation == Activation::silu_and_mul ||
1161+
ActivationOperation == Activation::swiglustep_and_mul ||
1162+
ActivationOperation == Activation::swiglu_oai_and_mul ||
1163+
ActivationOperation == Activation::gelu_tanh_and_mul,
1164+
"gridwise_moe_gemm only supports gelu_and_mul, silu_and_mul, "
1165+
"swiglustep_and_mul, swiglu_oai_and_mul and gelu_tanh_and_mul.");
11591166
ignore = b_element_op;
11601167
index_t BN0Shuffled = CalculateBN0Shuffled(problem.N);
11611168
index_t BK0Shuffled = CalculateBK0Shuffled(problem.K);
@@ -1505,6 +1512,26 @@ struct GridwiseMoeGemm : public GridwiseGemm_xdl_cshuffle_base<
15051512
tensor_operation::element_wise::Gelu{}(gate, gate);
15061513
c_thread_buf_fp32(cidx) = gate * up;
15071514
}
1515+
else if(ActivationOperation == Activation::gelu_tanh_and_mul)
1516+
{
1517+
const float scale_up =
1518+
p_scale_b[(n0 * NWave * NPerXdl + problem.N) *
1519+
PerTokenQuant];
1520+
float gate = scale_a * scale_b * c_thread_buf[cidx];
1521+
float up = scale_a * scale_up * c_thread_buf_up[cidx];
1522+
if constexpr(MulRoutedWeight)
1523+
{
1524+
gate = gate * topk_weights.template AsType<float>()[m4];
1525+
up = up * topk_weights.template AsType<float>()[m4];
1526+
}
1527+
if constexpr(is_same_v<remove_cvref_t<BDataType>, pk_i4_t>)
1528+
{
1529+
gate *= 16;
1530+
up *= 16;
1531+
}
1532+
tensor_operation::element_wise::FastGelu{}(gate, gate);
1533+
c_thread_buf_fp32(cidx) = gate * up;
1534+
}
15081535
else if constexpr(ActivationOperation ==
15091536
Activation::swiglustep_and_mul)
15101537
{
@@ -1606,6 +1633,18 @@ struct GridwiseMoeGemm : public GridwiseGemm_xdl_cshuffle_base<
16061633
tensor_operation::element_wise::Gelu{}(gate, gate);
16071634
c_thread_buf_fp32(cidx) = gate * up;
16081635
}
1636+
else if(ActivationOperation == Activation::gelu_tanh_and_mul)
1637+
{
1638+
float gate = c_thread_buf[cidx];
1639+
float up = c_thread_buf_up[cidx];
1640+
if constexpr(MulRoutedWeight)
1641+
{
1642+
gate = gate * topk_weights.template AsType<float>()[m4];
1643+
up = up * topk_weights.template AsType<float>()[m4];
1644+
}
1645+
tensor_operation::element_wise::FastGelu{}(gate, gate);
1646+
c_thread_buf_fp32(cidx) = gate * up;
1647+
}
16091648
else if constexpr(ActivationOperation ==
16101649
Activation::swiglustep_and_mul)
16111650
{
@@ -1686,6 +1725,13 @@ struct GridwiseMoeGemm : public GridwiseGemm_xdl_cshuffle_base<
16861725
BElementwiseOperation b_element_op,
16871726
CElementwiseOperation c_element_op)
16881727
{
1728+
static_assert(ActivationOperation == Activation::gelu_and_mul ||
1729+
ActivationOperation == Activation::silu_and_mul ||
1730+
ActivationOperation == Activation::swiglustep_and_mul ||
1731+
ActivationOperation == Activation::swiglu_oai_and_mul ||
1732+
ActivationOperation == Activation::gelu_tanh_and_mul,
1733+
"gridwise_moe_gemm only supports gelu_and_mul, silu_and_mul, "
1734+
"swiglustep_and_mul, swiglu_oai_and_mul and gelu_tanh_and_mul.");
16891735
ignore = b_element_op;
16901736
index_t BN0Shuffled = CalculateBN0Shuffled(problem.N);
16911737
index_t BK0Shuffled = CalculateBK0Shuffled(problem.K);
@@ -2042,6 +2088,26 @@ struct GridwiseMoeGemm : public GridwiseGemm_xdl_cshuffle_base<
20422088
tensor_operation::element_wise::Gelu{}(gate, gate);
20432089
c_thread_buf_fp32(cidx) = gate * up;
20442090
}
2091+
else if(ActivationOperation == Activation::gelu_tanh_and_mul)
2092+
{
2093+
const float scale_up =
2094+
p_scale_b[(n0 * NWave * NPerXdl + problem.N) *
2095+
PerTokenQuant];
2096+
float gate = scale_a * scale_b * c_thread_buf[cidx];
2097+
float up = scale_a * scale_up * c_thread_buf_up[cidx];
2098+
if constexpr(MulRoutedWeight)
2099+
{
2100+
gate = gate * topk_weights.template AsType<float>()[m4];
2101+
up = up * topk_weights.template AsType<float>()[m4];
2102+
}
2103+
if constexpr(is_same_v<remove_cvref_t<BDataType>, pk_i4_t>)
2104+
{
2105+
gate *= 16;
2106+
up *= 16;
2107+
}
2108+
tensor_operation::element_wise::FastGelu{}(gate, gate);
2109+
c_thread_buf_fp32(cidx) = gate * up;
2110+
}
20452111
else if constexpr(ActivationOperation ==
20462112
Activation::swiglustep_and_mul)
20472113
{
@@ -2143,6 +2209,18 @@ struct GridwiseMoeGemm : public GridwiseGemm_xdl_cshuffle_base<
21432209
tensor_operation::element_wise::Gelu{}(gate, gate);
21442210
c_thread_buf_fp32(cidx) = gate * up;
21452211
}
2212+
else if(ActivationOperation == Activation::gelu_tanh_and_mul)
2213+
{
2214+
float gate = c_thread_buf[cidx];
2215+
float up = c_thread_buf_up[cidx];
2216+
if constexpr(MulRoutedWeight)
2217+
{
2218+
gate = gate * topk_weights.template AsType<float>()[m4];
2219+
up = up * topk_weights.template AsType<float>()[m4];
2220+
}
2221+
tensor_operation::element_wise::FastGelu{}(gate, gate);
2222+
c_thread_buf_fp32(cidx) = gate * up;
2223+
}
21462224
else if constexpr(ActivationOperation ==
21472225
Activation::swiglustep_and_mul)
21482226
{

‎include/ck/tensor_operation/gpu/grid/gridwise_moe_gemm_blockscale.hpp‎

Lines changed: 46 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -1140,9 +1140,12 @@ struct GridwiseMoeGemmBlockScale
11401140
BElementwiseOperation b_element_op,
11411141
CElementwiseOperation c_element_op)
11421142
{
1143-
static_assert(ActivationOperation != Activation::swiglu_oai_and_mul,
1144-
"gridwise_moe_gemm_blockscale does not support swiglu_oai_and_mul; use the "
1145-
"non-blockscale gridwise_moe_gemm.");
1143+
static_assert(ActivationOperation == Activation::gelu_and_mul ||
1144+
ActivationOperation == Activation::silu_and_mul ||
1145+
ActivationOperation == Activation::swiglustep_and_mul ||
1146+
ActivationOperation == Activation::gelu_tanh_and_mul,
1147+
"gridwise_moe_gemm_blockscale only supports gelu_and_mul, silu_and_mul, "
1148+
"swiglustep_and_mul and gelu_tanh_and_mul.");
11461149
#if defined(__gfx942__) || defined(__gfx950__)
11471150
constexpr auto b_coherence_flag = NonTemporalLoadB
11481151
? AmdBufferCoherenceEnum::WAVE_NT1
@@ -1642,6 +1645,23 @@ struct GridwiseMoeGemmBlockScale
16421645
tensor_operation::element_wise::Gelu{}(gate, gate);
16431646
c_thread_buf(cidx) = gate * up;
16441647
}
1648+
else if(ActivationOperation == Activation::gelu_tanh_and_mul)
1649+
{
1650+
float gate = c_thread_buf[cidx];
1651+
float up = c_thread_buf_up[cidx];
1652+
if constexpr(MulRoutedWeight)
1653+
{
1654+
gate = gate * topk_weight;
1655+
up = up * topk_weight;
1656+
}
1657+
if constexpr(is_same_v<remove_cvref_t<BDataType>, pk_i4_t>)
1658+
{
1659+
gate *= 16;
1660+
up *= 16;
1661+
}
1662+
tensor_operation::element_wise::FastGelu{}(gate, gate);
1663+
c_thread_buf(cidx) = gate * up;
1664+
}
16451665
}
16461666
else
16471667
{
@@ -1697,9 +1717,12 @@ struct GridwiseMoeGemmBlockScale
16971717
BElementwiseOperation b_element_op,
16981718
CElementwiseOperation c_element_op)
16991719
{
1700-
static_assert(ActivationOperation != Activation::swiglu_oai_and_mul,
1701-
"gridwise_moe_gemm_blockscale does not support swiglu_oai_and_mul; use the "
1702-
"non-blockscale gridwise_moe_gemm.");
1720+
static_assert(ActivationOperation == Activation::gelu_and_mul ||
1721+
ActivationOperation == Activation::silu_and_mul ||
1722+
ActivationOperation == Activation::swiglustep_and_mul ||
1723+
ActivationOperation == Activation::gelu_tanh_and_mul,
1724+
"gridwise_moe_gemm_blockscale only supports gelu_and_mul, silu_and_mul, "
1725+
"swiglustep_and_mul and gelu_tanh_and_mul.");
17031726
#if defined(__gfx942__) || defined(__gfx950__)
17041727
constexpr auto b_coherence_flag = NonTemporalLoadB
17051728
? AmdBufferCoherenceEnum::WAVE_NT1
@@ -2190,6 +2213,23 @@ struct GridwiseMoeGemmBlockScale
21902213
tensor_operation::element_wise::Gelu{}(gate, gate);
21912214
c_thread_buf(cidx) = gate * up;
21922215
}
2216+
else if(ActivationOperation == Activation::gelu_tanh_and_mul)
2217+
{
2218+
float gate = c_thread_buf[cidx];
2219+
float up = c_thread_buf_up[cidx];
2220+
if constexpr(MulRoutedWeight)
2221+
{
2222+
gate = gate * topk_weight;
2223+
up = up * topk_weight;
2224+
}
2225+
if constexpr(is_same_v<remove_cvref_t<BDataType>, pk_i4_t>)
2226+
{
2227+
gate *= 16;
2228+
up *= 16;
2229+
}
2230+
tensor_operation::element_wise::FastGelu{}(gate, gate);
2231+
c_thread_buf(cidx) = gate * up;
2232+
}
21932233
}
21942234
else
21952235
{

0 commit comments

Comments
 (0)