24 "HARD KERNEL CONTRACT FAULT: llama.cpp attention gate "
25 "requires expf from libm.so.6\n");
30#if defined(__GNUC__) || defined(__clang__)
45 const int packed_dim = q_dim + gate_dim;
46 if (!packed_qg || !q || !gate || rows <= 0 || q_dim <= 0 || gate_dim <= 0) {
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);
62 q_dst + (
size_t) group * (
size_t) group_dim,
64 (
size_t) group_dim *
sizeof(
float));
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));
71 memcpy(q_dst, src, (
size_t) q_dim *
sizeof(
float));
72 memcpy(gate_dst, src + q_dim, (
size_t) gate_dim *
sizeof(
float));
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) {
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);
102 dq_src + (
size_t) group * (
size_t) group_dim,
103 (
size_t) group_dim *
sizeof(
float));
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));
110 memcpy(dst, dq_src, (
size_t) q_dim *
sizeof(
float));
111 memcpy(dst + q_dim, dg_src, (
size_t) gate_dim *
sizeof(
float));
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) {
139 if (!x || !gate || !out || rows <= 0 || num_heads <= 0 || state_dim <= 0) {
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;
151 gate_row + col, sigmoid, 1, width);
152 for (
int lane = 0; lane < width; ++lane) {
156 x_bf16 * sigmoid_bf16));
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
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;
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) {
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);
static uint16_t float_to_bf16(float f)
static float bf16_to_float(uint16_t v)
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)