#include "strata/core/conversation_snapshot.hpp" #include "strata/kernels/kv_q4.hpp" #include "conversation_checked.hpp" #include #include #include #include namespace strata::core { namespace { struct Layout { int format; int64_t cells, pooled_rows, page_size; size_t data, scales, pooled, value_data, value_scales; }; bool valid_extent(const QsaState& st, int64_t upto, std::string& error) { const int64_t page = strata::kernels::qsa_real_shapes().page_size; if (upto < 0 || upto > st.max_cells || upto > std::numeric_limits::max() - (page - 1)) { error = "conversation snapshot: invalid K/V extent"; return false; } return true; } bool layout(const QsaState& st, const ModelGeometry& g, int64_t upto, bool index, Layout& l, std::string& error) { if (!valid_extent(st, upto, error)) return false; if (st.kv_hybrid && (st.kv_mode != 0 || st.kv_q4 || st.kv_int8)) { error = "conversation snapshot: hybrid K8V4 requires an identity layout and distinct format flags"; return false; } const auto s = strata::kernels::qsa_real_shapes(); const bool int8_keys = st.kv_int8 || st.kv_hybrid; if (g.n_head_kv <= 0 || g.head_dim <= 0 || g.head_dim > INT32_MAX || g.idx_key_dim <= 0 || (st.kv_q4 && g.head_dim % 32) || (int8_keys && !st.kv_q4 && g.head_dim % 64)) { error = "conversation snapshot: invalid K/V geometry"; return false; } const int64_t cells = ((upto + s.page_size - 1) / s.page_size) * s.page_size; const size_t per = st.kv_q4 ? (size_t) strata::kernels::kv_q4_bytes_per_head((int) g.head_dim) : (size_t) g.head_dim * (int8_keys ? 1 : 2); // Include the moving spare row, not only completed blocks. The checkpoint // restore reconstructs that row when rewinding to an earlier prefix. const int64_t pooled = index && upto > 0 ? upto / s.idx_block + 1 : 0; // Format 3 identifies snapshot K8V4 only; never pass it to the block movers. Rotated INT8 K/V (#293, // STRATA_KV_ROT=1) is another format too (+16): its bytes mean nothing to a state that does not rotate, so a // snapshot or prompt checkpoint taken with the rotation never restores into one without it, nor the reverse. // (Q4_0 is always rotated: nothing to tell apart; without the rotation the format is the one it always was.) const int rotated = st.kv_rot && !st.kv_q4 && !st.kv_hybrid ? 16 : 0; l = {(st.kv_hybrid ? 3 : qsa_kv_format(st)) + rotated, cells, pooled, s.page_size, 0, 0, 0, 0, 0}; using conversation_detail::product; if (!product(l.data, {(uint64_t) cells, (uint64_t) g.n_head_kv, per}) || !product(l.scales, {(uint64_t) cells, (uint64_t) g.n_head_kv, int8_keys && !st.kv_q4 ? (uint64_t) (g.head_dim / 64) * 2 : 0}) || !product(l.pooled, {(uint64_t) pooled, (uint64_t) g.idx_key_dim, sizeof(float)})) { error = "conversation snapshot: K/V byte count overflow"; return false; } l.value_data = l.data; l.value_scales = l.scales; if (st.kv_hybrid) { if (!product(l.value_data, {(uint64_t) cells, (uint64_t) g.n_head_kv, (uint64_t) strata::kernels::kv_q4_bytes_per_head((int) g.head_dim)})) { error = "conversation snapshot: hybrid V byte count overflow"; return false; } l.value_scales = 0; } return true; } std::array pools(const QsaState& st, bool resident = false) { const bool host = st.kv_mode != 0 && !resident; if (st.kv_hybrid) return {st.k_q, st.v_q4, st.k_scale, nullptr, st.idx_pooled}; if (st.kv_q4) return {host ? st.host.k_q4 : st.k_q4, host ? st.host.v_q4 : st.v_q4, nullptr, nullptr, st.idx_pooled}; if (st.kv_int8) return {host ? st.host.k_q : st.k_q, host ? st.host.v_q : st.v_q, host ? st.host.k_scale : st.k_scale, host ? st.host.v_scale : st.v_scale, st.idx_pooled}; return {host ? st.host.k_pool : st.k_pool, host ? st.host.v_pool : st.v_pool, nullptr, nullptr, st.idx_pooled}; } bool valid(const QsaState& st, const Layout& l, int64_t upto, std::string& error) { if (upto < 0 || upto > st.max_cells || st.n_pages < 0 || l.cells / l.page_size > st.n_pages || st.kv_mode < 0 || st.kv_mode > 2 || st.n_slots <= 0 || (st.kv_mode == 0 && st.n_slots < st.n_pages) || l.pooled_rows > st.idx_pooled_rows || (st.kv_mode != 0 && !st.host.present())) { error = "conversation snapshot: invalid K/V extent or missing authoritative host pool"; return false; } if (st.kv_mode == 1) { const auto& m = st.map; if (m.n_blocks != st.n_pages || m.n_slots != st.n_slots || !m.page_table || !m.slot_block || !m.slot_stamp || !m.slot_ref || !m.ctl || !m.miss_block || !m.miss_slot) { error = "conversation snapshot: invalid streaming map"; return false; } } if (st.kv_mode == 2 && (!st.page_table || (st.kv_q4 ? !st.k_q4 || !st.v_q4 : st.kv_int8 ? !st.k_q || !st.v_q || !st.k_scale || !st.v_scale : !st.k_pool || !st.v_pool))) { error = "conversation snapshot: invalid draft ring"; return false; } return true; } bool transfer(void* dst, const void* src, size_t n, std::string& error) { if (!n) return true; if (!src || !dst) { error = "conversation snapshot: missing state buffer"; return false; } // Default handles both device allocations and device-mapped host pool aliases. const cudaError_t e = cudaMemcpy(dst, src, n, cudaMemcpyDefault); if (e == cudaSuccess) return true; error = std::string("conversation snapshot copy: ") + cudaGetErrorString(e); return false; } } // namespace bool conversation_kv_capture_bytes(const ConversationKv& image, const QsaState& st, const ModelGeometry& g, int64_t upto, bool index, size_t& bytes, std::string& error) { Layout l{}; if (!layout(st, g, upto, index, l, error) || !valid(st, l, upto, error)) return false; const std::array sizes = {l.data, l.value_data, l.scales, l.value_scales, l.pooled}; const std::array buffers = {&image.k, &image.v, &image.k_scale, &image.v_scale, &image.pooled}; bytes = 0; for (size_t i = 0; i < buffers.size(); ++i) { const size_t n = buffers[i]->allocation_peak(sizes[i]); if (n == SIZE_MAX || !conversation_detail::add(bytes, n)) { error = "conversation snapshot: segmented K/V allocation overflow"; return false; } } return true; } size_t conversation_kv_bytes(const QsaState& st, const ModelGeometry& g, int64_t upto, bool index) { size_t bytes = 0; std::string error; return conversation_kv_capture_bytes({}, st, g, upto, index, bytes, error) ? bytes : 0; } bool conversation_kv_save(ConversationKv& image, const QsaState& st, const ModelGeometry& g, int64_t upto, bool index, std::string& error, int64_t unchanged_tokens, size_t* reused_bytes) { Layout l{}; if (!layout(st, g, upto, index, l, error)) return false; if (!valid(st, l, upto, error)) return false; if (unchanged_tokens < 0 || unchanged_tokens > upto || unchanged_tokens > image.cells) { error = "conversation snapshot: invalid unchanged prefix"; return false; } const int64_t whole_cells = (unchanged_tokens / l.page_size) * l.page_size; image.format = l.format; image.cells = l.cells; image.heads = g.n_head_kv; image.head_dim = g.head_dim; image.page_size = l.page_size; image.pooled_rows = l.pooled_rows; image.idx_dim = g.idx_key_dim; const auto src = pools(st); const std::array sizes = {l.data, l.value_data, l.scales, l.value_scales, l.pooled}; const std::array dst = {&image.k, &image.v, &image.k_scale, &image.v_scale, &image.pooled}; for (size_t i = 0; i < dst.size(); ++i) { // Recopy the partial page and the indexer's moving spare row. Completed // pages/rows strictly before the first rewritten token remain identical. const size_t keep = i == 4 ? (index ? size_t(unchanged_tokens / strata::kernels::qsa_real_shapes().idx_block) * g.idx_key_dim * 4 : 0) : l.cells ? (sizes[i] / size_t(l.cells)) * size_t(whole_cells) : 0; if (keep > dst[i]->size()) { error = "conversation snapshot: missing reusable prefix"; return false; } dst[i]->resize(sizes[i]); if (!dst[i]->visit(keep, sizes[i] - keep, [&](uint8_t* p, size_t n, size_t at) { return transfer(p, src[i] ? static_cast(src[i]) + at : nullptr, n, error); })) return false; if (reused_bytes) *reused_bytes += keep; } return true; } bool conversation_kv_validate(const ConversationKv& image, const QsaState& st, const ModelGeometry& g, int64_t upto, bool index, std::string& error) { Layout l{}; if (!layout(st, g, upto, index, l, error)) return false; if (!valid(st, l, upto, error)) return false; const std::array sizes = {l.data, l.value_data, l.scales, l.value_scales, l.pooled}; const std::array src = {&image.k, &image.v, &image.k_scale, &image.v_scale, &image.pooled}; if (image.format != l.format || image.cells != l.cells || image.heads != g.n_head_kv || image.head_dim != g.head_dim || image.page_size != l.page_size || image.pooled_rows != l.pooled_rows || image.idx_dim != g.idx_key_dim) { error = "conversation snapshot: incompatible K/V geometry"; return false; } for (size_t i = 0; i < src.size(); ++i) if (src[i]->size() != sizes[i]) { error = "conversation snapshot: invalid K/V payload size"; return false; } const auto dst = pools(st); for (size_t i = 0; i < src.size(); ++i) if (sizes[i] && !dst[i]) { error = "conversation snapshot: missing target state buffer"; return false; } return true; } bool conversation_kv_restore(const ConversationKv& image, const QsaState& st, const ModelGeometry& g, int64_t upto, bool index, std::string& error) { if (!conversation_kv_validate(image, st, g, upto, index, error)) return false; const std::array src = {&image.k, &image.v, &image.k_scale, &image.v_scale, &image.pooled}; const auto dst = pools(st); for (size_t i = 0; i < src.size(); ++i) if (!src[i]->visit(0, src[i]->size(), [&](const uint8_t* p, size_t n, size_t at) { return transfer(static_cast(dst[i]) + at, p, n, error); })) return false; // VRAM slots still contain the outgoing conversation. Resolve must refill // them from the restored authoritative pools before any attention reads. if (st.kv_mode == 1) strata::kernels::kv_stream_reset(st.map, nullptr); if (st.kv_mode == 2 && upto > 0) { auto shapes = strata::kernels::qsa_real_shapes(); shapes.n_head_kv = g.n_head_kv; shapes.head_dim = g.head_dim; const int64_t end = (upto + shapes.page_size - 1) / shapes.page_size; strata::kernels::kv_ring_restore(qsa_attn_pools(st), st.host, qsa_kv_format(st), std::max(0, end - st.n_slots), end, st.n_slots, shapes, nullptr); } const auto status = cudaGetLastError(); if (status == cudaSuccess) return true; error = std::string("conversation snapshot residency restore: ") + cudaGetErrorString(status); return false; } bool conversation_kv_verify(const ConversationKv& image, const QsaState& st, const ModelGeometry& g, int64_t upto, bool index, uint64_t& fingerprint, std::string& error) { if (!conversation_kv_validate(image, st, g, upto, index, error)) return false; std::array buffer; uint64_t hash = 1469598103934665603ull; auto compare = [&](const void* device, const uint8_t* expected, size_t bytes, bool include_hash) { for (size_t offset = 0; offset < bytes;) { const size_t n = std::min(buffer.size(), bytes - offset); if (!transfer(buffer.data(), static_cast(device) + offset, n, error)) return false; if (std::memcmp(buffer.data(), expected + offset, n) != 0) { error = "conversation snapshot: restored K/V bytes differ"; return false; } if (include_hash) for (size_t i = 0; i < n; ++i) { hash ^= buffer[i]; hash *= 1099511628211ull; } offset += n; } return true; }; const std::array saved = {&image.k, &image.v, &image.k_scale, &image.v_scale, &image.pooled}; const auto authoritative = pools(st); for (size_t i = 0; i < saved.size(); ++i) if (!saved[i]->visit(0, saved[i]->size(), [&](const uint8_t* p, size_t n, size_t at) { return compare(static_cast(authoritative[i]) + at, p, n, true); })) return false; if (st.kv_mode == 2 && image.cells > 0) { const auto resident = pools(st, true); const int64_t end = image.cells / image.page_size; const int64_t begin = std::max(0, end - st.n_slots); for (size_t i = 0; i < 4; ++i) { if (saved[i]->empty()) continue; const size_t page_bytes = saved[i]->size() / size_t(end); for (int64_t page = begin; page < end; ++page) if (!saved[i]->visit(size_t(page) * page_bytes, page_bytes, [&](const uint8_t* p, size_t n, size_t at) { const size_t offset = size_t(page % st.n_slots) * page_bytes + at - size_t(page) * page_bytes; return compare(static_cast(resident[i]) + offset, p, n, false); })) return false; } } fingerprint = hash; return true; } } // namespace strata::core