]> git.djapps.eu Git - pkg/ggml/sources/llama.cpp/commitdiff
cuda : prevent integer truncation and overflow errors when using KQ mask strides...
authorfairydreaming <redacted>
Tue, 30 Jun 2026 18:47:05 +0000 (20:47 +0200)
committerGitHub <redacted>
Tue, 30 Jun 2026 18:47:05 +0000 (20:47 +0200)
Co-authored-by: Stanisław Szymczyk <redacted>
ggml/src/ggml-cuda/fattn-common.cuh

index 8dfa51ad1e8f7854613775a6e74e0893a4115f05..b76121fef3a9ab30c264cd3fb96becd449d60c0a 100644 (file)
@@ -664,7 +664,7 @@ constexpr __device__ dequantize_V_t get_dequantize_V() {
 template <int ncols1>
 __launch_bounds__(FATTN_KQ_STRIDE/2, 1)
 static __global__ void flash_attn_mask_to_KV_max(
-        const half2 * __restrict__ mask, int * __restrict__ KV_max, const int ne30, const int s31, const int s33) {
+        const half2 * __restrict__ mask, int * __restrict__ KV_max, const int ne30, const int64_t s31, const int64_t s33) {
     const int ne31     = gridDim.x;
     const int tid      = threadIdx.x;
     const int sequence = blockIdx.y;
@@ -1089,8 +1089,8 @@ void launch_fattn(
     // Only worth the overhead if there is at lease one FATTN_KQ_STRIDE x FATTN_KQ_STRIDE square to be skipped or
     //     multiple sequences of possibly different lengths.
     if (mask && K->ne[1] % FATTN_KQ_STRIDE == 0 && (Q->ne[1] >= 1024 || Q->ne[3] > 1)) {
-        const int s31 = mask->nb[1] / sizeof(half2);
-        const int s33 = mask->nb[3] / sizeof(half2);
+        const int64_t s31 = mask->nb[1] / sizeof(half2);
+        const int64_t s33 = mask->nb[3] / sizeof(half2);
 
         const dim3 blocks_num_KV_max(ntiles_x, Q->ne[3], 1);
         const dim3 block_dim_KV_max(FATTN_KQ_STRIDE/2, 1, 1);