Jelajahi Sumber

CUDA: fix broken oob check for FA vec f32 kernel (#7904)

Johannes Gäßler 1 tahun lalu
induk
melakukan
963552903f
1 mengubah file dengan 1 tambahan dan 1 penghapusan
  1. 1 1
      ggml-cuda/fattn-vec-f32.cuh

+ 1 - 1
ggml-cuda/fattn-vec-f32.cuh

@@ -149,7 +149,7 @@ static __global__ void flash_attn_vec_ext_f32(
             for (int i0 = 0; i0 < D/2; i0 += WARP_SIZE) {
                 const int i = i0 + threadIdx.x;
 
-                Q_f2[j][i0/WARP_SIZE]    = ncols <= 2 || ic0 + j ? Q_f2_j[i] : make_float2(0.0f, 0.0f);
+                Q_f2[j][i0/WARP_SIZE]    = ncols <= 2 || ic0 + j < ne01 ? Q_f2_j[i] : make_float2(0.0f, 0.0f);
                 Q_f2[j][i0/WARP_SIZE].x *= scale;
                 Q_f2[j][i0/WARP_SIZE].y *= scale;
             }