]> git.djapps.eu Git - pkg/ggml/sources/llama.cpp/commitdiff
llama_dsv4: write only used rows in state (#25325)
authorAman Gupta <redacted>
Mon, 20 Jul 2026 14:43:39 +0000 (22:43 +0800)
committerGitHub <redacted>
Mon, 20 Jul 2026 14:43:39 +0000 (22:43 +0800)
* llama_dsv4: write only used rows in state

* add TODO about conflating token pos with kv rows

src/llama-kv-cache-dsv4.cpp

index 7cb6cc18dac37b8240298956cf7ece61c503cf30..069da45f4ea34b66727b06449d35efc4b0c4443f 100644 (file)
@@ -22,7 +22,7 @@ static constexpr uint32_t DSV4_STATE_MAGIC         = 0x34565344; // DSV4
 static constexpr uint32_t DSV4_STATE_VERSION       = 1;
 static constexpr uint32_t DSV4_STATE_MODE_FULL     = 0;
 static constexpr uint32_t DSV4_STATE_MODE_PARTIAL  = 1;
-static constexpr uint32_t DSV4_K_CACHE_STATE_VER   = 1;
+static constexpr uint32_t DSV4_K_CACHE_STATE_VER   = 2;
 static constexpr uint32_t DSV4_COMP_STATE_VER      = 1;
 
 static uint32_t dsv4_comp_size(uint32_t kv_size, uint32_t ratio) {
@@ -38,6 +38,16 @@ static void dsv4_clear_tensor_stream(ggml_tensor * tensor, uint32_t stream) {
     ggml_backend_tensor_memset(tensor, 0, stream*stream_size, stream_size);
 }
 
+static uint32_t dsv4_state_n_used_k_rows(llama_pos pos_max, uint32_t ratio, uint32_t kv_size) {
+    if (pos_max < 0) {
+        return 0;
+    }
+
+    const uint64_t n_rows = ((uint64_t) pos_max + 1)/ratio;
+
+    return (uint32_t) std::min<uint64_t>(kv_size, n_rows);
+}
+
 static int64_t dsv4_stream_offset(uint32_t n_stream, llama_seq_id seq_id, uint32_t size) {
     if (n_stream <= 1) {
         return 0;
@@ -239,6 +249,7 @@ static void dsv4_state_dst_stream_range(
 static void dsv4_state_write_tensor_streams(
         llama_io_write_i & io,
         ggml_tensor      * tensor,
+        uint32_t           tensor_rows,
         uint32_t           n_rows,
         uint32_t           s0,
         uint32_t           ns) {
@@ -247,20 +258,31 @@ static void dsv4_state_write_tensor_streams(
     const uint64_t rows     = n_rows;
     const uint64_t row_size = ggml_row_size(tensor->type, tensor->ne[0]);
 
+    if (n_rows > tensor_rows) {
+        throw std::runtime_error("DSV4 state tensor row count exceeds storage");
+    }
+
     io.write(&type_i,   sizeof(type_i));
     io.write(&ne0,      sizeof(ne0));
     io.write(&rows,     sizeof(rows));
     io.write(&row_size, sizeof(row_size));
 
-    const size_t offset = (size_t) s0*n_rows*row_size;
-    const size_t size   = (size_t) ns*n_rows*row_size;
+    const size_t stream_stride = (size_t) tensor_rows*row_size;
+    const size_t size          = (size_t) n_rows*row_size;
+    if (size == 0) {
+        return;
+    }
 
-    io.write_tensor(tensor, offset, size);
+    for (uint32_t s = 0; s < ns; ++s) {
+        const size_t offset = (size_t) (s0 + s)*stream_stride;
+        io.write_tensor(tensor, offset, size);
+    }
 }
 
 static void dsv4_state_read_tensor_streams(
         llama_io_read_i & io,
         ggml_tensor     * tensor,
+        uint32_t          tensor_rows,
         uint32_t          n_rows,
         uint32_t          s0,
         uint32_t          ns) {
@@ -282,18 +304,28 @@ static void dsv4_state_read_tensor_streams(
     if (type_i != type_i_ref || ne0 != ne0_ref || rows != rows_ref || row_size != row_size_ref) {
         throw std::runtime_error("DSV4 state tensor metadata mismatch");
     }
+    if (n_rows > tensor_rows) {
+        throw std::runtime_error("DSV4 state tensor row count exceeds storage");
+    }
 
-    const size_t offset = (size_t) s0*n_rows*row_size;
-    const size_t size   = (size_t) ns*n_rows*row_size;
+    const size_t stream_stride = (size_t) tensor_rows*row_size;
+    const size_t size          = (size_t) n_rows*row_size;
+    if (size == 0) {
+        return;
+    }
 
-    io.read_tensor(tensor, offset, size);
+    for (uint32_t s = 0; s < ns; ++s) {
+        const size_t offset = (size_t) (s0 + s)*stream_stride;
+        io.read_tensor(tensor, offset, size);
+    }
 }
 
 static void dsv4_state_write_k_cache(
         llama_io_write_i    & io,
         const llama_kv_cache * kv,
         llama_seq_id          seq_id,
-        llama_state_seq_flags flags) {
+        llama_state_seq_flags flags,
+        uint32_t              n_rows) {
     GGML_UNUSED(flags);
 
     uint32_t s0;
@@ -305,14 +337,18 @@ static void dsv4_state_write_k_cache(
     const auto layer_ids = kv->get_layer_ids();
     const uint32_t n_layer = layer_ids.size();
 
+    if (n_rows > kv_size) {
+        throw std::runtime_error("DSV4 K-cache state row count exceeds cache size");
+    }
+
     io.write(&version, sizeof(version));
-    io.write(&kv_size, sizeof(kv_size));
+    io.write(&n_rows,  sizeof(n_rows));
     io.write(&ns,      sizeof(ns));
     io.write(&n_layer, sizeof(n_layer));
 
     for (uint32_t il : layer_ids) {
         io.write(&il, sizeof(il));
-        dsv4_state_write_tensor_streams(io, kv->get_k_storage(il), kv_size, s0, ns);
+        dsv4_state_write_tensor_streams(io, kv->get_k_storage(il), kv_size, n_rows, s0, ns);
     }
 }
 
@@ -324,19 +360,26 @@ static void dsv4_state_read_k_cache(
     GGML_UNUSED(flags);
 
     uint32_t version;
-    uint32_t kv_size_ref;
+    uint32_t n_rows_ref;
     uint32_t ns;
     uint32_t n_layer_ref;
 
     io.read(&version,     sizeof(version));
-    io.read(&kv_size_ref, sizeof(kv_size_ref));
+    io.read(&n_rows_ref,  sizeof(n_rows_ref));
     io.read(&ns,          sizeof(ns));
     io.read(&n_layer_ref, sizeof(n_layer_ref));
 
-    if (version != DSV4_K_CACHE_STATE_VER) {
+    if (version != 1 && version != DSV4_K_CACHE_STATE_VER) {
         throw std::runtime_error("DSV4 K-cache state version mismatch");
     }
-    if (kv_size_ref != kv->get_size()) {
+
+    const uint32_t kv_size = kv->get_size();
+    if (version == 1 && n_rows_ref != kv_size) {
+        LLAMA_LOG_INFO("kv size ref %d kv %d\n", n_rows_ref, kv_size);
+        throw std::runtime_error("DSV4 K-cache state size mismatch");
+    }
+    if (n_rows_ref > kv_size) {
+        LLAMA_LOG_INFO("kv rows ref %d kv %d\n", n_rows_ref, kv_size);
         throw std::runtime_error("DSV4 K-cache state size mismatch");
     }
 
@@ -355,7 +398,7 @@ static void dsv4_state_read_k_cache(
             throw std::runtime_error("DSV4 K-cache layer id mismatch");
         }
 
-        dsv4_state_read_tensor_streams(io, kv->get_k_storage(il), kv->get_size(), s0, ns);
+        dsv4_state_read_tensor_streams(io, kv->get_k_storage(il), kv_size, n_rows_ref, s0, ns);
     }
 }
 
@@ -882,8 +925,8 @@ void llama_dsv4_comp_state::state_write(llama_io_write_i & io, llama_seq_id seq_
     for (const auto & layer : layers) {
         io.write(&layer.il, sizeof(layer.il));
 
-        dsv4_state_write_tensor_streams(io, layer.kv,    state_size, s0, ns);
-        dsv4_state_write_tensor_streams(io, layer.score, state_size, s0, ns);
+        dsv4_state_write_tensor_streams(io, layer.kv,    state_size, state_size, s0, ns);
+        dsv4_state_write_tensor_streams(io, layer.score, state_size, state_size, s0, ns);
     }
 }
 
@@ -924,8 +967,8 @@ void llama_dsv4_comp_state::state_read(llama_io_read_i & io, llama_seq_id seq_id
             throw std::runtime_error("DSV4 compressor state layer id mismatch");
         }
 
-        dsv4_state_read_tensor_streams(io, layer.kv,    state_size, s0, ns);
-        dsv4_state_read_tensor_streams(io, layer.score, state_size, s0, ns);
+        dsv4_state_read_tensor_streams(io, layer.kv,    state_size, state_size, s0, ns);
+        dsv4_state_read_tensor_streams(io, layer.score, state_size, state_size, s0, ns);
     }
 }
 
@@ -1328,9 +1371,19 @@ void llama_kv_cache_dsv4::state_write(llama_io_write_i & io, llama_seq_id seq_id
     kv_raw->state_write(io, seq_id, flags);
 
     if (!partial_only) {
-        dsv4_state_write_k_cache(io, kv_csa.get(), seq_id, flags);
-        dsv4_state_write_k_cache(io, kv_hca.get(), seq_id, flags);
-        dsv4_state_write_k_cache(io, kv_lid.get(), seq_id, flags);
+        const llama_pos pos_max = seq_id >= 0 ? kv_raw->seq_pos_max(seq_id) : -1;
+
+        //FIXME : note that we conflate token positions with rows, which is not true for multi-modal case.
+        const uint32_t n_rows_csa = seq_id >= 0 ?
+            dsv4_state_n_used_k_rows(pos_max, DSV4_CSA_RATIO, kv_csa->get_size()) : kv_csa->get_size();
+        const uint32_t n_rows_hca = seq_id >= 0 ?
+            dsv4_state_n_used_k_rows(pos_max, DSV4_HCA_RATIO, kv_hca->get_size()) : kv_hca->get_size();
+        const uint32_t n_rows_lid = seq_id >= 0 ?
+            dsv4_state_n_used_k_rows(pos_max, DSV4_CSA_RATIO, kv_lid->get_size()) : kv_lid->get_size();
+
+        dsv4_state_write_k_cache(io, kv_csa.get(), seq_id, flags, n_rows_csa);
+        dsv4_state_write_k_cache(io, kv_hca.get(), seq_id, flags, n_rows_hca);
+        dsv4_state_write_k_cache(io, kv_lid.get(), seq_id, flags, n_rows_lid);
     }
 
     csa_state->state_write(io, seq_id, flags);
@@ -1366,6 +1419,10 @@ void llama_kv_cache_dsv4::state_read(llama_io_read_i & io, llama_seq_id seq_id,
     kv_raw->state_read(io, seq_id, flags);
 
     if (!partial_only) {
+        kv_csa->clear(true);
+        kv_hca->clear(true);
+        kv_lid->clear(true);
+
         dsv4_state_read_k_cache(io, kv_csa.get(), seq_id, flags);
         dsv4_state_read_k_cache(io, kv_hca.get(), seq_id, flags);
         dsv4_state_read_k_cache(io, kv_lid.get(), seq_id, flags);