1#ifndef INFINI_OPS_BASE_FLASH_ATTENTION_H_
2#define INFINI_OPS_BASE_FLASH_ATTENTION_H_
19 "Use `ScaledDotProductAttention` for standard attention "
23 std::optional<Tensor> cu_seqlens_q,
24 std::optional<Tensor> cu_seqlens_kv,
25 std::optional<Tensor> block_table, int64_t num_heads,
26 int64_t num_kv_heads, int64_t head_size,
double scale,
27 bool causal, int64_t window_left, int64_t window_right,
28 int64_t block_size,
Tensor output)
29 : num_tokens_{query.size(0)},
30 num_heads_{num_heads},
31 num_kv_heads_{num_kv_heads},
32 head_size_{head_size},
35 window_left_{window_left},
36 window_right_{window_right},
37 block_size_{block_size},
38 dtype_{query.dtype()},
39 query_shape_{query.shape()},
40 key_shape_{key.shape()},
41 value_shape_{value.shape()},
42 output_shape_{output.shape()},
43 query_strides_{query.strides()},
44 key_strides_{key.strides()},
45 value_strides_{value.strides()},
46 output_strides_{output.strides()},
47 has_cu_seqlens_q_{cu_seqlens_q.has_value()},
48 has_cu_seqlens_kv_{cu_seqlens_kv.has_value()},
49 has_block_table_{block_table.has_value()} {
50 assert(num_heads % num_kv_heads == 0 &&
51 "`FlashAttention` requires num_heads divisible by num_kv_heads");
52 assert(query.ndim() == 3 &&
53 "`FlashAttention` requires query to be 3D [T, N, D]");
58 std::optional<Tensor> cu_seqlens_q,
59 std::optional<Tensor> cu_seqlens_kv,
60 std::optional<Tensor> block_table, int64_t num_heads,
61 int64_t num_kv_heads, int64_t head_size,
double scale,
62 bool causal, int64_t window_left,
63 int64_t window_right, int64_t block_size,
67 Tensor::Size num_tokens_{0};
69 int64_t num_heads_{0};
71 int64_t num_kv_heads_{0};
73 int64_t head_size_{0};
79 int64_t window_left_{-1};
81 int64_t window_right_{-1};
83 int64_t block_size_{0};
103 bool has_cu_seqlens_q_{
false};
105 bool has_cu_seqlens_kv_{
false};
107 bool has_block_table_{
false};
Definition flash_attention.h:20
Tensor::Shape value_shape_
Definition flash_attention.h:91
const DataType dtype_
Definition flash_attention.h:85
Tensor::Shape output_shape_
Definition flash_attention.h:93
Tensor::Strides output_strides_
Definition flash_attention.h:101
FlashAttention(const Tensor query, const Tensor key, const Tensor value, std::optional< Tensor > cu_seqlens_q, std::optional< Tensor > cu_seqlens_kv, std::optional< Tensor > block_table, int64_t num_heads, int64_t num_kv_heads, int64_t head_size, double scale, bool causal, int64_t window_left, int64_t window_right, int64_t block_size, Tensor output)
Definition flash_attention.h:22
Tensor::Strides value_strides_
Definition flash_attention.h:99
virtual void operator()(const Tensor query, const Tensor key, const Tensor value, std::optional< Tensor > cu_seqlens_q, std::optional< Tensor > cu_seqlens_kv, std::optional< Tensor > block_table, int64_t num_heads, int64_t num_kv_heads, int64_t head_size, double scale, bool causal, int64_t window_left, int64_t window_right, int64_t block_size, Tensor output) const =0
Tensor::Shape query_shape_
Definition flash_attention.h:87
Tensor::Strides query_strides_
Definition flash_attention.h:95
Tensor::Shape key_shape_
Definition flash_attention.h:89
Tensor::Strides key_strides_
Definition flash_attention.h:97
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8