23 std::optional<Tensor> alibi_slopes,
float scale,
26 index_dtype_{block_tables.dtype()},
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));