| 105 | } |
| 106 | |
| 107 | void KVCache::update_from_graph(CactusGraph* gb, const std::vector<size_t>& k_nodes, |
| 108 | const std::vector<size_t>& v_nodes, size_t seq_len, |
| 109 | size_t layers) { |
| 110 | size_t old_seq_len = current_seq_len; |
| 111 | size_t new_total_len = old_seq_len + seq_len; |
| 112 | |
| 113 | total_seq_len += seq_len; |
| 114 | |
| 115 | size_t effective_seq_len; |
| 116 | bool use_sliding_window = (window_size > 0 && new_total_len > window_size); |
| 117 | |
| 118 | if (use_sliding_window) { |
| 119 | effective_seq_len = window_size; |
| 120 | } else { |
| 121 | effective_seq_len = new_total_len; |
| 122 | } |
| 123 | |
| 124 | bool any_layer_updated = false; |
| 125 | |
| 126 | for (size_t layer_idx = 0; layer_idx < layers; layer_idx++) { |
| 127 | if (k_nodes[layer_idx] == 0 || v_nodes[layer_idx] == 0) |
| 128 | continue; |
| 129 | |
| 130 | auto& cache = layer_caches[layer_idx]; |
| 131 | |
| 132 | void* k_output = gb->get_output(k_nodes[layer_idx]); |
| 133 | void* v_output = gb->get_output(v_nodes[layer_idx]); |
| 134 | |
| 135 | if (k_output && v_output) { |
| 136 | const auto& k_buffer = gb->get_output_buffer(k_nodes[layer_idx]); |
| 137 | const auto& v_buffer = gb->get_output_buffer(v_nodes[layer_idx]); |
| 138 | |
| 139 | size_t dim = get_layer_head_dim(layer_idx); |
| 140 | size_t kv_heads = get_layer_kv_heads(layer_idx); |
| 141 | size_t elements_per_token = kv_heads * dim; |
| 142 | size_t num_groups = (dim + KV_QUANT_GROUP_SIZE - 1) / KV_QUANT_GROUP_SIZE; |
| 143 | size_t scales_per_token = kv_heads * num_groups; |
| 144 | size_t bytes_per_token = elements_per_token * element_size; |
| 145 | |
| 146 | size_t expected_elements = new_total_len * elements_per_token; |
| 147 | |
| 148 | if (k_buffer.total_size == expected_elements && v_buffer.total_size == expected_elements) { |
| 149 | any_layer_updated = true; |
| 150 | |
| 151 | if (!use_sliding_window) { |
| 152 | size_t total_bytes = new_total_len * bytes_per_token; |
| 153 | cache.keys.resize(total_bytes); |
| 154 | cache.values.resize(total_bytes); |
| 155 | |
| 156 | if (precision == Precision::INT8) { |
| 157 | size_t num_scales = new_total_len * scales_per_token; |
| 158 | cache.key_scales.resize(num_scales); |
| 159 | cache.value_scales.resize(num_scales); |
| 160 | |
| 161 | cactus_quantize_kv_fp16_to_int8( |
| 162 | static_cast<const __fp16*>(k_output), |
| 163 | reinterpret_cast<int8_t*>(cache.keys.data()), |
| 164 | cache.key_scales.data(), |