InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
flash_attn_with_kvcache.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_FLASH_ATTN_WITH_KVCACHE_H_
2#define INFINI_OPS_BASE_FLASH_ATTN_WITH_KVCACHE_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// Inference attention with an optional in-place KV-cache update, aligned with
14// Dao-AILab FlashAttention's `flash_attn_with_kvcache` public interface.
15class FlashAttnWithKvcache : public Operator<FlashAttnWithKvcache> {
16 public:
17 FlashAttnWithKvcache(const Tensor q, Tensor k_cache, Tensor v_cache,
18 Tensor out)
20 k_cache,
21 v_cache,
22 std::nullopt,
23 std::nullopt,
24 std::nullopt,
25 std::nullopt,
26 std::optional<Tensor>{},
27 std::nullopt,
28 std::nullopt,
29 std::nullopt,
30 std::nullopt,
31 std::nullopt,
32 false,
33 {-1, -1},
34 0.0,
35 true,
36 0,
37 false,
38 out,
39 std::nullopt} {}
40
41 FlashAttnWithKvcache(const Tensor q, Tensor k_cache, Tensor v_cache,
42 const std::optional<Tensor> k,
43 const std::optional<Tensor> v,
44 const std::optional<Tensor> rotary_cos,
45 const std::optional<Tensor> rotary_sin,
46 const int64_t cache_seqlens,
47 const std::optional<Tensor> cache_batch_idx,
48 const std::optional<Tensor> cache_leftpad,
49 const std::optional<Tensor> block_table,
50 const std::optional<Tensor> alibi_slopes,
51 const std::optional<double> softmax_scale,
52 const bool causal,
53 const std::vector<int64_t> window_size,
54 const double softcap, const bool rotary_interleaved,
55 const int64_t num_splits, const bool return_softmax_lse,
56 Tensor out, std::optional<Tensor> softmax_lse)
58 k_cache,
59 v_cache,
60 k,
61 v,
62 rotary_cos,
63 rotary_sin,
64 std::optional<Tensor>{},
65 cache_batch_idx,
66 cache_leftpad,
67 block_table,
68 alibi_slopes,
69 softmax_scale,
70 causal,
71 window_size,
72 softcap,
73 rotary_interleaved,
74 num_splits,
75 return_softmax_lse,
76 out,
77 softmax_lse} {
78 assert(cache_seqlens >= 0 &&
79 "`FlashAttnWithKvcache` requires non-negative scalar "
80 "`cache_seqlens`");
81 }
82
83 FlashAttnWithKvcache(const Tensor q, Tensor k_cache, Tensor v_cache,
84 const std::optional<Tensor> k,
85 const std::optional<Tensor> v,
86 const std::optional<Tensor> rotary_cos,
87 const std::optional<Tensor> rotary_sin,
88 const std::optional<Tensor> cache_seqlens,
89 const std::optional<Tensor> cache_batch_idx,
90 const std::optional<Tensor> cache_leftpad,
91 const std::optional<Tensor> block_table,
92 const std::optional<Tensor> alibi_slopes,
93 const std::optional<double> softmax_scale,
94 const bool causal,
95 const std::vector<int64_t> window_size,
96 const double softcap, const bool rotary_interleaved,
97 const int64_t num_splits, const bool return_softmax_lse,
98 Tensor out, std::optional<Tensor> softmax_lse)
99 : q_shape_{q.shape()},
100 k_cache_shape_{k_cache.shape()},
101 v_cache_shape_{v_cache.shape()},
102 k_shape_{k.has_value() ? Tensor::Shape{k->shape()} : Tensor::Shape{}},
103 v_shape_{v.has_value() ? Tensor::Shape{v->shape()} : Tensor::Shape{}},
104 rotary_cos_shape_{rotary_cos.has_value()
105 ? Tensor::Shape{rotary_cos->shape()}
106 : Tensor::Shape{}},
107 rotary_sin_shape_{rotary_sin.has_value()
108 ? Tensor::Shape{rotary_sin->shape()}
109 : Tensor::Shape{}},
110 cache_seqlens_shape_{cache_seqlens.has_value()
111 ? Tensor::Shape{cache_seqlens->shape()}
112 : Tensor::Shape{}},
113 cache_batch_idx_shape_{cache_batch_idx.has_value()
114 ? Tensor::Shape{cache_batch_idx->shape()}
115 : Tensor::Shape{}},
116 cache_leftpad_shape_{cache_leftpad.has_value()
117 ? Tensor::Shape{cache_leftpad->shape()}
118 : Tensor::Shape{}},
119 block_table_shape_{block_table.has_value()
120 ? Tensor::Shape{block_table->shape()}
121 : Tensor::Shape{}},
122 alibi_slopes_shape_{alibi_slopes.has_value()
123 ? Tensor::Shape{alibi_slopes->shape()}
124 : Tensor::Shape{}},
125 out_shape_{out.shape()},
126 softmax_lse_shape_{softmax_lse.has_value()
127 ? Tensor::Shape{softmax_lse->shape()}
128 : Tensor::Shape{}},
129 q_strides_{q.strides()},
130 k_cache_strides_{k_cache.strides()},
131 v_cache_strides_{v_cache.strides()},
132 k_strides_{k.has_value() ? Tensor::Strides{k->strides()}
133 : Tensor::Strides{}},
134 v_strides_{v.has_value() ? Tensor::Strides{v->strides()}
135 : Tensor::Strides{}},
136 rotary_cos_strides_{rotary_cos.has_value()
137 ? Tensor::Strides{rotary_cos->strides()}
138 : Tensor::Strides{}},
139 rotary_sin_strides_{rotary_sin.has_value()
140 ? Tensor::Strides{rotary_sin->strides()}
141 : Tensor::Strides{}},
142 cache_seqlens_strides_{cache_seqlens.has_value()
143 ? Tensor::Strides{cache_seqlens->strides()}
144 : Tensor::Strides{}},
145 cache_batch_idx_strides_{
146 cache_batch_idx.has_value()
147 ? Tensor::Strides{cache_batch_idx->strides()}
148 : Tensor::Strides{}},
149 cache_leftpad_strides_{cache_leftpad.has_value()
150 ? Tensor::Strides{cache_leftpad->strides()}
151 : Tensor::Strides{}},
152 block_table_strides_{block_table.has_value()
153 ? Tensor::Strides{block_table->strides()}
154 : Tensor::Strides{}},
155 alibi_slopes_strides_{alibi_slopes.has_value()
156 ? Tensor::Strides{alibi_slopes->strides()}
157 : Tensor::Strides{}},
158 out_strides_{out.strides()},
159 softmax_lse_strides_{softmax_lse.has_value()
160 ? Tensor::Strides{softmax_lse->strides()}
161 : Tensor::Strides{}},
162 q_dtype_{q.dtype()},
163 k_cache_dtype_{k_cache.dtype()},
164 v_cache_dtype_{v_cache.dtype()},
165 k_dtype_{k.has_value() ? k->dtype() : q.dtype()},
166 v_dtype_{v.has_value() ? v->dtype() : q.dtype()},
167 rotary_cos_dtype_{rotary_cos.has_value() ? rotary_cos->dtype()
168 : q.dtype()},
169 rotary_sin_dtype_{rotary_sin.has_value() ? rotary_sin->dtype()
170 : q.dtype()},
171 cache_seqlens_dtype_{cache_seqlens.has_value() ? cache_seqlens->dtype()
172 : DataType::kInt32},
173 cache_batch_idx_dtype_{cache_batch_idx.has_value()
174 ? cache_batch_idx->dtype()
175 : DataType::kInt32},
176 cache_leftpad_dtype_{cache_leftpad.has_value() ? cache_leftpad->dtype()
177 : DataType::kInt32},
178 block_table_dtype_{block_table.has_value() ? block_table->dtype()
179 : DataType::kInt32},
180 alibi_slopes_dtype_{alibi_slopes.has_value() ? alibi_slopes->dtype()
181 : DataType::kFloat32},
182 out_dtype_{out.dtype()},
183 softmax_lse_dtype_{softmax_lse.has_value() ? softmax_lse->dtype()
184 : DataType::kFloat32},
185 has_k_{k.has_value()},
186 has_v_{v.has_value()},
187 has_rotary_cos_{rotary_cos.has_value()},
188 has_rotary_sin_{rotary_sin.has_value()},
189 has_cache_batch_idx_{cache_batch_idx.has_value()},
190 has_block_table_{block_table.has_value()},
191 has_alibi_slopes_{alibi_slopes.has_value()},
192 has_softmax_lse_{softmax_lse.has_value()},
193 batch_size_{q.ndim() > 0 ? q.size(0) : 0},
194 head_size_{q.ndim() == 4 ? q.size(3) : 0},
195 device_index_{q.device().index()} {
196 assert(q.ndim() == 4 && k_cache.ndim() == 4 && v_cache.ndim() == 4 &&
197 "`FlashAttnWithKvcache` requires 4D `q`, `k_cache`, and "
198 "`v_cache`");
199 assert(k_cache.shape() == v_cache.shape() &&
200 "`FlashAttnWithKvcache` requires matching K/V cache shapes");
201 assert(
202 (q_dtype_ == DataType::kFloat16 || q_dtype_ == DataType::kBFloat16) &&
203 q_dtype_ == k_cache_dtype_ && q_dtype_ == v_cache_dtype_ &&
204 q_dtype_ == out_dtype_ &&
205 "`FlashAttnWithKvcache` requires matching float16 or bfloat16 "
206 "Q, cache, and output dtypes");
207 assert(q.size(0) > 0 && q.size(1) > 0 && q.size(2) > 0 &&
208 k_cache.size(0) > 0 && k_cache.size(1) > 0 && k_cache.size(2) > 0 &&
209 "`FlashAttnWithKvcache` requires non-empty Q and KV cache "
210 "dimensions");
211 assert(q.size(2) % k_cache.size(2) == 0 && q.size(3) == k_cache.size(3) &&
212 "`FlashAttnWithKvcache` requires compatible Q and KV heads");
213 assert(head_size_ > 0 && head_size_ <= 256 &&
214 "`FlashAttnWithKvcache` requires a head dimension no greater than "
215 "256");
216 assert(out.shape() == q.shape() &&
217 "`FlashAttnWithKvcache` output must have the same shape as Q");
218 assert(return_softmax_lse == has_softmax_lse_ &&
219 "`FlashAttnWithKvcache` requires `softmax_lse` exactly when "
220 "`return_softmax_lse` is true");
221 if (has_softmax_lse_) {
222 assert((softmax_lse->shape() ==
223 Tensor::Shape{q.size(0), q.size(2), q.size(1)}) &&
224 softmax_lse_dtype_ == DataType::kFloat32 &&
225 "`FlashAttnWithKvcache` softmax LSE output must have shape "
226 "(batch_size, num_heads, seqlen) and dtype float32");
227 }
228 assert(q.stride(-1) == 1 && k_cache.stride(-1) == 1 &&
229 v_cache.stride(-1) == 1 && out.stride(-1) == 1 &&
230 "`FlashAttnWithKvcache` requires contiguous head dimensions");
231 assert(has_k_ == has_v_ &&
232 "`FlashAttnWithKvcache` requires `k` and `v` together");
233 if (has_k_) {
234 assert(k->ndim() == 4 && k->shape() == v->shape() &&
235 k->size(0) == batch_size_ && k->size(2) == k_cache.size(2) &&
236 k->size(3) == head_size_ &&
237 "`FlashAttnWithKvcache` received incompatible new K/V shapes");
238 assert(k_dtype_ == q_dtype_ && v_dtype_ == q_dtype_ &&
239 "`FlashAttnWithKvcache` requires matching new K/V dtypes");
240 assert(k->stride(-1) == 1 && v->stride(-1) == 1 &&
241 "`FlashAttnWithKvcache` requires contiguous new K/V head "
242 "dimensions");
243 }
244 assert(has_rotary_cos_ == has_rotary_sin_ &&
245 "`FlashAttnWithKvcache` requires rotary cosine and sine together");
246 if (has_rotary_cos_) {
247 assert(has_k_ && rotary_cos->ndim() == 2 &&
248 rotary_cos->shape() == rotary_sin->shape() &&
249 rotary_cos_dtype_ == q_dtype_ && rotary_sin_dtype_ == q_dtype_ &&
250 rotary_cos->size(1) > 0 && rotary_cos->size(1) * 2 <= head_size_ &&
251 (rotary_cos->size(1) * 2) % 16 == 0 &&
252 "`FlashAttnWithKvcache` received incompatible rotary tables");
253 }
254 ValidateIndexVector(cache_seqlens);
255 ValidateIndexVector(cache_batch_idx);
256 ValidateIndexVector(cache_leftpad);
257 if (has_block_table_) {
258 assert(block_table->ndim() == 2 && block_table->size(0) == batch_size_ &&
259 block_table_dtype_ == DataType::kInt32 &&
260 block_table->IsContiguous() && k_cache.size(1) % 256 == 0 &&
261 "`FlashAttnWithKvcache` requires a contiguous int32 block table "
262 "and page size divisible by 256");
263 } else {
264 assert((has_cache_batch_idx_ || k_cache.size(0) >= batch_size_) &&
265 "`FlashAttnWithKvcache` cache batch is too small");
266 }
267 if (has_alibi_slopes_) {
268 assert((alibi_slopes->ndim() == 1 || alibi_slopes->ndim() == 2) &&
269 alibi_slopes_dtype_ == DataType::kFloat32 &&
270 alibi_slopes->IsContiguous() &&
271 "`FlashAttnWithKvcache` requires contiguous float32 ALiBi "
272 "slopes");
273 assert(
274 ((alibi_slopes->ndim() == 1 && alibi_slopes->size(0) == q.size(2)) ||
275 (alibi_slopes->ndim() == 2 && alibi_slopes->size(0) == batch_size_ &&
276 alibi_slopes->size(1) == q.size(2))) &&
277 "`FlashAttnWithKvcache` received incompatible ALiBi slopes");
278 }
279 assert(window_size.size() == 2 && window_size[0] >= -1 &&
280 window_size[1] >= -1 &&
281 "`FlashAttnWithKvcache` `window_size` must contain two values >= "
282 "-1");
283 assert(softcap >= 0.0 &&
284 "`FlashAttnWithKvcache` requires non-negative `softcap`");
285 assert(num_splits >= 0 &&
286 "`FlashAttnWithKvcache` requires non-negative `num_splits`");
287 const auto same_device_as_q = [&](const Tensor tensor) {
288 return tensor.device().type() == q.device().type() &&
289 tensor.device().index() == q.device().index();
290 };
291 assert(
292 same_device_as_q(k_cache) && same_device_as_q(v_cache) &&
293 same_device_as_q(out) &&
294 (!softmax_lse.has_value() || same_device_as_q(*softmax_lse)) &&
295 (!k.has_value() || same_device_as_q(*k)) &&
296 (!v.has_value() || same_device_as_q(*v)) &&
297 (!rotary_cos.has_value() || same_device_as_q(*rotary_cos)) &&
298 (!rotary_sin.has_value() || same_device_as_q(*rotary_sin)) &&
299 (!cache_seqlens.has_value() || same_device_as_q(*cache_seqlens)) &&
300 (!cache_batch_idx.has_value() || same_device_as_q(*cache_batch_idx)) &&
301 (!cache_leftpad.has_value() || same_device_as_q(*cache_leftpad)) &&
302 (!block_table.has_value() || same_device_as_q(*block_table)) &&
303 (!alibi_slopes.has_value() || same_device_as_q(*alibi_slopes)) &&
304 "`FlashAttnWithKvcache` tensors must be on the same device");
305
306 (void)softmax_scale;
307 (void)causal;
308 (void)rotary_interleaved;
309 }
310
311 void operator()(const Tensor q, Tensor k_cache, Tensor v_cache,
312 Tensor out) const {
313 (*this)(q, k_cache, v_cache, std::nullopt, std::nullopt, std::nullopt,
314 std::nullopt, std::optional<Tensor>{}, std::nullopt, std::nullopt,
315 std::nullopt, std::nullopt, std::nullopt, false, {-1, -1}, 0.0,
316 true, 0, false, out, std::nullopt);
317 }
318
319 virtual void operator()(
320 const Tensor q, Tensor k_cache, Tensor v_cache,
321 const std::optional<Tensor> k, const std::optional<Tensor> v,
322 const std::optional<Tensor> rotary_cos,
323 const std::optional<Tensor> rotary_sin, const int64_t cache_seqlens,
324 const std::optional<Tensor> cache_batch_idx,
325 const std::optional<Tensor> cache_leftpad,
326 const std::optional<Tensor> block_table,
327 const std::optional<Tensor> alibi_slopes,
328 const std::optional<double> softmax_scale, const bool causal,
329 const std::vector<int64_t> window_size, const double softcap,
330 const bool rotary_interleaved, const int64_t num_splits,
331 const bool return_softmax_lse, Tensor out,
332 std::optional<Tensor> softmax_lse) const = 0;
333
334 virtual void operator()(const Tensor q, Tensor k_cache, Tensor v_cache,
335 const std::optional<Tensor> k,
336 const std::optional<Tensor> v,
337 const std::optional<Tensor> rotary_cos,
338 const std::optional<Tensor> rotary_sin,
339 const std::optional<Tensor> cache_seqlens,
340 const std::optional<Tensor> cache_batch_idx,
341 const std::optional<Tensor> cache_leftpad,
342 const std::optional<Tensor> block_table,
343 const std::optional<Tensor> alibi_slopes,
344 const std::optional<double> softmax_scale,
345 const bool causal,
346 const std::vector<int64_t> window_size,
347 const double softcap, const bool rotary_interleaved,
348 const int64_t num_splits,
349 const bool return_softmax_lse, Tensor out,
350 std::optional<Tensor> softmax_lse) const = 0;
351
352 protected:
353 void ValidateIndexVector(const std::optional<Tensor>& tensor) const {
354 if (!tensor.has_value()) {
355 return;
356 }
357 assert(tensor->ndim() == 1 && tensor->size(0) == batch_size_ &&
358 tensor->dtype() == DataType::kInt32 && tensor->IsContiguous() &&
359 "`FlashAttnWithKvcache` index metadata must be contiguous int32 "
360 "vectors with one value per query batch");
361 }
362
363 Tensor::Shape q_shape_;
364
365 Tensor::Shape k_cache_shape_;
366
367 Tensor::Shape v_cache_shape_;
368
369 Tensor::Shape k_shape_;
370
371 Tensor::Shape v_shape_;
372
373 Tensor::Shape rotary_cos_shape_;
374
375 Tensor::Shape rotary_sin_shape_;
376
377 Tensor::Shape cache_seqlens_shape_;
378
380
381 Tensor::Shape cache_leftpad_shape_;
382
383 Tensor::Shape block_table_shape_;
384
385 Tensor::Shape alibi_slopes_shape_;
386
387 Tensor::Shape out_shape_;
388
389 Tensor::Shape softmax_lse_shape_;
390
391 Tensor::Strides q_strides_;
392
393 Tensor::Strides k_cache_strides_;
394
395 Tensor::Strides v_cache_strides_;
396
397 Tensor::Strides k_strides_;
398
399 Tensor::Strides v_strides_;
400
401 Tensor::Strides rotary_cos_strides_;
402
403 Tensor::Strides rotary_sin_strides_;
404
405 Tensor::Strides cache_seqlens_strides_;
406
408
409 Tensor::Strides cache_leftpad_strides_;
410
411 Tensor::Strides block_table_strides_;
412
413 Tensor::Strides alibi_slopes_strides_;
414
415 Tensor::Strides out_strides_;
416
417 Tensor::Strides softmax_lse_strides_;
418
419 DataType q_dtype_;
420
422
424
425 DataType k_dtype_;
426
427 DataType v_dtype_;
428
430
432
434
436
438
440
442
443 DataType out_dtype_;
444
446
447 bool has_k_{false};
448
449 bool has_v_{false};
450
451 bool has_rotary_cos_{false};
452
453 bool has_rotary_sin_{false};
454
455 bool has_cache_batch_idx_{false};
456
457 bool has_block_table_{false};
458
459 bool has_alibi_slopes_{false};
460
461 bool has_softmax_lse_{false};
462
463 Tensor::Size batch_size_{0};
464
465 Tensor::Size head_size_{0};
466
467 int device_index_{0};
468};
469
470} // namespace infini::ops
471
472#endif // INFINI_OPS_BASE_FLASH_ATTN_WITH_KVCACHE_H_
Definition flash_attn_with_kvcache.h:15
DataType k_cache_dtype_
Definition flash_attn_with_kvcache.h:421
FlashAttnWithKvcache(const Tensor q, Tensor k_cache, Tensor v_cache, const std::optional< Tensor > k, const std::optional< Tensor > v, const std::optional< Tensor > rotary_cos, const std::optional< Tensor > rotary_sin, const std::optional< Tensor > cache_seqlens, const std::optional< Tensor > cache_batch_idx, const std::optional< Tensor > cache_leftpad, const std::optional< Tensor > block_table, const std::optional< Tensor > alibi_slopes, const std::optional< double > softmax_scale, const bool causal, const std::vector< int64_t > window_size, const double softcap, const bool rotary_interleaved, const int64_t num_splits, const bool return_softmax_lse, Tensor out, std::optional< Tensor > softmax_lse)
Definition flash_attn_with_kvcache.h:83
virtual void operator()(const Tensor q, Tensor k_cache, Tensor v_cache, const std::optional< Tensor > k, const std::optional< Tensor > v, const std::optional< Tensor > rotary_cos, const std::optional< Tensor > rotary_sin, const int64_t cache_seqlens, const std::optional< Tensor > cache_batch_idx, const std::optional< Tensor > cache_leftpad, const std::optional< Tensor > block_table, const std::optional< Tensor > alibi_slopes, const std::optional< double > softmax_scale, const bool causal, const std::vector< int64_t > window_size, const double softcap, const bool rotary_interleaved, const int64_t num_splits, const bool return_softmax_lse, Tensor out, std::optional< Tensor > softmax_lse) const =0
Tensor::Strides out_strides_
Definition flash_attn_with_kvcache.h:415
Tensor::Strides k_strides_
Definition flash_attn_with_kvcache.h:397
Tensor::Strides q_strides_
Definition flash_attn_with_kvcache.h:391
DataType q_dtype_
Definition flash_attn_with_kvcache.h:419
Tensor::Shape k_cache_shape_
Definition flash_attn_with_kvcache.h:365
Tensor::Strides v_strides_
Definition flash_attn_with_kvcache.h:399
Tensor::Shape alibi_slopes_shape_
Definition flash_attn_with_kvcache.h:385
DataType rotary_sin_dtype_
Definition flash_attn_with_kvcache.h:431
Tensor::Shape softmax_lse_shape_
Definition flash_attn_with_kvcache.h:389
Tensor::Strides rotary_cos_strides_
Definition flash_attn_with_kvcache.h:401
virtual void operator()(const Tensor q, Tensor k_cache, Tensor v_cache, const std::optional< Tensor > k, const std::optional< Tensor > v, const std::optional< Tensor > rotary_cos, const std::optional< Tensor > rotary_sin, const std::optional< Tensor > cache_seqlens, const std::optional< Tensor > cache_batch_idx, const std::optional< Tensor > cache_leftpad, const std::optional< Tensor > block_table, const std::optional< Tensor > alibi_slopes, const std::optional< double > softmax_scale, const bool causal, const std::vector< int64_t > window_size, const double softcap, const bool rotary_interleaved, const int64_t num_splits, const bool return_softmax_lse, Tensor out, std::optional< Tensor > softmax_lse) const =0
Tensor::Strides cache_seqlens_strides_
Definition flash_attn_with_kvcache.h:405
Tensor::Strides k_cache_strides_
Definition flash_attn_with_kvcache.h:393
Tensor::Shape k_shape_
Definition flash_attn_with_kvcache.h:369
DataType out_dtype_
Definition flash_attn_with_kvcache.h:443
Tensor::Strides cache_leftpad_strides_
Definition flash_attn_with_kvcache.h:409
DataType v_dtype_
Definition flash_attn_with_kvcache.h:427
FlashAttnWithKvcache(const Tensor q, Tensor k_cache, Tensor v_cache, const std::optional< Tensor > k, const std::optional< Tensor > v, const std::optional< Tensor > rotary_cos, const std::optional< Tensor > rotary_sin, const int64_t cache_seqlens, const std::optional< Tensor > cache_batch_idx, const std::optional< Tensor > cache_leftpad, const std::optional< Tensor > block_table, const std::optional< Tensor > alibi_slopes, const std::optional< double > softmax_scale, const bool causal, const std::vector< int64_t > window_size, const double softcap, const bool rotary_interleaved, const int64_t num_splits, const bool return_softmax_lse, Tensor out, std::optional< Tensor > softmax_lse)
Definition flash_attn_with_kvcache.h:41
DataType softmax_lse_dtype_
Definition flash_attn_with_kvcache.h:445
Tensor::Shape v_shape_
Definition flash_attn_with_kvcache.h:371
Tensor::Shape rotary_cos_shape_
Definition flash_attn_with_kvcache.h:373
Tensor::Shape block_table_shape_
Definition flash_attn_with_kvcache.h:383
DataType rotary_cos_dtype_
Definition flash_attn_with_kvcache.h:429
Tensor::Strides block_table_strides_
Definition flash_attn_with_kvcache.h:411
Tensor::Shape v_cache_shape_
Definition flash_attn_with_kvcache.h:367
Tensor::Strides rotary_sin_strides_
Definition flash_attn_with_kvcache.h:403
Tensor::Shape cache_batch_idx_shape_
Definition flash_attn_with_kvcache.h:379
Tensor::Shape cache_leftpad_shape_
Definition flash_attn_with_kvcache.h:381
DataType cache_leftpad_dtype_
Definition flash_attn_with_kvcache.h:437
DataType alibi_slopes_dtype_
Definition flash_attn_with_kvcache.h:441
Tensor::Strides v_cache_strides_
Definition flash_attn_with_kvcache.h:395
Tensor::Shape out_shape_
Definition flash_attn_with_kvcache.h:387
Tensor::Strides cache_batch_idx_strides_
Definition flash_attn_with_kvcache.h:407
DataType v_cache_dtype_
Definition flash_attn_with_kvcache.h:423
DataType block_table_dtype_
Definition flash_attn_with_kvcache.h:439
void ValidateIndexVector(const std::optional< Tensor > &tensor) const
Definition flash_attn_with_kvcache.h:353
DataType cache_seqlens_dtype_
Definition flash_attn_with_kvcache.h:433
Tensor::Shape rotary_sin_shape_
Definition flash_attn_with_kvcache.h:375
FlashAttnWithKvcache(const Tensor q, Tensor k_cache, Tensor v_cache, Tensor out)
Definition flash_attn_with_kvcache.h:17
Tensor::Shape q_shape_
Definition flash_attn_with_kvcache.h:363
Tensor::Strides alibi_slopes_strides_
Definition flash_attn_with_kvcache.h:413
Tensor::Shape cache_seqlens_shape_
Definition flash_attn_with_kvcache.h:377
DataType cache_batch_idx_dtype_
Definition flash_attn_with_kvcache.h:435
DataType k_dtype_
Definition flash_attn_with_kvcache.h:425
void operator()(const Tensor q, Tensor k_cache, Tensor v_cache, Tensor out) const
Definition flash_attn_with_kvcache.h:311
Tensor::Strides softmax_lse_strides_
Definition flash_attn_with_kvcache.h:417
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8