| 5252 | // ggml_flash_attn_ext |
| 5253 | |
| 5254 | struct ggml_tensor * ggml_flash_attn_ext( |
| 5255 | struct ggml_context * ctx, |
| 5256 | struct ggml_tensor * q, |
| 5257 | struct ggml_tensor * k, |
| 5258 | struct ggml_tensor * v, |
| 5259 | struct ggml_tensor * mask, |
| 5260 | float scale, |
| 5261 | float max_bias, |
| 5262 | float logit_softcap) { |
| 5263 | GGML_ASSERT(ggml_can_mul_mat(k, q)); |
| 5264 | // TODO: check if vT can be multiplied by (k*qT) |
| 5265 | |
| 5266 | GGML_ASSERT(q->ne[3] == k->ne[3]); |
| 5267 | GGML_ASSERT(q->ne[3] == v->ne[3]); |
| 5268 | |
| 5269 | if (mask) { |
| 5270 | GGML_ASSERT(ggml_is_contiguous(mask)); |
| 5271 | //GGML_ASSERT(ggml_can_repeat_rows(mask, qk)); |
| 5272 | |
| 5273 | GGML_ASSERT(q->ne[2] % mask->ne[2] == 0); |
| 5274 | GGML_ASSERT(q->ne[3] % mask->ne[3] == 0); |
| 5275 | } |
| 5276 | |
| 5277 | if (max_bias > 0.0f) { |
| 5278 | GGML_ASSERT(mask); |
| 5279 | } |
| 5280 | |
| 5281 | // permute(0, 2, 1, 3) |
| 5282 | int64_t ne[4] = { v->ne[0], q->ne[2], q->ne[1], q->ne[3] }; |
| 5283 | struct ggml_tensor * result = ggml_new_tensor(ctx, GGML_TYPE_F32, 4, ne); |
| 5284 | |
| 5285 | float params[] = { scale, max_bias, logit_softcap }; |
| 5286 | ggml_set_op_params(result, params, sizeof(params)); |
| 5287 | |
| 5288 | result->op = GGML_OP_FLASH_ATTN_EXT; |
| 5289 | result->src[0] = q; |
| 5290 | result->src[1] = k; |
| 5291 | result->src[2] = v; |
| 5292 | result->src[3] = mask; |
| 5293 | |
| 5294 | return result; |
| 5295 | } |
| 5296 | |
| 5297 | void ggml_flash_attn_ext_set_prec( |
| 5298 | struct ggml_tensor * a, |