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) {
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;
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) {
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) {
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;
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);
}
}
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");
}
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);
}
}
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);
}
}
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);
}
}
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);
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);