InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
flash_attn_varlen_func.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_FLASH_ATTN_VARLEN_FUNC_H_
2#define INFINI_OPS_BASE_FLASH_ATTN_VARLEN_FUNC_H_
3
4#include <cassert>
5#include <cstdint>
6#include <optional>
7#include <vector>
8
9#include "operator.h"
10
11namespace infini::ops {
12
13// Packed variable-length attention aligned with Dao-AILab FlashAttention's
14// `flash_attn_varlen_func` public interface.
15class FlashAttnVarlenFunc : public Operator<FlashAttnVarlenFunc> {
16 public:
17 FlashAttnVarlenFunc(const Tensor q, const Tensor k, const Tensor v,
18 const Tensor cu_seqlens_q, const Tensor cu_seqlens_k,
19 const int64_t max_seqlen_q, const int64_t max_seqlen_k,
20 Tensor out)
22 k,
23 v,
24 cu_seqlens_q,
25 cu_seqlens_k,
26 std::nullopt,
27 std::nullopt,
28 max_seqlen_q,
29 max_seqlen_k,
30 0.0,
31 std::nullopt,
32 false,
33 {-1, -1},
34 0.0,
35 false,
36 false,
37 out,
38 std::nullopt,
39 std::nullopt} {}
40
42 const Tensor q, const Tensor k, const Tensor v, const Tensor cu_seqlens_q,
43 const Tensor cu_seqlens_k, const std::optional<Tensor> alibi_slopes,
44 const std::optional<Tensor> block_table, const int64_t max_seqlen_q,
45 const int64_t max_seqlen_k, const double dropout_p,
46 const std::optional<double> softmax_scale, const bool causal,
47 const std::vector<int64_t> window_size, const double softcap,
48 const bool deterministic, const bool return_attn_probs, Tensor out,
49 std::optional<Tensor> softmax_lse, std::optional<Tensor> s_dmask)
50 : q_shape_{q.shape()},
51 k_shape_{k.shape()},
52 v_shape_{v.shape()},
53 cu_seqlens_q_shape_{cu_seqlens_q.shape()},
54 cu_seqlens_k_shape_{cu_seqlens_k.shape()},
55 alibi_slopes_shape_{alibi_slopes.has_value()
56 ? Tensor::Shape{alibi_slopes->shape()}
57 : Tensor::Shape{}},
58 block_table_shape_{block_table.has_value()
59 ? Tensor::Shape{block_table->shape()}
60 : Tensor::Shape{}},
61 out_shape_{out.shape()},
62 softmax_lse_shape_{softmax_lse.has_value()
63 ? Tensor::Shape{softmax_lse->shape()}
64 : Tensor::Shape{}},
65 s_dmask_shape_{s_dmask.has_value() ? Tensor::Shape{s_dmask->shape()}
66 : Tensor::Shape{}},
67 q_strides_{q.strides()},
68 k_strides_{k.strides()},
69 v_strides_{v.strides()},
70 cu_seqlens_q_strides_{cu_seqlens_q.strides()},
71 cu_seqlens_k_strides_{cu_seqlens_k.strides()},
72 alibi_slopes_strides_{alibi_slopes.has_value()
73 ? Tensor::Strides{alibi_slopes->strides()}
74 : Tensor::Strides{}},
75 block_table_strides_{block_table.has_value()
76 ? Tensor::Strides{block_table->strides()}
77 : Tensor::Strides{}},
78 out_strides_{out.strides()},
79 softmax_lse_strides_{softmax_lse.has_value()
80 ? Tensor::Strides{softmax_lse->strides()}
81 : Tensor::Strides{}},
82 s_dmask_strides_{s_dmask.has_value()
83 ? Tensor::Strides{s_dmask->strides()}
84 : Tensor::Strides{}},
85 q_dtype_{q.dtype()},
86 k_dtype_{k.dtype()},
87 v_dtype_{v.dtype()},
88 cu_seqlens_q_dtype_{cu_seqlens_q.dtype()},
89 cu_seqlens_k_dtype_{cu_seqlens_k.dtype()},
90 alibi_slopes_dtype_{alibi_slopes.has_value() ? alibi_slopes->dtype()
91 : DataType::kFloat32},
92 block_table_dtype_{block_table.has_value() ? block_table->dtype()
93 : DataType::kInt32},
94 out_dtype_{out.dtype()},
95 softmax_lse_dtype_{softmax_lse.has_value() ? softmax_lse->dtype()
96 : DataType::kFloat32},
97 s_dmask_dtype_{s_dmask.has_value() ? s_dmask->dtype() : q.dtype()},
98 has_auxiliary_outputs_{softmax_lse.has_value() && s_dmask.has_value()},
99 device_index_{q.device().index()} {
100 assert(q.ndim() == 3 &&
101 ((!block_table.has_value() && k.ndim() == 3 && v.ndim() == 3) ||
102 (block_table.has_value() && k.ndim() == 4 && v.ndim() == 4)) &&
103 "`FlashAttnVarlenFunc` requires packed 3D Q and either packed 3D "
104 "or paged 4D K and V tensors");
105 assert(k.shape() == v.shape() &&
106 "`FlashAttnVarlenFunc` requires K and V to have the same shape");
107 assert(q.size(1) > 0 && k.size(-2) > 0 && q.size(2) == k.size(-1) &&
108 q.size(1) % k.size(-2) == 0 &&
109 "`FlashAttnVarlenFunc` requires compatible Q and KV heads");
110 assert(q.size(2) > 0 && q.size(2) <= 256 && q.size(2) % 8 == 0 &&
111 "`FlashAttnVarlenFunc` requires a head dimension divisible by 8 "
112 "and no greater than 256");
113 assert(out.shape() == q.shape() &&
114 "`FlashAttnVarlenFunc` output must have the same shape as Q");
115 assert(softmax_lse.has_value() == s_dmask.has_value() &&
116 "`FlashAttnVarlenFunc` auxiliary outputs must be provided "
117 "together");
118 assert(return_attn_probs == has_auxiliary_outputs_ &&
119 "`FlashAttnVarlenFunc` requires auxiliary outputs exactly when "
120 "`return_attn_probs` is true");
121 if (has_auxiliary_outputs_) {
122 assert((softmax_lse->shape() == Tensor::Shape{q.size(1), q.size(0)} &&
123 softmax_lse_dtype_ == DataType::kFloat32 &&
124 "`FlashAttnVarlenFunc` softmax LSE output must have shape "
125 "(num_heads, total_q) and dtype float32"));
126 assert(s_dmask->shape() == Tensor::Shape{0} &&
127 s_dmask_dtype_ == q_dtype_ &&
128 "`FlashAttnVarlenFunc` inference attention mask output must be "
129 "empty and match the Q dtype");
130 }
131 assert(
132 (q_dtype_ == DataType::kFloat16 || q_dtype_ == DataType::kBFloat16) &&
133 q_dtype_ == k_dtype_ && q_dtype_ == v_dtype_ &&
134 q_dtype_ == out_dtype_ &&
135 "`FlashAttnVarlenFunc` requires matching float16 or bfloat16 Q, "
136 "K, V, and output dtypes");
137 assert(q.stride(-1) == 1 && k.stride(-1) == 1 && v.stride(-1) == 1 &&
138 out.stride(-1) == 1 &&
139 "`FlashAttnVarlenFunc` requires contiguous head dimensions");
140 assert(cu_seqlens_q.ndim() == 1 && cu_seqlens_k.ndim() == 1 &&
141 cu_seqlens_q.shape() == cu_seqlens_k.shape() &&
142 cu_seqlens_q.numel() >= 2 &&
143 "`FlashAttnVarlenFunc` cumulative sequence tensors must be "
144 "matching non-empty vectors");
145 assert(cu_seqlens_q_dtype_ == DataType::kInt32 &&
146 cu_seqlens_k_dtype_ == DataType::kInt32 &&
147 cu_seqlens_q.IsContiguous() && cu_seqlens_k.IsContiguous() &&
148 "`FlashAttnVarlenFunc` cumulative sequence tensors must be "
149 "contiguous int32 tensors");
150 assert(max_seqlen_q > 0 && max_seqlen_k > 0 &&
151 "`FlashAttnVarlenFunc` maximum sequence lengths must be positive");
152 assert(window_size.size() == 2 && window_size[0] >= -1 &&
153 window_size[1] >= -1 &&
154 "`FlashAttnVarlenFunc` `window_size` must contain two values >= -1");
155
156 assert(dropout_p == 0.0 &&
157 "`FlashAttnVarlenFunc` initially supports inference only");
158 assert(softcap == 0.0 &&
159 "`FlashAttnVarlenFunc` does not yet support softcap");
160 assert(!deterministic &&
161 "`FlashAttnVarlenFunc` does not yet support deterministic mode");
162 if (block_table.has_value()) {
163 assert(block_table->ndim() == 2 &&
164 block_table->size(0) + 1 == cu_seqlens_q.size(0) &&
165 block_table_dtype_ == DataType::kInt32 &&
166 block_table->IsContiguous() && k.size(1) % 256 == 0 &&
167 "`FlashAttnVarlenFunc` requires a contiguous int32 block table "
168 "and page size divisible by 256");
169 }
170 if (alibi_slopes.has_value()) {
171 assert(
172 (alibi_slopes->ndim() == 1 || alibi_slopes->ndim() == 2) &&
173 alibi_slopes_dtype_ == DataType::kFloat32 &&
174 alibi_slopes->IsContiguous() &&
175 ((alibi_slopes->ndim() == 1 && alibi_slopes->size(0) == q.size(1)) ||
176 (alibi_slopes->ndim() == 2 &&
177 alibi_slopes->size(0) + 1 == cu_seqlens_q.size(0) &&
178 alibi_slopes->size(1) == q.size(1))) &&
179 "`FlashAttnVarlenFunc` received incompatible ALiBi slopes");
180 }
181
182 const auto same_device_as_q = [&](const Tensor tensor) {
183 return tensor.device().type() == q.device().type() &&
184 tensor.device().index() == q.device().index();
185 };
186 assert(same_device_as_q(k) && same_device_as_q(v) &&
187 same_device_as_q(cu_seqlens_q) && same_device_as_q(cu_seqlens_k) &&
188 same_device_as_q(out) &&
189 (!alibi_slopes.has_value() || same_device_as_q(*alibi_slopes)) &&
190 (!block_table.has_value() || same_device_as_q(*block_table)) &&
191 (!softmax_lse.has_value() || same_device_as_q(*softmax_lse)) &&
192 (!s_dmask.has_value() || same_device_as_q(*s_dmask)) &&
193 "`FlashAttnVarlenFunc` tensors must be on the same device");
194
195 (void)softmax_scale;
196 (void)causal;
197 }
198
199 void operator()(const Tensor q, const Tensor k, const Tensor v,
200 const Tensor cu_seqlens_q, const Tensor cu_seqlens_k,
201 const int64_t max_seqlen_q, const int64_t max_seqlen_k,
202 Tensor out) const {
203 (*this)(q, k, v, cu_seqlens_q, cu_seqlens_k, std::nullopt, std::nullopt,
204 max_seqlen_q, max_seqlen_k, 0.0, std::nullopt, false, {-1, -1}, 0.0,
205 false, false, out, std::nullopt, std::nullopt);
206 }
207
208 virtual void operator()(
209 const Tensor q, const Tensor k, const Tensor v, const Tensor cu_seqlens_q,
210 const Tensor cu_seqlens_k, const std::optional<Tensor> alibi_slopes,
211 const std::optional<Tensor> block_table, const int64_t max_seqlen_q,
212 const int64_t max_seqlen_k, const double dropout_p,
213 const std::optional<double> softmax_scale, const bool causal,
214 const std::vector<int64_t> window_size, const double softcap,
215 const bool deterministic, const bool return_attn_probs, Tensor out,
216 std::optional<Tensor> softmax_lse,
217 std::optional<Tensor> s_dmask) const = 0;
218
219 protected:
220 Tensor::Shape q_shape_;
221
222 Tensor::Shape k_shape_;
223
224 Tensor::Shape v_shape_;
225
226 Tensor::Shape cu_seqlens_q_shape_;
227
228 Tensor::Shape cu_seqlens_k_shape_;
229
230 Tensor::Shape alibi_slopes_shape_;
231
232 Tensor::Shape block_table_shape_;
233
234 Tensor::Shape out_shape_;
235
236 Tensor::Shape softmax_lse_shape_;
237
238 Tensor::Shape s_dmask_shape_;
239
240 Tensor::Strides q_strides_;
241
242 Tensor::Strides k_strides_;
243
244 Tensor::Strides v_strides_;
245
246 Tensor::Strides cu_seqlens_q_strides_;
247
248 Tensor::Strides cu_seqlens_k_strides_;
249
250 Tensor::Strides alibi_slopes_strides_;
251
252 Tensor::Strides block_table_strides_;
253
254 Tensor::Strides out_strides_;
255
256 Tensor::Strides softmax_lse_strides_;
257
258 Tensor::Strides s_dmask_strides_;
259
260 DataType q_dtype_;
261
262 DataType k_dtype_;
263
264 DataType v_dtype_;
265
267
269
271
273
274 DataType out_dtype_;
275
277
279
280 bool has_auxiliary_outputs_{false};
281
282 int device_index_{0};
283};
284
285} // namespace infini::ops
286
287#endif // INFINI_OPS_BASE_FLASH_ATTN_VARLEN_FUNC_H_
Definition flash_attn_varlen_func.h:15
DataType s_dmask_dtype_
Definition flash_attn_varlen_func.h:278
DataType cu_seqlens_q_dtype_
Definition flash_attn_varlen_func.h:266
DataType softmax_lse_dtype_
Definition flash_attn_varlen_func.h:276
Tensor::Shape s_dmask_shape_
Definition flash_attn_varlen_func.h:238
Tensor::Strides cu_seqlens_q_strides_
Definition flash_attn_varlen_func.h:246
Tensor::Shape k_shape_
Definition flash_attn_varlen_func.h:222
DataType block_table_dtype_
Definition flash_attn_varlen_func.h:272
DataType cu_seqlens_k_dtype_
Definition flash_attn_varlen_func.h:268
DataType alibi_slopes_dtype_
Definition flash_attn_varlen_func.h:270
virtual void operator()(const Tensor q, const Tensor k, const Tensor v, const Tensor cu_seqlens_q, const Tensor cu_seqlens_k, const std::optional< Tensor > alibi_slopes, const std::optional< Tensor > block_table, const int64_t max_seqlen_q, const int64_t max_seqlen_k, const double dropout_p, const std::optional< double > softmax_scale, const bool causal, const std::vector< int64_t > window_size, const double softcap, const bool deterministic, const bool return_attn_probs, Tensor out, std::optional< Tensor > softmax_lse, std::optional< Tensor > s_dmask) const =0
Tensor::Strides cu_seqlens_k_strides_
Definition flash_attn_varlen_func.h:248
FlashAttnVarlenFunc(const Tensor q, const Tensor k, const Tensor v, const Tensor cu_seqlens_q, const Tensor cu_seqlens_k, const std::optional< Tensor > alibi_slopes, const std::optional< Tensor > block_table, const int64_t max_seqlen_q, const int64_t max_seqlen_k, const double dropout_p, const std::optional< double > softmax_scale, const bool causal, const std::vector< int64_t > window_size, const double softcap, const bool deterministic, const bool return_attn_probs, Tensor out, std::optional< Tensor > softmax_lse, std::optional< Tensor > s_dmask)
Definition flash_attn_varlen_func.h:41
Tensor::Strides block_table_strides_
Definition flash_attn_varlen_func.h:252
Tensor::Shape q_shape_
Definition flash_attn_varlen_func.h:220
DataType q_dtype_
Definition flash_attn_varlen_func.h:260
Tensor::Shape v_shape_
Definition flash_attn_varlen_func.h:224
DataType out_dtype_
Definition flash_attn_varlen_func.h:274
FlashAttnVarlenFunc(const Tensor q, const Tensor k, const Tensor v, const Tensor cu_seqlens_q, const Tensor cu_seqlens_k, const int64_t max_seqlen_q, const int64_t max_seqlen_k, Tensor out)
Definition flash_attn_varlen_func.h:17
Tensor::Strides q_strides_
Definition flash_attn_varlen_func.h:240
Tensor::Strides alibi_slopes_strides_
Definition flash_attn_varlen_func.h:250
DataType k_dtype_
Definition flash_attn_varlen_func.h:262
Tensor::Shape cu_seqlens_q_shape_
Definition flash_attn_varlen_func.h:226
Tensor::Strides s_dmask_strides_
Definition flash_attn_varlen_func.h:258
DataType v_dtype_
Definition flash_attn_varlen_func.h:264
Tensor::Shape alibi_slopes_shape_
Definition flash_attn_varlen_func.h:230
Tensor::Shape out_shape_
Definition flash_attn_varlen_func.h:234
Tensor::Strides v_strides_
Definition flash_attn_varlen_func.h:244
Tensor::Strides softmax_lse_strides_
Definition flash_attn_varlen_func.h:256
Tensor::Strides out_strides_
Definition flash_attn_varlen_func.h:254
Tensor::Strides k_strides_
Definition flash_attn_varlen_func.h:242
Tensor::Shape cu_seqlens_k_shape_
Definition flash_attn_varlen_func.h:228
Tensor::Shape block_table_shape_
Definition flash_attn_varlen_func.h:232
Tensor::Shape softmax_lse_shape_
Definition flash_attn_varlen_func.h:236
void operator()(const Tensor q, const Tensor k, const Tensor v, const Tensor cu_seqlens_q, const Tensor cu_seqlens_k, const int64_t max_seqlen_q, const int64_t max_seqlen_k, Tensor out) const
Definition flash_attn_varlen_func.h:199
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8