[MUSA][8/N] Port CUDA kernels that are compatible with MUSA (#17946)
Signed-off-by: yafeng.li <yafeng.li@mthreads.com> Co-authored-by: Alex Nails <alex.nails@radixark.ai>
This commit is contained in:
@@ -964,11 +964,11 @@ static __device__ __forceinline__ dst_t convert_from_half(half val) {
|
||||
|
||||
template <>
|
||||
__device__ __forceinline__ c10::BFloat16 convert_from_half<c10::BFloat16>(half val) {
|
||||
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 800
|
||||
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 800 || defined(USE_MUSA)
|
||||
return __float2bfloat16(__half2float(val));
|
||||
#else
|
||||
return __half2float(val);
|
||||
#endif // defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 800
|
||||
#endif // defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 800 || defined(USE_MUSA)
|
||||
}
|
||||
|
||||
template <>
|
||||
|
||||
@@ -48,7 +48,7 @@ static __device__ __forceinline__ int get_int_from_uint8_aligned(const uint8_t*
|
||||
template <int vdr>
|
||||
static __device__ __forceinline__ float
|
||||
vec_dot_q4_0_q8_1_impl(const int* v, const int* u, const float& d4, const half2& ds8) {
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA
|
||||
int sumi = 0;
|
||||
|
||||
#pragma unroll
|
||||
@@ -74,7 +74,7 @@ vec_dot_q4_0_q8_1_impl(const int* v, const int* u, const float& d4, const half2&
|
||||
template <int vdr>
|
||||
static __device__ __forceinline__ float
|
||||
vec_dot_q4_1_q8_1_impl(const int* v, const int* u, const half2& dm4, const half2& ds8) {
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA
|
||||
int sumi = 0;
|
||||
|
||||
#pragma unroll
|
||||
@@ -102,7 +102,7 @@ vec_dot_q4_1_q8_1_impl(const int* v, const int* u, const half2& dm4, const half2
|
||||
template <int vdr>
|
||||
static __device__ __forceinline__ float
|
||||
vec_dot_q5_0_q8_1_impl(const int* vl, const int* vh, const int* u, const float& d5, const half2& ds8) {
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA
|
||||
int sumi = 0;
|
||||
|
||||
#pragma unroll
|
||||
@@ -135,7 +135,7 @@ vec_dot_q5_0_q8_1_impl(const int* vl, const int* vh, const int* u, const float&
|
||||
template <int vdr>
|
||||
static __device__ __forceinline__ float
|
||||
vec_dot_q5_1_q8_1_impl(const int* vl, const int* vh, const int* u, const half2& dm5, const half2& ds8) {
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA
|
||||
int sumi = 0;
|
||||
|
||||
#pragma unroll
|
||||
@@ -170,7 +170,7 @@ vec_dot_q5_1_q8_1_impl(const int* vl, const int* vh, const int* u, const half2&
|
||||
template <int vdr>
|
||||
static __device__ __forceinline__ float
|
||||
vec_dot_q8_0_q8_1_impl(const int* v, const int* u, const float& d8_0, const float& d8_1) {
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA
|
||||
int sumi = 0;
|
||||
|
||||
#pragma unroll
|
||||
@@ -185,7 +185,7 @@ vec_dot_q8_0_q8_1_impl(const int* v, const int* u, const float& d8_0, const floa
|
||||
template <int vdr>
|
||||
static __device__ __forceinline__ float
|
||||
vec_dot_q8_1_q8_1_impl(const int* v, const int* u, const half2& dm8, const half2& ds8) {
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA
|
||||
|
||||
int sumi = 0;
|
||||
|
||||
@@ -214,7 +214,7 @@ static __device__ __forceinline__ float vec_dot_q2_K_q8_1_impl_mmvq(
|
||||
const uint8_t* __restrict__ scales,
|
||||
const half2& dm2,
|
||||
const float* __restrict__ d8) {
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA
|
||||
float sumf_d = 0.0f;
|
||||
float sumf_m = 0.0f;
|
||||
|
||||
@@ -245,7 +245,7 @@ static __device__ __forceinline__ float vec_dot_q2_K_q8_1_impl_mmq(
|
||||
const uint8_t* __restrict__ scales,
|
||||
const half2& dm2,
|
||||
const float& d8) {
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA
|
||||
int sumi_d = 0;
|
||||
int sumi_m = 0;
|
||||
|
||||
@@ -287,7 +287,7 @@ static __device__ __forceinline__ float vec_dot_q3_K_q8_1_impl_mmvq(
|
||||
const int& scale_offset,
|
||||
const float& d3,
|
||||
const float* __restrict__ d8) {
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA
|
||||
|
||||
float sumf = 0.0f;
|
||||
|
||||
@@ -324,7 +324,7 @@ static __device__ __forceinline__ float vec_dot_q3_K_q8_1_impl_mmq(
|
||||
const int8_t* __restrict__ scales,
|
||||
const float& d3,
|
||||
const float& d8) {
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA
|
||||
int sumi = 0;
|
||||
|
||||
#pragma unroll
|
||||
@@ -353,7 +353,7 @@ static __device__ __forceinline__ float vec_dot_q4_K_q8_1_impl_vmmq(
|
||||
const uint8_t* __restrict__ m,
|
||||
const half2& dm4,
|
||||
const float* __restrict__ d8) {
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA
|
||||
|
||||
float sumf_d = 0.0f;
|
||||
float sumf_m = 0.0f;
|
||||
@@ -382,7 +382,7 @@ static __device__ __forceinline__ float vec_dot_q4_K_q8_1_impl_mmq(
|
||||
const uint8_t* __restrict__ m,
|
||||
const half2& dm4,
|
||||
const half2* __restrict__ ds8) {
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA
|
||||
float sumf_d = 0.0f;
|
||||
float sumf_m = 0.0f;
|
||||
|
||||
@@ -418,7 +418,7 @@ static __device__ __forceinline__ float vec_dot_q5_K_q8_1_impl_vmmq(
|
||||
const uint8_t* __restrict__ m,
|
||||
const half2& dm5,
|
||||
const float* __restrict__ d8) {
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA
|
||||
|
||||
float sumf_d = 0.0f;
|
||||
float sumf_m = 0.0f;
|
||||
@@ -453,7 +453,7 @@ static __device__ __forceinline__ float vec_dot_q5_K_q8_1_impl_mmq(
|
||||
const uint8_t* __restrict__ m,
|
||||
const half2& dm4,
|
||||
const half2* __restrict__ ds8) {
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA
|
||||
float sumf_d = 0.0f;
|
||||
float sumf_m = 0.0f;
|
||||
|
||||
@@ -489,7 +489,7 @@ static __device__ __forceinline__ float vec_dot_q6_K_q8_1_impl_mmvq(
|
||||
const int8_t* __restrict__ scales,
|
||||
const float& d,
|
||||
const float* __restrict__ d8) {
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA
|
||||
float sumf = 0.0f;
|
||||
|
||||
#pragma unroll
|
||||
@@ -512,7 +512,7 @@ static __device__ __forceinline__ float vec_dot_q6_K_q8_1_impl_mmq(
|
||||
const int8_t* __restrict__ sc,
|
||||
const float& d6,
|
||||
const float* __restrict__ d8) {
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA
|
||||
float sumf_d = 0.0f;
|
||||
|
||||
#pragma unroll
|
||||
@@ -1807,7 +1807,7 @@ vec_dot_iq2_xs_q8_1(const void* __restrict__ vbq, const block_q8_1* __restrict__
|
||||
|
||||
static __device__ __forceinline__ float
|
||||
vec_dot_iq2_s_q8_1(const void* __restrict__ vbq, const block_q8_1* __restrict__ bq8_1, const int& iqs) {
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA
|
||||
const block_iq2_s* bq2 = (const block_iq2_s*)vbq;
|
||||
|
||||
const int ib32 = iqs;
|
||||
@@ -1846,7 +1846,7 @@ vec_dot_iq2_s_q8_1(const void* __restrict__ vbq, const block_q8_1* __restrict__
|
||||
|
||||
static __device__ __forceinline__ float
|
||||
vec_dot_iq3_xxs_q8_1(const void* __restrict__ vbq, const block_q8_1* __restrict__ bq8_1, const int& iqs) {
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA
|
||||
const block_iq3_xxs* bq2 = (const block_iq3_xxs*)vbq;
|
||||
|
||||
const int ib32 = iqs;
|
||||
@@ -1873,7 +1873,7 @@ vec_dot_iq3_xxs_q8_1(const void* __restrict__ vbq, const block_q8_1* __restrict_
|
||||
|
||||
static __device__ __forceinline__ float
|
||||
vec_dot_iq3_s_q8_1(const void* __restrict__ vbq, const block_q8_1* __restrict__ bq8_1, const int& iqs) {
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA
|
||||
const block_iq3_s* bq2 = (const block_iq3_s*)vbq;
|
||||
|
||||
const int ib32 = iqs;
|
||||
@@ -1899,7 +1899,7 @@ vec_dot_iq3_s_q8_1(const void* __restrict__ vbq, const block_q8_1* __restrict__
|
||||
|
||||
static __device__ __forceinline__ float
|
||||
vec_dot_iq1_s_q8_1(const void* __restrict__ vbq, const block_q8_1* __restrict__ bq8_1, const int& iqs) {
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA
|
||||
const block_iq1_s* bq1 = (const block_iq1_s*)vbq;
|
||||
|
||||
const int qs_packed = get_int_b2(bq1->qs, iqs);
|
||||
@@ -1931,7 +1931,7 @@ vec_dot_iq1_s_q8_1(const void* __restrict__ vbq, const block_q8_1* __restrict__
|
||||
|
||||
static __device__ __forceinline__ float
|
||||
vec_dot_iq1_m_q8_1(const void* __restrict__ vbq, const block_q8_1* __restrict__ bq8_1, const int& iqs) {
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA
|
||||
|
||||
const block_iq1_m* bq1 = (const block_iq1_m*)vbq;
|
||||
|
||||
@@ -1991,7 +1991,7 @@ get_int_from_table_16(const uint32_t& q4, const uint8_t* values, int& val1, int&
|
||||
|
||||
static __device__ __forceinline__ float
|
||||
vec_dot_iq4_nl_q8_1(const void* __restrict__ vbq, const block_q8_1* __restrict__ bq8_1, const int& iqs) {
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA
|
||||
|
||||
const block_iq4_nl* bq = (const block_iq4_nl*)vbq;
|
||||
|
||||
@@ -2015,7 +2015,7 @@ vec_dot_iq4_nl_q8_1(const void* __restrict__ vbq, const block_q8_1* __restrict__
|
||||
|
||||
static __device__ __forceinline__ float
|
||||
vec_dot_iq4_xs_q8_1(const void* __restrict__ vbq, const block_q8_1* __restrict__ bq8_1, const int& iqs) {
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA
|
||||
const block_iq4_xs* bq4 = (const block_iq4_xs*)vbq;
|
||||
const uint8_t* values = (const uint8_t*)kvalues_iq4nl;
|
||||
|
||||
|
||||
Reference in New Issue
Block a user