-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtransformer.cpp
More file actions
83 lines (65 loc) · 4.01 KB
/
Copy pathtransformer.cpp
File metadata and controls
83 lines (65 loc) · 4.01 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
#include "transformer.h"
#include "../gemv/gemv.h"
#include "../gemm/gemm.h"
#include "../ops/flash_decoding.h"
#include "../ops/residual_operations.h"
#include "../ops/kv_cache_ops.h"
#include <cmath>
#include <cuda_runtime.h>
void run_quantized_linear(const QuantizedLinear& proj, float* input, float* output, float* buffer, int num_tokens) {
if(num_tokens == 1) {
run_gemv_int4_optimized_kernel(proj.q_weight, proj.scales, input, output, proj.out_features, proj.in_features);
} else {
run_gemm_int4_splitk_partial_kernel(input, proj.q_weight, proj.scales, buffer, num_tokens, NUM_SPLITS, proj.in_features, proj.out_features);
run_gemm_int4_splitk_final_kernel(buffer, output, num_tokens, proj.out_features, NUM_SPLITS);
}
if(proj.bias != nullptr) {
run_add_bias_kernel(output, proj.bias, num_tokens, proj.out_features);
}
}
void forward_transformer_block(float* hidden_states, const TransformerBlockWeights& weights, LayerKVCache& kv_cache, LayerBuffers& buffers, const ModelDims& dims, int num_tokens) {
const int D = dims.hidden;
const int q_dim = dims.q_dim();
const int kv_dim = dims.kv_dim();
const int head_dim = dims.head_dim;
const int group = dims.group_size();
// ATTENTION
for(int t = 0; t < num_tokens; t++) run_RMSNorm_kernel(hidden_states + (t * D), weights.attn_norm_weight, buffers.norm_result + (t * D), D);
run_quantized_linear(weights.q_proj, buffers.norm_result, buffers.q_result, buffers.partial_O, num_tokens); // [T, q_dim]
run_quantized_linear(weights.k_proj, buffers.norm_result, buffers.k_result, buffers.partial_O, num_tokens); // [T, kv_dim]
run_quantized_linear(weights.v_proj, buffers.norm_result, buffers.v_result, buffers.partial_O, num_tokens); // [T, kv_dim]
for(int t = 0; t < num_tokens; t++) run_RoPE_kernel(buffers.q_result + (t * q_dim), kv_cache.current_seq_len + t, q_dim, head_dim);
for(int t = 0; t < num_tokens; t++) run_RoPE_kernel(buffers.k_result + (t * kv_dim), kv_cache.current_seq_len + t, kv_dim, head_dim);
run_append_kv_cache(buffers.k_result, buffers.v_result, kv_cache.k_cache, kv_cache.v_cache, kv_cache.current_seq_len, kv_cache.max_seq_len, head_dim, kv_dim, num_tokens);
for (int t = 0; t < num_tokens; t++) {
int S = kv_cache.current_seq_len + t + 1;
int num_chunks = (S + FLASH_CHUNK_SIZE - 1) / FLASH_CHUNK_SIZE;
for (int h = 0; h < dims.num_heads; h++) {
int kv_head = h / group;
int q_offset = t * q_dim + h * head_dim;
int kv_offset = kv_head * kv_cache.max_seq_len * head_dim;
run_flash_decoding_partial(
buffers.q_result + q_offset,
kv_cache.k_cache + kv_offset,
kv_cache.v_cache + kv_offset,
buffers.partial_O, buffers.partial_lse,
head_dim, S, FLASH_CHUNK_SIZE
);
run_flash_decoding_final(
buffers.partial_O, buffers.partial_lse,
buffers.attn_result + q_offset,
head_dim, num_chunks
);
}
}
run_quantized_linear(weights.o_proj, buffers.attn_result, buffers.o_result, buffers.partial_O, num_tokens);
run_add_residual_kernel(hidden_states, buffers.o_result, num_tokens * D);
// MLP
for(int t = 0; t < num_tokens; t++) run_RMSNorm_kernel(hidden_states + (t * D), weights.mlp_norm_weight, buffers.mlp_norm_result + (t * D), D);
run_quantized_linear(weights.gate_proj, buffers.mlp_norm_result, buffers.gate_result, buffers.partial_O, num_tokens);
run_quantized_linear(weights.up_proj, buffers.mlp_norm_result, buffers.up_result, buffers.partial_O, num_tokens);
run_swiglu_kernel(buffers.gate_result, buffers.up_result, buffers.swiglu_result, num_tokens * weights.gate_proj.out_features);
run_quantized_linear(weights.down_proj, buffers.swiglu_result, buffers.down_result, buffers.partial_O, num_tokens);
run_add_residual_kernel(hidden_states, buffers.down_result, num_tokens * D);
kv_cache.current_seq_len += num_tokens;
}