InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
flash_attention.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_FLASH_ATTENTION_H_
2#define INFINI_OPS_BASE_FLASH_ATTENTION_H_
3
4#include <cstddef>
5#include <optional>
6#include <vector>
7
8#include "operator.h"
9
10namespace infini::ops {
11
12// Legacy fused multi-head / grouped-query attention interface.
13// Layout: `query` / `key` / `value` are `[T, N, D]` (TND).
14// Prefill uses `cu_seqlens_q` / `cu_seqlens_kv` for variable-length packing.
15// Decode uses `block_table` for paged KV cache lookup.
18class [[deprecated(
19 "Use `ScaledDotProductAttention` for standard attention "
20 "semantics.")]] FlashAttention : public Operator<FlashAttention> {
21 public:
22 FlashAttention(const Tensor query, const Tensor key, const Tensor value,
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},
33 scale_{scale},
34 causal_{causal},
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]");
54 }
55
56 virtual void operator()(const Tensor query, const Tensor key,
57 const Tensor value,
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,
64 Tensor output) const = 0;
65
66 protected:
67 Tensor::Size num_tokens_{0};
68
69 int64_t num_heads_{0};
70
71 int64_t num_kv_heads_{0};
72
73 int64_t head_size_{0};
74
75 double scale_{0.0};
76
77 bool causal_{false};
78
79 int64_t window_left_{-1};
80
81 int64_t window_right_{-1};
82
83 int64_t block_size_{0};
84
85 const DataType dtype_;
86
87 Tensor::Shape query_shape_;
88
89 Tensor::Shape key_shape_;
90
91 Tensor::Shape value_shape_;
92
93 Tensor::Shape output_shape_;
94
95 Tensor::Strides query_strides_;
96
97 Tensor::Strides key_strides_;
98
99 Tensor::Strides value_strides_;
100
101 Tensor::Strides output_strides_;
102
103 bool has_cu_seqlens_q_{false};
104
105 bool has_cu_seqlens_kv_{false};
106
107 bool has_block_table_{false};
108};
109
110} // namespace infini::ops
111
112#endif
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