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