← Back to C-Kernel-Engine Docs Doxygen Source Documentation
 
Loading...
Searching...
No Matches
gemma4_per_layer_embed.c File Reference
#include <math.h>
#include <stdint.h>
#include <stddef.h>
#include <stdlib.h>
#include <string.h>
#include "ck_threadpool.h"
#include "ckernel_quant.h"

Go to the source code of this file.

Functions

void assistant_layer_scale_forward (float *hidden, const float *scale, int tokens, int embed_dim)
 
static float ck_bf16_to_f32 (uint16_t v)
 
static void ck_gemma4_dequant_q5_k_block (const ck_gemma4_block_q5_K *block, float *out)
 
static void ck_gemma4_embed_range (int begin, int end, void *opaque)
 
static float ck_gemma4_gelu (float x)
 
static void ck_gemma4_prepare_bf16_range (int begin, int end, void *opaque)
 
static void ck_gemma4_prepare_parallel (int tokens, ck_range_fn_t fn, ck_gemma4_prepare_args_t *args)
 
static void ck_gemma4_prepare_q5_range (int begin, int end, void *opaque)
 
static uint8_t ck_gemma4_q5_k_value (const ck_gemma4_block_q5_K *block, int subblock, int i)
 
static void ck_gemma4_rmsnorm_tmp (const float *x, const float *gamma, float *out, int n, float eps)
 
static void ck_gemma4_unpack_q5_k_scales (const uint8_t *scales, uint8_t *sc, uint8_t *m)
 
void gemma4_final_logit_softcap_forward (float *logits, int tokens, int vocab_size, float cap)
 
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)
 
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_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)
 

Function Documentation

◆ assistant_layer_scale_forward()

void assistant_layer_scale_forward ( float *  hidden,
const float *  scale,
int  tokens,
int  embed_dim 
)

Definition at line 389 of file gemma4_per_layer_embed.c.

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}

◆ ck_bf16_to_f32()

static float ck_bf16_to_f32 ( uint16_t  v)
inlinestatic

Definition at line 18 of file gemma4_per_layer_embed.c.

19{
20 uint32_t bits = ((uint32_t)v) << 16;
21 float out;
22 memcpy(&out, &bits, sizeof(out));
23 return out;
24}

Referenced by ck_gemma4_prepare_bf16_range(), and ck_gemma4_prepare_q5_range().

◆ ck_gemma4_dequant_q5_k_block()

static void ck_gemma4_dequant_q5_k_block ( const ck_gemma4_block_q5_K *  block,
float *  out 
)
static

Definition at line 64 of file gemma4_per_layer_embed.c.

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}
#define CK_FP16_TO_FP32(x)
static void ck_gemma4_unpack_q5_k_scales(const uint8_t *scales, uint8_t *sc, uint8_t *m)
static uint8_t ck_gemma4_q5_k_value(const ck_gemma4_block_q5_K *block, int subblock, int i)

References CK_FP16_TO_FP32, ck_gemma4_q5_k_value(), and ck_gemma4_unpack_q5_k_scales().

Referenced by ck_gemma4_prepare_q5_range().

◆ ck_gemma4_embed_range()

static void ck_gemma4_embed_range ( int  begin,
int  end,
void *  opaque 
)
static

Definition at line 306 of file gemma4_per_layer_embed.c.

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}
#define QK_K
static float ck_gemma4_gelu(float x)
static void ck_gemma4_rmsnorm_tmp(const float *x, const float *gamma, float *out, int n, float eps)
uint32_t end
Definition utf8.c:215

References ck_gemma4_gelu(), ck_gemma4_rmsnorm_tmp(), end, and QK_K.

Referenced by gemma4_per_layer_embed_forward().

◆ ck_gemma4_gelu()

static float ck_gemma4_gelu ( float  x)
inlinestatic

Definition at line 26 of file gemma4_per_layer_embed.c.

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}

Referenced by ck_gemma4_embed_range().

◆ ck_gemma4_prepare_bf16_range()

static void ck_gemma4_prepare_bf16_range ( int  begin,
int  end,
void *  opaque 
)
static

Definition at line 157 of file gemma4_per_layer_embed.c.

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}
static float ck_bf16_to_f32(uint16_t v)
const char * token
Definition tokenizer.h:307

References ck_bf16_to_f32(), ck_gemma4_rmsnorm_tmp(), end, QK_K, and token.

Referenced by gemma4_per_layer_prepare_bf16_forward().

◆ ck_gemma4_prepare_parallel()

static void ck_gemma4_prepare_parallel ( int  tokens,
ck_range_fn_t  fn,
ck_gemma4_prepare_args_t *  args 
)
static

Definition at line 207 of file gemma4_per_layer_embed.c.

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}
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)
int ck_threadpool_n_threads(const ck_threadpool_t *pool)

References ck_threadpool_global(), ck_threadpool_n_threads(), and ck_threadpool_parallel_for_n().

Referenced by gemma4_per_layer_prepare_bf16_forward(), and gemma4_per_layer_prepare_forward().

◆ ck_gemma4_prepare_q5_range()

static void ck_gemma4_prepare_q5_range ( int  begin,
int  end,
void *  opaque 
)
static

Definition at line 106 of file gemma4_per_layer_embed.c.

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}
static void ck_gemma4_dequant_q5_k_block(const ck_gemma4_block_q5_K *block, float *out)

References ck_bf16_to_f32(), ck_gemma4_dequant_q5_k_block(), ck_gemma4_rmsnorm_tmp(), end, QK_K, and token.

Referenced by gemma4_per_layer_prepare_forward().

◆ ck_gemma4_q5_k_value()

static uint8_t ck_gemma4_q5_k_value ( const ck_gemma4_block_q5_K *  block,
int  subblock,
int  i 
)
inlinestatic

Definition at line 56 of file gemma4_per_layer_embed.c.

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}

Referenced by ck_gemma4_dequant_q5_k_block().

◆ ck_gemma4_rmsnorm_tmp()

static void ck_gemma4_rmsnorm_tmp ( const float *  x,
const float *  gamma,
float *  out,
int  n,
float  eps 
)
static

Definition at line 80 of file gemma4_per_layer_embed.c.

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}

Referenced by ck_gemma4_embed_range(), ck_gemma4_prepare_bf16_range(), and ck_gemma4_prepare_q5_range().

◆ ck_gemma4_unpack_q5_k_scales()

static void ck_gemma4_unpack_q5_k_scales ( const uint8_t *  scales,
uint8_t *  sc,
uint8_t *  m 
)
inlinestatic

Definition at line 33 of file gemma4_per_layer_embed.c.

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}

Referenced by ck_gemma4_dequant_q5_k_block().

◆ gemma4_final_logit_softcap_forward()

void gemma4_final_logit_softcap_forward ( float *  logits,
int  tokens,
int  vocab_size,
float  cap 
)

Definition at line 405 of file gemma4_per_layer_embed.c.

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}
int vocab_size
Definition true_bpe.h:193

References vocab_size.

◆ gemma4_per_layer_embed_forward()

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 
)

Definition at line 348 of file gemma4_per_layer_embed.c.

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}
static void ck_gemma4_embed_range(int begin, int end, void *opaque)

References ck_gemma4_embed_range(), ck_threadpool_global(), ck_threadpool_n_threads(), ck_threadpool_parallel_for_n(), and QK_K.

◆ gemma4_per_layer_prepare_bf16_forward()

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 
)

Definition at line 254 of file gemma4_per_layer_embed.c.

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}
static void ck_gemma4_prepare_parallel(int tokens, ck_range_fn_t fn, ck_gemma4_prepare_args_t *args)
static void ck_gemma4_prepare_bf16_range(int begin, int end, void *opaque)

References ck_gemma4_prepare_bf16_range(), ck_gemma4_prepare_parallel(), QK_K, and vocab_size.

◆ gemma4_per_layer_prepare_forward()

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 
)

Definition at line 218 of file gemma4_per_layer_embed.c.

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}
static void ck_gemma4_prepare_q5_range(int begin, int end, void *opaque)

References ck_gemma4_prepare_parallel(), ck_gemma4_prepare_q5_range(), QK_K, and vocab_size.