diff --git a/src/native/cuda/metax/ops/fused_add_rms_norm/kernel.cuh b/src/native/cuda/metax/ops/fused_add_rms_norm/kernel.cuh new file mode 100644 index 000000000..829a81748 --- /dev/null +++ b/src/native/cuda/metax/ops/fused_add_rms_norm/kernel.cuh @@ -0,0 +1,77 @@ +#ifndef INFINI_OPS_METAX_FUSED_ADD_RMS_NORM_KERNEL_CUH_ +#define INFINI_OPS_METAX_FUSED_ADD_RMS_NORM_KERNEL_CUH_ + +#include + +#include "native/cuda/metax/caster.cuh" + +namespace infini::ops { + +template +struct alignas(16) MetaxNormPack { + T values[8]; +}; + +template +__global__ void MetaxFusedAddRmsNormVectorizedKernel( + T* input, int64_t input_stride, T* residual, int64_t residual_stride, + const T* weight, float epsilon) { + constexpr int kThreads = 256; + constexpr int kItems = 16; + constexpr auto kDevice = Device::Type::kMetax; + using Pack = MetaxNormPack; + auto x = reinterpret_cast(input + blockIdx.x * input_stride); + auto r = reinterpret_cast(residual + blockIdx.x * residual_stride); + auto w = reinterpret_cast(weight); + float values[kItems]; + float sum_squared = 0.0f; + +#pragma unroll + for (int pack = 0; pack < 2; ++pack) { + const int offset = threadIdx.x * 2 + pack; + Pack xv = x[offset]; + Pack rv = r[offset]; +#pragma unroll + for (int i = 0; i < 8; ++i) { + float merged = Caster::template Cast(xv.values[i]) + + Caster::template Cast(rv.values[i]); + rv.values[i] = Caster::template Cast(merged); + float rounded = Caster::template Cast(rv.values[i]); + values[pack * 8 + i] = rounded; + sum_squared += rounded * rounded; + } + r[offset] = rv; + } + + using Reduce = cub::BlockReduce; + __shared__ typename Reduce::TempStorage storage; + float total = Reduce(storage).Sum(sum_squared); + __shared__ float inverse_rms; + if (threadIdx.x == 0) { + inverse_rms = rsqrtf(total / 4096.0f + epsilon); + } + __syncthreads(); + +#pragma unroll + for (int pack = 0; pack < 2; ++pack) { + const int offset = threadIdx.x * 2 + pack; + Pack out; + Pack gamma; + if (weight != nullptr) { + gamma = w[offset]; + } +#pragma unroll + for (int i = 0; i < 8; ++i) { + float value = values[pack * 8 + i] * inverse_rms; + if (weight != nullptr) { + value *= Caster::template Cast(gamma.values[i]); + } + out.values[i] = Caster::template Cast(value); + } + x[offset] = out; + } +} + +} // namespace infini::ops + +#endif diff --git a/src/native/cuda/metax/ops/fused_add_rms_norm/kernel.h b/src/native/cuda/metax/ops/fused_add_rms_norm/kernel.h index c79d67df0..fd96fb5e0 100644 --- a/src/native/cuda/metax/ops/fused_add_rms_norm/kernel.h +++ b/src/native/cuda/metax/ops/fused_add_rms_norm/kernel.h @@ -4,6 +4,7 @@ #include #include "native/cuda/metax/caster.cuh" +#include "native/cuda/metax/ops/fused_add_rms_norm/kernel.cuh" #include "native/cuda/metax/runtime_.h" #include "native/cuda/ops/fused_add_rms_norm/kernel.h" @@ -14,6 +15,46 @@ class Operator : public CudaFusedAddRmsNorm> { public: using CudaFusedAddRmsNorm>::CudaFusedAddRmsNorm; + + void operator()(Tensor input, Tensor residual, + const std::optional weight, + float epsilon) const override { + if (num_tokens_ == 0) { + return; + } + + if (num_tokens_ > 128 || dim_ != 4096 || + input_strides_[input_strides_.size() - 2] % 8 != 0 || + residual_strides_[residual_strides_.size() - 2] % 8 != 0 || + reinterpret_cast(input.data()) % 16 != 0 || + reinterpret_cast(residual.data()) % 16 != 0 || + (weight.has_value() && + reinterpret_cast(weight->data()) % 16 != 0)) { + CudaFusedAddRmsNorm>::operator()( + input, residual, weight, epsilon); + return; + } + + auto cuda_stream = static_cast::Stream>( + stream_ ? stream_ : 0); + DispatchFunc, ReducedFloatTypes>>( + input.dtype(), + [&](auto tag) { + using T = typename decltype(tag)::type; + auto weight_data = weight.has_value() + ? reinterpret_cast(weight->data()) + : nullptr; + MetaxFusedAddRmsNormVectorizedKernel + <<(num_tokens_), 256, 0, cuda_stream>>>( + reinterpret_cast(input.data()), + input_strides_[input_strides_.size() - 2], + reinterpret_cast(residual.data()), + residual_strides_[residual_strides_.size() - 2], weight_data, + epsilon_); + }, + "MetaxFusedAddRmsNorm::operator()"); + } }; } // namespace infini::ops diff --git a/tests/test_fused_add_rms_norm.py b/tests/test_fused_add_rms_norm.py index 9cbd9e765..795b7dcc9 100644 --- a/tests/test_fused_add_rms_norm.py +++ b/tests/test_fused_add_rms_norm.py @@ -15,6 +15,11 @@ ((15, 3584), None, None), ((2, 32769), None, None), ((2, 3, 4, 128), (3072, 1024, 256, 1), (3840, 1280, 320, 1)), + ((1, 4096), None, None), + ((128, 4096), None, None), + ((129, 4096), None, None), + ((4, 4096), (4104, 1), (4112, 1)), + ((4, 4096), (4097, 1), (4099, 1)), )