Repository navigation
Expand file tree
/
Copy pathmemory.cpp
More file actions
74 lines (63 loc) · 2.65 KB
/
Copy pathmemory.cpp
File metadata and controls
74 lines (63 loc) · 2.65 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
#include "memory.h"
#include "transformer.h"
#include "cuda_check.h"
#include <cuda_runtime.h>
#include <algorithm>
#include <cstddef>
namespace {
float* alloc_floats(size_t count) {
float* ptr = nullptr;
CUDA_CHECK(cudaMalloc((void**)&ptr, count * sizeof(float)));
CUDA_CHECK(cudaMemset(ptr, 0, count * sizeof(float)));
return ptr;
}
}
void init_kv_cache(LayerKVCache& cache, int max_seq_len, const ModelDims& dims) {
cache.max_seq_len = max_seq_len;
cache.current_seq_len = 0;
const size_t elems = static_cast<size_t>(dims.kv_dim()) * max_seq_len;
cache.k_cache = alloc_floats(elems);
cache.v_cache = alloc_floats(elems);
}
void init_buffers(LayerBuffers& buffers, const ModelDims& dims, int /*max_seq_len*/) {
const size_t max_tokens = MAX_TOKENS_PER_FORWARD;
const size_t hidden = dims.hidden, inter = dims.intermediate, q_dim = dims.q_dim(), kv_dim = dims.kv_dim();
buffers.norm_result = alloc_floats(max_tokens * hidden);
buffers.q_result = alloc_floats(max_tokens * q_dim);
buffers.k_result = alloc_floats(max_tokens * kv_dim);
buffers.v_result = alloc_floats(max_tokens * kv_dim);
buffers.attn_result = alloc_floats(max_tokens * q_dim);
buffers.o_result = alloc_floats(max_tokens * hidden);
buffers.mlp_norm_result = alloc_floats(max_tokens * hidden);
buffers.down_result = alloc_floats(max_tokens * hidden);
buffers.gate_result = alloc_floats(max_tokens * inter);
buffers.up_result = alloc_floats(max_tokens * inter);
buffers.swiglu_result = alloc_floats(max_tokens * inter);
const size_t attn_rows = max_tokens * dims.num_heads * MAX_ATTN_CHUNKS;
const size_t widest_out = std::max({hidden, inter, q_dim, kv_dim});
const size_t splitk_elems = static_cast<size_t>(NUM_SPLITS) * max_tokens * widest_out;
buffers.partial_O = alloc_floats(std::max(attn_rows * dims.head_dim, splitk_elems));
buffers.partial_lse = alloc_floats(attn_rows);
}
void free_kv_cache(LayerKVCache& cache) {
cudaFree(cache.k_cache);
cudaFree(cache.v_cache);
cache.k_cache = nullptr;
cache.v_cache = nullptr;
}
void free_buffers(LayerBuffers& buffers) {
cudaFree(buffers.norm_result);
cudaFree(buffers.q_result);
cudaFree(buffers.k_result);
cudaFree(buffers.v_result);
cudaFree(buffers.attn_result);
cudaFree(buffers.o_result);
cudaFree(buffers.mlp_norm_result);
cudaFree(buffers.down_result);
cudaFree(buffers.gate_result);
cudaFree(buffers.up_result);
cudaFree(buffers.swiglu_result);
cudaFree(buffers.partial_O);
cudaFree(buffers.partial_lse);
buffers = LayerBuffers{};
}