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));
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) {
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) {
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) {
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;
55 for (
int row = 0; row < rows; ++row) {
56 const float *src = packed_qkv + (size_t) row * (
size_t) packed_dim;
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));
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));
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));
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)