InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
paged_attention_infinilm.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_PAGED_ATTENTION_INFINILM_H_
2#define INFINI_OPS_BASE_PAGED_ATTENTION_INFINILM_H_
3
4#include <cassert>
5#include <cstddef>
6#include <optional>
7
8#include "data_type.h"
9#include "operator.h"
10#include "tensor.h"
11
12namespace infini::ops {
13
16class [[deprecated(
17 "Migrate to an open-source-aligned operator when available.")]]
18PagedAttentionInfinilm : public Operator<PagedAttentionInfinilm> {
19 public:
20 PagedAttentionInfinilm(const Tensor q, const Tensor k_cache,
21 const Tensor v_cache, const Tensor block_tables,
22 const Tensor seq_lens,
23 std::optional<Tensor> alibi_slopes, float scale,
24 Tensor out)
25 : dtype_{q.dtype()},
26 index_dtype_{block_tables.dtype()},
27 scale_{scale},
28 num_seqs_{q.size(0)},
29 num_heads_{q.size(1)},
30 num_kv_heads_{k_cache.size(1)},
31 head_size_{q.size(2)},
32 block_size_{k_cache.size(2)},
33 max_num_blocks_per_seq_{block_tables.size(1)},
34 q_stride_{q.stride(0)},
35 q_head_stride_{q.stride(1)},
36 k_cache_block_stride_{k_cache.stride(0)},
37 k_cache_head_stride_{k_cache.stride(1)},
38 k_cache_slot_stride_{k_cache.stride(2)},
39 v_cache_block_stride_{v_cache.stride(0)},
40 v_cache_head_stride_{v_cache.stride(1)},
41 v_cache_slot_stride_{v_cache.stride(2)},
42 out_stride_{out.stride(0)},
43 out_head_stride_{out.stride(1)},
44 block_table_batch_stride_{block_tables.stride(0)},
45 seq_lens_stride_{seq_lens.stride(0)} {
46 assert(q.ndim() == 3 && out.ndim() == 3);
47 assert(k_cache.ndim() == 4 && v_cache.ndim() == 4);
48 assert(block_tables.ndim() == 2 && seq_lens.ndim() == 1);
49 assert((dtype_ == DataType::kFloat16 || dtype_ == DataType::kBFloat16) &&
50 "`PagedAttentionInfinilm` supports float16 and bfloat16");
51 assert(out.dtype() == dtype_ && k_cache.dtype() == dtype_ &&
52 v_cache.dtype() == dtype_);
53 assert(IsIndexDtype(index_dtype_) && seq_lens.dtype() == index_dtype_);
54 assert(q.shape() == out.shape());
55 assert(k_cache.shape() == v_cache.shape());
56 assert(block_tables.size(0) == num_seqs_ && seq_lens.size(0) == num_seqs_);
57 assert(k_cache.size(1) == num_kv_heads_ &&
58 v_cache.size(1) == num_kv_heads_);
59 assert(k_cache.size(3) == head_size_ && v_cache.size(3) == head_size_);
60 assert((head_size_ == 64 || head_size_ == 128) &&
61 "`PagedAttentionInfinilm` supports head sizes 64 and 128");
62 assert(num_heads_ % num_kv_heads_ == 0);
63 assert(q.stride(2) == 1 && out.stride(2) == 1);
64 assert(k_cache.stride(3) == 1 && v_cache.stride(3) == 1);
65 assert(block_tables.stride(1) == 1 && seq_lens.stride(0) == 1);
66 assert(!alibi_slopes.has_value() ||
67 (alibi_slopes->dtype() == DataType::kFloat32 &&
68 alibi_slopes->ndim() == 1 && alibi_slopes->size(0) == num_heads_ &&
69 alibi_slopes->stride(0) == 1));
70 }
71
72 virtual void operator()(const Tensor q, const Tensor k_cache,
73 const Tensor v_cache, const Tensor block_tables,
74 const Tensor seq_lens,
75 std::optional<Tensor> alibi_slopes, float scale,
76 Tensor out) const = 0;
77
78 protected:
79 static bool IsIndexDtype(DataType dtype) {
80 return dtype == DataType::kInt32 || dtype == DataType::kInt64 ||
81 dtype == DataType::kUInt32;
82 }
83
84 DataType dtype_;
85 DataType index_dtype_;
86 float scale_{1.0f};
87 std::size_t num_seqs_{0};
88 std::size_t num_heads_{0};
89 std::size_t num_kv_heads_{0};
90 std::size_t head_size_{0};
91 std::size_t block_size_{0};
92 std::size_t max_num_blocks_per_seq_{0};
93 Tensor::Stride q_stride_{0};
94 Tensor::Stride q_head_stride_{0};
95 Tensor::Stride k_cache_block_stride_{0};
96 Tensor::Stride k_cache_head_stride_{0};
97 Tensor::Stride k_cache_slot_stride_{0};
98 Tensor::Stride v_cache_block_stride_{0};
99 Tensor::Stride v_cache_head_stride_{0};
100 Tensor::Stride v_cache_slot_stride_{0};
101 Tensor::Stride out_stride_{0};
102 Tensor::Stride out_head_stride_{0};
103 Tensor::Stride block_table_batch_stride_{0};
104 Tensor::Stride seq_lens_stride_{0};
105};
106
107} // namespace infini::ops
108
109#endif
Definition generated/include/operator.h:282
Definition paged_attention_infinilm.h:18
DataType dtype_
Definition paged_attention_infinilm.h:84
DataType index_dtype_
Definition paged_attention_infinilm.h:85
virtual void operator()(const Tensor q, const Tensor k_cache, const Tensor v_cache, const Tensor block_tables, const Tensor seq_lens, std::optional< Tensor > alibi_slopes, float scale, Tensor out) const =0
PagedAttentionInfinilm(const Tensor q, const Tensor k_cache, const Tensor v_cache, const Tensor block_tables, const Tensor seq_lens, std::optional< Tensor > alibi_slopes, float scale, Tensor out)
Definition paged_attention_infinilm.h:20
static bool IsIndexDtype(DataType dtype)
Definition paged_attention_infinilm.h:79
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8