[Kernel] Dispatch exp/sin/cos through dtype_trait (#19798)
This commit is contained in:
@@ -1,8 +1,6 @@
|
||||
#pragma once
|
||||
#include <sgl_kernel/type.cuh>
|
||||
|
||||
#include <cmath>
|
||||
|
||||
namespace device::math {
|
||||
|
||||
inline constexpr float log2e = 1.44269504088896340736f;
|
||||
@@ -35,16 +33,19 @@ SGL_DEVICE T rsqrt(T a) {
|
||||
return dtype_trait<T>::rsqrt(a);
|
||||
}
|
||||
|
||||
SGL_DEVICE float exp(float a) {
|
||||
return ::expf(a);
|
||||
template <typename T>
|
||||
SGL_DEVICE T exp(T a) {
|
||||
return dtype_trait<T>::exp(a);
|
||||
}
|
||||
|
||||
SGL_DEVICE float sin(float a) {
|
||||
return ::sinf(a);
|
||||
template <typename T>
|
||||
SGL_DEVICE T sin(T a) {
|
||||
return dtype_trait<T>::sin(a);
|
||||
}
|
||||
|
||||
SGL_DEVICE float cos(float a) {
|
||||
return ::cosf(a);
|
||||
template <typename T>
|
||||
SGL_DEVICE T cos(T a) {
|
||||
return dtype_trait<T>::cos(a);
|
||||
}
|
||||
|
||||
} // namespace device::math
|
||||
|
||||
@@ -43,6 +43,9 @@ SGL_REGISTER_DTYPE_TRAIT(
|
||||
SGL_REGISTER_UNARY_FUNCTION(abs, fabsf);
|
||||
SGL_REGISTER_UNARY_FUNCTION(sqrt, sqrtf);
|
||||
SGL_REGISTER_UNARY_FUNCTION(rsqrt, rsqrtf);
|
||||
SGL_REGISTER_UNARY_FUNCTION(exp, expf);
|
||||
SGL_REGISTER_UNARY_FUNCTION(sin, sinf);
|
||||
SGL_REGISTER_UNARY_FUNCTION(cos, cosf);
|
||||
SGL_REGISTER_BINARY_FUNCTION(max, fmaxf);
|
||||
SGL_REGISTER_BINARY_FUNCTION(min, fminf););
|
||||
SGL_REGISTER_DTYPE_TRAIT(fp16_t, fp16x2_t);
|
||||
|
||||
Reference in New Issue
Block a user