| 5582 | // ggml_flash_attn |
| 5583 | |
| 5584 | struct ggml_tensor * ggml_flash_attn( |
| 5585 | struct ggml_context * ctx, |
| 5586 | struct ggml_tensor * q, |
| 5587 | struct ggml_tensor * k, |
| 5588 | struct ggml_tensor * v, |
| 5589 | bool masked) { |
| 5590 | GGML_ASSERT(ggml_can_mul_mat(k, q)); |
| 5591 | // TODO: check if vT can be multiplied by (k*qT) |
| 5592 | |
| 5593 | bool is_node = false; |
| 5594 | |
| 5595 | if (q->grad || k->grad || v->grad) { |
| 5596 | is_node = true; |
| 5597 | } |
| 5598 | |
| 5599 | //struct ggml_tensor * result = ggml_dup_tensor(ctx, q); |
| 5600 | struct ggml_tensor * result = ggml_new_tensor(ctx, GGML_TYPE_F32, q->n_dims, q->ne); |
| 5601 | |
| 5602 | int32_t t = masked ? 1 : 0; |
| 5603 | ggml_set_op_params(result, &t, sizeof(t)); |
| 5604 | |
| 5605 | result->op = GGML_OP_FLASH_ATTN; |
| 5606 | result->grad = is_node ? ggml_dup_tensor(ctx, result) : NULL; |
| 5607 | result->src[0] = q; |
| 5608 | result->src[1] = k; |
| 5609 | result->src[2] = v; |
| 5610 | |
| 5611 | return result; |
| 5612 | } |
| 5613 | |
| 5614 | // ggml_flash_ff |
| 5615 | |