10 const void *A,
const void *B,
const float *bias,
float *
C,
18 const void *A,
const void *B,
const float *bias,
float *
C,
21 const void *A,
const void *B,
const float *bias,
float *
C,
37 if (!input || !output || rows <= 0 || streams <= 0 || hidden_dim <= 0) {
40 for (
int row = 0; row < rows; ++row) {
41 const float *src = input + (size_t)row * (
size_t)hidden_dim;
42 float *dst = output + (size_t)row * (
size_t)streams * (size_t)hidden_dim;
43 for (
int stream = 0; stream < streams; ++stream) {
44 for (
int col = 0; col < hidden_dim; ++col) {
45 dst[(size_t)stream * (
size_t)hidden_dim + (size_t)col] = src[col];
56 if (!input || !output || rows <= 0 || streams <= 0 || hidden_dim <= 0) {
59 for (
int row = 0; row < rows; ++row) {
60 const float *src = input + (size_t)row * (
size_t)hidden_dim;
61 float *dst = output + (size_t)row * (
size_t)streams * (size_t)hidden_dim;
62 for (
int stream = 0; stream < streams; ++stream) {
63 for (
int col = 0; col < hidden_dim; ++col) {
64 dst[(size_t)stream * (
size_t)hidden_dim + (size_t)col] =
72 const float *norm_weight,
73 const uint16_t *mix_down_weight,
74 const uint16_t *mix_up_weight,
75 const uint16_t *inject_weight,
77 float *injection_output,
78 float *normalized_scratch,
79 float *dynamic_scratch,
87 if (!hyper_input || !norm_weight || !mix_down_weight || !mix_up_weight ||
88 !mixed_output || !normalized_scratch || !dynamic_scratch || !mix_scratch ||
89 rows <= 0 || streams <= 0 || hidden_dim <= 0 || dynamic_dim <= 0) {
92 if (emit_injection && (!inject_weight || !injection_output)) {
96 const int hyper_dim = streams * hidden_dim;
97 const float inv_streams = 1.0f / (float)streams;
99 for (
int row = 0; row < rows; ++row) {
100 const float *input_row =
101 hyper_input + (size_t)row * (
size_t)hyper_dim;
103 normalized_scratch + (size_t)row * (
size_t)hyper_dim;
105 dynamic_scratch + (size_t)row * (
size_t)dynamic_dim;
106 float *mix_row = mix_scratch + (size_t)row * (
size_t)hyper_dim;
108 for (
int stream = 0; stream < streams; ++stream) {
109 const int base = stream * hidden_dim;
122 for (
int out = 0; out < dynamic_dim; ++out) {
123 const uint16_t *weight_row =
124 mix_down_weight + (size_t)out * (
size_t)hyper_dim;
126 for (
int col = 0; col < hyper_dim; ++col) {
131 projected / (1.0f + expf(-projected)));
134 for (
int out = 0; out < hyper_dim; ++out) {
135 const uint16_t *weight_row =
136 mix_up_weight + (size_t)out * (
size_t)dynamic_dim;
138 for (
int col = 0; col < dynamic_dim; ++col) {
145 mixed_output + (size_t)row * (
size_t)hidden_dim;
146 for (
int col = 0; col < hidden_dim; ++col) {
148 for (
int stream = 0; stream < streams; ++stream) {
149 const int index = stream * hidden_dim + col;
155 if (emit_injection) {
156 float *injection_row =
157 injection_output + (size_t)row * (
size_t)streams;
158 for (
int stream = 0; stream < streams; ++stream) {
159 const uint16_t *weight_row =
160 inject_weight + (size_t)stream * (
size_t)hyper_dim;
162 for (
int col = 0; col < hyper_dim; ++col) {
173 const void *,
const void *,
const float *,
float *, int, int, int);
183 if (!input || !weight || !output || rows <= 0 || output_dim <= 0 ||
184 input_dim <= 0 || input_dim %
QK_K != 0) {
188 const size_t input_row_bytes =
190 for (
int row = 0; row < rows; ++row) {
191 float *output_row = output + (size_t)row * (
size_t)output_dim;
192 const void *input_row =
193 (
const uint8_t *)input + (
size_t)row * input_row_bytes;
195 output_row, weight, input_row, output_dim, input_dim);
197 for (
int col = 0; col < output_dim; ++col) {
198 output_row[col] += bias[col];
205 const float *norm_weight,
206 const void *mix_down_weight,
207 const void *mix_up_weight,
208 const void *inject_weight,
210 float *injection_output,
211 float *normalized_scratch,
212 float *dynamic_scratch,
222 if (!hyper_input || !norm_weight || !mix_down_weight || !mix_up_weight ||
223 !mixed_output || !normalized_scratch || !dynamic_scratch || !mix_scratch ||
224 !injection_gemm || !down_gemm || rows <= 0 || streams <= 0 || hidden_dim <= 0 ||
228 if (emit_injection && (!inject_weight || !injection_output)) {
232 const int hyper_dim = streams * hidden_dim;
233 const float inv_streams = 1.0f / (float)streams;
234 if (hyper_dim %
QK_K != 0 || dynamic_dim %
QK8_0 != 0) {
238 const size_t normalized_q8_row_bytes =
240 const size_t dynamic_q8_row_bytes =
245 for (
int row = 0; row < rows; ++row) {
246 const float *input_row = hyper_input + (size_t)row * (
size_t)hyper_dim;
247 float *norm_row = normalized_scratch + (size_t)row * (
size_t)hyper_dim;
248 for (
int stream = 0; stream < streams; ++stream) {
249 const int base = stream * hidden_dim;
251 for (
int col = 0; col < hidden_dim; ++col) {
252 const float value = input_row[base + col];
253 sum_sq += (double)(value * value);
255 const float mean = (float)(sum_sq / (
double)hidden_dim);
256 const float rstd = 1.0f / sqrtf(mean + eps);
257 for (
int col = 0; col < hidden_dim; ++col) {
258 const int index = base + col;
259 norm_row[index] = input_row[index] * rstd * norm_weight[index];
265 (uint8_t *)normalized_q8 + (
size_t)row * normalized_q8_row_bytes,
278 if (emit_injection) {
287 for (
int row = 0; row < rows; ++row) {
288 float *injection_row = injection_output + (size_t)row * (
size_t)streams;
289 for (
int stream = 0; stream < streams; ++stream) {
290 injection_row[stream] *= inv_streams;
293 injection_row, injection_row, 1, streams);
294 for (
int stream = 0; stream < streams; ++stream) {
295 injection_row[stream] *= 2.0f;
300 for (
int row = 0; row < rows; ++row) {
301 float *dynamic_row = dynamic_scratch + (size_t)row * (
size_t)dynamic_dim;
302 for (
int col = 0; col < dynamic_dim; ++col) {
303 dynamic_row[col] *= inv_streams;
306 dynamic_row, dynamic_row, 1, dynamic_dim);
310 (uint8_t *)dynamic_scratch + (
size_t)row * dynamic_q8_row_bytes,
312 dynamic_q8_row_bytes);
324 for (
int row = 0; row < rows; ++row) {
325 const float *norm_row =
326 normalized_scratch + (size_t)row * (
size_t)hyper_dim;
327 float *mix_row = mix_scratch + (size_t)row * (
size_t)hyper_dim;
329 mix_row, mix_row, 1, hyper_dim);
331 float *mixed_row = mixed_output + (size_t)row * (
size_t)hidden_dim;
332 for (
int col = 0; col < hidden_dim; ++col) {
334 for (
int stream = 0; stream < streams; ++stream) {
335 const int index = stream * hidden_dim + col;
336 sum += norm_row[index] * mix_row[index];
338 mixed_row[col] = sum * inv_streams;
345 const float *norm_weight,
346 const void *mix_down_weight,
347 const void *mix_up_weight,
348 const void *inject_weight,
350 float *injection_output,
351 float *normalized_scratch,
352 float *dynamic_scratch,
359 int emit_injection) {
361 hyper_input, norm_weight, mix_down_weight, mix_up_weight, inject_weight,
362 mixed_output, injection_output, normalized_scratch, dynamic_scratch,
363 mix_scratch, rows, streams, hidden_dim, dynamic_dim, eps,
369 const float *norm_weight,
370 const void *mix_down_weight,
371 const void *mix_up_weight,
372 const void *inject_weight,
374 float *injection_output,
375 float *normalized_scratch,
376 float *dynamic_scratch,
383 int emit_injection) {
385 hyper_input, norm_weight, mix_down_weight, mix_up_weight, inject_weight,
386 mixed_output, injection_output, normalized_scratch, dynamic_scratch,
387 mix_scratch, rows, streams, hidden_dim, dynamic_dim, eps,
393 const float *block_output,
394 const float *injection_weight,
399 if (!hyper_input || !block_output || !injection_weight || !output ||
400 rows <= 0 || streams <= 0 || hidden_dim <= 0) {
403 const int hyper_dim = streams * hidden_dim;
404 for (
int row = 0; row < rows; ++row) {
405 const float *hyper_row =
406 hyper_input + (size_t)row * (
size_t)hyper_dim;
407 const float *block_row =
408 block_output + (size_t)row * (
size_t)hidden_dim;
409 const float *inject_row =
410 injection_weight + (size_t)row * (
size_t)streams;
411 float *output_row = output + (size_t)row * (
size_t)hyper_dim;
412 for (
int stream = 0; stream < streams; ++stream) {
413 for (
int col = 0; col < hidden_dim; ++col) {
414 const int index = stream * hidden_dim + col;
424 const float *block_output,
425 const float *injection_weight,
430 if (!hyper_input || !block_output || !injection_weight || !output ||
431 rows <= 0 || streams <= 0 || hidden_dim <= 0) {
434 const int hyper_dim = streams * hidden_dim;
435 for (
int row = 0; row < rows; ++row) {
436 const float *hyper_row =
437 hyper_input + (size_t)row * (
size_t)hyper_dim;
438 const float *block_row =
439 block_output + (size_t)row * (
size_t)hidden_dim;
440 const float *inject_row =
441 injection_weight + (size_t)row * (
size_t)streams;
442 float *output_row = output + (size_t)row * (
size_t)hyper_dim;
443 for (
int stream = 0; stream < streams; ++stream) {
444 for (
int col = 0; col < hidden_dim; ++col) {
445 const int index = stream * hidden_dim + col;
446 volatile const float weighted =
447 block_row[col] * inject_row[stream];
448 output_row[index] = hyper_row[index] + weighted;
static uint16_t float_to_bf16(float f)
static float bf16_to_float(uint16_t v)
void rmsnorm_forward_qwen3next_pytorch_bf16_storage(const float *input, const float *gamma, float *output, float *rstd_cache, int tokens, int d_model, int aligned_embed_dim, float eps)
void recurrent_sigmoid_forward_ggml(const float *x, float *out, int rows, int dim)
void recurrent_silu_forward_ggml(const float *x, float *out, int rows, int dim)
void quantize_row_q8_k(const float *x, void *y, int k)
void quantize_row_q8_0(const float *x, void *y, int k)
Quantize FP32 to Q8_0 format (scalar reference)
void hyper_connection_mix_q4k_q5_0_q4k(const float *hyper_input, const float *norm_weight, const void *mix_down_weight, const void *mix_up_weight, const void *inject_weight, float *mixed_output, float *injection_output, float *normalized_scratch, float *dynamic_scratch, float *mix_scratch, int rows, int streams, int hidden_dim, int dynamic_dim, float eps, int emit_injection)
static float ck_sigmoid_bf16(float value)
void gemv_q4_k_q8_k_avx2(float *y, const void *W, const void *x_q8, int M, int K)
void hyper_connection_mix_bf16(const float *hyper_input, const float *norm_weight, const uint16_t *mix_down_weight, const uint16_t *mix_up_weight, const uint16_t *inject_weight, float *mixed_output, float *injection_output, float *normalized_scratch, float *dynamic_scratch, float *mix_scratch, int rows, int streams, int hidden_dim, int dynamic_dim, float eps, int emit_injection)
void hyper_stream_expand_f32(const float *input, float *output, int rows, int streams, int hidden_dim)
void hyper_stream_inject_bf16(const float *hyper_input, const float *block_output, const float *injection_weight, float *output, int rows, int streams, int hidden_dim)
void hyper_connection_mix_q6k_q5_0_q4k(const float *hyper_input, const float *norm_weight, const void *mix_down_weight, const void *mix_up_weight, const void *inject_weight, float *mixed_output, float *injection_output, float *normalized_scratch, float *dynamic_scratch, float *mix_scratch, int rows, int streams, int hidden_dim, int dynamic_dim, float eps, int emit_injection)
void gemm_nt_q6_k_q8_k_parallel_dispatch(const void *A, const void *B, const float *bias, float *C, int M, int N, int K)
static float ck_bf16_round(float value)
void hyper_stream_inject_f32(const float *hyper_input, const float *block_output, const float *injection_weight, float *output, int rows, int streams, int hidden_dim)
void(* ck_hyper_q8k_gemm_fn)(const void *, const void *, const float *, float *, int, int, int)
void gemm_nt_q4_k_q8_k_pairwise_split_min_parallel_dispatch(const void *A, const void *B, const float *bias, float *C, int M, int N, int K)
static void hyper_connection_mix_quantized(const float *hyper_input, const float *norm_weight, const void *mix_down_weight, const void *mix_up_weight, const void *inject_weight, float *mixed_output, float *injection_output, float *normalized_scratch, float *dynamic_scratch, float *mix_scratch, int rows, int streams, int hidden_dim, int dynamic_dim, float eps, int emit_injection, ck_hyper_q8k_gemm_fn injection_gemm, ck_hyper_q8k_gemm_fn down_gemm)
void hyper_stream_expand_bf16(const float *input, float *output, int rows, int streams, int hidden_dim)
static void hyper_injection_q4k_q8k_llama_dispatch(const void *input, const void *weight, const float *bias, float *output, int rows, int output_dim, int input_dim)
void gemm_nt_q5_0_q8_0_parallel_dispatch(const void *A, const void *B, const float *bias, float *C, int M, int N, int K)