| 1489 | } |
| 1490 | |
| 1491 | bool llama_kv_cache_unified::state_read_data(llama_io_read_i & io, uint32_t cell_count) { |
| 1492 | uint32_t v_trans; |
| 1493 | uint32_t n_layer; |
| 1494 | |
| 1495 | io.read_to(&v_trans, sizeof(v_trans)); |
| 1496 | io.read_to(&n_layer, sizeof(n_layer)); |
| 1497 | |
| 1498 | if (n_layer != layers.size()) { |
| 1499 | LLAMA_LOG_ERROR("%s: mismatched layer count (%u instead of %u)\n", __func__, n_layer, (uint32_t) layers.size()); |
| 1500 | return false; |
| 1501 | } |
| 1502 | |
| 1503 | if (cell_count > cells.size()) { |
| 1504 | LLAMA_LOG_ERROR("%s: not enough cells in kv cache to restore state (%u > %u)\n", __func__, cell_count, cells.size()); |
| 1505 | return false; |
| 1506 | } |
| 1507 | |
| 1508 | if (this->v_trans != (bool) v_trans) { |
| 1509 | LLAMA_LOG_ERROR("%s: incompatible V transposition\n", __func__); |
| 1510 | return false; |
| 1511 | } |
| 1512 | |
| 1513 | // For each layer, read the keys for each cell, one row is one cell, read as one contiguous block |
| 1514 | for (const auto & layer : layers) { |
| 1515 | const uint32_t il = layer.il; |
| 1516 | |
| 1517 | const uint32_t n_embd_k_gqa = hparams.n_embd_k_gqa(il) + hparams.n_embd_k_s(); |
| 1518 | |
| 1519 | // Read type of key |
| 1520 | int32_t k_type_i_ref; |
| 1521 | io.read_to(&k_type_i_ref, sizeof(k_type_i_ref)); |
| 1522 | const int32_t k_type_i = (int32_t) layer.k->type; |
| 1523 | if (k_type_i != k_type_i_ref) { |
| 1524 | LLAMA_LOG_ERROR("%s: mismatched key type (%d != %d, layer %d)\n", __func__, k_type_i, k_type_i_ref, il); |
| 1525 | return false; |
| 1526 | } |
| 1527 | |
| 1528 | // Read row size of key |
| 1529 | uint64_t k_size_row_ref; |
| 1530 | io.read_to(&k_size_row_ref, sizeof(k_size_row_ref)); |
| 1531 | const size_t k_size_row = ggml_row_size(layer.k->type, n_embd_k_gqa); |
| 1532 | if (k_size_row != k_size_row_ref) { |
| 1533 | LLAMA_LOG_ERROR("%s: mismatched key row size (%zu != %zu, layer %d)\n", __func__, k_size_row, (size_t) k_size_row_ref, il); |
| 1534 | return false; |
| 1535 | } |
| 1536 | |
| 1537 | if (cell_count) { |
| 1538 | // Read and set the keys for the whole cell range |
| 1539 | ggml_backend_tensor_set(layer.k, io.read(cell_count * k_size_row), head * k_size_row, cell_count * k_size_row); |
| 1540 | } |
| 1541 | } |
| 1542 | |
| 1543 | if (!this->v_trans) { |
| 1544 | for (const auto & layer : layers) { |
| 1545 | const uint32_t il = layer.il; |
| 1546 | |
| 1547 | const uint32_t n_embd_v_gqa = hparams.n_embd_v_gqa(il) + hparams.n_embd_v_s(); |
| 1548 |
nothing calls this directly
no test coverage detected