← Back to C-Kernel-Engine Docs Doxygen Source Documentation
 
Loading...
Searching...
No Matches
recurrent_split_kernels.c
Go to the documentation of this file.
1#include "ckernel_engine.h"
2
3#include <string.h>
4
5void recurrent_split_qkv_forward(const float *packed_qkv,
6 float *q,
7 float *k,
8 float *v,
9 int rows,
10 int q_dim,
11 int k_dim,
12 int v_dim) {
13 const int packed_dim = q_dim + k_dim + v_dim;
14 for (int row = 0; row < rows; ++row) {
15 const float *src = packed_qkv + (size_t) row * (size_t) packed_dim;
16 float *q_dst = q + (size_t) row * (size_t) q_dim;
17 float *k_dst = k + (size_t) row * (size_t) k_dim;
18 float *v_dst = v + (size_t) row * (size_t) v_dim;
19 memcpy(q_dst, src, (size_t) q_dim * sizeof(float));
20 memcpy(k_dst, src + q_dim, (size_t) k_dim * sizeof(float));
21 memcpy(v_dst, src + q_dim + k_dim, (size_t) v_dim * sizeof(float));
22 }
23}
24
25void split_qkv_packed_head_major_forward(const float *packed_qkv,
26 float *q,
27 float *k,
28 float *v,
29 int rows,
30 int q_dim,
31 int k_dim,
32 int v_dim,
33 int num_heads,
34 int num_kv_heads) {
35 if (!packed_qkv || !q || !k || !v || rows <= 0 || q_dim <= 0 || k_dim <= 0 || v_dim <= 0 ||
36 num_heads <= 0 || num_kv_heads <= 0) {
37 return;
38 }
39
40 const int q_head_dim = q_dim / num_heads;
41 const int k_head_dim = k_dim / num_kv_heads;
42 const int v_head_dim = v_dim / num_kv_heads;
43 if (q_head_dim <= 0 || k_head_dim <= 0 || v_head_dim <= 0) {
44 return;
45 }
46 if (q_head_dim * num_heads != q_dim || k_head_dim * num_kv_heads != k_dim || v_head_dim * num_kv_heads != v_dim) {
47 return;
48 }
49
50 const int packed_dim = q_dim + k_dim + v_dim;
51 const size_t q_head_stride = (size_t) rows * (size_t) q_head_dim;
52 const size_t k_head_stride = (size_t) rows * (size_t) k_head_dim;
53 const size_t v_head_stride = (size_t) rows * (size_t) v_head_dim;
54
55 for (int row = 0; row < rows; ++row) {
56 const float *src = packed_qkv + (size_t) row * (size_t) packed_dim;
57
58 for (int head = 0; head < num_heads; ++head) {
59 const float *src_q = src + (size_t) head * (size_t) q_head_dim;
60 float *dst_q = q + (size_t) head * q_head_stride + (size_t) row * (size_t) q_head_dim;
61 memcpy(dst_q, src_q, (size_t) q_head_dim * sizeof(float));
62 }
63
64 const float *src_k_base = src + (size_t) q_dim;
65 const float *src_v_base = src + (size_t) q_dim + (size_t) k_dim;
66 for (int head = 0; head < num_kv_heads; ++head) {
67 const float *src_k = src_k_base + (size_t) head * (size_t) k_head_dim;
68 const float *src_v = src_v_base + (size_t) head * (size_t) v_head_dim;
69 float *dst_k = k + (size_t) head * k_head_stride + (size_t) row * (size_t) k_head_dim;
70 float *dst_v = v + (size_t) head * v_head_stride + (size_t) row * (size_t) v_head_dim;
71 memcpy(dst_k, src_k, (size_t) k_head_dim * sizeof(float));
72 memcpy(dst_v, src_v, (size_t) v_head_dim * sizeof(float));
73 }
74 }
75}
76
77void recurrent_split_qkv_backward(const float *d_q,
78 const float *d_k,
79 const float *d_v,
80 float *d_packed_qkv,
81 int rows,
82 int q_dim,
83 int k_dim,
84 int v_dim) {
85 const int packed_dim = q_dim + k_dim + v_dim;
86 for (int row = 0; row < rows; ++row) {
87 const float *dq_src = d_q + (size_t) row * (size_t) q_dim;
88 const float *dk_src = d_k + (size_t) row * (size_t) k_dim;
89 const float *dv_src = d_v + (size_t) row * (size_t) v_dim;
90 float *dst = d_packed_qkv + (size_t) row * (size_t) packed_dim;
91 memcpy(dst, dq_src, (size_t) q_dim * sizeof(float));
92 memcpy(dst + q_dim, dk_src, (size_t) k_dim * sizeof(float));
93 memcpy(dst + q_dim + k_dim, dv_src, (size_t) v_dim * sizeof(float));
94 }
95}
96
97void recurrent_split_conv_qkv_forward(const float *packed_qkv,
98 float *q,
99 float *k,
100 float *v,
101 int rows,
102 int q_dim,
103 int k_dim,
104 int v_dim) {
105 recurrent_split_qkv_forward(packed_qkv, q, k, v, rows, q_dim, k_dim, v_dim);
106}
107
109 const float *d_k,
110 const float *d_v,
111 float *d_packed_qkv,
112 int rows,
113 int q_dim,
114 int k_dim,
115 int v_dim) {
116 recurrent_split_qkv_backward(d_q, d_k, d_v, d_packed_qkv, rows, q_dim, k_dim, v_dim);
117}
void recurrent_split_conv_qkv_backward(const float *d_q, const float *d_k, const float *d_v, float *d_packed_qkv, int rows, int q_dim, int k_dim, int v_dim)
void recurrent_split_qkv_backward(const float *d_q, const float *d_k, const float *d_v, float *d_packed_qkv, int rows, int q_dim, int k_dim, int v_dim)
void recurrent_split_conv_qkv_forward(const float *packed_qkv, float *q, float *k, float *v, int rows, int q_dim, int k_dim, int v_dim)
void split_qkv_packed_head_major_forward(const float *packed_qkv, float *q, float *k, float *v, int rows, int q_dim, int k_dim, int v_dim, int num_heads, int num_kv_heads)
void recurrent_split_qkv_forward(const float *packed_qkv, float *q, float *k, float *v, int rows, int q_dim, int k_dim, int v_dim)