| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| #define LAL_PONDER_IMPLEMENTATION |
| #include "lal_runtime.h" |
| #include "lal_whitebox_probe.h" |
| #include "lal_concept_gen.h" |
| #include "lal_concept_attn.h" |
|
|
| |
| #ifdef _WIN32 |
| #include <malloc.h> |
| #else |
| #include <stdlib.h> |
| #endif |
|
|
| |
| |
| |
| #ifdef HAVE_OPENBLAS |
| #include <cblas.h> |
| |
| |
| |
| static int g_openblas_inited = 0; |
| static void openblas_init_single_thread(void) { |
| if (!g_openblas_inited) { |
| openblas_set_num_threads(1); |
| g_openblas_inited = 1; |
| } |
| } |
| #endif |
|
|
| |
| |
| |
| |
|
|
| |
| |
| |
| #ifdef _WIN32 |
| #define WIN32_LEAN_AND_MEAN |
| #include <windows.h> |
| #include <io.h> |
|
|
| |
| #ifndef MAP_FAILED |
| #define MAP_FAILED ((void *)-1) |
| #endif |
| #ifndef PROT_READ |
| #define PROT_READ 0x1 |
| #define MAP_PRIVATE 0x2 |
| #endif |
| static inline void *mmap(void *addr, size_t length, int prot, int flags, int fd, long long offset) { |
| (void)addr; (void)prot; (void)flags; |
| HANDLE h = CreateFileMappingA((HANDLE)_get_osfhandle(fd), NULL, PAGE_READONLY, 0, 0, NULL); |
| if (!h) return MAP_FAILED; |
| void *p = MapViewOfFile(h, FILE_MAP_READ, 0, 0, length); |
| CloseHandle(h); |
| return p ? p : MAP_FAILED; |
| } |
| static inline int munmap(void *addr, size_t length) { |
| (void)length; |
| UnmapViewOfFile(addr); |
| return 0; |
| } |
|
|
| |
| static inline int rand_r(unsigned int *seedp) { |
| *seedp = *seedp * 1103515245u + 12345u; |
| return (int)((*seedp / 65536u) % 32768u); |
| } |
|
|
| |
| #include <sys/stat.h> |
| #define fstat _fstat |
| #define stat _stat |
| #else |
| #include <sys/mman.h> |
| #include <sys/stat.h> |
| #include <unistd.h> |
| #endif |
|
|
| |
| |
| |
| float cross_entropy_full(const float *hidden, const float *wte, |
| int target, int vocab_size, int n_embd, |
| float *logits_scratch); |
| void cross_entropy_full_grad(float *grad_hidden, const float *hidden, const float *wte, |
| int target, int vocab_size, int n_embd, |
| float *logits_scratch); |
|
|
| |
| |
| |
|
|
| void rms_norm(float *out, const float *x, const float *w, int n) { |
| float ms = 0; |
| for (int i = 0; i < n; i++) ms += x[i] * x[i]; |
| ms = 1.0f / sqrtf(ms / n + 1e-5f); |
| for (int i = 0; i < n; i++) out[i] = x[i] * ms * w[i]; |
| } |
|
|
| void rms_norm_backward(float *grad_x, const float *grad_y, const float *x, |
| const float *w, int n, float *grad_w) { |
| float ms = 0; |
| for (int i = 0; i < n; i++) ms += x[i] * x[i]; |
| ms = 1.0f / sqrtf(ms / n + 1e-5f); |
| for (int i = 0; i < n; i++) { |
| grad_x[i] = grad_y[i] * w[i] * ms; |
| if (grad_w) grad_w[i] += grad_y[i] * x[i] * ms; |
| } |
| } |
|
|
| float silu(float x) { return x / (1.0f + expf(-x)); } |
| float silu_grad(float x) { |
| float s = 1.0f / (1.0f + expf(-x)); |
| return s + x * s * (1.0f - s); |
| } |
|
|
| void norm_forward(float *out, const float *x, const float *w, const float *b, |
| NormType type, int n) { |
| if (type == NORM_RMS) rms_norm(out, x, w, n); |
| else layer_norm(out, x, w, b, n); |
| } |
|
|
| void norm_backward(float *grad_x, const float *grad_y, const float *x, |
| const float *w, const float *cached, NormType type, int n, |
| float *grad_w, float *grad_b) { |
| if (type == NORM_RMS) rms_norm_backward(grad_x, grad_y, x, w, n, grad_w); |
| else layer_norm_backward(grad_x, grad_y, x, w, cached[0], cached[1], n, grad_w, grad_b); |
| } |
|
|
| float act_forward(float x, ActType type) { |
| switch (type) { |
| case ACT_GELU: return gelu(x); |
| case ACT_SWIGLU: return silu(x); |
| case ACT_SILU: return silu(x); |
| default: return x; |
| } |
| } |
|
|
| float act_grad(float x, ActType type) { |
| switch (type) { |
| case ACT_GELU: return gelu_grad(x); |
| case ACT_SWIGLU: return silu_grad(x); |
| case ACT_SILU: return silu_grad(x); |
| default: return 1.0f; |
| } |
| } |
|
|
| void apply_rope(float *q, float *k, int seq_len, int n_head, int head_dim, int n_embd) { |
| |
| for (int h = 0; h < n_head; h++) { |
| float *qh = q + h * head_dim; |
| float *kh = k + h * head_dim; |
| for (int d = 0; d < head_dim / 2; d++) { |
| float angle = (float)seq_len / powf(10000.0f, (float)(2 * d) / head_dim); |
| float c = cosf(angle), s = sinf(angle); |
| float q0 = qh[d], q1 = qh[d + head_dim / 2]; |
| float k0 = kh[d], k1 = kh[d + head_dim / 2]; |
| qh[d] = q0 * c - q1 * s; |
| qh[d + head_dim / 2] = q0 * s + q1 * c; |
| kh[d] = k0 * c - k1 * s; |
| kh[d + head_dim / 2] = k0 * s + k1 * c; |
| } |
| } |
| } |
|
|
| |
| |
| |
|
|
| void trans_layer_init(TransLayer *tl, Tensor *tensors, int n_tensors, |
| ModelConfig *cfg, int layer_idx, |
| const char *qkv_key, const char *q_key, const char *k_key, |
| const char *v_key, const char *o_key, |
| const char *gate_key, const char *up_key, const char *down_key, |
| const char *norm1_w_key, const char *norm1_b_key, |
| const char *norm2_w_key, const char *norm2_b_key) { |
| tl->layer_idx = layer_idx; |
| int n = cfg->n_embd, m = cfg->mlp_dim; |
| tl->_kv_k = NULL; |
| tl->_kv_v = NULL; |
| char full_key[256]; |
|
|
| if (cfg->qkv_merged) { |
| |
| sprintf(full_key, qkv_key, layer_idx); |
| float *W = tensor_get(tensors, n_tensors, full_key); |
| |
| |
| |
| |
| |
| char bias_key[256]; |
| strncpy(bias_key, full_key, sizeof(bias_key) - 1); |
| bias_key[sizeof(bias_key) - 1] = '\0'; |
| char *dot = strstr(bias_key, ".weight"); |
| if (dot) { *dot = 0; strncat(bias_key, ".bias", sizeof(bias_key) - strlen(bias_key) - 1); } |
| float *b = tensor_get(tensors, n_tensors, bias_key); |
| bin_layer_init(&tl->attn_q, W, b, n, 3 * n); |
| } else { |
| |
| sprintf(full_key, q_key, layer_idx); |
| char bias_key[256]; |
| float *Wq = tensor_get(tensors, n_tensors, full_key); |
| strncpy(bias_key, full_key, sizeof(bias_key)); |
| char *dot = strstr(bias_key, ".weight"); if (dot) { *dot=0; strcat(bias_key, ".bias"); } |
| float *bq = tensor_get(tensors, n_tensors, bias_key); |
| bin_layer_init(&tl->attn_q, Wq, bq, n, n); |
|
|
| sprintf(full_key, k_key, layer_idx); |
| float *Wk = tensor_get(tensors, n_tensors, full_key); |
| strncpy(bias_key, full_key, sizeof(bias_key)); |
| dot = strstr(bias_key, ".weight"); if (dot) { *dot=0; strcat(bias_key, ".bias"); } |
| float *bk = tensor_get(tensors, n_tensors, bias_key); |
| bin_layer_init(&tl->attn_k, Wk, bk, n, n); |
|
|
| sprintf(full_key, v_key, layer_idx); |
| float *Wv = tensor_get(tensors, n_tensors, full_key); |
| strncpy(bias_key, full_key, sizeof(bias_key)); |
| dot = strstr(bias_key, ".weight"); if (dot) { *dot=0; strcat(bias_key, ".bias"); } |
| float *bv = tensor_get(tensors, n_tensors, bias_key); |
| bin_layer_init(&tl->attn_v, Wv, bv, n, n); |
| } |
|
|
| |
| sprintf(full_key, o_key, layer_idx); |
| float *Wo = tensor_get(tensors, n_tensors, full_key); |
| char bias_key[256]; strncpy(bias_key, full_key, sizeof(bias_key)); |
| char *dot = strstr(bias_key, ".weight"); if (dot) { *dot=0; strcat(bias_key, ".bias"); } |
| float *bo = tensor_get(tensors, n_tensors, bias_key); |
| bin_layer_init(&tl->attn_o, Wo, bo, n, n); |
|
|
| |
| if (cfg->act_type == ACT_SWIGLU) { |
| sprintf(full_key, gate_key, layer_idx); |
| float *Wg = tensor_get(tensors, n_tensors, full_key); |
| strncpy(bias_key, full_key, sizeof(bias_key)); |
| dot = strstr(bias_key, ".weight"); if (dot) { *dot=0; strcat(bias_key, ".bias"); } |
| float *bg = tensor_get(tensors, n_tensors, bias_key); |
| bin_layer_init(&tl->mlp_gate, Wg, bg, n, m); |
|
|
| sprintf(full_key, up_key, layer_idx); |
| float *Wu = tensor_get(tensors, n_tensors, full_key); |
| strncpy(bias_key, full_key, sizeof(bias_key)); |
| dot = strstr(bias_key, ".weight"); if (dot) { *dot=0; strcat(bias_key, ".bias"); } |
| float *bu = tensor_get(tensors, n_tensors, bias_key); |
| bin_layer_init(&tl->mlp_up, Wu, bu, n, m); |
| } else { |
| |
| sprintf(full_key, gate_key, layer_idx); |
| float *Wg = tensor_get(tensors, n_tensors, full_key); |
| strncpy(bias_key, full_key, sizeof(bias_key)); |
| dot = strstr(bias_key, ".weight"); if (dot) { *dot=0; strcat(bias_key, ".bias"); } |
| float *bg = tensor_get(tensors, n_tensors, bias_key); |
| bin_layer_init(&tl->mlp_gate, Wg, bg, n, m); |
| } |
|
|
| sprintf(full_key, down_key, layer_idx); |
| float *Wd = tensor_get(tensors, n_tensors, full_key); |
| strncpy(bias_key, full_key, sizeof(bias_key)); |
| dot = strstr(bias_key, ".weight"); if (dot) { *dot=0; strcat(bias_key, ".bias"); } |
| float *bd = tensor_get(tensors, n_tensors, bias_key); |
| bin_layer_init(&tl->mlp_down, Wd, bd, m, n); |
|
|
| |
| sprintf(full_key, norm1_w_key, layer_idx); |
| tl->norm1_w = tensor_get(tensors, n_tensors, full_key); |
| sprintf(full_key, norm1_b_key, layer_idx); |
| tl->norm1_b = tensor_get(tensors, n_tensors, full_key); |
| sprintf(full_key, norm2_w_key, layer_idx); |
| tl->norm2_w = tensor_get(tensors, n_tensors, full_key); |
| sprintf(full_key, norm2_b_key, layer_idx); |
| tl->norm2_b = tensor_get(tensors, n_tensors, full_key); |
| } |
|
|
| void trans_layer_free(TransLayer *tl, ModelConfig *cfg) { |
| bin_layer_free(&tl->attn_q); |
| if (!cfg->qkv_merged) { bin_layer_free(&tl->attn_k); bin_layer_free(&tl->attn_v); } |
| bin_layer_free(&tl->attn_o); |
| bin_layer_free(&tl->mlp_gate); |
| if (cfg->act_type == ACT_SWIGLU) bin_layer_free(&tl->mlp_up); |
| bin_layer_free(&tl->mlp_down); |
| } |
|
|
| |
| static inline void bin_fwd(float *y, const float *x, const BinLayer *bl) { |
| if (g_use_pure_float) bin_forward_pure_float(y, x, bl); |
| else bin_forward(y, x, bl); |
| } |
|
|
|
|
| |
| |
| |
| |
| |
|
|
| |
| |
| |
| |
| |
| |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| #include <omp.h> |
| #define LAL_MAX_THREADS 16 |
| int g_cur_tid = 0; |
| #pragma omp threadprivate(g_cur_tid) |
|
|
| typedef struct { |
| TransAct *acts; |
| TransAct *scratch; |
| |
| PonderBuf ponder; |
| float *ponder_mix; |
| float *ponder_state; |
| float *ponder_kv0k; |
| float *ponder_kv0v; |
| TransAct *rec_acts; |
| int ponder_first_rec_step; |
| int ponder_ready; |
| float *mlp, *hidden, *norm2, *proj, *attn, *qkv, *norm1, *pre; |
| float *gate, *up, *norm2_gate, *norm2_up; |
| float *n1k, *n1v; |
| float *xc, *x, *gh; |
| float *g_pre4; |
| float *full_logits; int full_logits_vocab; |
| int forward_done; |
| float *final_ln; |
| float *x_before_final; float final_mean, final_std_inv; |
| |
| float ***grad_w; |
| float ***grad_b; |
| float *grad_wte, *grad_wpe, *grad_lnfw, *grad_lnfb; |
| |
| float **grad_norm1_w, **grad_norm1_b, **grad_norm2_w, **grad_norm2_b; |
| int n_layer, n_bl_max; |
| } ThrRes; |
| static ThrRes g_thr[LAL_MAX_THREADS]; |
| static int g_thr_inited = 0; |
| static int g_nthr = 1; |
|
|
| void thr_res_alloc(Model *m); |
| void thr_res_free(void); |
|
|
| |
| |
| |
| |
| static float g_sign_lut[256][8]; |
| static int g_sign_lut_init = 0; |
| static void sign_lut_ensure(void) { |
| if (g_sign_lut_init) return; |
| for (int b = 0; b < 256; b++) |
| for (int i = 0; i < 8; i++) |
| g_sign_lut[b][i] = (b >> i) & 1 ? 1.0f : -1.0f; |
| g_sign_lut_init = 1; |
| } |
| void thr_grad_reduce(Model *m); |
|
|
| int g_use_ste = 1; |
| float g_attn_residual_scale = 1.0f; |
| int g_use_logic_binarization = 1; |
|
|
| |
| |
| |
| |
| float g_logic_core_ratio = 0.0f; |
| float g_logic_prune_ratio = 0.0f; |
|
|
| |
| |
| |
| int g_use_adam = 0; |
| int g_opt_step = 0; |
| float g_adam_beta1 = 0.9f; |
| float g_adam_beta2 = 0.999f; |
| float g_adam_eps = 1e-8f; |
|
|
| |
| |
| |
| |
| |
| |
| int g_use_ternary = 1; |
| float g_ternary_delta_factor = 0.7f; |
|
|
| |
| int g_merge_mode = 0; |
| float g_merge_beta_lo = 0.5f; |
| float g_merge_beta_hi = 0.9f; |
|
|
| |
| |
| |
| |
| |
| |
| |
| float lr_schedule(int step, int warmup_steps, int total_steps, float base_lr) { |
| if (warmup_steps > 0 && step < warmup_steps) { |
| return base_lr * (float)(step + 1) / (float)warmup_steps; |
| } |
| if (total_steps <= warmup_steps) return base_lr; |
| if (step >= total_steps) return base_lr * 0.01f; |
| float progress = (float)(step - warmup_steps) / (float)(total_steps - warmup_steps); |
| return base_lr * 0.5f * (1.0f + cosf((float)M_PI * progress)); |
| } |
|
|
| |
| |
| |
| |
| |
| |
| void bin_forward_pure_float(float *y, const float *x, const BinLayer *bl) { |
| int in = bl->in_dim, out = bl->out_dim, nw = bl->n_words; |
| if (bl->logic_mask) { |
| |
| |
| |
| int *cidx = (int *)alloca(out * sizeof(int)); |
| { int c = 0; for (int j = 0; j < out; j++) { cidx[j] = c; if (bl->logic_mask[j] == 0) c++; } } |
| if (out >= 64) { |
| #pragma omp parallel for schedule(guided, 8) |
| for (int j = 0; j < out; j++) { |
| uint8_t m = bl->logic_mask[j]; |
| if (m == 0) { |
| const float *wc = &bl->w_core[cidx[j] * in]; |
| float s = bl->bias[j]; |
| for (int i = 0; i < in; i++) s += wc[i] * x[i]; |
| y[j] = s; |
| } else if (m == 1) { |
| const float *wf = &bl->w_float[j * in]; |
| const uint64_t *zb = (g_use_ternary && bl->zbits) |
| ? &bl->zbits[j * nw] : NULL; |
| if (g_use_ternary && bl->zbits) { |
| |
| |
| |
| |
| |
| |
| float abs_sum = 0.0f; |
| for (int i = 0; i < in; i++) abs_sum += fabsf(x[i]); |
| float K = abs_sum / in; |
| float s = 0.0f; |
| for (int i = 0; i < in; i++) { |
| if (zb && (((zb[i >> 6] >> (i & 63)) & 1))) continue; |
| float w_sign = (wf[i] > 0.0f) ? 1.0f : (wf[i] < 0.0f ? -1.0f : 0.0f); |
| s += w_sign * x[i]; |
| } |
| y[j] = s * bl->alpha[j] * K * g_binary_scale + bl->bias[j]; |
| } else { |
| |
| float s = bl->bias[j]; |
| for (int i = 0; i < in; i++) s += wf[i] * x[i]; |
| y[j] = bl->bias[j] + (s - bl->bias[j]) * g_binary_scale; |
| } |
| } else { |
| y[j] = 0.0f; |
| } |
| } |
| } else { |
| |
| for (int j = 0; j < out; j++) { |
| uint8_t m = bl->logic_mask[j]; |
| if (m == 0) { |
| const float *wc = &bl->w_core[cidx[j] * in]; |
| float s = bl->bias[j]; |
| for (int i = 0; i < in; i++) s += wc[i] * x[i]; |
| y[j] = s; |
| } else if (m == 1) { |
| const float *wf = &bl->w_float[j * in]; |
| const uint64_t *zb = (g_use_ternary && bl->zbits) |
| ? &bl->zbits[j * nw] : NULL; |
| if (g_use_ternary && bl->zbits) { |
| float abs_sum = 0.0f; |
| for (int i = 0; i < in; i++) abs_sum += fabsf(x[i]); |
| float K = abs_sum / in; |
| float s = 0.0f; |
| for (int i = 0; i < in; i++) { |
| if (zb && (((zb[i >> 6] >> (i & 63)) & 1))) continue; |
| float w_sign = (wf[i] > 0.0f) ? 1.0f : (wf[i] < 0.0f ? -1.0f : 0.0f); |
| s += w_sign * x[i]; |
| } |
| y[j] = s * bl->alpha[j] * K * g_binary_scale + bl->bias[j]; |
| } else { |
| float s = bl->bias[j]; |
| for (int i = 0; i < in; i++) s += wf[i] * x[i]; |
| y[j] = bl->bias[j] + (s - bl->bias[j]) * g_binary_scale; |
| } |
| } else { |
| y[j] = 0.0f; |
| } |
| } |
| } |
| } else { |
| |
| if (out >= 64) { |
| #pragma omp parallel for schedule(guided, 8) |
| for (int j = 0; j < out; j++) { |
| const float *wf = &bl->w_float[j * in]; |
| float s = bl->bias[j]; |
| for (int i = 0; i < in; i++) s += wf[i] * x[i]; |
| y[j] = s; |
| } |
| } else { |
| for (int j = 0; j < out; j++) { |
| const float *wf = &bl->w_float[j * in]; |
| float s = bl->bias[j]; |
| for (int i = 0; i < in; i++) s += wf[i] * x[i]; |
| y[j] = s; |
| } |
| } |
| } |
| } |
|
|
| |
| |
| |
| |
| static void compute_norm_mask(const float *W, int in_dim, int out_dim, uint8_t *mask) { |
| |
| float *norms = malloc(out_dim * sizeof(float)); |
| for (int j = 0; j < out_dim; j++) { |
| float s = 0; |
| for (int i = 0; i < in_dim; i++) { |
| float w = W[i * out_dim + j]; |
| s += w * w; |
| } |
| norms[j] = sqrtf(s); |
| } |
| |
| float *sorted = malloc(out_dim * sizeof(float)); |
| memcpy(sorted, norms, out_dim * sizeof(float)); |
| |
| for (int i = 1; i < out_dim; i++) { |
| float v = sorted[i]; int k = i - 1; |
| while (k >= 0 && sorted[k] > v) { sorted[k+1] = sorted[k]; k--; } |
| sorted[k+1] = v; |
| } |
|
|
| |
| float core_r = (g_logic_core_ratio > 0.0f) ? g_logic_core_ratio : 0.20f; |
| float prune_r = (g_logic_prune_ratio > 0.0f) ? g_logic_prune_ratio : 0.10f; |
|
|
| int core_count = (int)(out_dim * core_r); |
| int prune_count = (int)(out_dim * prune_r); |
| if (core_count < 1) core_count = 1; |
| if (core_count + prune_count > out_dim) prune_count = out_dim - core_count; |
|
|
| |
| float core_threshold = sorted[out_dim - core_count]; |
| float prune_threshold = (prune_count > 0) ? sorted[prune_count - 1] : -1.0f; |
|
|
| int n_core = 0, n_binary = 0, n_prune = 0; |
| for (int j = 0; j < out_dim; j++) { |
| if (norms[j] >= core_threshold && n_core < core_count) { |
| mask[j] = 0; |
| n_core++; |
| } else if (norms[j] <= prune_threshold && n_prune < prune_count) { |
| mask[j] = 2; |
| n_prune++; |
| } else { |
| mask[j] = 1; |
| n_binary++; |
| } |
| } |
|
|
| static int first_call = 1; |
| if (first_call) { |
| printf(" [logic] CORE=%d (%.0f%%), BINARY=%d (%.0f%%), PRUNE=%d (%.0f%%)\n", |
| n_core, 100.0f * n_core / out_dim, |
| n_binary, 100.0f * n_binary / out_dim, |
| n_prune, 100.0f * n_prune / out_dim); |
| first_call = 0; |
| } |
|
|
| free(norms); free(sorted); |
| } |
|
|
| float g_core_lr_multiplier = 3.0f; |
| int g_use_lal_adam = 1; |
| float g_prune_decay = 0.01f; |
| float g_prune_freeze_thresh = 0.001f; |
|
|
| |
| |
| |
| int g_use_real_attention = 0; |
| int g_skip_wv = 0; |
| |
| |
| |
| |
| |
| |
| |
| |
| int g_attn_window = 1024; |
| int g_attn_sink = 64; |
| int g_use_pure_float = 0; |
| |
| float g_binary_scale = 0.25f; |
| |
| float g_wte_lr_scale = 0.5f; |
| |
| |
| |
| float g_logit_scale = 1.0f; |
| int g_accumulate_gradients = 0; |
|
|
| TransAct *trans_act_alloc(ModelConfig *cfg) { |
| int n = cfg->n_embd, m = cfg->mlp_dim; |
| TransAct *acts = malloc(cfg->n_layer * sizeof(TransAct)); |
| for (int l = 0; l < cfg->n_layer; l++) { |
| acts[l].x_pre_norm1 = malloc(n * sizeof(float)); |
| acts[l].norm1_out = malloc(n * sizeof(float)); |
| acts[l].q = malloc(3 * n * sizeof(float)); |
| |
| |
| |
| acts[l].k = acts[l].q + n; |
| acts[l].v = acts[l].q + 2 * n; |
| acts[l].attn_out = malloc(n * sizeof(float)); |
| acts[l].proj_out = malloc(n * sizeof(float)); |
| acts[l].x_pre_norm2 = malloc(n * sizeof(float)); |
| acts[l].norm2_out = malloc(n * sizeof(float)); |
| acts[l].mlp_hidden = malloc(m * sizeof(float)); |
| acts[l].mlp_out = malloc(n * sizeof(float)); |
| |
| if (cfg->act_type == ACT_SWIGLU) { |
| acts[l].swiglu_gate = malloc(m * sizeof(float)); |
| acts[l].swiglu_up = malloc(m * sizeof(float)); |
| } else { |
| acts[l].swiglu_gate = NULL; |
| acts[l].swiglu_up = NULL; |
| } |
| } |
| return acts; |
| } |
|
|
| void trans_act_free(TransAct *acts, int n_layer) { |
| for (int l = 0; l < n_layer; l++) { |
| free(acts[l].x_pre_norm1); free(acts[l].norm1_out); |
| free(acts[l].q); free(acts[l].attn_out); free(acts[l].proj_out); |
| free(acts[l].x_pre_norm2); free(acts[l].norm2_out); |
| free(acts[l].mlp_hidden); free(acts[l].mlp_out); |
| free(acts[l].swiglu_gate); free(acts[l].swiglu_up); |
| } |
| free(acts); |
| } |
|
|
| |
| |
| |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| void model_kv_cache_alloc(Model *m) { |
| if (m->k_cache) return; |
| int n_layer = m->cfg.n_layer; |
| size_t per_layer = (size_t)m->cfg.n_ctx * m->cfg.n_embd * sizeof(float); |
| m->k_cache = calloc(n_layer, sizeof(float *)); |
| m->v_cache = calloc(n_layer, sizeof(float *)); |
| for (int l = 0; l < n_layer; l++) { |
| m->k_cache[l] = calloc(1, per_layer); |
| m->v_cache[l] = calloc(1, per_layer); |
| |
| if (m->layers) { |
| m->layers[l]._kv_k = m->k_cache[l]; |
| m->layers[l]._kv_v = m->v_cache[l]; |
| } |
| } |
| } |
|
|
| void model_kv_cache_free(Model *m) { |
| if (!m->k_cache) return; |
| for (int l = 0; l < m->cfg.n_layer; l++) { |
| free(m->k_cache[l]); |
| free(m->v_cache[l]); |
| } |
| free(m->k_cache); |
| free(m->v_cache); |
| m->k_cache = NULL; |
| m->v_cache = NULL; |
| } |
|
|
| |
| |
| |
| |
| |
| static TransAct *get_scratch_acts(Model *m) { |
| static TransAct *scratch = NULL; |
| static int scratch_n_embd = 0; |
| static int scratch_n_layer = 0; |
| if (!scratch || scratch_n_embd != m->cfg.n_embd || scratch_n_layer != m->cfg.n_layer) { |
| if (scratch) { |
| trans_act_free(scratch, scratch_n_layer); |
| scratch = NULL; |
| } |
| scratch = trans_act_alloc(&m->cfg); |
| scratch_n_embd = m->cfg.n_embd; |
| scratch_n_layer = m->cfg.n_layer; |
| } |
| return scratch; |
| } |
|
|
| void model_load(Model *m, const char *weight_path, ModelConfig cfg, |
| const char *layer_prefix, int qkv_merged) { |
| m->cfg = cfg; |
| m->cfg.qkv_merged = qkv_merged; |
| |
| |
| |
| |
| |
| |
| m->cfg.sliding_window = g_attn_window; |
| m->cfg.n_sinks = g_attn_sink; |
| m->tensors = tensor_load_all(weight_path, &m->n_tensors); |
| if (!m->tensors) { fprintf(stderr, "failed to load %s\n", weight_path); exit(1); } |
| printf("[*] loaded %d tensors\n", m->n_tensors); |
|
|
| m->wte = tensor_get(m->tensors, m->n_tensors, "wte.weight"); |
| m->wpe = (cfg.attn_type == ATTN_LEARNED) |
| ? tensor_get(m->tensors, m->n_tensors, "wpe.weight") : NULL; |
| m->ln_f_w = tensor_get(m->tensors, m->n_tensors, "ln_f.weight"); |
| m->ln_f_b = tensor_get(m->tensors, m->n_tensors, "ln_f.bias"); |
|
|
| printf("[*] binarizing %d layers%s...\n", cfg.n_layer, |
| g_use_logic_binarization ? " (logic-guided)" : ""); |
| m->layers = calloc(cfg.n_layer, sizeof(TransLayer)); |
| m->acts = trans_act_alloc(&cfg); |
|
|
| |
| char key[256], bk[256]; |
| for (int l = 0; l < cfg.n_layer; l++) { |
| TransLayer *tl = &m->layers[l]; |
| int n = cfg.n_embd, mm = cfg.mlp_dim; |
|
|
| |
| #define BIN_INIT(bl, W, b, in, out) do { \ |
| if (g_use_logic_binarization) { \ |
| uint8_t *mask = malloc(out); \ |
| compute_norm_mask(W, in, out, mask); \ |
| bin_layer_init_logic(bl, W, b, in, out, mask); \ |
| free(mask); \ |
| } else { \ |
| bin_layer_init(bl, W, b, in, out); \ |
| } \ |
| } while(0) |
| |
| |
| |
| |
| |
| #define BIN_INIT_NO_LOGIC(bl, W, b, in, out) do { \ |
| bin_layer_init(bl, W, b, in, out); \ |
| } while(0) |
|
|
| if (qkv_merged) { |
| sprintf(key, "h.%d.attn.c_attn.weight", l); |
| char bk[256]; strncpy(bk, key, sizeof(bk)); |
| char *dot = strstr(bk, ".weight"); if(dot){*dot=0;strcat(bk,".bias");} |
| BIN_INIT(&tl->attn_q, tensor_get(m->tensors, m->n_tensors, key), |
| tensor_get(m->tensors, m->n_tensors, bk), n, 3*n); |
| } else { |
| sprintf(key, "model.layers.%d.self_attn.q_proj.weight", l); |
| char bk[256]; strncpy(bk, key, sizeof(bk)); |
| char *dot = strstr(bk, ".weight"); if(dot){*dot=0;strcat(bk,".bias");} |
| BIN_INIT(&tl->attn_q, tensor_get(m->tensors, m->n_tensors, key), |
| tensor_get(m->tensors, m->n_tensors, bk), n, n); |
| sprintf(key, "model.layers.%d.self_attn.k_proj.weight", l); |
| strncpy(bk, key, sizeof(bk)); dot=strstr(bk,".weight"); if(dot){*dot=0;strcat(bk,".bias");} |
| BIN_INIT(&tl->attn_k, tensor_get(m->tensors, m->n_tensors, key), |
| tensor_get(m->tensors, m->n_tensors, bk), n, n); |
| sprintf(key, "model.layers.%d.self_attn.v_proj.weight", l); |
| strncpy(bk, key, sizeof(bk)); dot=strstr(bk,".weight"); if(dot){*dot=0;strcat(bk,".bias");} |
| BIN_INIT(&tl->attn_v, tensor_get(m->tensors, m->n_tensors, key), |
| tensor_get(m->tensors, m->n_tensors, bk), n, n); |
| } |
|
|
| sprintf(key, qkv_merged ? "h.%d.attn.c_proj.weight" : "model.layers.%d.self_attn.o_proj.weight", l); |
| char bk[256]; strncpy(bk, key, sizeof(bk)); |
| char *dot = strstr(bk, ".weight"); if(dot){*dot=0;strcat(bk,".bias");} |
| BIN_INIT(&tl->attn_o, tensor_get(m->tensors, m->n_tensors, key), |
| tensor_get(m->tensors, m->n_tensors, bk), n, n); |
|
|
| if (cfg.act_type == ACT_SWIGLU) { |
| sprintf(key, "model.layers.%d.mlp.gate_proj.weight", l); |
| strncpy(bk, key, sizeof(bk)); dot=strstr(bk,".weight"); if(dot){*dot=0;strcat(bk,".bias");} |
| BIN_INIT(&tl->mlp_gate, tensor_get(m->tensors, m->n_tensors, key), |
| tensor_get(m->tensors, m->n_tensors, bk), n, mm); |
| sprintf(key, "model.layers.%d.mlp.up_proj.weight", l); |
| strncpy(bk, key, sizeof(bk)); dot=strstr(bk,".weight"); if(dot){*dot=0;strcat(bk,".bias");} |
| BIN_INIT(&tl->mlp_up, tensor_get(m->tensors, m->n_tensors, key), |
| tensor_get(m->tensors, m->n_tensors, bk), n, mm); |
| } else { |
| sprintf(key, "h.%d.mlp.c_fc.weight", l); |
| strncpy(bk, key, sizeof(bk)); dot=strstr(bk,".weight"); if(dot){*dot=0;strcat(bk,".bias");} |
| BIN_INIT(&tl->mlp_gate, tensor_get(m->tensors, m->n_tensors, key), |
| tensor_get(m->tensors, m->n_tensors, bk), n, mm); |
| } |
|
|
| sprintf(key, qkv_merged ? "h.%d.mlp.c_proj.weight" : "model.layers.%d.mlp.down_proj.weight", l); |
| strncpy(bk, key, sizeof(bk)); dot=strstr(bk,".weight"); if(dot){*dot=0;strcat(bk,".bias");} |
| BIN_INIT(&tl->mlp_down, tensor_get(m->tensors, m->n_tensors, key), |
| tensor_get(m->tensors, m->n_tensors, bk), mm, n); |
| #undef BIN_INIT |
|
|
| |
| if (qkv_merged) { |
| sprintf(key, "h.%d.ln_1.weight", l); tl->norm1_w = tensor_get(m->tensors, m->n_tensors, key); |
| sprintf(key, "h.%d.ln_1.bias", l); tl->norm1_b = tensor_get(m->tensors, m->n_tensors, key); |
| sprintf(key, "h.%d.ln_2.weight", l); tl->norm2_w = tensor_get(m->tensors, m->n_tensors, key); |
| sprintf(key, "h.%d.ln_2.bias", l); tl->norm2_b = tensor_get(m->tensors, m->n_tensors, key); |
| } else { |
| sprintf(key, "model.layers.%d.input_layernorm.weight", l); tl->norm1_w = tensor_get(m->tensors, m->n_tensors, key); |
| tl->norm1_b = NULL; |
| sprintf(key, "model.layers.%d.post_attention_layernorm.weight", l); tl->norm2_w = tensor_get(m->tensors, m->n_tensors, key); |
| tl->norm2_b = NULL; |
| } |
| } |
| printf("[*] done\n"); |
|
|
| |
| |
| |
| { |
| int freed = 0; |
| size_t freed_bytes = 0; |
| for (int l = 0; l < cfg.n_layer; l++) { |
| char wk[256]; |
| const char *mats[] = { |
| qkv_merged ? "h.%d.attn.c_attn.weight" : "model.layers.%d.self_attn.q_proj.weight", |
| qkv_merged ? "h.%d.attn.c_proj.weight" : "model.layers.%d.self_attn.o_proj.weight", |
| qkv_merged ? "h.%d.mlp.c_fc.weight" : "model.layers.%d.mlp.gate_proj.weight", |
| qkv_merged ? "h.%d.mlp.c_proj.weight" : "model.layers.%d.mlp.down_proj.weight", |
| }; |
| for (int mi = 0; mi < 4; mi++) { |
| sprintf(wk, mats[mi], l); |
| for (int i = 0; i < m->n_tensors; i++) { |
| if (m->tensors[i].data && strcmp(m->tensors[i].key, wk) == 0) { |
| int n2 = 1; |
| for (int d = 0; d < m->tensors[i].ndim; d++) n2 *= m->tensors[i].shape[d]; |
| freed_bytes += (size_t)n2 * sizeof(float); |
| free(m->tensors[i].data); |
| m->tensors[i].data = NULL; |
| freed++; |
| break; |
| } |
| } |
| } |
| } |
| printf("[*] freed %d weight tensors (%.0f MB) after binarization\n", |
| freed, freed_bytes / 1e6); |
| } |
|
|
| m->final_ln = malloc(cfg.n_embd * sizeof(float)); |
| m->x_before_final = malloc(cfg.n_embd * sizeof(float)); |
| m->k_cache = NULL; |
| m->v_cache = NULL; |
| |
| |
| if (g_use_real_attention) model_kv_cache_alloc(m); |
| thr_res_alloc(m); |
|
|
| |
| |
| |
| |
| if (g_use_real_attention) { |
| ConceptAttnConfig cca = concept_attn_default_config(); |
| if (getenv("LAL_CONCEPT_ATTN") && atoi(getenv("LAL_CONCEPT_ATTN")) == 0) |
| cca.enable = 0; |
| model_set_concept_attn(m, &cca); |
| } |
| } |
|
|
|
|
|
|
| |
| |
| |
| |
| void model_forward_float_logits(Model *m, const int *tokens, int n_tokens, |
| float *logits_out) { |
| |
| |
| |
| model_forward_sliding(m, tokens, n_tokens); |
| int n = m->cfg.n_embd, vocab = m->cfg.vocab_size; |
| int tid = g_cur_tid; |
| float *ln = g_thr[tid].final_ln; |
| for (int j = 0; j < vocab; j++) { |
| const float *w = &m->wte[(size_t)j * n]; |
| float s = 0; |
| for (int i = 0; i + 7 < n; i += 8) |
| s += ln[i+0]*w[i+0] + ln[i+1]*w[i+1] + ln[i+2]*w[i+2] + ln[i+3]*w[i+3] |
| + ln[i+4]*w[i+4] + ln[i+5]*w[i+5] + ln[i+6]*w[i+6] + ln[i+7]*w[i+7]; |
| for (int i = (n/8)*8; i < n; i++) s += ln[i] * w[i]; |
| logits_out[j] = s; |
| } |
| } |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| void model_free(Model *m) { |
| for (int l = 0; l < m->cfg.n_layer; l++) |
| trans_layer_free(&m->layers[l], &m->cfg); |
| free(m->layers); |
| trans_act_free(m->acts, m->cfg.n_layer); |
| free(m->final_ln); |
| free(m->x_before_final); |
| model_kv_cache_free(m); |
| ponder_model_free(m); |
| thr_res_free(); |
| tensor_free_all(m->tensors, m->n_tensors); |
| } |
|
|
| |
| |
| |
| int ponder_step_count(const Model *m) { |
| |
| |
| |
| int block_steps = (g_ponder_cfg.rec_iters >= 2) ? g_ponder_cfg.rec_iters : 1; |
| return (g_ponder_cfg.layer_halt ? m->cfg.n_layer - 1 : 0) + block_steps; |
| } |
|
|
| void ponder_model_alloc(Model *m) { |
| if (!g_ponder_cfg.enable || m->ponder_ready) return; |
| int n = m->cfg.n_embd; |
| m->ph = calloc(m->cfg.n_layer, sizeof(PonderLayer)); |
| for (int l = 0; l < m->cfg.n_layer; l++) { |
| ponder_layer_alloc(&m->ph[l], n); |
| ponder_layer_init(&m->ph[l]); |
| } |
| ponder_layer_alloc(&m->ph_rec, n); |
| ponder_layer_init(&m->ph_rec); |
| |
| m->rec_acts = trans_act_alloc(&m->cfg); |
| m->n_rec_acts = g_ponder_cfg.rec_iters; |
| m->ponder_ready = 1; |
| printf("[PONDER] model alloc: %d layer units + 1 rec unit, n_embd=%d, steps=%d\n", |
| m->cfg.n_layer, n, ponder_step_count(m)); |
| } |
|
|
| void ponder_model_free(Model *m) { |
| if (!m || !m->ponder_ready) return; |
| for (int l = 0; l < m->cfg.n_layer; l++) |
| ponder_layer_free(&m->ph[l]); |
| free(m->ph); |
| ponder_layer_free(&m->ph_rec); |
| if (m->rec_acts) trans_act_free(m->rec_acts, m->cfg.n_layer); |
| m->rec_acts = NULL; |
| m->ponder_ready = 0; |
| } |
|
|
| void ponder_apply(Model *m, float lr, int batch_size, int opt_step) { |
| |
| |
| if (!g_ponder_cfg.enable || !m->ponder_ready) return; |
| float plr = lr * g_ponder_cfg.lr_scale; |
| float inv_batch = 1.0f / (float)batch_size; |
| float bc1 = 1.0f - powf(g_adam_beta1, (float)opt_step); |
| float bc2 = 1.0f - powf(g_adam_beta2, (float)opt_step); |
| if (bc1 < 1e-8f) bc1 = 1e-8f; |
| if (bc2 < 1e-8f) bc2 = 1e-8f; |
|
|
| for (int u = 0; u <= m->cfg.n_layer; u++) { |
| PonderLayer *pu = (u < m->cfg.n_layer) ? &m->ph[u] : &m->ph_rec; |
| |
| for (int i = 0; i < pu->in_dim; i++) pu->grad_w[i] *= inv_batch; |
| pu->grad_b *= inv_batch; |
| |
| float gnorm = pu->grad_b * pu->grad_b; |
| for (int i = 0; i < pu->in_dim; i++) gnorm += pu->grad_w[i] * pu->grad_w[i]; |
| gnorm = sqrtf(gnorm); |
| if (gnorm > 1.0f) { |
| float clip = 1.0f / gnorm; |
| for (int i = 0; i < pu->in_dim; i++) pu->grad_w[i] *= clip; |
| pu->grad_b *= clip; |
| } |
| |
| for (int i = 0; i < pu->in_dim; i++) { |
| float g = pu->grad_w[i]; |
| pu->m_w[i] = g_adam_beta1 * pu->m_w[i] + (1.0f - g_adam_beta1) * g; |
| pu->v_w[i] = g_adam_beta2 * pu->v_w[i] + (1.0f - g_adam_beta2) * g * g; |
| float mh = pu->m_w[i] / bc1; |
| float vh = pu->v_w[i] / bc2; |
| pu->w[i] -= plr * mh / (sqrtf(vh) + g_adam_eps); |
| pu->grad_w[i] = 0.0f; |
| } |
| pu->m_b = g_adam_beta1 * pu->m_b + (1.0f - g_adam_beta1) * pu->grad_b; |
| pu->v_b = g_adam_beta2 * pu->v_b + (1.0f - g_adam_beta2) * pu->grad_b * pu->grad_b; |
| pu->b -= plr * (pu->m_b / bc1) / (sqrtf(pu->v_b / bc2) + g_adam_eps); |
| pu->grad_b = 0.0f; |
| } |
| } |
|
|
| |
| |
| |
| void bin_layer_init(BinLayer *bl, const float *W, const float *bias, |
| int in_dim, int out_dim) { |
| bl->in_dim = in_dim; |
| bl->out_dim = out_dim; |
| bl->n_words = (in_dim + 63) / 64; |
| bl->n_words_T = (out_dim + 63) / 64; |
| bl->wbits = calloc(out_dim * bl->n_words, sizeof(uint64_t)); |
| bl->wbits_T = calloc(in_dim * bl->n_words_T, sizeof(uint64_t)); |
| bl->zbits = NULL; |
| bl->w_core = NULL; |
| bl->logic_mask = NULL; |
| bl->n_core = 0; |
| bl->n_prune = 0; |
| bl->alpha = calloc(out_dim, sizeof(float)); |
| bl->bias = bias ? malloc(out_dim * sizeof(float)) : calloc(out_dim, sizeof(float)); |
| bl->w_float = malloc((size_t)in_dim * out_dim * sizeof(float)); |
| bl->m_adam = g_use_adam ? calloc((size_t)in_dim * out_dim, sizeof(float)) : NULL; |
| bl->v_adam = g_use_adam ? calloc((size_t)in_dim * out_dim, sizeof(float)) : NULL; |
| bl->grad_accum = calloc((size_t)in_dim * out_dim, sizeof(float)); |
| bl->bias_grad_accum = calloc((size_t)out_dim, sizeof(float)); |
| bl->ternary_delta = 0.0f; |
|
|
| |
| |
| |
| |
| for (int j = 0; j < out_dim; j++) |
| for (int i = 0; i < in_dim; i++) |
| bl->w_float[j * in_dim + i] = W[i * out_dim + j]; |
|
|
| |
| for (int j = 0; j < out_dim; j++) { |
| float abs_sum = 0; |
| for (int i = 0; i < in_dim; i++) abs_sum += fabsf(W[i * out_dim + j]); |
| bl->alpha[j] = abs_sum / in_dim; |
| if (bias) bl->bias[j] = bias[j]; |
| for (int wi = 0; wi < bl->n_words; wi++) { |
| uint64_t word = 0; |
| for (int bi = 0; bi < 64; bi++) { |
| int idx = wi * 64 + bi; |
| if (idx < in_dim && W[idx * out_dim + j] > 0.0f) word |= (1ULL << bi); |
| } |
| bl->wbits[j * bl->n_words + wi] = word; |
| } |
| } |
|
|
| |
| for (int i = 0; i < in_dim; i++) { |
| for (int wi = 0; wi < bl->n_words_T; wi++) { |
| uint64_t word = 0; |
| for (int bi = 0; bi < 64; bi++) { |
| int j = wi * 64 + bi; |
| if (j < out_dim && W[i * out_dim + j] > 0.0f) word |= (1ULL << bi); |
| } |
| bl->wbits_T[i * bl->n_words_T + wi] = word; |
| } |
| } |
| } |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| void bin_layer_init_logic(BinLayer *bl, const float *W, const float *bias, |
| int in_dim, int out_dim, const uint8_t *logic_mask) { |
| bl->in_dim = in_dim; |
| bl->out_dim = out_dim; |
| bl->n_words = (in_dim + 63) / 64; |
| bl->n_words_T = (out_dim + 63) / 64; |
| bl->wbits = calloc(out_dim * bl->n_words, sizeof(uint64_t)); |
| bl->wbits_T = calloc(in_dim * bl->n_words_T, sizeof(uint64_t)); |
| bl->alpha = calloc(out_dim, sizeof(float)); |
| bl->bias = bias ? malloc(out_dim * sizeof(float)) : calloc(out_dim, sizeof(float)); |
| bl->w_float = malloc((size_t)out_dim * in_dim * sizeof(float)); |
| bl->m_adam = NULL; |
| bl->v_adam = NULL; |
| bl->w_core = NULL; |
| bl->logic_mask = NULL; |
| bl->n_core = 0; |
| bl->n_prune = 0; |
| bl->zbits = NULL; |
| bl->ternary_delta = 0.0f; |
| bl->grad_accum = NULL; |
| bl->bias_grad_accum = NULL; |
|
|
| if (!logic_mask) { |
| |
| free(bl->w_float); bl->w_float = NULL; |
| bin_layer_init(bl, W, bias, in_dim, out_dim); |
| return; |
| } |
| |
| |
| |
| bl->m_adam = calloc((size_t)out_dim * in_dim, sizeof(float)); |
| bl->v_adam = calloc((size_t)out_dim * in_dim, sizeof(float)); |
|
|
| |
| bl->logic_mask = malloc(out_dim); |
| memcpy(bl->logic_mask, logic_mask, out_dim); |
| for (int j = 0; j < out_dim; j++) { |
| if (logic_mask[j] == 0) bl->n_core++; |
| else if (logic_mask[j] == 2) bl->n_prune++; |
| } |
|
|
| |
| if (bl->n_core > 0) { |
| bl->w_core = malloc((size_t)bl->n_core * in_dim * sizeof(float)); |
| } |
|
|
| |
| int core_idx = 0; |
| for (int j = 0; j < out_dim; j++) { |
| const float *wj = &W[j * in_dim]; |
|
|
| switch (logic_mask[j]) { |
| case 0: |
| memcpy(&bl->w_core[core_idx * in_dim], wj, in_dim * sizeof(float)); |
| bl->alpha[j] = 0.0f; |
| if (bias) bl->bias[j] = bias[j]; |
| |
| core_idx++; |
| break; |
|
|
| case 1: |
| { |
| float abs_sum = 0; |
| for (int i = 0; i < in_dim; i++) abs_sum += fabsf(wj[i]); |
| bl->alpha[j] = abs_sum / in_dim; |
| if (bias) bl->bias[j] = bias[j]; |
| for (int wi = 0; wi < bl->n_words; wi++) { |
| uint64_t word = 0; |
| for (int bi = 0; bi < 64; bi++) { |
| int idx = wi * 64 + bi; |
| if (idx < in_dim && wj[idx] > 0.0f) word |= (1ULL << bi); |
| } |
| bl->wbits[j * bl->n_words + wi] = word; |
| } |
| } |
| break; |
|
|
| case 2: |
| bl->alpha[j] = 0.0f; |
| bl->bias[j] = 0.0f; |
| |
| break; |
| } |
|
|
| |
| memcpy(&bl->w_float[j * in_dim], wj, in_dim * sizeof(float)); |
| } |
|
|
| |
| for (int i = 0; i < in_dim; i++) { |
| for (int wi = 0; wi < bl->n_words_T; wi++) { |
| uint64_t word = 0; |
| for (int bi = 0; bi < 64; bi++) { |
| int j = wi * 64 + bi; |
| if (j < out_dim && logic_mask[j] == 1 && W[j * in_dim + i] > 0.0f) |
| word |= (1ULL << bi); |
| } |
| bl->wbits_T[i * bl->n_words_T + wi] = word; |
| } |
| } |
|
|
| |
| |
| |
| if (g_use_ternary) { |
| bl->zbits = calloc((size_t)out_dim * bl->n_words, sizeof(uint64_t)); |
| bin_layer_repack_ternary(bl); |
| } |
| } |
|
|
| void bin_layer_free(BinLayer *bl) { |
| free(bl->wbits); free(bl->wbits_T); free(bl->zbits); free(bl->alpha); free(bl->bias); |
| free(bl->w_float); free(bl->w_core); free(bl->logic_mask); |
| free(bl->m_adam); free(bl->v_adam); free(bl->grad_accum); free(bl->bias_grad_accum); |
| bl->wbits = NULL; bl->wbits_T = NULL; bl->zbits = NULL; |
| bl->alpha = NULL; bl->bias = NULL; |
| bl->w_float = NULL; bl->w_core = NULL; bl->logic_mask = NULL; |
| bl->m_adam = NULL; bl->v_adam = NULL; |
| bl->grad_accum = NULL; bl->bias_grad_accum = NULL; |
| } |
|
|
| |
| |
| |
| |
| |
| |
| |
| void bin_layer_repack(BinLayer *bl) { |
| int in = bl->in_dim, out = bl->out_dim; |
|
|
| |
| |
| |
| |
| if (bl->w_core && bl->logic_mask) { |
| int core_idx = 0; |
| for (int j = 0; j < out; j++) { |
| if (bl->logic_mask[j] == 0) { |
| memcpy(&bl->w_core[core_idx * in], |
| &bl->w_float[j * in], |
| in * sizeof(float)); |
| core_idx++; |
| } |
| } |
| } |
|
|
| |
| for (int j = 0; j < out; j++) { |
| const float *wf = &bl->w_float[j * in]; |
| for (int wi = 0; wi < bl->n_words; wi++) { |
| uint64_t word = 0; |
| int base = wi * 64; |
| for (int grp = 0; grp < 8; grp++) { |
| int idx = base + grp * 8; |
| if (idx + 7 < in) { |
| |
| if (wf[idx+0] > 0.0f) word |= (1ULL << (grp*8 + 0)); |
| if (wf[idx+1] > 0.0f) word |= (1ULL << (grp*8 + 1)); |
| if (wf[idx+2] > 0.0f) word |= (1ULL << (grp*8 + 2)); |
| if (wf[idx+3] > 0.0f) word |= (1ULL << (grp*8 + 3)); |
| if (wf[idx+4] > 0.0f) word |= (1ULL << (grp*8 + 4)); |
| if (wf[idx+5] > 0.0f) word |= (1ULL << (grp*8 + 5)); |
| if (wf[idx+6] > 0.0f) word |= (1ULL << (grp*8 + 6)); |
| if (wf[idx+7] > 0.0f) word |= (1ULL << (grp*8 + 7)); |
| } else { |
| for (int bi = 0; bi < 8; bi++) { |
| int i = idx + bi; |
| if (i < in && wf[i] > 0.0f) |
| word |= (1ULL << (grp * 8 + bi)); |
| } |
| } |
| } |
| bl->wbits[j * bl->n_words + wi] = word; |
| } |
| } |
|
|
| |
| |
| |
| for (int j = 0; j < out; j++) { |
| const float *wf = &bl->w_float[j * in]; |
| float abs_sum = 0; |
| for (int i = 0; i + 7 < in; i += 8) { |
| abs_sum += fabsf(wf[i+0]) + fabsf(wf[i+1]) + fabsf(wf[i+2]) + fabsf(wf[i+3]); |
| abs_sum += fabsf(wf[i+4]) + fabsf(wf[i+5]) + fabsf(wf[i+6]) + fabsf(wf[i+7]); |
| } |
| for (int i = (in / 8) * 8; i < in; i++) |
| abs_sum += fabsf(wf[i]); |
| bl->alpha[j] = abs_sum / in; |
| } |
| } |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| void bin_layer_repack_ternary(BinLayer *bl) { |
| if (!bl->zbits || !bl->w_float) return; |
| int in = bl->in_dim, out = bl->out_dim, nw = bl->n_words; |
|
|
| for (int j = 0; j < out; j++) { |
| |
| if (bl->logic_mask && bl->logic_mask[j] != 1) continue; |
|
|
| const float *wf = &bl->w_float[j * in]; |
| |
| float abs_sum = 0.0f; |
| for (int i = 0; i + 7 < in; i += 8) { |
| abs_sum += fabsf(wf[i+0]) + fabsf(wf[i+1]) + fabsf(wf[i+2]) + fabsf(wf[i+3]); |
| abs_sum += fabsf(wf[i+4]) + fabsf(wf[i+5]) + fabsf(wf[i+6]) + fabsf(wf[i+7]); |
| } |
| for (int i = (in/8)*8; i < in; i++) abs_sum += fabsf(wf[i]); |
| float mean_abs = abs_sum / in; |
| float delta = g_ternary_delta_factor * mean_abs; |
| bl->ternary_delta = delta; |
|
|
| |
| uint64_t *zb = &bl->zbits[(size_t)j * nw]; |
| float active_abs_sum = 0.0f; |
| int n_active = 0; |
| for (int wi = 0; wi < nw; wi++) { |
| uint64_t word = 0; |
| int base = wi * 64; |
| for (int grp = 0; grp < 8; grp++) { |
| int idx = base + grp * 8; |
| if (idx + 7 < in) { |
| for (int k = 0; k < 8; k++) { |
| int i = idx + k; |
| if (fabsf(wf[i]) <= delta) { |
| word |= (1ULL << (grp*8 + k)); |
| } else { |
| active_abs_sum += fabsf(wf[i]); |
| n_active++; |
| } |
| } |
| } else { |
| for (int k = 0; k < 8; k++) { |
| int i = idx + k; |
| if (i >= in) break; |
| if (fabsf(wf[i]) <= delta) { |
| word |= (1ULL << (grp*8 + k)); |
| } else { |
| active_abs_sum += fabsf(wf[i]); |
| n_active++; |
| } |
| } |
| } |
| } |
| zb[wi] = word; |
| } |
| |
| |
| bl->alpha[j] = (n_active > 0) ? (active_abs_sum / n_active) : mean_abs; |
| } |
| } |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| void bin_forward(float *y, const float *x, const BinLayer *bl) { |
| int in = bl->in_dim, out = bl->out_dim, nw = bl->n_words; |
|
|
| |
| if (bl->logic_mask) { |
| |
| float abs_sum = 0.0f; |
| for (int i = 0; i < in; i++) abs_sum += fabsf(x[i]); |
| float K = abs_sum / in; |
|
|
| sign_lut_ensure(); |
|
|
| |
| int *cidx = (int *)alloca(out * sizeof(int)); |
| { int c = 0; for (int j = 0; j < out; j++) { cidx[j] = c; if (bl->logic_mask[j] == 0) c++; } } |
|
|
| #ifdef HAVE_OPENBLAS |
| |
| |
| |
| |
| |
| |
| int n_core = cidx[out > 0 ? out - 1 : 0] + (out > 0 && bl->logic_mask[out-1] == 0 ? 1 : 0); |
| if (n_core > 0 && bl->w_core) { |
| openblas_init_single_thread(); |
| |
| static __thread float *core_dots = NULL; |
| static __thread int core_dots_n = 0; |
| if (core_dots_n < n_core) { |
| free(core_dots); |
| core_dots = (float *)malloc(n_core * sizeof(float)); |
| core_dots_n = n_core; |
| } |
| |
| |
| |
| cblas_sgemv(CblasRowMajor, CblasNoTrans, |
| n_core, in, |
| 1.0f, bl->w_core, in, |
| x, 1, |
| 0.0f, core_dots, 1); |
| |
| #pragma omp parallel for schedule(static) |
| for (int j = 0; j < out; j++) { |
| if (bl->logic_mask[j] != 0) continue; |
| float core_gain = 1.0f / (bl->alpha[j] + 1e-8f); |
| if (core_gain > 5.0f) core_gain = 5.0f; |
| y[j] = core_dots[cidx[j]] * core_gain * K + bl->bias[j]; |
| } |
| |
| #pragma omp parallel for schedule(static) |
| for (int j = 0; j < out; j++) { |
| uint8_t m = bl->logic_mask[j]; |
| if (m == 0) continue; |
| if (m == 1) { |
| const uint64_t *wb = &bl->wbits[j * nw]; |
| const uint64_t *zb = bl->zbits ? &bl->zbits[j * nw] : NULL; |
| float s = 0.0f; |
| for (int wi = 0; wi < nw; wi++) { |
| uint64_t w = wb[wi]; |
| uint64_t z = zb ? zb[wi] : 0; |
| int base = wi * 64; |
| for (int bi = 0; bi < 8; bi++) { |
| int idx = base + bi * 8; |
| uint8_t byte = (uint8_t)((w >> (bi * 8)) & 0xFF); |
| uint8_t zbyte = (uint8_t)((z >> (bi * 8)) & 0xFF); |
| const float *sw = g_sign_lut[byte]; |
| if (idx + 7 < in) { |
| if (zbyte == 0) { |
| s += x[idx+0]*sw[0] + x[idx+1]*sw[1] + x[idx+2]*sw[2] + x[idx+3]*sw[3]; |
| s += x[idx+4]*sw[4] + x[idx+5]*sw[5] + x[idx+6]*sw[6] + x[idx+7]*sw[7]; |
| } else { |
| s += (zbyte & 0x01) ? 0 : x[idx+0]*sw[0]; |
| s += (zbyte & 0x02) ? 0 : x[idx+1]*sw[1]; |
| s += (zbyte & 0x04) ? 0 : x[idx+2]*sw[2]; |
| s += (zbyte & 0x08) ? 0 : x[idx+3]*sw[3]; |
| s += (zbyte & 0x10) ? 0 : x[idx+4]*sw[4]; |
| s += (zbyte & 0x20) ? 0 : x[idx+5]*sw[5]; |
| s += (zbyte & 0x40) ? 0 : x[idx+6]*sw[6]; |
| s += (zbyte & 0x80) ? 0 : x[idx+7]*sw[7]; |
| } |
| } else { |
| for (int i = idx; i < in; i++) { |
| if (zb && ((z >> (i & 63)) & 1)) continue; |
| s += (((w >> (i & 63)) & 1) ? sw[i-idx] : -sw[i-idx]); |
| } |
| } |
| } |
| } |
| y[j] = s * bl->alpha[j] * K * g_binary_scale + bl->bias[j]; |
| } else { |
| y[j] = 0.0f; |
| } |
| } |
| return; |
| } |
| #endif |
| |
| #pragma omp parallel for schedule(static) |
| for (int j = 0; j < out; j++) { |
| switch (bl->logic_mask[j]) { |
| case 0: { |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| const float *wc = &bl->w_core[cidx[j] * in]; |
| float s = 0.0f; |
| for (int i = 0; i + 7 < in; i += 8) { |
| s += x[i+0]*wc[i+0] + x[i+1]*wc[i+1] + x[i+2]*wc[i+2] + x[i+3]*wc[i+3]; |
| s += x[i+4]*wc[i+4] + x[i+5]*wc[i+5] + x[i+6]*wc[i+6] + x[i+7]*wc[i+7]; |
| } |
| for (int i = (in/8)*8; i < in; i++) s += x[i] * wc[i]; |
| float core_gain = 1.0f / (bl->alpha[j] + 1e-8f); |
| if (core_gain > 5.0f) core_gain = 5.0f; |
| y[j] = s * core_gain * K + bl->bias[j]; |
| break; |
| } |
| case 1: { |
| const uint64_t *wb = &bl->wbits[j * nw]; |
| const uint64_t *zb = bl->zbits ? &bl->zbits[j * nw] : NULL; |
| float s = 0.0f; |
| for (int wi = 0; wi < nw; wi++) { |
| uint64_t w = wb[wi]; |
| uint64_t z = zb ? zb[wi] : 0; |
| int base = wi * 64; |
| for (int bi = 0; bi < 8; bi++) { |
| int idx = base + bi * 8; |
| uint8_t byte = (uint8_t)((w >> (bi * 8)) & 0xFF); |
| uint8_t zbyte = (uint8_t)((z >> (bi * 8)) & 0xFF); |
| const float *sw = g_sign_lut[byte]; |
| if (idx + 7 < in) { |
| |
| |
| |
| if (zbyte == 0) { |
| |
| s += x[idx+0]*sw[0] + x[idx+1]*sw[1] + x[idx+2]*sw[2] + x[idx+3]*sw[3]; |
| s += x[idx+4]*sw[4] + x[idx+5]*sw[5] + x[idx+6]*sw[6] + x[idx+7]*sw[7]; |
| } else { |
| |
| s += (zbyte & 0x01) ? 0 : x[idx+0]*sw[0]; |
| s += (zbyte & 0x02) ? 0 : x[idx+1]*sw[1]; |
| s += (zbyte & 0x04) ? 0 : x[idx+2]*sw[2]; |
| s += (zbyte & 0x08) ? 0 : x[idx+3]*sw[3]; |
| s += (zbyte & 0x10) ? 0 : x[idx+4]*sw[4]; |
| s += (zbyte & 0x20) ? 0 : x[idx+5]*sw[5]; |
| s += (zbyte & 0x40) ? 0 : x[idx+6]*sw[6]; |
| s += (zbyte & 0x80) ? 0 : x[idx+7]*sw[7]; |
| } |
| } else { |
| for (int k = 0; k < 8; k++) { |
| int i = idx + k; |
| if (i < in && !((zbyte >> k) & 1)) s += x[i] * sw[k]; |
| } |
| } |
| } |
| } |
| y[j] = s * bl->alpha[j] * K * g_binary_scale + bl->bias[j]; |
| break; |
| } |
| default: |
| y[j] = 0.0f; |
| break; |
| } |
| } |
| return; |
| } |
|
|
| |
| float abs_sum = 0.0f; |
| for (int i = 0; i < in; i++) abs_sum += fabsf(x[i]); |
| float K = abs_sum / in; |
|
|
| sign_lut_ensure(); |
|
|
| for (int j = 0; j < out; j++) { |
| const uint64_t *wb = &bl->wbits[j * nw]; |
| float s = 0.0f; |
| for (int wi = 0; wi < nw; wi++) { |
| uint64_t w = wb[wi]; |
| int base = wi * 64; |
| |
| for (int bi = 0; bi < 8; bi++) { |
| int idx = base + bi * 8; |
| uint8_t byte = (uint8_t)((w >> (bi * 8)) & 0xFF); |
| const float *sw = g_sign_lut[byte]; |
| if (idx + 7 < in) { |
| |
| s += x[idx+0] * sw[0]; |
| s += x[idx+1] * sw[1]; |
| s += x[idx+2] * sw[2]; |
| s += x[idx+3] * sw[3]; |
| s += x[idx+4] * sw[4]; |
| s += x[idx+5] * sw[5]; |
| s += x[idx+6] * sw[6]; |
| s += x[idx+7] * sw[7]; |
| } else { |
| |
| for (int k = 0; k < 8; k++) { |
| int i = idx + k; |
| if (i < in) s += x[i] * sw[k]; |
| } |
| } |
| } |
| } |
| y[j] = s * bl->alpha[j] * K + bl->bias[j]; |
| } |
| } |
|
|
| |
| |
| |
| |
| void bin_forward_bnn(float *y, const float *x, const BinLayer *bl) { |
| int in = bl->in_dim, out = bl->out_dim, nw = bl->n_words; |
|
|
| |
| |
| float abs_sum = 0.0f; |
| for (int i = 0; i < in; i++) abs_sum += fabsf(x[i]); |
| float K = abs_sum / in; |
|
|
| |
| uint64_t xbits[64]; |
| for (int wi = 0; wi < nw; wi++) { |
| uint64_t word = 0; |
| for (int bi = 0; bi < 64; bi++) { |
| int idx = wi * 64 + bi; |
| if (idx < in && x[idx] > 0.0f) word |= (1ULL << bi); |
| } |
| xbits[wi] = word; |
| } |
| |
| for (int j = 0; j < out; j++) { |
| int pc = 0; |
| const uint64_t *wb = &bl->wbits[j * nw]; |
| for (int wi = 0; wi < nw; wi++) |
| pc += __builtin_popcountll(~(xbits[wi] ^ wb[wi])); |
| y[j] = (float)(2 * pc - in) * bl->alpha[j] * K + bl->bias[j]; |
| } |
| } |
|
|
| |
| |
| void bin_forward_float(float *y, const float *x, const BinLayer *bl) { |
| int in = bl->in_dim, out = bl->out_dim, nw = bl->n_words; |
| for (int j = 0; j < out; j++) { |
| float s = bl->bias[j]; |
| const uint64_t *wb = &bl->wbits[j * nw]; |
| float a = bl->alpha[j]; |
| for (int wi = 0; wi < nw; wi++) { |
| uint64_t w = wb[wi]; |
| for (int bi = 0; bi < 64; bi++) { |
| int idx = wi * 64 + bi; |
| if (idx >= in) break; |
| s += x[idx] * ((w >> bi) & 1 ? 1.0f : -1.0f) * a; |
| } |
| } |
| y[j] = s; |
| } |
| } |
|
|
| |
| |
| |
| void bin_backward(float *grad_x, const float *grad_y, const float *x, |
| BinLayer *bl, float lr) { |
| int in = bl->in_dim, out = bl->out_dim; |
| int nw_T = bl->n_words_T; |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| if (bl->logic_mask) { |
| for (int i = 0; i < in; i++) grad_x[i] = 0.0f; |
| |
| |
| float abs_sum = 0.0f; |
| for (int i = 0; i < in; i++) abs_sum += fabsf(x[i]); |
| float K = abs_sum / in; |
| int core_idx = 0; |
| for (int j = 0; j < out; j++) { |
| float gy = grad_y[j]; |
| if (bl->logic_mask[j] == 0) { |
| |
| if (fabsf(gy) >= 1e-8f) { |
| const float *wc = &bl->w_core[core_idx * in]; |
| float core_gain = 1.0f / (bl->alpha[j] + 1e-8f); |
| if (core_gain > 5.0f) core_gain = 5.0f; |
| float scale = gy * core_gain * K; |
| for (int i = 0; i + 7 < in; i += 8) { |
| grad_x[i+0] += scale * wc[i+0]; |
| grad_x[i+1] += scale * wc[i+1]; |
| grad_x[i+2] += scale * wc[i+2]; |
| grad_x[i+3] += scale * wc[i+3]; |
| grad_x[i+4] += scale * wc[i+4]; |
| grad_x[i+5] += scale * wc[i+5]; |
| grad_x[i+6] += scale * wc[i+6]; |
| grad_x[i+7] += scale * wc[i+7]; |
| } |
| for (int i = (in/8)*8; i < in; i++) grad_x[i] += scale * wc[i]; |
| } |
| core_idx++; |
| bl->bias[j] -= lr * gy; |
| } else if (bl->logic_mask[j] == 1) { |
| |
| if (fabsf(gy) >= 1e-8f) { |
| const uint64_t *wb = &bl->wbits[j * bl->n_words]; |
| float scale = gy * bl->alpha[j]; |
| for (int wi = 0; wi < bl->n_words; wi++) { |
| uint64_t w = wb[wi]; |
| int base = wi * 64; |
| for (int bi = 0; bi < 8; bi++) { |
| int idx = base + bi * 8; |
| if (idx + 7 < in) { |
| grad_x[idx+0] += scale * ((w >> (bi*8+0)) & 1 ? 1.0f : -1.0f); |
| grad_x[idx+1] += scale * ((w >> (bi*8+1)) & 1 ? 1.0f : -1.0f); |
| grad_x[idx+2] += scale * ((w >> (bi*8+2)) & 1 ? 1.0f : -1.0f); |
| grad_x[idx+3] += scale * ((w >> (bi*8+3)) & 1 ? 1.0f : -1.0f); |
| grad_x[idx+4] += scale * ((w >> (bi*8+4)) & 1 ? 1.0f : -1.0f); |
| grad_x[idx+5] += scale * ((w >> (bi*8+5)) & 1 ? 1.0f : -1.0f); |
| grad_x[idx+6] += scale * ((w >> (bi*8+6)) & 1 ? 1.0f : -1.0f); |
| grad_x[idx+7] += scale * ((w >> (bi*8+7)) & 1 ? 1.0f : -1.0f); |
| } else { |
| for (int k = 0; k < 8; k++) { |
| int i = idx + k; |
| if (i < in) grad_x[i] += scale * ((w >> (bi*8+k)) & 1 ? 1.0f : -1.0f); |
| } |
| } |
| } |
| } |
| } |
| bl->bias[j] -= lr * gy; |
| } |
| |
| } |
| return; |
| } |
|
|
| |
|
|
| |
| uint64_t gybits[64]; |
| for (int wi = 0; wi < nw_T; wi++) { |
| uint64_t word = 0; |
| for (int bi = 0; bi < 64; bi++) { |
| int j = wi * 64 + bi; |
| if (j < out && grad_y[j] > 0.0f) word |= (1ULL << bi); |
| } |
| gybits[wi] = word; |
| } |
| float mean_abs_gy = 0, mean_alpha = 0; |
| for (int j = 0; j < out; j++) mean_abs_gy += fabsf(grad_y[j]); |
| mean_abs_gy /= out; |
| for (int j = 0; j < out; j++) mean_alpha += bl->alpha[j]; |
| mean_alpha /= out; |
| for (int i = 0; i < in; i++) { |
| int pc = 0; |
| const uint64_t *wbT = &bl->wbits_T[i * nw_T]; |
| for (int wi = 0; wi < nw_T; wi++) |
| pc += __builtin_popcountll(~(gybits[wi] ^ wbT[wi])); |
| grad_x[i] = (float)(2 * pc - out) * mean_alpha * mean_abs_gy; |
| } |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| float mean_abs_x = 0; |
| for (int i = 0; i < in; i++) mean_abs_x += fabsf(x[i]); |
| mean_abs_x /= in; |
| uint64_t xbits[64]; |
| for (int wi = 0; wi < bl->n_words; wi++) { |
| uint64_t word = 0; |
| for (int bi = 0; bi < 64; bi++) { |
| int idx = wi * 64 + bi; |
| if (idx < in && x[idx] > 0.0f) word |= (1ULL << bi); |
| } |
| xbits[wi] = word; |
| } |
| for (int j = 0; j < out; j++) { |
| float gy = grad_y[j]; |
| if (fabsf(gy) < 1e-6f) continue; |
| int pc = 0; |
| const uint64_t *wb = &bl->wbits[j * bl->n_words]; |
| for (int wi = 0; wi < bl->n_words; wi++) |
| pc += __builtin_popcountll(~(xbits[wi] ^ wb[wi])); |
| float grad_alpha = (float)(2 * pc - in) * mean_abs_x; |
| bl->alpha[j] -= lr * grad_alpha * gy; |
| |
| |
| if (bl->alpha[j] < 0.0f) bl->alpha[j] = 0.0f; |
| bl->bias[j] -= lr * gy; |
| } |
| } |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| void bin_backward_ste(float *grad_x, const float *grad_y, const float *x, |
| BinLayer *bl, float lr, int layer_idx, int bl_slot) { |
| int in = bl->in_dim, out = bl->out_dim; |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| if (bl->w_float) { |
| |
| for (int i = 0; i < in; i++) grad_x[i] = 0.0f; |
| |
| |
| |
| float K = 1.0f; |
| int nw_gx = (in + 63) / 64; |
| sign_lut_ensure(); |
| if (g_use_ternary && bl->zbits) { |
| K = 0.0f; |
| for (int i = 0; i < in; i++) K += fabsf(x[i]); |
| K = K / in; |
| } |
| |
| for (int j = 0; j < out; j++) { |
| float gy = grad_y[j]; |
| if (fabsf(gy) < 1e-8f) continue; |
| |
| |
| |
| |
| if (bl->logic_mask && bl->logic_mask[j] == 2) continue; |
| const float *wf = &bl->w_float[j * in]; |
| if (g_use_ternary && bl->zbits) { |
| |
| |
| |
| |
| |
| float scale = gy * bl->alpha[j] * K; |
| const uint64_t *zb = &bl->zbits[(size_t)j * nw_gx]; |
| const uint64_t *wb = &bl->wbits[(size_t)j * nw_gx]; |
| for (int i = 0; i + 7 < in; i += 8) { |
| uint8_t zbyte = (uint8_t)((zb[i/64] >> (i%64)) & 0xFF); |
| uint8_t wbyte = (uint8_t)((wb[i/64] >> (i%64)) & 0xFF); |
| const float *sw = g_sign_lut[wbyte]; |
| if (zbyte == 0) { |
| grad_x[i+0] += scale * sw[0]; grad_x[i+1] += scale * sw[1]; |
| grad_x[i+2] += scale * sw[2]; grad_x[i+3] += scale * sw[3]; |
| grad_x[i+4] += scale * sw[4]; grad_x[i+5] += scale * sw[5]; |
| grad_x[i+6] += scale * sw[6]; grad_x[i+7] += scale * sw[7]; |
| } else { |
| grad_x[i+0] += scale * ((zbyte & 0x01) ? 0.0f : sw[0]); |
| grad_x[i+1] += scale * ((zbyte & 0x02) ? 0.0f : sw[1]); |
| grad_x[i+2] += scale * ((zbyte & 0x04) ? 0.0f : sw[2]); |
| grad_x[i+3] += scale * ((zbyte & 0x08) ? 0.0f : sw[3]); |
| grad_x[i+4] += scale * ((zbyte & 0x10) ? 0.0f : sw[4]); |
| grad_x[i+5] += scale * ((zbyte & 0x20) ? 0.0f : sw[5]); |
| grad_x[i+6] += scale * ((zbyte & 0x40) ? 0.0f : sw[6]); |
| grad_x[i+7] += scale * ((zbyte & 0x80) ? 0.0f : sw[7]); |
| } |
| } |
| for (int i = (in/8)*8; i < in; i++) { |
| uint64_t z = (zb[i/64] >> (i%64)) & 1; |
| uint64_t wbit = (wb[i/64] >> (i%64)) & 1; |
| grad_x[i] += scale * (z ? 0.0f : (wbit ? 1.0f : -1.0f)); |
| } |
| } else if (g_use_pure_float) { |
| |
| float scale = gy; |
| for (int i = 0; i + 7 < in; i += 8) { |
| grad_x[i+0] += scale * wf[i+0]; |
| grad_x[i+1] += scale * wf[i+1]; |
| grad_x[i+2] += scale * wf[i+2]; |
| grad_x[i+3] += scale * wf[i+3]; |
| grad_x[i+4] += scale * wf[i+4]; |
| grad_x[i+5] += scale * wf[i+5]; |
| grad_x[i+6] += scale * wf[i+6]; |
| grad_x[i+7] += scale * wf[i+7]; |
| } |
| for (int i = (in / 8) * 8; i < in; i++) |
| grad_x[i] += scale * wf[i]; |
| } else { |
| |
| float scale = gy * bl->alpha[j]; |
| for (int i = 0; i + 7 < in; i += 8) { |
| grad_x[i+0] += scale * (wf[i+0] > 0.0f ? 1.0f : -1.0f); |
| grad_x[i+1] += scale * (wf[i+1] > 0.0f ? 1.0f : -1.0f); |
| grad_x[i+2] += scale * (wf[i+2] > 0.0f ? 1.0f : -1.0f); |
| grad_x[i+3] += scale * (wf[i+3] > 0.0f ? 1.0f : -1.0f); |
| grad_x[i+4] += scale * (wf[i+4] > 0.0f ? 1.0f : -1.0f); |
| grad_x[i+5] += scale * (wf[i+5] > 0.0f ? 1.0f : -1.0f); |
| grad_x[i+6] += scale * (wf[i+6] > 0.0f ? 1.0f : -1.0f); |
| grad_x[i+7] += scale * (wf[i+7] > 0.0f ? 1.0f : -1.0f); |
| } |
| for (int i = (in / 8) * 8; i < in; i++) |
| grad_x[i] += scale * (wf[i] > 0.0f ? 1.0f : -1.0f); |
| } |
| } |
| } else { |
| |
| int nw_T = bl->n_words_T; |
| uint64_t gybits[64]; |
| for (int wi = 0; wi < nw_T; wi++) { |
| uint64_t word = 0; |
| for (int bi = 0; bi < 64; bi++) { |
| int j = wi * 64 + bi; |
| if (j < out && grad_y[j] > 0.0f) word |= (1ULL << bi); |
| } |
| gybits[wi] = word; |
| } |
| float mean_abs_gy = 0, mean_alpha = 0; |
| for (int j = 0; j < out; j++) mean_abs_gy += fabsf(grad_y[j]); |
| mean_abs_gy /= out; |
| for (int j = 0; j < out; j++) mean_alpha += bl->alpha[j]; |
| mean_alpha /= out; |
| for (int i = 0; i < in; i++) { |
| int pc = 0; |
| const uint64_t *wbT = &bl->wbits_T[i * nw_T]; |
| for (int wi = 0; wi < nw_T; wi++) |
| pc += __builtin_popcountll(~(gybits[wi] ^ wbT[wi])); |
| grad_x[i] = (float)(2 * pc - out) * mean_alpha * mean_abs_gy; |
| } |
| } |
|
|
| |
| |
| |
| |
| |
| if (g_accumulate_gradients && bl->grad_accum) { |
| float *ga = (bl_slot >= 0) ? g_thr[g_cur_tid].grad_w[layer_idx][bl_slot] : bl->grad_accum; |
| float *gba = (bl_slot >= 0) ? g_thr[g_cur_tid].grad_b[layer_idx][bl_slot] : bl->bias_grad_accum; |
| int nw_gw = (in + 63) / 64; |
| for (int j = 0; j < out; j++) { |
| float gy = grad_y[j]; |
| if (fabsf(gy) < 1e-8f) continue; |
| if (bl->logic_mask && bl->logic_mask[j] == 2) continue; |
| float *ga_j = &ga[j * in]; |
| if (g_use_ternary && bl->zbits) { |
| |
| |
| |
| const uint64_t *zb = &bl->zbits[(size_t)j * nw_gw]; |
| const uint64_t *wb = &bl->wbits[(size_t)j * nw_gw]; |
| for (int i = 0; i + 7 < in; i += 8) { |
| uint8_t zbyte = (uint8_t)((zb[i/64] >> (i%64)) & 0xFF); |
| uint8_t wbyte = (uint8_t)((wb[i/64] >> (i%64)) & 0xFF); |
| const float *sw = g_sign_lut[wbyte]; |
| if (zbyte == 0) { |
| ga_j[i+0] += gy * x[i+0] * sw[0]; ga_j[i+1] += gy * x[i+1] * sw[1]; |
| ga_j[i+2] += gy * x[i+2] * sw[2]; ga_j[i+3] += gy * x[i+3] * sw[3]; |
| ga_j[i+4] += gy * x[i+4] * sw[4]; ga_j[i+5] += gy * x[i+5] * sw[5]; |
| ga_j[i+6] += gy * x[i+6] * sw[6]; ga_j[i+7] += gy * x[i+7] * sw[7]; |
| } else { |
| ga_j[i+0] += gy * x[i+0] * ((zbyte & 0x01) ? 0.0f : sw[0]); |
| ga_j[i+1] += gy * x[i+1] * ((zbyte & 0x02) ? 0.0f : sw[1]); |
| ga_j[i+2] += gy * x[i+2] * ((zbyte & 0x04) ? 0.0f : sw[2]); |
| ga_j[i+3] += gy * x[i+3] * ((zbyte & 0x08) ? 0.0f : sw[3]); |
| ga_j[i+4] += gy * x[i+4] * ((zbyte & 0x10) ? 0.0f : sw[4]); |
| ga_j[i+5] += gy * x[i+5] * ((zbyte & 0x20) ? 0.0f : sw[5]); |
| ga_j[i+6] += gy * x[i+6] * ((zbyte & 0x40) ? 0.0f : sw[6]); |
| ga_j[i+7] += gy * x[i+7] * ((zbyte & 0x80) ? 0.0f : sw[7]); |
| } |
| } |
| for (int i = (in/8)*8; i < in; i++) { |
| uint64_t z = (zb[i/64] >> (i%64)) & 1; |
| uint64_t wbit = (wb[i/64] >> (i%64)) & 1; |
| ga_j[i] += gy * x[i] * (z ? 0.0f : (wbit ? 1.0f : -1.0f)); |
| } |
| } else { |
| for (int i = 0; i + 7 < in; i += 8) { |
| ga_j[i+0] += gy * x[i+0]; ga_j[i+1] += gy * x[i+1]; |
| ga_j[i+2] += gy * x[i+2]; ga_j[i+3] += gy * x[i+3]; |
| ga_j[i+4] += gy * x[i+4]; ga_j[i+5] += gy * x[i+5]; |
| ga_j[i+6] += gy * x[i+6]; ga_j[i+7] += gy * x[i+7]; |
| } |
| for (int i = (in / 8) * 8; i < in; i++) |
| ga_j[i] += gy * x[i]; |
| } |
| gba[j] += gy; |
| } |
| return; |
| } |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| if (bl->w_float) { |
| int t = g_opt_step + 1; |
| float bc1 = 1.0f - powf(g_adam_beta1, (float)t); |
| float bc2 = 1.0f - powf(g_adam_beta2, (float)t); |
| for (int j = 0; j < out; j++) { |
| float gy = grad_y[j]; |
| if (fabsf(gy) < 1e-8f) continue; |
| |
| if (bl->logic_mask && bl->logic_mask[j] == 2) continue; |
| float *wf = &bl->w_float[j * in]; |
| if (g_use_adam && bl->m_adam) { |
| float *m = &bl->m_adam[j * in]; |
| float *v = &bl->v_adam[j * in]; |
| |
| for (int i = 0; i + 7 < in; i += 8) { |
| float g0 = gy * x[i+0], g1 = gy * x[i+1], g2 = gy * x[i+2], g3 = gy * x[i+3]; |
| float g4 = gy * x[i+4], g5 = gy * x[i+5], g6 = gy * x[i+6], g7 = gy * x[i+7]; |
| m[i+0] = g_adam_beta1*m[i+0] + (1.0f-g_adam_beta1)*g0; |
| m[i+1] = g_adam_beta1*m[i+1] + (1.0f-g_adam_beta1)*g1; |
| m[i+2] = g_adam_beta1*m[i+2] + (1.0f-g_adam_beta1)*g2; |
| m[i+3] = g_adam_beta1*m[i+3] + (1.0f-g_adam_beta1)*g3; |
| m[i+4] = g_adam_beta1*m[i+4] + (1.0f-g_adam_beta1)*g4; |
| m[i+5] = g_adam_beta1*m[i+5] + (1.0f-g_adam_beta1)*g5; |
| m[i+6] = g_adam_beta1*m[i+6] + (1.0f-g_adam_beta1)*g6; |
| m[i+7] = g_adam_beta1*m[i+7] + (1.0f-g_adam_beta1)*g7; |
| v[i+0] = g_adam_beta2*v[i+0] + (1.0f-g_adam_beta2)*g0*g0; |
| v[i+1] = g_adam_beta2*v[i+1] + (1.0f-g_adam_beta2)*g1*g1; |
| v[i+2] = g_adam_beta2*v[i+2] + (1.0f-g_adam_beta2)*g2*g2; |
| v[i+3] = g_adam_beta2*v[i+3] + (1.0f-g_adam_beta2)*g3*g3; |
| v[i+4] = g_adam_beta2*v[i+4] + (1.0f-g_adam_beta2)*g4*g4; |
| v[i+5] = g_adam_beta2*v[i+5] + (1.0f-g_adam_beta2)*g5*g5; |
| v[i+6] = g_adam_beta2*v[i+6] + (1.0f-g_adam_beta2)*g6*g6; |
| v[i+7] = g_adam_beta2*v[i+7] + (1.0f-g_adam_beta2)*g7*g7; |
| float mh0=m[i+0]/bc1, mh1=m[i+1]/bc1, mh2=m[i+2]/bc1, mh3=m[i+3]/bc1; |
| float mh4=m[i+4]/bc1, mh5=m[i+5]/bc1, mh6=m[i+6]/bc1, mh7=m[i+7]/bc1; |
| float vh0=sqrtf(v[i+0]/bc2)+g_adam_eps, vh1=sqrtf(v[i+1]/bc2)+g_adam_eps; |
| float vh2=sqrtf(v[i+2]/bc2)+g_adam_eps, vh3=sqrtf(v[i+3]/bc2)+g_adam_eps; |
| float vh4=sqrtf(v[i+4]/bc2)+g_adam_eps, vh5=sqrtf(v[i+5]/bc2)+g_adam_eps; |
| float vh6=sqrtf(v[i+6]/bc2)+g_adam_eps, vh7=sqrtf(v[i+7]/bc2)+g_adam_eps; |
| wf[i+0] -= lr * mh0/vh0; wf[i+1] -= lr * mh1/vh1; |
| wf[i+2] -= lr * mh2/vh2; wf[i+3] -= lr * mh3/vh3; |
| wf[i+4] -= lr * mh4/vh4; wf[i+5] -= lr * mh5/vh5; |
| wf[i+6] -= lr * mh6/vh6; wf[i+7] -= lr * mh7/vh7; |
| } |
| for (int i = (in / 8) * 8; i < in; i++) { |
| float g = gy * x[i]; |
| m[i] = g_adam_beta1*m[i] + (1.0f-g_adam_beta1)*g; |
| v[i] = g_adam_beta2*v[i] + (1.0f-g_adam_beta2)*g*g; |
| wf[i] -= lr * (m[i]/bc1) / (sqrtf(v[i]/bc2) + g_adam_eps); |
| } |
| } else { |
| |
| float scale = lr * gy; |
| for (int i = 0; i + 7 < in; i += 8) { |
| wf[i+0] -= scale * x[i+0]; wf[i+1] -= scale * x[i+1]; |
| wf[i+2] -= scale * x[i+2]; wf[i+3] -= scale * x[i+3]; |
| wf[i+4] -= scale * x[i+4]; wf[i+5] -= scale * x[i+5]; |
| wf[i+6] -= scale * x[i+6]; wf[i+7] -= scale * x[i+7]; |
| } |
| for (int i = (in / 8) * 8; i < in; i++) |
| wf[i] -= scale * x[i]; |
| } |
| |
| bl->bias[j] -= lr * gy; |
| } |
| |
| |
| |
| |
| |
| |
| |
| if (!g_use_pure_float) { |
| #define W_CLIP 1.0f |
| for (int i = 0; i < in * out; i++) { |
| float w = bl->w_float[i]; |
| if (w > W_CLIP) bl->w_float[i] = W_CLIP; |
| else if (w < -W_CLIP) bl->w_float[i] = -W_CLIP; |
| } |
| |
| bin_layer_repack(bl); |
| } else { |
| |
| #define W_CLIP_FLOAT 2.0f |
| for (int i = 0; i < in * out; i++) { |
| float w = bl->w_float[i]; |
| if (w > W_CLIP_FLOAT) bl->w_float[i] = W_CLIP_FLOAT; |
| else if (w < -W_CLIP_FLOAT) bl->w_float[i] = -W_CLIP_FLOAT; |
| } |
| } |
| |
| |
| |
| |
| |
| if (bl->zbits) bin_layer_repack_ternary(bl); |
| } else { |
| |
| float mean_abs_x = 0; |
| for (int i = 0; i < in; i++) mean_abs_x += fabsf(x[i]); |
| mean_abs_x /= in; |
| uint64_t xbits[64]; |
| for (int wi = 0; wi < bl->n_words; wi++) { |
| uint64_t word = 0; |
| for (int bi = 0; bi < 64; bi++) { |
| int idx = wi * 64 + bi; |
| if (idx < in && x[idx] > 0.0f) word |= (1ULL << bi); |
| } |
| xbits[wi] = word; |
| } |
| for (int j = 0; j < out; j++) { |
| float gy = grad_y[j]; |
| if (fabsf(gy) < 1e-6f) continue; |
| int pc = 0; |
| const uint64_t *wb = &bl->wbits[j * bl->n_words]; |
| for (int wi = 0; wi < bl->n_words; wi++) |
| pc += __builtin_popcountll(~(xbits[wi] ^ wb[wi])); |
| float grad_alpha = (float)(2 * pc - in) * mean_abs_x; |
| bl->alpha[j] -= lr * grad_alpha * gy; |
| if (bl->alpha[j] < 0.0f) bl->alpha[j] = 0.0f; |
| bl->bias[j] -= lr * gy; |
| } |
| } |
| } |
|
|
| |
| |
| |
| void layer_norm(float *out, const float *x, const float *w, const float *b, int n) { |
| float mean = 0; |
| for (int i = 0; i < n; i++) mean += x[i]; |
| mean /= n; |
| float var = 0; |
| for (int i = 0; i < n; i++) { float d = x[i] - mean; var += d * d; } |
| var /= n; |
| float is = 1.0f / sqrtf(var + 1e-5f); |
| for (int i = 0; i < n; i++) out[i] = (x[i] - mean) * is * w[i] + b[i]; |
| } |
|
|
| void layer_norm_backward(float *grad_x, const float *grad_y, const float *x, |
| const float *w, float mean, float std_inv, int n, |
| float *grad_w, float *grad_b) { |
| float sum_grad = 0; |
| for (int i = 0; i < n; i++) sum_grad += grad_y[i] * w[i] * (x[i] - mean); |
| float common = std_inv / n * sum_grad; |
| float scale = (1.0f - 1.0f / n); |
| for (int i = 0; i < n; i++) { |
| grad_x[i] = grad_y[i] * w[i] * std_inv * scale - common; |
| if (grad_w) grad_w[i] += grad_y[i] * (x[i] - mean) * std_inv; |
| if (grad_b) grad_b[i] += grad_y[i]; |
| } |
| } |
|
|
| float gelu(float x) { |
| return 0.5f * x * (1.0f + tanhf(0.7978845608f * (x + 0.044715f * x * x * x))); |
| } |
|
|
| float gelu_grad(float x) { |
| float inner = 0.7978845608f * (x + 0.044715f * x * x * x); |
| float t = tanhf(inner); |
| return 0.5f * (1.0f + t) + 0.5f * x * (1.0f - t * t) * 0.7978845608f * (1.0f + 0.134145f * x * x); |
| } |
|
|
| void softmax(float *x, int n) { |
| float mx = x[0]; |
| for (int i = 1; i < n; i++) if (x[i] > mx) mx = x[i]; |
| float sum = 0; |
| for (int i = 0; i < n; i++) { x[i] = expf(x[i] - mx); sum += x[i]; } |
| for (int i = 0; i < n; i++) x[i] /= sum; |
| } |
|
|
| float cross_entropy_sampled(const float *hidden, const float *wte, |
| int target, int vocab_size, int n_embd, |
| int n_samples, unsigned int *seed) { |
| float tl = 0; |
| for (int i = 0; i < n_embd; i++) tl += hidden[i] * wte[target * n_embd + i]; |
| float mx = tl; |
| float neg[256]; |
| for (int k = 0; k < n_samples && k < 256; k++) { |
| int v = rand_r(seed) % vocab_size; |
| float s = 0; |
| for (int i = 0; i < n_embd; i++) s += hidden[i] * wte[v * n_embd + i]; |
| neg[k] = s; |
| if (s > mx) mx = s; |
| } |
| float se = expf(tl - mx); |
| for (int k = 0; k < n_samples && k < 256; k++) se += expf(neg[k] - mx); |
| return -logf(expf(tl - mx) / se + 1e-7f); |
| } |
|
|
| void cross_entropy_grad(float *grad_hidden, const float *hidden, const float *wte, |
| int target, int vocab_size, int n_embd, |
| int n_samples, unsigned int *seed) { |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| float tl = 0; |
| for (int i = 0; i < n_embd; i++) tl += hidden[i] * wte[target * n_embd + i]; |
|
|
| |
| |
| |
| float neg[256]; |
| int actual = n_samples < 256 ? n_samples : 256; |
| float mx = tl; |
| for (int k = 0; k < actual; k++) { |
| int v = rand_r(seed) % vocab_size; |
| float s = 0; |
| for (int i = 0; i < n_embd; i++) s += hidden[i] * wte[v * n_embd + i]; |
| neg[k] = s; |
| if (s > mx) mx = s; |
| } |
|
|
| float se = expf(tl - mx); |
| for (int k = 0; k < actual; k++) se += expf(neg[k] - mx); |
|
|
| |
| |
| float prob = expf(tl - mx) / (se + 1e-7f); |
|
|
| |
| |
| |
| |
| float coef = -(1.0f - prob); |
| const float *wt = &wte[target * n_embd]; |
| for (int i = 0; i < n_embd; i++) |
| grad_hidden[i] = coef * wt[i]; |
| } |
|
|
| void clip_array(float *x, int n, float clip_val) { |
| for (int i = 0; i < n; i++) { |
| if (x[i] > clip_val) x[i] = clip_val; |
| if (x[i] < -clip_val) x[i] = -clip_val; |
| } |
| } |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| void normalize_residual(float *x, int n, float target_norm) { |
| float norm_sq = 0; |
| for (int i = 0; i < n; i++) norm_sq += x[i] * x[i]; |
| float norm = sqrtf(norm_sq) + 1e-8f; |
| |
| if (norm > target_norm) { |
| float scale = target_norm / norm; |
| for (int i = 0; i < n; i++) x[i] *= scale; |
| } |
| } |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| void scale_to_norm(float *v, int n, float target_norm) { |
| float norm_sq = 0; |
| for (int i = 0; i < n; i++) norm_sq += v[i] * v[i]; |
| float norm = sqrtf(norm_sq) + 1e-8f; |
| float scale = target_norm / norm; |
| for (int i = 0; i < n; i++) v[i] *= scale; |
| } |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| static void compute_full_logits(const float *hidden, const float *wte, |
| float *logits_out, int vocab, int n_embd) { |
| #pragma omp parallel for schedule(static) |
| for (int j = 0; j < vocab; j++) { |
| const float *w = &wte[(size_t)j * n_embd]; |
| float s = 0; |
| for (int i = 0; i + 7 < n_embd; i += 8) |
| s += hidden[i+0]*w[i+0] + hidden[i+1]*w[i+1] |
| + hidden[i+2]*w[i+2] + hidden[i+3]*w[i+3] |
| + hidden[i+4]*w[i+4] + hidden[i+5]*w[i+5] |
| + hidden[i+6]*w[i+6] + hidden[i+7]*w[i+7]; |
| for (int i = (n_embd/8)*8; i < n_embd; i++) s += hidden[i] * w[i]; |
| logits_out[j] = s * g_logit_scale; |
| } |
| } |
|
|
| float cross_entropy_full(const float *hidden, const float *wte, |
| int target, int vocab_size, int n_embd, |
| float *logits_scratch) { |
| compute_full_logits(hidden, wte, logits_scratch, vocab_size, n_embd); |
| |
| float mx = logits_scratch[0]; |
| for (int j = 1; j < vocab_size; j++) |
| if (logits_scratch[j] > mx) mx = logits_scratch[j]; |
| float sum = 0; |
| for (int j = 0; j < vocab_size; j++) { |
| logits_scratch[j] = expf(logits_scratch[j] - mx); |
| sum += logits_scratch[j]; |
| } |
| |
| float p_target = logits_scratch[target] / sum; |
| return -logf(p_target + 1e-12f); |
| } |
|
|
| void cross_entropy_full_grad(float *grad_hidden, const float *hidden, const float *wte, |
| int target, int vocab_size, int n_embd, |
| float *logits_scratch) { |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| const float *wt_target = &wte[(size_t)target * n_embd]; |
| for (int i = 0; i < n_embd; i++) |
| grad_hidden[i] = -wt_target[i]; |
|
|
| |
| |
| |
| float psum = 0; |
| for (int j = 0; j < vocab_size; j++) psum += logits_scratch[j]; |
| float inv_psum = 1.0f / (psum + 1e-12f); |
|
|
| |
| |
| for (int j = 0; j < vocab_size; j++) { |
| float p = logits_scratch[j] * inv_psum; |
| if (p < 1e-7f) continue; |
| const float *w = &wte[(size_t)j * n_embd]; |
| float coef = p; |
| for (int i = 0; i + 7 < n_embd; i += 8) { |
| grad_hidden[i+0] += coef * w[i+0]; grad_hidden[i+1] += coef * w[i+1]; |
| grad_hidden[i+2] += coef * w[i+2]; grad_hidden[i+3] += coef * w[i+3]; |
| grad_hidden[i+4] += coef * w[i+4]; grad_hidden[i+5] += coef * w[i+5]; |
| grad_hidden[i+6] += coef * w[i+6]; grad_hidden[i+7] += coef * w[i+7]; |
| } |
| for (int i = (n_embd/8)*8; i < n_embd; i++) |
| grad_hidden[i] += coef * w[i]; |
| } |
| |
| if (g_logit_scale != 1.0f) |
| for (int i = 0; i < n_embd; i++) grad_hidden[i] *= g_logit_scale; |
| } |
|
|
| void compute_mean_std(const float *x, int n, float *mean, float *std_inv) { |
| float m = 0; |
| for (int i = 0; i < n; i++) m += x[i]; |
| m /= n; |
| float var = 0; |
| for (int i = 0; i < n; i++) { float d = x[i] - m; var += d * d; } |
| var /= n; |
| *mean = m; |
| *std_inv = 1.0f / sqrtf(var + 1e-5f); |
| } |
|
|
| |
| |
| |
| Tensor *tensor_load_all(const char *path, int *n_tensors) { |
| FILE *f = fopen(path, "rb"); |
| if (!f) { fprintf(stderr, "cannot open %s\n", path); return NULL; } |
| char magic[4]; |
| fread(magic, 1, 4, f); |
| if (memcmp(magic, "GPW2", 4) != 0) { fprintf(stderr, "bad magic\n"); fclose(f); return NULL; } |
| fread(n_tensors, 4, 1, f); |
| Tensor *t = calloc(*n_tensors, sizeof(Tensor)); |
| for (int i = 0; i < *n_tensors; i++) { |
| int klen; |
| fread(&klen, 4, 1, f); |
| fread(t[i].key, 1, klen, f); |
| t[i].key[klen] = '\0'; |
| fread(&t[i].ndim, 4, 1, f); |
| int n = 1; |
| for (int d = 0; d < t[i].ndim; d++) { |
| fread(&t[i].shape[d], 4, 1, f); |
| n *= t[i].shape[d]; |
| } |
| t[i].data = malloc(n * sizeof(float)); |
| fread(t[i].data, 4, n, f); |
| } |
| fclose(f); |
| return t; |
| } |
|
|
| float *tensor_get(Tensor *tensors, int n, const char *key) { |
| for (int i = 0; i < n; i++) |
| if (strcmp(tensors[i].key, key) == 0) return tensors[i].data; |
| fprintf(stderr, "tensor not found: %s\n", key); |
| return NULL; |
| } |
|
|
| void tensor_free_all(Tensor *tensors, int n) { |
| for (int i = 0; i < n; i++) free(tensors[i].data); |
| free(tensors); |
| } |
|
|
| |
| |
| |
| void tensor_free_data_by_key(Tensor *tensors, int n, const char *key) { |
| for (int i = 0; i < n; i++) { |
| if (tensors[i].data && strcmp(tensors[i].key, key) == 0) { |
| free(tensors[i].data); |
| tensors[i].data = NULL; |
| return; |
| } |
| } |
| } |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| #ifndef _WIN32 |
| #include <sys/mman.h> |
| #include <sys/stat.h> |
| #endif |
| #include <fcntl.h> |
| #ifndef _WIN32 |
| #include <unistd.h> |
| #else |
| |
| #define open _open |
| #define close _close |
| #define read _read |
| #define lseek _lseek |
| #define O_RDONLY _O_RDONLY |
| #endif |
|
|
| typedef struct { |
| Tensor *tensors; |
| int n_tensors; |
| void *mmap_base; |
| size_t mmap_size; |
| int fd; |
| } MmapedTensors; |
|
|
| static MmapedTensors g_mmap_state = {NULL, 0, NULL, 0, -1}; |
|
|
| |
| |
| |
| |
| |
| |
| |
| #ifndef M_PI |
| #define M_PI 3.14159265358979323846f |
| #endif |
| typedef struct { char key[64]; int ndim; int shape[4]; } TEntry; |
| static float bin_randn(void) { |
| static int has = 0; static float spare = 0.0f; |
| if (has) { has = 0; return spare; } |
| float u = (rand() + 1.0f) / (RAND_MAX + 2.0f); |
| float v = (rand() + 1.0f) / (RAND_MAX + 2.0f); |
| float mag = sqrtf(-2.0f * logf(u)); |
| spare = mag * sinf(2.0f * M_PI * v); has = 1; |
| return mag * cosf(2.0f * M_PI * v); |
| } |
| static void bin_push(TEntry **E, int *cnt, const char *key, int ndim, int s0, int s1, int s2, int s3) { |
| TEntry *e = &(*E)[(*cnt)++]; |
| strncpy(e->key, key, sizeof(e->key) - 1); e->key[sizeof(e->key) - 1] = '\0'; |
| e->ndim = ndim; e->shape[0] = s0; e->shape[1] = s1; e->shape[2] = s2; e->shape[3] = s3; |
| } |
| static void bin_gpw2_put(FILE *f, const char *key, int ndim, const int *shape) { |
| int klen = (int)strlen(key), n = 1; |
| for (int d = 0; d < ndim; d++) n *= shape[d]; |
| fwrite(&klen, 4, 1, f); fwrite(key, 1, (size_t)klen, f); |
| fwrite(&ndim, 4, 1, f); |
| for (int d = 0; d < ndim; d++) fwrite(&shape[d], 4, 1, f); |
| for (int i = 0; i < n; i++) { float g = bin_randn() * 0.02f; fwrite(&g, 4, 1, f); } |
| } |
| |
| static void bin_gpw2_put_init(FILE *f, const char *key, int ndim, const int *shape, float scale, int init_mode) { |
| |
| int klen = (int)strlen(key), n = 1; |
| for (int d = 0; d < ndim; d++) n *= shape[d]; |
| fwrite(&klen, 4, 1, f); fwrite(key, 1, (size_t)klen, f); |
| fwrite(&ndim, 4, 1, f); |
| for (int d = 0; d < ndim; d++) fwrite(&shape[d], 4, 1, f); |
| if (init_mode == 1) { |
| |
| for (int i = 0; i < n; i++) fwrite(&scale, 4, 1, f); |
| } else if (init_mode == 2) { |
| |
| float z = 0.0f; |
| for (int i = 0; i < n; i++) fwrite(&z, 4, 1, f); |
| } else if (init_mode == 3) { |
| |
| float fan_in = (float)shape[ndim - 1]; |
| float std_val = sqrtf(2.0f / fan_in); |
| for (int i = 0; i < n; i++) { float g = bin_randn() * std_val; fwrite(&g, 4, 1, f); } |
| } else { |
| |
| for (int i = 0; i < n; i++) { float g = bin_randn() * scale; fwrite(&g, 4, 1, f); } |
| } |
| } |
|
|
| void gen_random_gpw2(const char *path, ModelConfig cfg) { |
| int n = cfg.n_embd, m = cfg.mlp_dim, V = cfg.vocab_size, C = cfg.n_ctx; |
| int cnt = 0; char kb[64]; |
| FILE *f = fopen(path, "wb"); |
| if (!f) { fprintf(stderr, "gen_random_gpw2: cannot write %s\n", path); exit(1); } |
|
|
| |
| int per_layer; |
| if (cfg.qkv_merged) { |
| per_layer = (cfg.act_type == ACT_SWIGLU) ? 11 : 12; |
| } else { |
| per_layer = 9; |
| } |
| int n_tensors = 4 + cfg.n_layer * per_layer; |
|
|
| fwrite("GPW2", 1, 4, f); |
| fwrite(&n_tensors, 4, 1, f); |
|
|
| |
| float emb_scale = 1.0f / sqrtf((float)n); |
| bin_gpw2_put_init(f, "wte.weight", 2, (int[]){V, n}, emb_scale, 0); |
| if (cfg.attn_type == ATTN_LEARNED) |
| bin_gpw2_put_init(f, "wpe.weight", 2, (int[]){C, n}, emb_scale, 0); |
|
|
| |
| bin_gpw2_put_init(f, "ln_f.weight", 1, (int[]){n}, 1.0f, 1); |
| bin_gpw2_put_init(f, "ln_f.bias", 1, (int[]){n}, 0.0f, 2); |
|
|
| for (int l = 0; l < cfg.n_layer; l++) { |
| if (cfg.qkv_merged) { |
| |
| snprintf(kb, sizeof kb, "h.%d.attn.c_attn.weight", l); |
| bin_gpw2_put_init(f, kb, 2, (int[]){3*n, n}, 0.0f, 3); |
| snprintf(kb, sizeof kb, "h.%d.attn.c_attn.bias", l); |
| bin_gpw2_put_init(f, kb, 1, (int[]){3*n}, 0.0f, 2); |
| snprintf(kb, sizeof kb, "h.%d.attn.c_proj.weight", l); |
| bin_gpw2_put_init(f, kb, 2, (int[]){n, n}, 0.0f, 3); |
| snprintf(kb, sizeof kb, "h.%d.attn.c_proj.bias", l); |
| bin_gpw2_put_init(f, kb, 1, (int[]){n}, 0.0f, 2); |
| if (cfg.act_type == ACT_SWIGLU) { |
| snprintf(kb, sizeof kb, "h.%d.mlp.gate_proj.weight", l); |
| bin_gpw2_put_init(f, kb, 2, (int[]){m, n}, 0.0f, 3); |
| snprintf(kb, sizeof kb, "h.%d.mlp.up_proj.weight", l); |
| bin_gpw2_put_init(f, kb, 2, (int[]){m, n}, 0.0f, 3); |
| snprintf(kb, sizeof kb, "h.%d.mlp.down_proj.weight", l); |
| bin_gpw2_put_init(f, kb, 2, (int[]){n, m}, 0.0f, 3); |
| } else { |
| snprintf(kb, sizeof kb, "h.%d.mlp.c_fc.weight", l); |
| bin_gpw2_put_init(f, kb, 2, (int[]){m, n}, 0.0f, 3); |
| snprintf(kb, sizeof kb, "h.%d.mlp.c_fc.bias", l); |
| bin_gpw2_put_init(f, kb, 1, (int[]){m}, 0.0f, 2); |
| snprintf(kb, sizeof kb, "h.%d.mlp.c_proj.weight", l); |
| bin_gpw2_put_init(f, kb, 2, (int[]){n, m}, 0.0f, 3); |
| snprintf(kb, sizeof kb, "h.%d.mlp.c_proj.bias", l); |
| bin_gpw2_put_init(f, kb, 1, (int[]){n}, 0.0f, 2); |
| } |
| |
| snprintf(kb, sizeof kb, "h.%d.ln_1.weight", l); |
| bin_gpw2_put_init(f, kb, 1, (int[]){n}, 1.0f, 1); |
| snprintf(kb, sizeof kb, "h.%d.ln_1.bias", l); |
| bin_gpw2_put_init(f, kb, 1, (int[]){n}, 0.0f, 2); |
| snprintf(kb, sizeof kb, "h.%d.ln_2.weight", l); |
| bin_gpw2_put_init(f, kb, 1, (int[]){n}, 1.0f, 1); |
| snprintf(kb, sizeof kb, "h.%d.ln_2.bias", l); |
| bin_gpw2_put_init(f, kb, 1, (int[]){n}, 0.0f, 2); |
| } else { |
| snprintf(kb, sizeof kb, "model.layers.%d.self_attn.q_proj.weight", l); |
| bin_gpw2_put_init(f, kb, 2, (int[]){n, n}, 0.0f, 3); |
| snprintf(kb, sizeof kb, "model.layers.%d.self_attn.k_proj.weight", l); |
| bin_gpw2_put_init(f, kb, 2, (int[]){n, n}, 0.0f, 3); |
| snprintf(kb, sizeof kb, "model.layers.%d.self_attn.v_proj.weight", l); |
| bin_gpw2_put_init(f, kb, 2, (int[]){n, n}, 0.0f, 3); |
| snprintf(kb, sizeof kb, "model.layers.%d.self_attn.o_proj.weight", l); |
| bin_gpw2_put_init(f, kb, 2, (int[]){n, n}, 0.0f, 3); |
| snprintf(kb, sizeof kb, "model.layers.%d.mlp.gate_proj.weight", l); |
| bin_gpw2_put_init(f, kb, 2, (int[]){m, n}, 0.0f, 3); |
| snprintf(kb, sizeof kb, "model.layers.%d.mlp.up_proj.weight", l); |
| bin_gpw2_put_init(f, kb, 2, (int[]){m, n}, 0.0f, 3); |
| snprintf(kb, sizeof kb, "model.layers.%d.mlp.down_proj.weight", l); |
| bin_gpw2_put_init(f, kb, 2, (int[]){n, m}, 0.0f, 3); |
| snprintf(kb, sizeof kb, "model.layers.%d.input_layernorm.weight", l); |
| bin_gpw2_put_init(f, kb, 1, (int[]){n}, 1.0f, 1); |
| snprintf(kb, sizeof kb, "model.layers.%d.post_attention_layernorm.weight", l); |
| bin_gpw2_put_init(f, kb, 1, (int[]){n}, 1.0f, 1); |
| } |
| cnt++; |
| } |
| fclose(f); |
| printf("[*] generated random weights (Xavier init, LN=1.0): %d tensors -> %s\n", n_tensors, path); |
| } |
|
|
| Tensor *tensor_load_all_mmap(const char *path, int *n_tensors) { |
| int fd = open(path, O_RDONLY); |
| if (fd < 0) { fprintf(stderr, "cannot open %s\n", path); return NULL; } |
| struct stat st; |
| if (fstat(fd, &st) < 0) { fprintf(stderr, "fstat failed\n"); close(fd); return NULL; } |
| size_t file_size = st.st_size; |
| void *base = mmap(NULL, file_size, PROT_READ, MAP_PRIVATE, fd, 0); |
| if (base == MAP_FAILED) { fprintf(stderr, "mmap failed\n"); close(fd); return NULL; } |
|
|
| const unsigned char *p = (const unsigned char *)base; |
| if (memcmp(p, "GPW2", 4) != 0) { fprintf(stderr, "bad magic\n"); munmap(base, file_size); close(fd); return NULL; } |
| p += 4; |
| int n = *(const int *)p; p += 4; |
| *n_tensors = n; |
| Tensor *t = calloc(n, sizeof(Tensor)); |
| for (int i = 0; i < n; i++) { |
| int klen = *(const int *)p; p += 4; |
| memcpy(t[i].key, p, klen); t[i].key[klen] = '\0'; p += klen; |
| t[i].ndim = *(const int *)p; p += 4; |
| int sz = 1; |
| for (int d = 0; d < t[i].ndim; d++) { |
| t[i].shape[d] = *(const int *)p; p += 4; |
| sz *= t[i].shape[d]; |
| } |
| |
| t[i].data = (float *)p; |
| p += sz * sizeof(float); |
| } |
| g_mmap_state.tensors = t; |
| g_mmap_state.n_tensors = n; |
| g_mmap_state.mmap_base = base; |
| g_mmap_state.mmap_size = file_size; |
| g_mmap_state.fd = fd; |
| return t; |
| } |
|
|
| void tensor_free_all_mmap(Tensor *tensors, int n) { |
| |
| free(tensors); |
| if (g_mmap_state.mmap_base) { |
| munmap(g_mmap_state.mmap_base, g_mmap_state.mmap_size); |
| g_mmap_state.mmap_base = NULL; |
| } |
| if (g_mmap_state.fd >= 0) { |
| close(g_mmap_state.fd); |
| g_mmap_state.fd = -1; |
| } |
| } |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| |
| StatefulContext g_sctx = {0}; |
|
|
| |
| void attention_forward_sliding(float *attn_out, const float *qkv, |
| int n_embd, int n_head, |
| int seq_pos, |
| float *k_cache_layer, float *v_cache_layer, |
| int n_ctx, int window_size, int n_sinks) { |
| int head_dim = n_embd / n_head; |
| float scale = 1.0f / sqrtf((float)head_dim); |
|
|
| const float *Q = qkv; |
| const float *K_new = qkv + n_embd; |
| const float *V_new = qkv + 2 * n_embd; |
|
|
| |
| int cache_pos = seq_pos % n_ctx; |
| memcpy(k_cache_layer + (size_t)cache_pos * n_embd, K_new, n_embd * sizeof(float)); |
| memcpy(v_cache_layer + (size_t)cache_pos * n_embd, V_new, n_embd * sizeof(float)); |
|
|
| |
| int n_sink = (seq_pos < n_sinks) ? seq_pos : n_sinks; |
| int win_start = seq_pos - window_size + 1; |
| if (win_start < n_sinks) win_start = n_sinks; |
| if (win_start > seq_pos) win_start = 0; |
| int n_win = seq_pos - win_start + 1; |
| if (n_win < 0) n_win = 0; |
| int n_attend = n_sink + n_win; |
| if (n_attend < 1) n_attend = seq_pos + 1; |
| if (n_attend > n_ctx) n_attend = n_ctx; |
|
|
| |
| |
| static __thread int *tl_pos_list = NULL; |
| static __thread float *tl_scores = NULL; |
| static __thread float *tl_attn_w = NULL; |
| static __thread int tl_n = 0; |
| if (tl_n < n_attend) { |
| free(tl_pos_list); free(tl_scores); free(tl_attn_w); |
| tl_pos_list = (int *)malloc(n_attend * sizeof(int)); |
| tl_scores = (float *)malloc(n_attend * sizeof(float)); |
| tl_attn_w = (float *)malloc(n_attend * sizeof(float)); |
| tl_n = n_attend; |
| } |
| int *pos_list = tl_pos_list; |
| float *scores = tl_scores; |
| float *attn_w = tl_attn_w; |
|
|
| int idx = 0; |
| for (int j = 0; j < n_sink && idx < n_attend; j++) pos_list[idx++] = j; |
| for (int j = win_start; j <= seq_pos && idx < n_attend; j++) pos_list[idx++] = j; |
|
|
| |
| |
| |
| if (n_attend > 256) { |
| #pragma omp parallel for schedule(static) |
| for (int h = 0; h < n_head; h++) { |
| const float *Q_h = Q + h * head_dim; |
|
|
| |
| float max_score = -1e30f; |
| for (int i = 0; i < n_attend; i++) { |
| int j = pos_list[i]; |
| int phys_j = j % n_ctx; |
| const float *K_jh = k_cache_layer + (size_t)phys_j * n_embd + h * head_dim; |
| float dot = 0.0f; |
| for (int d = 0; d + 7 < head_dim; d += 8) |
| dot += Q_h[d]*K_jh[d] + Q_h[d+1]*K_jh[d+1] + Q_h[d+2]*K_jh[d+2] + Q_h[d+3]*K_jh[d+3] |
| + Q_h[d+4]*K_jh[d+4] + Q_h[d+5]*K_jh[d+5] + Q_h[d+6]*K_jh[d+6] + Q_h[d+7]*K_jh[d+7]; |
| for (int d = (head_dim/8)*8; d < head_dim; d++) dot += Q_h[d] * K_jh[d]; |
| dot *= scale; |
| scores[i] = dot; |
| if (dot > max_score) max_score = dot; |
| } |
|
|
| |
| float sum_exp = 0.0f; |
| for (int i = 0; i < n_attend; i++) { |
| float e = expf(scores[i] - max_score); |
| attn_w[i] = e; |
| sum_exp += e; |
| } |
| float inv_sum = 1.0f / (sum_exp + 1e-12f); |
| for (int i = 0; i < n_attend; i++) attn_w[i] *= inv_sum; |
|
|
| |
| float *out_h = attn_out + h * head_dim; |
| for (int d = 0; d < head_dim; d++) out_h[d] = 0.0f; |
| for (int i = 0; i < n_attend; i++) { |
| int j = pos_list[i]; |
| int phys_j = j % n_ctx; |
| float w = attn_w[i]; |
| const float *V_jh = v_cache_layer + (size_t)phys_j * n_embd + h * head_dim; |
| for (int d = 0; d + 7 < head_dim; d += 8) { |
| out_h[d+0] += w * V_jh[d+0]; out_h[d+1] += w * V_jh[d+1]; |
| out_h[d+2] += w * V_jh[d+2]; out_h[d+3] += w * V_jh[d+3]; |
| out_h[d+4] += w * V_jh[d+4]; out_h[d+5] += w * V_jh[d+5]; |
| out_h[d+6] += w * V_jh[d+6]; out_h[d+7] += w * V_jh[d+7]; |
| } |
| for (int d = (head_dim/8)*8; d < head_dim; d++) out_h[d] += w * V_jh[d]; |
| } |
| } |
| } else { |
| |
| for (int h = 0; h < n_head; h++) { |
| const float *Q_h = Q + h * head_dim; |
| float max_score = -1e30f; |
| for (int i = 0; i < n_attend; i++) { |
| int j = pos_list[i]; |
| int phys_j = j % n_ctx; |
| const float *K_jh = k_cache_layer + (size_t)phys_j * n_embd + h * head_dim; |
| float dot = 0.0f; |
| for (int d = 0; d + 7 < head_dim; d += 8) |
| dot += Q_h[d]*K_jh[d] + Q_h[d+1]*K_jh[d+1] + Q_h[d+2]*K_jh[d+2] + Q_h[d+3]*K_jh[d+3] |
| + Q_h[d+4]*K_jh[d+4] + Q_h[d+5]*K_jh[d+5] + Q_h[d+6]*K_jh[d+6] + Q_h[d+7]*K_jh[d+7]; |
| for (int d = (head_dim/8)*8; d < head_dim; d++) dot += Q_h[d] * K_jh[d]; |
| dot *= scale; |
| scores[i] = dot; |
| if (dot > max_score) max_score = dot; |
| } |
| float sum_exp = 0.0f; |
| for (int i = 0; i < n_attend; i++) { |
| float e = expf(scores[i] - max_score); |
| attn_w[i] = e; |
| sum_exp += e; |
| } |
| float inv_sum = 1.0f / (sum_exp + 1e-12f); |
| for (int i = 0; i < n_attend; i++) attn_w[i] *= inv_sum; |
| float *out_h = attn_out + h * head_dim; |
| for (int d = 0; d < head_dim; d++) out_h[d] = 0.0f; |
| for (int i = 0; i < n_attend; i++) { |
| int j = pos_list[i]; |
| int phys_j = j % n_ctx; |
| float w = attn_w[i]; |
| const float *V_jh = v_cache_layer + (size_t)phys_j * n_embd + h * head_dim; |
| for (int d = 0; d + 7 < head_dim; d += 8) { |
| out_h[d+0] += w * V_jh[d+0]; out_h[d+1] += w * V_jh[d+1]; |
| out_h[d+2] += w * V_jh[d+2]; out_h[d+3] += w * V_jh[d+3]; |
| out_h[d+4] += w * V_jh[d+4]; out_h[d+5] += w * V_jh[d+5]; |
| out_h[d+6] += w * V_jh[d+6]; out_h[d+7] += w * V_jh[d+7]; |
| } |
| for (int d = (head_dim/8)*8; d < head_dim; d++) out_h[d] += w * V_jh[d]; |
| } |
| } |
| } |
| } |
|
|
| |
| |
| |
| |
| typedef struct ConceptAttnStats { |
| long forwards; |
| long forwards_ctx; |
| long candidates; |
| long full_equiv; |
| long gate_pairs; |
| long gate_blocked; |
| long msg_candidates; |
| double msg_mass; |
| int last_n_filled; |
| long ctx_memory_slots_used; |
| long ctx_total_attend; |
| |
| double msg_inter_cos; |
| double msg_norm; |
| long msg_segments; |
| } ConceptAttnStats; |
| ConceptAttnStats g_ca_stats = {0}; |
| void concept_attn_stats_reset(void); |
| void concept_attn_stats_reset(void) { |
| int keep = g_ca_stats.last_n_filled; |
| ConceptAttnStats z = {0}; |
| g_ca_stats = z; |
| g_ca_stats.last_n_filled = keep; |
| } |
|
|
| |
| |
| |
| |
| void attention_forward_concept_ctx(float *attn_out, const float *qkv, |
| int n_embd, int n_head, |
| int seq_pos, |
| float *k_cache_layer, float *v_cache_layer, |
| int n_ctx, int window_size, int n_sinks, |
| const float *wte, |
| const float *cctx_k, const float *cctx_v, |
| const int *cctx_cnt, const int *cctx_anchor, |
| int n_slots, float mem_scale) { |
| int head_dim = n_embd / n_head; |
| float scale = 1.0f / sqrtf((float)head_dim); |
|
|
| |
| g_ca_stats.forwards_ctx++; |
| g_ca_stats.full_equiv += (long)(seq_pos + 1) * n_head; |
|
|
| const float *Q = qkv; |
| const float *K_new = qkv + n_embd; |
| const float *V_new = qkv + 2 * n_embd; |
|
|
| int cache_pos = seq_pos % n_ctx; |
| memcpy(k_cache_layer + (size_t)cache_pos * n_embd, K_new, n_embd * sizeof(float)); |
| memcpy(v_cache_layer + (size_t)cache_pos * n_embd, V_new, n_embd * sizeof(float)); |
|
|
| int n_sink = (seq_pos < n_sinks) ? seq_pos : n_sinks; |
| int win_start = seq_pos - window_size + 1; |
| if (win_start < n_sinks) win_start = n_sinks; |
| if (win_start > seq_pos) win_start = 0; |
| int n_win = seq_pos - win_start + 1; |
| if (n_win < 0) n_win = 0; |
| int n_attend = n_sink + n_win; |
|
|
| int n_mem = 0; |
| for (int s = 0; s < n_slots; s++) if (cctx_cnt[s] > 0) n_mem++; |
| int n_total = n_attend + n_mem; |
| if (n_total < 1) n_total = seq_pos + 1; |
| if (n_total > n_ctx + n_slots) n_total = n_ctx + n_slots; |
|
|
| |
| g_ca_stats.ctx_total_attend += (long)n_total * n_head; |
| g_ca_stats.ctx_memory_slots_used += (long)n_mem * n_head; |
| g_ca_stats.candidates += (long)n_total * n_head; |
|
|
| int pos_idx[8192]; |
| int *pos_list = (n_total <= 8192) ? pos_idx : malloc(n_total * sizeof(int)); |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| char is_mem[8192]; |
| char *mem_flag = (n_total <= 8192) ? is_mem : malloc(n_total * sizeof(char)); |
| int idx = 0; |
| for (int j = 0; j < n_sink && idx < n_attend; j++) { pos_list[idx] = j; mem_flag[idx] = 0; idx++; } |
| for (int j = win_start; j <= seq_pos && idx < n_attend; j++) { pos_list[idx] = j; mem_flag[idx] = 0; idx++; } |
| int mem_written = 0; |
| for (int s = 0; s < n_slots && mem_written < n_mem; s++) { |
| if (cctx_cnt[s] > 0) { pos_list[idx] = s; mem_flag[idx] = 1; idx++; mem_written++; } |
| } |
| int n_final = idx; |
|
|
| float scores_stack[8192]; |
| float *scores = (n_total <= 8192) ? scores_stack : malloc(n_total * sizeof(float)); |
| float *attn_w = (n_total <= 8192) ? scores_stack : malloc(n_total * sizeof(float)); |
|
|
| for (int h = 0; h < n_head; h++) { |
| const float *Q_h = Q + h * head_dim; |
| float max_score = -1e30f; |
| for (int i = 0; i < n_final; i++) { |
| float dot; |
| if (mem_flag[i]) { |
| |
| const float *K_mh = cctx_k + (size_t)pos_list[i] * n_embd + h * head_dim; |
| dot = 0.0f; |
| for (int d = 0; d < head_dim; d++) dot += Q_h[d] * K_mh[d]; |
| dot *= scale * mem_scale; |
| (void)wte; (void)cctx_anchor; |
| } else { |
| int j = pos_list[i]; |
| int phys_j = j % n_ctx; |
| const float *K_jh = k_cache_layer + (size_t)phys_j * n_embd + h * head_dim; |
| dot = 0.0f; |
| for (int d = 0; d < head_dim; d++) dot += Q_h[d] * K_jh[d]; |
| dot *= scale; |
| } |
| scores[i] = dot; |
| if (dot > max_score) max_score = dot; |
| } |
| float sum_exp = 0.0f; |
| for (int i = 0; i < n_final; i++) { |
| float e = expf(scores[i] - max_score); |
| attn_w[i] = e; |
| sum_exp += e; |
| } |
| float inv_sum = 1.0f / (sum_exp + 1e-12f); |
| for (int i = 0; i < n_final; i++) attn_w[i] *= inv_sum; |
|
|
| float *out_h = attn_out + h * head_dim; |
| for (int d = 0; d < head_dim; d++) out_h[d] = 0.0f; |
| for (int i = 0; i < n_final; i++) { |
| const float *V_h; |
| if (mem_flag[i]) V_h = cctx_v + (size_t)pos_list[i] * n_embd + h * head_dim; |
| else { |
| int j = pos_list[i]; |
| int phys_j = j % n_ctx; |
| V_h = v_cache_layer + (size_t)phys_j * n_embd + h * head_dim; |
| } |
| float w = attn_w[i]; |
| for (int d = 0; d < head_dim; d++) out_h[d] += w * V_h[d]; |
| } |
| } |
|
|
| if (n_total > 8192) { free(scores); free(attn_w); } |
| if (n_total > 8192) { free(pos_list); free(mem_flag); } |
| } |
|
|
| |
| void attention_backward_sliding(float *grad_qkv, const float *grad_attn_out, |
| const float *qkv, int n_embd, int n_head, |
| int seq_pos, |
| const float *k_cache_layer, const float *v_cache_layer, |
| int n_ctx, int window_size, int n_sinks) { |
| int head_dim = n_embd / n_head; |
| float scale = 1.0f / sqrtf((float)head_dim); |
|
|
| |
| int n_sink = (seq_pos < n_sinks) ? seq_pos : n_sinks; |
| int win_start = seq_pos - window_size + 1; |
| if (win_start < n_sinks) win_start = n_sinks; |
| int n_win = seq_pos - win_start + 1; |
| if (n_win < 0) n_win = 0; |
| int n_attend = n_sink + n_win; |
|
|
| const float *Q = qkv; |
| float *gQ = grad_qkv; |
| float *gK = grad_qkv + n_embd; |
| float *gV = grad_qkv + 2 * n_embd; |
| memset(grad_qkv, 0, 3 * n_embd * sizeof(float)); |
|
|
| |
| static __thread int *tl_pos_list = NULL; |
| static __thread float *tl_scores = NULL; |
| static __thread float *tl_w = NULL; |
| static __thread float *tl_gw = NULL; |
| static __thread int tl_n = 0; |
| if (tl_n < n_attend) { |
| free(tl_pos_list); free(tl_scores); free(tl_w); free(tl_gw); |
| tl_pos_list = (int *)malloc(n_attend * sizeof(int)); |
| tl_scores = (float *)malloc(n_attend * sizeof(float)); |
| tl_w = (float *)malloc(n_attend * sizeof(float)); |
| tl_gw = (float *)malloc(n_attend * sizeof(float)); |
| tl_n = n_attend; |
| } |
| int *pos_list = tl_pos_list; |
| float *scores = tl_scores; |
| float *w = tl_w; |
| float *g_w = tl_gw; |
|
|
| int idx = 0; |
| for (int j = 0; j < n_sink; j++) pos_list[idx++] = j; |
| for (int j = win_start; j <= seq_pos; j++) pos_list[idx++] = j; |
|
|
| int have_self = 0; |
| int self_idx = -1; |
| for (int i = 0; i < n_attend; i++) { |
| if (pos_list[i] == seq_pos) { have_self = 1; self_idx = i; break; } |
| } |
|
|
| |
| if (n_attend > 256) { |
| #pragma omp parallel for schedule(static) |
| for (int h = 0; h < n_head; h++) { |
| const float *Q_h = Q + h * head_dim; |
| const float *g_out_h = grad_attn_out + h * head_dim; |
|
|
| |
| float max_score = -1e30f; |
| for (int i = 0; i < n_attend; i++) { |
| int j = pos_list[i]; |
| int phys_j = j % n_ctx; |
| const float *K_jh = k_cache_layer + (size_t)phys_j * n_embd + h * head_dim; |
| float dot = 0.0f; |
| for (int d = 0; d + 7 < head_dim; d += 8) |
| dot += Q_h[d]*K_jh[d] + Q_h[d+1]*K_jh[d+1] + Q_h[d+2]*K_jh[d+2] + Q_h[d+3]*K_jh[d+3] |
| + Q_h[d+4]*K_jh[d+4] + Q_h[d+5]*K_jh[d+5] + Q_h[d+6]*K_jh[d+6] + Q_h[d+7]*K_jh[d+7]; |
| for (int d = (head_dim/8)*8; d < head_dim; d++) dot += Q_h[d] * K_jh[d]; |
| dot *= scale; |
| scores[i] = dot; |
| if (dot > max_score) max_score = dot; |
| } |
| float sum_exp = 0.0f; |
| for (int i = 0; i < n_attend; i++) { |
| float e = expf(scores[i] - max_score); |
| w[i] = e; sum_exp += e; |
| } |
| float inv = 1.0f / (sum_exp + 1e-12f); |
| for (int i = 0; i < n_attend; i++) w[i] *= inv; |
|
|
| |
| float dot_gw_w = 0.0f; |
| for (int i = 0; i < n_attend; i++) { |
| int j = pos_list[i]; |
| int phys_j = j % n_ctx; |
| const float *V_jh = v_cache_layer + (size_t)phys_j * n_embd + h * head_dim; |
| float s = 0.0f; |
| for (int d = 0; d + 7 < head_dim; d += 8) |
| s += g_out_h[d]*V_jh[d] + g_out_h[d+1]*V_jh[d+1] + g_out_h[d+2]*V_jh[d+2] + g_out_h[d+3]*V_jh[d+3] |
| + g_out_h[d+4]*V_jh[d+4] + g_out_h[d+5]*V_jh[d+5] + g_out_h[d+6]*V_jh[d+6] + g_out_h[d+7]*V_jh[d+7]; |
| for (int d = (head_dim/8)*8; d < head_dim; d++) s += g_out_h[d] * V_jh[d]; |
| g_w[i] = s; |
| dot_gw_w += w[i] * s; |
| } |
| for (int i = 0; i < n_attend; i++) g_w[i] = w[i] * (g_w[i] - dot_gw_w); |
|
|
| |
| float *gQ_h = gQ + h * head_dim; |
| for (int d = 0; d < head_dim; d++) { |
| float s = 0.0f; |
| for (int i = 0; i < n_attend; i++) { |
| int j = pos_list[i]; |
| int phys_j = j % n_ctx; |
| const float *K_jh = k_cache_layer + (size_t)phys_j * n_embd + h * head_dim; |
| s += g_w[i] * K_jh[d]; |
| } |
| gQ_h[d] += s * scale; |
| } |
|
|
| if (have_self) { |
| float gs_cur = g_w[self_idx] * scale; |
| float *gK_h = gK + h * head_dim; |
| for (int d = 0; d < head_dim; d++) gK_h[d] += gs_cur * Q_h[d]; |
| float w_cur = w[self_idx]; |
| float *gV_h = gV + h * head_dim; |
| for (int d = 0; d < head_dim; d++) gV_h[d] += w_cur * g_out_h[d]; |
| } |
| } |
| } else { |
| |
| for (int h = 0; h < n_head; h++) { |
| const float *Q_h = Q + h * head_dim; |
| const float *g_out_h = grad_attn_out + h * head_dim; |
| float max_score = -1e30f; |
| for (int i = 0; i < n_attend; i++) { |
| int j = pos_list[i]; |
| int phys_j = j % n_ctx; |
| const float *K_jh = k_cache_layer + (size_t)phys_j * n_embd + h * head_dim; |
| float dot = 0.0f; |
| for (int d = 0; d + 7 < head_dim; d += 8) |
| dot += Q_h[d]*K_jh[d] + Q_h[d+1]*K_jh[d+1] + Q_h[d+2]*K_jh[d+2] + Q_h[d+3]*K_jh[d+3] |
| + Q_h[d+4]*K_jh[d+4] + Q_h[d+5]*K_jh[d+5] + Q_h[d+6]*K_jh[d+6] + Q_h[d+7]*K_jh[d+7]; |
| for (int d = (head_dim/8)*8; d < head_dim; d++) dot += Q_h[d] * K_jh[d]; |
| dot *= scale; |
| scores[i] = dot; |
| if (dot > max_score) max_score = dot; |
| } |
| float sum_exp = 0.0f; |
| for (int i = 0; i < n_attend; i++) { |
| float e = expf(scores[i] - max_score); |
| w[i] = e; sum_exp += e; |
| } |
| float inv = 1.0f / (sum_exp + 1e-12f); |
| for (int i = 0; i < n_attend; i++) w[i] *= inv; |
| float dot_gw_w = 0.0f; |
| for (int i = 0; i < n_attend; i++) { |
| int j = pos_list[i]; |
| int phys_j = j % n_ctx; |
| const float *V_jh = v_cache_layer + (size_t)phys_j * n_embd + h * head_dim; |
| float s = 0.0f; |
| for (int d = 0; d + 7 < head_dim; d += 8) |
| s += g_out_h[d]*V_jh[d] + g_out_h[d+1]*V_jh[d+1] + g_out_h[d+2]*V_jh[d+2] + g_out_h[d+3]*V_jh[d+3] |
| + g_out_h[d+4]*V_jh[d+4] + g_out_h[d+5]*V_jh[d+5] + g_out_h[d+6]*V_jh[d+6] + g_out_h[d+7]*V_jh[d+7]; |
| for (int d = (head_dim/8)*8; d < head_dim; d++) s += g_out_h[d] * V_jh[d]; |
| g_w[i] = s; |
| dot_gw_w += w[i] * s; |
| } |
| for (int i = 0; i < n_attend; i++) g_w[i] = w[i] * (g_w[i] - dot_gw_w); |
| float *gQ_h = gQ + h * head_dim; |
| for (int d = 0; d < head_dim; d++) { |
| float s = 0.0f; |
| for (int i = 0; i < n_attend; i++) { |
| int j = pos_list[i]; |
| int phys_j = j % n_ctx; |
| const float *K_jh = k_cache_layer + (size_t)phys_j * n_embd + h * head_dim; |
| s += g_w[i] * K_jh[d]; |
| } |
| gQ_h[d] += s * scale; |
| } |
| if (have_self) { |
| float gs_cur = g_w[self_idx] * scale; |
| float *gK_h = gK + h * head_dim; |
| for (int d = 0; d < head_dim; d++) gK_h[d] += gs_cur * Q_h[d]; |
| float w_cur = w[self_idx]; |
| float *gV_h = gV + h * head_dim; |
| for (int d = 0; d < head_dim; d++) gV_h[d] += w_cur * g_out_h[d]; |
| } |
| } |
| } |
| } |
|
|
| void trans_layer_forward_sliding(float *x, TransLayer *tl, TransAct *act, |
| ModelConfig *cfg, int cache_pos, int abs_pos, |
| int window, int n_sinks, int n_ctx); |
|
|
| |
| |
| |
| |
| void trans_layer_forward_kv_only_sliding(float *x, TransLayer *tl, TransAct *act, |
| ModelConfig *cfg, int cache_pos, int abs_pos) { |
| int n = cfg->n_embd; |
| act->seq_pos = abs_pos; |
| act->n_ctx = cfg->n_ctx; |
|
|
| |
| norm_forward(act->norm1_out, x, tl->norm1_w, tl->norm1_b, cfg->norm_type, n); |
|
|
| if (g_skip_wv) { |
| if (tl->_kv_k && tl->_kv_v) { |
| memset(tl->_kv_k + (size_t)cache_pos * n, 0, n * sizeof(float)); |
| memset(tl->_kv_v + (size_t)cache_pos * n, 0, n * sizeof(float)); |
| } |
| return; |
| } |
|
|
| |
| |
| if (cfg->qkv_merged) { |
| bin_fwd(act->q, act->norm1_out, &tl->attn_q); |
| act->k = act->q + n; |
| act->v = act->q + 2 * n; |
| } else { |
| bin_fwd(act->k, act->norm1_out, &tl->attn_k); |
| bin_fwd(act->v, act->norm1_out, &tl->attn_v); |
| } |
| if (tl->_kv_k && tl->_kv_v) { |
| memcpy(tl->_kv_k + (size_t)cache_pos * n, act->k, n * sizeof(float)); |
| memcpy(tl->_kv_v + (size_t)cache_pos * n, act->v, n * sizeof(float)); |
| } |
| (void)x; |
| } |
|
|
| |
| void trans_layer_forward_sliding(float *x, TransLayer *tl, TransAct *act, |
| ModelConfig *cfg, int cache_pos, int abs_pos, |
| int window, int n_sinks, int n_ctx) { |
| int n = cfg->n_embd, m = cfg->mlp_dim; |
| float rs = cfg->residual_scale; |
| act->seq_pos = abs_pos; |
| act->n_ctx = n_ctx; |
|
|
| |
| memcpy(act->x_pre_norm1, x, n * sizeof(float)); |
| norm_forward(act->norm1_out, x, tl->norm1_w, tl->norm1_b, cfg->norm_type, n); |
| compute_mean_std(act->x_pre_norm1, n, &act->norm1_cache[0], &act->norm1_cache[1]); |
|
|
| |
| if (g_skip_wv) { |
| memcpy(act->attn_out, act->norm1_out, n * sizeof(float)); |
| memset(act->q, 0, 3 * n * sizeof(float)); |
| } else { |
| if (cfg->qkv_merged) { |
| bin_fwd(act->q, act->norm1_out, &tl->attn_q); |
| act->k = act->q + n; |
| act->v = act->q + 2 * n; |
| } else { |
| bin_fwd(act->q, act->norm1_out, &tl->attn_q); |
| bin_fwd(act->k, act->norm1_out, &tl->attn_k); |
| bin_fwd(act->v, act->norm1_out, &tl->attn_v); |
| } |
|
|
| |
| if (cfg->attn_type == ATTN_ROPE) |
| apply_rope(act->q, act->k, abs_pos, cfg->n_head, n / cfg->n_head, n); |
|
|
| |
| if (tl->_kv_k && tl->_kv_v) { |
| |
| |
| if (g_cctx_cfg.enable && g_runtime_cg && g_sctx.cctx_k) { |
| int layer = tl->layer_idx; |
| const float *ck = g_sctx.cctx_k + (size_t)layer * LCTX_SLOTS * n; |
| const float *cv = g_sctx.cctx_v + (size_t)layer * LCTX_SLOTS * n; |
| const int *cc = g_sctx.cctx_cnt + (size_t)layer * LCTX_SLOTS; |
| const int *ca = g_sctx.cctx_anchor + (size_t)layer * LCTX_SLOTS; |
| attention_forward_concept_ctx(act->attn_out, act->q, n, cfg->n_head, |
| cache_pos, tl->_kv_k, tl->_kv_v, |
| cfg->n_ctx, window, n_sinks, |
| NULL, ck, cv, cc, ca, LCTX_SLOTS, g_cctx_cfg.mem_scale); |
| } else if (g_concept_attn_cfg.enable && g_messenger_caches) { |
| |
| attention_forward_concept(act->attn_out, act->q, n, cfg->n_head, |
| abs_pos, tl->_kv_k, tl->_kv_v, n_ctx, |
| &g_concept_attn_cfg, &g_messenger_caches[tl->layer_idx]); |
| } else { |
| attention_forward_sliding(act->attn_out, act->q, n, cfg->n_head, |
| cache_pos, tl->_kv_k, tl->_kv_v, |
| cfg->n_ctx, window, n_sinks); |
| } |
| } else { |
| |
| memcpy(act->attn_out, act->v, n * sizeof(float)); |
| } |
| } |
|
|
| |
| bin_fwd(act->proj_out, act->attn_out, &tl->attn_o); |
| |
| { |
| float xn = 0, pn = 0; |
| for (int i = 0; i < n; i++) { xn += x[i] * x[i]; pn += act->proj_out[i] * act->proj_out[i]; } |
| xn = sqrtf(xn) + 1e-8f; |
| pn = sqrtf(pn) + 1e-8f; |
| float target = g_attn_res_scale * xn; |
| act->attn_scale = target / pn; |
| for (int i = 0; i < n; i++) act->proj_out[i] *= act->attn_scale; |
| } |
| for (int i = 0; i < n; i++) x[i] += rs * act->proj_out[i]; |
|
|
| |
| memcpy(act->x_pre_norm2, x, n * sizeof(float)); |
| norm_forward(act->norm2_out, x, tl->norm2_w, tl->norm2_b, cfg->norm_type, n); |
| compute_mean_std(act->x_pre_norm2, n, &act->norm2_cache[0], &act->norm2_cache[1]); |
|
|
| if (cfg->act_type == ACT_SWIGLU) { |
| |
| static float *sgate = NULL, *sup = NULL; |
| static int sg_m = 0; |
| if (sg_m != m) { |
| free(sgate); free(sup); |
| sgate = malloc(m * sizeof(float)); |
| sup = malloc(m * sizeof(float)); |
| sg_m = m; |
| } |
| bin_fwd(sgate, act->norm2_out, &tl->mlp_gate); |
| bin_fwd(sup, act->norm2_out, &tl->mlp_up); |
| for (int i = 0; i < m; i++) act->mlp_hidden[i] = silu(sgate[i]) * sup[i]; |
| } else { |
| bin_fwd(act->mlp_hidden, act->norm2_out, &tl->mlp_gate); |
| for (int i = 0; i < m; i++) act->mlp_hidden[i] = gelu(act->mlp_hidden[i]); |
| } |
|
|
| bin_fwd(act->mlp_out, act->mlp_hidden, &tl->mlp_down); |
| |
| { |
| float xn = 0, mlp_norm_sq = 0; |
| for (int i = 0; i < n; i++) { xn += x[i] * x[i]; mlp_norm_sq += act->mlp_out[i] * act->mlp_out[i]; } |
| float xn_norm = sqrtf(xn) + 1e-8f; |
| float mlp_norm = sqrtf(mlp_norm_sq) + 1e-8f; |
| float mlp_cap = 0.25f * xn_norm; |
| act->mlp_scale = (mlp_norm > mlp_cap) ? (mlp_cap / mlp_norm) : 1.0f; |
| for (int i = 0; i < n; i++) x[i] += rs * act->mlp_scale * act->mlp_out[i]; |
| } |
| normalize_residual(x, n, 6.0f); |
| } |
|
|
| |
| |
| |
| |
| |
| void trans_layer_backward_sliding(float *grad_x, TransLayer *tl, TransAct *act, |
| ModelConfig *cfg, int window, int n_sinks, |
| float lr) { |
| int n = cfg->n_embd, m = cfg->mlp_dim; |
| float rs = cfg->residual_scale; |
| int tid = g_cur_tid; |
| float *g_mlp = g_thr[tid].mlp, *g_hidden = g_thr[tid].hidden, |
| *g_norm2 = g_thr[tid].norm2, *g_proj = g_thr[tid].proj; |
| float *g_attn = g_thr[tid].attn, *g_qkv = g_thr[tid].qkv, |
| *g_norm1 = g_thr[tid].norm1, *g_pre = g_thr[tid].pre; |
|
|
| #define BIN_BW(gx, gy, x, bl, lr, slot) \ |
| (g_use_ste ? bin_backward_ste(gx, gy, x, bl, lr, tl->layer_idx, slot) \ |
| : bin_backward(gx, gy, x, bl, lr)) |
|
|
| |
| for (int i = 0; i < n; i++) g_mlp[i] = grad_x[i] * rs * act->mlp_scale; |
| BIN_BW(g_hidden, g_mlp, act->mlp_hidden, &tl->mlp_down, lr, 3); |
| if (cfg->act_type == ACT_SWIGLU) { |
| float *g_gate = g_thr[tid].gate, *g_up = g_thr[tid].up; |
| float *g_norm2_gate = g_thr[tid].norm2_gate, *g_norm2_up = g_thr[tid].norm2_up; |
| for (int i = 0; i < m; i++) { |
| float sv = silu(act->swiglu_gate[i]); |
| float sg = silu_grad(act->swiglu_gate[i]); |
| g_gate[i] = g_hidden[i] * sg * act->swiglu_up[i]; |
| g_up[i] = g_hidden[i] * sv; |
| } |
| BIN_BW(g_norm2_gate, g_gate, act->norm2_out, &tl->mlp_gate, lr, 2); |
| BIN_BW(g_norm2_up, g_up, act->norm2_out, &tl->mlp_up, lr, 6); |
| for (int i = 0; i < n; i++) g_norm2[i] = g_norm2_gate[i] + g_norm2_up[i]; |
| } else { |
| for (int i = 0; i < m; i++) g_hidden[i] *= gelu_grad(act->mlp_hidden[i]); |
| BIN_BW(g_norm2, g_hidden, act->norm2_out, &tl->mlp_gate, lr, 2); |
| } |
| norm_backward(g_pre, g_norm2, act->x_pre_norm2, tl->norm2_w, |
| act->norm2_cache, cfg->norm_type, n, |
| g_thr[tid].grad_norm2_w[tl->layer_idx], g_thr[tid].grad_norm2_b[tl->layer_idx]); |
| for (int i = 0; i < n; i++) grad_x[i] += g_pre[i] * rs; |
|
|
| |
| for (int i = 0; i < n; i++) g_proj[i] = grad_x[i] * rs * act->attn_scale; |
| BIN_BW(g_attn, g_proj, act->attn_out, &tl->attn_o, lr, 1); |
| if (g_skip_wv) { |
| memcpy(g_norm1, g_attn, n * sizeof(float)); |
| memset(g_qkv, 0, 3 * n * sizeof(float)); |
| } else if (g_use_real_attention && tl->_kv_k && tl->_kv_v) { |
| |
| |
| attention_backward_sliding(g_qkv, g_attn, act->q, n, cfg->n_head, |
| act->seq_pos, tl->_kv_k, tl->_kv_v, |
| act->n_ctx, window, n_sinks); |
| } else { |
| memset(g_qkv, 0, 3 * n * sizeof(float)); |
| memcpy(g_qkv + 2 * n, g_attn, n * sizeof(float)); |
| } |
| if (!g_skip_wv) { |
| if (cfg->qkv_merged) { |
| BIN_BW(g_norm1, g_qkv, act->norm1_out, &tl->attn_q, lr, 0); |
| } else { |
| float *g_n1k = g_thr[tid].n1k, *g_n1v = g_thr[tid].n1v; |
| BIN_BW(g_norm1, g_qkv, act->norm1_out, &tl->attn_q, lr, 0); |
| BIN_BW(g_n1k, g_qkv + n, act->norm1_out, &tl->attn_k, lr, 4); |
| BIN_BW(g_n1v, g_qkv + 2*n, act->norm1_out, &tl->attn_v, lr, 5); |
| for (int i = 0; i < n; i++) g_norm1[i] += g_n1k[i] + g_n1v[i]; |
| } |
| } |
| norm_backward(g_pre, g_norm1, act->x_pre_norm1, tl->norm1_w, |
| act->norm1_cache, cfg->norm_type, n, |
| g_thr[tid].grad_norm1_w[tl->layer_idx], g_thr[tid].grad_norm1_b[tl->layer_idx]); |
| for (int i = 0; i < n; i++) grad_x[i] += g_pre[i] * rs; |
| } |
|
|
| |
|
|
| |
| |
| |
|
|
| void model_batch_alloc(Model *m) { |
| |
| |
| for (int l = 0; l < m->cfg.n_layer; l++) { |
| TransLayer *tl = &m->layers[l]; |
| BinLayer *bls[8] = {&tl->attn_q, &tl->attn_o, &tl->mlp_gate, &tl->mlp_down}; |
| int n_bl = 4; |
| if (!m->cfg.qkv_merged) { |
| |
| bls[4] = &tl->attn_k; bls[5] = &tl->attn_v; |
| n_bl = 6; |
| } |
| if (m->cfg.act_type == ACT_SWIGLU) { |
| bls[n_bl] = &tl->mlp_up; |
| n_bl++; |
| } |
| for (int b = 0; b < n_bl; b++) { |
| BinLayer *bl = bls[b]; |
| if (!bl->grad_accum && bl->w_float) { |
| bl->grad_accum = calloc((size_t)bl->in_dim * bl->out_dim, sizeof(float)); |
| } |
| if (!bl->bias_grad_accum && bl->w_float) { |
| bl->bias_grad_accum = calloc((size_t)bl->out_dim, sizeof(float)); |
| } |
| } |
| } |
|
|
| |
| if (!m->grad_wte_accum) { |
| size_t wte_size = (size_t)m->cfg.vocab_size * m->cfg.n_embd; |
| |
| #if defined(_WIN32) |
| m->grad_wte_accum = _aligned_malloc(wte_size * sizeof(float), 64); |
| memset(m->grad_wte_accum, 0, wte_size * sizeof(float)); |
| m->m_wte = _aligned_malloc(wte_size * sizeof(float), 64); |
| memset(m->m_wte, 0, wte_size * sizeof(float)); |
| m->v_wte = _aligned_malloc(wte_size * sizeof(float), 64); |
| memset(m->v_wte, 0, wte_size * sizeof(float)); |
| #else |
| posix_memalign((void**)&m->grad_wte_accum, 64, wte_size * sizeof(float)); |
| memset(m->grad_wte_accum, 0, wte_size * sizeof(float)); |
| posix_memalign((void**)&m->m_wte, 64, wte_size * sizeof(float)); |
| memset(m->m_wte, 0, wte_size * sizeof(float)); |
| posix_memalign((void**)&m->v_wte, 64, wte_size * sizeof(float)); |
| memset(m->v_wte, 0, wte_size * sizeof(float)); |
| #endif |
| } |
| if (m->wpe && !m->grad_wpe_accum) { |
| size_t wpe_size = (size_t)m->cfg.n_ctx * m->cfg.n_embd; |
| m->grad_wpe_accum = calloc(wpe_size, sizeof(float)); |
| m->m_wpe = calloc(wpe_size, sizeof(float)); |
| m->v_wpe = calloc(wpe_size, sizeof(float)); |
| } |
| if (!m->grad_ln_f_w_accum) { |
| m->grad_ln_f_w_accum = calloc(m->cfg.n_embd, sizeof(float)); |
| m->grad_ln_f_b_accum = calloc(m->cfg.n_embd, sizeof(float)); |
| m->m_ln_f_w = calloc(m->cfg.n_embd, sizeof(float)); |
| m->v_ln_f_w = calloc(m->cfg.n_embd, sizeof(float)); |
| m->m_ln_f_b = calloc(m->cfg.n_embd, sizeof(float)); |
| m->v_ln_f_b = calloc(m->cfg.n_embd, sizeof(float)); |
| } |
|
|
| |
| for (int l = 0; l < m->cfg.n_layer; l++) { |
| TransLayer *tl = &m->layers[l]; |
| if (!tl->grad_norm1_w) { |
| tl->grad_norm1_w = calloc(m->cfg.n_embd, sizeof(float)); |
| tl->grad_norm1_b = calloc(m->cfg.n_embd, sizeof(float)); |
| tl->grad_norm2_w = calloc(m->cfg.n_embd, sizeof(float)); |
| tl->grad_norm2_b = calloc(m->cfg.n_embd, sizeof(float)); |
| } |
| |
| if (!tl->m_norm1_w) { |
| tl->m_norm1_w = calloc(m->cfg.n_embd, sizeof(float)); |
| tl->v_norm1_w = calloc(m->cfg.n_embd, sizeof(float)); |
| tl->m_norm1_b = calloc(m->cfg.n_embd, sizeof(float)); |
| tl->v_norm1_b = calloc(m->cfg.n_embd, sizeof(float)); |
| tl->m_norm2_w = calloc(m->cfg.n_embd, sizeof(float)); |
| tl->v_norm2_w = calloc(m->cfg.n_embd, sizeof(float)); |
| tl->m_norm2_b = calloc(m->cfg.n_embd, sizeof(float)); |
| tl->v_norm2_b = calloc(m->cfg.n_embd, sizeof(float)); |
| } |
| } |
| thr_res_alloc(m); |
| } |
|
|
| |
| void thr_res_alloc(Model *m) { |
| g_nthr = omp_get_max_threads(); |
| if (g_nthr > LAL_MAX_THREADS) g_nthr = LAL_MAX_THREADS; |
| if (g_thr_inited) return; |
| for (int t = 0; t < g_nthr; t++) { |
| ThrRes *r = &g_thr[t]; |
| r->n_layer = m->cfg.n_layer; |
| r->acts = trans_act_alloc(&m->cfg); |
| r->scratch = trans_act_alloc(&m->cfg); |
| r->mlp = calloc(16384, sizeof(float)); |
| r->hidden = calloc(16384, sizeof(float)); |
| r->norm2 = calloc(16384, sizeof(float)); |
| r->proj = calloc(16384, sizeof(float)); |
| r->attn = calloc(16384, sizeof(float)); |
| r->qkv = calloc(16384*3, sizeof(float)); |
| r->norm1 = calloc(16384, sizeof(float)); |
| r->pre = calloc(16384, sizeof(float)); |
| r->gate = calloc(16384, sizeof(float)); |
| r->up = calloc(16384, sizeof(float)); |
| r->norm2_gate = calloc(16384, sizeof(float)); |
| r->norm2_up = calloc(16384, sizeof(float)); |
| r->n1k = calloc(4096, sizeof(float)); |
| r->n1v = calloc(4096, sizeof(float)); |
| r->xc = calloc(4096, sizeof(float)); |
| r->x = calloc(4096, sizeof(float)); |
| r->gh = calloc(4096, sizeof(float)); |
| r->g_pre4 = calloc(4096, sizeof(float)); |
| r->full_logits = calloc((size_t)m->cfg.vocab_size, sizeof(float)); |
| r->full_logits_vocab = m->cfg.vocab_size; |
| r->forward_done = 0; |
| r->x_before_final = calloc(m->cfg.n_embd, sizeof(float)); |
| r->final_ln = calloc(m->cfg.n_embd, sizeof(float)); |
| r->n_bl_max = 7; |
| r->grad_w = calloc(m->cfg.n_layer, sizeof(float**)); |
| r->grad_b = calloc(m->cfg.n_layer, sizeof(float**)); |
| for (int l = 0; l < m->cfg.n_layer; l++) { |
| TransLayer *tl = &m->layers[l]; |
| BinLayer *bls[8] = {&tl->attn_q, &tl->attn_o, &tl->mlp_gate, &tl->mlp_down}; |
| int n_bl = 4; |
| if (!m->cfg.qkv_merged) { bls[4] = &tl->attn_k; bls[5] = &tl->attn_v; n_bl = 6; } |
| if (m->cfg.act_type == ACT_SWIGLU) { bls[n_bl] = &tl->mlp_up; n_bl++; } |
| |
| |
| r->grad_w[l] = calloc(r->n_bl_max, sizeof(float*)); |
| r->grad_b[l] = calloc(r->n_bl_max, sizeof(float*)); |
| for (int b = 0; b < n_bl; b++) { |
| BinLayer *bl = bls[b]; |
| r->grad_w[l][b] = calloc((size_t)bl->in_dim * bl->out_dim, sizeof(float)); |
| r->grad_b[l][b] = calloc(bl->out_dim, sizeof(float)); |
| } |
| } |
| size_t wte_size = (size_t)m->cfg.vocab_size * m->cfg.n_embd; |
| r->grad_wte = calloc(wte_size, sizeof(float)); |
| size_t wpe_size = (size_t)m->cfg.n_ctx * m->cfg.n_embd; |
| r->grad_wpe = calloc(wpe_size, sizeof(float)); |
| r->grad_lnfw = calloc(m->cfg.n_embd, sizeof(float)); |
| r->grad_lnfb = calloc(m->cfg.n_embd, sizeof(float)); |
| r->grad_norm1_w = calloc(m->cfg.n_layer, sizeof(float*)); |
| r->grad_norm1_b = calloc(m->cfg.n_layer, sizeof(float*)); |
| r->grad_norm2_w = calloc(m->cfg.n_layer, sizeof(float*)); |
| r->grad_norm2_b = calloc(m->cfg.n_layer, sizeof(float*)); |
| for (int l = 0; l < m->cfg.n_layer; l++) { |
| r->grad_norm1_w[l] = calloc(m->cfg.n_embd, sizeof(float)); |
| r->grad_norm1_b[l] = calloc(m->cfg.n_embd, sizeof(float)); |
| r->grad_norm2_w[l] = calloc(m->cfg.n_embd, sizeof(float)); |
| r->grad_norm2_b[l] = calloc(m->cfg.n_embd, sizeof(float)); |
| } |
| |
| r->ponder_ready = 0; |
| if (g_ponder_cfg.enable) { |
| memset(&r->ponder, 0, sizeof(r->ponder)); |
| r->ponder_mix = calloc(m->cfg.n_embd, sizeof(float)); |
| r->ponder_state = calloc((size_t)LAL_PONDER_MAX_STEPS * m->cfg.n_embd, sizeof(float)); |
| r->ponder_kv0k = calloc(m->cfg.n_embd, sizeof(float)); |
| r->ponder_kv0v = calloc(m->cfg.n_embd, sizeof(float)); |
| r->rec_acts = trans_act_alloc(&m->cfg); |
| r->ponder_first_rec_step = g_ponder_cfg.layer_halt ? (m->cfg.n_layer - 1) : 0; |
| r->ponder_ready = 1; |
| } |
| } |
| g_thr_inited = 1; |
| } |
|
|
| void thr_res_free(void) { |
| if (!g_thr_inited) return; |
| for (int t = 0; t < g_nthr; t++) { |
| ThrRes *r = &g_thr[t]; |
| trans_act_free(r->acts, r->n_layer); |
| trans_act_free(r->scratch, r->n_layer); |
| free(r->mlp); free(r->hidden); free(r->norm2); free(r->proj); |
| free(r->attn); free(r->qkv); free(r->norm1); free(r->pre); |
| free(r->gate); free(r->up); free(r->norm2_gate); free(r->norm2_up); |
| free(r->n1k); free(r->n1v); free(r->xc); free(r->x); free(r->gh); |
| free(r->g_pre4); free(r->full_logits); free(r->x_before_final); free(r->final_ln); |
| for (int l = 0; l < r->n_layer; l++) { |
| for (int b = 0; b < r->n_bl_max; b++) { |
| if (r->grad_w[l]) free(r->grad_w[l][b]); |
| if (r->grad_b[l]) free(r->grad_b[l][b]); |
| } |
| free(r->grad_w[l]); free(r->grad_b[l]); |
| } |
| free(r->grad_w); free(r->grad_b); |
| free(r->grad_wte); free(r->grad_wpe); free(r->grad_lnfw); free(r->grad_lnfb); |
| for (int l = 0; l < r->n_layer; l++) { |
| free(r->grad_norm1_w[l]); free(r->grad_norm1_b[l]); |
| free(r->grad_norm2_w[l]); free(r->grad_norm2_b[l]); |
| } |
| free(r->grad_norm1_w); free(r->grad_norm1_b); |
| free(r->grad_norm2_w); free(r->grad_norm2_b); |
| |
| if (r->ponder_ready) { |
| free(r->ponder_mix); free(r->ponder_state); |
| free(r->ponder_kv0k); free(r->ponder_kv0v); |
| trans_act_free(r->rec_acts, r->n_layer); |
| r->ponder_ready = 0; |
| } |
| } |
| g_thr_inited = 0; |
| } |
|
|
| |
| |
| |
| void thr_grad_reduce(Model *m) { |
| int n = m->cfg.n_embd; |
| for (int t = 0; t < g_nthr; t++) { |
| ThrRes *r = &g_thr[t]; |
| for (int l = 0; l < m->cfg.n_layer; l++) { |
| TransLayer *tl = &m->layers[l]; |
| BinLayer *bls[8] = {&tl->attn_q, &tl->attn_o, &tl->mlp_gate, &tl->mlp_down}; |
| int n_bl = 4; |
| if (!m->cfg.qkv_merged) { bls[4] = &tl->attn_k; bls[5] = &tl->attn_v; n_bl = 6; } |
| if (m->cfg.act_type == ACT_SWIGLU) { bls[n_bl] = &tl->mlp_up; n_bl++; } |
| for (int b = 0; b < n_bl; b++) { |
| BinLayer *bl = bls[b]; |
| long sz = (long)bl->in_dim * bl->out_dim; |
| if (bl->grad_accum && r->grad_w[l][b]) { |
| |
| if (sz > 4096) { |
| #pragma omp parallel for schedule(static) |
| for (long i = 0; i < sz; i++) bl->grad_accum[i] += r->grad_w[l][b][i]; |
| } else { |
| for (long i = 0; i < sz; i++) bl->grad_accum[i] += r->grad_w[l][b][i]; |
| } |
| } |
| long osz = bl->out_dim; |
| if (bl->bias_grad_accum && r->grad_b[l][b]) |
| for (long i = 0; i < osz; i++) bl->bias_grad_accum[i] += r->grad_b[l][b][i]; |
| } |
| for (int i = 0; i < n; i++) { |
| tl->grad_norm1_w[i] += r->grad_norm1_w[l][i]; |
| tl->grad_norm1_b[i] += r->grad_norm1_b[l][i]; |
| tl->grad_norm2_w[i] += r->grad_norm2_w[l][i]; |
| tl->grad_norm2_b[i] += r->grad_norm2_b[l][i]; |
| } |
| } |
| |
| if (m->grad_wte_accum) { |
| size_t wte_size = (size_t)m->cfg.vocab_size * m->cfg.n_embd; |
| #pragma omp parallel for schedule(static) |
| for (long i = 0; i < (long)wte_size; i++) m->grad_wte_accum[i] += r->grad_wte[i]; |
| } |
| if (m->grad_wpe_accum) { |
| size_t wpe_size = (size_t)m->cfg.n_ctx * m->cfg.n_embd; |
| #pragma omp parallel for schedule(static) |
| for (long i = 0; i < (long)wpe_size; i++) m->grad_wpe_accum[i] += r->grad_wpe[i]; |
| } |
| if (m->grad_ln_f_w_accum) { |
| for (int i = 0; i < n; i++) { |
| m->grad_ln_f_w_accum[i] += r->grad_lnfw[i]; |
| m->grad_ln_f_b_accum[i] += r->grad_lnfb[i]; |
| } |
| } |
| } |
| } |
|
|
| void model_batch_begin(Model *m) { |
| for (int l = 0; l < m->cfg.n_layer; l++) { |
| TransLayer *tl = &m->layers[l]; |
| |
| if (tl->grad_norm1_w) { |
| memset(tl->grad_norm1_w, 0, m->cfg.n_embd * sizeof(float)); |
| memset(tl->grad_norm1_b, 0, m->cfg.n_embd * sizeof(float)); |
| memset(tl->grad_norm2_w, 0, m->cfg.n_embd * sizeof(float)); |
| memset(tl->grad_norm2_b, 0, m->cfg.n_embd * sizeof(float)); |
| } |
| BinLayer *bls[8] = {&tl->attn_q, &tl->attn_o, &tl->mlp_gate, &tl->mlp_down}; |
| int n_bl = 4; |
| if (!m->cfg.qkv_merged) { bls[4] = &tl->attn_k; bls[5] = &tl->attn_v; n_bl = 6; } |
| if (m->cfg.act_type == ACT_SWIGLU) { bls[n_bl] = &tl->mlp_up; n_bl++; } |
| for (int b = 0; b < n_bl; b++) { |
| BinLayer *bl = bls[b]; |
| if (bl->grad_accum) |
| memset(bl->grad_accum, 0, (size_t)bl->in_dim * bl->out_dim * sizeof(float)); |
| if (bl->bias_grad_accum) |
| memset(bl->bias_grad_accum, 0, (size_t)bl->out_dim * sizeof(float)); |
| } |
| } |
| |
| if (m->grad_wte_accum) |
| memset(m->grad_wte_accum, 0, (size_t)m->cfg.vocab_size * m->cfg.n_embd * sizeof(float)); |
| if (m->grad_wpe_accum) |
| memset(m->grad_wpe_accum, 0, (size_t)m->cfg.n_ctx * m->cfg.n_embd * sizeof(float)); |
| if (m->grad_ln_f_w_accum) { |
| memset(m->grad_ln_f_w_accum, 0, m->cfg.n_embd * sizeof(float)); |
| memset(m->grad_ln_f_b_accum, 0, m->cfg.n_embd * sizeof(float)); |
| } |
| |
| |
| |
| if (g_thr_inited) { |
| for (int t = 0; t < g_nthr; t++) { |
| ThrRes *r = &g_thr[t]; |
| for (int l = 0; l < m->cfg.n_layer; l++) { |
| TransLayer *tl = &m->layers[l]; |
| int n_bl = 4; |
| if (!m->cfg.qkv_merged) n_bl = 6; |
| if (m->cfg.act_type == ACT_SWIGLU) n_bl++; |
| for (int b = 0; b < n_bl; b++) { |
| BinLayer *bl = NULL; |
| if (b == 0) bl = &tl->attn_q; else if (b == 1) bl = &tl->attn_o; |
| else if (b == 2) bl = &tl->mlp_gate; else if (b == 3) bl = &tl->mlp_down; |
| else if (b == 4) bl = &tl->attn_k; else if (b == 5) bl = &tl->attn_v; |
| else if (b == 6) bl = &tl->mlp_up; |
| if (bl && r->grad_w[l][b]) |
| memset(r->grad_w[l][b], 0, (size_t)bl->in_dim * bl->out_dim * sizeof(float)); |
| if (bl && r->grad_b[l][b]) |
| memset(r->grad_b[l][b], 0, (size_t)bl->out_dim * sizeof(float)); |
| } |
| if (r->grad_norm1_w[l]) { |
| memset(r->grad_norm1_w[l], 0, m->cfg.n_embd * sizeof(float)); |
| memset(r->grad_norm1_b[l], 0, m->cfg.n_embd * sizeof(float)); |
| memset(r->grad_norm2_w[l], 0, m->cfg.n_embd * sizeof(float)); |
| memset(r->grad_norm2_b[l], 0, m->cfg.n_embd * sizeof(float)); |
| } |
| } |
| if (r->grad_wte) memset(r->grad_wte, 0, (size_t)m->cfg.vocab_size * m->cfg.n_embd * sizeof(float)); |
| if (r->grad_wpe) memset(r->grad_wpe, 0, (size_t)m->cfg.n_ctx * m->cfg.n_embd * sizeof(float)); |
| if (r->grad_lnfw) { |
| memset(r->grad_lnfw, 0, m->cfg.n_embd * sizeof(float)); |
| memset(r->grad_lnfb, 0, m->cfg.n_embd * sizeof(float)); |
| } |
| } |
| } |
| } |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| |
| |
| |
| static void ponder_train_forward(Model *m, int tid, int cache_pos, int abs_pos, |
| int window, int n_sinks, int ctx) { |
| int n = m->cfg.n_embd; |
| int nL = m->cfg.n_layer; |
| int R = g_ponder_cfg.rec_iters; |
| int last_block = nL - 1; |
| ThrRes *r = &g_thr[tid]; |
| float *x = r->x; |
| PonderBuf *pb = &r->ponder; |
| float *mix = r->ponder_mix; |
| float halts[LAL_PONDER_MAX_STEPS]; |
| int step = 0; |
|
|
| |
| if (g_ponder_cfg.layer_halt) { |
| for (int l = 0; l < last_block; l++) { |
| trans_layer_forward_sliding(x, &m->layers[l], &r->acts[l], &m->cfg, |
| cache_pos, abs_pos, window, n_sinks, ctx); |
| float *s = r->ponder_state + (size_t)step * n; |
| memcpy(s, x, n * sizeof(float)); |
| halts[step] = ponder_halt(&m->ph[l], s, n); |
| step++; |
| } |
| } else { |
| for (int l = 0; l < last_block; l++) |
| trans_layer_forward_sliding(x, &m->layers[l], &r->acts[l], &m->cfg, |
| cache_pos, abs_pos, window, n_sinks, ctx); |
| } |
|
|
| |
| if (R >= 2) { |
| int first_rec = step; |
| for (int it = 0; it < R; it++) { |
| TransAct *act = (it == 0) ? &r->acts[last_block] : &r->rec_acts[it]; |
| trans_layer_forward_sliding(x, &m->layers[last_block], act, &m->cfg, |
| cache_pos, abs_pos, window, n_sinks, ctx); |
| if (it == 0) { |
| |
| memcpy(r->ponder_kv0k, m->k_cache[last_block] + (size_t)cache_pos * n, n * sizeof(float)); |
| memcpy(r->ponder_kv0v, m->v_cache[last_block] + (size_t)cache_pos * n, n * sizeof(float)); |
| } else { |
| |
| memcpy(m->k_cache[last_block] + (size_t)cache_pos * n, r->ponder_kv0k, n * sizeof(float)); |
| memcpy(m->v_cache[last_block] + (size_t)cache_pos * n, r->ponder_kv0v, n * sizeof(float)); |
| } |
| float *s = r->ponder_state + (size_t)step * n; |
| memcpy(s, x, n * sizeof(float)); |
| if (it < R - 1) { |
| halts[step] = ponder_halt(&m->ph_rec, s, n); |
| } |
| step++; |
| } |
| r->ponder_first_rec_step = first_rec; |
| } else { |
| trans_layer_forward_sliding(x, &m->layers[last_block], &r->acts[last_block], &m->cfg, |
| cache_pos, abs_pos, window, n_sinks, ctx); |
| float *s = r->ponder_state + (size_t)step * n; |
| memcpy(s, x, n * sizeof(float)); |
| step++; |
| r->ponder_first_rec_step = step; |
| } |
|
|
| int n_param = step - 1; |
| ponder_dist_fill(pb, halts, n_param); |
|
|
| |
| memset(mix, 0, n * sizeof(float)); |
| for (int s2 = 0; s2 < pb->n_steps; s2++) { |
| const float *s = r->ponder_state + (size_t)s2 * n; |
| float c = pb->c[s2]; |
| for (int i = 0; i < n; i++) mix[i] += c * s[i]; |
| } |
| memcpy(x, mix, n * sizeof(float)); |
|
|
| |
| g_ponder_last_al = pb->loss_al; |
| g_ponder_last_p = pb->loss_p; |
| g_ponder_last_mean = pb->mean_step; |
| } |
|
|
| static void ponder_train_backward(Model *m, int tid, int window, int n_sinks) { |
| int n = m->cfg.n_embd; |
| int nL = m->cfg.n_layer; |
| int last_block = nL - 1; |
| ThrRes *r = &g_thr[tid]; |
| PonderBuf *pb = &r->ponder; |
| float *G = r->gh; |
| float *V = r->g_pre4; |
|
|
| |
| for (int s = 0; s < pb->n_steps; s++) { |
| const float *st = r->ponder_state + (size_t)s * n; |
| float d = 0.0f; |
| for (int i = 0; i < n; i++) d += G[i] * st[i]; |
| pb->gdot[s] = d; |
| } |
| |
| float dpre[LAL_PONDER_MAX_STEPS]; |
| ponder_grad(pb, dpre); |
| for (int s = 0; s < pb->n_param; s++) { |
| PonderLayer *u = (s < r->ponder_first_rec_step) ? &m->ph[s] : &m->ph_rec; |
| const float *st = r->ponder_state + (size_t)s * n; |
| for (int i = 0; i < n; i++) u->grad_w[i] += dpre[s] * st[i]; |
| u->grad_b += dpre[s]; |
| } |
|
|
| |
| |
| memset(V, 0, n * sizeof(float)); |
| for (int s = pb->n_steps - 1; s >= 0; s--) { |
| float c = pb->c[s]; |
| for (int i = 0; i < n; i++) V[i] += c * G[i]; |
| if (s < r->ponder_first_rec_step) { |
| trans_layer_backward_sliding(V, &m->layers[s], &r->acts[s], &m->cfg, |
| window, n_sinks, 0.0f); |
| } else { |
| int it = s - r->ponder_first_rec_step; |
| TransAct *act = (it == 0) ? &r->acts[last_block] : &r->rec_acts[it]; |
| int cp = act->seq_pos % m->cfg.n_ctx; |
| |
| memcpy(m->k_cache[last_block] + (size_t)cp * n, act->k, n * sizeof(float)); |
| memcpy(m->v_cache[last_block] + (size_t)cp * n, act->v, n * sizeof(float)); |
| trans_layer_backward_sliding(V, &m->layers[last_block], act, &m->cfg, |
| window, n_sinks, 0.0f); |
| } |
| } |
| |
| memcpy(r->gh, V, n * sizeof(float)); |
| } |
|
|
| |
| static PonderBuf g_ponder_ibuf; |
| static float *g_ponder_imix = NULL, *g_ponder_ikv0k = NULL, *g_ponder_ikv0v = NULL; |
| static int g_ponder_ibuf_n = 0; |
|
|
| static void ponder_infer_forward(Model *m, int pos, int abs_pos, |
| int window, int n_sinks) { |
| int n = m->cfg.n_embd; |
| int nL = m->cfg.n_layer; |
| int R = g_ponder_cfg.rec_iters; |
| int last_block = nL - 1; |
| float *x = g_sctx.x; |
| if (!g_ponder_imix || g_ponder_ibuf_n != n) { |
| free(g_ponder_imix); free(g_ponder_ikv0k); free(g_ponder_ikv0v); |
| g_ponder_imix = malloc(n * sizeof(float)); |
| g_ponder_ikv0k = malloc(n * sizeof(float)); |
| g_ponder_ikv0v = malloc(n * sizeof(float)); |
| g_ponder_ibuf_n = n; |
| } |
| PonderBuf *pb = &g_ponder_ibuf; |
| memset(pb, 0, sizeof(*pb)); |
| float *mix = g_ponder_imix; |
| memset(mix, 0, n * sizeof(float)); |
| float Rm = 1.0f, ms = 0.0f; |
| int step = 0, early_exit = 0, exited_layer = -1; |
|
|
| if (g_ponder_cfg.layer_halt) { |
| for (int l = 0; l < last_block; l++) { |
| trans_layer_forward_sliding(x, &m->layers[l], &m->acts[l], &m->cfg, |
| pos, abs_pos, window, n_sinks, pos + 1); |
| float p = ponder_halt(&m->ph[l], x, n); |
| float c = Rm * p; |
| for (int i = 0; i < n; i++) mix[i] += c * x[i]; |
| ms += (float)step * c; |
| Rm *= (1.0f - p); |
| pb->c[step] = c; |
| step++; |
| if (step >= g_ponder_cfg.infer_min_layer && Rm <= 1.0f - g_ponder_cfg.threshold) { |
| early_exit = 1; exited_layer = l; break; |
| } |
| } |
| } else { |
| for (int l = 0; l < last_block; l++) |
| trans_layer_forward_sliding(x, &m->layers[l], &m->acts[l], &m->cfg, |
| pos, abs_pos, window, n_sinks, pos + 1); |
| } |
|
|
| if (!early_exit) { |
| if (R >= 2) { |
| for (int it = 0; it < R; it++) { |
| TransAct *act = (it == 0) ? &m->acts[last_block] : &m->rec_acts[it]; |
| trans_layer_forward_sliding(x, &m->layers[last_block], act, &m->cfg, |
| pos, abs_pos, window, n_sinks, pos + 1); |
| if (it == 0) { |
| memcpy(g_ponder_ikv0k, m->k_cache[last_block] + (size_t)pos * n, n * sizeof(float)); |
| memcpy(g_ponder_ikv0v, m->v_cache[last_block] + (size_t)pos * n, n * sizeof(float)); |
| } else { |
| memcpy(m->k_cache[last_block] + (size_t)pos * n, g_ponder_ikv0k, n * sizeof(float)); |
| memcpy(m->v_cache[last_block] + (size_t)pos * n, g_ponder_ikv0v, n * sizeof(float)); |
| } |
| float p = ponder_halt(&m->ph_rec, x, n); |
| float c = Rm * p; |
| for (int i = 0; i < n; i++) mix[i] += c * x[i]; |
| ms += (float)step * c; |
| Rm *= (1.0f - p); |
| pb->c[step] = c; |
| step++; |
| if (step >= g_ponder_cfg.infer_min_layer && Rm <= 1.0f - g_ponder_cfg.threshold) { |
| early_exit = 1; |
| break; |
| } |
| } |
| } else { |
| trans_layer_forward_sliding(x, &m->layers[last_block], &m->acts[last_block], &m->cfg, |
| pos, abs_pos, window, n_sinks, pos + 1); |
| } |
| } |
|
|
| if (early_exit && exited_layer >= 0) { |
| |
| for (int lb = exited_layer + 1; lb <= last_block; lb++) |
| trans_layer_forward_kv_only_sliding(x, &m->layers[lb], &m->acts[lb], &m->cfg, pos, abs_pos); |
| } |
| |
| for (int i = 0; i < n; i++) mix[i] += Rm * x[i]; |
| ms += (float)step * Rm; |
| pb->c[step] = Rm; |
| step++; |
| pb->n_steps = step; |
| pb->n_param = step - 1; |
| pb->mean_step = ms; |
|
|
| memcpy(x, mix, n * sizeof(float)); |
| ponder_stats_record(pb, early_exit); |
| } |
|
|
| float model_forward_sliding(Model *m, const int *tokens, int n_tokens) { |
| int n = m->cfg.n_embd; |
| int nL = m->cfg.n_layer; |
| int tid = g_cur_tid; |
| int ctx = m->cfg.n_ctx; |
| int window = g_attn_window > 0 ? g_attn_window : ctx; |
| int n_sinks = g_attn_sink; |
| int t = n_tokens - 1; |
|
|
| |
| |
| if (!m->k_cache) model_kv_cache_alloc(m); |
|
|
| |
| |
| |
| |
| |
| static __thread int last_prefill_to = -1; |
| static __thread const int *last_tokens = NULL; |
| int prefill_from = 0; |
| int need_clear = 0; |
| if (tokens != last_tokens) { |
| |
| need_clear = 1; |
| prefill_from = 0; |
| last_prefill_to = -1; |
| } else if (t > last_prefill_to) { |
| |
| prefill_from = last_prefill_to + 1; |
| if (prefill_from > t) prefill_from = t; |
| } else { |
| |
| need_clear = 1; |
| prefill_from = 0; |
| last_prefill_to = -1; |
| } |
| if (need_clear) { |
| size_t per_layer = (size_t)ctx * n * sizeof(float); |
| for (int l = 0; l < nL; l++) { |
| memset(m->k_cache[l], 0, per_layer); |
| memset(m->v_cache[l], 0, per_layer); |
| } |
| if (g_concept_attn_cfg.enable && g_messenger_caches) model_messenger_caches_reset(); |
| } |
|
|
| |
| |
| |
| |
| |
| |
| float *x = g_thr[tid].x; |
| for (int p = prefill_from; p <= t; p++) { |
| |
| int pe_pos = p % ctx; |
| for (int i = 0; i < n; i++) { |
| x[i] = m->wte[(size_t)tokens[p] * n + i]; |
| if (m->wpe) x[i] += m->wpe[(size_t)pe_pos * n + i]; |
| } |
| int cache_pos = p % ctx; |
| if (p < t) { |
| |
| for (int l = 0; l < nL; l++) |
| trans_layer_forward_kv_only_sliding(x, &m->layers[l], &g_thr[tid].acts[l], &m->cfg, |
| cache_pos, p); |
| } else if (g_ponder_cfg.enable && m->ponder_ready) { |
| |
| ponder_train_forward(m, tid, cache_pos, p, window, n_sinks, ctx); |
| } else { |
| |
| for (int l = 0; l < nL; l++) |
| trans_layer_forward_sliding(x, &m->layers[l], &g_thr[tid].acts[l], &m->cfg, |
| cache_pos, p, window, n_sinks, ctx); |
| } |
| } |
| last_prefill_to = t; |
| last_tokens = tokens; |
|
|
| |
| memcpy(g_thr[tid].x_before_final, x, n * sizeof(float)); |
| norm_forward(g_thr[tid].final_ln, x, m->ln_f_w, m->ln_f_b, m->cfg.norm_type, n); |
| compute_mean_std(g_thr[tid].x_before_final, n, &g_thr[tid].final_mean, &g_thr[tid].final_std_inv); |
| |
| if (m->final_ln) memcpy(m->final_ln, g_thr[tid].final_ln, n * sizeof(float)); |
|
|
| int target = tokens[n_tokens]; |
| float *g_full_logits = g_thr[tid].full_logits; |
| g_thr[tid].forward_done = 1; |
| return cross_entropy_full(g_thr[tid].final_ln, m->wte, target, m->cfg.vocab_size, n, g_full_logits); |
| } |
|
|
| void model_backward_sliding(Model *m, const int *tokens, int n_tokens) { |
| int prev = g_accumulate_gradients; |
| g_accumulate_gradients = 1; |
|
|
| int n = m->cfg.n_embd; |
| int nL = m->cfg.n_layer; |
| int ctx = m->cfg.n_ctx; |
| int window = g_attn_window > 0 ? g_attn_window : ctx; |
| int n_sinks = g_attn_sink; |
| int target = tokens[n_tokens]; |
| int tid = g_cur_tid; |
| int t = n_tokens - 1; |
| float *gh = g_thr[tid].gh; |
| float *g_full_logits = g_thr[tid].full_logits; |
|
|
| |
| if (!g_thr[tid].forward_done) { |
| cross_entropy_full(g_thr[tid].final_ln, m->wte, target, m->cfg.vocab_size, n, g_full_logits); |
| } |
| cross_entropy_full_grad(gh, g_thr[tid].final_ln, m->wte, target, m->cfg.vocab_size, n, g_full_logits); |
| g_thr[tid].forward_done = 0; |
|
|
| |
| float gnorm = 0; |
| for (int i = 0; i < n; i++) gnorm += gh[i] * gh[i]; |
| gnorm = sqrtf(gnorm); |
| if (gnorm > 1.0f) { float clip = 1.0f / gnorm; for (int i = 0; i < n; i++) gh[i] *= clip; } |
|
|
| |
| float *g_pre = g_thr[tid].g_pre4; |
| norm_backward(g_pre, gh, g_thr[tid].x_before_final, m->ln_f_w, |
| (float[]){g_thr[tid].final_mean, g_thr[tid].final_std_inv}, m->cfg.norm_type, n, |
| g_thr[tid].grad_lnfw, g_thr[tid].grad_lnfb); |
| memcpy(gh, g_pre, n * sizeof(float)); |
|
|
| |
| |
| |
| if (g_ponder_cfg.enable && m->ponder_ready) { |
| |
| ponder_train_backward(m, tid, window, n_sinks); |
| } else { |
| for (int l = nL - 1; l >= 0; l--) |
| trans_layer_backward_sliding(gh, &m->layers[l], &g_thr[tid].acts[l], &m->cfg, |
| window, n_sinks, 0.0f); |
| } |
|
|
| |
| if (m->grad_wte_accum) { |
| int input_token = tokens[n_tokens - 1]; |
| if (input_token >= 0 && input_token < m->cfg.vocab_size) { |
| float *gw = &g_thr[tid].grad_wte[(size_t)input_token * n]; |
| for (int i = 0; i < n; i++) |
| gw[i] += gh[i]; |
| } |
| } |
|
|
| |
| if (m->grad_wpe_accum) { |
| int n_pos = n_tokens; |
| if (n_pos > 0 && n_pos <= m->cfg.n_ctx) { |
| float inv_npos = 1.0f / (float)n_pos; |
| for (int pos = 0; pos < n_pos; pos++) { |
| float *gw = &g_thr[tid].grad_wpe[(size_t)pos * n]; |
| for (int i = 0; i < n; i++) |
| gw[i] += gh[i] * inv_npos; |
| } |
| } |
| } |
|
|
| g_accumulate_gradients = prev; |
| } |
|
|
| void model_batch_apply(Model *m, float lr, int batch_size) { |
| |
| |
| int n = m->cfg.n_embd; |
| float inv_batch = 1.0f / (float)batch_size; |
|
|
| for (int l = 0; l < m->cfg.n_layer; l++) { |
| TransLayer *tl = &m->layers[l]; |
| BinLayer *bls[8] = {&tl->attn_q, &tl->attn_o, &tl->mlp_gate, &tl->mlp_down}; |
| int n_bl = 4; |
| if (!m->cfg.qkv_merged) { bls[4] = &tl->attn_k; bls[5] = &tl->attn_v; n_bl = 6; } |
| if (m->cfg.act_type == ACT_SWIGLU) { bls[n_bl] = &tl->mlp_up; n_bl++; } |
|
|
| for (int b = 0; b < n_bl; b++) { |
| BinLayer *bl = bls[b]; |
| if (!bl->grad_accum || !bl->w_float) continue; |
|
|
| int in = bl->in_dim, out = bl->out_dim; |
| int t = g_opt_step + 1; |
| float bc1 = 1.0f - powf(g_adam_beta1, (float)t); |
| float bc2 = 1.0f - powf(g_adam_beta2, (float)t); |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| if (g_use_lal_adam && bl->logic_mask) { |
| |
| float core_g_sq = 0, bin_g_sq = 0; |
| int n_core = 0, n_bin = 0; |
| for (int j = 0; j < out; j++) { |
| uint8_t m = bl->logic_mask[j]; |
| if (m == 2) continue; |
| const float *ga = &bl->grad_accum[j * in]; |
| float row_sq = 0; |
| for (int i = 0; i < in; i++) row_sq += ga[i] * ga[i]; |
| row_sq /= in; |
| if (m == 0) { core_g_sq += row_sq; n_core++; } |
| else { bin_g_sq += row_sq; n_bin++; } |
| } |
| core_g_sq = n_core > 0 ? core_g_sq / n_core : 0; |
| bin_g_sq = n_bin > 0 ? bin_g_sq / n_bin : 0; |
|
|
| |
| |
| |
| |
| int core_first = -1, bin_first = -1; |
| for (int j = 0; j < out; j++) { |
| uint8_t m = bl->logic_mask[j]; |
| if (m == 0 && core_first < 0) core_first = j; |
| if (m == 1 && bin_first < 0) bin_first = j; |
| } |
|
|
| float core_v, bin_v; |
| if (g_use_adam && bl->v_adam) { |
| if (core_first >= 0) { |
| bl->v_adam[core_first * in] = |
| g_adam_beta2 * bl->v_adam[core_first * in] + |
| (1.0f - g_adam_beta2) * core_g_sq; |
| core_v = bl->v_adam[core_first * in] / bc2; |
| } else core_v = 1e-8f; |
| if (bin_first >= 0) { |
| bl->v_adam[bin_first * in] = |
| g_adam_beta2 * bl->v_adam[bin_first * in] + |
| (1.0f - g_adam_beta2) * bin_g_sq; |
| bin_v = bl->v_adam[bin_first * in] / bc2; |
| } else bin_v = 1e-8f; |
| } else { |
| core_v = core_g_sq; |
| bin_v = bin_g_sq; |
| } |
| float core_sqrt_v = sqrtf(core_v) + g_adam_eps; |
| float bin_sqrt_v = sqrtf(bin_v) + g_adam_eps; |
|
|
| |
| |
| #pragma omp parallel for schedule(static) |
| for (int j = 0; j < out; j++) { |
| uint8_t m = bl->logic_mask[j]; |
| float *wf = &bl->w_float[j * in]; |
|
|
| if (m == 2) { |
| |
| for (int i = 0; i < in; i++) { |
| float w = wf[i] * (1.0f - g_prune_decay); |
| if (fabsf(w) < g_prune_freeze_thresh) w = 0.0f; |
| wf[i] = w; |
| } |
| continue; |
| } |
|
|
| float lr_j = (m == 0) ? lr * g_core_lr_multiplier : lr; |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| float sqrt_v = (m == 0) ? core_sqrt_v : bin_sqrt_v; |
| float *ga = &bl->grad_accum[j * in]; |
|
|
| if (g_use_adam && bl->m_adam) { |
| float *ma = &bl->m_adam[j * in]; |
| for (int i = 0; i + 7 < in; i += 8) { |
| float g0=ga[i+0]*inv_batch, g1=ga[i+1]*inv_batch; |
| float g2=ga[i+2]*inv_batch, g3=ga[i+3]*inv_batch; |
| float g4=ga[i+4]*inv_batch, g5=ga[i+5]*inv_batch; |
| float g6=ga[i+6]*inv_batch, g7=ga[i+7]*inv_batch; |
| ma[i+0]=g_adam_beta1*ma[i+0]+(1.0f-g_adam_beta1)*g0; |
| ma[i+1]=g_adam_beta1*ma[i+1]+(1.0f-g_adam_beta1)*g1; |
| ma[i+2]=g_adam_beta1*ma[i+2]+(1.0f-g_adam_beta1)*g2; |
| ma[i+3]=g_adam_beta1*ma[i+3]+(1.0f-g_adam_beta1)*g3; |
| ma[i+4]=g_adam_beta1*ma[i+4]+(1.0f-g_adam_beta1)*g4; |
| ma[i+5]=g_adam_beta1*ma[i+5]+(1.0f-g_adam_beta1)*g5; |
| ma[i+6]=g_adam_beta1*ma[i+6]+(1.0f-g_adam_beta1)*g6; |
| ma[i+7]=g_adam_beta1*ma[i+7]+(1.0f-g_adam_beta1)*g7; |
| |
| wf[i+0]-=lr_j*(ma[i+0]/bc1)/sqrt_v; |
| wf[i+1]-=lr_j*(ma[i+1]/bc1)/sqrt_v; |
| wf[i+2]-=lr_j*(ma[i+2]/bc1)/sqrt_v; |
| wf[i+3]-=lr_j*(ma[i+3]/bc1)/sqrt_v; |
| wf[i+4]-=lr_j*(ma[i+4]/bc1)/sqrt_v; |
| wf[i+5]-=lr_j*(ma[i+5]/bc1)/sqrt_v; |
| wf[i+6]-=lr_j*(ma[i+6]/bc1)/sqrt_v; |
| wf[i+7]-=lr_j*(ma[i+7]/bc1)/sqrt_v; |
| } |
| for (int i = (in/8)*8; i < in; i++) { |
| float g = ga[i]*inv_batch; |
| ma[i]=g_adam_beta1*ma[i]+(1.0f-g_adam_beta1)*g; |
| wf[i]-=lr_j*(ma[i]/bc1)/sqrt_v; |
| } |
| } else { |
| float scale = lr_j * inv_batch / sqrt_v; |
| for (int i = 0; i + 7 < in; i += 8) { |
| wf[i+0]-=scale*ga[i+0]; wf[i+1]-=scale*ga[i+1]; |
| wf[i+2]-=scale*ga[i+2]; wf[i+3]-=scale*ga[i+3]; |
| wf[i+4]-=scale*ga[i+4]; wf[i+5]-=scale*ga[i+5]; |
| wf[i+6]-=scale*ga[i+6]; wf[i+7]-=scale*ga[i+7]; |
| } |
| for (int i = (in/8)*8; i < in; i++) |
| wf[i] -= scale * ga[i]; |
| } |
| bl->bias[j] -= lr_j * bl->bias_grad_accum[j] * inv_batch; |
| } |
| |
| goto layer_done; |
| } |
|
|
| |
| #pragma omp parallel for schedule(static) |
| for (int j = 0; j < out; j++) { |
| if (bl->logic_mask && bl->logic_mask[j] == 2) continue; |
| |
| float lr_j = lr; |
| if (bl->logic_mask && bl->logic_mask[j] == 0) |
| lr_j = lr * g_core_lr_multiplier; |
| float *wf = &bl->w_float[j * in]; |
| float *ga = &bl->grad_accum[j * in]; |
|
|
| if (g_use_adam && bl->m_adam) { |
| float *ma = &bl->m_adam[j * in]; |
| float *va = &bl->v_adam[j * in]; |
| for (int i = 0; i + 7 < in; i += 8) { |
| |
| float g0 = ga[i+0]*inv_batch, g1 = ga[i+1]*inv_batch; |
| float g2 = ga[i+2]*inv_batch, g3 = ga[i+3]*inv_batch; |
| float g4 = ga[i+4]*inv_batch, g5 = ga[i+5]*inv_batch; |
| float g6 = ga[i+6]*inv_batch, g7 = ga[i+7]*inv_batch; |
| |
| ma[i+0]=g_adam_beta1*ma[i+0]+(1.0f-g_adam_beta1)*g0; |
| ma[i+1]=g_adam_beta1*ma[i+1]+(1.0f-g_adam_beta1)*g1; |
| ma[i+2]=g_adam_beta1*ma[i+2]+(1.0f-g_adam_beta1)*g2; |
| ma[i+3]=g_adam_beta1*ma[i+3]+(1.0f-g_adam_beta1)*g3; |
| ma[i+4]=g_adam_beta1*ma[i+4]+(1.0f-g_adam_beta1)*g4; |
| ma[i+5]=g_adam_beta1*ma[i+5]+(1.0f-g_adam_beta1)*g5; |
| ma[i+6]=g_adam_beta1*ma[i+6]+(1.0f-g_adam_beta1)*g6; |
| ma[i+7]=g_adam_beta1*ma[i+7]+(1.0f-g_adam_beta1)*g7; |
| va[i+0]=g_adam_beta2*va[i+0]+(1.0f-g_adam_beta2)*g0*g0; |
| va[i+1]=g_adam_beta2*va[i+1]+(1.0f-g_adam_beta2)*g1*g1; |
| va[i+2]=g_adam_beta2*va[i+2]+(1.0f-g_adam_beta2)*g2*g2; |
| va[i+3]=g_adam_beta2*va[i+3]+(1.0f-g_adam_beta2)*g3*g3; |
| va[i+4]=g_adam_beta2*va[i+4]+(1.0f-g_adam_beta2)*g4*g4; |
| va[i+5]=g_adam_beta2*va[i+5]+(1.0f-g_adam_beta2)*g5*g5; |
| va[i+6]=g_adam_beta2*va[i+6]+(1.0f-g_adam_beta2)*g6*g6; |
| va[i+7]=g_adam_beta2*va[i+7]+(1.0f-g_adam_beta2)*g7*g7; |
| |
| float mh0=ma[i+0]/bc1, mh1=ma[i+1]/bc1; |
| float mh2=ma[i+2]/bc1, mh3=ma[i+3]/bc1; |
| float mh4=ma[i+4]/bc1, mh5=ma[i+5]/bc1; |
| float mh6=ma[i+6]/bc1, mh7=ma[i+7]/bc1; |
| float vh0=sqrtf(va[i+0]/bc2)+g_adam_eps; |
| float vh1=sqrtf(va[i+1]/bc2)+g_adam_eps; |
| float vh2=sqrtf(va[i+2]/bc2)+g_adam_eps; |
| float vh3=sqrtf(va[i+3]/bc2)+g_adam_eps; |
| float vh4=sqrtf(va[i+4]/bc2)+g_adam_eps; |
| float vh5=sqrtf(va[i+5]/bc2)+g_adam_eps; |
| float vh6=sqrtf(va[i+6]/bc2)+g_adam_eps; |
| float vh7=sqrtf(va[i+7]/bc2)+g_adam_eps; |
| wf[i+0]-=lr_j*mh0/vh0; wf[i+1]-=lr_j*mh1/vh1; |
| wf[i+2]-=lr_j*mh2/vh2; wf[i+3]-=lr_j*mh3/vh3; |
| wf[i+4]-=lr_j*mh4/vh4; wf[i+5]-=lr_j*mh5/vh5; |
| wf[i+6]-=lr_j*mh6/vh6; wf[i+7]-=lr_j*mh7/vh7; |
| } |
| for (int i = (in/8)*8; i < in; i++) { |
| float g = ga[i]*inv_batch; |
| ma[i]=g_adam_beta1*ma[i]+(1.0f-g_adam_beta1)*g; |
| va[i]=g_adam_beta2*va[i]+(1.0f-g_adam_beta2)*g*g; |
| wf[i]-=lr_j*(ma[i]/bc1)/(sqrtf(va[i]/bc2)+g_adam_eps); |
| } |
| } else { |
| |
| float scale = lr_j * inv_batch; |
| for (int i = 0; i + 7 < in; i += 8) { |
| wf[i+0]-=scale*ga[i+0]; wf[i+1]-=scale*ga[i+1]; |
| wf[i+2]-=scale*ga[i+2]; wf[i+3]-=scale*ga[i+3]; |
| wf[i+4]-=scale*ga[i+4]; wf[i+5]-=scale*ga[i+5]; |
| wf[i+6]-=scale*ga[i+6]; wf[i+7]-=scale*ga[i+7]; |
| } |
| for (int i = (in/8)*8; i < in; i++) |
| wf[i] -= scale * ga[i]; |
| } |
| |
| bl->bias[j] -= lr_j * bl->bias_grad_accum[j] * inv_batch; |
| } |
|
|
| layer_done: |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| if (b == 0 && m->cfg.qkv_merged) { |
| int n = m->cfg.n_embd; |
| int in = bl->in_dim; |
| |
| float frob_sq = 0, max_row_sq = 0; |
| for (int j = 2*n; j < 3*n; j++) { |
| float *wf = &bl->w_float[(size_t)j * in]; |
| float row_sq = 0; |
| for (int i = 0; i < in; i++) row_sq += wf[i] * wf[i]; |
| frob_sq += row_sq; |
| if (row_sq > max_row_sq) max_row_sq = row_sq; |
| } |
| float frob = sqrtf(frob_sq); |
| float max_row = sqrtf(max_row_sq); |
| float approx_rank = frob / (max_row + 1e-12f); |
| if (g_opt_step % 50 == 49) { |
| printf(" [plan-I] step %d W_v approx_rank=%.1f (frob=%.2f max_row=%.2f)\n", |
| g_opt_step, approx_rank, frob, max_row); |
| } |
|
|
| |
| |
| if (!g_skip_wv) { |
| for (int j = 2*n; j < 3*n; j++) { |
| float *wf = &bl->w_float[(size_t)j * in]; |
| for (int i = 0; i < in; i++) { |
| wf[i] *= 0.999f; |
| wf[i] += 0.001f * ((float)rand() / RAND_MAX * 2.0f - 1.0f); |
| } |
| } |
| } |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| if (!g_skip_wv) { |
| float lambda_ortho = 0.05f; |
| |
| |
| static float *G = NULL; |
| static int G_n = 0; |
| if (G_n != n) { |
| free(G); |
| G = (float *)malloc((size_t)n * n * sizeof(float)); |
| G_n = n; |
| } |
|
|
| |
| for (int i = 0; i < n; i++) { |
| for (int j = i; j < n; j++) { |
| float dot = 0; |
| for (int k = 0; k < n; k++) { |
| float *wf_row = &bl->w_float[(size_t)(2*n + k) * in]; |
| dot += wf_row[i] * wf_row[j]; |
| } |
| G[i * n + j] = dot; |
| G[j * n + i] = dot; |
| } |
| } |
|
|
| |
| |
| |
| |
| float tr_G = 0, tr_G2 = 0; |
| for (int i = 0; i < n; i++) tr_G += G[i * n + i]; |
| for (int i = 0; i < n; i++) { |
| for (int j = 0; j < n; j++) { |
| tr_G2 += G[i * n + j] * G[i * n + j]; |
| } |
| } |
| float eff_rank = (tr_G * tr_G) / (tr_G2 + 1e-12f); |
|
|
| for (int i = 0; i < n; i++) |
| G[i * n + i] -= 1.0f; |
|
|
| |
| |
| float scale = 4.0f * lambda_ortho; |
| for (int k = 0; k < n; k++) { |
| float *wf_row = &bl->w_float[(size_t)(2*n + k) * in]; |
| for (int i = 0; i < n; i++) { |
| float grad = 0; |
| for (int j = 0; j < n; j++) |
| grad += wf_row[j] * G[j * n + i]; |
| wf_row[i] -= scale * grad; |
| } |
| } |
|
|
| |
| if (g_opt_step % 50 == 49) { |
| |
| float off_diag = 0, diag_dev = 0; |
| for (int i = 0; i < n; i++) { |
| diag_dev += G[i * n + i] * G[i * n + i]; |
| for (int j = 0; j < n; j++) { |
| if (i != j) off_diag += G[i * n + j] * G[i * n + j]; |
| } |
| } |
| printf(" [ortho] step %d L%d W_v off_diag=%.2f diag_dev=%.4f eff_rank=%.1f/%d\n", |
| g_opt_step, l, off_diag, diag_dev, eff_rank, n); |
| } |
| } |
|
|
| |
| |
| |
| if (b == 1) { |
| float lambda_ortho = 0.05f; |
| static float *Go = NULL; |
| static int Go_n = 0; |
| if (Go_n != n) { |
| free(Go); |
| Go = (float *)malloc((size_t)n * n * sizeof(float)); |
| Go_n = n; |
| } |
| int in_o = bl->in_dim; |
| |
| for (int i = 0; i < n; i++) { |
| for (int j = i; j < n; j++) { |
| float dot = 0; |
| for (int k = 0; k < n; k++) { |
| float *wf_row = &bl->w_float[(size_t)k * in_o]; |
| dot += wf_row[i] * wf_row[j]; |
| } |
| Go[i * n + j] = dot; |
| Go[j * n + i] = dot; |
| } |
| } |
| |
| for (int i = 0; i < n; i++) Go[i * n + i] -= 1.0f; |
| |
| float scale_o = 4.0f * lambda_ortho; |
| for (int k = 0; k < n; k++) { |
| float *wf_row = &bl->w_float[(size_t)k * in_o]; |
| for (int i = 0; i < n; i++) { |
| float grad = 0; |
| for (int j = 0; j < n; j++) |
| grad += wf_row[j] * Go[j * n + i]; |
| wf_row[i] -= scale_o * grad; |
| } |
| } |
| if (g_opt_step % 50 == 49) { |
| float off_diag = 0, diag_dev = 0; |
| for (int i = 0; i < n; i++) { |
| diag_dev += Go[i * n + i] * Go[i * n + i]; |
| for (int j = 0; j < n; j++) { |
| if (i != j) off_diag += Go[i * n + j] * Go[i * n + j]; |
| } |
| } |
| printf(" [ortho] step %d L%d W_o off_diag=%.2f diag_dev=%.4f\n", |
| g_opt_step, l, off_diag, diag_dev); |
| } |
| } |
| } |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| if (b == 1) { |
| int n = m->cfg.n_embd; |
| int in = bl->in_dim; |
| float lambda_ortho_o = 0.05f; |
|
|
| static float *Go = NULL; |
| static int Go_n = 0; |
| if (Go_n != n) { |
| free(Go); |
| Go = (float *)malloc((size_t)n * n * sizeof(float)); |
| Go_n = n; |
| } |
|
|
| |
| for (int i = 0; i < n; i++) { |
| for (int j = i; j < n; j++) { |
| float dot = 0; |
| for (int k = 0; k < n; k++) { |
| float *wf_row = &bl->w_float[(size_t)k * in]; |
| dot += wf_row[i] * wf_row[j]; |
| } |
| Go[i * n + j] = dot; |
| Go[j * n + i] = dot; |
| } |
| } |
|
|
| |
| for (int i = 0; i < n; i++) |
| Go[i * n + i] -= 1.0f; |
|
|
| |
| float scale_o = 4.0f * lambda_ortho_o; |
| for (int k = 0; k < n; k++) { |
| float *wf_row = &bl->w_float[(size_t)k * in]; |
| for (int i = 0; i < n; i++) { |
| float grad = 0; |
| for (int j = 0; j < n; j++) |
| grad += wf_row[j] * Go[j * n + i]; |
| wf_row[i] -= scale_o * grad; |
| } |
| } |
|
|
| |
| if (g_opt_step % 50 == 49) { |
| float off_diag_o = 0, diag_dev_o = 0; |
| for (int i = 0; i < n; i++) { |
| diag_dev_o += Go[i * n + i] * Go[i * n + i]; |
| for (int j = 0; j < n; j++) { |
| if (i != j) off_diag_o += Go[i * n + j] * Go[i * n + j]; |
| } |
| } |
| printf(" [ortho] step %d L%d W_o off_diag=%.2f diag_dev=%.4f\n", |
| g_opt_step, l, off_diag_o, diag_dev_o); |
| } |
| } |
|
|
| |
| |
| |
| |
| if (!g_use_pure_float) { |
| for (int j = 0; j < out; j++) { |
| float clip_val = 1.0f; |
| if (bl->logic_mask && bl->logic_mask[j] == 0) |
| clip_val = 2.0f; |
| |
| float *wf_row = &bl->w_float[j * in]; |
| for (int i = 0; i < in; i++) { |
| if (wf_row[i] > clip_val) wf_row[i] = clip_val; |
| else if (wf_row[i] < -clip_val) wf_row[i] = -clip_val; |
| } |
| } |
| bin_layer_repack(bl); |
| } else { |
| #define W_CLIP_BF 2.0f |
| for (int i = 0; i < in * out; i++) { |
| float w = bl->w_float[i]; |
| if (w > W_CLIP_BF) bl->w_float[i] = W_CLIP_BF; |
| else if (w < -W_CLIP_BF) bl->w_float[i] = -W_CLIP_BF; |
| } |
| #undef W_CLIP_BF |
| } |
| } |
| } |
|
|
| |
| |
| |
| if (m->grad_wte_accum && m->m_wte && m->v_wte && g_use_adam) { |
| int vocab = m->cfg.vocab_size; |
| int t = g_opt_step + 1; |
| float bc1 = 1.0f - powf(g_adam_beta1, (float)t); |
| float bc2 = 1.0f - powf(g_adam_beta2, (float)t); |
| float inv_batch = 1.0f / (float)batch_size; |
|
|
| |
| #pragma omp parallel for schedule(static) |
| for (int v = 0; v < vocab; v++) { |
| float *w = &m->wte[(size_t)v * n]; |
| float *gw = &m->grad_wte_accum[(size_t)v * n]; |
| float *ma = &m->m_wte[(size_t)v * n]; |
| float *va = &m->v_wte[(size_t)v * n]; |
| |
| |
| |
| |
| int has_grad = 0; |
| for (int i = 0; i < n; i++) { |
| float g = gw[i] * inv_batch; |
| if (fabsf(g) >= 1e-12f) { |
| has_grad = 1; |
| ma[i] = g_adam_beta1 * ma[i] + (1.0f - g_adam_beta1) * g; |
| va[i] = g_adam_beta2 * va[i] + (1.0f - g_adam_beta2) * g * g; |
| float mh = ma[i] / bc1; |
| float vh = sqrtf(va[i] / bc2) + g_adam_eps; |
| |
| |
| |
| |
| |
| |
| |
| if (vh < 1e-4f) vh = 1e-4f; |
| w[i] -= lr * g_wte_lr_scale * mh / vh; |
| } |
| } |
| |
| if (!has_grad) { |
| |
| float norm_sq = 0.0f; |
| for (int i = 0; i < n; i++) norm_sq += w[i] * w[i]; |
| float norm = sqrtf(norm_sq); |
| |
| if (norm > 0.5f) { |
| float decay = 0.9999f; |
| for (int i = 0; i < n; i++) { |
| w[i] *= decay; |
| } |
| } |
| } |
| } |
| } |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| if (m->grad_wpe_accum && m->m_wpe && m->v_wpe && g_use_adam && m->wpe) { |
| int n_ctx = m->cfg.n_ctx; |
| int t = g_opt_step + 1; |
| float bc1 = 1.0f - powf(g_adam_beta1, (float)t); |
| float bc2 = 1.0f - powf(g_adam_beta2, (float)t); |
| float inv_batch = 1.0f / (float)batch_size; |
| float wpe_lr = lr * 0.3f; |
| float wpe_max_norm = 1.0f; |
|
|
| for (int pos = 0; pos < n_ctx; pos++) { |
| float *w = &m->wpe[(size_t)pos * n]; |
| float *gw = &m->grad_wpe_accum[(size_t)pos * n]; |
| float *ma = &m->m_wpe[(size_t)pos * n]; |
| float *va = &m->v_wpe[(size_t)pos * n]; |
| for (int i = 0; i < n; i++) { |
| float g = gw[i] * inv_batch; |
| if (fabsf(g) < 1e-12f) continue; |
| ma[i] = g_adam_beta1 * ma[i] + (1.0f - g_adam_beta1) * g; |
| va[i] = g_adam_beta2 * va[i] + (1.0f - g_adam_beta2) * g * g; |
| float mh = ma[i] / bc1; |
| float vh = sqrtf(va[i] / bc2) + g_adam_eps; |
| w[i] -= wpe_lr * mh / vh; |
| } |
| |
| float norm_sq = 0.0f; |
| for (int i = 0; i < n; i++) norm_sq += w[i] * w[i]; |
| float norm = sqrtf(norm_sq); |
| if (norm > wpe_max_norm) { |
| float scale = wpe_max_norm / norm; |
| for (int i = 0; i < n; i++) w[i] *= scale; |
| } |
| } |
| } |
|
|
| |
| |
| |
| |
| |
| for (int l = 0; l < m->cfg.n_layer; l++) { |
| TransLayer *tl = &m->layers[l]; |
| if (tl->grad_norm1_w && tl->m_norm1_w && g_use_adam) { |
| |
| int t = g_opt_step + 1; |
| float bc1 = 1.0f - powf(g_adam_beta1, (float)t); |
| float bc2 = 1.0f - powf(g_adam_beta2, (float)t); |
| float lr_norm = lr; |
| for (int i = 0; i < n; i++) { |
| |
| |
| |
| float g1w = tl->grad_norm1_w[i] * inv_batch; |
| if (fabsf(g1w) > 1e-12f) { |
| tl->m_norm1_w[i] = g_adam_beta1 * tl->m_norm1_w[i] + (1.0f - g_adam_beta1) * g1w; |
| tl->v_norm1_w[i] = g_adam_beta2 * tl->v_norm1_w[i] + (1.0f - g_adam_beta2) * g1w * g1w; |
| float mh = tl->m_norm1_w[i] / bc1; |
| float vh = sqrtf(tl->v_norm1_w[i] / bc2) + g_adam_eps; |
| tl->norm1_w[i] -= lr_norm * mh / vh; |
| |
| |
| |
| |
| if (tl->norm1_w[i] < 0.5f) tl->norm1_w[i] = 0.5f; |
| if (tl->norm1_w[i] > 2.0f) tl->norm1_w[i] = 2.0f; |
| } |
|
|
| |
| float g1b = tl->grad_norm1_b[i] * inv_batch; |
| if (fabsf(g1b) > 1e-12f) { |
| tl->m_norm1_b[i] = g_adam_beta1 * tl->m_norm1_b[i] + (1.0f - g_adam_beta1) * g1b; |
| tl->v_norm1_b[i] = g_adam_beta2 * tl->v_norm1_b[i] + (1.0f - g_adam_beta2) * g1b * g1b; |
| float mh = tl->m_norm1_b[i] / bc1; |
| float vh = sqrtf(tl->v_norm1_b[i] / bc2) + g_adam_eps; |
| tl->norm1_b[i] -= lr_norm * mh / vh; |
| } |
| |
| float g2w = tl->grad_norm2_w[i] * inv_batch; |
| if (fabsf(g2w) > 1e-12f) { |
| tl->m_norm2_w[i] = g_adam_beta1 * tl->m_norm2_w[i] + (1.0f - g_adam_beta1) * g2w; |
| tl->v_norm2_w[i] = g_adam_beta2 * tl->v_norm2_w[i] + (1.0f - g_adam_beta2) * g2w * g2w; |
| float mh = tl->m_norm2_w[i] / bc1; |
| float vh = sqrtf(tl->v_norm2_w[i] / bc2) + g_adam_eps; |
| tl->norm2_w[i] -= lr_norm * mh / vh; |
| |
| if (tl->norm2_w[i] < 0.5f) tl->norm2_w[i] = 0.5f; |
| if (tl->norm2_w[i] > 2.0f) tl->norm2_w[i] = 2.0f; |
| } |
|
|
| |
| float g2b = tl->grad_norm2_b[i] * inv_batch; |
| if (fabsf(g2b) > 1e-12f) { |
| tl->m_norm2_b[i] = g_adam_beta1 * tl->m_norm2_b[i] + (1.0f - g_adam_beta1) * g2b; |
| tl->v_norm2_b[i] = g_adam_beta2 * tl->v_norm2_b[i] + (1.0f - g_adam_beta2) * g2b * g2b; |
| float mh = tl->m_norm2_b[i] / bc1; |
| float vh = sqrtf(tl->v_norm2_b[i] / bc2) + g_adam_eps; |
| tl->norm2_b[i] -= lr_norm * mh / vh; |
| } |
| } |
| } |
| } |
| |
| if (m->grad_ln_f_w_accum && m->m_ln_f_w && g_use_adam) { |
| int t = g_opt_step + 1; |
| float bc1 = 1.0f - powf(g_adam_beta1, (float)t); |
| float bc2 = 1.0f - powf(g_adam_beta2, (float)t); |
| for (int i = 0; i < n; i++) { |
| float gw = m->grad_ln_f_w_accum[i] * inv_batch; |
| float gb = m->grad_ln_f_b_accum[i] * inv_batch; |
| if (fabsf(gw) < 1e-12f && fabsf(gb) < 1e-12f) continue; |
| m->m_ln_f_w[i] = g_adam_beta1 * m->m_ln_f_w[i] + (1.0f - g_adam_beta1) * gw; |
| m->v_ln_f_w[i] = g_adam_beta2 * m->v_ln_f_w[i] + (1.0f - g_adam_beta2) * gw * gw; |
| m->m_ln_f_b[i] = g_adam_beta1 * m->m_ln_f_b[i] + (1.0f - g_adam_beta1) * gb; |
| m->v_ln_f_b[i] = g_adam_beta2 * m->v_ln_f_b[i] + (1.0f - g_adam_beta2) * gb * gb; |
| float mhw = m->m_ln_f_w[i] / bc1, vhw = sqrtf(m->v_ln_f_w[i] / bc2) + g_adam_eps; |
| float mhb = m->m_ln_f_b[i] / bc1, vhb = sqrtf(m->v_ln_f_b[i] / bc2) + g_adam_eps; |
| m->ln_f_w[i] -= lr * mhw / vhw; |
| m->ln_f_b[i] -= lr * mhb / vhb; |
| |
| if (m->ln_f_w[i] < 0.5f) m->ln_f_w[i] = 0.5f; |
| if (m->ln_f_w[i] > 2.0f) m->ln_f_w[i] = 2.0f; |
| } |
| } |
|
|
| |
| |
| |
| |
| |
| |
| for (int l = 0; l < m->cfg.n_layer; l++) { |
| TransLayer *tl = &m->layers[l]; |
| BinLayer *bls[8] = {&tl->attn_q, &tl->attn_o, |
| &tl->mlp_gate, &tl->mlp_down}; |
| int n_bl = 4; |
| if (!m->cfg.qkv_merged) { bls[4] = &tl->attn_k; bls[5] = &tl->attn_v; n_bl = 6; } |
| if (m->cfg.act_type == ACT_SWIGLU) { bls[n_bl] = &tl->mlp_up; n_bl++; } |
| for (int b = 0; b < n_bl; b++) { |
| if (bls[b]->w_float && bls[b]->logic_mask) |
| bin_layer_repack(bls[b]); |
| } |
| } |
|
|
| |
| if (g_use_adam) { |
| |
| ponder_apply(m, lr, batch_size, g_opt_step + 1); |
| g_opt_step++; |
| } |
|
|
| |
| |
| |
| |
| |
| } |
|
|
| void model_stateful_begin(Model *m) { |
| |
| if (!m->k_cache) model_kv_cache_alloc(m); |
|
|
| |
| int n_layer = m->cfg.n_layer; |
| size_t per_layer = (size_t)m->cfg.n_ctx * m->cfg.n_embd * sizeof(float); |
| for (int l = 0; l < n_layer; l++) { |
| memset(m->k_cache[l], 0, per_layer); |
| memset(m->v_cache[l], 0, per_layer); |
| } |
|
|
| |
| if (!g_sctx.x) g_sctx.x = malloc(m->cfg.n_embd * sizeof(float)); |
| if (!g_sctx.logits) g_sctx.logits = malloc(m->cfg.vocab_size * sizeof(float)); |
|
|
| g_sctx.kv_pos = 0; |
| g_sctx.total_pos = 0; |
| g_sctx.active = 1; |
| model_messenger_caches_reset(); |
|
|
| |
| |
| { |
| int nL = m->cfg.n_layer, nE = m->cfg.n_embd; |
| size_t bytes = (size_t)nL * LCTX_SLOTS * nE * sizeof(float); |
| if (g_runtime_cg) { |
| if (!g_sctx.cctx_k) g_sctx.cctx_k = (float *)malloc(bytes); |
| if (!g_sctx.cctx_v) g_sctx.cctx_v = (float *)malloc(bytes); |
| if (!g_sctx.cctx_cnt) g_sctx.cctx_cnt = (int *)calloc((size_t)nL * LCTX_SLOTS, sizeof(int)); |
| if (!g_sctx.cctx_anchor) g_sctx.cctx_anchor = (int *)calloc((size_t)nL * LCTX_SLOTS, sizeof(int)); |
| memset(g_sctx.cctx_k, 0, bytes); |
| memset(g_sctx.cctx_v, 0, bytes); |
| memset(g_sctx.cctx_cnt, 0, (size_t)nL * LCTX_SLOTS * sizeof(int)); |
| memset(g_sctx.cctx_anchor, 0, (size_t)nL * LCTX_SLOTS * sizeof(int)); |
| g_sctx.cctx_n_layer = nL; |
| g_sctx.cctx_n_embd = nE; |
| } else { |
| g_sctx.cctx_k = g_sctx.cctx_v = NULL; |
| g_sctx.cctx_cnt = g_sctx.cctx_anchor = NULL; |
| g_sctx.cctx_n_layer = g_sctx.cctx_n_embd = 0; |
| } |
| } |
|
|
| |
| |
| |
| |
| |
| |
| int window = g_attn_window > 0 ? g_attn_window : m->cfg.n_ctx; |
| int sinks = g_attn_sink; |
| printf("[*] stateful inference started: window=%d, sinks=%d, ctx=%d " |
| "(from global g_attn_window/g_attn_sink)\n", |
| window, sinks, m->cfg.n_ctx); |
| } |
|
|
| |
| void model_stateful_reset(Model *m) { |
| if (!m->k_cache) return; |
| int n_layer = m->cfg.n_layer; |
| size_t per_layer = (size_t)m->cfg.n_ctx * m->cfg.n_embd * sizeof(float); |
| for (int l = 0; l < n_layer; l++) { |
| memset(m->k_cache[l], 0, per_layer); |
| memset(m->v_cache[l], 0, per_layer); |
| } |
| g_sctx.kv_pos = 0; |
| g_sctx.total_pos = 0; |
| model_messenger_caches_reset(); |
| |
| if (g_sctx.cctx_k && g_sctx.cctx_v && g_sctx.cctx_cnt) { |
| int nL = g_sctx.cctx_n_layer, nE = g_sctx.cctx_n_embd; |
| size_t bytes = (size_t)nL * LCTX_SLOTS * nE * sizeof(float); |
| memset(g_sctx.cctx_k, 0, bytes); |
| memset(g_sctx.cctx_v, 0, bytes); |
| memset(g_sctx.cctx_cnt, 0, (size_t)nL * LCTX_SLOTS * sizeof(int)); |
| memset(g_sctx.cctx_anchor, 0, (size_t)nL * LCTX_SLOTS * sizeof(int)); |
| } |
| } |
|
|
| |
| const float *model_stateful_forward_sliding(Model *m, int token) { |
| if (!g_sctx.active || !m->k_cache) { |
| fprintf(stderr, "[!] stateful mode not active — call model_stateful_begin() first\n"); |
| return NULL; |
| } |
| int n = m->cfg.n_embd, nL = m->cfg.n_layer, ctx = m->cfg.n_ctx; |
| |
| |
| |
| int window = g_attn_window > 0 ? g_attn_window : ctx; |
| int n_sinks = g_attn_sink; |
|
|
| |
| int pos = g_sctx.kv_pos; |
| int abs_pos = g_sctx.total_pos; |
| int pe_pos = (m->cfg.attn_type == ATTN_LEARNED) ? (abs_pos % ctx) : abs_pos; |
|
|
| float *x = g_sctx.x; |
| |
| for (int i = 0; i < n; i++) { |
| x[i] = m->wte[(size_t)token * n + i]; |
| if (m->wpe) x[i] += m->wpe[(size_t)pe_pos * n + i]; |
| } |
|
|
| |
| |
| |
| |
| int eject = abs_pos - window; |
| if (g_cctx_cfg.enable && g_runtime_cg && g_sctx.cctx_k && eject >= (int)n_sinks) { |
| int K = g_runtime_cg->K; |
| int anchor = g_runtime_cg->neighbor[(size_t)eject * K]; |
| if (anchor >= 0) { |
| int slot = anchor % LCTX_SLOTS; |
| int eject_phys = eject % ctx; |
| for (int l = 0; l < nL; l++) { |
| const float *k_e = m->k_cache[l] + (size_t)eject_phys * n; |
| const float *v_e = m->v_cache[l] + (size_t)eject_phys * n; |
| float *ck = g_sctx.cctx_k + ((size_t)l * LCTX_SLOTS + slot) * n; |
| float *cv = g_sctx.cctx_v + ((size_t)l * LCTX_SLOTS + slot) * n; |
| int cnt = g_sctx.cctx_cnt[(size_t)l * LCTX_SLOTS + slot]; |
| float inv = (cnt > 0) ? (1.0f / (cnt + 1)) : 1.0f; |
| for (int i = 0; i < n; i++) { |
| ck[i] = ck[i] * (cnt * inv) + k_e[i] * inv; |
| cv[i] = cv[i] * (cnt * inv) + v_e[i] * inv; |
| } |
| g_sctx.cctx_cnt[(size_t)l * LCTX_SLOTS + slot] = cnt + 1; |
| g_sctx.cctx_anchor[(size_t)l * LCTX_SLOTS + slot] = anchor; |
| } |
| } |
| } |
|
|
| |
| if (g_ponder_cfg.enable && m->ponder_ready) { |
| |
| ponder_infer_forward(m, pos, abs_pos, window, n_sinks); |
| } else { |
| for (int l = 0; l < nL; l++) |
| trans_layer_forward_sliding(x, &m->layers[l], &m->acts[l], &m->cfg, |
| pos, abs_pos, window, n_sinks, pos + 1); |
| } |
|
|
| |
| memcpy(m->x_before_final, x, n * sizeof(float)); |
| norm_forward(m->final_ln, x, m->ln_f_w, m->ln_f_b, m->cfg.norm_type, n); |
| compute_mean_std(m->x_before_final, n, &m->final_mean, &m->final_std_inv); |
|
|
| int V = m->cfg.vocab_size; |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| for (int j = 0; j < V; j++) { |
| const float *w = &m->wte[(size_t)j * n]; |
| float s = 0; |
| for (int k = 0; k + 7 < n; k += 8) |
| s += m->final_ln[k+0]*w[k+0] + m->final_ln[k+1]*w[k+1] |
| + m->final_ln[k+2]*w[k+2] + m->final_ln[k+3]*w[k+3] |
| + m->final_ln[k+4]*w[k+4] + m->final_ln[k+5]*w[k+5] |
| + m->final_ln[k+6]*w[k+6] + m->final_ln[k+7]*w[k+7]; |
| for (int k = (n/8)*8; k < n; k++) |
| s += m->final_ln[k] * w[k]; |
| g_sctx.logits[j] = s * g_logit_scale; |
| } |
|
|
| |
| g_sctx.kv_pos = (g_sctx.kv_pos + 1) % ctx; |
| g_sctx.total_pos++; |
| return g_sctx.logits; |
| } |
|
|
| |
| void model_set_sliding_window(Model *m, int window, int n_sinks) { |
| m->cfg.sliding_window = window; |
| m->cfg.n_sinks = n_sinks; |
| printf("[*] sliding window configured: W=%d, sinks=%d (effective context: %d)\n", |
| window, n_sinks, window + n_sinks); |
| } |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| |
| ConceptAttnConfig g_concept_attn_cfg = {0}; |
| |
| ConceptCtxConfig g_cctx_cfg = {0, 1.0f}; |
| ConceptGraph *g_runtime_cg = NULL; |
| void model_set_concept_ctx(const ConceptCtxConfig *cfg, ConceptGraph *cg) { |
| g_runtime_cg = cg; |
| g_cctx_cfg.enable = (cg != NULL); |
| if (cfg) g_cctx_cfg.mem_scale = cfg->mem_scale; |
| printf("[*] C3 概念图驱动长上下文记忆: %s (mem_scale=%.2f)\n", |
| g_cctx_cfg.enable ? "ENABLED" : "disabled", g_cctx_cfg.mem_scale); |
| } |
| |
| |
| float g_attn_res_scale = 0.15f; |
|
|
| |
| |
| |
|
|
| |
| |
| |
| #define GATE_RING_SIZE 512 |
| static float g_gate_ring[GATE_RING_SIZE]; |
| static int g_gate_head = 0; |
| static int g_gate_n_samples = 0; |
| float g_gate_p25 = 0.1f; |
|
|
| |
| static void gate_score_record(float score) { |
| g_gate_ring[g_gate_head] = score; |
| g_gate_head = (g_gate_head + 1) % GATE_RING_SIZE; |
| if (g_gate_n_samples < GATE_RING_SIZE) g_gate_n_samples++; |
|
|
| |
| if ((g_gate_n_samples & 63) == 0 && g_gate_n_samples >= 64) { |
| |
| float tmp[GATE_RING_SIZE]; |
| int n = g_gate_n_samples; |
| memcpy(tmp, g_gate_ring, n * sizeof(float)); |
| |
| for (int i = 1; i < n; i++) { |
| float key = tmp[i]; |
| int j = i - 1; |
| while (j >= 0 && tmp[j] > key) { tmp[j+1] = tmp[j]; j--; } |
| tmp[j+1] = key; |
| } |
| |
| int idx = (int)(0.25f * (n - 1)); |
| g_gate_p25 = tmp[idx]; |
| } |
| } |
|
|
| |
| MessengerCache *g_messenger_caches = NULL; |
| static int g_messenger_caches_n_layer = 0; |
|
|
| |
|
|
| void messenger_cache_alloc(MessengerCache *mc, int segment_capacity, |
| int num_messengers, int n_embd) { |
| if (!mc || segment_capacity <= 0 || num_messengers <= 0 || n_embd <= 0) { |
| if (mc) memset(mc, 0, sizeof(*mc)); |
| return; |
| } |
| mc->segment_capacity = segment_capacity; |
| mc->num_messengers = num_messengers; |
| mc->n_embd = n_embd; |
| mc->n_filled = 0; |
| size_t total = (size_t)segment_capacity * num_messengers * n_embd; |
| mc->messenger_k = (float *)calloc(total, sizeof(float)); |
| mc->messenger_v = (float *)calloc(total, sizeof(float)); |
| mc->segment_filled = (uint8_t *)calloc(segment_capacity, sizeof(uint8_t)); |
| if (!mc->messenger_k || !mc->messenger_v || !mc->segment_filled) { |
| fprintf(stderr, "[!] messenger_cache_alloc: OOM (cap=%d, S=%d, d=%d)\n", |
| segment_capacity, num_messengers, n_embd); |
| messenger_cache_free(mc); |
| } |
| } |
|
|
| void messenger_cache_free(MessengerCache *mc) { |
| if (!mc) return; |
| free(mc->messenger_k); |
| free(mc->messenger_v); |
| free(mc->segment_filled); |
| memset(mc, 0, sizeof(*mc)); |
| } |
|
|
| void messenger_cache_reset(MessengerCache *mc) { |
| if (!mc || mc->segment_capacity <= 0) return; |
| size_t total = (size_t)mc->segment_capacity * mc->num_messengers * mc->n_embd; |
| if (mc->messenger_k) memset(mc->messenger_k, 0, total * sizeof(float)); |
| if (mc->messenger_v) memset(mc->messenger_v, 0, total * sizeof(float)); |
| if (mc->segment_filled) memset(mc->segment_filled, 0, mc->segment_capacity); |
| mc->n_filled = 0; |
| } |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| void generate_segment_messengers(const float *v_seg, int seg_len, int n_embd, |
| int num_messengers, |
| float *out_k, float *out_v) { |
| if (!v_seg || !out_k || !out_v || seg_len <= 0 || n_embd <= 0 || num_messengers <= 0) |
| return; |
|
|
| |
| |
| |
| |
| |
| float *global_mean = (float *)calloc(n_embd, sizeof(float)); |
| if (!global_mean) { |
| |
| for (int m = 0; m < num_messengers; m++) { |
| int bucket_start = (int)((long long)m * seg_len / num_messengers); |
| int bucket_end = (int)((long long)(m + 1) * seg_len / num_messengers); |
| if (bucket_end <= bucket_start) bucket_end = bucket_start + 1; |
| if (bucket_end > seg_len) bucket_end = seg_len; |
| int bucket_size = bucket_end - bucket_start; |
| if (bucket_size <= 0) bucket_size = 1; |
| float *k_dst = out_k + (size_t)m * n_embd; |
| float *v_dst = out_v + (size_t)m * n_embd; |
| float inv = 1.0f / (float)bucket_size; |
| for (int d = 0; d < n_embd; d++) { |
| float sum = 0.0f; |
| for (int t = bucket_start; t < bucket_end; t++) |
| sum += v_seg[(size_t)t * n_embd + d]; |
| float val = sum * inv; |
| v_dst[d] = val; |
| k_dst[d] = val; |
| } |
| } |
| return; |
| } |
|
|
| |
| float inv_seg = 1.0f / (float)seg_len; |
| for (int d = 0; d < n_embd; d++) { |
| float sum = 0.0f; |
| for (int t = 0; t < seg_len; t++) |
| sum += v_seg[(size_t)t * n_embd + d]; |
| global_mean[d] = sum * inv_seg; |
| } |
|
|
| |
| float max_norm = sqrtf((float)n_embd) * 0.5f; |
|
|
| |
| for (int m = 0; m < num_messengers; m++) { |
| int bucket_start = (int)((long long)m * seg_len / num_messengers); |
| int bucket_end = (int)((long long)(m + 1) * seg_len / num_messengers); |
| if (bucket_end <= bucket_start) bucket_end = bucket_start + 1; |
| if (bucket_end > seg_len) bucket_end = seg_len; |
| int bucket_size = bucket_end - bucket_start; |
| if (bucket_size <= 0) bucket_size = 1; |
|
|
| float *k_dst = out_k + (size_t)m * n_embd; |
| float *v_dst = out_v + (size_t)m * n_embd; |
| float inv = 1.0f / (float)bucket_size; |
| for (int d = 0; d < n_embd; d++) { |
| float sum = 0.0f; |
| for (int t = bucket_start; t < bucket_end; t++) |
| sum += v_seg[(size_t)t * n_embd + d]; |
| float val = sum * inv - global_mean[d]; |
| v_dst[d] = val; |
| k_dst[d] = val; |
| } |
|
|
| |
| float norm_sq = 0.0f; |
| for (int d = 0; d < n_embd; d++) |
| norm_sq += v_dst[d] * v_dst[d]; |
| float norm = sqrtf(norm_sq + 1e-8f); |
| if (norm > max_norm) { |
| float scale_factor = max_norm / norm; |
| for (int d = 0; d < n_embd; d++) { |
| v_dst[d] *= scale_factor; |
| k_dst[d] *= scale_factor; |
| } |
| } |
| } |
|
|
| free(global_mean); |
| } |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| int concept_boundary_gate(const float *q_i, const float *k_j, |
| int head_dim, int distance, |
| const ConceptAttnConfig *cfg) { |
| if (!cfg->gate_enable) return 1; |
|
|
| |
| int coarse_dim = head_dim > 16 ? head_dim / 4 : head_dim; |
| float q_norm = 0.0f, k_norm = 0.0f, dot = 0.0f; |
| for (int d = 0; d < coarse_dim; d++) { |
| dot += q_i[d] * k_j[d]; |
| q_norm += q_i[d] * q_i[d]; |
| k_norm += k_j[d] * k_j[d]; |
| } |
| q_norm = sqrtf(q_norm + 1e-8f); |
| k_norm = sqrtf(k_norm + 1e-8f); |
| float sim_coarse = dot / (q_norm * k_norm + 1e-8f); |
|
|
| |
| float dist_prior = 1.0f; |
| if (cfg->gate_distance_prior && distance > 0) { |
| float tau = (float)(cfg->segment_len > 0 ? cfg->segment_len : 64); |
| dist_prior = expf(-(float)distance / tau); |
| } |
|
|
| |
| float gate_score = sim_coarse; |
| if (cfg->gate_distance_prior) { |
| gate_score += 0.5f * dist_prior; |
| } |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| float effective_threshold = cfg->gate_threshold; |
|
|
| |
| if (g_gate_n_samples > 50) { |
| |
| effective_threshold = g_gate_p25; |
| } |
|
|
| |
| if (gate_score < effective_threshold) { |
| |
| |
| unsigned int hash = (unsigned int)(distance * 2654435761u); |
| hash ^= (unsigned int)((sim_coarse + 1000.0f) * 10000.0f); |
| hash = (hash * 40503u) ^ (hash >> 7); |
| float r = (float)(hash & 0xFFFF) / 65535.0f; |
| if (r < cfg->gate_fallback_prob) { |
| return 1; |
| } |
| return 0; |
| } |
| return 1; |
| } |
|
|
| |
| |
| |
| |
| |
| |
| HeadAccessType get_head_access_type(int head_idx, int n_head, |
| const ConceptAttnConfig *cfg) { |
| if (!cfg->hetero_enable) return HEAD_GLOBAL; |
|
|
| int n_local = cfg->n_local_heads; |
| int n_messenger = cfg->n_messenger_heads; |
| |
| if (n_local < 0) n_local = n_head / 2; |
| if (n_messenger < 0) n_messenger = (n_head - n_local) / 2; |
| if (n_local + n_messenger > n_head) n_local = n_head / 2; |
|
|
| if (head_idx < n_local) return HEAD_LOCAL; |
| if (head_idx < n_local + n_messenger) return HEAD_MESSENGER; |
| return HEAD_GLOBAL; |
| } |
|
|
| int get_head_window(int head_idx, int n_head, int base_window, |
| const ConceptAttnConfig *cfg) { |
| if (!cfg->hetero_enable) return base_window; |
| HeadAccessType t = get_head_access_type(head_idx, n_head, cfg); |
| switch (t) { |
| case HEAD_LOCAL: return base_window; |
| case HEAD_MESSENGER: return base_window; |
| case HEAD_GLOBAL: return base_window * 2; |
| default: return base_window; |
| } |
| } |
|
|
| int head_can_access_messenger(int head_idx, int n_head, |
| const ConceptAttnConfig *cfg) { |
| if (!cfg->hetero_enable) return 1; |
| HeadAccessType t = get_head_access_type(head_idx, n_head, cfg); |
| return (t == HEAD_MESSENGER || t == HEAD_GLOBAL); |
| } |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| void attention_forward_concept(float *attn_out, const float *qkv, |
| int n_embd, int n_head, int seq_pos, |
| float *k_cache, float *v_cache, |
| int n_ctx, |
| const ConceptAttnConfig *cfg, |
| MessengerCache *mc) { |
| |
| if (!cfg || !cfg->enable) { |
| attention_forward_sliding(attn_out, qkv, n_embd, n_head, seq_pos, |
| k_cache, v_cache, n_ctx, |
| g_attn_window > 0 ? g_attn_window : n_ctx, |
| g_attn_sink); |
| return; |
| } |
| (void)n_ctx; |
|
|
| int head_dim = n_embd / n_head; |
| float scale = 1.0f / sqrtf((float)head_dim); |
| g_ca_stats.forwards++; |
| g_ca_stats.full_equiv += (long)(seq_pos + 1) * n_head; |
|
|
| const float *Q = qkv; |
| const float *K_new = qkv + n_embd; |
| const float *V_new = qkv + 2 * n_embd; |
|
|
| |
| int eff_ctx = (n_ctx > 0) ? n_ctx : (seq_pos + 1); |
| int cache_pos = seq_pos % eff_ctx; |
| memcpy(k_cache + (size_t)cache_pos * n_embd, K_new, n_embd * sizeof(float)); |
| memcpy(v_cache + (size_t)cache_pos * n_embd, V_new, n_embd * sizeof(float)); |
|
|
| |
| |
| |
| int seg_len = cfg->segment_len > 0 ? cfg->segment_len : n_ctx; |
| int seg_idx = seq_pos / seg_len; |
| int seg_start = seg_idx * seg_len; |
| int seg_end = seg_start + seg_len; |
| if (seg_end > seq_pos + 1) seg_end = seq_pos + 1; |
| if (seg_end > n_ctx) seg_end = n_ctx; |
| int actual_seg_len = seg_end - seg_start; |
|
|
| |
| |
| |
| |
| int min_seg = cfg->min_seg_len > 0 ? cfg->min_seg_len : 1; |
| int tail_complete = (seq_pos == n_ctx - 1) && (actual_seg_len >= min_seg); |
| int is_seg_complete = (actual_seg_len >= seg_len) || |
| tail_complete || |
| ((seq_pos + 1) % seg_len == 0 && seq_pos > 0); |
|
|
| if (mc && cfg->num_messengers > 0 && is_seg_complete && actual_seg_len > 0) { |
| if (seg_idx < mc->segment_capacity && !mc->segment_filled[seg_idx]) { |
| |
| float *v_seg = v_cache + (size_t)seg_start * n_embd; |
| float *mk = mc->messenger_k + (size_t)seg_idx * cfg->num_messengers * n_embd; |
| float *mv = mc->messenger_v + (size_t)seg_idx * cfg->num_messengers * n_embd; |
| generate_segment_messengers(v_seg, actual_seg_len, n_embd, |
| cfg->num_messengers, mk, mv); |
| mc->segment_filled[seg_idx] = 1; |
| if (seg_idx + 1 > mc->n_filled) mc->n_filled = seg_idx + 1; |
| g_ca_stats.last_n_filled = mc->n_filled; |
|
|
| |
| |
| |
| { |
| int S = cfg->num_messengers; |
| float seg_cos_sum = 0.0f; int seg_cos_pairs = 0; |
| float seg_norm_sum = 0.0f; |
| for (int a = 0; a < S; a++) { |
| const float *ma = mv + (size_t)a * n_embd; |
| float na = 0.0f; |
| for (int d = 0; d < n_embd; d++) na += ma[d] * ma[d]; |
| na = sqrtf(na); |
| seg_norm_sum += na; |
| for (int b = a + 1; b < S; b++) { |
| const float *mb = mv + (size_t)b * n_embd; |
| float dot = 0.0f, nb = 0.0f; |
| for (int d = 0; d < n_embd; d++) { |
| dot += ma[d] * mb[d]; |
| nb += mb[d] * mb[d]; |
| } |
| nb = sqrtf(nb); |
| float cos = (na > 1e-6f && nb > 1e-6f) ? dot / (na * nb) : 0.0f; |
| seg_cos_sum += cos; |
| seg_cos_pairs++; |
| } |
| } |
| if (seg_cos_pairs > 0) |
| g_ca_stats.msg_inter_cos += seg_cos_sum / seg_cos_pairs; |
| g_ca_stats.msg_norm += seg_norm_sum / (float)S; |
| g_ca_stats.msg_segments++; |
| } |
| } |
| } |
|
|
| |
| |
| |
| |
| |
| |
| |
| int n_attend = 0; |
| int pos_list[10240]; |
| int is_messenger[10240]; |
|
|
| for (int h = 0; h < n_head; h++) { |
| HeadAccessType htype = get_head_access_type(h, n_head, cfg); |
| int window = get_head_window(h, n_head, |
| cfg->gate_window > 0 ? cfg->gate_window : 64, |
| cfg); |
| if (window < 1) window = 1; |
|
|
| |
| n_attend = 0; |
| |
| int win_start = seq_pos - window + 1; |
| if (win_start < 0) win_start = 0; |
| for (int j = win_start; j <= seq_pos && n_attend < 10240; j++) { |
| pos_list[n_attend] = j; |
| is_messenger[n_attend] = -1; |
| n_attend++; |
| } |
|
|
| |
| int can_msg = head_can_access_messenger(h, n_head, cfg); |
| if (can_msg && mc && cfg->num_messengers > 0) { |
| int n_neighbor = cfg->messenger_neighbors > 0 ? cfg->messenger_neighbors : 2; |
| |
| int neighbor_start = seg_idx - n_neighbor; |
| if (neighbor_start < 0) neighbor_start = 0; |
| for (int s = neighbor_start; s <= seg_idx && n_attend < 10240; s++) { |
| if (s >= mc->segment_capacity) break; |
| if (!mc->segment_filled[s]) continue; |
| |
| for (int m = 0; m < cfg->num_messengers && n_attend < 10240; m++) { |
| |
| pos_list[n_attend] = -1; |
| is_messenger[n_attend] = s * cfg->num_messengers + m; |
| n_attend++; |
| } |
| } |
| } |
|
|
| |
| if (htype == HEAD_GLOBAL && mc && cfg->num_messengers > 0) { |
| for (int s = 0; s < mc->n_filled && n_attend < 10240; s++) { |
| if (s >= mc->segment_capacity) break; |
| if (!mc->segment_filled[s]) continue; |
| |
| int n_neighbor = cfg->messenger_neighbors > 0 ? cfg->messenger_neighbors : 2; |
| int neighbor_start = seg_idx - n_neighbor; |
| if (neighbor_start < 0) neighbor_start = 0; |
| if (s >= neighbor_start && s <= seg_idx) continue; |
| for (int m = 0; m < cfg->num_messengers && n_attend < 10240; m++) { |
| pos_list[n_attend] = -1; |
| is_messenger[n_attend] = s * cfg->num_messengers + m; |
| n_attend++; |
| } |
| } |
| } |
|
|
| if (n_attend == 0) { |
| |
| pos_list[0] = seq_pos; |
| is_messenger[0] = -1; |
| n_attend = 1; |
| } |
| |
| |
| |
| |
| |
| |
| int token_cands = 0; |
| for (int i = 0; i < n_attend; i++) { |
| if (is_messenger[i] >= 0) g_ca_stats.msg_candidates++; |
| else token_cands++; |
| } |
| g_ca_stats.candidates += token_cands; |
|
|
| |
| const float *Q_h = Q + h * head_dim; |
| float scores[10240]; |
| float max_score = -1e30f; |
|
|
| for (int i = 0; i < n_attend; i++) { |
| const float *K_jh; |
| if (is_messenger[i] >= 0) { |
| |
| int msg_idx = is_messenger[i]; |
| K_jh = mc->messenger_k + (size_t)msg_idx * n_embd + h * head_dim; |
| } else { |
| |
| int j = pos_list[i]; |
| int phys_j = j % n_ctx; |
| K_jh = k_cache + (size_t)phys_j * n_embd + h * head_dim; |
| } |
|
|
| |
| if (is_messenger[i] < 0 && cfg->gate_enable) { |
| int j = pos_list[i]; |
| int distance = seq_pos - j; |
| g_ca_stats.gate_pairs++; |
| |
| { |
| int coarse_dim2 = head_dim > 16 ? head_dim / 4 : head_dim; |
| float qn = 0, kn = 0, dt = 0; |
| for (int d = 0; d < coarse_dim2; d++) { |
| dt += Q_h[d] * K_jh[d]; |
| qn += Q_h[d] * Q_h[d]; |
| kn += K_jh[d] * K_jh[d]; |
| } |
| float sim = dt / (sqrtf(qn + 1e-8f) * sqrtf(kn + 1e-8f) + 1e-8f); |
| float dp = 1.0f; |
| if (cfg->gate_distance_prior && distance > 0) { |
| float tau = (float)(cfg->segment_len > 0 ? cfg->segment_len : 64); |
| dp = expf(-(float)distance / tau); |
| } |
| float gs = sim + (cfg->gate_distance_prior ? 0.5f * dp : 0.0f); |
| gate_score_record(gs); |
| } |
| if (!concept_boundary_gate(Q_h, K_jh, head_dim, distance, cfg)) { |
| scores[i] = -1e30f; |
| g_ca_stats.gate_blocked++; |
| continue; |
| } |
| } |
|
|
| float dot = 0.0f; |
| for (int d = 0; d < head_dim; d++) dot += Q_h[d] * K_jh[d]; |
| dot *= scale; |
| scores[i] = dot; |
| if (dot > max_score) max_score = dot; |
| } |
|
|
| |
| float sum_exp = 0.0f; |
| float attn_w[10240]; |
| for (int i = 0; i < n_attend; i++) { |
| float e = expf(scores[i] - max_score); |
| attn_w[i] = e; |
| sum_exp += e; |
| } |
| float inv_sum = 1.0f / (sum_exp + 1e-12f); |
| for (int i = 0; i < n_attend; i++) attn_w[i] *= inv_sum; |
| |
| for (int i = 0; i < n_attend; i++) |
| if (is_messenger[i] >= 0) g_ca_stats.msg_mass += attn_w[i]; |
|
|
| |
| float *out_h = attn_out + h * head_dim; |
| for (int d = 0; d < head_dim; d++) out_h[d] = 0.0f; |
| for (int i = 0; i < n_attend; i++) { |
| if (scores[i] <= -1e29f) continue; |
| const float *V_jh; |
| if (is_messenger[i] >= 0) { |
| int msg_idx = is_messenger[i]; |
| V_jh = mc->messenger_v + (size_t)msg_idx * n_embd + h * head_dim; |
| } else { |
| int j = pos_list[i]; |
| int phys_j = j % n_ctx; |
| V_jh = v_cache + (size_t)phys_j * n_embd + h * head_dim; |
| } |
| float w = attn_w[i]; |
| for (int d = 0; d < head_dim; d++) out_h[d] += w * V_jh[d]; |
| } |
| } |
| } |
|
|
| |
| |
| |
| |
| |
| void attention_backward_concept(float *grad_qkv, const float *grad_attn_out, |
| const float *qkv, int n_embd, int n_head, |
| int seq_pos, |
| const float *k_cache, const float *v_cache, |
| int n_ctx, |
| const ConceptAttnConfig *cfg, |
| MessengerCache *mc) { |
| |
| if (!cfg || !cfg->enable) { |
| attention_backward_sliding(grad_qkv, grad_attn_out, qkv, n_embd, n_head, |
| seq_pos, k_cache, v_cache, n_ctx, |
| g_attn_window > 0 ? g_attn_window : n_ctx, |
| g_attn_sink); |
| return; |
| } |
| (void)n_ctx; |
|
|
| int head_dim = n_embd / n_head; |
| float scale = 1.0f / sqrtf((float)head_dim); |
|
|
| const float *Q = qkv; |
| float *gQ = grad_qkv; |
| float *gK = grad_qkv + n_embd; |
| float *gV = grad_qkv + 2 * n_embd; |
| memset(grad_qkv, 0, 3 * n_embd * sizeof(float)); |
|
|
| int seg_len = cfg->segment_len > 0 ? cfg->segment_len : n_ctx; |
| int seg_idx = seq_pos / seg_len; |
|
|
| for (int h = 0; h < n_head; h++) { |
| HeadAccessType htype = get_head_access_type(h, n_head, cfg); |
| int window = get_head_window(h, n_head, |
| cfg->gate_window > 0 ? cfg->gate_window : 64, |
| cfg); |
| if (window < 1) window = 1; |
| const float *Q_h = Q + h * head_dim; |
| const float *g_out_h = grad_attn_out + h * head_dim; |
|
|
| |
| int n_attend = 0; |
| int pos_list[10240]; |
| int is_messenger[10240]; |
|
|
| int win_start = seq_pos - window + 1; |
| if (win_start < 0) win_start = 0; |
| for (int j = win_start; j <= seq_pos && n_attend < 10240; j++) { |
| pos_list[n_attend] = j; |
| is_messenger[n_attend] = -1; |
| n_attend++; |
| } |
|
|
| int can_msg = head_can_access_messenger(h, n_head, cfg); |
| if (can_msg && mc && cfg->num_messengers > 0) { |
| int n_neighbor = cfg->messenger_neighbors > 0 ? cfg->messenger_neighbors : 2; |
| int neighbor_start = seg_idx - n_neighbor; |
| if (neighbor_start < 0) neighbor_start = 0; |
| for (int s = neighbor_start; s <= seg_idx && n_attend < 10240; s++) { |
| if (s >= mc->segment_capacity) break; |
| if (!mc->segment_filled[s]) continue; |
| for (int m = 0; m < cfg->num_messengers && n_attend < 10240; m++) { |
| pos_list[n_attend] = -1; |
| is_messenger[n_attend] = s * cfg->num_messengers + m; |
| n_attend++; |
| } |
| } |
| } |
|
|
| if (htype == HEAD_GLOBAL && mc && cfg->num_messengers > 0) { |
| for (int s = 0; s < mc->n_filled && n_attend < 10240; s++) { |
| if (s >= mc->segment_capacity) break; |
| if (!mc->segment_filled[s]) continue; |
| int n_neighbor = cfg->messenger_neighbors > 0 ? cfg->messenger_neighbors : 2; |
| int neighbor_start = seg_idx - n_neighbor; |
| if (neighbor_start < 0) neighbor_start = 0; |
| if (s >= neighbor_start && s <= seg_idx) continue; |
| for (int m = 0; m < cfg->num_messengers && n_attend < 10240; m++) { |
| pos_list[n_attend] = -1; |
| is_messenger[n_attend] = s * cfg->num_messengers + m; |
| n_attend++; |
| } |
| } |
| } |
|
|
| if (n_attend == 0) { |
| pos_list[0] = seq_pos; |
| is_messenger[0] = -1; |
| n_attend = 1; |
| } |
|
|
| |
| float scores[10240], w[10240], g_w[10240]; |
| float max_score = -1e30f; |
| for (int i = 0; i < n_attend; i++) { |
| const float *K_jh; |
| if (is_messenger[i] >= 0) { |
| int msg_idx = is_messenger[i]; |
| K_jh = mc->messenger_k + (size_t)msg_idx * n_embd + h * head_dim; |
| } else { |
| int j = pos_list[i]; |
| int phys_j = j % n_ctx; |
| K_jh = k_cache + (size_t)phys_j * n_embd + h * head_dim; |
| } |
| if (is_messenger[i] < 0 && cfg->gate_enable) { |
| int j = pos_list[i]; |
| int distance = seq_pos - j; |
| if (!concept_boundary_gate(Q_h, K_jh, head_dim, distance, cfg)) { |
| scores[i] = -1e30f; |
| continue; |
| } |
| } |
| float dot = 0.0f; |
| for (int d = 0; d < head_dim; d++) dot += Q_h[d] * K_jh[d]; |
| dot *= scale; |
| scores[i] = dot; |
| if (dot > max_score) max_score = dot; |
| } |
|
|
| float sum_exp = 0.0f; |
| for (int i = 0; i < n_attend; i++) { |
| float e = expf(scores[i] - max_score); |
| w[i] = e; sum_exp += e; |
| } |
| float inv = 1.0f / (sum_exp + 1e-12f); |
| for (int i = 0; i < n_attend; i++) w[i] *= inv; |
|
|
| |
| float dot_gw_w = 0.0f; |
| for (int i = 0; i < n_attend; i++) { |
| if (scores[i] <= -1e29f) { g_w[i] = 0.0f; continue; } |
| const float *V_jh; |
| if (is_messenger[i] >= 0) { |
| int msg_idx = is_messenger[i]; |
| V_jh = mc->messenger_v + (size_t)msg_idx * n_embd + h * head_dim; |
| } else { |
| int j = pos_list[i]; |
| int phys_j = j % n_ctx; |
| V_jh = v_cache + (size_t)phys_j * n_embd + h * head_dim; |
| } |
| float g = 0.0f; |
| for (int d = 0; d < head_dim; d++) g += g_out_h[d] * V_jh[d]; |
| g_w[i] = g; |
| dot_gw_w += g * w[i]; |
| } |
|
|
| |
| float g_scores[10240]; |
| for (int i = 0; i < n_attend; i++) { |
| g_scores[i] = (scores[i] <= -1e29f) ? 0.0f : w[i] * (g_w[i] - dot_gw_w); |
| } |
|
|
| |
| float *gQ_h = gQ + h * head_dim; |
| for (int i = 0; i < n_attend; i++) { |
| if (g_scores[i] == 0.0f) continue; |
| const float *K_jh; |
| if (is_messenger[i] >= 0) { |
| int msg_idx = is_messenger[i]; |
| K_jh = mc->messenger_k + (size_t)msg_idx * n_embd + h * head_dim; |
| } else { |
| int j = pos_list[i]; |
| int phys_j = j % n_ctx; |
| K_jh = k_cache + (size_t)phys_j * n_embd + h * head_dim; |
| } |
| float gs = g_scores[i] * scale; |
| for (int d = 0; d < head_dim; d++) gQ_h[d] += gs * K_jh[d]; |
| } |
|
|
| |
| int self_idx = -1; |
| for (int i = 0; i < n_attend; i++) { |
| if (is_messenger[i] < 0 && pos_list[i] == seq_pos) { |
| self_idx = i; |
| break; |
| } |
| } |
| if (self_idx >= 0) { |
| |
| float *gK_h = gK + h * head_dim; |
| float gs = g_scores[self_idx] * scale; |
| for (int d = 0; d < head_dim; d++) gK_h[d] += gs * Q_h[d]; |
|
|
| |
| float *gV_h = gV + h * head_dim; |
| float w_self = w[self_idx]; |
| for (int d = 0; d < head_dim; d++) gV_h[d] += w_self * g_out_h[d]; |
| } |
| } |
| } |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| |
| |
| |
| |
| void model_messenger_caches_alloc(Model *m, const ConceptAttnConfig *cfg) { |
| if (!m || !cfg || !cfg->enable) return; |
| int n_layer = m->cfg.n_layer; |
| if (n_layer <= 0) return; |
|
|
| |
| if (g_messenger_caches) { |
| for (int i = 0; i < g_messenger_caches_n_layer; i++) |
| messenger_cache_free(&g_messenger_caches[i]); |
| free(g_messenger_caches); |
| } |
|
|
| int seg_len = cfg->segment_len > 0 ? cfg->segment_len : m->cfg.n_ctx; |
| int seg_capacity = (m->cfg.n_ctx / seg_len) + 2; |
| g_messenger_caches = (MessengerCache *)calloc(n_layer, sizeof(MessengerCache)); |
| if (!g_messenger_caches) { |
| fprintf(stderr, "[!] model_messenger_caches_alloc: OOM\n"); |
| return; |
| } |
| g_messenger_caches_n_layer = n_layer; |
| for (int i = 0; i < n_layer; i++) { |
| messenger_cache_alloc(&g_messenger_caches[i], seg_capacity, |
| cfg->num_messengers, m->cfg.n_embd); |
| } |
| printf("[*] messenger caches allocated: %d layers, cap=%d segments/layer, S=%d messengers/segment\n", |
| n_layer, seg_capacity, cfg->num_messengers); |
| } |
|
|
| void model_messenger_caches_free(void) { |
| if (!g_messenger_caches) return; |
| for (int i = 0; i < g_messenger_caches_n_layer; i++) |
| messenger_cache_free(&g_messenger_caches[i]); |
| free(g_messenger_caches); |
| g_messenger_caches = NULL; |
| g_messenger_caches_n_layer = 0; |
| } |
|
|
| void model_messenger_caches_reset(void) { |
| if (!g_messenger_caches) return; |
| for (int i = 0; i < g_messenger_caches_n_layer; i++) |
| messenger_cache_reset(&g_messenger_caches[i]); |
| } |
|
|
| |
| void model_set_concept_attn(Model *m, const ConceptAttnConfig *cfg) { |
| if (!m || !cfg) return; |
| g_concept_attn_cfg = *cfg; |
| if (cfg->enable) { |
| model_messenger_caches_alloc(m, cfg); |
| printf("[*] concept-aware attention enabled: seg_len=%d, S=%d, neighbors=%d, " |
| "gate=%d(threshold=%.3f, fallback=%.4f), hetero=%d(local=%d, msg=%d)\n", |
| cfg->segment_len, cfg->num_messengers, cfg->messenger_neighbors, |
| cfg->gate_enable, cfg->gate_threshold, cfg->gate_fallback_prob, |
| cfg->hetero_enable, cfg->n_local_heads, cfg->n_messenger_heads); |
| } else { |
| model_messenger_caches_free(); |
| printf("[*] concept-aware attention disabled (fallback to standard attention)\n"); |
| } |
| } |
|
|
| |
| void concept_attn_probe_print(void) { |
| ConceptAttnStats *s = &g_ca_stats; |
| |
| |
| |
| |
| long total_fwd = s->forwards + s->forwards_ctx; |
| if (total_fwd == 0) { |
| printf(" [CATTN] fwd=0 ⚠ 概念注意力未参与前向 | enable=%d caches=%s cctx_enable=%d cg=%s\n", |
| g_concept_attn_cfg.enable, g_messenger_caches ? "OK" : "NULL", |
| g_cctx_cfg.enable, g_runtime_cg ? "loaded" : "NULL"); |
| concept_attn_stats_reset(); |
| return; |
| } |
| double reduction = s->full_equiv > 0 ? |
| 100.0 * (1.0 - (double)s->candidates / (double)s->full_equiv) : 0.0; |
| double gate_rate = s->gate_pairs > 0 ? |
| 100.0 * (double)s->gate_blocked / (double)s->gate_pairs : 0.0; |
| double msg_share = s->candidates > 0 ? |
| 100.0 * (double)s->msg_candidates / (double)s->candidates : 0.0; |
| double msg_mass_per_head = (s->forwards + s->forwards_ctx) > 0 ? |
| s->msg_mass / (double)(s->forwards + s->forwards_ctx) : 0.0; |
| double msg_cos = s->msg_segments > 0 ? |
| s->msg_inter_cos / (double)s->msg_segments : 0.0; |
| double msg_norm = s->msg_segments > 0 ? |
| s->msg_norm / (double)s->msg_segments : 0.0; |
| |
| const char *cos_tag = msg_cos < 0.2f ? "OK" : (msg_cos < 0.35f ? "改善中" : "同质化!"); |
| printf(" [CATTN] fwd_simple=%ld fwd_ctx=%ld 候选精简=%.1f%% 门控屏蔽=%.1f%% 信使候选=%.1f%% 信使质量=%.3f/头 片段=%d\n", |
| s->forwards, s->forwards_ctx, reduction, gate_rate, msg_share, msg_mass_per_head, s->last_n_filled); |
| |
| if (s->forwards_ctx > 0) { |
| double avg_attend = (double)s->ctx_total_attend / (double)s->forwards_ctx; |
| double avg_slots = (double)s->ctx_memory_slots_used / (double)s->forwards_ctx; |
| printf(" [CATTN-CTX] 平均候选/前向=%.1f (含概念槽=%.2f, sink+window=%.1f) 槽命中率=%.1f%%\n", |
| avg_attend, avg_slots, avg_attend - avg_slots, |
| avg_attend > 0 ? 100.0 * avg_slots / avg_attend : 0.0); |
| } |
| printf(" [CATTN-PROBE] 信使间余弦=%.3f(%s,目标<0.2) 信使均范数=%.2f(钳制4.0) 统计片段=%ld\n", |
| msg_cos, cos_tag, msg_norm, s->msg_segments); |
| concept_attn_stats_reset(); |
| } |
|
|
|
|