-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathkv_cache_ops.cpp
More file actions
127 lines (106 loc) · 4.43 KB
/
Copy pathkv_cache_ops.cpp
File metadata and controls
127 lines (106 loc) · 4.43 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
/*
* Stage 5: KV layout optimization with contiguous cache
* Optimizes KV cache memory layout for better cache utilization
*/
#include <torch/extension.h>
#include <omp.h>
#include <cmath>
// Optimized attention decode pass with contiguous KV cache
// For single token generation, we only compute attention for the new token
torch::Tensor attention_decode_contiguous_kv(
torch::Tensor query, // (num_heads, head_dim) - new token
torch::Tensor key_cache, // (num_heads, seq_len, head_dim) - contiguous
torch::Tensor value_cache,// (num_heads, seq_len, head_dim) - contiguous
float scale,
int num_threads = 0
) {
const int64_t num_heads = query.size(0);
const int64_t head_dim = query.size(1);
const int64_t seq_len = key_cache.size(1);
auto q_contig = query.contiguous();
auto k_contig = key_cache.contiguous();
auto v_contig = value_cache.contiguous();
auto q_data = q_contig.data_ptr<float>();
auto k_data = k_contig.data_ptr<float>();
auto v_data = v_contig.data_ptr<float>();
// Output: (num_heads, head_dim)
auto output = torch::zeros({num_heads, head_dim}, query.options());
auto out_data = output.data_ptr<float>();
if (num_threads > 0) {
omp_set_num_threads(num_threads);
}
// Parallel over heads
#pragma omp parallel for schedule(static)
for (int64_t h = 0; h < num_heads; h++) {
const float* q_head = q_data + h * head_dim;
const float* k_head = k_data + h * seq_len * head_dim;
const float* v_head = v_data + h * seq_len * head_dim;
float* out_head = out_data + h * head_dim;
// Compute attention scores for this head
std::vector<float> scores(seq_len);
float max_score = -INFINITY;
// QK^T for all keys
for (int64_t s = 0; s < seq_len; s++) {
const float* k_vec = k_head + s * head_dim;
float score = 0.0f;
for (int64_t d = 0; d < head_dim; d++) {
score += q_head[d] * k_vec[d];
}
score *= scale;
scores[s] = score;
max_score = std::max(max_score, score);
}
// Softmax: exp and sum
float exp_sum = 0.0f;
for (int64_t s = 0; s < seq_len; s++) {
scores[s] = std::exp(scores[s] - max_score);
exp_sum += scores[s];
}
// Normalize
for (int64_t s = 0; s < seq_len; s++) {
scores[s] /= exp_sum;
}
// Weighted sum of values
for (int64_t d = 0; d < head_dim; d++) {
float sum = 0.0f;
for (int64_t s = 0; s < seq_len; s++) {
const float* v_vec = v_head + s * head_dim;
sum += scores[s] * v_vec[d];
}
out_head[d] = sum;
}
}
return output;
}
// Append new key/value to cache in contiguous format
void append_to_kv_cache(
torch::Tensor key_cache, // (num_heads, max_seq_len, head_dim)
torch::Tensor value_cache, // (num_heads, max_seq_len, head_dim)
torch::Tensor new_key, // (num_heads, head_dim)
torch::Tensor new_value, // (num_heads, head_dim)
int64_t current_pos
) {
const int64_t num_heads = new_key.size(0);
const int64_t head_dim = new_key.size(1);
const int64_t max_seq_len = key_cache.size(1);
TORCH_CHECK(current_pos < max_seq_len, "KV cache overflow");
auto k_data = key_cache.data_ptr<float>();
auto v_data = value_cache.data_ptr<float>();
auto new_k_data = new_key.data_ptr<float>();
auto new_v_data = new_value.data_ptr<float>();
// Copy new key/value into cache at position
for (int64_t h = 0; h < num_heads; h++) {
float* k_cache_pos = k_data + h * max_seq_len * head_dim + current_pos * head_dim;
float* v_cache_pos = v_data + h * max_seq_len * head_dim + current_pos * head_dim;
const float* new_k = new_k_data + h * head_dim;
const float* new_v = new_v_data + h * head_dim;
std::copy(new_k, new_k + head_dim, k_cache_pos);
std::copy(new_v, new_v + head_dim, v_cache_pos);
}
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("attention_decode_contiguous_kv", &attention_decode_contiguous_kv,
"Optimized attention for decode with contiguous KV cache");
m.def("append_to_kv_cache", &append_to_kv_cache,
"Append new key/value to contiguous cache");
}