← Back to C-Kernel-Engine Docs Doxygen Source Documentation
 
Loading...
Searching...
No Matches
gemma4_per_layer_embed.c
Go to the documentation of this file.
1#include <math.h>
2#include <stdint.h>
3#include <stddef.h>
4#include <stdlib.h>
5#include <string.h>
6
7#include "ck_threadpool.h"
8#include "ckernel_quant.h"
9
10typedef struct {
11 ck_half d;
12 ck_half dmin;
13 uint8_t scales[K_SCALE_SIZE];
14 uint8_t qh[QK_K / 8];
15 uint8_t qs[QK_K / 2];
16} ck_gemma4_block_q5_K;
17
18static inline float ck_bf16_to_f32(uint16_t v)
19{
20 uint32_t bits = ((uint32_t)v) << 16;
21 float out;
22 memcpy(&out, &bits, sizeof(out));
23 return out;
24}
25
26static inline float ck_gemma4_gelu(float x)
27{
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)));
31}
32
33static inline void ck_gemma4_unpack_q5_k_scales(const uint8_t *scales, uint8_t *sc, uint8_t *m)
34{
35 sc[0] = scales[0] & 0x3F;
36 sc[1] = scales[1] & 0x3F;
37 sc[2] = scales[2] & 0x3F;
38 sc[3] = scales[3] & 0x3F;
39
40 m[0] = scales[4] & 0x3F;
41 m[1] = scales[5] & 0x3F;
42 m[2] = scales[6] & 0x3F;
43 m[3] = scales[7] & 0x3F;
44
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);
49
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);
54}
55
56static inline uint8_t ck_gemma4_q5_k_value(const ck_gemma4_block_q5_K *block, int subblock, int i)
57{
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);
62}
63
64static void ck_gemma4_dequant_q5_k_block(const ck_gemma4_block_q5_K *block, float *out)
65{
66 uint8_t sc[8];
67 uint8_t m[8];
68 ck_gemma4_unpack_q5_k_scales(block->scales, sc, m);
69 const float d = CK_FP16_TO_FP32(block->d);
70 const float dmin = CK_FP16_TO_FP32(block->dmin);
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) {
75 out[s * 32 + i] = scale * (float)ck_gemma4_q5_k_value(block, s, i) - minv;
76 }
77 }
78}
79
80static void ck_gemma4_rmsnorm_tmp(const float *x, const float *gamma, float *out, int n, float eps)
81{
82 double ss = 0.0;
83 for (int i = 0; i < n; ++i) {
84 ss += (double)x[i] * (double)x[i];
85 }
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];
89 }
90}
91
92typedef struct {
93 float *per_layer_input;
94 const float *hidden;
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;
99 int num_layers;
100 int embed_dim;
101 int per_layer_dim;
102 int vocab_size;
103 float eps;
104} ck_gemma4_prepare_args_t;
105
106static void ck_gemma4_prepare_q5_range(int begin, int end, void *opaque)
107{
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];
119
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;
130 ck_gemma4_dequant_q5_k_block(tok_block, token_vec);
131 for (int i = 0; i < args->per_layer_dim; ++i) {
132 token_vec[i] *= token_scale;
133 }
134
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;
141 float acc = 0.0f;
142 for (int j = 0; j < args->embed_dim; ++j) {
143 acc += ck_bf16_to_f32(row[j]) * h[j];
144 }
145 proj_vec[i] = acc * model_scale;
146 }
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;
152 }
153 }
154 }
155}
156
157static void ck_gemma4_prepare_bf16_range(int begin, int end, void *opaque)
158{
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];
169
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) {
182 token_vec[i] = ck_bf16_to_f32(tok_row[i]) * token_scale;
183 }
184
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;
191 float acc = 0.0f;
192 for (int j = 0; j < args->embed_dim; ++j) {
193 acc += ck_bf16_to_f32(row[j]) * h[j];
194 }
195 proj_vec[i] = acc * model_scale;
196 }
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;
202 }
203 }
204 }
205}
206
208 int tokens, ck_range_fn_t fn, ck_gemma4_prepare_args_t *args)
209{
210 ck_threadpool_t *pool = ck_threadpool_global();
211 int active = pool ? ck_threadpool_n_threads(pool) : 1;
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;
215 ck_threadpool_parallel_for_n(pool, active, 0, tokens, 1, fn, args);
216}
217
218void gemma4_per_layer_prepare_forward(float *per_layer_input,
219 const float *hidden,
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,
224 int tokens,
225 int num_layers,
226 int embed_dim,
227 int per_layer_dim,
228 int vocab_size,
229 float eps)
230{
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) {
234 return;
235 }
236
237 ck_gemma4_prepare_args_t args = {
238 .per_layer_input = per_layer_input,
239 .hidden = hidden,
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,
247 .vocab_size = vocab_size,
248 .eps = eps,
249 };
251}
252
253
254void gemma4_per_layer_prepare_bf16_forward(float *per_layer_input,
255 const float *hidden,
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,
260 int tokens,
261 int num_layers,
262 int embed_dim,
263 int per_layer_dim,
264 int vocab_size,
265 float eps)
266{
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) {
270 return;
271 }
272
273 if (per_layer_dim > QK_K) {
274 return;
275 }
276 ck_gemma4_prepare_args_t args = {
277 .per_layer_input = per_layer_input,
278 .hidden = hidden,
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,
286 .vocab_size = vocab_size,
287 .eps = eps,
288 };
290}
291
292typedef struct {
293 float *hidden;
294 const float *per_layer_input;
295 const float *inp_gate;
296 const float *proj;
297 const float *post_norm;
298 const float *out_scale;
299 int layer;
300 int num_layers;
301 int embed_dim;
302 int per_layer_dim;
303 float eps;
304} ck_gemma4_embed_args_t;
305
306static void ck_gemma4_embed_range(int begin, int end, void *opaque)
307{
308 const ck_gemma4_embed_args_t *args =
309 (const ck_gemma4_embed_args_t *)opaque;
310 float gate_vec[QK_K];
311 float branch[4096];
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;
318
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;
322 float acc = 0.0f;
323 for (int j = 0; j < args->embed_dim; ++j) {
324 acc += row[j] * h[j];
325 }
326 gate_vec[i] = ck_gemma4_gelu(acc) * inp_vec[i];
327 }
328
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;
332 float acc = 0.0f;
333 for (int i = 0; i < args->per_layer_dim; ++i) {
334 acc += row[i] * gate_vec[i];
335 }
336 branch[j] = acc;
337 }
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;
344 }
345 }
346}
347
349 const float *per_layer_input,
350 const float *inp_gate,
351 const float *proj,
352 const float *post_norm,
353 const float *out_scale,
354 int tokens,
355 int layer,
356 int num_layers,
357 int embed_dim,
358 int per_layer_dim,
359 float eps)
360{
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) {
364 return;
365 }
366
367 ck_gemma4_embed_args_t args = {
368 .hidden = hidden,
369 .per_layer_input = per_layer_input,
370 .inp_gate = inp_gate,
371 .proj = proj,
372 .post_norm = post_norm,
373 .out_scale = out_scale,
374 .layer = layer,
375 .num_layers = num_layers,
376 .embed_dim = embed_dim,
377 .per_layer_dim = per_layer_dim,
378 .eps = eps,
379 };
380 ck_threadpool_t *pool = ck_threadpool_global();
381 int active = pool ? ck_threadpool_n_threads(pool) : 1;
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;
386 pool, active, 0, tokens, 1, ck_gemma4_embed_range, &args);
387}
388
390 const float *scale,
391 int tokens,
392 int embed_dim)
393{
394 if (!hidden || !scale || tokens <= 0 || embed_dim <= 0) {
395 return;
396 }
397
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) {
401 hidden[i] *= s;
402 }
403}
404
406 int tokens,
407 int vocab_size,
408 float cap)
409{
410 if (!logits || tokens <= 0 || vocab_size <= 0 || cap <= 0.0f) {
411 return;
412 }
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;
417 }
418}
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 K_SCALE_SIZE
uint16_t ck_half
#define CK_FP16_TO_FP32(x)
#define QK_K
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)
const char * token
Definition tokenizer.h:307
int vocab_size
Definition true_bpe.h:193
uint32_t end
Definition utf8.c:215