Skip to content
332 changes: 332 additions & 0 deletions ggml/src/ggml-cuda/mmq-load-tiles.cuh
Original file line number Diff line number Diff line change
Expand Up @@ -263,6 +263,90 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
}
}

#if defined(GGML_USE_HIP) && defined(RDNA3_5)
template <ggml_type type, int J, bool fallback>
static __device__ __forceinline__ void ggml_cuda_mmq_prefetch_tiles_q4_0_rdna35(
const char * __restrict__ x, const int kbx0, const int i_max, const int stride,
int (&qs_cache)[ggml_cuda_mmq_get_I(type, J, fallback)/
(ggml_cuda_mmq_get_nthreads(type, J, fallback)/ggml_cuda_get_physical_warp_size())],
float (&d_cache)[4]) {
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
constexpr int nthreads = ggml_cuda_mmq_get_nthreads(type, J, fallback);
constexpr int nwarps = nthreads / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
static_assert(warp_size == 32 && nthreads == 128 && I == 64, "unexpected RDNA3.5 Q4_0 MMQ configuration");

const int kbx = threadIdx.x / QI4_0;
const int kqsx = threadIdx.x % QI4_0;

#pragma unroll
for (int i0 = 0; i0 < I; i0 += nwarps) {
int i = i0 + threadIdx.y;
if constexpr (fallback) {
i = min(i, i_max);
}

const block_q4_0 * bxi = (const block_q4_0 *) x + kbx0 + i*stride + kbx;
qs_cache[i0/nwarps] = get_int_b2(bxi->qs, kqsx);
}

constexpr int blocks_per_tile_x_row = MMQ_TILE_NE_K / QI4_0;
constexpr int scale_rows_per_warp = warp_size / blocks_per_tile_x_row;
const int kbxd = threadIdx.x % blocks_per_tile_x_row;
int d_idx = 0;
#pragma unroll
for (int i0 = 0; i0 < I; i0 += nwarps * scale_rows_per_warp) {
int i = i0 + threadIdx.y * scale_rows_per_warp + threadIdx.x / blocks_per_tile_x_row;
if constexpr (fallback) {
i = min(i, i_max);
}

const block_q4_0 * bxi = (const block_q4_0 *) x + kbx0 + i*stride + kbxd;
d_cache[d_idx++] = bxi->d;
}

asm volatile("" ::: "memory");
}

template <ggml_type type, int J, bool fallback>
static __device__ __forceinline__ void ggml_cuda_mmq_store_tiles_q4_0_rdna35(
int * __restrict__ x_tile,
const int (&qs_cache)[ggml_cuda_mmq_get_I(type, J, fallback)/
(ggml_cuda_mmq_get_nthreads(type, J, fallback)/ggml_cuda_get_physical_warp_size())],
const float (&d_cache)[4]) {
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
constexpr int nthreads = ggml_cuda_mmq_get_nthreads(type, J, fallback);
constexpr int nwarps = nthreads / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
static_assert(warp_size == 32 && nthreads == 128 && I == 64, "unexpected RDNA3.5 Q4_0 MMQ configuration");

int * x_qs = x_tile;
float * x_df = (float *) (x_qs + 2*MMQ_TILE_NE_K);
const int txi = threadIdx.x;
const int kbx = txi / QI4_0;
const int kqsx = txi % QI4_0;

#pragma unroll
for (int i0 = 0; i0 < I; i0 += nwarps) {
const int i = i0 + threadIdx.y;
const int qs0 = qs_cache[i0/nwarps];
x_qs[i*sram_stride + kbx*(2*QI4_0) + kqsx + 0] = __vsubss4((qs0 >> 0) & 0x0F0F0F0F, 0x08080808);
x_qs[i*sram_stride + kbx*(2*QI4_0) + kqsx + QI4_0] = __vsubss4((qs0 >> 4) & 0x0F0F0F0F, 0x08080808);
}

constexpr int blocks_per_tile_x_row = MMQ_TILE_NE_K / QI4_0;
constexpr int scale_rows_per_warp = warp_size / blocks_per_tile_x_row;
const int kbxd = threadIdx.x % blocks_per_tile_x_row;
int d_idx = 0;
#pragma unroll
for (int i0 = 0; i0 < I; i0 += nwarps * scale_rows_per_warp) {
const int i = i0 + threadIdx.y * scale_rows_per_warp + threadIdx.x / blocks_per_tile_x_row;
x_df[i*sram_stride + kbxd] = d_cache[d_idx++];
}
}
#endif // defined(GGML_USE_HIP) && defined(RDNA3_5)

template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_q4_1(
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
Expand Down Expand Up @@ -548,6 +632,89 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
}
}

#if defined(GGML_USE_HIP) && defined(RDNA3_5)
template <ggml_type type, int J, bool fallback>
static __device__ __forceinline__ void ggml_cuda_mmq_prefetch_tiles_q8_0_rdna35(
const char * __restrict__ x, const int kbx0, const int i_max, const int stride,
int (&qs_cache)[2 * ggml_cuda_mmq_get_I(type, J, fallback)/
(ggml_cuda_mmq_get_nthreads(type, J, fallback)/ggml_cuda_get_physical_warp_size())],
float (&d_cache)[4]) {
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
constexpr int nthreads = ggml_cuda_mmq_get_nthreads(type, J, fallback);
constexpr int nwarps = nthreads / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
static_assert(warp_size == 32 && nthreads == 128 && I == 64, "unexpected RDNA3.5 Q8_0 MMQ configuration");

const int txi = threadIdx.x;
const int kbx = txi / QI8_0;
const int kqsx = txi % QI8_0;

#pragma unroll
for (int i0 = 0; i0 < I; i0 += nwarps) {
int i = i0 + threadIdx.y;
if constexpr (fallback) {
i = min(i, i_max);
}

const block_q8_0 * bxi = (const block_q8_0 *) x + kbx0 + i*stride + kbx;
qs_cache[2*(i0/nwarps) + 0] = get_int_b2(bxi[0].qs, kqsx);
qs_cache[2*(i0/nwarps) + 1] = get_int_b2(bxi[MMQ_TILE_NE_K/QI8_0].qs, kqsx);
}

constexpr int blocks_per_tile_x_row = 2*MMQ_TILE_NE_K / QI8_0;
constexpr int scale_rows_per_warp = warp_size / blocks_per_tile_x_row;
const int kbxd = threadIdx.x % blocks_per_tile_x_row;
int d_idx = 0;
#pragma unroll
for (int i0 = 0; i0 < I; i0 += nwarps * scale_rows_per_warp) {
int i = i0 + threadIdx.y * scale_rows_per_warp + threadIdx.x / blocks_per_tile_x_row;
if constexpr (fallback) {
i = min(i, i_max);
}

const block_q8_0 * bxi = (const block_q8_0 *) x + kbx0 + i*stride + kbxd;
d_cache[d_idx++] = bxi->d;
}

asm volatile("" ::: "memory");
}

template <ggml_type type, int J, bool fallback>
static __device__ __forceinline__ void ggml_cuda_mmq_store_tiles_q8_0_rdna35(
int * __restrict__ x_tile,
const int (&qs_cache)[2 * ggml_cuda_mmq_get_I(type, J, fallback)/
(ggml_cuda_mmq_get_nthreads(type, J, fallback)/ggml_cuda_get_physical_warp_size())],
const float (&d_cache)[4]) {
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
constexpr int nthreads = ggml_cuda_mmq_get_nthreads(type, J, fallback);
constexpr int nwarps = nthreads / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
static_assert(warp_size == 32 && nthreads == 128 && I == 64, "unexpected RDNA3.5 Q8_0 MMQ configuration");

int * x_qs = x_tile;
float * x_df = (float *) (x_tile + 2*MMQ_TILE_NE_K);
const int txi = threadIdx.x;

#pragma unroll
for (int i0 = 0; i0 < I; i0 += nwarps) {
const int i = i0 + threadIdx.y;
x_qs[i*sram_stride + 0 + txi] = qs_cache[2*(i0/nwarps) + 0];
x_qs[i*sram_stride + MMQ_TILE_NE_K + txi] = qs_cache[2*(i0/nwarps) + 1];
}

constexpr int blocks_per_tile_x_row = 2*MMQ_TILE_NE_K / QI8_0;
constexpr int scale_rows_per_warp = warp_size / blocks_per_tile_x_row;
const int kbxd = threadIdx.x % blocks_per_tile_x_row;
int d_idx = 0;
#pragma unroll
for (int i0 = 0; i0 < I; i0 += nwarps * scale_rows_per_warp) {
const int i = i0 + threadIdx.y * scale_rows_per_warp + threadIdx.x / blocks_per_tile_x_row;
x_df[i*sram_stride + kbxd] = d_cache[d_idx++];
}
}
#endif // defined(GGML_USE_HIP) && defined(RDNA3_5)

// ---------------------------------------------------------------------------------------------

template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_q2_K(
Expand Down Expand Up @@ -845,6 +1012,166 @@ static __device__ __forceinline__ void ggml_cuda_mmq_store_tiles_q4_K_rdna35(
const uint8_t * m8 = (const uint8_t *) &m32;
const half2 dm = dm_cache * make_half2(1.0f, -1.0f);

#pragma unroll
for (int l = 0; l < int(sizeof(int)); ++l) {
x_dm[i*sram_stride + sizeof(int)*ksc + l] = dm*make_half2(sc8[l], m8[l]);
}
}

template <ggml_type type, int J, bool fallback>
static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_q5_K_rdna35(
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
constexpr int nthreads = ggml_cuda_mmq_get_nthreads(type, J, fallback);
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
static_assert(warp_size == 32 && nthreads == 128 && I == 64, "unexpected RDNA3.5 Q5_K MMQ configuration");

int * x_qs = (int *) x_tile;
half2 * x_dm = (half2 *) (x_qs + 2*MMQ_TILE_NE_K);

const int linear_tid = threadIdx.y*warp_size + threadIdx.x;
int i = linear_tid/2;
if (fallback) {
i = min(i, i_max);
}

const block_q5_K * bxi = (const block_q5_K *) x + kbx0 + i*stride;
const int * scales = (const int *) bxi->scales;
const int ksc = linear_tid % 2;
const int sc32 = unpack_scales_q45_K(scales, ksc);
const int m32 = unpack_scales_q45_K(scales, ksc + 2);
const uint8_t * sc8 = (const uint8_t *) &sc32;
const uint8_t * m8 = (const uint8_t *) &m32;
const half2 dm = bxi->dm * make_half2(1.0f, -1.0f);

#pragma unroll
for (int l = 0; l < int(sizeof(int)); ++l) {
x_dm[i*sram_stride + sizeof(int)*ksc + l] = dm*make_half2(sc8[l], m8[l]);
}

const int txi = threadIdx.x;
const int kqs = 16*(txi/8) + txi%8;
const int qh_shift0 = 2*(txi/8);

#pragma unroll
for (int i0 = 0; i0 < I; i0 += nthreads/warp_size) {
int row = i0 + threadIdx.y;
if (fallback) {
row = min(row, i_max);
}

const block_q5_K * bxq = (const block_q5_K *) x + kbx0 + row*stride;
const int qs = ((const int *) bxq->qs)[txi];
const int qh = ((const int *) bxq->qh)[txi % (QI5_K/4)];
int * row_qs = x_qs + row*sram_stride;
row_qs[kqs] = (qs & 0x0F0F0F0F) | (((qh >> (qh_shift0 + 0)) << 4) & 0x10101010);
row_qs[kqs+8] = ((qs >> 4) & 0x0F0F0F0F) | (((qh >> (qh_shift0 + 1)) << 4) & 0x10101010);
}
}

template <ggml_type type, int J, bool fallback>
static __device__ __forceinline__ void ggml_cuda_mmq_prefetch_tiles_q5_K_rdna35(
const char * __restrict__ x, const int kbx0, const int i_max, const int stride,
int (&qs_cache)[ggml_cuda_mmq_get_I(type, J, fallback)/
(ggml_cuda_mmq_get_nthreads(type, J, fallback)/ggml_cuda_get_physical_warp_size())],
int (&qh_cache)[ggml_cuda_mmq_get_I(type, J, fallback)*(QI5_K/4)/
ggml_cuda_mmq_get_nthreads(type, J, fallback)],
int (&scales_cache)[3], half2 & dm_cache) {
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
constexpr int nthreads = ggml_cuda_mmq_get_nthreads(type, J, fallback);
constexpr int nwarps = nthreads / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
constexpr int qh_words_per_row = QI5_K/4;
constexpr int qh_cache_size = I*qh_words_per_row/nthreads;
constexpr int rows_per_warp = I/nwarps;
static_assert(warp_size == 32 && nthreads == 128 && I == 64, "unexpected RDNA3.5 Q5_K MMQ configuration");
static_assert(qh_cache_size*warp_size == rows_per_warp*qh_words_per_row,
"Q5_K high bits must be distributed evenly across the warp");

#pragma unroll
for (int i0 = 0; i0 < I; i0 += nwarps) {
int i = i0 + threadIdx.y;
if constexpr (fallback) {
i = min(i, i_max);
}

const block_q5_K * bxi = (const block_q5_K *) x + kbx0 + i*stride;
qs_cache[i0/nwarps] = ((const int *) bxi->qs)[threadIdx.x];
}

#pragma unroll
for (int l = 0; l < qh_cache_size; ++l) {
const int qh_linear = l*warp_size + threadIdx.x;
int i = (qh_linear/qh_words_per_row)*nwarps + threadIdx.y;
if constexpr (fallback) {
i = min(i, i_max);
}

const block_q5_K * bxi = (const block_q5_K *) x + kbx0 + i*stride;
qh_cache[l] = ((const int *) bxi->qh)[qh_linear % qh_words_per_row];
}

int i = (threadIdx.y*warp_size + threadIdx.x)/2;
if constexpr (fallback) {
i = min(i, i_max);
}

const block_q5_K * bxi = (const block_q5_K *) x + kbx0 + i*stride;
#pragma unroll
for (int l = 0; l < 3; ++l) {
scales_cache[l] = ((const int *) bxi->scales)[l];
}
dm_cache = bxi->dm;

asm volatile("" ::: "memory");
}

template <ggml_type type, int J, bool fallback>
static __device__ __forceinline__ void ggml_cuda_mmq_store_tiles_q5_K_rdna35(
int * __restrict__ x_tile,
const int (&qs_cache)[ggml_cuda_mmq_get_I(type, J, fallback)/
(ggml_cuda_mmq_get_nthreads(type, J, fallback)/ggml_cuda_get_physical_warp_size())],
const int (&qh_cache)[ggml_cuda_mmq_get_I(type, J, fallback)*(QI5_K/4)/
ggml_cuda_mmq_get_nthreads(type, J, fallback)],
const int (&scales_cache)[3], const half2 dm_cache) {
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
constexpr int nthreads = ggml_cuda_mmq_get_nthreads(type, J, fallback);
constexpr int nwarps = nthreads / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
constexpr int qh_words_per_row = QI5_K/4;
constexpr int qh_rows_per_slot = warp_size/qh_words_per_row;
static_assert(warp_size == 32 && nthreads == 128 && I == 64, "unexpected RDNA3.5 Q5_K MMQ configuration");

int * x_qs = x_tile;
const int txi = threadIdx.x;
const int kqs = 16*(txi/8) + txi%8;
const int qh_shift0 = 2*(txi/8);

#pragma unroll
for (int i0 = 0; i0 < I; i0 += nwarps) {
const int row_in_warp = i0/nwarps;
const int qh_slot = row_in_warp/qh_rows_per_slot;
const int qh_src_lane = (row_in_warp % qh_rows_per_slot)*qh_words_per_row + txi%qh_words_per_row;
const int qs = qs_cache[row_in_warp];
const int qh = __shfl_sync(0xFFFFFFFF, qh_cache[qh_slot], qh_src_lane, warp_size);
const int i = i0 + threadIdx.y;
int * row_qs = x_qs + i*sram_stride;
row_qs[kqs] = (qs & 0x0F0F0F0F) | (((qh >> (qh_shift0 + 0)) << 4) & 0x10101010);
row_qs[kqs+8] = ((qs >> 4) & 0x0F0F0F0F) | (((qh >> (qh_shift0 + 1)) << 4) & 0x10101010);
}

half2 * x_dm = (half2 *) (x_qs + 2*MMQ_TILE_NE_K);
const int linear_tid = threadIdx.y*warp_size + threadIdx.x;
const int i = linear_tid/2;
const int ksc = linear_tid%2;
const int sc32 = unpack_scales_q45_K(scales_cache, ksc);
const int m32 = unpack_scales_q45_K(scales_cache, ksc + 2);
const uint8_t * sc8 = (const uint8_t *) &sc32;
const uint8_t * m8 = (const uint8_t *) &m32;
const half2 dm = dm_cache * make_half2(1.0f, -1.0f);

#pragma unroll
for (int l = 0; l < int(sizeof(int)); ++l) {
x_dm[i*sram_stride + sizeof(int)*ksc + l] = dm*make_half2(sc8[l], m8[l]);
Expand Down Expand Up @@ -970,6 +1297,11 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_

template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_q5_K(
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
#if defined(RDNA3_5)
ggml_cuda_mmq_load_tiles_q5_K_rdna35<type, J, fallback>(x, x_tile, kbx0, i_max, stride);
return;
#endif

constexpr int warp_size = ggml_cuda_get_physical_warp_size();
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
Expand Down
Loading
Loading