diff --git a/src/runtime/relax_vm/kv_cache.h b/src/runtime/relax_vm/kv_cache.h index e42c444d7a8e..4f4b538cb3b4 100644 --- a/src/runtime/relax_vm/kv_cache.h +++ b/src/runtime/relax_vm/kv_cache.h @@ -137,6 +137,19 @@ class AttentionKVCache : public Object { virtual void Attention(int64_t layer_id, NDArray q_data, NDArray k_data, NDArray v_data, Optional mask, NDArray o_data) = 0; + /*! + * \brief Compute attention with Q/K/V data which are concatenated along + * the head dimension. + * \param layer_id The model layer where the attention compute happens. + * \param qkv_data The input Q/K/V data, in layout + * `(total_length, num_qo_heads + 2 * num_kv_heads, head_dim)`. + * \param mask The input mask data, in layout `(total_sqr_length)`. + * \param o_data The output O data, in layout `(total_length, num_qo_heads, head_dim)`. + * \sa AttentionKVCache::Attention + */ + virtual void AttentionWithFusedQKV(int64_t layer_id, NDArray qkv_data, Optional mask, + NDArray o_data) = 0; + /************** Debug Helpers **************/ /*! diff --git a/src/runtime/relax_vm/paged_kv_cache.cc b/src/runtime/relax_vm/paged_kv_cache.cc index 20e68a9d3300..8e126b057f4e 100644 --- a/src/runtime/relax_vm/paged_kv_cache.cc +++ b/src/runtime/relax_vm/paged_kv_cache.cc @@ -148,6 +148,17 @@ struct Sequence { } }; +/*! + * \brief The rotary embedding mode adopted by the paged KV cache + * when computing attention. + * "Normal" means RoPE is computed in a standalone kernel. + * "Inline" means RoPE is computed on-the-fly in attention kernels. + */ +enum class RoPEMode : int { + kNormal = 0, + kInline = 1, +}; + /*! * \brief The paged KV cache for attention. * - It supports managing the K/V data of **multiple sequences**. @@ -184,7 +195,11 @@ class PagedAttentionKVCacheObj : public AttentionKVCache { const int64_t head_dim_; /*! \brief The number of total pages allocated in KV cache. */ const int64_t num_total_pages_; + /*! \brief The maximum total sequence length in a prefill. */ + const int64_t prefill_chunk_size_; + /*! \brief The RoPE application mode of KV cache.*/ + const RoPEMode rope_mode_; /*! \brief The RoPE scale. */ const double rotary_scale_; /*! \brief The RoPE theta. */ @@ -258,6 +273,9 @@ class PagedAttentionKVCacheObj : public AttentionKVCache { NDArray append_position_map_device_; // Temporary arrays to store intermediate attention results. + NDArray temp_attn_q_device_; + NDArray temp_attn_k_device_; + NDArray temp_attn_v_device_; NDArray temp_attn_output_device_; NDArray temp_attn_scores_device_; NDArray merged_attn_scores_device_; @@ -293,10 +311,14 @@ class PagedAttentionKVCacheObj : public AttentionKVCache { Optional f_attention_decode_begin_forward_; Optional f_attention_decode_end_forward_; Optional f_merge_inplace_; + PackedFunc f_split_rotary_; + PackedFunc f_rotary_inplace_; Optional f_debug_get_kv_; /*! \brief Number of fork depth in the current round of forward. */ int num_depths_; + /*! \brief Whether to compute attention after appending KV into cache or not. */ + bool append_before_attn_; /*! \brief Whether to use decode kernel for each depth. (see GetChunkedBlockIds) */ std::vector use_decode_kernel_; /*! \brief Whether the attention request is a decode request, set in BeginForwardFunction. */ @@ -304,29 +326,29 @@ class PagedAttentionKVCacheObj : public AttentionKVCache { public: /*! \brief Constructor. Take the cache configuration and initialize the NDArrays. */ - explicit PagedAttentionKVCacheObj(int64_t page_size, // - int64_t num_layers, int64_t num_qo_heads, int64_t num_kv_heads, - int64_t head_dim, // - int64_t reserved_num_seqs, int64_t num_total_pages, // - double rotary_scale, double rotary_theta, // - DLDataType dtype, DLDevice device, - PackedFunc f_transpose_append, PackedFunc f_attention_prefill, - PackedFunc f_attention_decode, - Optional f_attention_prefill_ragged, - Optional f_attention_prefill_ragged_begin_forward, - Optional f_attention_prefill_ragged_end_forward, - Optional f_attention_prefill_begin_forward, - Optional f_attention_prefill_end_forward, - Optional f_attention_decode_begin_forward, - Optional f_attention_decode_end_forward, - Optional f_merge_inplace, - Optional f_debug_get_kv) + explicit PagedAttentionKVCacheObj( + int64_t page_size, // + int64_t num_layers, int64_t num_qo_heads, int64_t num_kv_heads, int64_t head_dim, + int64_t reserved_num_seqs, int64_t num_total_pages, int64_t prefill_chunk_size, // + RoPEMode rope_mode, double rotary_scale, double rotary_theta, // + DLDataType dtype, DLDevice device, PackedFunc f_transpose_append, + PackedFunc f_attention_prefill, PackedFunc f_attention_decode, + Optional f_attention_prefill_ragged, + Optional f_attention_prefill_ragged_begin_forward, + Optional f_attention_prefill_ragged_end_forward, + Optional f_attention_prefill_begin_forward, + Optional f_attention_prefill_end_forward, + Optional f_attention_decode_begin_forward, + Optional f_attention_decode_end_forward, Optional f_merge_inplace, + PackedFunc f_split_rotary, PackedFunc f_rotary_inplace, Optional f_debug_get_kv) : page_size_(page_size), num_layers_(num_layers), num_qo_heads_(num_qo_heads), num_kv_heads_(num_kv_heads), head_dim_(head_dim), num_total_pages_(num_total_pages), + prefill_chunk_size_(prefill_chunk_size), + rope_mode_(rope_mode), rotary_scale_(rotary_scale), rotary_theta_(rotary_theta), f_transpose_append_(std::move(f_transpose_append)), @@ -341,6 +363,8 @@ class PagedAttentionKVCacheObj : public AttentionKVCache { f_attention_decode_begin_forward_(std::move(f_attention_decode_begin_forward)), f_attention_decode_end_forward_(std::move(f_attention_decode_end_forward)), f_merge_inplace_(std::move(f_merge_inplace)), + f_split_rotary_(std::move(f_split_rotary)), + f_rotary_inplace_(std::move(f_rotary_inplace)), f_debug_get_kv_(std::move(f_debug_get_kv)) { pages_.reserve(num_layers); for (int i = 0; i < num_layers; ++i) { @@ -365,15 +389,21 @@ class PagedAttentionKVCacheObj : public AttentionKVCache { } cur_append_length_indptr_device_ = NDArray::Empty({reserved_num_seqs + 1}, dtype_aux_, device); k_ragged_rope_pos_offset_device_ = NDArray::Empty({reserved_num_seqs}, dtype_aux_, device); - q_rope_position_map_device_ = NDArray::Empty({num_total_pages * page_size}, dtype_aux_, device); - append_position_map_device_ = NDArray::Empty({num_total_pages * page_size}, dtype_aux_, device); - + q_rope_position_map_device_ = NDArray::Empty({prefill_chunk_size_}, dtype_aux_, device); + append_position_map_device_ = NDArray::Empty({prefill_chunk_size_}, dtype_aux_, device); + + temp_attn_q_device_ = + NDArray::Empty({prefill_chunk_size_, num_qo_heads, head_dim}, dtype, device); + temp_attn_k_device_ = + NDArray::Empty({prefill_chunk_size_, num_kv_heads, head_dim}, dtype, device); + temp_attn_v_device_ = + NDArray::Empty({prefill_chunk_size_, num_kv_heads, head_dim}, dtype, device); temp_attn_output_device_ = - NDArray::Empty({num_total_pages * page_size, num_qo_heads, head_dim}, dtype, device); + NDArray::Empty({prefill_chunk_size_, num_qo_heads, head_dim}, dtype, device); temp_attn_scores_device_ = - NDArray::Empty({num_total_pages * page_size, num_qo_heads}, DataType::Float(32), device); + NDArray::Empty({prefill_chunk_size_, num_qo_heads}, DataType::Float(32), device); merged_attn_scores_device_ = - NDArray::Empty({num_total_pages * page_size, num_qo_heads}, DataType::Float(32), device); + NDArray::Empty({prefill_chunk_size_, num_qo_heads}, DataType::Float(32), device); for (int64_t page_id = num_total_pages - 1; page_id >= 0; --page_id) { free_page_ids_.push_back(page_id); } @@ -501,7 +531,17 @@ class PagedAttentionKVCacheObj : public AttentionKVCache { num_depths_ = block_ids_on_depths.size(); ICHECK_LE(num_depths_, kPagedKVCacheMaxBlockDepth); - if (num_depths_ == 1) { + std::vector>> chunked_block_ids_arr; + chunked_block_ids_arr.reserve(num_depths_); + use_decode_kernel_.clear(); + for (int d = 0; d < num_depths_; ++d) { + auto [chunked_block_ids, use_decode_kernel] = GetChunkedBlockIds(block_ids_on_depths[d]); + chunked_block_ids_arr.push_back(chunked_block_ids); + use_decode_kernel_.push_back(use_decode_kernel); + } + + append_before_attn_ = num_depths_ == 1 && use_decode_kernel_[0]; + if (append_before_attn_) { // Right now we use different kernels when depth is 1 or not 1. // For the case where maximum depth is 1, we create the auxiliary // data structure with regard to the page table after appending. @@ -515,17 +555,14 @@ class PagedAttentionKVCacheObj : public AttentionKVCache { std::vector> page_indices_on_depths; std::vector> last_page_len_on_depths; std::vector> k_rope_pos_offset_on_depths; - use_decode_kernel_.clear(); - for (int d = 0; d < num_depths_; ++d) { - auto [chunked_block_ids, use_decode_kernel] = GetChunkedBlockIds(block_ids_on_depths[d]); - use_decode_kernel_.push_back(use_decode_kernel); + for (int d = 0; d < num_depths_; ++d) { std::vector qo_indptr_h{0}; std::vector page_indptr_h{0}; std::vector page_indices_h; std::vector last_page_len_h; std::vector k_rope_pos_offset_h; - for (const auto& [block_id, chunk_append_length] : chunked_block_ids) { + for (const auto& [block_id, chunk_append_length] : chunked_block_ids_arr[d]) { qo_indptr_h.push_back(qo_indptr_h.back() + chunk_append_length); if (block_id == -1) { page_indptr_h.push_back(page_indptr_h.back()); @@ -547,7 +584,7 @@ class PagedAttentionKVCacheObj : public AttentionKVCache { k_rope_pos_offset_on_depths.push_back(k_rope_pos_offset_h); } - if (num_depths_ > 1) { + if (!append_before_attn_) { // Right now we use different kernels when depth is 1 or not 1. // For the case where maximum depth is not 1, we create the auxiliary // data structure with regard to the page table before appending. @@ -606,21 +643,19 @@ class PagedAttentionKVCacheObj : public AttentionKVCache { CHECK(v_data.DataType() == pages.DataType()); CHECK(o_data.DataType() == pages.DataType()); - // Case 1. q/o_data: (cur_batch_size, 1, num_qo_heads, head_dim) - // k/v_data: (cur_batch_size, 1, num_kv_heads, head_dim) - // Case 2. q/o_data: (1, num_total_length, num_qo_heads, head_dim) - // k/v_data: (1, num_total_length, num_kv_heads, head_dim) - - CHECK_EQ(q_data->ndim, 4); - CHECK_EQ(k_data->ndim, 4); - CHECK_EQ(v_data->ndim, 4); - CHECK_EQ(o_data->ndim, 4); - for (int dim = 0; dim < 4; ++dim) { - if (dim == 2) { - CHECK_EQ(q_data->shape[2], num_qo_heads_); - CHECK_EQ(k_data->shape[2], num_kv_heads_); - CHECK_EQ(v_data->shape[2], num_kv_heads_); - CHECK_EQ(o_data->shape[2], num_qo_heads_); + // q/o_data: (num_total_length, num_qo_heads, head_dim) + // k/v_data: (num_total_length, num_kv_heads, head_dim) + + CHECK_EQ(q_data->ndim, 3); + CHECK_EQ(k_data->ndim, 3); + CHECK_EQ(v_data->ndim, 3); + CHECK_EQ(o_data->ndim, 3); + for (int dim = 0; dim < 3; ++dim) { + if (dim == 1) { + CHECK_EQ(q_data->shape[1], num_qo_heads_); + CHECK_EQ(k_data->shape[1], num_kv_heads_); + CHECK_EQ(v_data->shape[1], num_kv_heads_); + CHECK_EQ(o_data->shape[1], num_qo_heads_); } else { CHECK_EQ(k_data->shape[dim], q_data->shape[dim]); CHECK_EQ(v_data->shape[dim], q_data->shape[dim]); @@ -628,43 +663,76 @@ class PagedAttentionKVCacheObj : public AttentionKVCache { } } - CHECK_GT(q_data->shape[1], 0); - CHECK_EQ(q_data->shape[3], head_dim_); + CHECK_GT(q_data->shape[0], 0); + CHECK_EQ(q_data->shape[2], head_dim_); + int64_t total_seq_length = 0; + for (int64_t seq_id = 0; seq_id < cur_batch_size_; ++seq_id) { + total_seq_length += cur_append_lengths_[seq_id]; + } + CHECK_EQ(total_seq_length, q_data->shape[0]); + // The auxiliary data structure on device must have been synchronized. + CHECK(!dirty_aux_data_device_) + << "The auxiliary arrays are not synchronized to device. Please call " + "`BeginForward` to synchronize before calling `Attention`."; - if (q_data->shape[0] > 1) { - // Case 1. - CHECK_EQ(q_data->shape[0], cur_batch_size_); - CHECK_EQ(q_data->shape[1], 1); - } else { - // Case 2. - CHECK_EQ(q_data->shape[0], 1); + if (rope_mode_ == RoPEMode::kNormal) { + // Apply rotary embedding to q/k data. + f_rotary_inplace_(q_data, k_data, cur_append_length_indptr_view_, + k_ragged_rope_pos_offset_view_, cur_batch_size_, num_qo_heads_, + num_kv_heads_, head_dim_, /*qkv_layout=*/0, rotary_scale_, rotary_theta_); + } + + // Part 3: append k/v data to kv-cache + f_transpose_append_(pages_[layer_id], k_data, v_data, append_position_map_view_); + // Part 4: perform attention + AttentionInternal(layer_id, q_data, k_data, v_data, o_data); + } + + void AttentionWithFusedQKV(int64_t layer_id, NDArray qkv_data, Optional mask, + NDArray o_data) final { + // Part 1. Shape and dtype check. + NDArray pages = pages_[layer_id]; + CHECK(qkv_data.DataType() == pages.DataType()); + CHECK(o_data.DataType() == pages.DataType()); + + // qkv_data: (num_total_length, num_qo_heads + 2 * num_kv_heads, head_dim) + // o_data: (num_total_length, num_qo_heads, head_dim) + + CHECK_EQ(qkv_data->ndim, 3); + CHECK_EQ(o_data->ndim, 3); + for (int dim = 0; dim < 3; ++dim) { + if (dim == 1) { + CHECK_EQ(qkv_data->shape[1], num_qo_heads_ + 2 * num_kv_heads_); + CHECK_EQ(o_data->shape[1], num_qo_heads_); + } else { + CHECK_EQ(o_data->shape[dim], qkv_data->shape[dim]); + } } + CHECK_EQ(qkv_data->shape[2], head_dim_); int64_t total_seq_length = 0; for (int64_t seq_id = 0; seq_id < cur_batch_size_; ++seq_id) { total_seq_length += cur_append_lengths_[seq_id]; - if (q_data->shape[0] > 1) { - CHECK_EQ(cur_append_lengths_[seq_id], 1); - } } - CHECK_EQ(total_seq_length, q_data->shape[0] * q_data->shape[1]); - q_data = - q_data.CreateView({total_seq_length, q_data->shape[2], q_data->shape[3]}, q_data->dtype); - k_data = - k_data.CreateView({total_seq_length, k_data->shape[2], k_data->shape[3]}, k_data->dtype); - v_data = - v_data.CreateView({total_seq_length, v_data->shape[2], v_data->shape[3]}, v_data->dtype); - o_data = - o_data.CreateView({total_seq_length, o_data->shape[2], o_data->shape[3]}, o_data->dtype); - + CHECK_EQ(total_seq_length, qkv_data->shape[0]); // The auxiliary data structure on device must have been synchronized. CHECK(!dirty_aux_data_device_) << "The auxiliary arrays are not synchronized to device. Please call " "`BeginForward` to synchronize before calling `Attention`."; - // Part 2: append k/v data to kv-cache + NDArray q_data = temp_attn_q_device_.CreateView({total_seq_length, num_qo_heads_, head_dim_}, + qkv_data->dtype); + NDArray k_data = temp_attn_k_device_.CreateView({total_seq_length, num_kv_heads_, head_dim_}, + qkv_data->dtype); + NDArray v_data = temp_attn_v_device_.CreateView({total_seq_length, num_kv_heads_, head_dim_}, + qkv_data->dtype); + // Part 2. Split fused qkv and apply rotary embedding to q/k data. + f_split_rotary_(qkv_data, q_rope_position_map_view_, q_data, k_data, v_data, + rope_mode_ == RoPEMode::kNormal); + + // Part 3: append k/v data to kv-cache f_transpose_append_(pages_[layer_id], k_data, v_data, append_position_map_view_); - // Part 3: perform attention + // Part 4: perform attention AttentionInternal(layer_id, q_data, k_data, v_data, o_data); } @@ -875,15 +943,11 @@ class PagedAttentionKVCacheObj : public AttentionKVCache { return; } - if (num_depths_ == 1) { - if (use_decode_kernel_[0]) { - f_attention_decode_begin_forward_.value()( - /*depth=*/0, page_indptr_on_depths_view_[0], last_page_len_on_depths_view_[0], - num_qo_heads_, num_kv_heads_, head_dim_, page_size_, /*rotary_mode=*/1); - } else { - f_attention_prefill_begin_forward_.value()(/*depth=*/0, qo_indptr_on_depths_view_[0], - cur_batch_size_, num_qo_heads_, num_kv_heads_); - } + if (append_before_attn_) { + f_attention_decode_begin_forward_.value()( + /*depth=*/0, page_indptr_on_depths_view_[0], last_page_len_on_depths_view_[0], + num_qo_heads_, num_kv_heads_, head_dim_, page_size_, + /*rotary_mode=*/rope_mode_ == RoPEMode::kInline); } else { f_attention_prefill_ragged_begin_forward_.value()( cur_append_length_indptr_view_, cur_batch_size_, num_qo_heads_, num_kv_heads_); @@ -894,7 +958,8 @@ class PagedAttentionKVCacheObj : public AttentionKVCache { if (use_decode_kernel_[d]) { f_attention_decode_begin_forward_.value()( d, page_indptr_on_depths_view_[d], last_page_len_on_depths_view_[d], num_qo_heads_, - num_kv_heads_, head_dim_, page_size_, /*rotary_mode=*/1); + num_kv_heads_, head_dim_, page_size_, + /*rotary_mode=*/rope_mode_ == RoPEMode::kInline); } else { f_attention_prefill_begin_forward_.value()(/*depth=*/d, qo_indptr_on_depths_view_[d], last_page_len_on_depths_view_[d]->shape[0], @@ -911,29 +976,20 @@ class PagedAttentionKVCacheObj : public AttentionKVCache { void AttentionInternal(int64_t layer_id, NDArray q_data, NDArray k_data, NDArray v_data, NDArray output) { CHECK_GE(num_depths_, 1) << "The number of effective depths must be greater or equal to 1."; - if (num_depths_ == 1) { - if (use_decode_kernel_[0]) { - f_attention_decode_(/*depth=*/0, q_data, pages_[layer_id], page_indptr_on_depths_view_[0], - page_indices_on_depths_view_[0], last_page_len_on_depths_view_[0], - k_rope_pos_offset_view_[0], q_rope_position_map_view_, output, - merged_attn_scores_view_, - /*rotary_mode=*/1, rotary_scale_, rotary_theta_); - } else { - f_attention_prefill_(/*depth=*/0, q_data, qo_indptr_on_depths_view_[0], pages_[layer_id], - page_indptr_on_depths_view_[0], page_indices_on_depths_view_[0], - last_page_len_on_depths_view_[0], k_rope_pos_offset_view_[0], - q_rope_position_map_view_, output, merged_attn_scores_view_, - /*causal=*/1, - /*rotary_mode=*/1, rotary_scale_, rotary_theta_); - } + if (append_before_attn_) { + f_attention_decode_( + /*depth=*/0, q_data, pages_[layer_id], page_indptr_on_depths_view_[0], + page_indices_on_depths_view_[0], last_page_len_on_depths_view_[0], + k_rope_pos_offset_view_[0], q_rope_position_map_view_, output, merged_attn_scores_view_, + /*rotary_mode=*/rope_mode_ == RoPEMode::kInline, rotary_scale_, rotary_theta_); } else { // Compute appended text self-attention - f_attention_prefill_ragged_.value()(q_data, cur_append_length_indptr_view_, k_data, v_data, - cur_append_length_indptr_view_, q_rope_position_map_view_, - k_ragged_rope_pos_offset_view_, output, - merged_attn_scores_view_, - /*causal=*/1, - /*rotary_mode=*/1, rotary_scale_, rotary_theta_); + f_attention_prefill_ragged_.value()( + q_data, cur_append_length_indptr_view_, k_data, v_data, cur_append_length_indptr_view_, + q_rope_position_map_view_, k_ragged_rope_pos_offset_view_, output, + merged_attn_scores_view_, + /*causal=*/1, + /*rotary_mode=*/rope_mode_ == RoPEMode::kInline, rotary_scale_, rotary_theta_); for (int d = 0; d < num_depths_; ++d) { if (page_indices_on_depths_view_[d]->shape[0] == 0) { @@ -945,16 +1001,17 @@ class PagedAttentionKVCacheObj : public AttentionKVCache { page_indices_on_depths_view_[d], last_page_len_on_depths_view_[d], k_rope_pos_offset_view_[d], q_rope_position_map_view_, temp_attn_output_view_, temp_attn_scores_view_, - /*rotary_mode=*/1, rotary_scale_, rotary_theta_); + /*rotary_mode=*/rope_mode_ == RoPEMode::kInline, rotary_scale_, + rotary_theta_); } else { // Use prefill kernel for depth d - f_attention_prefill_(/*depth=*/d, q_data, qo_indptr_on_depths_view_[d], pages_[layer_id], - page_indptr_on_depths_view_[d], page_indices_on_depths_view_[d], - last_page_len_on_depths_view_[d], k_rope_pos_offset_view_[d], - q_rope_position_map_view_, temp_attn_output_view_, - temp_attn_scores_view_, - /*causal=*/0, - /*rotary_mode=*/1, rotary_scale_, rotary_theta_); + f_attention_prefill_( + /*depth=*/d, q_data, qo_indptr_on_depths_view_[d], pages_[layer_id], + page_indptr_on_depths_view_[d], page_indices_on_depths_view_[d], + last_page_len_on_depths_view_[d], k_rope_pos_offset_view_[d], + q_rope_position_map_view_, temp_attn_output_view_, temp_attn_scores_view_, + /*causal=*/0, + /*rotary_mode=*/rope_mode_ == RoPEMode::kInline, rotary_scale_, rotary_theta_); } f_merge_inplace_.value()(output, merged_attn_scores_view_, temp_attn_output_view_, temp_attn_scores_view_); @@ -1092,7 +1149,7 @@ TVM_REGISTER_OBJECT_TYPE(PagedAttentionKVCacheObj); TVM_REGISTER_GLOBAL("vm.builtin.paged_attention_kv_cache_create") .set_body_typed([](ShapeTuple cache_config, int64_t num_layers, int64_t num_qo_heads, - int64_t num_kv_heads, int64_t head_dim, double rotary_scale, + int64_t num_kv_heads, int64_t head_dim, int rope_mode, double rotary_scale, double rotary_theta, NDArray init, PackedFunc f_transpose_append, PackedFunc f_attention_prefill, PackedFunc f_attention_decode, PackedFunc f_attention_prefill_ragged, @@ -1102,44 +1159,50 @@ TVM_REGISTER_GLOBAL("vm.builtin.paged_attention_kv_cache_create") PackedFunc f_attention_prefill_end_forward, PackedFunc f_attention_decode_begin_forward, PackedFunc f_attention_decode_end_forward, PackedFunc f_merge_inplace, + PackedFunc f_split_rotary, PackedFunc f_rotary_inplace, Optional f_debug_get_kv) { - CHECK_EQ(cache_config.size(), 3); + CHECK_EQ(cache_config.size(), 4); int64_t reserved_num_seqs = cache_config[0]; int64_t total_token_capacity = cache_config[1]; - int64_t page_size = cache_config[2]; + int64_t prefill_chunk_size = cache_config[2]; + int64_t page_size = cache_config[3]; int64_t num_total_pages = (total_token_capacity + page_size - 1) / page_size; ObjectPtr n = make_object( page_size, num_layers, num_qo_heads, num_kv_heads, head_dim, reserved_num_seqs, - num_total_pages, rotary_scale, rotary_theta, init->dtype, init->device, - std::move(f_transpose_append), std::move(f_attention_prefill), + num_total_pages, prefill_chunk_size, RoPEMode(rope_mode), rotary_scale, rotary_theta, + init->dtype, init->device, std::move(f_transpose_append), std::move(f_attention_prefill), std::move(f_attention_decode), std::move(f_attention_prefill_ragged), std::move(f_attention_prefill_ragged_begin_forward), std::move(f_attention_prefill_ragged_end_forward), std::move(f_attention_prefill_begin_forward), std::move(f_attention_prefill_end_forward), std::move(f_attention_decode_begin_forward), std::move(f_attention_decode_end_forward), - std::move(f_merge_inplace), std::move(f_debug_get_kv)); + std::move(f_merge_inplace), std::move(f_split_rotary), std::move(f_rotary_inplace), + std::move(f_debug_get_kv)); return PagedAttentionKVCache(std::move(n)); }); TVM_REGISTER_GLOBAL("vm.builtin.paged_attention_kv_cache_create_reduced") .set_body_typed([](ShapeTuple cache_config, int64_t num_layers, int64_t num_qo_heads, - int64_t num_kv_heads, int64_t head_dim, double rotary_scale, + int64_t num_kv_heads, int64_t head_dim, int rope_mode, double rotary_scale, double rotary_theta, NDArray init, PackedFunc f_transpose_append, PackedFunc f_attention_prefill, PackedFunc f_attention_decode, PackedFunc f_attention_prefill_ragged, PackedFunc f_merge_inplace, + PackedFunc f_split_rotary, PackedFunc f_rotary_inplace, Optional f_debug_get_kv) { - CHECK_EQ(cache_config.size(), 3); + CHECK_EQ(cache_config.size(), 4); int64_t reserved_num_seqs = cache_config[0]; int64_t total_token_capacity = cache_config[1]; - int64_t page_size = cache_config[2]; + int64_t prefill_chunk_size = cache_config[2]; + int64_t page_size = cache_config[3]; int64_t num_total_pages = (total_token_capacity + page_size - 1) / page_size; ObjectPtr n = make_object( page_size, num_layers, num_qo_heads, num_kv_heads, head_dim, reserved_num_seqs, - num_total_pages, rotary_scale, rotary_theta, init->dtype, init->device, - std::move(f_transpose_append), std::move(f_attention_prefill), + num_total_pages, prefill_chunk_size, RoPEMode(rope_mode), rotary_scale, rotary_theta, + init->dtype, init->device, std::move(f_transpose_append), std::move(f_attention_prefill), std::move(f_attention_decode), std::move(f_attention_prefill_ragged), // NullOpt, NullOpt, NullOpt, NullOpt, NullOpt, NullOpt, // - std::move(f_merge_inplace), std::move(f_debug_get_kv)); + std::move(f_merge_inplace), std::move(f_split_rotary), std::move(f_rotary_inplace), + std::move(f_debug_get_kv)); return PagedAttentionKVCache(std::move(n)); }); @@ -1167,6 +1230,11 @@ TVM_REGISTER_GLOBAL("vm.builtin.paged_attention_kv_cache_attention") kv_cache->Attention(layer_id, std::move(q_data), std::move(k_data), std::move(v_data), NullOpt, std::move(o_data)); }); +TVM_REGISTER_GLOBAL("vm.builtin.paged_attention_kv_cache_attention_with_fused_qkv") + .set_body_typed([](PagedAttentionKVCache kv_cache, int64_t layer_id, NDArray qkv_data, + NDArray o_data) { + kv_cache->AttentionWithFusedQKV(layer_id, std::move(qkv_data), NullOpt, std::move(o_data)); + }); } // namespace relax_vm } // namespace runtime diff --git a/tests/python/relax/test_runtime_builtin_paged_attention_kv_cache_flashinfer.py b/tests/python/relax/test_runtime_builtin_paged_attention_kv_cache_flashinfer.py index 69b7a15793b5..8a40f43ab760 100644 --- a/tests/python/relax/test_runtime_builtin_paged_attention_kv_cache_flashinfer.py +++ b/tests/python/relax/test_runtime_builtin_paged_attention_kv_cache_flashinfer.py @@ -23,11 +23,13 @@ import tvm import tvm.testing from tvm import dlight as dl +from tvm import tir from tvm.runtime import ShapeTuple from tvm.script import tir as T reserved_nseq = 32 maximum_total_seq_length = 1024 +prefill_chunk_size = 512 page_size = 16 num_layers = 4 num_qo_heads = 32 @@ -47,6 +49,7 @@ fbegin_forward = None fend_forward = None fattention = None +fattention_with_fuse_qkv = None fdebug_get_kv = None fattention_prefill = None @@ -59,6 +62,11 @@ fattention_prefill_ragged_begin_forward = None fattention_prefill_ragged_end_forward = None fattention_merge_state = None +fattention_rotary = None + +ftranspose_append = None +fsplit_rotary = None +fcopy_cache = None @T.prim_func @@ -100,6 +108,84 @@ def kv_cache_transpose_append( ] = v_data[vgpos, vh, vf] +def llama_rope_with_position_map( # pylint: disable=too-many-arguments + theta: float, + scale: float, + head_dim: int, + num_q_heads: int, + num_kv_heads: int, + dtype: float = "float16", + rotary_dim: int = None, +): + fused_heads = num_q_heads + num_kv_heads * 2 + if rotary_dim is None: + rotary_dim = head_dim + scale = tir.const(scale, dtype) + + def _rope_freq(s: tir.Var, d: tir.Var, d_range: int, theta: float, dtype: str): + freq = s / tir.power(theta, d * 2 % d_range / tir.const(d_range, "float32")) + cos_freq = tir.cos(freq).astype(dtype) + sin_freq = tir.sin(freq).astype(dtype) + return cos_freq, sin_freq + + def _rope( # pylint: disable=too-many-arguments + x: T.Buffer, + s: tir.Var, + h: tir.Var, + d: tir.Var, + pos: tir.Var, + ): + cos_freq, sin_freq = _rope_freq(pos * scale, d, rotary_dim, theta, dtype) + cos = cos_freq * x[s, h, d] + sin = sin_freq * tir.if_then_else( + d < rotary_dim // 2, + -x[s, h, d + rotary_dim // 2], + x[s, h, d - rotary_dim // 2], + ) + return cos + sin + + @T.prim_func(private=True) + def fused_rope( # pylint: disable=too-many-locals + var_qkv: T.handle, + var_position_map: T.handle, + var_q: T.handle, + var_k: T.handle, + var_v: T.handle, + apply_rope: T.int32, + ): + T.func_attr( + { + "op_pattern": 8, # 2 means injective, 8 means opaque + "tir.noalias": T.bool(True), + } + ) + seq_len = T.int64() + qkv = T.match_buffer(var_qkv, (seq_len, fused_heads, head_dim), dtype) + q = T.match_buffer(var_q, (seq_len, num_q_heads, head_dim), dtype) + k = T.match_buffer(var_k, (seq_len, num_kv_heads, head_dim), dtype) + v = T.match_buffer(var_v, (seq_len, num_kv_heads, head_dim), dtype) + position_map = T.match_buffer(var_position_map, (seq_len,), "int32") + for iters in T.grid(seq_len, fused_heads, head_dim): + with T.block("llama_fused_rope"): + s, h, d = T.axis.remap("SSS", iters) + if h < num_q_heads: + q[s, h, d] = T.if_then_else( + apply_rope > 0 and d < rotary_dim, + _rope(qkv, s, h, d, position_map[s]), + qkv[s, h, d], + ) + elif h < num_q_heads + num_kv_heads: + k[s, h - num_q_heads, d] = T.if_then_else( + apply_rope > 0 and d < rotary_dim, + _rope(qkv, s, h, d, position_map[s]), + qkv[s, h, d], + ) + else: + v[s, h - (num_q_heads + num_kv_heads), d] = qkv[s, h, d] + + return fused_rope + + @T.prim_func def copy_cache( var_pages: T.handle, @@ -138,13 +224,14 @@ def copy_cache( def set_global_func(): global fclear, fcreate, fadd_sequence, fremove_sequence, ffork_sequence, fpopn - global fbegin_forward, fend_forward, fattention, fdebug_get_kv + global fbegin_forward, fend_forward, fattention, fattention_with_fuse_qkv, fdebug_get_kv global fattention_prefill, fattention_prefill_begin_forward, fattention_prefill_end_forward global fattention_decode, fattention_decode_begin_forward, fattention_decode_end_forward global fattention_prefill_ragged global fattention_prefill_ragged_begin_forward global fattention_prefill_ragged_end_forward - global fattention_merge_state + global fattention_merge_state, fsplit_rotary, fattention_rotary + global ftranspose_append, fcopy_cache fclear = tvm.get_global_func("vm.builtin.paged_attention_kv_cache_clear") fcreate = tvm.get_global_func("vm.builtin.paged_attention_kv_cache_create") @@ -155,6 +242,9 @@ def set_global_func(): fbegin_forward = tvm.get_global_func("vm.builtin.paged_attention_kv_cache_begin_forward") fend_forward = tvm.get_global_func("vm.builtin.paged_attention_kv_cache_end_forward") fattention = tvm.get_global_func("vm.builtin.paged_attention_kv_cache_attention") + fattention_with_fuse_qkv = tvm.get_global_func( + "vm.builtin.paged_attention_kv_cache_attention_with_fused_qkv" + ) fdebug_get_kv = tvm.get_global_func("vm.builtin.paged_attention_kv_cache_debug_get_kv") fattention_prefill = tvm.get_global_func("paged_kv_cache.attention_kernel_prefill") @@ -181,26 +271,36 @@ def set_global_func(): "flashinfer.attention_kernel_prefill_with_ragged_kv_cache_end_forward" ) fattention_merge_state = tvm.get_global_func("flashinfer.merge_state_in_place") + fattention_rotary = tvm.get_global_func("flashinfer.batch_qk_apply_rotary_in_place") - -def create_kv_cache(): - set_global_func() target = tvm.target.Target("nvidia/geforce-rtx-3090-ti") builts = [] - for tir_func in [kv_cache_transpose_append, copy_cache]: + for tir_func in [ + kv_cache_transpose_append, + llama_rope_with_position_map( + rope_theta, rope_scale, head_dim, num_qo_heads, num_kv_heads, dtype + ), + copy_cache, + ]: mod = tvm.IRModule({"main": tir_func}) with target: mod = dl.ApplyDefaultSchedule(dl.gpu.Fallback())(mod) f = tvm.build(mod["main"], target=target) builts.append(f.entry_func) - ftranspose_append, fcopy_cache = builts + ftranspose_append, fsplit_rotary, fcopy_cache = builts + + +def create_kv_cache(rope_mode): cache = fcreate( - tvm.runtime.ShapeTuple([reserved_nseq, maximum_total_seq_length, page_size]), + tvm.runtime.ShapeTuple( + [reserved_nseq, maximum_total_seq_length, prefill_chunk_size, page_size] + ), num_layers, num_qo_heads, num_kv_heads, head_dim, + rope_mode, rope_scale, rope_theta, tvm.nd.empty((), dtype, device=device), @@ -215,14 +315,17 @@ def create_kv_cache(): fattention_decode_begin_forward, fattention_decode_end_forward, fattention_merge_state, + fsplit_rotary, + fattention_rotary, fcopy_cache, ) return cache -@pytest.fixture() -def kv_cache(): - return create_kv_cache() +@pytest.fixture(params=[0, 1]) +def kv_cache_and_rope_mode(request): + set_global_func() + return create_kv_cache(request.param), request.param def verify_cached_kv(kv_cache, seq_ids, expected_k, expected_v): @@ -258,9 +361,11 @@ def f_apply_rotary(x, offset, scale, theta): def apply_attention( kv_cache, + rope_mode: int, batch: List[Tuple[Union[int, Tuple[int, int]], int]], cached_k: Dict[int, np.ndarray], cached_v: Dict[int, np.ndarray], + fuse_qkv: bool, ) -> None: seq_ids = [] append_lengths = [] @@ -283,7 +388,6 @@ def apply_attention( cached_k[seq_id] = np.zeros((num_layers, 0, num_kv_heads, head_dim), dtype) cached_v[seq_id] = np.zeros((num_layers, 0, num_kv_heads, head_dim), dtype) - use_decode_shape = all(append_length == 1 for _, append_length in batch) fbegin_forward(kv_cache, ShapeTuple(seq_ids), ShapeTuple(append_lengths)) global_new_q = np.zeros((num_layers, 0, num_qo_heads, head_dim), dtype) @@ -300,7 +404,17 @@ def apply_attention( cached_k[seq_id] = np.concatenate( [ cached_k[seq_id], - np.stack([new_k[l] for l in range(num_layers)], axis=0), + np.stack( + [ + new_k[l] + if rope_mode == 1 + else f_apply_rotary( + new_k[l], cached_k[seq_id].shape[1], rope_scale, rope_theta + ) + for l in range(num_layers) + ], + axis=0, + ), ], axis=1, ) @@ -310,23 +424,22 @@ def apply_attention( global_new_v = np.concatenate([global_new_v, new_v], axis=1) for layer_id in range(num_layers): - queries_np = global_new_q[layer_id : layer_id + 1] - keys_np = global_new_k[layer_id : layer_id + 1] - values_np = global_new_v[layer_id : layer_id + 1] - if use_decode_shape: - queries_np = queries_np.transpose(1, 0, 2, 3) - keys_np = keys_np.transpose(1, 0, 2, 3) - values_np = values_np.transpose(1, 0, 2, 3) - queries = tvm.nd.array(queries_np, device=device) - keys = tvm.nd.array(keys_np, device=device) - values = tvm.nd.array(values_np, device=device) - outputs = tvm.nd.empty(queries.shape, dtype, device=device) - fattention(kv_cache, layer_id, queries, keys, values, outputs) + queries_np = global_new_q[layer_id] + keys_np = global_new_k[layer_id] + values_np = global_new_v[layer_id] + if not fuse_qkv: + queries = tvm.nd.array(queries_np, device=device) + keys = tvm.nd.array(keys_np, device=device) + values = tvm.nd.array(values_np, device=device) + outputs = tvm.nd.empty(queries.shape, dtype, device=device) + fattention(kv_cache, layer_id, queries, keys, values, outputs) + else: + qkv = tvm.nd.array(np.concatenate([queries_np, keys_np, values_np], axis=1), device) + outputs = tvm.nd.empty(queries_np.shape, dtype, device=device) + fattention_with_fuse_qkv(kv_cache, layer_id, qkv, outputs) # Compute attention expected results. - outputs = outputs.numpy() - if use_decode_shape: - outputs = outputs.transpose(1, 0, 2, 3) + outputs = np.expand_dims(outputs.numpy(), axis=0) sum_length = 0 for i, (seq_id, append_length) in enumerate(batch): assert cached_k[seq_id].shape[1] == cached_v[seq_id].shape[1] >= append_length @@ -338,9 +451,11 @@ def apply_attention( rope_scale, rope_theta, ).transpose(1, 0, 2) - k_seq = f_apply_rotary(cached_k[seq_id][layer_id], 0, rope_scale, rope_theta).transpose( - 1, 2, 0 - ) + k_seq = ( + cached_k[seq_id][layer_id] + if rope_mode == 0 + else f_apply_rotary(cached_k[seq_id][layer_id], 0, rope_scale, rope_theta) + ).transpose(1, 2, 0) v_seq = cached_v[seq_id][layer_id].transpose(1, 0, 2) k_seq = np.repeat(k_seq, num_qo_heads // num_kv_heads, axis=0) @@ -375,7 +490,9 @@ def apply_attention( @pytest.mark.skip(reason="Require FlashInfer enabled") -def test_paged_attention_kv_cache_prefill_and_decode(kv_cache): +@pytest.mark.parametrize("fuse_qkv", [False, True]) +def test_paged_attention_kv_cache_prefill_and_decode(kv_cache_and_rope_mode, fuse_qkv): + kv_cache, rope_mode = kv_cache_and_rope_mode fclear(kv_cache) # Prefill. @@ -391,11 +508,13 @@ def test_paged_attention_kv_cache_prefill_and_decode(kv_cache): cached_k = {} cached_v = {} for batch in operation_seq: - apply_attention(kv_cache, batch, cached_k, cached_v) + apply_attention(kv_cache, rope_mode, batch, cached_k, cached_v, fuse_qkv) @pytest.mark.skip(reason="Require FlashInfer enabled") -def test_paged_attention_kv_cache_remove_sequence(kv_cache): +@pytest.mark.parametrize("fuse_qkv", [False, True]) +def test_paged_attention_kv_cache_remove_sequence(kv_cache_and_rope_mode, fuse_qkv): + kv_cache, rope_mode = kv_cache_and_rope_mode fclear(kv_cache) num_sequences = 5 @@ -403,7 +522,7 @@ def test_paged_attention_kv_cache_remove_sequence(kv_cache): cached_k = {} cached_v = {} for seq_id_to_remove in range(num_sequences): - apply_attention(kv_cache, batch, cached_k, cached_v) + apply_attention(kv_cache, rope_mode, batch, cached_k, cached_v, fuse_qkv) # Remove sequence. fremove_sequence(kv_cache, seq_id_to_remove) cached_k.pop(seq_id_to_remove) @@ -417,20 +536,22 @@ def test_paged_attention_kv_cache_remove_sequence(kv_cache): @pytest.mark.skip(reason="Require FlashInfer enabled") -def test_paged_attention_kv_cache_fork_sequence(kv_cache): +@pytest.mark.parametrize("fuse_qkv", [False, True]) +def test_paged_attention_kv_cache_fork_sequence(kv_cache_and_rope_mode, fuse_qkv): + kv_cache, rope_mode = kv_cache_and_rope_mode fclear(kv_cache) cached_k = {} cached_v = {} batch = [(0, 60), (1, 88), (2, 17), (3, 4)] - apply_attention(kv_cache, batch, cached_k, cached_v) + apply_attention(kv_cache, rope_mode, batch, cached_k, cached_v, fuse_qkv) # Fork existing sequences. - apply_attention(kv_cache, [((4, 3), 35)], cached_k, cached_v) - apply_attention(kv_cache, [((5, 0), 20)], cached_k, cached_v) - apply_attention(kv_cache, [((6, 5), 102)], cached_k, cached_v) - apply_attention(kv_cache, [((7, 0), 3)], cached_k, cached_v) - apply_attention(kv_cache, [((8, 5), 71)], cached_k, cached_v) - apply_attention(kv_cache, [((9, 5), 20)], cached_k, cached_v) + apply_attention(kv_cache, rope_mode, [((4, 3), 35)], cached_k, cached_v, fuse_qkv) + apply_attention(kv_cache, rope_mode, [((5, 0), 20)], cached_k, cached_v, fuse_qkv) + apply_attention(kv_cache, rope_mode, [((6, 5), 102)], cached_k, cached_v, fuse_qkv) + apply_attention(kv_cache, rope_mode, [((7, 0), 3)], cached_k, cached_v, fuse_qkv) + apply_attention(kv_cache, rope_mode, [((8, 5), 71)], cached_k, cached_v, fuse_qkv) + apply_attention(kv_cache, rope_mode, [((9, 5), 20)], cached_k, cached_v, fuse_qkv) # Mixture of decode and prefill. operation_seq = [ [(2, 1), (4, 1), (7, 1), (6, 1), (8, 1), (9, 1)], @@ -439,18 +560,20 @@ def test_paged_attention_kv_cache_fork_sequence(kv_cache): [(7, 10), (6, 2), (8, 3), (9, 4)], ] for batch in operation_seq: - apply_attention(kv_cache, batch, cached_k, cached_v) + apply_attention(kv_cache, rope_mode, batch, cached_k, cached_v, fuse_qkv) @pytest.mark.skip(reason="Require FlashInfer enabled") -def test_paged_attention_kv_cache_popn(kv_cache): +@pytest.mark.parametrize("fuse_qkv", [False, True]) +def test_paged_attention_kv_cache_popn(kv_cache_and_rope_mode, fuse_qkv): + kv_cache, rope_mode = kv_cache_and_rope_mode fclear(kv_cache) cached_k = {} cached_v = {} batch = [(0, 35), (1, 88), (2, 17), (3, 4)] - apply_attention(kv_cache, batch, cached_k, cached_v) - apply_attention(kv_cache, [((4, 3), 35)], cached_k, cached_v) + apply_attention(kv_cache, rope_mode, batch, cached_k, cached_v, fuse_qkv) + apply_attention(kv_cache, rope_mode, [((4, 3), 35)], cached_k, cached_v, fuse_qkv) popn_operations = [(0, 17), (1, 57), (2, 16), (3, 0), (4, 19)] for seq_id, pop_length in popn_operations: @@ -462,8 +585,11 @@ def test_paged_attention_kv_cache_popn(kv_cache): if __name__ == "__main__": - cache = create_kv_cache() - test_paged_attention_kv_cache_prefill_and_decode(cache) - test_paged_attention_kv_cache_remove_sequence(cache) - test_paged_attention_kv_cache_fork_sequence(cache) - test_paged_attention_kv_cache_popn(cache) + set_global_func() + for rope_mode in [0, 1]: + cache = create_kv_cache(rope_mode) + for fuse_qkv in [False, True]: + test_paged_attention_kv_cache_prefill_and_decode((cache, rope_mode), fuse_qkv) + test_paged_attention_kv_cache_remove_sequence((cache, rope_mode), fuse_qkv) + test_paged_attention_kv_cache_fork_sequence((cache, rope_mode), fuse_qkv) + test_paged_attention_kv_cache_popn((cache, rope_mode), fuse_qkv) diff --git a/tests/python/relax/test_runtime_builtin_paged_attention_kv_cache_tir.py b/tests/python/relax/test_runtime_builtin_paged_attention_kv_cache_tir.py index bd667292ea44..e4c066342b65 100644 --- a/tests/python/relax/test_runtime_builtin_paged_attention_kv_cache_tir.py +++ b/tests/python/relax/test_runtime_builtin_paged_attention_kv_cache_tir.py @@ -30,6 +30,7 @@ reserved_nseq = 32 maximum_total_seq_length = 1024 +prefill_chunk_size = 512 page_size = 16 num_layers = 4 num_qo_heads = 32 @@ -48,8 +49,18 @@ fbegin_forward = None fend_forward = None fattention = None +fattention_with_fuse_qkv = None fdebug_get_kv = None +ftranspose_append = None +fcopy_cache = None +fattn_prefill = None +fattn_decode = None +fattn_prefill_ragged = None +fmerge_state = None +fsplit_rotary = None +fattention_rotary = None + @T.prim_func def kv_cache_transpose_append( @@ -123,7 +134,9 @@ def copy_cache( def set_global_func(): global fclear, fadd_sequence, fremove_sequence, ffork_sequence, fpopn - global fbegin_forward, fend_forward, fattention, fdebug_get_kv + global fbegin_forward, fend_forward, fattention, fattention_with_fuse_qkv, fdebug_get_kv + global ftranspose_append, fcopy_cache, fattn_prefill, fattn_decode, fattn_prefill_ragged + global fmerge_state, fsplit_rotary, fattention_rotary fclear = tvm.get_global_func("vm.builtin.paged_attention_kv_cache_clear") fadd_sequence = tvm.get_global_func("vm.builtin.paged_attention_kv_cache_add_sequence") @@ -133,11 +146,11 @@ def set_global_func(): fbegin_forward = tvm.get_global_func("vm.builtin.paged_attention_kv_cache_begin_forward") fend_forward = tvm.get_global_func("vm.builtin.paged_attention_kv_cache_end_forward") fattention = tvm.get_global_func("vm.builtin.paged_attention_kv_cache_attention") + fattention_with_fuse_qkv = tvm.get_global_func( + "vm.builtin.paged_attention_kv_cache_attention_with_fused_qkv" + ) fdebug_get_kv = tvm.get_global_func("vm.builtin.paged_attention_kv_cache_debug_get_kv") - -def create_kv_cache(): - set_global_func() target = tvm.target.Target("cuda") builts = [] for tir_func in [ @@ -147,6 +160,10 @@ def create_kv_cache(): _attention_decode(num_kv_heads, num_qo_heads, head_dim, dtype), _attention_prefill_ragged(num_kv_heads, num_qo_heads, head_dim, dtype), _merge_state_inplace(num_qo_heads, head_dim, dtype), + llama_rope_with_position_map( + rope_theta, rope_scale, head_dim, num_qo_heads, num_kv_heads, dtype + ), + _inplace_rope(rope_theta, rope_scale, head_dim, num_qo_heads, num_kv_heads, dtype), ]: mod = tvm.IRModule({"main": tir_func}) with target: @@ -161,14 +178,22 @@ def create_kv_cache(): fattn_decode, fattn_prefill_ragged, fmerge_state, + fsplit_rotary, + fattention_rotary, ) = builts + + +def create_kv_cache(rope_mode): fcreate = tvm.get_global_func("vm.builtin.paged_attention_kv_cache_create_reduced") cache = fcreate( - tvm.runtime.ShapeTuple([reserved_nseq, maximum_total_seq_length, page_size]), + tvm.runtime.ShapeTuple( + [reserved_nseq, maximum_total_seq_length, prefill_chunk_size, page_size] + ), num_layers, num_qo_heads, num_kv_heads, head_dim, + rope_mode, rope_scale, rope_theta, tvm.nd.empty((), dtype, device=device), @@ -177,14 +202,17 @@ def create_kv_cache(): fattn_decode, fattn_prefill_ragged, fmerge_state, + fsplit_rotary, + fattention_rotary, fcopy_cache, ) return cache -@pytest.fixture() -def kv_cache(): - return create_kv_cache() +@pytest.fixture(params=[0, 1]) +def kv_cache_and_rope_mode(request): + set_global_func() + return create_kv_cache(request.param), request.param def verify_cached_kv(kv_cache, seq_ids, expected_k, expected_v): @@ -220,9 +248,11 @@ def f_apply_rotary(x, offset, scale, theta): def apply_attention( kv_cache, + rope_mode: int, batch: List[Tuple[Union[int, Tuple[int, int]], int]], cached_k: Dict[int, np.ndarray], cached_v: Dict[int, np.ndarray], + fuse_qkv: bool, ) -> None: seq_ids = [] append_lengths = [] @@ -245,7 +275,6 @@ def apply_attention( cached_k[seq_id] = np.zeros((num_layers, 0, num_kv_heads, head_dim), dtype) cached_v[seq_id] = np.zeros((num_layers, 0, num_kv_heads, head_dim), dtype) - use_decode_shape = all(append_length == 1 for _, append_length in batch) fbegin_forward(kv_cache, ShapeTuple(seq_ids), ShapeTuple(append_lengths)) global_new_q = np.zeros((num_layers, 0, num_qo_heads, head_dim), dtype) @@ -262,7 +291,17 @@ def apply_attention( cached_k[seq_id] = np.concatenate( [ cached_k[seq_id], - np.stack([new_k[l] for l in range(num_layers)], axis=0), + np.stack( + [ + new_k[l] + if rope_mode == 1 + else f_apply_rotary( + new_k[l], cached_k[seq_id].shape[1], rope_scale, rope_theta + ) + for l in range(num_layers) + ], + axis=0, + ), ], axis=1, ) @@ -272,23 +311,22 @@ def apply_attention( global_new_v = np.concatenate([global_new_v, new_v], axis=1) for layer_id in range(num_layers): - queries_np = global_new_q[layer_id : layer_id + 1] - keys_np = global_new_k[layer_id : layer_id + 1] - values_np = global_new_v[layer_id : layer_id + 1] - if use_decode_shape: - queries_np = queries_np.transpose(1, 0, 2, 3) - keys_np = keys_np.transpose(1, 0, 2, 3) - values_np = values_np.transpose(1, 0, 2, 3) - queries = tvm.nd.array(queries_np, device=device) - keys = tvm.nd.array(keys_np, device=device) - values = tvm.nd.array(values_np, device=device) - outputs = tvm.nd.empty(queries.shape, dtype, device=device) - fattention(kv_cache, layer_id, queries, keys, values, outputs) + queries_np = global_new_q[layer_id] + keys_np = global_new_k[layer_id] + values_np = global_new_v[layer_id] + if not fuse_qkv: + queries = tvm.nd.array(queries_np, device=device) + keys = tvm.nd.array(keys_np, device=device) + values = tvm.nd.array(values_np, device=device) + outputs = tvm.nd.empty(queries.shape, dtype, device=device) + fattention(kv_cache, layer_id, queries, keys, values, outputs) + else: + qkv = tvm.nd.array(np.concatenate([queries_np, keys_np, values_np], axis=1), device) + outputs = tvm.nd.empty(queries_np.shape, dtype, device=device) + fattention_with_fuse_qkv(kv_cache, layer_id, qkv, outputs) # Compute attention expected results. - outputs = outputs.numpy() - if use_decode_shape: - outputs = outputs.transpose(1, 0, 2, 3) + outputs = np.expand_dims(outputs.numpy(), axis=0) sum_length = 0 for i, (seq_id, append_length) in enumerate(batch): assert cached_k[seq_id].shape[1] == cached_v[seq_id].shape[1] >= append_length @@ -300,9 +338,11 @@ def apply_attention( rope_scale, rope_theta, ).transpose(1, 0, 2) - k_seq = f_apply_rotary(cached_k[seq_id][layer_id], 0, rope_scale, rope_theta).transpose( - 1, 2, 0 - ) + k_seq = ( + cached_k[seq_id][layer_id] + if rope_mode == 0 + else f_apply_rotary(cached_k[seq_id][layer_id], 0, rope_scale, rope_theta) + ).transpose(1, 2, 0) v_seq = cached_v[seq_id][layer_id].transpose(1, 0, 2) k_seq = np.repeat(k_seq, num_qo_heads // num_kv_heads, axis=0) @@ -338,7 +378,9 @@ def apply_attention( @tvm.testing.requires_gpu @tvm.testing.requires_cuda -def test_paged_attention_kv_cache_prefill_and_decode(kv_cache): +@pytest.mark.parametrize("fuse_qkv", [False, True]) +def test_paged_attention_kv_cache_prefill_and_decode(kv_cache_and_rope_mode, fuse_qkv): + kv_cache, rope_mode = kv_cache_and_rope_mode fclear(kv_cache) # Prefill. @@ -354,12 +396,14 @@ def test_paged_attention_kv_cache_prefill_and_decode(kv_cache): cached_k = {} cached_v = {} for batch in operation_seq: - apply_attention(kv_cache, batch, cached_k, cached_v) + apply_attention(kv_cache, rope_mode, batch, cached_k, cached_v, fuse_qkv) @tvm.testing.requires_gpu @tvm.testing.requires_cuda -def test_paged_attention_kv_cache_remove_sequence(kv_cache): +@pytest.mark.parametrize("fuse_qkv", [False, True]) +def test_paged_attention_kv_cache_remove_sequence(kv_cache_and_rope_mode, fuse_qkv): + kv_cache, rope_mode = kv_cache_and_rope_mode fclear(kv_cache) num_sequences = 5 @@ -367,7 +411,7 @@ def test_paged_attention_kv_cache_remove_sequence(kv_cache): cached_k = {} cached_v = {} for seq_id_to_remove in range(num_sequences): - apply_attention(kv_cache, batch, cached_k, cached_v) + apply_attention(kv_cache, rope_mode, batch, cached_k, cached_v, fuse_qkv) # Remove sequence. fremove_sequence(kv_cache, seq_id_to_remove) cached_k.pop(seq_id_to_remove) @@ -382,20 +426,22 @@ def test_paged_attention_kv_cache_remove_sequence(kv_cache): @tvm.testing.requires_gpu @tvm.testing.requires_cuda -def test_paged_attention_kv_cache_fork_sequence(kv_cache): +@pytest.mark.parametrize("fuse_qkv", [False, True]) +def test_paged_attention_kv_cache_fork_sequence(kv_cache_and_rope_mode, fuse_qkv): + kv_cache, rope_mode = kv_cache_and_rope_mode fclear(kv_cache) cached_k = {} cached_v = {} batch = [(0, 60), (1, 88), (2, 17), (3, 4)] - apply_attention(kv_cache, batch, cached_k, cached_v) + apply_attention(kv_cache, rope_mode, batch, cached_k, cached_v, fuse_qkv) # Fork existing sequences. - apply_attention(kv_cache, [((4, 3), 35)], cached_k, cached_v) - apply_attention(kv_cache, [((5, 0), 20)], cached_k, cached_v) - apply_attention(kv_cache, [((6, 5), 102)], cached_k, cached_v) - apply_attention(kv_cache, [((7, 0), 3)], cached_k, cached_v) - apply_attention(kv_cache, [((8, 5), 71)], cached_k, cached_v) - apply_attention(kv_cache, [((9, 5), 20)], cached_k, cached_v) + apply_attention(kv_cache, rope_mode, [((4, 3), 35)], cached_k, cached_v, fuse_qkv) + apply_attention(kv_cache, rope_mode, [((5, 0), 20)], cached_k, cached_v, fuse_qkv) + apply_attention(kv_cache, rope_mode, [((6, 5), 102)], cached_k, cached_v, fuse_qkv) + apply_attention(kv_cache, rope_mode, [((7, 0), 3)], cached_k, cached_v, fuse_qkv) + apply_attention(kv_cache, rope_mode, [((8, 5), 71)], cached_k, cached_v, fuse_qkv) + apply_attention(kv_cache, rope_mode, [((9, 5), 20)], cached_k, cached_v, fuse_qkv) # Mixture of decode and prefill. operation_seq = [ [(2, 1), (4, 1), (7, 1), (6, 1), (8, 1), (9, 1)], @@ -404,18 +450,21 @@ def test_paged_attention_kv_cache_fork_sequence(kv_cache): [(7, 10), (6, 2), (8, 3), (9, 4)], ] for batch in operation_seq: - apply_attention(kv_cache, batch, cached_k, cached_v) + apply_attention(kv_cache, rope_mode, batch, cached_k, cached_v, fuse_qkv) @tvm.testing.requires_gpu @tvm.testing.requires_cuda -def test_paged_attention_kv_cache_popn(kv_cache): +@pytest.mark.parametrize("fuse_qkv", [False, True]) +def test_paged_attention_kv_cache_popn(kv_cache_and_rope_mode, fuse_qkv): + kv_cache, rope_mode = kv_cache_and_rope_mode fclear(kv_cache) cached_k = {} cached_v = {} batch = [(0, 35), (1, 88), (2, 17), (3, 4)] - apply_attention(kv_cache, batch, cached_k, cached_v) + apply_attention(kv_cache, rope_mode, batch, cached_k, cached_v, fuse_qkv) + apply_attention(kv_cache, rope_mode, [((4, 3), 35)], cached_k, cached_v, fuse_qkv) popn_operations = [(0, 17), (1, 57), (2, 16), (3, 0)] for seq_id, pop_length in popn_operations: @@ -426,6 +475,154 @@ def test_paged_attention_kv_cache_popn(kv_cache): verify_cached_kv(kv_cache, seq_ids=list(range(4)), expected_k=cached_k, expected_v=cached_v) +def _inplace_rope( + theta: float, + scale: float, + head_dim: int, + num_q_heads: int, + num_kv_heads: int, + dtype: str, +): + assert head_dim <= 128, "Rotary embedding currently only supports head_dim <= 128" + rotary_dim = head_dim + + def _rope( + x: T.Buffer, + s: tir.Var, + h: tir.Var, + d: tir.Var, + rope_offset: tir.Var, + instance_offset: tir.Var, + ): + cos_freq, sin_freq = rope_freq((s + rope_offset) * scale, d, rotary_dim, theta, dtype) + cos = cos_freq * x[s + instance_offset, h, d] + sin = sin_freq * tir.if_then_else( + d < rotary_dim // 2, + -x[s + instance_offset, h, d + rotary_dim // 2], + x[s + instance_offset, h, d - rotary_dim // 2], + ) + return cos + sin + + # fmt: off + @T.prim_func + def tir_rotary( + var_q: T.handle, + var_k: T.handle, + var_append_len_indptr: T.handle, + var_rope_offsets: T.handle, + _0: T.int32, + _1: T.int32, + _2: T.int32, + _3: T.int32, + _4: T.int32, + _5: T.float32, + _6: T.float32, + ): + T.func_attr({"tir.is_scheduled": 1}) + total_len = T.int32() + batch_size = T.int32() + q = T.match_buffer(var_q, (total_len, num_q_heads, head_dim), dtype) + k = T.match_buffer(var_k, (total_len, num_kv_heads, head_dim), dtype) + rope_offsets = T.match_buffer(var_rope_offsets, (batch_size,), "int32") + append_len_indptr = T.match_buffer(var_append_len_indptr, (batch_size + 1,), "int32") + for b_h in T.thread_binding(batch_size * (num_q_heads + num_kv_heads), thread="blockIdx.x"): + b: T.int32 = b_h // (num_q_heads + num_kv_heads) + h: T.int32 = b_h % (num_q_heads + num_kv_heads) + instance_offset: T.int32 = append_len_indptr[b] + rope_offset: T.int32 = rope_offsets[b] + append_len: T.int32 = append_len_indptr[b + 1] - append_len_indptr[b] + for s0 in range(T.ceildiv(append_len, 32)): + for s1 in T.thread_binding(32, thread="threadIdx.y"): + for d0 in T.thread_binding(T.ceildiv(head_dim, 4), thread="threadIdx.x"): + for d1 in T.vectorized(4): + s: T.int32 = s0 * 32 + s1 + d: T.int32 = d0 * 4 + d1 + if s < append_len and d < head_dim: + if h < num_q_heads: + q[s + instance_offset, h, d] = _rope(q, s, h, d, rope_offset, instance_offset) + else: + k[s + instance_offset, h - num_q_heads, d] = _rope(k, s, h - num_q_heads, d, rope_offset, instance_offset) + return tir_rotary + + +def llama_rope_with_position_map( # pylint: disable=too-many-arguments + theta: float, + scale: float, + head_dim: int, + num_q_heads: int, + num_kv_heads: int, + dtype: float = "float16", + rotary_dim: int = None, +): + fused_heads = num_q_heads + num_kv_heads * 2 + if rotary_dim is None: + rotary_dim = head_dim + scale = tir.const(scale, dtype) + + def _rope_freq(s: tir.Var, d: tir.Var, d_range: int, theta: float, dtype: str): + freq = s / tir.power(theta, d * 2 % d_range / tir.const(d_range, "float32")) + cos_freq = tir.cos(freq).astype(dtype) + sin_freq = tir.sin(freq).astype(dtype) + return cos_freq, sin_freq + + def _rope( # pylint: disable=too-many-arguments + x: T.Buffer, + s: tir.Var, + h: tir.Var, + d: tir.Var, + pos: tir.Var, + ): + cos_freq, sin_freq = _rope_freq(pos * scale, d, rotary_dim, theta, dtype) + cos = cos_freq * x[s, h, d] + sin = sin_freq * tir.if_then_else( + d < rotary_dim // 2, + -x[s, h, d + rotary_dim // 2], + x[s, h, d - rotary_dim // 2], + ) + return cos + sin + + @T.prim_func(private=True) + def fused_rope( # pylint: disable=too-many-locals + var_qkv: T.handle, + var_position_map: T.handle, + var_q: T.handle, + var_k: T.handle, + var_v: T.handle, + apply_rope: T.int32, + ): + T.func_attr( + { + "op_pattern": 8, # 2 means injective, 8 means opaque + "tir.noalias": T.bool(True), + } + ) + seq_len = T.int64() + qkv = T.match_buffer(var_qkv, (seq_len, fused_heads, head_dim), dtype) + q = T.match_buffer(var_q, (seq_len, num_q_heads, head_dim), dtype) + k = T.match_buffer(var_k, (seq_len, num_kv_heads, head_dim), dtype) + v = T.match_buffer(var_v, (seq_len, num_kv_heads, head_dim), dtype) + position_map = T.match_buffer(var_position_map, (seq_len,), "int32") + for iters in T.grid(seq_len, fused_heads, head_dim): + with T.block("llama_fused_rope"): + s, h, d = T.axis.remap("SSS", iters) + if h < num_q_heads: + q[s, h, d] = T.if_then_else( + apply_rope > 0 and d < rotary_dim, + _rope(qkv, s, h, d, position_map[s]), + qkv[s, h, d], + ) + elif h < num_q_heads + num_kv_heads: + k[s, h - num_q_heads, d] = T.if_then_else( + apply_rope > 0 and d < rotary_dim, + _rope(qkv, s, h, d, position_map[s]), + qkv[s, h, d], + ) + else: + v[s, h - (num_q_heads + num_kv_heads), d] = qkv[s, h, d] + + return fused_rope + + def rope_freq(s: tir.Var, d: tir.Var, d_range: int, theta: float, dtype: str): """Compute the inverse frequency of RoPE and then return the cosine and sine of it. @@ -1449,8 +1646,11 @@ def merge_state_inplace( if __name__ == "__main__": - cache = create_kv_cache() - test_paged_attention_kv_cache_prefill_and_decode(cache) - test_paged_attention_kv_cache_remove_sequence(cache) - test_paged_attention_kv_cache_fork_sequence(cache) - test_paged_attention_kv_cache_popn(cache) + set_global_func() + for rope_mode in [0, 1]: + cache = create_kv_cache(rope_mode) + for fuse_qkv in [False, True]: + test_paged_attention_kv_cache_prefill_and_decode((cache, rope_mode), fuse_qkv) + test_paged_attention_kv_cache_remove_sequence((cache, rope_mode), fuse_qkv) + test_paged_attention_kv_cache_fork_sequence((cache, rope_mode), fuse_qkv) + test_paged_attention_kv_cache_popn((cache, rope_mode), fuse_qkv)