16} ck_gemma4_block_q5_K;
20 uint32_t bits = ((uint32_t)v) << 16;
22 memcpy(&out, &bits,
sizeof(out));
28 const float c0 = 0.7978845608028654f;
29 const float c1 = 0.044715f;
30 return 0.5f * x * (1.0f + tanhf(c0 * x * (1.0f + c1 * x * x)));
35 sc[0] = scales[0] & 0x3F;
36 sc[1] = scales[1] & 0x3F;
37 sc[2] = scales[2] & 0x3F;
38 sc[3] = scales[3] & 0x3F;
40 m[0] = scales[4] & 0x3F;
41 m[1] = scales[5] & 0x3F;
42 m[2] = scales[6] & 0x3F;
43 m[3] = scales[7] & 0x3F;
45 sc[4] = (scales[8] & 0x0F) | ((scales[0] >> 6) << 4);
46 sc[5] = (scales[9] & 0x0F) | ((scales[1] >> 6) << 4);
47 sc[6] = (scales[10] & 0x0F) | ((scales[2] >> 6) << 4);
48 sc[7] = (scales[11] & 0x0F) | ((scales[3] >> 6) << 4);
50 m[4] = (scales[8] >> 4) | ((scales[4] >> 6) << 4);
51 m[5] = (scales[9] >> 4) | ((scales[5] >> 6) << 4);
52 m[6] = (scales[10] >> 4) | ((scales[6] >> 6) << 4);
53 m[7] = (scales[11] >> 4) | ((scales[7] >> 6) << 4);
58 const uint8_t *ql = block->qs + (subblock / 2) * 32;
59 const uint8_t low = (subblock & 1) ? (uint8_t)(ql[i] >> 4) : (uint8_t)(ql[i] & 0x0F);
60 const uint8_t high = (block->qh[i] & (uint8_t)(1u << subblock)) ? 16u : 0u;
61 return (uint8_t)(low | high);
71 for (
int s = 0; s < 8; ++s) {
72 const float scale = d * (float)sc[s];
73 const float minv = dmin * (float)m[s];
74 for (
int i = 0; i < 32; ++i) {
83 for (
int i = 0; i < n; ++i) {
84 ss += (double)x[i] * (
double)x[i];
86 const float scale = 1.0f / sqrtf((
float)(ss / (
double)n) + eps);
87 for (
int i = 0; i < n; ++i) {
88 out[i] = x[i] * scale * gamma[i];
93 float *per_layer_input;
95 const int32_t *token_ids;
96 const void *per_layer_token_emb;
97 const uint16_t *per_layer_model_proj;
98 const float *per_layer_proj_norm;
104} ck_gemma4_prepare_args_t;
108 const ck_gemma4_prepare_args_t *args =
109 (
const ck_gemma4_prepare_args_t *)opaque;
110 const ck_gemma4_block_q5_K *token_blocks =
111 (
const ck_gemma4_block_q5_K *)args->per_layer_token_emb;
112 const size_t token_blocks_per_row = (size_t)args->num_layers;
113 const float token_scale = sqrtf((
float)args->per_layer_dim);
114 const float model_scale = 1.0f / sqrtf((
float)args->embed_dim);
115 const float mix_scale = 0.7071067811865475f;
116 float token_vec[
QK_K];
117 float proj_vec[
QK_K];
118 float proj_normed[
QK_K];
120 for (
int t = begin; t <
end; ++t) {
121 const int token = args->token_ids[t];
122 if (token < 0 || token >= args->vocab_size)
continue;
123 const float *h = args->hidden + (size_t)t * (
size_t)args->embed_dim;
124 for (
int layer = 0; layer < args->num_layers; ++layer) {
125 float *dst = args->per_layer_input +
126 ((size_t)t * (
size_t)args->num_layers + (size_t)layer) *
127 (size_t)args->per_layer_dim;
128 const ck_gemma4_block_q5_K *tok_block = token_blocks +
129 (size_t)
token * token_blocks_per_row + (
size_t)layer;
131 for (
int i = 0; i < args->per_layer_dim; ++i) {
132 token_vec[i] *= token_scale;
135 const uint16_t *model_proj_base = args->per_layer_model_proj +
136 (size_t)layer * (
size_t)args->per_layer_dim *
137 (size_t)args->embed_dim;
138 for (
int i = 0; i < args->per_layer_dim; ++i) {
139 const uint16_t *row = model_proj_base +
140 (size_t)i * (
size_t)args->embed_dim;
142 for (
int j = 0; j < args->embed_dim; ++j) {
145 proj_vec[i] = acc * model_scale;
148 proj_vec, args->per_layer_proj_norm, proj_normed,
149 args->per_layer_dim, args->eps);
150 for (
int i = 0; i < args->per_layer_dim; ++i) {
151 dst[i] = (token_vec[i] + proj_normed[i]) * mix_scale;
159 const ck_gemma4_prepare_args_t *args =
160 (
const ck_gemma4_prepare_args_t *)opaque;
161 const uint16_t *token_embeddings =
162 (
const uint16_t *)args->per_layer_token_emb;
163 const float token_scale = sqrtf((
float)args->per_layer_dim);
164 const float model_scale = 1.0f / sqrtf((
float)args->embed_dim);
165 const float mix_scale = 0.7071067811865475f;
166 float token_vec[
QK_K];
167 float proj_vec[
QK_K];
168 float proj_normed[
QK_K];
170 for (
int t = begin; t <
end; ++t) {
171 const int token = args->token_ids[t];
172 if (token < 0 || token >= args->vocab_size)
continue;
173 const float *h = args->hidden + (size_t)t * (
size_t)args->embed_dim;
174 for (
int layer = 0; layer < args->num_layers; ++layer) {
175 float *dst = args->per_layer_input +
176 ((size_t)t * (
size_t)args->num_layers + (size_t)layer) *
177 (size_t)args->per_layer_dim;
178 const uint16_t *tok_row = token_embeddings +
179 ((size_t)
token * (
size_t)args->num_layers + (size_t)layer) *
180 (size_t)args->per_layer_dim;
181 for (
int i = 0; i < args->per_layer_dim; ++i) {
185 const uint16_t *model_proj_base = args->per_layer_model_proj +
186 (size_t)layer * (
size_t)args->per_layer_dim *
187 (size_t)args->embed_dim;
188 for (
int i = 0; i < args->per_layer_dim; ++i) {
189 const uint16_t *row = model_proj_base +
190 (size_t)i * (
size_t)args->embed_dim;
192 for (
int j = 0; j < args->embed_dim; ++j) {
195 proj_vec[i] = acc * model_scale;
198 proj_vec, args->per_layer_proj_norm, proj_normed,
199 args->per_layer_dim, args->eps);
200 for (
int i = 0; i < args->per_layer_dim; ++i) {
201 dst[i] = (token_vec[i] + proj_normed[i]) * mix_scale;
208 int tokens,
ck_range_fn_t fn, ck_gemma4_prepare_args_t *args)
212 const char *disabled = getenv(
"CK_DISABLE_GEMMA4_PREPARE_PARALLEL");
213 if (disabled && disabled[0] && strcmp(disabled,
"0") != 0) active = 1;
214 if (active > tokens) active = tokens;
220 const int32_t *token_ids,
221 const void *per_layer_token_emb,
222 const uint16_t *per_layer_model_proj,
223 const float *per_layer_proj_norm,
231 if (!per_layer_input || !hidden || !token_ids || !per_layer_token_emb ||
232 !per_layer_model_proj || !per_layer_proj_norm || tokens <= 0 ||
233 num_layers <= 0 || embed_dim <= 0 || per_layer_dim !=
QK_K ||
vocab_size <= 0) {
237 ck_gemma4_prepare_args_t args = {
238 .per_layer_input = per_layer_input,
240 .token_ids = token_ids,
241 .per_layer_token_emb = per_layer_token_emb,
242 .per_layer_model_proj = per_layer_model_proj,
243 .per_layer_proj_norm = per_layer_proj_norm,
244 .num_layers = num_layers,
245 .embed_dim = embed_dim,
246 .per_layer_dim = per_layer_dim,
256 const int32_t *token_ids,
257 const uint16_t *per_layer_token_emb,
258 const uint16_t *per_layer_model_proj,
259 const float *per_layer_proj_norm,
267 if (!per_layer_input || !hidden || !token_ids || !per_layer_token_emb ||
268 !per_layer_model_proj || !per_layer_proj_norm || tokens <= 0 ||
269 num_layers <= 0 || embed_dim <= 0 || per_layer_dim <= 0 ||
vocab_size <= 0) {
273 if (per_layer_dim >
QK_K) {
276 ck_gemma4_prepare_args_t args = {
277 .per_layer_input = per_layer_input,
279 .token_ids = token_ids,
280 .per_layer_token_emb = per_layer_token_emb,
281 .per_layer_model_proj = per_layer_model_proj,
282 .per_layer_proj_norm = per_layer_proj_norm,
283 .num_layers = num_layers,
284 .embed_dim = embed_dim,
285 .per_layer_dim = per_layer_dim,
294 const float *per_layer_input;
295 const float *inp_gate;
297 const float *post_norm;
298 const float *out_scale;
304} ck_gemma4_embed_args_t;
308 const ck_gemma4_embed_args_t *args =
309 (
const ck_gemma4_embed_args_t *)opaque;
310 float gate_vec[
QK_K];
312 float branch_normed[4096];
313 for (
int t = begin; t <
end; ++t) {
314 float *h = args->hidden + (size_t)t * (
size_t)args->embed_dim;
315 const float *inp_vec = args->per_layer_input +
316 ((size_t)t * (
size_t)args->num_layers + (size_t)args->layer) *
317 (size_t)args->per_layer_dim;
319 for (
int i = 0; i < args->per_layer_dim; ++i) {
320 const float *row = args->inp_gate +
321 (size_t)i * (
size_t)args->embed_dim;
323 for (
int j = 0; j < args->embed_dim; ++j) {
324 acc += row[j] * h[j];
329 for (
int j = 0; j < args->embed_dim; ++j) {
330 const float *row = args->proj +
331 (size_t)j * (
size_t)args->per_layer_dim;
333 for (
int i = 0; i < args->per_layer_dim; ++i) {
334 acc += row[i] * gate_vec[i];
339 branch, args->post_norm, branch_normed,
340 args->embed_dim, args->eps);
341 const float layer_scale = args->out_scale ? args->out_scale[0] : 1.0f;
342 for (
int j = 0; j < args->embed_dim; ++j) {
343 h[j] = (h[j] + branch_normed[j]) * layer_scale;
349 const float *per_layer_input,
350 const float *inp_gate,
352 const float *post_norm,
353 const float *out_scale,
361 if (!hidden || !per_layer_input || !inp_gate || !proj || !post_norm ||
362 tokens <= 0 || layer < 0 || layer >= num_layers || embed_dim <= 0 ||
363 per_layer_dim !=
QK_K || embed_dim > 4096) {
367 ck_gemma4_embed_args_t args = {
369 .per_layer_input = per_layer_input,
370 .inp_gate = inp_gate,
372 .post_norm = post_norm,
373 .out_scale = out_scale,
375 .num_layers = num_layers,
376 .embed_dim = embed_dim,
377 .per_layer_dim = per_layer_dim,
382 const char *disabled = getenv(
"CK_DISABLE_GEMMA4_EMBED_PARALLEL");
383 if (disabled && disabled[0] && strcmp(disabled,
"0") != 0) active = 1;
384 if (active > tokens) active = tokens;
394 if (!hidden || !scale || tokens <= 0 || embed_dim <= 0) {
398 const float s = scale[0];
399 const size_t n = (size_t)tokens * (
size_t)embed_dim;
400 for (
size_t i = 0; i < n; ++i) {
410 if (!logits || tokens <= 0 ||
vocab_size <= 0 || cap <= 0.0f) {
413 const float inv_cap = 1.0f / cap;
414 const size_t total = (size_t)tokens * (
size_t)
vocab_size;
415 for (
size_t i = 0; i < total; ++i) {
416 logits[i] = tanhf(logits[i] * inv_cap) * cap;
Persistent pthread thread pool for CK-Engine inference.
void ck_threadpool_parallel_for_n(ck_threadpool_t *pool, int active_threads, int begin, int end, int grain_size, ck_range_fn_t fn, void *args)
ck_threadpool_t * ck_threadpool_global(void)
void(* ck_range_fn_t)(int begin, int end, void *args)
int ck_threadpool_n_threads(const ck_threadpool_t *pool)
Quantization block structures for weight-only quantization.
#define CK_FP16_TO_FP32(x)
void gemma4_per_layer_prepare_forward(float *per_layer_input, const float *hidden, const int32_t *token_ids, const void *per_layer_token_emb, const uint16_t *per_layer_model_proj, const float *per_layer_proj_norm, int tokens, int num_layers, int embed_dim, int per_layer_dim, int vocab_size, float eps)
void gemma4_per_layer_embed_forward(float *hidden, const float *per_layer_input, const float *inp_gate, const float *proj, const float *post_norm, const float *out_scale, int tokens, int layer, int num_layers, int embed_dim, int per_layer_dim, float eps)
static void ck_gemma4_dequant_q5_k_block(const ck_gemma4_block_q5_K *block, float *out)
static void ck_gemma4_prepare_parallel(int tokens, ck_range_fn_t fn, ck_gemma4_prepare_args_t *args)
static void ck_gemma4_embed_range(int begin, int end, void *opaque)
static void ck_gemma4_unpack_q5_k_scales(const uint8_t *scales, uint8_t *sc, uint8_t *m)
static float ck_bf16_to_f32(uint16_t v)
static float ck_gemma4_gelu(float x)
void assistant_layer_scale_forward(float *hidden, const float *scale, int tokens, int embed_dim)
void gemma4_per_layer_prepare_bf16_forward(float *per_layer_input, const float *hidden, const int32_t *token_ids, const uint16_t *per_layer_token_emb, const uint16_t *per_layer_model_proj, const float *per_layer_proj_norm, int tokens, int num_layers, int embed_dim, int per_layer_dim, int vocab_size, float eps)
void gemma4_final_logit_softcap_forward(float *logits, int tokens, int vocab_size, float cap)
static void ck_gemma4_prepare_bf16_range(int begin, int end, void *opaque)
static void ck_gemma4_prepare_q5_range(int begin, int end, void *opaque)
static void ck_gemma4_rmsnorm_tmp(const float *x, const float *gamma, float *out, int n, float eps)
static uint8_t ck_gemma4_q5_k_value(const ck_gemma4_block_q5_K *block, int subblock, int i)