llama_dsv4: write only used rows in state (#25325)

* llama_dsv4: write only used rows in state

* add TODO about conflating token pos with kv rows
This commit is contained in:
Aman Gupta
2026-07-20 22:43:39 +08:00
committed by GitHub
parent 4ee6a9af71
commit 91d2fc3875

View 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);