← Back to C-Kernel-Engine Docs Doxygen Source Documentation
 
Loading...
Searching...
No Matches
hybrid_attention_kernels.c File Reference
#include "ckernel_engine.h"
#include "bf16_utils.h"
#include <dlfcn.h>
#include <math.h>
#include <pthread.h>
#include <stdio.h>
#include <stdlib.h>
#include <string.h>

Go to the source code of this file.

Typedefs

typedef float(* ck_hybrid_libm_f32_fn) (float)
 

Functions

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)
 
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_sigmoid_mul_pytorch_bf16_storage (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_bind_hybrid_llama_libm (void)
 
static float hybrid_sigmoid (float x)
 
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 split_q_gate_forward (const float *packed_qg, float *q, float *gate, int rows, int q_dim, int gate_dim, int group_dim)
 

Variables

static void * ck_hybrid_libm_handle = NULL
 
static pthread_once_t ck_hybrid_libm_once = PTHREAD_ONCE_INIT
 
static ck_hybrid_libm_f32_fn ck_hybrid_llama_expf = NULL
 

Typedef Documentation

◆ ck_hybrid_libm_f32_fn

typedef float(* ck_hybrid_libm_f32_fn) (float)

Definition at line 11 of file hybrid_attention_kernels.c.

Function Documentation

◆ attn_gate_sigmoid_mul_backward()

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 
)

Definition at line 187 of file hybrid_attention_kernels.c.

194 {
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 float hybrid_sigmoid(float x)

References hybrid_sigmoid().

◆ attn_gate_sigmoid_mul_forward()

void attn_gate_sigmoid_mul_forward ( const float *  x,
const float *  gate,
float *  out,
int  rows,
int  num_heads,
int  state_dim 
)

Definition at line 116 of file hybrid_attention_kernels.c.

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

References hybrid_sigmoid().

Referenced by ck_test_attn_gate_sigmoid_mul().

◆ attn_gate_sigmoid_mul_pytorch_bf16_storage()

void attn_gate_sigmoid_mul_pytorch_bf16_storage ( const float *  x,
const float *  gate,
float *  out,
int  rows,
int  num_heads,
int  state_dim 
)

Definition at line 133 of file hybrid_attention_kernels.c.

138 {
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}
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)

References bf16_to_float(), float_to_bf16(), and recurrent_sigmoid_forward_pytorch_bf16_input_fp32_output().

◆ attn_gate_softplus_mul_forward()

void attn_gate_softplus_mul_forward ( const float *  x,
const float *  gate,
float *  out,
int  rows,
int  num_heads,
int  state_dim 
)

Definition at line 162 of file hybrid_attention_kernels.c.

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

◆ ck_bind_hybrid_llama_libm()

static void ck_bind_hybrid_llama_libm ( void  )
static

Definition at line 16 of file hybrid_attention_kernels.c.

16 {
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}
static void * ck_hybrid_libm_handle
static ck_hybrid_libm_f32_fn ck_hybrid_llama_expf
float(* ck_hybrid_libm_f32_fn)(float)

References ck_hybrid_libm_handle, and ck_hybrid_llama_expf.

Referenced by hybrid_sigmoid().

◆ hybrid_sigmoid()

static float hybrid_sigmoid ( float  x)
static

Definition at line 33 of file hybrid_attention_kernels.c.

33 {
35 return 1.0f / (1.0f + ck_hybrid_llama_expf(-x));
36}
static void ck_bind_hybrid_llama_libm(void)
static pthread_once_t ck_hybrid_libm_once

References ck_bind_hybrid_llama_libm(), ck_hybrid_libm_once, and ck_hybrid_llama_expf.

Referenced by attn_gate_sigmoid_mul_backward(), and attn_gate_sigmoid_mul_forward().

◆ split_q_gate_backward()

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 
)

Definition at line 77 of file hybrid_attention_kernels.c.

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

◆ split_q_gate_forward()

void split_q_gate_forward ( const float *  packed_qg,
float *  q,
float *  gate,
int  rows,
int  q_dim,
int  gate_dim,
int  group_dim 
)

Definition at line 38 of file hybrid_attention_kernels.c.

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

Referenced by ck_test_split_q_gate().

Variable Documentation

◆ ck_hybrid_libm_handle

void* ck_hybrid_libm_handle = NULL
static

Definition at line 13 of file hybrid_attention_kernels.c.

Referenced by ck_bind_hybrid_llama_libm().

◆ ck_hybrid_libm_once

pthread_once_t ck_hybrid_libm_once = PTHREAD_ONCE_INIT
static

Definition at line 14 of file hybrid_attention_kernels.c.

Referenced by hybrid_sigmoid().

◆ ck_hybrid_llama_expf

ck_hybrid_libm_f32_fn ck_hybrid_llama_expf = NULL
static

Definition at line 12 of file hybrid_attention_kernels.c.

Referenced by ck_bind_hybrid_llama_libm(), and hybrid_sigmoid().