MCPcopy Create free account
hub / github.com/cactus-compute/cactus / update_from_graph

Method update_from_graph

cactus/engine/engine_cache.cpp:107–349  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

105}
106
107void 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(),

Callers 6

update_kv_cacheMethod · 0.80
post_execute_updatesMethod · 0.80
test_incremental_updatesFunction · 0.80
test_reset_functionalityFunction · 0.80
test_large_windowFunction · 0.80

Calls 5

get_output_bufferMethod · 0.80
dataMethod · 0.80
sizeMethod · 0.80
get_outputMethod · 0.45

Tested by 4

test_incremental_updatesFunction · 0.64
test_reset_functionalityFunction · 0.64
test_large_windowFunction · 0.64