← Back to C-Kernel-Engine Docs Doxygen Source Documentation
 
Loading...
Searching...
No Matches
recurrent_split_kernels.c File Reference
#include "ckernel_engine.h"
#include <string.h>

Go to the source code of this file.

Functions

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_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 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_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)
 

Function Documentation

◆ recurrent_split_conv_qkv_backward()

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 
)

Definition at line 108 of file recurrent_split_kernels.c.

115 {
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_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)

References recurrent_split_qkv_backward().

◆ recurrent_split_conv_qkv_forward()

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 
)

Definition at line 97 of file recurrent_split_kernels.c.

104 {
105 recurrent_split_qkv_forward(packed_qkv, q, k, v, rows, q_dim, k_dim, v_dim);
106}
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)

References recurrent_split_qkv_forward().

Referenced by ck_test_recurrent_split_conv_qkv().

◆ recurrent_split_qkv_backward()

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 
)

Definition at line 77 of file recurrent_split_kernels.c.

84 {
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}

Referenced by recurrent_split_conv_qkv_backward().

◆ recurrent_split_qkv_forward()

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 
)

Definition at line 5 of file recurrent_split_kernels.c.

12 {
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}

Referenced by ck_test_recurrent_split_qkv(), and recurrent_split_conv_qkv_forward().

◆ split_qkv_packed_head_major_forward()

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 
)

Definition at line 25 of file recurrent_split_kernels.c.

34 {
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}