diff --git a/ggml/src/ggml-cuda/fattn-tile.cuh b/ggml/src/ggml-cuda/fattn-tile.cuh index 0a099810e14..3eb81f7385d 100644 --- a/ggml/src/ggml-cuda/fattn-tile.cuh +++ b/ggml/src/ggml-cuda/fattn-tile.cuh @@ -370,6 +370,27 @@ static constexpr __device__ int ggml_cuda_fattn_tile_get_nbatch_K(const int DKQ, return (ggml_cuda_fattn_tile_get_config(DKQ, DV, ncols) >> 23) & ((1 << 9) - 1); } +// fp32 storage roughly doubles the tile kernel's LDS; shrink nbatch_K until it fits one block. +static constexpr __device__ int ggml_cuda_fattn_tile_get_nbatch_K_f32( + const int DKQ, const int DV, const int ncols, const int nbatch_fa, const int nbatch_K_f16) { + constexpr int warp_size = 32; + const int DVp = (DV + 2*warp_size - 1) & ~(2*warp_size - 1); + const int cpy_ne = ggml_cuda_get_max_cpy_bytes() / 4; + const int budget = (60*1024) / 4; // floats; stay safely below 64 KiB LDS for one block + // 0 means fp32 won't fit this shape, so the caller keeps fp16. + const int step = 8; + for (int nk = (nbatch_K_f16 / step) * step; nk >= step; nk -= step) { + if ((DV % nk) != 0) { + continue; + } + const int floats = ncols*DKQ + (nbatch_fa*(nk + cpy_ne) + DVp - DV) + ncols*nbatch_fa; + if (floats <= budget) { + return nk; + } + } + return 0; +} + // TODO: deduplicate with mma-f16 template static __device__ __forceinline__ void flash_attn_tile_load_tile( @@ -499,39 +520,25 @@ static __device__ __forceinline__ void flash_attn_tile_iter_KQ( (K_h2 + int64_t(k_VKQ_0)*stride_K2 + k_KQ_0/2, KV_tmp, stride_K2, k_VKQ_sup); __syncthreads(); -#ifdef FAST_FP16_AVAILABLE - static_assert((nbatch_K/2) % cpy_ne == 0, "bad nbatch_K"); -#pragma unroll - for (int k_KQ_1 = 0; k_KQ_1 < nbatch_K/2; k_KQ_1 += cpy_ne) { - __align__(16) half2 K_k[nbatch_fa/(np*warp_size)][cpy_ne]; - __align__(16) half2 Q_k[cpw][cpy_ne]; -#else - static_assert(nbatch_K % cpy_ne == 0, "bad nbatch_K"); + constexpr bool use_fp16 = std::is_same::value; + constexpr int kpack = use_fp16 ? 2 : 1; // scalar values packed per T_vec_dot element + static_assert((nbatch_K/kpack) % cpy_ne == 0, "bad nbatch_K"); #pragma unroll - for (int k_KQ_1 = 0; k_KQ_1 < nbatch_K; k_KQ_1 += cpy_ne) { - __align__(16) float K_k[nbatch_fa/(np*warp_size)][cpy_ne]; - __align__(16) float Q_k[cpw][cpy_ne]; -#endif // FAST_FP16_AVAILABLE + for (int k_KQ_1 = 0; k_KQ_1 < nbatch_K/kpack; k_KQ_1 += cpy_ne) { + __align__(16) T_vec_dot K_k[nbatch_fa/(np*warp_size)][cpy_ne]; + __align__(16) T_vec_dot Q_k[cpw][cpy_ne]; #pragma unroll for (int i_KQ_0 = 0; i_KQ_0 < nbatch_fa; i_KQ_0 += np*warp_size) { const int i_KQ = i_KQ_0 + (threadIdx.y % np)*warp_size + threadIdx.x; -#ifdef FAST_FP16_AVAILABLE - ggml_cuda_memcpy_1(&K_k[i_KQ_0/(np*warp_size)], &KV_tmp[i_KQ*(nbatch_K/2 + cpy_ne) + k_KQ_1]); -#else - ggml_cuda_memcpy_1(&K_k[i_KQ_0/(np*warp_size)], &KV_tmp[i_KQ*(nbatch_K + cpy_ne) + k_KQ_1]); -#endif // FAST_FP16_AVAILABLE + ggml_cuda_memcpy_1(&K_k[i_KQ_0/(np*warp_size)], &KV_tmp[i_KQ*(nbatch_K/kpack + cpy_ne) + k_KQ_1]); } #pragma unroll for (int jc0 = 0; jc0 < cpw; ++jc0) { const int jc = jc0 + (threadIdx.y / np)*cpw; -#ifdef FAST_FP16_AVAILABLE - ggml_cuda_memcpy_1(&Q_k[jc0], &Q_tmp[jc*(DKQ/2) + k_KQ_0/2 + k_KQ_1]); -#else - ggml_cuda_memcpy_1(&Q_k[jc0], &Q_tmp[jc* DKQ + k_KQ_0 + k_KQ_1]); -#endif // FAST_FP16_AVAILABLE + ggml_cuda_memcpy_1(&Q_k[jc0], &Q_tmp[jc*(DKQ/kpack) + k_KQ_0/kpack + k_KQ_1]); } #pragma unroll @@ -584,11 +591,8 @@ static __device__ __forceinline__ void flash_attn_tile_iter( // KQ_cs == KQ chunk size, number of KQ values in j direction to store as one contiguous chunk in memory. // KQ is originally 2D but uses a Z-shaped 3D memory pattern like KQ[ncols/KQ_cs][DVp][KQ_cs]. -#ifdef FAST_FP16_AVAILABLE - constexpr int KQ_cs = cpw < 2*cpy_ne ? cpw : 2*cpy_ne; -#else - constexpr int KQ_cs = cpw < 1*cpy_ne ? cpw : 1*cpy_ne; -#endif // FAST_FP16_AVAILABLE + constexpr bool use_fp16 = std::is_same::value; + constexpr int KQ_cs = cpw < (use_fp16 ? 2 : 1)*cpy_ne ? cpw : (use_fp16 ? 2 : 1)*cpy_ne; static_assert(cpw % KQ_cs == 0, "bad KQ_cs"); const int k_VKQ_sup = k_VKQ_max - k_VKQ_0; // k supremum, only smaller k values have valid KV data @@ -623,9 +627,11 @@ static __device__ __forceinline__ void flash_attn_tile_iter( const int i_KQ = i_KQ_0 + (threadIdx.y % np)*warp_size + threadIdx.x; #if defined(FAST_FP16_AVAILABLE) && !defined(V_DOT2_F32_F16_AVAILABLE) - // Without the v_dot2_f32_f16 instruction there is a higher risk of numerical overflow in the KQ calculation. - // Therefore, scale down Q values and apply the inverse scale the FP32 KQ values afterwards again. - KQ_acc[i_KQ_0/(np*warp_size)*cpw + jc0] *= 4.0f; + if constexpr (use_fp16) { + // Without the v_dot2_f32_f16 instruction there is a higher risk of numerical overflow in the KQ calculation. + // Therefore, scale down Q values and apply the inverse scale the FP32 KQ values afterwards again. + KQ_acc[i_KQ_0/(np*warp_size)*cpw + jc0] *= 4.0f; + } #endif // defined(FAST_FP16_AVAILABLE) && !defined(V_DOT2_F32_F16_AVAILABLE) if (use_logit_softcap) { @@ -659,11 +665,7 @@ static __device__ __forceinline__ void flash_attn_tile_iter( // Calculate KQ softmax, write to shared KQ buffer, re-scale VKQ accumulators: #pragma unroll for (int jc0 = 0; jc0 < cpw; jc0 += KQ_cs) { -#ifdef FAST_FP16_AVAILABLE - __align__(16) half tmp[nbatch_fa/(np*warp_size)][KQ_cs]; -#else - __align__(16) float tmp[nbatch_fa/(np*warp_size)][KQ_cs]; -#endif // FAST_FP16_AVAILABLE + __align__(16) T_KQ tmp[nbatch_fa/(np*warp_size)][KQ_cs]; #pragma unroll for (int jc1 = 0; jc1 < KQ_cs; ++jc1) { @@ -682,19 +684,19 @@ static __device__ __forceinline__ void flash_attn_tile_iter( } KQ_sum[jc] = KQ_sum[jc]*KQ_max_scale + KQ_sum_add; -#ifdef FAST_FP16_AVAILABLE - const half2 KQ_max_scale_h2 = make_half2(KQ_max_scale, KQ_max_scale); + if constexpr (use_fp16) { + const half2 KQ_max_scale_h2 = make_half2(KQ_max_scale, KQ_max_scale); #pragma unroll - for (int i0 = 0; i0 < DVp/2; i0 += warp_size) { - VKQ[jc*((DVp/2)/warp_size) + i0/warp_size] *= KQ_max_scale_h2; - } -#else + for (int i0 = 0; i0 < DVp/2; i0 += warp_size) { + VKQ[jc*((DVp/2)/warp_size) + i0/warp_size] *= KQ_max_scale_h2; + } + } else { #pragma unroll - for (int i0 = 0; i0 < DVp/2; i0 += warp_size) { - VKQ[jc*((DVp/2)/warp_size) + i0/warp_size].x *= KQ_max_scale; - VKQ[jc*((DVp/2)/warp_size) + i0/warp_size].y *= KQ_max_scale; + for (int i0 = 0; i0 < DVp/2; i0 += warp_size) { + VKQ[jc*((DVp/2)/warp_size) + i0/warp_size].x *= KQ_max_scale; + VKQ[jc*((DVp/2)/warp_size) + i0/warp_size].y *= KQ_max_scale; + } } -#endif // FAST_FP16_AVAILABLE } #pragma unroll @@ -719,74 +721,74 @@ static __device__ __forceinline__ void flash_attn_tile_iter( (V_h2 + int64_t(k_VKQ_0 + k0)*stride_V2, KV_tmp, stride_V2, k_VKQ_sup - k0); __syncthreads(); -#ifdef FAST_FP16_AVAILABLE + if constexpr (use_fp16) { #pragma unroll - for (int k1 = 0; k1 < nbatch_V; k1 += np) { - __align__(16) half2 V_k[(DVp/2)/warp_size]; - __align__(16) half2 KQ_k[cpw]; + for (int k1 = 0; k1 < nbatch_V; k1 += np) { + __align__(16) half2 V_k[(DVp/2)/warp_size]; + __align__(16) half2 KQ_k[cpw]; - constexpr int cpy_ne_D = cpy_ne/2 < (DVp/2)/warp_size ? cpy_ne/2 : (DVp/2)/warp_size; + constexpr int cpy_ne_D = cpy_ne/2 < (DVp/2)/warp_size ? cpy_ne/2 : (DVp/2)/warp_size; #pragma unroll - for (int i0 = 0; i0 < DVp/2; i0 += warp_size*cpy_ne_D) { - ggml_cuda_memcpy_1(&V_k[i0/warp_size], &KV_tmp[(k1 + threadIdx.y % np)*(DV/2) + i0 + threadIdx.x*cpy_ne_D]); - } + for (int i0 = 0; i0 < DVp/2; i0 += warp_size*cpy_ne_D) { + ggml_cuda_memcpy_1(&V_k[i0/warp_size], &KV_tmp[(k1 + threadIdx.y % np)*(DV/2) + i0 + threadIdx.x*cpy_ne_D]); + } #pragma unroll - for (int jc_VKQ_0 = 0; jc_VKQ_0 < cpw; jc_VKQ_0 += KQ_cs) { - const int jc_KQ = jc_VKQ_0/KQ_cs + (threadIdx.y / np)*(cpw/KQ_cs); + for (int jc_VKQ_0 = 0; jc_VKQ_0 < cpw; jc_VKQ_0 += KQ_cs) { + const int jc_KQ = jc_VKQ_0/KQ_cs + (threadIdx.y / np)*(cpw/KQ_cs); - __align__(16) half tmp[KQ_cs]; - ggml_cuda_memcpy_1( - &tmp, KQ + jc_KQ*(nbatch_fa*KQ_cs) + (k0 + k1 + threadIdx.y % np)*KQ_cs); + __align__(16) half tmp[KQ_cs]; + ggml_cuda_memcpy_1( + &tmp, KQ + jc_KQ*(nbatch_fa*KQ_cs) + (k0 + k1 + threadIdx.y % np)*KQ_cs); #pragma unroll - for (int jc_VKQ_1 = 0; jc_VKQ_1 < KQ_cs; ++jc_VKQ_1) { - KQ_k[jc_VKQ_0+jc_VKQ_1] = __half2half2(tmp[jc_VKQ_1]); + for (int jc_VKQ_1 = 0; jc_VKQ_1 < KQ_cs; ++jc_VKQ_1) { + KQ_k[jc_VKQ_0+jc_VKQ_1] = __half2half2(tmp[jc_VKQ_1]); + } } - } #pragma unroll - for (int i0 = 0; i0 < DVp/2; i0 += warp_size) { + for (int i0 = 0; i0 < DVp/2; i0 += warp_size) { #pragma unroll - for (int jc_VKQ_0 = 0; jc_VKQ_0 < cpw; ++jc_VKQ_0) { - VKQ[jc_VKQ_0*((DVp/2)/warp_size) + i0/warp_size] += V_k[i0/warp_size]*KQ_k[jc_VKQ_0]; + for (int jc_VKQ_0 = 0; jc_VKQ_0 < cpw; ++jc_VKQ_0) { + VKQ[jc_VKQ_0*((DVp/2)/warp_size) + i0/warp_size] += V_k[i0/warp_size]*KQ_k[jc_VKQ_0]; + } } } - } -#else + } else { #pragma unroll - for (int k1 = 0; k1 < nbatch_V; k1 += np) { - __align__(16) float2 V_k[(DVp/2)/warp_size]; - __align__(16) float KQ_k[cpw]; + for (int k1 = 0; k1 < nbatch_V; k1 += np) { + __align__(16) float2 V_k[(DVp/2)/warp_size]; + __align__(16) float KQ_k[cpw]; - constexpr int cpy_ne_D = cpy_ne < DVp/warp_size ? cpy_ne : DVp/warp_size; + constexpr int cpy_ne_D = cpy_ne < DVp/warp_size ? cpy_ne : DVp/warp_size; #pragma unroll - for (int i0 = 0; i0 < DVp; i0 += warp_size*cpy_ne_D) { - ggml_cuda_memcpy_1(&V_k[i0/(2*warp_size)], &KV_tmp[(k1 + threadIdx.y % np)*DV + i0 + threadIdx.x*cpy_ne_D]); - } + for (int i0 = 0; i0 < DVp; i0 += warp_size*cpy_ne_D) { + ggml_cuda_memcpy_1(&V_k[i0/(2*warp_size)], &KV_tmp[(k1 + threadIdx.y % np)*DV + i0 + threadIdx.x*cpy_ne_D]); + } #pragma unroll - for (int jc_VKQ_0 = 0; jc_VKQ_0 < cpw; jc_VKQ_0 += KQ_cs) { - const int jc_KQ = jc_VKQ_0/KQ_cs + (threadIdx.y / np)*(cpw/KQ_cs); + for (int jc_VKQ_0 = 0; jc_VKQ_0 < cpw; jc_VKQ_0 += KQ_cs) { + const int jc_KQ = jc_VKQ_0/KQ_cs + (threadIdx.y / np)*(cpw/KQ_cs); - ggml_cuda_memcpy_1( - &KQ_k[jc_VKQ_0], KQ + jc_KQ*(nbatch_fa*KQ_cs) + (k0 + k1 + threadIdx.y % np)*KQ_cs); - } + ggml_cuda_memcpy_1( + &KQ_k[jc_VKQ_0], KQ + jc_KQ*(nbatch_fa*KQ_cs) + (k0 + k1 + threadIdx.y % np)*KQ_cs); + } #pragma unroll - for (int i0 = 0; i0 < DVp/2; i0 += warp_size) { + for (int i0 = 0; i0 < DVp/2; i0 += warp_size) { #pragma unroll - for (int jc_VKQ_0 = 0; jc_VKQ_0 < cpw; ++jc_VKQ_0) { - VKQ[jc_VKQ_0*((DVp/2)/warp_size) + i0/warp_size].x += V_k[i0/warp_size].x*KQ_k[jc_VKQ_0]; - VKQ[jc_VKQ_0*((DVp/2)/warp_size) + i0/warp_size].y += V_k[i0/warp_size].y*KQ_k[jc_VKQ_0]; + for (int jc_VKQ_0 = 0; jc_VKQ_0 < cpw; ++jc_VKQ_0) { + VKQ[jc_VKQ_0*((DVp/2)/warp_size) + i0/warp_size].x += V_k[i0/warp_size].x*KQ_k[jc_VKQ_0]; + VKQ[jc_VKQ_0*((DVp/2)/warp_size) + i0/warp_size].y += V_k[i0/warp_size].y*KQ_k[jc_VKQ_0]; + } } } } -#endif // FAST_FP16_AVAILABLE __syncthreads(); } } -template // D == head size -__launch_bounds__(ggml_cuda_fattn_tile_get_nthreads(DKQ, DV, ncols1*ncols2), ggml_cuda_fattn_tile_get_occupancy(DKQ, DV, ncols1*ncols2)) +template // D == head size +__launch_bounds__(ggml_cuda_fattn_tile_get_nthreads(DKQ, DV, ncols1*ncols2), prec_f32 ? 1 : ggml_cuda_fattn_tile_get_occupancy(DKQ, DV, ncols1*ncols2)) static __global__ void flash_attn_tile( const char * Q_ptr, const char * K_ptr, @@ -846,7 +848,16 @@ static __global__ void flash_attn_tile( constexpr int warp_size = 32; constexpr int nwarps = ggml_cuda_fattn_tile_get_nthreads (DKQ, DV, ncols1*ncols2) / warp_size; constexpr int nbatch_fa = ggml_cuda_fattn_tile_get_nbatch_fa(DKQ, DV, ncols1*ncols2); - constexpr int nbatch_K = ggml_cuda_fattn_tile_get_nbatch_K (DKQ, DV, ncols1*ncols2); +#ifdef FAST_FP16_AVAILABLE + constexpr bool fast_fp16_avail = true; +#else + constexpr bool fast_fp16_avail = false; +#endif // FAST_FP16_AVAILABLE + constexpr int nbatch_K_f16cfg = ggml_cuda_fattn_tile_get_nbatch_K(DKQ, DV, ncols1*ncols2); + constexpr int nbatch_K_f32fit = ggml_cuda_fattn_tile_get_nbatch_K_f32(DKQ, DV, ncols1*ncols2, nbatch_fa, nbatch_K_f16cfg); + constexpr bool force_f32 = fast_fp16_avail && prec_f32 && (nbatch_K_f32fit > 0); + constexpr bool use_fp16 = fast_fp16_avail && !force_f32; + constexpr int nbatch_K = force_f32 ? nbatch_K_f32fit : nbatch_K_f16cfg; // In this kernel Q, K, V are matrices while i, j, k are matrix indices. @@ -883,17 +894,13 @@ static __global__ void flash_attn_tile( // KV_tmp is padded to avoid memory conflicts for K (cpy_ne) and OOB accesses for V (DVp-DV). // KQ == SRAM buffer to hold KQ fragments between KQ and VKQ matrix multiplications. // VKQ == Accumulators in registers for the final VKQ result. -#ifdef FAST_FP16_AVAILABLE - __shared__ half2 Q_tmp[ncols * DKQ/2]; - __shared__ half2 KV_tmp[nbatch_fa * (nbatch_K/2 + cpy_ne) + DVp-DV]; - __shared__ half KQ[ncols * nbatch_fa]; - __align__(16) half2 VKQ[cpw * ((DVp/2)/warp_size)] = {{0.0f, 0.0f}}; -#else - __shared__ float Q_tmp[ncols * DKQ]; - __shared__ float KV_tmp[nbatch_fa * (nbatch_K + cpy_ne) + DVp-DV]; - __shared__ float KQ[ncols * nbatch_fa]; - __align__(16) float2 VKQ[cpw * ((DVp/2)/warp_size)] = {{0.0f, 0.0f}}; -#endif // FAST_FP16_AVAILABLE + using TQ = std::conditional_t; // Q_tmp / KV_tmp element + using TKQ = std::conditional_t; // KQ element + using TVKQ = std::conditional_t; // VKQ accumulator element + __shared__ TQ Q_tmp[ncols * (use_fp16 ? DKQ/2 : DKQ)]; + __shared__ TQ KV_tmp[nbatch_fa * ((use_fp16 ? nbatch_K/2 : nbatch_K) + cpy_ne) + DVp-DV]; + __shared__ TKQ KQ[ncols * nbatch_fa]; + __align__(16) TVKQ VKQ[cpw * ((DVp/2)/warp_size)] = {{0.0f, 0.0f}}; float KQ_max[cpw]; #pragma unroll @@ -927,25 +934,25 @@ static __global__ void flash_attn_tile( tmp_f[i1] *= scale; } -#ifdef FAST_FP16_AVAILABLE - __align__(16) half2 tmp_h2[cpy_ne_D/2]; + if constexpr (use_fp16) { + __align__(16) half2 tmp_h2[cpy_ne_D/2]; #pragma unroll - for (int i1 = 0; i1 < cpy_ne_D; i1 += 2) { - tmp_h2[i1/2] = make_half2(tmp_f[i1 + 0], tmp_f[i1 + 1]); + for (int i1 = 0; i1 < cpy_ne_D; i1 += 2) { + tmp_h2[i1/2] = make_half2(tmp_f[i1 + 0], tmp_f[i1 + 1]); #if defined(FAST_FP16_AVAILABLE) && !defined(V_DOT2_F32_F16_AVAILABLE) - // Without the v_dot2_f32_f16 instruction there is a higher risk of numerical overflow in the KQ calculation. - // Therefore, scale down Q values and apply the inverse scale the FP32 KQ values afterwards again. - tmp_h2[i1/2] *= make_half2(0.25f, 0.25f); + // Without the v_dot2_f32_f16 instruction there is a higher risk of numerical overflow in the KQ calculation. + // Therefore, scale down Q values and apply the inverse scale the FP32 KQ values afterwards again. + tmp_h2[i1/2] *= make_half2(0.25f, 0.25f); #endif // defined(FAST_FP16_AVAILABLE) && !defined(V_DOT2_F32_F16_AVAILABLE) + } + ggml_cuda_memcpy_1( + &Q_tmp[jc*(DKQ/2) + i0/2 + (threadIdx.y % np)*(warp_size*cpy_ne_D/2) + threadIdx.x*(cpy_ne_D/2)], + tmp_h2); + } else { + ggml_cuda_memcpy_1( + &Q_tmp[jc* DKQ + i0 + (threadIdx.y % np)*(warp_size*cpy_ne_D) + threadIdx.x* cpy_ne_D], + tmp_f); } - ggml_cuda_memcpy_1( - &Q_tmp[jc*(DKQ/2) + i0/2 + (threadIdx.y % np)*(warp_size*cpy_ne_D/2) + threadIdx.x*(cpy_ne_D/2)], - tmp_h2); -#else - ggml_cuda_memcpy_1( - &Q_tmp[jc* DKQ + i0 + (threadIdx.y % np)*(warp_size*cpy_ne_D) + threadIdx.x* cpy_ne_D], - tmp_f); -#endif // FAST_FP16_AVAILABLE } } } @@ -989,28 +996,24 @@ static __global__ void flash_attn_tile( static_assert(cpw == 1, "bad cpw"); static_assert(nbatch_fa*nbatch_K >= nwarps*DVp, "KV_tmp too small"); -#ifdef FAST_FP16_AVAILABLE - half2 * VKQ_combine = (half2 *) KV_tmp; -#else - float * VKQ_combine = (float *) KV_tmp; -#endif // FAST_FP16_AVAILABLE + TQ * VKQ_combine = KV_tmp; float * KQ_sum_combine = (float *) Q_tmp; if (threadIdx.y % np != 0) { -#ifdef FAST_FP16_AVAILABLE - constexpr int cpy_ne_D = cpy_ne < (DVp/2)/warp_size ? cpy_ne : (DVp/2)/warp_size; + if constexpr (use_fp16) { + constexpr int cpy_ne_D = cpy_ne < (DVp/2)/warp_size ? cpy_ne : (DVp/2)/warp_size; #pragma unroll - for (int i0 = 0; i0 < DVp/2; i0 += warp_size*cpy_ne_D) { - ggml_cuda_memcpy_1(&VKQ_combine[threadIdx.y*(DVp/2) + i0 + threadIdx.x*cpy_ne_D], &VKQ[i0/warp_size]); - } -#else - constexpr int cpy_ne_D = cpy_ne < DVp/warp_size ? cpy_ne : DVp/warp_size; + for (int i0 = 0; i0 < DVp/2; i0 += warp_size*cpy_ne_D) { + ggml_cuda_memcpy_1(&VKQ_combine[threadIdx.y*(DVp/2) + i0 + threadIdx.x*cpy_ne_D], &VKQ[i0/warp_size]); + } + } else { + constexpr int cpy_ne_D = cpy_ne < DVp/warp_size ? cpy_ne : DVp/warp_size; #pragma unroll - for (int i0 = 0; i0 < DVp; i0 += warp_size*cpy_ne_D) { - ggml_cuda_memcpy_1( - &VKQ_combine[threadIdx.y*DVp + i0 + threadIdx.x*cpy_ne_D], ((const float *) VKQ) + i0/warp_size); + for (int i0 = 0; i0 < DVp; i0 += warp_size*cpy_ne_D) { + ggml_cuda_memcpy_1( + &VKQ_combine[threadIdx.y*DVp + i0 + threadIdx.x*cpy_ne_D], ((const float *) VKQ) + i0/warp_size); + } } -#endif // FAST_FP16_AVAILABLE if (threadIdx.x == 0) { KQ_sum_combine[threadIdx.y] = KQ_sum[0]; @@ -1023,29 +1026,29 @@ static __global__ void flash_attn_tile( #pragma unroll for (int ip = 1; ip < np; ++ip) { -#ifdef FAST_FP16_AVAILABLE - constexpr int cpy_ne_D = cpy_ne < (DVp/2)/warp_size ? cpy_ne : (DVp/2)/warp_size; + if constexpr (use_fp16) { + constexpr int cpy_ne_D = cpy_ne < (DVp/2)/warp_size ? cpy_ne : (DVp/2)/warp_size; #pragma unroll - for (int i0 = 0; i0 < DVp/2; i0 += warp_size*cpy_ne_D) { - __align__(16) half2 tmp[cpy_ne_D]; - ggml_cuda_memcpy_1(tmp, &VKQ_combine[(threadIdx.y + ip)*(DVp/2) + i0 + threadIdx.x*cpy_ne_D]); + for (int i0 = 0; i0 < DVp/2; i0 += warp_size*cpy_ne_D) { + __align__(16) half2 tmp[cpy_ne_D]; + ggml_cuda_memcpy_1(tmp, &VKQ_combine[(threadIdx.y + ip)*(DVp/2) + i0 + threadIdx.x*cpy_ne_D]); #pragma unroll - for (int i1 = 0; i1 < cpy_ne_D; ++i1) { - VKQ[i0/warp_size + i1] += tmp[i1]; + for (int i1 = 0; i1 < cpy_ne_D; ++i1) { + VKQ[i0/warp_size + i1] += tmp[i1]; + } } - } -#else - constexpr int cpy_ne_D = cpy_ne < DVp/warp_size ? cpy_ne : DVp/warp_size; + } else { + constexpr int cpy_ne_D = cpy_ne < DVp/warp_size ? cpy_ne : DVp/warp_size; #pragma unroll - for (int i0 = 0; i0 < DVp; i0 += warp_size*cpy_ne_D) { - __align__(16) float tmp[cpy_ne_D]; - ggml_cuda_memcpy_1(tmp, &VKQ_combine[(threadIdx.y + ip)*DVp + i0 + threadIdx.x*cpy_ne_D]); + for (int i0 = 0; i0 < DVp; i0 += warp_size*cpy_ne_D) { + __align__(16) float tmp[cpy_ne_D]; + ggml_cuda_memcpy_1(tmp, &VKQ_combine[(threadIdx.y + ip)*DVp + i0 + threadIdx.x*cpy_ne_D]); #pragma unroll - for (int i1 = 0; i1 < cpy_ne_D; ++i1) { - ((float *)VKQ)[i0/warp_size + i1] += tmp[i1]; + for (int i1 = 0; i1 < cpy_ne_D; ++i1) { + ((float *)VKQ)[i0/warp_size + i1] += tmp[i1]; + } } } -#endif // FAST_FP16_AVAILABLE KQ_sum[0] += KQ_sum_combine[threadIdx.y + ip]; } @@ -1065,19 +1068,19 @@ static __global__ void flash_attn_tile( const float val = expf(sink - KQ_max[jc0]); KQ_sum[jc0] = KQ_sum[jc0]*KQ_max_scale + val; -#ifdef FAST_FP16_AVAILABLE - const half2 KQ_max_scale_h2 = make_half2(KQ_max_scale, KQ_max_scale); + if constexpr (use_fp16) { + const half2 KQ_max_scale_h2 = make_half2(KQ_max_scale, KQ_max_scale); #pragma unroll - for (int i0 = 0; i0 < DVp/2; i0 += warp_size) { - VKQ[jc0*((DVp/2)/warp_size) + i0/warp_size] *= KQ_max_scale_h2; - } -#else + for (int i0 = 0; i0 < DVp/2; i0 += warp_size) { + VKQ[jc0*((DVp/2)/warp_size) + i0/warp_size] *= KQ_max_scale_h2; + } + } else { #pragma unroll - for (int i0 = 0; i0 < DVp/2; i0 += warp_size) { - VKQ[jc0*((DVp/2)/warp_size) + i0/warp_size].x *= KQ_max_scale; - VKQ[jc0*((DVp/2)/warp_size) + i0/warp_size].y *= KQ_max_scale; + for (int i0 = 0; i0 < DVp/2; i0 += warp_size) { + VKQ[jc0*((DVp/2)/warp_size) + i0/warp_size].x *= KQ_max_scale; + VKQ[jc0*((DVp/2)/warp_size) + i0/warp_size].y *= KQ_max_scale; + } } -#endif // FAST_FP16_AVAILABLE } } @@ -1097,37 +1100,37 @@ static __global__ void flash_attn_tile( const int j_dst_unrolled = ((sequence*int(ne01.z) + col_Q_0 + j)*ne02 + head0 + c)*gridDim.y + blockIdx.y; -#ifdef FAST_FP16_AVAILABLE - constexpr int cpy_ne_D = cpy_ne/2 < (DVp/2)/warp_size ? cpy_ne/2 : (DVp/2)/warp_size; + if constexpr (use_fp16) { + constexpr int cpy_ne_D = cpy_ne/2 < (DVp/2)/warp_size ? cpy_ne/2 : (DVp/2)/warp_size; #pragma unroll - for (int i0 = 0; i0 < DVp/2; i0 += warp_size*cpy_ne_D) { - __align__(16) float2 tmp[cpy_ne_D]; + for (int i0 = 0; i0 < DVp/2; i0 += warp_size*cpy_ne_D) { + __align__(16) float2 tmp[cpy_ne_D]; #pragma unroll - for (int i1 = 0; i1 < cpy_ne_D; ++i1) { - tmp[i1] = __half22float2(VKQ[jc0*((DVp/2)/warp_size) + i0/warp_size + i1]); - tmp[i1].x *= scale; - tmp[i1].y *= scale; - } - if (i0 + warp_size*cpy_ne_D <= DV/2 || i0 + threadIdx.x*cpy_ne_D < DV/2) { - ggml_cuda_memcpy_1(&dst[j_dst_unrolled*DV + 2*i0 + threadIdx.x*(2*cpy_ne_D)], tmp); + for (int i1 = 0; i1 < cpy_ne_D; ++i1) { + tmp[i1] = __half22float2(VKQ[jc0*((DVp/2)/warp_size) + i0/warp_size + i1]); + tmp[i1].x *= scale; + tmp[i1].y *= scale; + } + if (i0 + warp_size*cpy_ne_D <= DV/2 || i0 + threadIdx.x*cpy_ne_D < DV/2) { + ggml_cuda_memcpy_1(&dst[j_dst_unrolled*DV + 2*i0 + threadIdx.x*(2*cpy_ne_D)], tmp); + } } - } -#else - constexpr int cpy_ne_D = cpy_ne < DVp/warp_size ? cpy_ne : DVp/warp_size; + } else { + constexpr int cpy_ne_D = cpy_ne < DVp/warp_size ? cpy_ne : DVp/warp_size; #pragma unroll - for (int i0 = 0; i0 < DVp; i0 += warp_size*cpy_ne_D) { - if (i0 + warp_size*cpy_ne_D <= DV || i0 + threadIdx.x*cpy_ne_D < DV) { + for (int i0 = 0; i0 < DVp; i0 += warp_size*cpy_ne_D) { + if (i0 + warp_size*cpy_ne_D <= DV || i0 + threadIdx.x*cpy_ne_D < DV) { #pragma unroll - for (int i1 = 0; i1 < cpy_ne_D/2; ++i1) { - VKQ[jc0*((DVp/2)/warp_size) + i0/(2*warp_size) + i1].x *= scale; - VKQ[jc0*((DVp/2)/warp_size) + i0/(2*warp_size) + i1].y *= scale; + for (int i1 = 0; i1 < cpy_ne_D/2; ++i1) { + VKQ[jc0*((DVp/2)/warp_size) + i0/(2*warp_size) + i1].x *= scale; + VKQ[jc0*((DVp/2)/warp_size) + i0/(2*warp_size) + i1].y *= scale; + } + ggml_cuda_memcpy_1( + &dst[j_dst_unrolled*DV + i0 + threadIdx.x*cpy_ne_D], + &VKQ[jc0*((DVp/2)/warp_size) + i0/(2*warp_size)]); } - ggml_cuda_memcpy_1( - &dst[j_dst_unrolled*DV + i0 + threadIdx.x*cpy_ne_D], - &VKQ[jc0*((DVp/2)/warp_size) + i0/(2*warp_size)]); } } -#endif // FAST_FP16_AVAILABLE if (gridDim.y != 1 && threadIdx.x == 0) { dst_meta[j_dst_unrolled] = make_float2(KQ_max[jc0], KQ_sum[jc0]); @@ -1151,6 +1154,13 @@ template static void launch_fattn_tile_switch_ncols1(ggml_backend_cuda_context & ctx, ggml_tensor * dst) { const ggml_tensor * Q = dst->src[0]; +#ifdef GGML_USE_HIP + // AMD non-GQA fallback: honor GGML_PREC_F32 here, or fp16 corrupts the result. + const bool prec_f32 = ggml_flash_attn_ext_get_prec(dst) == GGML_PREC_F32; +#else + const bool prec_f32 = false; +#endif // GGML_USE_HIP + const int id = ggml_cuda_get_device(); const int cc = ggml_cuda_info().devices[id].cc; const int warp_size = 32; @@ -1163,7 +1173,9 @@ static void launch_fattn_tile_switch_ncols1(ggml_backend_cuda_context & ctx, ggm constexpr int cols_per_block = 64; const int nwarps = ggml_cuda_fattn_tile_get_nthreads (DKQ, DV, cols_per_block, cc) / warp_size; const int nbatch_fa = ggml_cuda_fattn_tile_get_nbatch_fa(DKQ, DV, cols_per_block, cc); - fattn_kernel_t fattn_kernel = flash_attn_tile; + fattn_kernel_t fattn_kernel = prec_f32 + ? flash_attn_tile + : flash_attn_tile; launch_fattn (ctx, dst, fattn_kernel, nwarps, nbytes_shared, nbatch_fa, true, true, false, warp_size); return; @@ -1179,7 +1191,9 @@ static void launch_fattn_tile_switch_ncols1(ggml_backend_cuda_context & ctx, ggm constexpr int cols_per_block = 32; const int nwarps = ggml_cuda_fattn_tile_get_nthreads (DKQ, DV, cols_per_block, cc) / warp_size; const int nbatch_fa = ggml_cuda_fattn_tile_get_nbatch_fa(DKQ, DV, cols_per_block, cc); - fattn_kernel_t fattn_kernel = flash_attn_tile; + fattn_kernel_t fattn_kernel = prec_f32 + ? flash_attn_tile + : flash_attn_tile; launch_fattn (ctx, dst, fattn_kernel, nwarps, nbytes_shared, nbatch_fa, true, true, false, warp_size); return; @@ -1191,7 +1205,9 @@ static void launch_fattn_tile_switch_ncols1(ggml_backend_cuda_context & ctx, ggm constexpr int cols_per_block = 16; const int nwarps = ggml_cuda_fattn_tile_get_nthreads (DKQ, DV, cols_per_block, cc) / warp_size; const int nbatch_fa = ggml_cuda_fattn_tile_get_nbatch_fa(DKQ, DV, cols_per_block, cc); - fattn_kernel_t fattn_kernel = flash_attn_tile; + fattn_kernel_t fattn_kernel = prec_f32 + ? flash_attn_tile + : flash_attn_tile; launch_fattn (ctx, dst, fattn_kernel, nwarps, nbytes_shared, nbatch_fa, true, true, false, warp_size); return; @@ -1203,7 +1219,9 @@ static void launch_fattn_tile_switch_ncols1(ggml_backend_cuda_context & ctx, ggm constexpr int cols_per_block = 8; const int nwarps = ggml_cuda_fattn_tile_get_nthreads (DKQ, DV, cols_per_block, cc) / warp_size; const int nbatch_fa = ggml_cuda_fattn_tile_get_nbatch_fa(DKQ, DV, cols_per_block, cc); - fattn_kernel_t fattn_kernel = flash_attn_tile; + fattn_kernel_t fattn_kernel = prec_f32 + ? flash_attn_tile + : flash_attn_tile; launch_fattn (ctx, dst, fattn_kernel, nwarps, nbytes_shared, nbatch_fa, true, true, false, warp_size); return; @@ -1215,7 +1233,9 @@ static void launch_fattn_tile_switch_ncols1(ggml_backend_cuda_context & ctx, ggm constexpr int cols_per_block = 4; const int nwarps = ggml_cuda_fattn_tile_get_nthreads (DKQ, DV, cols_per_block, cc) / warp_size; const int nbatch_fa = ggml_cuda_fattn_tile_get_nbatch_fa(DKQ, DV, cols_per_block, cc); - fattn_kernel_t fattn_kernel = flash_attn_tile; + fattn_kernel_t fattn_kernel = prec_f32 + ? flash_attn_tile + : flash_attn_tile; launch_fattn (ctx, dst, fattn_kernel, nwarps, nbytes_shared, nbatch_fa, true, true, false, warp_size); return;