← Back to C-Kernel-Engine Docs Doxygen Source Documentation
 
Loading...
Searching...
No Matches
hybrid_attention_kernels.c
Go to the documentation of this file.
1#include "ckernel_engine.h"
2#include "bf16_utils.h"
3
4#include <dlfcn.h>
5#include <math.h>
6#include <pthread.h>
7#include <stdio.h>
8#include <stdlib.h>
9#include <string.h>
10
11typedef float (*ck_hybrid_libm_f32_fn)(float);
13static void *ck_hybrid_libm_handle = NULL;
14static pthread_once_t ck_hybrid_libm_once = PTHREAD_ONCE_INIT;
15
16static void ck_bind_hybrid_llama_libm(void) {
17 ck_hybrid_libm_handle = dlopen("libm.so.6", RTLD_NOW | RTLD_LOCAL);
21 }
23 fprintf(stderr,
24 "HARD KERNEL CONTRACT FAULT: llama.cpp attention gate "
25 "requires expf from libm.so.6\n");
26 abort();
27 }
28}
29
30#if defined(__GNUC__) || defined(__clang__)
31__attribute__((noinline))
32#endif
33static float hybrid_sigmoid(float x) {
35 return 1.0f / (1.0f + ck_hybrid_llama_expf(-x));
36}
37
38void split_q_gate_forward(const float *packed_qg,
39 float *q,
40 float *gate,
41 int rows,
42 int q_dim,
43 int gate_dim,
44 int group_dim) {
45 const int packed_dim = q_dim + gate_dim;
46 if (!packed_qg || !q || !gate || rows <= 0 || q_dim <= 0 || gate_dim <= 0) {
47 return;
48 }
49 if (group_dim <= 0) {
50 group_dim = q_dim;
51 }
52 const int q_groups = q_dim / group_dim;
53 const int gate_group_dim = (q_groups > 0 && gate_dim % q_groups == 0) ? (gate_dim / q_groups) : gate_dim;
54 for (int row = 0; row < rows; ++row) {
55 const float *src = packed_qg + (size_t) row * (size_t) packed_dim;
56 float *q_dst = q + (size_t) row * (size_t) q_dim;
57 float *gate_dst = gate + (size_t) row * (size_t) gate_dim;
58 if (q_groups > 0 && q_groups * group_dim == q_dim && q_groups * gate_group_dim == gate_dim) {
59 for (int group = 0; group < q_groups; ++group) {
60 const size_t src_group_off = (size_t) group * (size_t) (group_dim + gate_group_dim);
61 memcpy(
62 q_dst + (size_t) group * (size_t) group_dim,
63 src + src_group_off,
64 (size_t) group_dim * sizeof(float));
65 memcpy(
66 gate_dst + (size_t) group * (size_t) gate_group_dim,
67 src + src_group_off + (size_t) group_dim,
68 (size_t) gate_group_dim * sizeof(float));
69 }
70 } else {
71 memcpy(q_dst, src, (size_t) q_dim * sizeof(float));
72 memcpy(gate_dst, src + q_dim, (size_t) gate_dim * sizeof(float));
73 }
74 }
75}
76
77void split_q_gate_backward(const float *d_q,
78 const float *d_gate,
79 float *d_packed_qg,
80 int rows,
81 int q_dim,
82 int gate_dim,
83 int group_dim) {
84 const int packed_dim = q_dim + gate_dim;
85 if (!d_q || !d_gate || !d_packed_qg || rows <= 0 || q_dim <= 0 || gate_dim <= 0) {
86 return;
87 }
88 if (group_dim <= 0) {
89 group_dim = q_dim;
90 }
91 const int q_groups = q_dim / group_dim;
92 const int gate_group_dim = (q_groups > 0 && gate_dim % q_groups == 0) ? (gate_dim / q_groups) : gate_dim;
93 for (int row = 0; row < rows; ++row) {
94 const float *dq_src = d_q + (size_t) row * (size_t) q_dim;
95 const float *dg_src = d_gate + (size_t) row * (size_t) gate_dim;
96 float *dst = d_packed_qg + (size_t) row * (size_t) packed_dim;
97 if (q_groups > 0 && q_groups * group_dim == q_dim && q_groups * gate_group_dim == gate_dim) {
98 for (int group = 0; group < q_groups; ++group) {
99 const size_t dst_group_off = (size_t) group * (size_t) (group_dim + gate_group_dim);
100 memcpy(
101 dst + dst_group_off,
102 dq_src + (size_t) group * (size_t) group_dim,
103 (size_t) group_dim * sizeof(float));
104 memcpy(
105 dst + dst_group_off + (size_t) group_dim,
106 dg_src + (size_t) group * (size_t) gate_group_dim,
107 (size_t) gate_group_dim * sizeof(float));
108 }
109 } else {
110 memcpy(dst, dq_src, (size_t) q_dim * sizeof(float));
111 memcpy(dst + q_dim, dg_src, (size_t) gate_dim * sizeof(float));
112 }
113 }
114}
115
117 const float *gate,
118 float *out,
119 int rows,
120 int num_heads,
121 int state_dim) {
122 const int dim = num_heads * state_dim;
123 for (int row = 0; row < rows; ++row) {
124 const float *x_row = x + (size_t) row * (size_t) dim;
125 const float *gate_row = gate + (size_t) row * (size_t) dim;
126 float *out_row = out + (size_t) row * (size_t) dim;
127 for (int col = 0; col < dim; ++col) {
128 out_row[col] = x_row[col] * hybrid_sigmoid(gate_row[col]);
129 }
130 }
131}
132
134 const float *gate,
135 float *out,
136 int rows,
137 int num_heads,
138 int state_dim) {
139 if (!x || !gate || !out || rows <= 0 || num_heads <= 0 || state_dim <= 0) {
140 return;
141 }
142 const int dim = num_heads * state_dim;
143 for (int row = 0; row < rows; ++row) {
144 const float *x_row = x + (size_t)row * (size_t)dim;
145 const float *gate_row = gate + (size_t)row * (size_t)dim;
146 float *out_row = out + (size_t)row * (size_t)dim;
147 for (int col = 0; col < dim; col += 16) {
148 const int width = dim - col < 16 ? dim - col : 16;
149 float sigmoid[16];
151 gate_row + col, sigmoid, 1, width);
152 for (int lane = 0; lane < width; ++lane) {
153 const float x_bf16 = bf16_to_float(float_to_bf16(x_row[col + lane]));
154 const float sigmoid_bf16 = bf16_to_float(float_to_bf16(sigmoid[lane]));
155 out_row[col + lane] = bf16_to_float(float_to_bf16(
156 x_bf16 * sigmoid_bf16));
157 }
158 }
159 }
160}
161
163 const float *gate,
164 float *out,
165 int rows,
166 int num_heads,
167 int state_dim) {
168 const int dim = num_heads * state_dim;
169 for (int row = 0; row < rows; ++row) {
170 const float *x_row = x + (size_t) row * (size_t) dim;
171 const float *gate_row = gate + (size_t) row * (size_t) num_heads;
172 float *out_row = out + (size_t) row * (size_t) dim;
173 for (int head = 0; head < num_heads; ++head) {
174 const float value = gate_row[head];
175 const float scale = value > 20.0f
176 ? value
177 : log1pf(expf(value));
178 const size_t base = (size_t) head * (size_t) state_dim;
179 for (int col = 0; col < state_dim; ++col) {
180 out_row[base + (size_t) col] =
181 x_row[base + (size_t) col] * scale;
182 }
183 }
184 }
185}
186
187void attn_gate_sigmoid_mul_backward(const float *d_out,
188 const float *x,
189 const float *gate,
190 float *d_x,
191 float *d_gate,
192 int rows,
193 int num_heads,
194 int state_dim) {
195 const int dim = num_heads * state_dim;
196 for (int row = 0; row < rows; ++row) {
197 const float *d_out_row = d_out + (size_t) row * (size_t) dim;
198 const float *x_row = x + (size_t) row * (size_t) dim;
199 const float *gate_row = gate + (size_t) row * (size_t) dim;
200 float *d_x_row = d_x + (size_t) row * (size_t) dim;
201 float *d_gate_row = d_gate + (size_t) row * (size_t) dim;
202 for (int col = 0; col < dim; ++col) {
203 const float sig = hybrid_sigmoid(gate_row[col]);
204 d_x_row[col] = d_out_row[col] * sig;
205 d_gate_row[col] = d_out_row[col] * x_row[col] * sig * (1.0f - sig);
206 }
207 }
208}
static uint16_t float_to_bf16(float f)
Definition bf16_utils.h:90
static float bf16_to_float(uint16_t v)
Definition bf16_utils.h:38
void recurrent_sigmoid_forward_pytorch_bf16_input_fp32_output(const float *x, float *out, int rows, int dim)
static void ck_bind_hybrid_llama_libm(void)
void split_q_gate_backward(const float *d_q, const float *d_gate, float *d_packed_qg, int rows, int q_dim, int gate_dim, int group_dim)
void attn_gate_sigmoid_mul_pytorch_bf16_storage(const float *x, const float *gate, float *out, int rows, int num_heads, int state_dim)
void attn_gate_sigmoid_mul_forward(const float *x, const float *gate, float *out, int rows, int num_heads, int state_dim)
void attn_gate_softplus_mul_forward(const float *x, const float *gate, float *out, int rows, int num_heads, int state_dim)
static void * ck_hybrid_libm_handle
static pthread_once_t ck_hybrid_libm_once
static ck_hybrid_libm_f32_fn ck_hybrid_llama_expf
float(* ck_hybrid_libm_f32_fn)(float)
void attn_gate_sigmoid_mul_backward(const float *d_out, const float *x, const float *gate, float *d_x, float *d_gate, int rows, int num_heads, int state_dim)
static float hybrid_sigmoid(float x)
void split_q_gate_forward(const float *packed_qg, float *q, float *gate, int rows, int q_dim, int gate_dim, int group_dim)
__attribute__((visibility("default"))) CKTokenizer *ck_tokenizer_create(CKTokenizerType type)