diff --git a/src/runtime/vm/kv_state.cc b/src/runtime/vm/kv_state.cc index 05f951dc74d0..a912d1995ac2 100644 --- a/src/runtime/vm/kv_state.cc +++ b/src/runtime/vm/kv_state.cc @@ -70,6 +70,18 @@ TVM_FFI_STATIC_INIT_BLOCK() { &AttentionKVCacheObj::GetNumAvailablePages) .def_method("vm.builtin.attention_kv_cache_get_total_sequence_length", &AttentionKVCacheObj::GetTotalSequenceLength) + .def_method("vm.builtin.attention_kv_cache_get_checkpoint_metadata", + &AttentionKVCacheObj::GetCheckpointMetadata) + .def_method("vm.builtin.attention_kv_cache_get_layout_hash", + &AttentionKVCacheObj::GetLayoutHash) + .def_method("vm.builtin.attention_kv_cache_export_page_group", + &AttentionKVCacheObj::ExportPageGroup) + .def_method("vm.builtin.attention_kv_cache_prepare_import", + &AttentionKVCacheObj::PrepareImport) + .def_method("vm.builtin.attention_kv_cache_import_page_group", + &AttentionKVCacheObj::ImportPageGroup) + .def_method("vm.builtin.attention_kv_cache_get_sequence_length", + &AttentionKVCacheObj::GetSequenceLength) .def_method("vm.builtin.attention_kv_cache_get_query_positions", &AttentionKVCacheObj::GetQueryPositions) .def_method("vm.builtin.attention_kv_cache_debug_get_kv", &AttentionKVCacheObj::DebugGetKV) diff --git a/src/runtime/vm/kv_state.h b/src/runtime/vm/kv_state.h index 198bd18d979d..beda4329f974 100644 --- a/src/runtime/vm/kv_state.h +++ b/src/runtime/vm/kv_state.h @@ -23,6 +23,7 @@ #include #include #include +#include #include #include @@ -135,6 +136,49 @@ class AttentionKVCacheObj : public KVStateObj { /*! \brief Get the current total sequence length in the KV cache. */ virtual int32_t GetTotalSequenceLength() const = 0; + /*! + * \brief Get checkpoint metadata for a sequence. + * \param seq_id The id of the sequence whose checkpoint metadata is requested. + * \return JSON string describing runtime layout and logical pages. + */ + virtual ffi::String GetCheckpointMetadata(int64_t seq_id) const = 0; + + /*! + * \brief Get a stable hash over runtime layout-defining fields. + * \return Hex string hash of the checkpoint layout metadata. + */ + virtual ffi::String GetLayoutHash() const = 0; + + /*! + * \brief Export a checkpoint page group for a sequence. + * \param seq_id The id of the sequence whose page group is exported. + * \param group_id The checkpoint group id to export. + * \param dst The destination tensor for the exported page group. + */ + virtual void ExportPageGroup(int64_t seq_id, int64_t group_id, Tensor dst) = 0; + + /*! + * \brief Prepare sequence state and page tables for checkpoint import. + * \param seq_id The id of the sequence being imported. + * \param metadata_json The checkpoint metadata JSON string. + */ + virtual void PrepareImport(int64_t seq_id, ffi::String metadata_json) = 0; + + /*! + * \brief Import a checkpoint page group into a prepared sequence. + * \param seq_id The id of the sequence being imported. + * \param group_id The checkpoint group id to import. + * \param src The source tensor containing the page group. + */ + virtual void ImportPageGroup(int64_t seq_id, int64_t group_id, Tensor src) = 0; + + /*! + * \brief Get the sequence length for a sequence in the KV cache. + * \param seq_id The id of the sequence whose length is requested. + * \return The sequence length. + */ + virtual int32_t GetSequenceLength(int64_t seq_id) const = 0; + /************** Sequence Management **************/ /*! diff --git a/src/runtime/vm/paged_kv_cache.cc b/src/runtime/vm/paged_kv_cache.cc index efd099e82616..62d94fdf8350 100644 --- a/src/runtime/vm/paged_kv_cache.cc +++ b/src/runtime/vm/paged_kv_cache.cc @@ -22,6 +22,7 @@ */ #include #include +#include #include #include #include @@ -31,7 +32,11 @@ #include #include +#include +#include +#include #include +#include #include #include #include @@ -44,6 +49,71 @@ namespace tvm { namespace runtime { namespace vm { +namespace { + +constexpr const char* kPagedKVCacheCheckpointRuntime = "relax.vm.PagedAttentionKVCache"; + +const char* AttnKindToString(AttnKind attn_kind) { + switch (attn_kind) { + case AttnKind::kMHA: + return "mha"; + case AttnKind::kMLA: + return "mla"; + case AttnKind::kLinearAttn: + return "linear"; + case AttnKind::kMHASliding: + return "mha_sliding"; + } + TVM_FFI_ICHECK(false) << "Unknown attention kind: " << static_cast(attn_kind); + return "unknown"; +} + +const char* RoPEModeToString(RoPEMode rope_mode) { + switch (rope_mode) { + case RoPEMode::kNone: + return "none"; + case RoPEMode::kNormal: + return "normal"; + case RoPEMode::kInline: + return "inline"; + } + TVM_FFI_ICHECK(false) << "Unknown RoPE mode: " << static_cast(rope_mode); + return "unknown"; +} + +ffi::json::Array ShapeToJSON(const int64_t* shape, int ndim) { + ffi::json::Array result; + for (int i = 0; i < ndim; ++i) { + result.push_back(shape[i]); + } + return result; +} + +ffi::json::Array IntArrayToJSON(const std::vector& values) { + ffi::json::Array result; + for (int32_t value : values) { + result.push_back(static_cast(value)); + } + return result; +} + +std::string Uint64ToHex(uint64_t value) { + std::ostringstream os; + os << std::hex << std::setw(16) << std::setfill('0') << value; + return os.str(); +} + +uint64_t FNV1a64(const std::string& value) { + uint64_t hash = 14695981039346656037ULL; + for (unsigned char c : value) { + hash ^= c; + hash *= 1099511628211ULL; + } + return hash; +} + +} // namespace + //------------------------------------------- // We keep the implementation private as // they may subject to future changes. @@ -858,6 +928,138 @@ class PagedAttentionKVCacheObj : public AttentionKVCacheObj { return total_seq_len; } + ffi::String GetCheckpointMetadata(int64_t seq_id) const final { + CheckCheckpointSequenceSupported(seq_id); + const Sequence& seq = seq_map_.at(seq_id); + + namespace json = tvm::ffi::json; + json::Object metadata = MakeLayoutMetadata(); + metadata.Set("layoutHash", GetLayoutHash()); + metadata.Set("seqId", seq_id); + metadata.Set("seqLength", static_cast(seq.seq_length)); + metadata.Set("blocks", MakeBlockMetadata(seq)); + metadata.Set("logicalPages", MakeLogicalPageMetadata(seq)); + metadata.Set("groups", MakePageGroupMetadata(seq)); + return json::Stringify(metadata); + } + + ffi::String GetLayoutHash() const final { + CheckCheckpointLayoutSupported(); + return ffi::String(Uint64ToHex(FNV1a64(GetLayoutDescriptor()))); + } + + void ExportPageGroup(int64_t seq_id, int64_t group_id, Tensor dst) final { + CheckCheckpointSequenceSupported(seq_id); + const Sequence& seq = seq_map_.at(seq_id); + TVM_FFI_ICHECK_GE(group_id, 0) + << "PagedAttentionKVCache checkpoint export got invalid group id " << group_id << "."; + TVM_FFI_ICHECK_LT(group_id, num_layers_) + << "PagedAttentionKVCache checkpoint export got invalid group id " << group_id + << ", but only " << num_layers_ << " groups are available."; + + int64_t num_logical_pages = GetNumLogicalPages(seq); + CheckExportPageGroupTensor(dst, num_logical_pages); + + if (copy_stream_ != nullptr) { + DeviceAPI::Get(device_)->SyncStreamFromTo(device_, copy_stream_, compute_stream_); + } + + int64_t bytes_per_page = 2 * num_kv_heads_ * page_size_ * qk_head_dim_ * + ((static_cast(kv_dtype_.bits) * kv_dtype_.lanes + 7) / 8); + Tensor layer_pages = pages_[group_id]; + int64_t logical_page_index = 0; + for (int32_t block_id : seq.GetBlockTrace(global_block_pool_)) { + const Block& block = global_block_pool_[block_id]; + for (int32_t page_id : block.page_ids) { + TVM_FFI_ICHECK_GE(page_id, 0) + << "PagedAttentionKVCache checkpoint export found invalid page id " << page_id << "."; + TVM_FFI_ICHECK_LT(page_id, num_total_pages_) + << "PagedAttentionKVCache checkpoint export found out-of-range page id " << page_id + << "."; + Tensor src_page = + layer_pages.CreateView({2, num_kv_heads_, page_size_, qk_head_dim_}, layer_pages->dtype, + static_cast(page_id * bytes_per_page)); + Tensor dst_page = + dst.CreateView({2, num_kv_heads_, page_size_, qk_head_dim_}, dst->dtype, + static_cast(logical_page_index * bytes_per_page)); + DLTensor dst_page_view = *dst_page.operator->(); + Tensor::CopyFromTo(src_page.operator->(), &dst_page_view, compute_stream_); + ++logical_page_index; + } + } + TVM_FFI_ICHECK_EQ(logical_page_index, num_logical_pages); + } + + void PrepareImport(int64_t seq_id, ffi::String metadata_json) final { + CheckCheckpointLayoutSupported(); + TVM_FFI_ICHECK_EQ(seq_id, 0) + << "PagedAttentionKVCache checkpoint import only supports sequence id 0, got " << seq_id + << "."; + ffi::json::Object metadata = ParseCheckpointMetadata(metadata_json); + int64_t seq_length = CheckCheckpointImportMetadata(seq_id, metadata); + int64_t num_logical_pages = GetExpectedNumLogicalPages(seq_length); + + Clear(); + int32_t block_idx = GetFreeBlock(); + Block& block = global_block_pool_[block_idx]; + block.start_pos = 0; + block.seq_length = static_cast(seq_length); + for (int64_t page_index = 0; page_index < num_logical_pages; ++page_index) { + block.page_ids.push_back(GetFreePage()); + } + seq_map_.insert({seq_id, Sequence(&global_block_pool_, block_idx)}); + dirty_aux_data_device_ = true; + } + + void ImportPageGroup(int64_t seq_id, int64_t group_id, Tensor src) final { + CheckCheckpointSequenceSupported(seq_id); + const Sequence& seq = seq_map_.at(seq_id); + TVM_FFI_ICHECK_GE(group_id, 0) + << "PagedAttentionKVCache checkpoint import got invalid group id " << group_id << "."; + TVM_FFI_ICHECK_LT(group_id, num_layers_) + << "PagedAttentionKVCache checkpoint import got invalid group id " << group_id + << ", but only " << num_layers_ << " groups are available."; + + int64_t num_logical_pages = GetNumLogicalPages(seq); + CheckImportPageGroupTensor(src, num_logical_pages); + + if (copy_stream_ != nullptr) { + DeviceAPI::Get(device_)->SyncStreamFromTo(device_, copy_stream_, compute_stream_); + } + + int64_t bytes_per_page = 2 * num_kv_heads_ * page_size_ * qk_head_dim_ * + ((static_cast(kv_dtype_.bits) * kv_dtype_.lanes + 7) / 8); + Tensor layer_pages = pages_[group_id]; + int64_t logical_page_index = 0; + for (int32_t block_id : seq.GetBlockTrace(global_block_pool_)) { + const Block& block = global_block_pool_[block_id]; + for (int32_t page_id : block.page_ids) { + TVM_FFI_ICHECK_GE(page_id, 0) + << "PagedAttentionKVCache checkpoint import found invalid page id " << page_id << "."; + TVM_FFI_ICHECK_LT(page_id, num_total_pages_) + << "PagedAttentionKVCache checkpoint import found out-of-range page id " << page_id + << "."; + Tensor src_page = + src.CreateView({2, num_kv_heads_, page_size_, qk_head_dim_}, src->dtype, + static_cast(logical_page_index * bytes_per_page)); + Tensor dst_page = + layer_pages.CreateView({2, num_kv_heads_, page_size_, qk_head_dim_}, layer_pages->dtype, + static_cast(page_id * bytes_per_page)); + DLTensor dst_page_view = *dst_page.operator->(); + Tensor::CopyFromTo(src_page.operator->(), &dst_page_view, compute_stream_); + ++logical_page_index; + } + } + TVM_FFI_ICHECK_EQ(logical_page_index, num_logical_pages); + } + + int32_t GetSequenceLength(int64_t seq_id) const final { + auto it = seq_map_.find(seq_id); + TVM_FFI_ICHECK(it != seq_map_.end()) + << "The sequence \"" << seq_id << "\" cannot be found in KV cache."; + return it->second.seq_length; + } + /************** Attention **************/ void BeginForward(const ffi::Shape& seq_ids, const ffi::Shape& append_lengths, @@ -1721,6 +1923,390 @@ class PagedAttentionKVCacheObj : public AttentionKVCacheObj { AttentionKVCacheObj); private: + void CheckCheckpointLayoutSupported() const { + for (int64_t layer_id = layer_id_begin_offset_; layer_id < layer_id_end_offset_; ++layer_id) { + AttnKind attn_kind = attn_kinds_[layer_id]; + TVM_FFI_ICHECK(attn_kind == AttnKind::kMHA) + << "PagedAttentionKVCache checkpointing only supports full-context MHA/GQA layers."; + } + TVM_FFI_ICHECK_EQ(qk_head_dim_, v_head_dim_) + << "PagedAttentionKVCache checkpointing requires qk_head_dim to equal v_head_dim."; + TVM_FFI_ICHECK(!support_sliding_window_ && !support_layer_sliding_window_) + << "PagedAttentionKVCache checkpointing does not support sliding-window cache layouts."; + TVM_FFI_ICHECK(!f_transfer_kv_.has_value() && !f_transfer_kv_page_to_page_.has_value()) + << "PagedAttentionKVCache checkpointing does not support KV transfer/disaggregation."; + } + + void CheckCheckpointSequenceSupported(int64_t seq_id) const { + CheckCheckpointLayoutSupported(); + TVM_FFI_ICHECK_EQ(seq_id, 0) + << "PagedAttentionKVCache checkpointing only supports sequence id 0, got " << seq_id << "."; + auto it = seq_map_.find(seq_id); + TVM_FFI_ICHECK(it != seq_map_.end()) + << "The sequence \"" << seq_id << "\" cannot be found in KV cache."; + const Sequence& seq = it->second; + TVM_FFI_ICHECK(seq.accepted_indices_committed && seq.is_chain) + << "PagedAttentionKVCache checkpointing requires committed token-chain state."; + TVM_FFI_ICHECK_EQ(seq.sliding_window_size, -1) + << "PagedAttentionKVCache checkpointing does not support sequences with sliding window."; + } + + int64_t GetExpectedNumLogicalPages(int64_t seq_length) const { + return (seq_length + page_size_ - 1) / page_size_; + } + + ffi::json::Object ParseCheckpointMetadata(const ffi::String& metadata_json) const { + ffi::String error_msg; + ffi::json::Value json_info = ffi::json::Parse(metadata_json, &error_msg); + TVM_FFI_ICHECK(error_msg.empty()) + << "Failed to parse PagedAttentionKVCache checkpoint metadata JSON: " << error_msg << "."; + TVM_FFI_ICHECK(json_info.as()) + << "PagedAttentionKVCache checkpoint metadata should be a JSON object."; + return json_info.cast(); + } + + ffi::json::Value GetJSONField(const ffi::json::Object& object, const char* field, + const char* context) const { + auto it = object.find(field); + TVM_FFI_ICHECK(it != object.end()) << context << " missing field \"" << field << "\"."; + return (*it).second; + } + + int64_t GetJSONIntegerField(const ffi::json::Object& object, const char* field, + const char* context) const { + return GetJSONField(object, field, context).cast(); + } + + bool GetJSONBoolField(const ffi::json::Object& object, const char* field, + const char* context) const { + return GetJSONField(object, field, context).cast(); + } + + std::string GetJSONStringField(const ffi::json::Object& object, const char* field, + const char* context) const { + return std::string(GetJSONField(object, field, context).cast()); + } + + ffi::json::Array GetJSONArrayField(const ffi::json::Object& object, const char* field, + const char* context) const { + return GetJSONField(object, field, context).cast(); + } + + void CheckJSONIntegerField(const ffi::json::Object& object, const char* field, int64_t expected, + const char* context) const { + int64_t value = GetJSONIntegerField(object, field, context); + TVM_FFI_ICHECK_EQ(value, expected) + << context << " field \"" << field << "\" mismatch: expected " << expected << ", got " + << value << "."; + } + + void CheckJSONStringField(const ffi::json::Object& object, const char* field, + const std::string& expected, const char* context) const { + std::string value = GetJSONStringField(object, field, context); + TVM_FFI_ICHECK_EQ(value, expected) + << context << " field \"" << field << "\" mismatch: expected " << expected << ", got " + << value << "."; + } + + void CheckJSONBoolField(const ffi::json::Object& object, const char* field, bool expected, + const char* context) const { + bool value = GetJSONBoolField(object, field, context); + TVM_FFI_ICHECK_EQ(value, expected) + << context << " field \"" << field << "\" mismatch: expected " << expected << ", got " + << value << "."; + } + + void CheckCheckpointImportLayout(const ffi::json::Object& metadata) const { + static constexpr const char* context = "PagedAttentionKVCache checkpoint import metadata"; + CheckJSONStringField(metadata, "cacheType", kPagedKVCacheCheckpointRuntime, context); + CheckJSONIntegerField(metadata, "pageSize", page_size_, context); + CheckJSONIntegerField(metadata, "numLayers", num_layers_, context); + CheckJSONIntegerField(metadata, "layerBegin", layer_id_begin_offset_, context); + CheckJSONIntegerField(metadata, "layerEnd", layer_id_end_offset_, context); + CheckJSONIntegerField(metadata, "numQOHeads", num_qo_heads_, context); + CheckJSONIntegerField(metadata, "numKVHeads", num_kv_heads_, context); + CheckJSONIntegerField(metadata, "qkHeadDim", qk_head_dim_, context); + CheckJSONIntegerField(metadata, "vHeadDim", v_head_dim_, context); + CheckJSONIntegerField(metadata, "numTotalPages", num_total_pages_, context); + CheckJSONIntegerField(metadata, "prefillChunkSize", prefill_chunk_size_, context); + CheckJSONStringField(metadata, "dtype", std::string(ffi::DLDataTypeToString(kv_dtype_)), + context); + CheckJSONStringField(metadata, "ropeMode", RoPEModeToString(rope_mode_), context); + CheckJSONBoolField(metadata, "hasRopeExtFactors", rope_ext_factors_.has_value(), context); + CheckJSONBoolField(metadata, "supportSlidingWindow", support_sliding_window_, context); + CheckJSONBoolField(metadata, "supportLayerSlidingWindow", support_layer_sliding_window_, + context); + CheckJSONStringField(metadata, "pageTensorLayout", + "num_total_pages,2,num_kv_heads,page_size,qk_head_dim", context); + + ffi::String expected_layout_hash = GetLayoutHash(); + std::string layout_hash = GetJSONStringField(metadata, "layoutHash", context); + TVM_FFI_ICHECK_EQ(layout_hash, std::string(expected_layout_hash)) + << "PagedAttentionKVCache checkpoint import layout hash mismatch: expected " + << expected_layout_hash << ", got " << layout_hash << "."; + + ffi::json::Array attn_kinds = GetJSONArrayField(metadata, "attnKinds", context); + TVM_FFI_ICHECK_EQ(attn_kinds.size(), num_layers_) + << context << " field \"attnKinds\" size mismatch."; + for (int64_t local_layer = 0; local_layer < num_layers_; ++local_layer) { + std::string attn_kind = std::string(attn_kinds[local_layer].cast()); + std::string expected = AttnKindToString(attn_kinds_[layer_id_begin_offset_ + local_layer]); + TVM_FFI_ICHECK_EQ(attn_kind, expected) + << context << " field \"attnKinds\" mismatch at local layer " << local_layer + << ": expected " << expected << ", got " << attn_kind << "."; + } + } + + void CheckCheckpointImportLogicalPages(const ffi::json::Object& metadata, + int64_t seq_length) const { + static constexpr const char* context = "PagedAttentionKVCache checkpoint import metadata"; + int64_t num_logical_pages = GetExpectedNumLogicalPages(seq_length); + ffi::json::Array logical_pages = GetJSONArrayField(metadata, "logicalPages", context); + TVM_FFI_ICHECK_EQ(static_cast(logical_pages.size()), num_logical_pages) + << "PagedAttentionKVCache checkpoint import sequence length mismatch: seqLength " + << seq_length << " requires " << num_logical_pages << " logical pages, but metadata has " + << logical_pages.size() << "."; + + int64_t expected_start = 0; + for (int64_t i = 0; i < num_logical_pages; ++i) { + ffi::json::Object page = logical_pages[i].cast(); + int64_t expected_length = std::min(page_size_, seq_length - expected_start); + CheckJSONIntegerField(page, "logicalPageIndex", i, context); + CheckJSONIntegerField(page, "startPos", expected_start, context); + CheckJSONIntegerField(page, "length", expected_length, context); + expected_start += expected_length; + } + TVM_FFI_ICHECK_EQ(expected_start, seq_length) + << "PagedAttentionKVCache checkpoint import logical pages do not cover seqLength " + << seq_length << "."; + } + + void CheckPageGroupMetadataShape(const ffi::json::Object& group, int64_t num_logical_pages, + const char* context) const { + ffi::json::Array shape = GetJSONArrayField(group, "shape", context); + std::vector expected_shape = {1, num_logical_pages, 2, num_kv_heads_, + page_size_, qk_head_dim_}; + TVM_FFI_ICHECK_EQ(shape.size(), expected_shape.size()) + << context << " field \"shape\" rank mismatch."; + for (int64_t i = 0; i < static_cast(expected_shape.size()); ++i) { + int64_t dim = shape[i].cast(); + TVM_FFI_ICHECK_EQ(dim, expected_shape[i]) + << context << " field \"shape\" mismatch at dim " << i << ": expected " + << expected_shape[i] << ", got " << dim << "."; + } + } + + void CheckCheckpointImportGroups(const ffi::json::Object& metadata, + int64_t num_logical_pages) const { + static constexpr const char* context = "PagedAttentionKVCache checkpoint import group metadata"; + ffi::json::Array groups = GetJSONArrayField(metadata, "groups", context); + TVM_FFI_ICHECK_EQ(static_cast(groups.size()), num_layers_) + << context << " size mismatch: expected " << num_layers_ << ", got " << groups.size() + << "."; + for (int64_t local_layer = 0; local_layer < num_layers_; ++local_layer) { + ffi::json::Object group = groups[local_layer].cast(); + CheckJSONIntegerField(group, "groupIndex", local_layer, context); + CheckJSONIntegerField(group, "layerBegin", layer_id_begin_offset_ + local_layer, context); + CheckJSONIntegerField(group, "layerEnd", layer_id_begin_offset_ + local_layer + 1, context); + CheckJSONIntegerField(group, "numLogicalPages", num_logical_pages, context); + CheckJSONStringField(group, "dtype", std::string(ffi::DLDataTypeToString(kv_dtype_)), + context); + CheckPageGroupMetadataShape(group, num_logical_pages, context); + } + } + + int64_t CheckCheckpointImportMetadata(int64_t seq_id, const ffi::json::Object& metadata) const { + CheckCheckpointImportLayout(metadata); + CheckJSONIntegerField(metadata, "seqId", seq_id, + "PagedAttentionKVCache checkpoint import metadata"); + int64_t seq_length = GetJSONIntegerField(metadata, "seqLength", + "PagedAttentionKVCache checkpoint import metadata"); + TVM_FFI_ICHECK_GE(seq_length, 0) + << "PagedAttentionKVCache checkpoint import seqLength cannot be negative."; + TVM_FFI_ICHECK_LE(seq_length, std::numeric_limits::max()) + << "PagedAttentionKVCache checkpoint import seqLength exceeds int32 range."; + int64_t num_logical_pages = GetExpectedNumLogicalPages(seq_length); + TVM_FFI_ICHECK_LE(num_logical_pages, num_total_pages_) + << "PagedAttentionKVCache checkpoint import requires " << num_logical_pages + << " pages, but this cache only has " << num_total_pages_ << " pages."; + CheckCheckpointImportLogicalPages(metadata, seq_length); + CheckCheckpointImportGroups(metadata, num_logical_pages); + return seq_length; + } + + void CheckPageGroupTensor(const Tensor& tensor, int64_t num_logical_pages, + const char* api_name) const { + std::string error_msg = std::string(api_name) + + " expects the tensor in layout " + "(1,num_logical_pages,2,num_kv_heads,page_size,qk_head_dim)."; + TVM_FFI_ICHECK(tensor.defined()) << error_msg; + TVM_FFI_ICHECK(tensor.DataType() == kv_dtype_) + << error_msg << " The dtype mismatches, expected " << kv_dtype_ << ", got " + << tensor.DataType() << "."; + TVM_FFI_ICHECK_EQ(tensor->ndim, 6) << error_msg; + TVM_FFI_ICHECK_EQ(tensor->shape[0], 1) << error_msg << " The group count mismatches."; + TVM_FFI_ICHECK_EQ(tensor->shape[1], num_logical_pages) + << error_msg << " The number of logical pages mismatches."; + TVM_FFI_ICHECK_EQ(tensor->shape[2], 2) << error_msg << " The K/V axis mismatches."; + TVM_FFI_ICHECK_EQ(tensor->shape[3], num_kv_heads_) + << error_msg << " The number of KV heads mismatches."; + TVM_FFI_ICHECK_EQ(tensor->shape[4], page_size_) << error_msg << " The page size mismatches."; + TVM_FFI_ICHECK_EQ(tensor->shape[5], qk_head_dim_) + << error_msg << " The head dimension mismatches."; + } + + void CheckExportPageGroupTensor(const Tensor& dst, int64_t num_logical_pages) const { + CheckPageGroupTensor(dst, num_logical_pages, "ExportPageGroup"); + } + + void CheckImportPageGroupTensor(const Tensor& src, int64_t num_logical_pages) const { + CheckPageGroupTensor(src, num_logical_pages, "ImportPageGroup"); + } + + ffi::json::Array MakeAttnKindsMetadata() const { + ffi::json::Array result; + for (int64_t layer_id = layer_id_begin_offset_; layer_id < layer_id_end_offset_; ++layer_id) { + AttnKind attn_kind = attn_kinds_[layer_id]; + result.push_back(ffi::String(AttnKindToString(attn_kind))); + } + return result; + } + + ffi::json::Object MakeLayoutMetadata() const { + namespace json = tvm::ffi::json; + json::Object metadata; + metadata.Set("cacheType", ffi::String(kPagedKVCacheCheckpointRuntime)); + metadata.Set("pageSize", page_size_); + metadata.Set("numLayers", num_layers_); + metadata.Set("layerBegin", layer_id_begin_offset_); + metadata.Set("layerEnd", layer_id_end_offset_); + metadata.Set("numQOHeads", num_qo_heads_); + metadata.Set("numKVHeads", num_kv_heads_); + metadata.Set("qkHeadDim", qk_head_dim_); + metadata.Set("vHeadDim", v_head_dim_); + metadata.Set("numTotalPages", num_total_pages_); + metadata.Set("prefillChunkSize", prefill_chunk_size_); + metadata.Set("dtype", ffi::DLDataTypeToString(kv_dtype_)); + metadata.Set("attnKinds", MakeAttnKindsMetadata()); + metadata.Set("ropeMode", ffi::String(RoPEModeToString(rope_mode_))); + metadata.Set("rotaryScale", rotary_scale_); + metadata.Set("rotaryTheta", rotary_theta_); + metadata.Set("hasRopeExtFactors", rope_ext_factors_.has_value()); + metadata.Set("supportSlidingWindow", support_sliding_window_); + metadata.Set("supportLayerSlidingWindow", support_layer_sliding_window_); + metadata.Set("pageTensorLayout", + ffi::String("num_total_pages,2,num_kv_heads,page_size,qk_head_dim")); + if (!pages_.empty()) { + metadata.Set("pageTensorShape", ShapeToJSON(pages_[0]->shape, pages_[0]->ndim)); + } + return metadata; + } + + std::string GetLayoutDescriptor() const { + std::ostringstream os; + os.imbue(std::locale::classic()); + os << std::setprecision(std::numeric_limits::max_digits10); + os << "cacheType=" << kPagedKVCacheCheckpointRuntime << ";"; + os << "pageSize=" << page_size_ << ";"; + os << "numLayers=" << num_layers_ << ";"; + os << "layerBegin=" << layer_id_begin_offset_ << ";"; + os << "layerEnd=" << layer_id_end_offset_ << ";"; + os << "numQOHeads=" << num_qo_heads_ << ";"; + os << "numKVHeads=" << num_kv_heads_ << ";"; + os << "qkHeadDim=" << qk_head_dim_ << ";"; + os << "vHeadDim=" << v_head_dim_ << ";"; + os << "numTotalPages=" << num_total_pages_ << ";"; + os << "prefillChunkSize=" << prefill_chunk_size_ << ";"; + os << "dtype=" << std::string(ffi::DLDataTypeToString(kv_dtype_)) << ";"; + os << "ropeMode=" << RoPEModeToString(rope_mode_) << ";"; + os << "rotaryScale=" << rotary_scale_ << ";"; + os << "rotaryTheta=" << rotary_theta_ << ";"; + os << "hasRopeExtFactors=" << rope_ext_factors_.has_value() << ";"; + os << "supportSlidingWindow=" << support_sliding_window_ << ";"; + os << "supportLayerSlidingWindow=" << support_layer_sliding_window_ << ";"; + os << "pageTensorLayout=num_total_pages,2,num_kv_heads,page_size,qk_head_dim;"; + os << "attnKinds="; + for (int64_t layer_id = layer_id_begin_offset_; layer_id < layer_id_end_offset_; ++layer_id) { + if (layer_id != layer_id_begin_offset_) { + os << ","; + } + os << AttnKindToString(attn_kinds_[layer_id]); + } + return os.str(); + } + + ffi::json::Array MakeBlockMetadata(const Sequence& seq) const { + namespace json = tvm::ffi::json; + json::Array blocks; + for (int32_t block_id : seq.GetBlockTrace(global_block_pool_)) { + const Block& block = global_block_pool_[block_id]; + json::Object block_json; + block_json.Set("blockIndex", static_cast(block.index)); + block_json.Set("parentBlockIndex", static_cast(block.parent_idx)); + block_json.Set("startPos", static_cast(block.start_pos)); + block_json.Set("seqLength", static_cast(block.seq_length)); + block_json.Set("sinkLength", static_cast(block.sink_length)); + block_json.Set("slidingWindowOffset", static_cast(block.sliding_window_offset)); + block_json.Set("pageIds", IntArrayToJSON(block.page_ids)); + blocks.push_back(block_json); + } + return blocks; + } + + ffi::json::Array MakeLogicalPageMetadata(const Sequence& seq) const { + namespace json = tvm::ffi::json; + json::Array pages; + int64_t logical_page_index = 0; + for (int32_t block_id : seq.GetBlockTrace(global_block_pool_)) { + const Block& block = global_block_pool_[block_id]; + for (int64_t page_index = 0; page_index < static_cast(block.page_ids.size()); + ++page_index) { + int64_t page_start = block.start_pos + page_index * page_size_; + int64_t page_length = + std::min(page_size_, block.seq_length - page_index * page_size_); + TVM_FFI_ICHECK_GT(page_length, 0); + json::Object page_json; + page_json.Set("logicalPageIndex", logical_page_index++); + page_json.Set("blockIndex", static_cast(block.index)); + page_json.Set("pageIndexInBlock", page_index); + page_json.Set("pageId", static_cast(block.page_ids[page_index])); + page_json.Set("startPos", page_start); + page_json.Set("length", page_length); + pages.push_back(page_json); + } + } + return pages; + } + + int64_t GetNumLogicalPages(const Sequence& seq) const { + int64_t num_pages = 0; + for (int32_t block_id : seq.GetBlockTrace(global_block_pool_)) { + num_pages += global_block_pool_[block_id].page_ids.size(); + } + return num_pages; + } + + ffi::json::Array MakePageGroupMetadata(const Sequence& seq) const { + namespace json = tvm::ffi::json; + json::Array groups; + int64_t num_logical_pages = GetNumLogicalPages(seq); + int64_t bytes_per_scalar = (static_cast(kv_dtype_.bits) * kv_dtype_.lanes + 7) / 8; + for (int64_t local_layer = 0; local_layer < num_layers_; ++local_layer) { + json::Object group; + group.Set("groupIndex", local_layer); + group.Set("layerBegin", layer_id_begin_offset_ + local_layer); + group.Set("layerEnd", layer_id_begin_offset_ + local_layer + 1); + group.Set("numLogicalPages", num_logical_pages); + group.Set("dtype", ffi::DLDataTypeToString(kv_dtype_)); + group.Set("shape", + json::Array{1, num_logical_pages, 2, num_kv_heads_, page_size_, qk_head_dim_}); + group.Set("nbytes", num_logical_pages * 2 * num_kv_heads_ * page_size_ * qk_head_dim_ * + bytes_per_scalar); + groups.push_back(group); + } + return groups; + } + /*! \brief Get a new free page and return its id. */ int32_t GetFreePage() { // Find a page from the free page pools. diff --git a/tests/python/relax/test_runtime_builtin_paged_attention_kv_cache_cpu.py b/tests/python/relax/test_runtime_builtin_paged_attention_kv_cache_cpu.py index dfd01351789d..9e7996163955 100644 --- a/tests/python/relax/test_runtime_builtin_paged_attention_kv_cache_cpu.py +++ b/tests/python/relax/test_runtime_builtin_paged_attention_kv_cache_cpu.py @@ -73,6 +73,11 @@ fattention_with_fuse_qkv = None fis_empty = None fdebug_get_kv = None +fget_checkpoint_metadata = None +fexport_page_group = None +fprepare_import = None +fimport_page_group = None +fget_sequence_length = None ftranspose_append = None fcopy_cache = None @@ -96,6 +101,8 @@ def set_global_func(head_dim, dtype): global fclear, fadd_sequence, fremove_sequence, ffork_sequence, fenable_sliding_window_for_seq global fpopn, fbegin_forward, fend_forward, fcommit_accepted_token_tree_nodes global fattention_with_fuse_qkv, fis_empty, fdebug_get_kv + global fget_checkpoint_metadata, fexport_page_group + global fprepare_import, fimport_page_group, fget_sequence_length global ftranspose_append, fcopy_cache, fattn_prefill, fattn_decode global \ fattn_prefill_ragged, \ @@ -122,6 +129,13 @@ def set_global_func(head_dim, dtype): ) fis_empty = tvm.get_global_func("vm.builtin.attention_kv_cache_empty") fdebug_get_kv = tvm.get_global_func("vm.builtin.attention_kv_cache_debug_get_kv") + fget_checkpoint_metadata = tvm.get_global_func( + "vm.builtin.attention_kv_cache_get_checkpoint_metadata" + ) + fexport_page_group = tvm.get_global_func("vm.builtin.attention_kv_cache_export_page_group") + fprepare_import = tvm.get_global_func("vm.builtin.attention_kv_cache_prepare_import") + fimport_page_group = tvm.get_global_func("vm.builtin.attention_kv_cache_import_page_group") + fget_sequence_length = tvm.get_global_func("vm.builtin.attention_kv_cache_get_sequence_length") target = tvm.target.Target.from_device(device) cache_key = ( @@ -279,6 +293,66 @@ def verify_cached_kv(kv_cache, seq_ids, expected_k, expected_v): tvm.testing.assert_allclose(values.numpy(), values_expected, rtol=1e-3, atol=1e-3) +def verify_exported_page_groups(kv_cache, seq_id): + metadata = json.loads(fget_checkpoint_metadata(kv_cache, seq_id)) + seq_length = metadata["seqLength"] + keys = tvm.runtime.empty( + (num_layers, seq_length, num_kv_heads, head_dim), dtype=dtype, device=device + ) + values = tvm.runtime.empty( + (num_layers, seq_length, num_kv_heads, head_dim), dtype=dtype, device=device + ) + fdebug_get_kv(kv_cache, seq_id, 0, seq_length, keys, values) + keys_np = keys.numpy() + values_np = values.numpy() + + for group in metadata["groups"]: + group_data = tvm.runtime.empty(tuple(group["shape"]), dtype=dtype, device=device) + fexport_page_group(kv_cache, seq_id, group["groupIndex"], group_data) + group_np = group_data.numpy() + layer = group["groupIndex"] + for page in metadata["logicalPages"]: + logical_page_index = page["logicalPageIndex"] + start_pos = page["startPos"] + length = page["length"] + exported_k = group_np[0, logical_page_index, 0, :, :length, :].transpose(1, 0, 2) + exported_v = group_np[0, logical_page_index, 1, :, :length, :].transpose(1, 0, 2) + tvm.testing.assert_allclose( + exported_k, keys_np[layer, start_pos : start_pos + length], rtol=1e-3, atol=1e-3 + ) + tvm.testing.assert_allclose( + exported_v, values_np[layer, start_pos : start_pos + length], rtol=1e-3, atol=1e-3 + ) + + +def export_page_groups(kv_cache, metadata): + groups = [] + for group in metadata["groups"]: + group_data = tvm.runtime.empty(tuple(group["shape"]), dtype=dtype, device=device) + fexport_page_group(kv_cache, metadata["seqId"], group["groupIndex"], group_data) + groups.append(group_data) + return groups + + +def verify_debug_kv_equal(lhs_cache, rhs_cache, seq_id, seq_length): + lhs_keys = tvm.runtime.empty( + (num_layers, seq_length, num_kv_heads, head_dim), dtype=dtype, device=device + ) + lhs_values = tvm.runtime.empty( + (num_layers, seq_length, num_kv_heads, head_dim), dtype=dtype, device=device + ) + rhs_keys = tvm.runtime.empty( + (num_layers, seq_length, num_kv_heads, head_dim), dtype=dtype, device=device + ) + rhs_values = tvm.runtime.empty( + (num_layers, seq_length, num_kv_heads, head_dim), dtype=dtype, device=device + ) + fdebug_get_kv(lhs_cache, seq_id, 0, seq_length, lhs_keys, lhs_values) + fdebug_get_kv(rhs_cache, seq_id, 0, seq_length, rhs_keys, rhs_values) + tvm.testing.assert_allclose(lhs_keys.numpy(), rhs_keys.numpy(), rtol=1e-3, atol=1e-3) + tvm.testing.assert_allclose(lhs_values.numpy(), rhs_values.numpy(), rtol=1e-3, atol=1e-3) + + def f_apply_rotary(x, offset, scale, theta, offset_list: list[int] | None = None): # x: (N, H, D) assert len(x.shape) == 3 @@ -593,6 +667,45 @@ def test_paged_attention_kv_cache_prefill_and_decode(kv_cache_and_config): apply_attention(kv_cache, rope_mode, batch, cached_k, cached_v) +def test_paged_attention_kv_cache_export_page_group(): + global head_dim, sm_scale, dtype + head_dim = 64 + dtype = "float32" + sm_scale = head_dim ** (-0.5) + set_global_func(head_dim, dtype) + kv_cache = create_kv_cache(head_dim, dtype, RopeMode.NONE, False) + + cached_k = {} + cached_v = {} + apply_attention(kv_cache, RopeMode.NONE, [(0, page_size * 2 + 3)], cached_k, cached_v) + verify_exported_page_groups(kv_cache, 0) + + +def test_paged_attention_kv_cache_import_page_group_round_trip(): + global head_dim, sm_scale, dtype + head_dim = 64 + dtype = "float32" + sm_scale = head_dim ** (-0.5) + set_global_func(head_dim, dtype) + src_cache = create_kv_cache(head_dim, dtype, RopeMode.NONE, False) + + cached_k = {} + cached_v = {} + apply_attention(src_cache, RopeMode.NONE, [(0, page_size * 2 + 3)], cached_k, cached_v) + + metadata_json = fget_checkpoint_metadata(src_cache, 0) + metadata = json.loads(metadata_json) + groups = export_page_groups(src_cache, metadata) + + dst_cache = create_kv_cache(head_dim, dtype, RopeMode.NONE, False) + fprepare_import(dst_cache, 0, metadata_json) + assert fget_sequence_length(dst_cache, 0) == metadata["seqLength"] + for group, group_data in zip(metadata["groups"], groups): + fimport_page_group(dst_cache, 0, group["groupIndex"], group_data) + + verify_debug_kv_equal(src_cache, dst_cache, 0, metadata["seqLength"]) + + def test_paged_attention_kv_cache_remove_sequence(kv_cache_and_config): kv_cache, rope_mode, support_sliding_window = kv_cache_and_config if support_sliding_window and rope_mode == RopeMode.NORMAL: diff --git a/tests/python/relax/test_runtime_builtin_paged_attention_kv_cache_metadata.py b/tests/python/relax/test_runtime_builtin_paged_attention_kv_cache_metadata.py new file mode 100644 index 000000000000..c6b88774b492 --- /dev/null +++ b/tests/python/relax/test_runtime_builtin_paged_attention_kv_cache_metadata.py @@ -0,0 +1,336 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +import json + +import pytest +import tvm_ffi +from tvm_ffi import Shape + +import tvm +import tvm.testing +from tvm.error import InternalError +from tvm.relax.frontend.nn.llm.kv_cache import AttnKind, RopeMode + +reserved_nseq = 4 +maximum_total_seq_length = 128 +prefill_chunk_size = 64 +page_size = 16 +num_layers = 4 +num_qo_heads = 32 +num_kv_heads = 4 +head_dim = 64 +rope_scale = 1.0 +rope_theta = 1e4 +device = tvm.cpu() + + +def _nop(*args): + return None + + +def create_kv_cache( + *, + dtype="float16", + head_dim_value=head_dim, + v_head_dim_value=None, + page_size_value=page_size, + num_layers_value=num_layers, + rope_mode=RopeMode.NORMAL, + attn_kind=AttnKind.MHA, + support_sliding_window=False, +): + fcreate = tvm.get_global_func("vm.builtin.paged_attention_kv_cache_create") + dummy_func = tvm.runtime.convert(_nop) + return fcreate( + tvm_ffi.Shape( + [ + reserved_nseq, + maximum_total_seq_length, + prefill_chunk_size, + page_size_value, + int(support_sliding_window), + ] + ), + tvm_ffi.Shape([0, num_layers_value]), + num_qo_heads, + num_kv_heads, + head_dim_value, + head_dim_value if v_head_dim_value is None else v_head_dim_value, + tvm_ffi.Shape([int(attn_kind) for _ in range(num_layers_value)]), + False, # enable_kv_transfer + int(rope_mode), + rope_scale, + rope_theta, + None, # rope_ext_factors + tvm.runtime.empty((), dtype, device=device), + dummy_func, # f_transpose_append_mha + None, # f_transpose_append_mla + [], # f_attention_prefill_ragged + [], # f_attention_prefill + [], # f_attention_decode + [], # f_attention_prefill_sliding_window + [], # f_attention_decode_sliding_window + [], # f_attention_prefill_with_tree_mask_paged_kv + [], # f_attention_prefill_with_tree_mask + [], # f_mla_prefill + [dummy_func], # f_merge_inplace + dummy_func, # f_split_rotary + dummy_func, # f_copy_single_page + dummy_func, # f_debug_get_kv + dummy_func, # f_compact_copy + ) + + +def append_tokens(kv_cache, seq_id=0, append_length=page_size + 1): + fadd_sequence = tvm.get_global_func("vm.builtin.kv_state_add_sequence") + fbegin_forward = tvm.get_global_func("vm.builtin.kv_state_begin_forward") + fend_forward = tvm.get_global_func("vm.builtin.kv_state_end_forward") + fadd_sequence(kv_cache, seq_id) + fbegin_forward(kv_cache, Shape([seq_id]), Shape([append_length]), None) + fend_forward(kv_cache) + + +def test_checkpoint_metadata_reports_layout_pages_and_groups(): + fget_checkpoint_metadata = tvm.get_global_func( + "vm.builtin.attention_kv_cache_get_checkpoint_metadata" + ) + fget_layout_hash = tvm.get_global_func("vm.builtin.attention_kv_cache_get_layout_hash") + fexport_page_group = tvm.get_global_func("vm.builtin.attention_kv_cache_export_page_group") + fprepare_import = tvm.get_global_func("vm.builtin.attention_kv_cache_prepare_import") + fimport_page_group = tvm.get_global_func("vm.builtin.attention_kv_cache_import_page_group") + fget_sequence_length = tvm.get_global_func("vm.builtin.attention_kv_cache_get_sequence_length") + + kv_cache = create_kv_cache() + append_tokens(kv_cache) + metadata_json = fget_checkpoint_metadata(kv_cache, 0) + metadata = json.loads(metadata_json) + + assert metadata["cacheType"] == "relax.vm.PagedAttentionKVCache" + assert metadata["layoutHash"] == fget_layout_hash(kv_cache) + assert metadata["seqId"] == 0 + assert metadata["seqLength"] == page_size + 1 + assert metadata["pageSize"] == page_size + assert metadata["dtype"] == "float16" + assert metadata["layerBegin"] == 0 + assert metadata["layerEnd"] == num_layers + assert metadata["numKVHeads"] == num_kv_heads + assert metadata["qkHeadDim"] == head_dim + assert metadata["vHeadDim"] == head_dim + assert metadata["attnKinds"] == ["mha"] * num_layers + assert metadata["pageTensorLayout"] == "num_total_pages,2,num_kv_heads,page_size,qk_head_dim" + assert metadata["pageTensorShape"] == [ + metadata["numTotalPages"], + 2, + num_kv_heads, + page_size, + head_dim, + ] + assert len(metadata["blocks"]) == 1 + assert metadata["blocks"][0]["seqLength"] == page_size + 1 + assert len(metadata["logicalPages"]) == 2 + assert metadata["logicalPages"][0]["startPos"] == 0 + assert metadata["logicalPages"][0]["length"] == page_size + assert metadata["logicalPages"][1]["startPos"] == page_size + assert metadata["logicalPages"][1]["length"] == 1 + assert len(metadata["groups"]) == num_layers + assert metadata["groups"][0]["layerBegin"] == 0 + assert metadata["groups"][0]["layerEnd"] == 1 + assert metadata["groups"][0]["numLogicalPages"] == 2 + assert metadata["groups"][0]["dtype"] == "float16" + assert metadata["groups"][0]["shape"] == [1, 2, 2, num_kv_heads, page_size, head_dim] + + group = tvm.runtime.empty(tuple(metadata["groups"][0]["shape"]), "float16", device=device) + fexport_page_group(kv_cache, 0, 0, group) + import_cache = create_kv_cache() + fprepare_import(import_cache, 0, metadata_json) + assert fget_sequence_length(import_cache, 0) == page_size + 1 + fimport_page_group(import_cache, 0, 0, group) + + +def test_checkpoint_layout_hash_is_stable_and_layout_sensitive(): + fget_layout_hash = tvm.get_global_func("vm.builtin.attention_kv_cache_get_layout_hash") + + kv_cache = create_kv_cache() + same_layout = create_kv_cache() + different_page_size = create_kv_cache(page_size_value=page_size * 2) + different_num_layers = create_kv_cache(num_layers_value=2) + different_head_dim = create_kv_cache(head_dim_value=128) + different_dtype = create_kv_cache(dtype="float32") + different_rope = create_kv_cache(rope_mode=RopeMode.NONE) + + layout_hash = fget_layout_hash(kv_cache) + assert layout_hash == fget_layout_hash(kv_cache) + assert layout_hash == fget_layout_hash(same_layout) + assert layout_hash != fget_layout_hash(different_page_size) + assert layout_hash != fget_layout_hash(different_num_layers) + assert layout_hash != fget_layout_hash(different_head_dim) + assert layout_hash != fget_layout_hash(different_dtype) + assert layout_hash != fget_layout_hash(different_rope) + + +def test_checkpoint_metadata_rejects_unsupported_layouts_and_sequence_ids(): + fget_checkpoint_metadata = tvm.get_global_func( + "vm.builtin.attention_kv_cache_get_checkpoint_metadata" + ) + fget_layout_hash = tvm.get_global_func("vm.builtin.attention_kv_cache_get_layout_hash") + fexport_page_group = tvm.get_global_func("vm.builtin.attention_kv_cache_export_page_group") + + sliding_cache = create_kv_cache(support_sliding_window=True) + append_tokens(sliding_cache) + mla_cache = create_kv_cache(attn_kind=AttnKind.MLA) + asymmetric_cache = create_kv_cache(v_head_dim_value=head_dim // 2) + dst = tvm.runtime.empty((1, 1, 2, num_kv_heads, page_size, head_dim), "float16", device=device) + + with pytest.raises(InternalError, match="sliding-window"): + fget_layout_hash(sliding_cache) + with pytest.raises(InternalError, match="sliding-window"): + fget_checkpoint_metadata(sliding_cache, 0) + with pytest.raises(InternalError, match="sliding-window"): + fexport_page_group(sliding_cache, 0, 0, dst) + with pytest.raises(InternalError, match="sequence id 0"): + fget_checkpoint_metadata(create_kv_cache(), 1) + with pytest.raises(InternalError, match="full-context MHA/GQA"): + fget_layout_hash(mla_cache) + with pytest.raises(InternalError, match="full-context MHA/GQA"): + fexport_page_group(mla_cache, 0, 0, dst) + with pytest.raises(InternalError, match="qk_head_dim to equal v_head_dim"): + fget_layout_hash(asymmetric_cache) + + tree_cache = create_kv_cache() + tvm.get_global_func("vm.builtin.kv_state_add_sequence")(tree_cache, 0) + tvm.get_global_func("vm.builtin.kv_state_begin_forward")( + tree_cache, Shape([0]), Shape([2]), Shape([-1, 0]) + ) + with pytest.raises(InternalError, match="committed token-chain state"): + fexport_page_group(tree_cache, 0, 0, dst) + + +def test_checkpoint_export_page_group_validates_group_shape(): + fget_checkpoint_metadata = tvm.get_global_func( + "vm.builtin.attention_kv_cache_get_checkpoint_metadata" + ) + fexport_page_group = tvm.get_global_func("vm.builtin.attention_kv_cache_export_page_group") + + kv_cache = create_kv_cache() + append_tokens(kv_cache) + metadata = json.loads(fget_checkpoint_metadata(kv_cache, 0)) + shape = metadata["groups"][0]["shape"] + + with pytest.raises(InternalError, match="group id"): + fexport_page_group( + kv_cache, + 0, + num_layers, + tvm.runtime.empty(tuple(shape), "float16", device=device), + ) + + bad_shape = shape.copy() + bad_shape[-1] += 1 + with pytest.raises(InternalError, match="ExportPageGroup expects"): + fexport_page_group( + kv_cache, + 0, + 0, + tvm.runtime.empty(tuple(bad_shape), "float16", device=device), + ) + + with pytest.raises(InternalError, match="dtype mismatches"): + fexport_page_group( + kv_cache, + 0, + 0, + tvm.runtime.empty(tuple(shape), "float32", device=device), + ) + + +def test_checkpoint_prepare_import_validates_metadata(): + fget_checkpoint_metadata = tvm.get_global_func( + "vm.builtin.attention_kv_cache_get_checkpoint_metadata" + ) + fprepare_import = tvm.get_global_func("vm.builtin.attention_kv_cache_prepare_import") + + kv_cache = create_kv_cache() + append_tokens(kv_cache) + metadata_json = fget_checkpoint_metadata(kv_cache, 0) + metadata = json.loads(metadata_json) + + with pytest.raises(InternalError, match="sequence id 0"): + fprepare_import(create_kv_cache(), 1, metadata_json) + + with pytest.raises(InternalError, match="dtype"): + fprepare_import(create_kv_cache(dtype="float32"), 0, metadata_json) + + bad_length = json.loads(metadata_json) + bad_length["seqLength"] = page_size * 2 + 1 + with pytest.raises(InternalError, match="sequence length"): + fprepare_import(create_kv_cache(), 0, json.dumps(bad_length)) + + bad_group = json.loads(metadata_json) + bad_group["groups"][0]["shape"][-1] += 1 + with pytest.raises(InternalError, match="shape"): + fprepare_import(create_kv_cache(), 0, json.dumps(bad_group)) + + metadata["layoutHash"] = "bad-layout-hash" + with pytest.raises(InternalError, match="layout hash mismatch"): + fprepare_import(create_kv_cache(), 0, json.dumps(metadata)) + + +def test_checkpoint_import_page_group_validates_group_shape(): + fget_checkpoint_metadata = tvm.get_global_func( + "vm.builtin.attention_kv_cache_get_checkpoint_metadata" + ) + fprepare_import = tvm.get_global_func("vm.builtin.attention_kv_cache_prepare_import") + fimport_page_group = tvm.get_global_func("vm.builtin.attention_kv_cache_import_page_group") + + kv_cache = create_kv_cache() + append_tokens(kv_cache) + metadata_json = fget_checkpoint_metadata(kv_cache, 0) + metadata = json.loads(metadata_json) + shape = metadata["groups"][0]["shape"] + + import_cache = create_kv_cache() + fprepare_import(import_cache, 0, metadata_json) + + with pytest.raises(InternalError, match="group id"): + fimport_page_group( + import_cache, + 0, + num_layers, + tvm.runtime.empty(tuple(shape), "float16", device=device), + ) + + bad_shape = shape.copy() + bad_shape[-1] += 1 + with pytest.raises(InternalError, match="ImportPageGroup expects"): + fimport_page_group( + import_cache, + 0, + 0, + tvm.runtime.empty(tuple(bad_shape), "float16", device=device), + ) + + with pytest.raises(InternalError, match="dtype mismatches"): + fimport_page_group( + import_cache, + 0, + 0, + tvm.runtime.empty(tuple(shape), "float32", device=device), + ) + + +if __name__ == "__main__": + tvm.testing.main()