InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
paged_attention_prefill_infinilm.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_PAGED_ATTENTION_PREFILL_INFINILM_H_
2#define INFINI_OPS_BASE_PAGED_ATTENTION_PREFILL_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.")]]
18PagedAttentionPrefillInfinilm : public Operator<PagedAttentionPrefillInfinilm> {
19 public:
21 const Tensor v_cache, const Tensor block_tables,
22 const Tensor seq_lens,
23 const Tensor cum_seq_lens_q,
24 std::optional<Tensor> alibi_slopes, float scale,
25 Tensor out)
26 : dtype_{q.dtype()},
27 index_dtype_{block_tables.dtype()},
28 scale_{scale},
29 num_seqs_{seq_lens.size(0)},
30 total_q_tokens_{q.size(0)},
31 num_heads_{q.size(1)},
32 num_kv_heads_{k_cache.size(1)},
33 head_size_{q.size(2)},
34 block_size_{k_cache.size(2)},
35 max_num_blocks_per_seq_{block_tables.size(1)},
36 q_stride_{q.stride(0)},
37 q_head_stride_{q.stride(1)},
38 k_cache_block_stride_{k_cache.stride(0)},
39 k_cache_head_stride_{k_cache.stride(1)},
40 k_cache_slot_stride_{k_cache.stride(2)},
41 v_cache_block_stride_{v_cache.stride(0)},
42 v_cache_head_stride_{v_cache.stride(1)},
43 v_cache_slot_stride_{v_cache.stride(2)},
44 out_stride_{out.stride(0)},
45 out_head_stride_{out.stride(1)},
46 block_table_batch_stride_{block_tables.stride(0)} {
47 assert(q.ndim() == 3 && out.ndim() == 3);
48 assert(k_cache.ndim() == 4 && v_cache.ndim() == 4);
49 assert(block_tables.ndim() == 2 && seq_lens.ndim() == 1 &&
50 cum_seq_lens_q.ndim() == 1);
51 assert((dtype_ == DataType::kFloat16 || dtype_ == DataType::kBFloat16) &&
52 "`PagedAttentionPrefillInfinilm` supports float16 and bfloat16");
53 assert(out.dtype() == dtype_ && k_cache.dtype() == dtype_ &&
54 v_cache.dtype() == dtype_);
55 assert(IsIndexDtype(index_dtype_) && seq_lens.dtype() == index_dtype_ &&
56 cum_seq_lens_q.dtype() == index_dtype_);
57 assert(cum_seq_lens_q.size(0) == num_seqs_ + 1);
58 assert(q.shape() == out.shape());
59 assert(k_cache.shape() == v_cache.shape());
60 assert(block_tables.size(0) == num_seqs_);
61 assert(k_cache.size(1) == num_kv_heads_ &&
62 v_cache.size(1) == num_kv_heads_);
63 assert(k_cache.size(3) == head_size_ && v_cache.size(3) == head_size_);
64 assert((head_size_ == 64 || head_size_ == 128) &&
65 "`PagedAttentionPrefillInfinilm` supports head sizes 64 and 128");
66 assert(num_heads_ % num_kv_heads_ == 0);
67 assert(q.stride(2) == 1 && out.stride(2) == 1);
68 assert(k_cache.stride(3) == 1 && v_cache.stride(3) == 1);
69 assert(!alibi_slopes.has_value() ||
70 (alibi_slopes->dtype() == DataType::kFloat32 &&
71 alibi_slopes->ndim() == 1 && alibi_slopes->size(0) == num_heads_ &&
72 alibi_slopes->stride(0) == 1));
73 }
74
75 virtual void operator()(const Tensor q, const Tensor k_cache,
76 const Tensor v_cache, const Tensor block_tables,
77 const Tensor seq_lens, const Tensor cum_seq_lens_q,
78 std::optional<Tensor> alibi_slopes, float scale,
79 Tensor out) const = 0;
80
81 protected:
82 static bool IsIndexDtype(DataType dtype) {
83 return dtype == DataType::kInt32 || dtype == DataType::kInt64 ||
84 dtype == DataType::kUInt32;
85 }
86
87 DataType dtype_;
88
89 DataType index_dtype_;
90
91 float scale_{1.0f};
92
93 std::size_t num_seqs_{0};
94
95 std::size_t total_q_tokens_{0};
96
97 std::size_t num_heads_{0};
98
99 std::size_t num_kv_heads_{0};
100
101 std::size_t head_size_{0};
102
103 std::size_t block_size_{0};
104
105 std::size_t max_num_blocks_per_seq_{0};
106
107 Tensor::Stride q_stride_{0};
108
109 Tensor::Stride q_head_stride_{0};
110
111 Tensor::Stride k_cache_block_stride_{0};
112
113 Tensor::Stride k_cache_head_stride_{0};
114
115 Tensor::Stride k_cache_slot_stride_{0};
116
117 Tensor::Stride v_cache_block_stride_{0};
118
119 Tensor::Stride v_cache_head_stride_{0};
120
121 Tensor::Stride v_cache_slot_stride_{0};
122
123 Tensor::Stride out_stride_{0};
124
125 Tensor::Stride out_head_stride_{0};
126
127 Tensor::Stride block_table_batch_stride_{0};
128};
129
130} // namespace infini::ops
131
132#endif
Definition generated/include/operator.h:282
Definition paged_attention_prefill_infinilm.h:18
PagedAttentionPrefillInfinilm(const Tensor q, const Tensor k_cache, const Tensor v_cache, const Tensor block_tables, const Tensor seq_lens, const Tensor cum_seq_lens_q, std::optional< Tensor > alibi_slopes, float scale, Tensor out)
Definition paged_attention_prefill_infinilm.h:20
virtual void operator()(const Tensor q, const Tensor k_cache, const Tensor v_cache, const Tensor block_tables, const Tensor seq_lens, const Tensor cum_seq_lens_q, std::optional< Tensor > alibi_slopes, float scale, Tensor out) const =0
DataType dtype_
Definition paged_attention_prefill_infinilm.h:87
static bool IsIndexDtype(DataType dtype)
Definition paged_attention_prefill_infinilm.h:82
DataType index_dtype_
Definition paged_attention_prefill_infinilm.h:89
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8