lalmodel-code / src /runtime /lal_runtime.c
gasschina's picture
feat: PonderNet 循环思考 (逐层停机 + 末块循环 双重混合)
1423b78
Raw
History Blame Contribute Delete
300 kB
/* lal_runtime.c — LAL Universal Runtime implementation
*
* Three API levels:
* Level 1: operators (bin_forward, norm, gelu, etc.)
* Level 2: transformer layer (trans_layer_forward/backward)
* Level 3: full model (model_load/forward/backward)
*
* Models only need Level 3 — just config + weight key patterns.
*/
/* === PonderNet 循环思考: 实现体唯一定义在本翻译单元 ===
* 注意: 必须在 include lal_runtime.h (间接 include lal_ponder.h) 之前定义,
* 否则 include guard 会把实现段挡掉 (单头库规则) */
#define LAL_PONDER_IMPLEMENTATION
#include "lal_runtime.h"
#include "lal_whitebox_probe.h"
#include "lal_concept_gen.h"
#include "lal_concept_attn.h"
/* v2: 对齐分配 */
#ifdef _WIN32
#include <malloc.h>
#else
#include <stdlib.h>
#endif
/* [加速] OpenBLAS 条件编译: Makefile 检测到 libopenblas 时定义 HAVE_OPENBLAS,
* CORE 路径的 matmul 用 cblas_sgemm 一次性算所有 CORE 行 (AVX2/AVX-512 + 多线程).
* 没装 OpenBLAS 时退回原 OpenMP + 8 倍展开循环. */
#ifdef HAVE_OPENBLAS
#include <cblas.h>
/* OpenBLAS 默认用自己的线程池, 和 OpenMP 的线程池冲突会 segfault.
* 强制 OpenBLAS 单线程, 只用 OpenMP 并行 (避免线程竞争).
* 用 static flag 在第一次 bin_forward 调用时初始化 (constructor 在 MSYS2 不稳). */
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
/* This project is pure-CPU, no GPU. The old LAL_CUDA backend
* (runtime/lal_cuda.cu / lal_cuda.h) has been removed. Any remaining
* '#ifdef LAL_CUDA' blocks below are dead code and never compiled
* (the Makefile / build.ps1 never define LAL_CUDA). */
/* === Windows/MinGW compatibility ===
* MinGW lacks sys/mman.h and rand_r(). We provide shims so the same
* lal_runtime.c compiles on both Linux and Windows/MinGW64. */
#ifdef _WIN32
#define WIN32_LEAN_AND_MEAN
#include <windows.h>
#include <io.h>
/* mmap shim: use CreateFileMapping on Windows */
#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;
}
/* rand_r shim: MinGW lacks it, use rand() with thread-local seed */
static inline int rand_r(unsigned int *seedp) {
*seedp = *seedp * 1103515245u + 12345u;
return (int)((*seedp / 65536u) % 32768u);
}
/* fstat/stat shim: MinGW has them in sys/stat.h but with different struct */
#include <sys/stat.h>
#define fstat _fstat
#define stat _stat
#else
#include <sys/mman.h>
#include <sys/stat.h>
#include <unistd.h>
#endif
/* Forward declarations for full-vocab softmax (defined later in this file,
* but model_forward/model_backward call them — declared here to avoid
* implicit-declaration errors since the definitions sit after the callers). */
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);
/* ========================================================================
* Level 1 additions: RMSNorm, SiLU, dispatch functions, RoPE
* ======================================================================== */
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); /* gate * silu(up), caller handles gate */
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) {
/* Simplified RoPE: rotate pairs by position-dependent angle */
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;
}
}
}
/* ========================================================================
* Level 2: Transformer Layer (building block)
* ======================================================================== */
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; /* v16: 概念注意力信使缓存索引 */
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) {
/* GPT-2: merged QKV [n → 3n] */
sprintf(full_key, qkv_key, layer_idx);
float *W = tensor_get(tensors, n_tensors, full_key);
/* Key format is like "h.%d.attn.c_attn.weight"; bias key replaces the
* ".weight" suffix with ".bias".
* (Removed a `sprintf(full_key, "%s.bias", full_key)` here: src and dst
* overlapped, which is undefined behaviour, and it was dead anyway
* because the suffix swap below rebuilds the key from scratch.) */
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 {
/* LLaMA/Qwen: separate Q, K, V, O */
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);
}
/* Output projection */
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);
/* MLP */
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 {
/* GELU: single c_fc */
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);
/* Norm weights */
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);
}
/* Dispatch: pure float 为唯一前向路径 (BNN 快速路径已移除) */
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);
}
/* KV-cache-only forward: fill _kv_k/_kv_v for a CONTEXT position without
* computing attention output, output projection, or MLP. Context positions
* only need their K/V stored in the cache (they are constants during backward,
* their output is discarded), so skipping the ~50% of FLOPs spent on attn_o +
* MLP yields a large speedup in model_forward's context prefill loop. */
/* Pure-float forward: same as trans_layer_forward but uses bin_forward_pure_float
* for every matmul (no sign binarization anywhere). Used by the teacher model
* in distillation — w_float holds original GPT-2 weights, never updated.
* Activations cache is shared with the student's structure (same shape) so we
* can reuse m->acts. NOTE: this does NOT overwrite student activations if
* called on a separate teacher Model (m->acts is per-model). */
/* Global flag: use STE backward (updates w_float + repacks wbits) */
/* ========================================================================
* v: Data-parallel (per-thread) batch training support.
* The batch loop (for b in batch_size) is parallelized across OpenMP
* threads. Every thread-local transient buffer previously declared as a
* function-scope `static` is moved into a per-thread ThrRes slot indexed
* by g_cur_tid, so concurrent samples never clobber each other. Gradient
* accumulators are also kept per-thread and reduced into the real
* grad_accum after the parallel region.
* ======================================================================== */
#include <omp.h>
#define LAL_MAX_THREADS 16
int g_cur_tid = 0;
#pragma omp threadprivate(g_cur_tid)
typedef struct {
TransAct *acts; /* n_layer activation buffers */
TransAct *scratch; /* context-prefill scratch (replaces get_scratch_acts) */
/* === PonderNet 循环思考 per-thread 缓冲 (g_ponder_cfg.enable 时分配) === */
PonderBuf ponder; /* 停机分布/损失缓冲 */
float *ponder_mix; /* [n_embd] 混合读出累积 */
float *ponder_state; /* [LAL_PONDER_MAX_STEPS][n_embd] 各步状态缓存 */
float *ponder_kv0k; /* [n_embd] 末块迭代0 K 快照 (cache 恢复用) */
float *ponder_kv0v; /* [n_embd] 末块迭代0 V 快照 */
TransAct *rec_acts; /* [rec_iters] 末块迭代 act 快照 (训练反向用) */
int ponder_first_rec_step; /* 末块循环步的起始 step 索引 */
int ponder_ready;
float *mlp, *hidden, *norm2, *proj, *attn, *qkv, *norm1, *pre; /* 16384 */
float *gate, *up, *norm2_gate, *norm2_up; /* 16384 */
float *n1k, *n1v; /* 4096 */
float *xc, *x, *gh; /* 4096 */
float *g_pre4; /* 4096 (model_backward) */
float *full_logits; int full_logits_vocab;
int forward_done; /* v21: forward 已写入 full_logits(softmax probs), backward 可复用, 省一次重算 */
float *final_ln;
float *x_before_final; float final_mean, final_std_inv;
/* per-layer per-binlayer gradient pools (parallel accumulators) */
float ***grad_w; /* [n_layer][n_bl] -> float[in*out] */
float ***grad_b; /* [n_layer][n_bl] -> float[out] */
float *grad_wte, *grad_wpe, *grad_lnfw, *grad_lnfb;
/* per-layer norm gradients */
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);
/* Shared sign lookup table: maps an 8-bit sign word (bit i set => +1) to the
* 8 float signs. Used by both bin_forward (ternary/BWN) and bin_backward_ste
* so forward and backward agree on the quantized weights they differentiate
* through. Initialized once. */
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; /* 固化: STE 模式 (train=infer, 直接学二值逻辑) */
float g_attn_residual_scale = 1.0f; /* 固化: 注意力残差缩放 = 1.0 (真实注意力) */
int g_use_logic_binarization = 1; /* 固化: 逻辑引导层 (CORE/BINARY/PRUNE 语义结构) */
/* Semantic logic mask ratios (set by training script per curriculum phase).
* When g_logic_core_ratio > 0, compute_norm_mask uses these instead of
* the hardcoded 20%/70%/10% split. This enables progressive activation:
* early stages are sparse (high PRUNE), later stages are dense. */
float g_logic_core_ratio = 0.0f; /* 0 = use default 20% */
float g_logic_prune_ratio = 0.0f; /* 0 = use default 10% */
/* Adam optimizer globals (used inside bin_backward_ste when g_use_adam=1).
* Defaults are standard Adam (Kingma & Ba 2015).
* g_opt_step is incremented per model_backward call to drive bias correction. */
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;
/* Ternary Weight Network (TWN) globals.
* When g_use_ternary=1, BINARY rows use {-1,0,+1}: |W|<=Δ is zeroed (Δ stored
* per-layer in BinLayer.ternary_delta). Triples capacity vs BWN at ~1.58 bits. */
/* 固化: 三值权重默认开启 — 当前 ckpt (model_dialogue.ste / ckpt_mp_*) 全部是
* ternary 训练产物。默认关掉会导致按 BWN 解释权重 → 输出乱码。
* 若确需浮点/BWN, 显式传 --no-ternary。 */
int g_use_ternary = 1;
float g_ternary_delta_factor = 0.7f; /* Δ = factor * mean(|W_row|), TWN default */
/* Checkpoint fusion strategy for --merge (see merge_models in ste_train.c). */
int g_merge_mode = 0; /* 0 = step-weighted avg, 1 = EMA by step */
float g_merge_beta_lo = 0.5f; /* EMA fold-in coef for first (least-trained) model */
float g_merge_beta_hi = 0.9f; /* EMA fold-in coef for last (most-trained) model */
/* Cosine LR with linear warmup.
* step < warmup : lr = base * (step+1) / warmup (linear ramp from 0)
* warmup <= step < total : lr = base * 0.5 * (1 + cos(pi * progress)) (cosine)
* step >= total : lr = base * 0.01 (floor — keep updating)
* Warmup tames the early-step gradient explosion (STE on bit-space is noisy).
* Cosine decay reduces late-step oscillation for convergence.
* Pass warmup=0 to disable warmup, total=0 to disable decay. */
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; /* degenerate: no decay */
if (step >= total_steps) return base_lr * 0.01f; /* floor */
float progress = (float)(step - warmup_steps) / (float)(total_steps - warmup_steps);
return base_lr * 0.5f * (1.0f + cosf((float)M_PI * progress));
}
/* Pure float forward: y[j] = sum_i w_float[j*in+i] * x[i] + bias[j].
* Skips sign binarization entirely. Used for the teacher model in
* distillation — the teacher's w_float holds the original GPT-2 weights
* and is never updated, so this is a faithful full-precision matmul.
* Logic-guided layers: CORE uses w_core (already float), BINARY uses w_float,
* PRUNE outputs 0 (skipped). */
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) {
/* v16-perf: 前缀索引 + OpenMP 并行
* v2-perf: 小矩阵 (out < 64) 串行, 避免 fork/join 开销
* 大矩阵用 guided schedule 动态负载均衡 */
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) { /* CORE: float dot with w_core[cidx[j]] (never quantized) */
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) { /* BINARY */
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) {
/* Ternary QAT forward — MUST match bin_forward() case 1 exactly:
* y = (Σ sign(w_float[i]) * x[i] over NON-ZEROED positions) *
* alpha[j] * K * g_binary_scale + bias[j]
* The old pure-float path used w_float directly (un-sign, no
* alpha, no K) which made BINARY outputs ~1/alpha times too
* large → generation gibberish on ternary checkpoints. */
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; /* ternary-0 */
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 {
/* Plain BWN / pure-float BINARY path (non-ternary teacher). */
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; /* PRUNE */
}
}
} 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 {
/* 无 logic_mask: 全 float matmul, 大矩阵并行 */
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;
}
}
}
}
/* Auto-generate per-output logic mask based on weight norms.
* W is [in, out] (GPT-2 Conv1D format). We compute per-output column norms.
* top 20% → CORE (0), bottom 10% → PRUNE (2), middle 70% → BINARY (1).
* mask: [out_dim] bytes, 0=CORE, 1=BINARY, 2=PRUNE. */
static void compute_norm_mask(const float *W, int in_dim, int out_dim, uint8_t *mask) {
/* Compute per-output norms (W is [in, out] row-major) */
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);
}
/* Find thresholds via partial sort (simple: sort a copy) */
float *sorted = malloc(out_dim * sizeof(float));
memcpy(sorted, norms, out_dim * sizeof(float));
/* Simple insertion sort (out_dim ≤ 3072, OK) */
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;
}
/* Use semantic ratios when set, otherwise default 20%/10% */
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;
/* sorted[0] = smallest norm, sorted[out_dim-1] = largest */
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; /* CORE */
n_core++;
} else if (norms[j] <= prune_threshold && n_prune < prune_count) {
mask[j] = 2; /* PRUNE */
n_prune++;
} else {
mask[j] = 1; /* BINARY */
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; /* CORE neurons learn 3x faster than BINARY */
int g_use_lal_adam = 1; /* 1=group-wise Adam (LAL-aware), 0=standard per-param Adam */
float g_prune_decay = 0.01f; /* PRUNE weight decay per step (pulls toward 0) */
float g_prune_freeze_thresh = 0.001f; /* PRUNE neurons below this are frozen */
/* Global flag: use real causal multi-head self-attention with KV cache.
* Off by default — backward compat with V-copy. When on, trans_layer_forward
* calls attention_forward() instead of memcpy(act->attn_out, act->v, n). */
int g_use_real_attention = 0;
int g_skip_wv = 0; /* v13l: skip W_v projection, use norm1_out as attn_out */
/* 滑动窗口注意力 (长上下文训练固化): 默认 0 = 全因果(兼容旧行为).
* >0 时每个 token 只看前 window 个 token + 前 sink 个全局锚点 token.
* 这是支撑 8192-token 长文训练的唯一可行路径(全因果 O(n^2) 会 OOM). */
/* === 固化: 滑动窗口注意力默认开启 ===
* 长上下文路径 (max_pos=8192) 必须靠滑动窗口把 O(n^2) 降到 O(n*w),
* 否则 8192 全注意力既慢又爆内存。默认值必须与 train_4core.ps1 一致,
* 且训练/推理两侧窗口必须相同, 否则生成结果乱码 (历史 Bug: 推理 9996 vs 训练 1024)。
* 若确需全注意力, 显式传 --attn-window 0。 */
int g_attn_window = 1024;
int g_attn_sink = 64;
int g_use_pure_float = 0;
/* v16: BINARY 共模抑制系数 (白盒: BINARY 能量~18x CORE 但区分度≈0, 0.25=能量均衡) */
float g_binary_scale = 0.25f;
/* v16: wte/wpe 更新速率系数 (对齐泵降速) */
float g_wte_lr_scale = 0.5f; /* v16 原 0.1: wte 对齐泵降速过度 → embedding 几乎不更新 →
所有 token 向量趋同 (VDIVERSE n_collapsed=42/42)。
提到 0.5 让 embedding 有效分化, 修复生成乱码塌缩。 */
/* v16: logit 缩放 (残差范数小→softmax近均匀→CE梯度稀释, 放大锐化分布) */
float g_logit_scale = 1.0f;
int g_accumulate_gradients = 0; /* 1 = accumulate grads, don't update weights */
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));
/* k/v alias into the contiguous Q|K|V buffer so both merged (GPT-2)
* and separate (LLaMA/Qwen) paths share one [3n] layout. Previously
* k/v were left NULL for the separate path → segfault. */
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));
/* BUG #45 FIX: allocate SwiGLU gate/up cache (NULL for GELU mode) */
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);
}
/* ========================================================================
* Level 3: Full Model
* ======================================================================== */
/* ----- Causal Multi-Head Self-Attention (KV cache) -----
* Replaces the degenerate V-copy in trans_layer_forward.
* Mirrors tools/server/gpt2_server.c:real_attention (scalar version).
*
* Layout:
* qkv: [3 * n_embd] — Q | K | V concatenated, single token
* k_cache_layer / v_cache_layer: [n_ctx * n_embd] — filled position-by-position
* attn_out: [n_embd] — output, weighted sum of V across heads
*
* Causal: position seq_pos attends only to positions 0..seq_pos (inclusive).
* Multi-head: n_head heads, head_dim = n_embd / n_head (must divide evenly).
*/
/* ----- Attention backward (dQ/dK/dV) -----
* Computes gradients for the current token's Q, K, V. Cached K/V at positions
* 0..seq_pos-1 are treated as constants (they are context, not learned here —
* only the current token's QKV projection receives gradient, matching the
* single-position activation cache used by model_forward/backward).
*
* Per head h (head_dim d, scale = 1/sqrt(head_dim)):
* forward: scores[j]=Q·K_j*scale; w=softmax(scores); out=sum_j w[j]*V_j
* backward:
* g_w[j] = <g_out, V_j> (grad w.r.t. weight j)
* g_scores[j] = w[j] * (g_w[j] - <g_w, w>) (softmax bwd)
* g_Q[d] += sum_j g_scores[j] * K_j[d] * scale
* g_K_cur[d] += g_scores[seq_pos] * Q[d] * scale (current K only)
* g_V_cur[d] += w[seq_pos] * g_out[d] (current V only)
*/
void model_kv_cache_alloc(Model *m) {
if (m->k_cache) return; /* idempotent */
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);
/* Wire into TransLayer so trans_layer_forward can find them */
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;
}
/* FIX: get-or-realloc a thread-local scratch TransAct buffer that tracks
* the model's current config. Previously this was a static pointer
* allocated once for the first model and never updated — on phase switch
* (n_embd change) the scratch was too small, causing heap-buffer-overflow
* in trans_layer_forward's memcpy. */
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); /* frees inner arrays + scratch itself */
scratch = NULL; /* trans_act_free already freed scratch; avoid double-free */
}
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;
/* Single source of truth for the attention window: the GLOBAL
* g_attn_window / g_attn_sink flags (set by --attn-window/--attn-sink,
* default 1024/64). Sync them into cfg so any code reading
* ModelConfig.sliding_window / n_sinks (e.g. stateful inference) matches
* training exactly. ModelConfig.sliding_window defaults to 9996 and is NOT
* a valid inference window, so we overwrite it here. */
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)); /* FIX: calloc (not malloc) zero-inits BinLayer fields like grad_accum so model_batch_alloc's NULL check works */
m->acts = trans_act_alloc(&cfg);
/* Build keys and binarize each layer */
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;
/* Helper: bin_layer_init or bin_layer_init_logic depending on flag */
#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)
/* BUG #54 FIX: Attention 层不做 logic binarization!
* 根因: QKV merged 模式下, Q/K 的梯度淹没 V 的梯度, 导致 W_v 退化为 rank-1.
* (SVD: 最大奇异值 4.548 vs 第二大 1.116, effective rank 276/530)
* 修复: attention 的 Q/K/V/O 用普通 bin_layer_init (无 CORE/BINARY/PRUNE),
* 只有 MLP 层用 logic binarization. 这样 W_v 能正常学习. */
#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
/* Norm weights */
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");
/* Free large weight matrix tensor data after binarization to save ~3.6GB.
* Small tensors (wte, wpe, ln_f, per-layer norms) are kept for forward pass.
* bin_layer_init copies all needed data into w_float/wbits/alpha/bias. */
{
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;
/* Auto-allocate KV cache if real attention is requested at load time.
* Callers can also call model_kv_cache_alloc() later to enable it. */
if (g_use_real_attention) model_kv_cache_alloc(m);
thr_res_alloc(m); /* per-thread training buffers (also used by diagnostics) */
/* 唯一路线底座:概念感知注意力默认随模型加载即启用(CORE/BINARY/PRUNE + 浮点 + 概念注意力)。
* 修复根因:推理侧此前 g_messenger_caches==NULL, 永远走标准 attention (探针 fwd=0)。
* LAL_CONCEPT_ATTN=0 作为调试逃生口强制关闭;其余环境变量(LAL_CA_SEG_LEN/LAL_CA_MSG)
* 由调用方在 model_set_concept_attn 时覆盖。训练侧 ste_train.c 会二次调用并覆盖分段长度。 */
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; /* 显式 0 才关,否则默认开 */
model_set_concept_attn(m, &cca);
}
}
/* Compute full vocab logits at target position using pure float forward.
* Replaces bin_forward with bin_forward_pure_float for one pass (no
* binarization anywhere). The result is the teacher signal for distillation.
* Caller must allocate logits_out[vocab_size]. */
void model_forward_float_logits(Model *m, const int *tokens, int n_tokens,
float *logits_out) {
/* 端到端统一: 诊断也用 sliding window forward (与训练/推理同路径).
* 旧版用 trans_layer_forward_pure_float (独立路径) 已废弃.
* 返回 full vocab logits 供诊断打印. */
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;
}
}
/* Backward with distillation: hard CE (target) + soft KL(teacher || student).
* The KL gradient w.r.t. student logits[j] is:
* d_KL/d_s[j] = T * (softmax(s/T)[j] - softmax(t/T)[j])
* Then w.r.t. final_ln[i]:
* d_KL/d_final_ln[i] = sum_j (T * (ps[j]-pt[j])) * wte[j*n+i]
*
* Combined grad on final_ln:
* gh[i] = alpha * CE_grad[i] + (1-alpha) * T^2 * KL_grad[i]
* (T^2 because KL of T-scaled soft targets is conventionally multiplied by T^2
* to keep gradient magnitude roughly constant across T.)
*
* Teacher logits must be full vocab (computed by model_forward_float_logits).
* Memory cost: ~3*vocab*sizeof(float) = 600KB scratch (heap-allocated here). */
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);
}
/* ========================================================================
* PonderNet 循环思考 — Model 级实现
* ======================================================================== */
int ponder_step_count(const Model *m) {
/* 总步数 (含末步 remainder):
* layer_halt: 前 (n_layer-1) 个层步 + 末块步(1 或 rec_iters)
* !layer_halt: 仅末块步 (rec_iters ≥ 2, 否则 enable=0) */
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);
/* 推理侧末块迭代 act 快照 (训练侧用 g_thr[tid].rec_acts) */
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) {
/* 停机单元 Adam 更新 (与 BinLayer 同式的 bias-correction Adam)
* lr = CE lr × g_ponder_cfg.lr_scale; 梯度已在前向/反向中直接累加 (串行训练) */
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;
/* 梯度范数钳制 (停机单元敏感, 单元级 clip 1.0) */
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;
}
/* Adam 步 */
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;
}
}
/* ========================================================================
* Binary Weight Layer
* ======================================================================== */
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; /* allocated only in ternary mode (logic-guided path) */
bl->w_core = NULL; /* allocated only in logic-guided path; NULL → free() is a no-op */
bl->logic_mask = NULL; /* set only in logic-guided path; NULL → free() is a no-op */
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)); /* STE */
bl->m_adam = g_use_adam ? calloc((size_t)in_dim * out_dim, sizeof(float)) : NULL; /* Adam m (conditional) */
bl->v_adam = g_use_adam ? calloc((size_t)in_dim * out_dim, sizeof(float)) : NULL; /* Adam v (conditional) */
bl->grad_accum = calloc((size_t)in_dim * out_dim, sizeof(float)); /* batch grad accumulation */
bl->bias_grad_accum = calloc((size_t)out_dim, sizeof(float)); /* batch bias grad accumulation */
bl->ternary_delta = 0.0f; /* BWN by default; set by bin_layer_repack_ternary */
/* Copy float weights for STE updates — TRANSPOSE to [out, in] layout!
* W is [in, out] row-major (GPT-2 Conv1D format). We store w_float as
* [out, in] so that w_float[j*in + i] is contiguous per output j.
* This makes repack/alpha/update loops all contiguous → SIMD-friendly. */
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];
/* Row-major: pack sign(w[j][i]) per output 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;
}
}
/* Col-major (transposed): pack sign(w[j][i]) per input i */
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;
}
}
}
/* Logic-guided binarization: initialize with per-output logic_mask.
* mask[j]: 0=CORE (keep float), 1=BINARY (sign+alpha), 2=PRUNE (zero).
*
* This implements PHONE's "logic extraction at binarization time":
* - CORE outputs: weights stored as float in w_core, NOT binarized
* - BINARY outputs: sign(w) packed into wbits, alpha = mean(|w|)
* - PRUNE outputs: wbits all zero, alpha=0, bias=0 (effectively removed)
*
* The forward pass (bin_forward) checks logic_mask per output:
* - CORE: y[j] = x @ w_core[j] (float matmul, no binarization)
* - BINARY: y[j] = alpha * (2*popcount - N) + bias (XNOR+popcount)
* - PRUNE: y[j] = 0 (skipped entirely)
*/
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; /* allocated below only if we keep this layer */
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; /* FIX: must NULL-init; model_batch_alloc checks !grad_accum */
bl->bias_grad_accum = NULL; /* FIX: same — otherwise random heap value passes the check */
if (!logic_mask) {
/* No logic mask → free w_float (bin_layer_init will re-alloc) and delegate. */
free(bl->w_float); bl->w_float = NULL;
bin_layer_init(bl, W, bias, in_dim, out_dim);
return;
}
/* We're keeping this layer — allocate Adam state for the BINARY/CORE
* STE-update path. PRUNE outputs contribute zero, but they still own a
* slot in w_float (so the [out*in] indexing stays uniform). */
bl->m_adam = calloc((size_t)out_dim * in_dim, sizeof(float));
bl->v_adam = calloc((size_t)out_dim * in_dim, sizeof(float));
/* Copy logic mask + count categories */
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++;
}
/* Allocate w_core for CORE outputs (float weights, [n_core, in_dim]) */
if (bl->n_core > 0) {
bl->w_core = malloc((size_t)bl->n_core * in_dim * sizeof(float));
}
/* Process each output based on its logic category */
int core_idx = 0;
for (int j = 0; j < out_dim; j++) {
const float *wj = &W[j * in_dim]; /* W is [out, in] (transposed) */
switch (logic_mask[j]) {
case 0: /* CORE: keep float */
memcpy(&bl->w_core[core_idx * in_dim], wj, in_dim * sizeof(float));
bl->alpha[j] = 0.0f; /* not used for CORE */
if (bias) bl->bias[j] = bias[j];
/* wbits for CORE: all zero (not used, but keep for indexing) */
core_idx++;
break;
case 1: /* BINARY: sign(w) + alpha */
{
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: /* PRUNE: zero out */
bl->alpha[j] = 0.0f;
bl->bias[j] = 0.0f;
/* wbits already zero from calloc */
break;
}
/* Copy to w_float (transposed [out, in] for STE compatibility) */
memcpy(&bl->w_float[j * in_dim], wj, in_dim * sizeof(float));
}
/* Build wbits_T (transposed) only for BINARY outputs */
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;
}
}
/* Ternary mode: allocate zbits (zero mask, same shape as wbits) and
* compute initial ternary binarization from w_float. When ternary is off,
* zbits stays NULL — bin_forward uses the pure BWN ±1 path. */
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; /* FIX: was missing, causing wild-pointer crash in model_batch_begin after model_free + model_batch_alloc on phase switch */
}
/* Re-pack wbits and wbits_T from sign(w_float).
* w_float is [out, in] (transposed from Conv1D's [in, out] for contiguous
* per-output access). All loops here are now contiguous → auto-vectorizable.
*
* Key optimization: wbits[j] packs sign(w_float[j*in + 0..in-1]) which is
* contiguous memory. The compiler auto-vectorizes the 8x unrolled comparison
* into SIMD compare + movemask-style bit extraction. */
void bin_layer_repack(BinLayer *bl) {
int in = bl->in_dim, out = bl->out_dim;
/* === CRITICAL FIX: Sync w_core from w_float for CORE neurons ===
* CORE neurons use w_core (float) in forward pass, but model_batch_apply
* updates w_float. Without this sync, CORE weights are FROZEN and
* CORE/BINARY differentiation never improves. */
if (bl->w_core && bl->logic_mask) {
int core_idx = 0;
for (int j = 0; j < out; j++) {
if (bl->logic_mask[j] == 0) { /* CORE */
memcpy(&bl->w_core[core_idx * in],
&bl->w_float[j * in],
in * sizeof(float));
core_idx++;
}
}
}
/* Pack wbits[j][wi] from sign(w_float[j*in + i]) — CONTIGUOUS in i! */
for (int j = 0; j < out; j++) {
const float *wf = &bl->w_float[j * in]; /* contiguous [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) {
/* 8 contiguous floats — compiler auto-vectorizes to SIMD */
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;
}
}
/* Skip wbits_T repack in STE mode — grad_x is now computed from w_float
* directly (float arithmetic), so wbits_T is never read during STE training.
* This saves the strided wbits_T repack loop (the slowest part). */
for (int j = 0; j < out; j++) {
const float *wf = &bl->w_float[j * in]; /* contiguous [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;
}
}
/* Ternary repack: recompute zbits (zero mask) from |w_float| vs Δ.
* For each BINARY output row j:
* Δ_j = g_ternary_delta_factor * mean(|w_float[j]|)
* zbits[j][i] = 1 if |w_float[j*in+i]| <= Δ_j (weight zeroed → ternary 0)
* zbits[j][i] = 0 otherwise (weight active → ±1)
* Also updates alpha[j] = mean(|w|) over ACTIVE weights only (standard TWN
* scaling: the zeroed weights contribute nothing, so scaling reflects the
* active subset). CORE/PRUNE rows are skipped (zbits stays 0 there).
*
* wbits (sign) is NOT recomputed here — bin_layer_repack (called separately
* for STE) keeps sign in sync. This function only updates the zero mask.
* Called at init and after every STE step when g_use_ternary is on. */
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++) {
/* Skip non-BINARY rows — they don't use ternary (CORE=float, PRUNE=0). */
if (bl->logic_mask && bl->logic_mask[j] != 1) continue;
const float *wf = &bl->w_float[j * in];
/* Compute per-row Δ = factor * mean(|w|). */
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; /* store per-layer (last row wins, used for stats) */
/* Pack zbits[j]: 1 where |w| <= delta. Also sum active |w| for alpha. */
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;
}
/* TWN alpha: mean(|w|) over active weights. Falls back to mean over all
* if everything got zeroed (degenerate row). */
bl->alpha[j] = (n_active > 0) ? (active_abs_sum / n_active) : mean_abs;
}
}
/* ========================================================================
* Binary Forward — BWN (default, matches Python STE training)
* ========================================================================
* x stays float. Only W is binarized (sign(W) * alpha).
* Adds XNOR-Net K-norm: K = ||x||_1 / in_dim, preserves input magnitude.
*
* y[j] = (sum_i sign(W[j,i]) * x[i]) * alpha[j] * K + bias[j]
*
* This is the mathematically correct BWN forward. The old bin_forward was
* BNN (binarized x too) which diverged from training and caused quality
* collapse. BNN is retained as bin_forward_bnn() for opt-in fast mode.
* ======================================================================== */
void bin_forward(float *y, const float *x, const BinLayer *bl) {
int in = bl->in_dim, out = bl->out_dim, nw = bl->n_words;
/* Logic-guided: if logic_mask exists, dispatch per-output */
if (bl->logic_mask) {
/* K-norm for BINARY outputs */
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();
/* v16-perf: 预计算 CORE 行索引前缀, 解除循环依赖后可 OpenMP 并行 */
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
/* [加速] CORE 路径用 cblas_sgemm 一次性算所有 CORE 行.
* w_core 是 [n_core, in] 行主序, x 是 [in], 结果是 [n_core].
* 用 sgemm: C[1, n_core] = 1.0 * A[1, in] @ B[in, n_core] + 0.0
* 其中 A = x (1×in), B = w_core^T (in×n_core, 但 w_core 是 n_core×in 行主序,
* 所以 B = w_core 用 CblasTrans), C = raw_dots (1×n_core).
* 然后 per-row: y[j] = raw_dots[cidx[j]] * core_gain[j] * K + bias[j]. */
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(); /* 首次调用: 强制 OpenBLAS 单线程, 避免与 OpenMP 冲突 */
/* 用 static buffer 避免 per-call malloc; 大小 = n_core * sizeof(float) */
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 算 y = alpha * A @ x + beta * y
* A = w_core [n_core, in] 行主序, x = x [in], y = core_dots [n_core]
* sgemm M=1 时 sgemv 更高效 (专门优化的 GEMV 路径, 无需转置) */
cblas_sgemv(CblasRowMajor, CblasNoTrans,
n_core, in,
1.0f, bl->w_core, in,
x, 1,
0.0f, core_dots, 1);
/* per-row 后处理: core_gain * K + bias */
#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];
}
/* BINARY + PRUNE 路径仍用原循环 (位运算, sgemm 不适合) */
#pragma omp parallel for schedule(static)
for (int j = 0; j < out; j++) {
uint8_t m = bl->logic_mask[j];
if (m == 0) continue; /* CORE 已用 sgemm 算完 */
if (m == 1) { /* BINARY */
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 { /* PRUNE */
y[j] = 0.0f;
}
}
return;
}
#endif
/* 无 OpenBLAS 或 n_core=0: 走原 OpenMP + 8 倍展开路径 */
#pragma omp parallel for schedule(static)
for (int j = 0; j < out; j++) {
switch (bl->logic_mask[j]) {
case 0: { /* CORE: float matmul * core_gain * K
* Whitebox circuit trace: CORE was 22x weaker than BINARY.
* BINARY: sign(w)*alpha*K — sign() amplifies every weight to ±1.
* CORE: w*K — raw float weights (~0.02), no amplification.
*
* Fix: core_gain = 1/alpha normalizes CORE's effective weight
* magnitude to ~1 (like BINARY's sign). Capped at 5 to prevent
* explosion when alpha is tiny. This makes CORE's signal O(1)
* like BINARY, so the circuit can actually use CORE's precision.
*
* alpha[j] = mean(|w_float[j]|), recalculated in bin_layer_repack.
* For CORE neurons, alpha is set in init then recalculated in repack. */
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; /* moderate cap: 5x boost */
y[j] = s * core_gain * K + bl->bias[j];
break;
}
case 1: { /* BINARY: sign(w) * alpha * K + bias (ternary if zbits set) */
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; /* zero mask: 1=skip */
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) {
/* Ternary: zeroed positions contribute 0.
* contribution = sign * (1 - zbit) * x.
* (1 - zbit) ∈ {0,1} acts as an enable mask. */
if (zbyte == 0) {
/* No zeros in this byte — full 8x dot product. */
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 {
/* Mixed: check each bit. zbyte bit set = skip. */
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: /* PRUNE: zero */
y[j] = 0.0f;
break;
}
}
return;
}
/* Standard BWN path (no 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();
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;
/* Process 8 bytes (8×8=64 bits) per word, 8 floats at a time */
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) {
/* 8x unrolled dot product — auto-vectorizes to SIMD */
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 {
/* Tail: handle remaining elements (< 8) */
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];
}
}
/* BNN fast path: XNOR + popcount, binarizes BOTH x and W.
* ~64x faster than BWN. With K-norm input scaling (XNOR-Net, Rastegari 2016),
* the output magnitude is restored: y = (2*pc-in) * alpha * K + bias, where
* K = mean(|x|). Without K, outputs have wrong magnitude → garbled generation. */
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;
/* Compute input scale K = mean(|x|) — restores magnitude lost by sign(x).
* O(in) cost is negligible vs the O(in*out) XNOR+popcount matmul. */
float abs_sum = 0.0f;
for (int i = 0; i < in; i++) abs_sum += fabsf(x[i]);
float K = abs_sum / in;
/* Binarize input */
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;
}
/* XNOR + popcount per output, scaled by alpha * K */
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];
}
}
/* Legacy bin_forward_float: BWN without K-norm. Kept for callers that
* explicitly don't want input magnitude scaling. */
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;
}
}
/* ========================================================================
* Binary Backward: popcount for grad_x, popcount for alpha update
* ======================================================================== */
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;
/* Logic-guided: if logic_mask exists, dispatch per-output.
* Without this, the non-logic path uses mean_alpha = sum(alpha)/out,
* but PRUNE (alpha=0) and CORE (alpha=0) dilute mean_alpha → wrong
* grad_x → NaN divergence. This was the root cause of ai_4116f587's
* step-100 NaN in --logic + --real-attention testing.
* CORE: grad_x += grad_y * w_core (proper float gradient)
* BINARY: grad_x += grad_y * sign(wbits) * alpha (original logic)
* PRUNE: skip (zero gradient, output is zeroed in forward) */
if (bl->logic_mask) {
for (int i = 0; i < in; i++) grad_x[i] = 0.0f;
/* K-norm for BINARY/CORE outputs (must match bin_forward's logic_mask path,
* otherwise CORE/BINARY grad_x is off by a factor of K = mean(|x|)). */
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) {
/* CORE: float gradient through w_core * core_gain * K */
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) {
/* BINARY: gradient through sign(wbits) * alpha */
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;
}
/* PRUNE (case 2): no gradient, skip entirely */
}
return;
}
/* Non-logic path: original bin_backward */
/* Part 1: grad_x via XNOR+popcount using transposed weights */
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;
}
/* Part 2: alpha + bias update via popcount (reuse x_bits)
*
* [FIX 严重4] alpha update direction: was '+=', now '-'= to match bias.
* Old code: bl->alpha[j] += lr * grad_alpha * gy / in; (WRONG: ascends loss)
* New code: bl->alpha[j] -= lr * grad_alpha * gy; (correct: descends loss)
* Also dropped spurious '/in' that shrank alpha's effective LR by in_dim.
* [FIX 严重6] Removed alpha clamp to [0.001, 1.0] — it prevented alpha from
* converging to its natural magnitude and forced a fake floor. */
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; /* FIXED: direction + no /in */
/* Removed: if (bl->alpha[j] < 0.001f) bl->alpha[j] = 0.001f;
* if (bl->alpha[j] > 1.0f) bl->alpha[j] = 1.0f; */
if (bl->alpha[j] < 0.0f) bl->alpha[j] = 0.0f; /* only non-negativity */
bl->bias[j] -= lr * gy;
}
}
/* STE (Straight-Through Estimator) backward pass.
*
* Key difference from bin_backward: this updates w_float (the full-precision
* weights) using the gradient, treating sign() as identity. After the update,
* wbits is re-packed from sign(w_float) via bin_layer_repack().
*
* This allows the binary weights to actually change during training, which
* is impossible with bin_backward (it only updates alpha and bias).
*
* STE gradient: d(loss)/d(w_float) = d(loss)/d(sign(w)) * d(sign(w))/d(w)
* = grad_y * x * 1 (STE: sign'(w) = 1)
* So: w_float[i,j] -= lr * grad_y[j] * x[i]
*
* Memory note: w_float is [in_dim, out_dim] row-major, same as input W.
* This adds ~in*out*4 bytes per layer (e.g. 768*2304*4 = 7MB for c_attn).
* Total for 12 layers × 4 matrices ≈ 339 MB extra during training.
* For inference, w_float can be freed (set to NULL after training). */
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;
/* Part 1: grad_x computation.
* In STE mode, skip wbits_T repack entirely — compute grad_x directly
* from w_float using float arithmetic. This avoids the strided wbits_T
* repack (50% of repack cost) at the expense of float mul-adds.
*
* grad_x[i] = sum_j grad_y[j] * sign(w_float[j*in+i]) * alpha[j]
*
* w_float is [out, in], so w_float[j*in+i] has i contiguous per j.
* But we need i fixed, j varying — that's strided. So we compute
* per-i by accumulating over j. With [out,in] layout, w_float[j*in+i]
* for fixed i has stride=in. This is still strided but avoids repack.
*
* Alternative: compute grad_x = sign(w_float)^T @ (grad_y * alpha)
* which is a matrix-vector product. We can do it per-output j and
* accumulate into grad_x (since w_float[j*in+i] is contiguous in i). */
if (bl->w_float) {
/* Zero grad_x first */
for (int i = 0; i < in; i++) grad_x[i] = 0.0f;
/* For each output j: grad_x += grad_y[j] * alpha[j] * sign(w_float[j*in+i])
* w_float[j*in + 0..in-1] is contiguous → SIMD-friendly! */
/* Ternary QAT grad_x: scale = gy * alpha * K, weight = ternary_w(wf). */
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;
}
/* bin_backward_ste grad_x — 串行 (OpenMP reduction 在某些 BinLayer 配置下产生 NaN) */
for (int j = 0; j < out; j++) {
float gy = grad_y[j];
if (fabsf(gy) < 1e-8f) continue;
/* Skip PRUNE in grad_x: PRUNE outputs 0 in forward, so it must
* NOT contribute gradient to the input. Without this skip,
* sign(w_float) of dead neurons leaks gradient upstream,
* causing PRUNE activations to grow instead of staying silent. */
if (bl->logic_mask && bl->logic_mask[j] == 2) continue;
const float *wf = &bl->w_float[j * in]; /* contiguous [in] */
if (g_use_ternary && bl->zbits) {
/* Ternary STE: grad_x += gy * alpha * K * t(wf[i]),
* t = zero-masked sign. Use the SAME packed wbits/zbits as
* forward so the straight-through gradient differentiates
* exactly the weights used in the forward pass (bit-packed
* XNOR path, no per-element float sign compare). */
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) {
/* Pure float: grad_x uses w_float directly (not sign) */
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 {
/* BWN: grad_x uses sign(w_float) * alpha */
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 {
/* No w_float — use popcount on existing wbits_T (original path) */
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;
}
}
/* Part 2: Gradient accumulation or STE update.
* When g_accumulate_gradients is set (batch training), we accumulate
* grad_y[j]*x[i] into grad_accum and grad_y[j] into bias_grad_accum
* instead of updating weights. model_batch_apply() later applies the
* accumulated (averaged) gradient with Adam. */
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) {
/* Ternary STE: dL/dw_float ~ gy * x * t(wf). The zero-mask makes
* zeroed weights receive no gradient (they're already at ternary 0).
* Use the packed wbits (same as forward) for the sign. */
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; /* Don't update weights yet — wait for model_batch_apply */
}
/* Part 2: STE update — w_float[j*in + i] -= lr * grad_y[j] * x[i]
* w_float is [out, in] (transposed), so w_float[j*in + i] is CONTIGUOUS in i!
* This means the inner loop over i is a contiguous SAXPY: w_float[j] -= scale * x
* The compiler auto-vectorizes this to SIMD FMA (8 floats per iteration).
*
* When g_use_adam is set, we use the Adam update instead of plain SGD:
* m[i] = b1*m[i] + (1-b1)*g (1st moment)
* v[i] = b2*v[i] + (1-b2)*g*g (2nd moment)
* w -= lr * m_hat / (sqrt(v_hat) + eps)
* where m_hat, v_hat are bias-corrected with g_opt_step+1. Adam dramatically
* stabilizes STE on bit-space: per-param adaptive lr counters the
* extreme gradient variance that causes SGD to mode-collapse to "servers". */
if (bl->w_float) {
int t = g_opt_step + 1; /* Adam timestep (1-indexed) */
float bc1 = 1.0f - powf(g_adam_beta1, (float)t); /* bias correction 1 */
float bc2 = 1.0f - powf(g_adam_beta2, (float)t); /* bias correction 2 */
for (int j = 0; j < out; j++) {
float gy = grad_y[j];
if (fabsf(gy) < 1e-8f) continue;
/* Skip PRUNE rows — they're zeroed and contribute nothing. */
if (bl->logic_mask && bl->logic_mask[j] == 2) continue;
float *wf = &bl->w_float[j * in]; /* contiguous [in] */
if (g_use_adam && bl->m_adam) {
float *m = &bl->m_adam[j * in];
float *v = &bl->v_adam[j * in];
/* Adam: per-param adaptive update. 8x unrolled for SIMD. */
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 {
/* SGD: wf[i] -= lr * gy * x[i] (original path). */
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];
}
/* Update bias (SGD always — bias is a scalar, Adam benefit marginal). */
bl->bias[j] -= lr * gy;
}
/* Weight clipping: bound w_float to [-W_CLIP, W_CLIP].
* Standard technique for BWN training (cf. XNOR-Net, Real-to-Binary
* Networks). Prevents w_float from drifting to extreme values where
* Adam's adaptive update becomes numerically unstable and sign(w)
* starts flipping chaotically. The bound [-1, 1] is natural because
* alpha = mean(|w|) is in [0, 1] for typical GPT-2 weight rows.
* Without this, STE+Adam diverges to NaN around step 400-850. */
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;
}
/* Re-pack binary weights from updated (and clipped) w_float */
bin_layer_repack(bl);
} else {
/* Pure float: clip to larger bound to prevent gradient explosion */
#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;
}
}
/* Ternary: refresh zbits (zero mask) from updated |w_float| vs Δ.
* This is the TWN STE dynamic — zeroed weights that received enough
* gradient to cross Δ "wake up" (become ±1), and active weights whose
* |w| dropped below Δ get zeroed. alpha is also recomputed over the
* new active set. Skipped automatically if zbits is NULL (BWN mode). */
if (bl->zbits) bin_layer_repack_ternary(bl);
} else {
/* No w_float — fall back to alpha-only update */
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; /* FIXED: direction + no /in */
if (bl->alpha[j] < 0.0f) bl->alpha[j] = 0.0f;
bl->bias[j] -= lr * gy;
}
}
}
/* ========================================================================
* Standard Neural Network Operations
* ======================================================================== */
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) {
/* Sampled-softmax cross-entropy gradient.
*
* Loss: L = -log( exp(tl) / (exp(tl) + sum_k exp(neg_k)) )
* = -log( prob ), prob = exp(tl-mx) / (exp(tl-mx) + sum_k exp(neg_k-mx))
*
* Gradient w.r.t. hidden[i] (treating sampled negatives as constants —
* standard sampled-softmax approximation, drops the second-order term
* sum_k prob_k * wte[k, i]):
*
* dL/d(hidden[i]) = dL/d(tl) * d(tl)/d(hidden[i])
* = -(1 - prob) * wte[target, i]
*
* ---------------------------------------------------------------------
* BUGFIX (gibberish-output root cause):
*
* The previous implementation returned
*
* grad_hidden[i] = +(1 - prob) * wte[target, i] * 0.001f
*
* which had THREE bugs that together made binary training diverge into
* mode-collapse / gibberish:
*
* (1) WRONG SIGN. Returned +grad instead of -grad. Combined with the
* optimizer's `w -= lr * grad`, this flipped descent into ascent
* on -log(p_target): the model was trained to *lower* p_target,
* i.e. to actively avoid predicting the correct token. After a few
* hundred steps the logits collapse and generation produces
* constant-token gibberish.
*
* (2) grad_scale = 0.001f shrank the learning signal by 1000x. Even
* after fixing the sign, with lr=0.05 the effective step on
* `hidden` was 5e-5 — far too small to escape random init in any
* reasonable number of steps. Removed.
*
* (3) `se += 1.0f` per sampled negative (instead of the true
* exp(neg_k - mx)) inflated the denominator systematically,
* forcing prob -> 0 and (1-prob) -> 1, which (combined with the
* wrong sign) made every step push hidden AWAY from wte[target]
* at maximum magnitude. Now we use the actual exp(neg_k - mx).
* ---------------------------------------------------------------------
*/
float tl = 0;
for (int i = 0; i < n_embd; i++) tl += hidden[i] * wte[target * n_embd + i];
/* Sample negatives and remember their logits so we can build the
* correct softmax denominator. Cap at 256 to keep the stack buffer
* bounded (n_samples=100 in practice). */
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);
/* prob = P(target) under the sampled softmax. +1e-7f guards against
* logf(0) in the caller (cross_entropy_sampled) and div-by-zero here. */
float prob = expf(tl - mx) / (se + 1e-7f);
/* Correct gradient of L = -log(prob) w.r.t. hidden[i].
* Optimizer does `w -= lr * grad`, so a NEGATIVE grad here means
* hidden moves TOWARD wte[target], which INCREASES prob and
* DECREASES loss — i.e. true gradient descent. */
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;
}
}
/* BUG #48 FIX: Normalize residual stream ||x|| to target_norm.
* The residual stream accumulates: x = wte + sum(attn_residual + mlp_residual).
* With residual_scale=1.0 and 8+ layers, ||x|| grows from ~1 (L0) to ~210 (L7).
* This causes logits = dot(final_ln, wte) to explode, making sampling degenerate.
*
* LayerNorm normalizes the INPUT to each sublayer, but NOT the residual x itself.
* So x grows unboundedly between layers (or collapses to near-zero).
*
* Fix: after each residual addition, scale x so ||x|| ≈ target_norm.
* This is similar to "RMSNorm on residual stream" used in some architectures.
*
* v13m: CRITICAL FIX — previously only capped at target_norm (6.0), never
* boosted. Since ||x|| starts at ~3.5 (below 6.0) and shrinks ~25% per
* layer due to binary projections, normalize_residual NEVER fired, creating
* a death spiral: ||x|| 3.54 → 2.64 → 1.95 → 1.46 → ... → 0.76.
* The proportional scaling (0.15*||x||) amplified this: as ||x|| shrank,
* attention/MLP contributions shrank too, unable to maintain signal.
*
* Fix: enforce BOTH minimum and maximum. If ||x|| < target_min, scale UP
* to target_min. If ||x|| > target_max, scale DOWN to target_max.
* This keeps the residual in a healthy range [3.0, 6.0] across all layers. */
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;
/* v16: 移除放大分支 — 白盒: 放大操作把共模方向等比放大, 逐层推高相似度. 只保留向下封顶 */
if (norm > target_norm) {
float scale = target_norm / norm;
for (int i = 0; i < n; i++) x[i] *= scale;
}
}
/* v13b: Scale a sublayer output to a small target norm before adding to
* the residual stream. This makes attn/mlp outputs small perturbations
* rather than dominant signals, preventing representation collapse.
*
* Without this, ||proj_out|| ~ 27 (since norm1_out has ||.||~24 from
* LayerNorm over 512 dims), while ||x|| ~ 1.8 (embedding). The sublayer
* output completely overwrites the residual direction, causing all inputs
* to converge to the same representation after 1-2 layers.
*
* With target_norm=0.3, the sublayer contributes a 0.3-magnitude
* perturbation on top of the ~1.0-norm residual, preserving input
* diversity while still allowing the model to transform representations. */
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;
}
/* ========================================================================
* Full-vocab softmax cross-entropy (replaces sampled softmax for training)
*
* The sampled-softmax path (cross_entropy_sampled / cross_entropy_grad)
* uses 100 random negatives per step. When the training data is heavily
* skewed (91% of sentences end in token 764='.'), the model can trivially
* win against 100 random negatives by always outputting 764 — collapsing
* to a single-token predictor. The full-softmax path computes the true
* gradient over all 50257 vocab tokens, so the model is forced to actually
* learn the distribution (token 764 gets probability mass only when the
* context genuinely predicts it).
*
* Cost: 50257 * 768 = ~38M FMA per forward, ~76M per backward. Negligible
* vs the per-layer binary matmul (12 layers * ~3M FMA = 36M).
* ======================================================================== */
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; /* v16 */
}
}
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);
/* numerically stable softmax + cross-entropy */
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];
}
/* logits_scratch now holds softmax probabilities; loss = -log(p_target) */
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) {
/* grad_hidden[i] = (softmax(logits)[target_or_not] - one_hot[target]) * wte[i]
* = (p[j] - 1{j==target}) * wte[j, i] summed over j.
*
* Equivalent to: grad_hidden = wte[target] - sum_j p[j] * wte[j]
* But computing it as wte[target] - sum_j p[j]*wte[j] is O(vocab*n_embd)
* and avoids materializing a per-(j,i) gradient.
*
* logits_scratch must already hold the softmax probabilities from
* cross_entropy_full (caller reuses it to avoid recomputing logits). */
/* Start with wte[target] (the +1 in dL/d_logit = p - one_hot, multiplied
* by -1 because we want dL/d_hidden, and the chain rule gives a negative
* sign through the loss). Actually:
* L = -log(p_target), p = softmax(logits), logits[j] = hidden . wte[j]
* dL/d_logits[j] = p[j] - 1{j==target}
* dL/d_hidden[i] = sum_j (p[j] - 1{j==target}) * wte[j, i]
* = sum_j p[j]*wte[j,i] - wte[target, i]
* The optimizer does w -= lr * grad, so we return dL/d_hidden directly.
* (Previously the sign bug was here; now correct.) */
const float *wt_target = &wte[(size_t)target * n_embd];
for (int i = 0; i < n_embd; i++)
grad_hidden[i] = -wt_target[i];
/* Add sum_j p[j] * wte[j, i]. p[j] is in logits_scratch (already
* normalized to sum=1 by cross_entropy_full, but we re-normalize
* defensively in case the caller passed un-normalized logits). */
float psum = 0;
for (int j = 0; j < vocab_size; j++) psum += logits_scratch[j];
float inv_psum = 1.0f / (psum + 1e-12f);
/* CE backward: 串行 (每步调用 12 次, per-thread partials 的 memset+combine
* 开销 > 并行收益. 串行更稳定, 无 NaN 风险.) */
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];
}
/* v16: logits 缩放了 g_logit_scale, 梯度按链式法则同乘 */
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 File Loading (GPW2 format)
* ======================================================================== */
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);
}
/* Free a single tensor's data by key (sets data to NULL so tensor_free_all
* won't double-free). Used to reclaim memory from large weight matrices
* after they've been binarized into BinLayer. */
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;
}
}
}
/* mmap-based tensor loader: maps the GPW2 file into memory and points
* each tensor->data at the corresponding offset. The OS pages in data
* on demand, so startup is ~10x faster on cold cache and peak RSS is
* lower (only touched pages count).
*
* Trade-off: cannot free individual tensors (they live in the mmap region),
* so the free-float-weights optimization is disabled in mmap mode. Use
* this when startup time matters more than steady-state RSS.
*
* The returned Tensor array must be freed with tensor_free_all_mmap(). */
/* sys/mman.h and sys/stat.h are included at the top (with Windows shims) */
#ifndef _WIN32
#include <sys/mman.h>
#include <sys/stat.h>
#endif
#include <fcntl.h>
#ifndef _WIN32
#include <unistd.h>
#else
/* Windows: open/close/read lseek shims via _io.h */
#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; /* the mmap'd region, for munmap later */
size_t mmap_size;
int fd;
} MmapedTensors;
static MmapedTensors g_mmap_state = {NULL, 0, NULL, 0, -1};
/* ========================================================================
* Random-weight GPW2 generator — train an arbitrary-size model from scratch
* (no pretrained checkpoint needed). Writes Gaussian-init weights in the same
* "GPW2" layout that tensor_load_all / model_load expect, for any ModelConfig.
* Keys follow the GPT-2 (qkv_merged) or LLaMA (separate Q/K/V, SwiGLU) layout
* selected by cfg.qkv_merged / cfg.act_type.
* ======================================================================== */
#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); }
}
/* Write a tensor with custom initialization to GPW2 file */
static void bin_gpw2_put_init(FILE *f, const char *key, int ndim, const int *shape, float scale, int init_mode) {
/* init_mode: 0 = N(0, scale), 1 = constant scale, 2 = zeros, 3 = Xavier(sqrt(2/fan_in)) */
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) {
/* constant value (for LayerNorm weight = 1.0) */
for (int i = 0; i < n; i++) fwrite(&scale, 4, 1, f);
} else if (init_mode == 2) {
/* zeros (for biases) */
float z = 0.0f;
for (int i = 0; i < n; i++) fwrite(&z, 4, 1, f);
} else if (init_mode == 3) {
/* Xavier/He: std = sqrt(2.0 / fan_in) for ReLU/GELU, fan_in = shape[ndim-1] */
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 {
/* Normal(0, scale) */
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); }
/* Count tensors: base(4) + per_layer depends on config */
int per_layer;
if (cfg.qkv_merged) {
per_layer = (cfg.act_type == ACT_SWIGLU) ? 11 : 12;
} else {
per_layer = 9; /* q/k/v/o + gate/up/down + 2 layernorms */
}
int n_tensors = 4 + cfg.n_layer * per_layer;
fwrite("GPW2", 1, 4, f);
fwrite(&n_tensors, 4, 1, f);
/* Embeddings: N(0, 1/sqrt(n_embd)) for proper scale */
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);
/* Final LayerNorm: weight=1.0, bias=0.0 (CRITICAL for convergence) */
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) {
/* Weight matrices: Xavier init (large enough for meaningful alpha after binarization) */
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);
}
/* LayerNorm: weight=1.0, bias=0.0 */
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];
}
/* Point data at the mmap'd region (no copy) */
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) {
/* Don't free individual data pointers — they live in the mmap region */
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;
}
}
/* ========================================================================
* Sparse Sliding Window Attention + Stateful Continuous Inference
* ========================================================================
* Implements:
* 1. attention_forward_sliding() — sparse attention with configurable window
* 2. attention_backward_sliding() — gradient computation for sparse attention
* 3. Circular buffer KV cache management (no memcpy shifting)
* 4. Attention sinks (StreamingLLM-style: keep first N tokens stable)
* 5. trans_layer_forward_sliding() — layer forward with sparse attention
* 6. Stateful inference context (g_sctx) for token-by-token generation
*
* Design:
* - Sliding window: each token attends to last W tokens + first S sink tokens
* - Circular buffer: KV cache uses ring buffer, write pointer wraps around
* - Attention sinks: first S positions are always in the attention window
* - Configurable via ModelConfig.sliding_window and ModelConfig.n_sinks
* ======================================================================== */
/* Global stateful inference context */
StatefulContext g_sctx = {0};
/* ─── Sliding Window Attention Forward ──────────────────────────── */
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;
/* Circular buffer: store at seq_pos % n_ctx */
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));
/* Build attended position list: sinks + sliding window */
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;
/* [加速] thread-local 预分配缓冲区, 避免每调用 malloc/free (参考 llama.cpp).
* 之前每 token × 每层都 malloc 3 个数组, 8192 token × 10 层 = 81920 次 malloc/step. */
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;
/* [加速] head 循环并行 — 但只在 n_attend > 256 时开启 (大上下文才值得 fork/join).
* 小 n_attend 时 OpenMP fork/join 开销 > 计算收益 (8192 token × 10 层 = 81920 次调用).
* SIMD 8 倍展开点积始终启用. */
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;
/* Compute attention scores + max (SIMD 8 倍展开点积) */
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;
}
/* Softmax */
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;
/* Weighted sum of V (SIMD 8 倍展开) */
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 {
/* 小 n_attend: 串行 (避免 fork/join 开销) */
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];
}
}
}
}
/* ─── C3 概念图驱动长上下文记忆注意力 (推理端, 与 --concept-graph 一体) ─── */
/* 概念注意力探针统计结构 (定义在此处, 因为 attention_forward_concept_ctx 在下方使用).
* 修复 (2026-08-17): 旧版 ctx 函数没统计, 导致 [CATTN] fwd=0 假警报, 团队持续
* 误以为概念注意力没参与前向. 现在两个版本都统计. */
typedef struct ConceptAttnStats {
long forwards; /* 前向调用次数 (attention_forward_concept 简单版) */
long forwards_ctx; /* 前向调用次数 (attention_forward_concept_ctx 长上下文记忆版) */
long candidates; /* 累计候选对数 (实际计算量) */
long full_equiv; /* 累计等效全注意力对数 (seq_pos+1) */
long gate_pairs; /* 参与门控判断的对数 */
long gate_blocked; /* 被门控屏蔽的对数 */
long msg_candidates; /* 信使候选数 */
double msg_mass; /* 信使获得的注意力质量累计 */
int last_n_filled; /* 最近一次前向的已填充片段数 */
long ctx_memory_slots_used; /* ctx 版本: 实际命中的概念槽总数(累加) */
long ctx_total_attend; /* ctx 版本: 实际 attend 总候选数(累加) */
/* 审查建议的核心验证项: 信使是否携带"差异"而非"共识均值" */
double msg_inter_cos; /* 信使间平均余弦 (越低越好, 目标 < 0.2 说明去同质化生效) */
double msg_norm; /* 信使平均范数 (验证范数钳制 MSG_NORM_CAP=4.0 是否生效) */
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;
}
/* ─── C3 概念图驱动长上下文记忆注意力 (定义) ─── */
/* 在 sinks+window 之外, 额外 attend 一组"概念状态槽". 槽由被窗口挤出的中间段 token
* 按其在概念图里的概念归属(neighbor[i*K+0])聚合而成 — 即同一份概念图既引导生成,
* 又驱动长上下文记忆, 远端信息以概念压缩态回流. 零额外训练. */
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, /* [vocab*n_embd] 概念锚点 */
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);
/* 探针统计: ctx 版本前向计数 + 等效全注意力对数 */
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));
/* BUG FIX (2026-08-17): is_mem was 256 bytes but the check used
* `n_total <= 8192` (matching pos_idx size), so if n_total > 256 the
* code wrote past is_mem and triggered "stack smashing detected" on
* glibc 2.35 (zzai). The crash happened after step 50 when
* inference_trace_compact fired model_stateful_begin, which allocated
* g_sctx.cctx_k and caused subsequent training forward passes to route
* through this function with n_total > 256 (since n_attend = n_sink + n_win
* = 64 + min(seq_pos+1, 1024) can reach 1088, well past 256). Fix: size
* is_mem to match pos_idx (8192) so the (n_total <= 8192) check holds. */
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]) {
/* 概念状态槽: K 用聚合 cctx_k; query 与槽锚点(概念图里的概念 token)算分 */
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; /* 锚点已在聚合时决定槽归属, 此处用聚合K即可 */
} 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); }
}
/* ─── Sliding Window Attention Backward ─────────────────────────── */
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);
/* Determine attended positions (same as forward) */
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));
/* [加速] thread-local 预分配缓冲区 */
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; }
}
/* [加速] head 循环并行 — 只在 n_attend > 256 时开启 (避免小 n_attend fork/join 开销) */
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;
/* Recompute scores + softmax (SIMD 展开) */
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;
/* g_w[i] = <g_out, V_{pos_list[i]}> (SIMD 展开) */
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);
/* g_Q[d] += sum_i g_scores[i] * K_{pos_list[i]}[d] * scale */
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 {
/* 小 n_attend: 串行 */
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);
/* [加速] KV-only 快速 prefill: 只算 norm1 + Q/K/V + 存 cache, 跳过 attention/attn_o/MLP.
* 用于训练 forward 的中间 token (p < t), backward 只对最后一个 token 做,
* 所以中间 token 不需要 attn_out/proj_out/mlp_hidden 等 act, 只需 K/V 进 cache.
* 节省 ~60% 计算量 (attn_o + MLP 占层 forward 的大头). */
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;
/* Norm1 */
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;
}
/* kv_only: 用 bin_fwd 算 Q/K/V (纯 float 投影在 Windows 上产生 NaN, 回退).
* K/V 在 backward 是常量, 但 bin_fwd 保证和完整 forward 一致. */
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;
}
/* ─── Transformer Layer Forward with Sliding Window ─────────────── */
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; /* REAL sequence length for concept-attn closing */
/* Norm1 + QKV projection */
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]);
/* v13l: Skip W_v projection — use norm1_out directly as attention output */
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);
}
/* Apply RoPE if configured */
if (cfg->attn_type == ATTN_ROPE)
apply_rope(act->q, act->k, abs_pos, cfg->n_head, n / cfg->n_head, n);
/* Sliding window attention */
if (tl->_kv_k && tl->_kv_v) {
/* C3 概念图驱动长上下文记忆: 与 --concept-graph 一体. 仅当概念图已加载时启用,
* 把被窗口挤出的中间段按概念聚合进概念状态槽, 额外 attend. 概念图主线不旁落. */
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) {
/* v16: 概念感知注意力推理接入 — 与训练同通路, 支持同权重 A/B 对比 */
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 {
/* Fallback: V-copy (legacy) */
memcpy(act->attn_out, act->v, n * sizeof(float));
}
} /* end !g_skip_wv */
/* Output projection */
bin_fwd(act->proj_out, act->attn_out, &tl->attn_o);
/* v13j: Proportional attention scaling (same as standard forward) */
{
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];
/* Norm2 + MLP */
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) {
/* BUG #44 FIX: static buffer instead of malloc/free per call */
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);
/* v13j: Proportional MLP scaling + normalize_residual(6.0) (same as standard) */
{
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);
}
/* ─── Transformer Layer Backward with Sliding Window ────────────────
* 与 trans_layer_backward 完全对称, 唯一区别: attention 反向用
* attention_backward_sliding (与推理端 attention_forward_sliding 配对).
* 训练端用这个, 训练/推理 attention 窗口完全一致 (sinks + window 两段式).
* act->seq_pos = abs_pos (推理端存的), act->n_ctx = n_ctx (传入的物理 cache 大小). */
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))
/* MLP backward (与 trans_layer_backward 完全一致) */
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;
/* Attention backward — 用 sliding 版本, 与推理 forward 配对 */
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, 与推理 attention_forward_sliding 完全配对.
* act->seq_pos = abs_pos, act->n_ctx = 物理 cache 大小 (传给 n_ctx 参数). */
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;
}
/* ─── Stateful Inference: Begin New Session ─────────────────────── */
/* ========================================================================
* Batch Training Implementation
* ======================================================================== */
void model_batch_alloc(Model *m) {
/* Buffers are already allocated in bin_layer_init, but this ensures
* they exist for models loaded without g_use_adam. */
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) {
/* Separate Q/K/V — attn_k and attn_v also need buffers */
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));
}
}
}
/* Allocate gradient accumulation + Adam state for wte, wpe and ln_f */
if (!m->grad_wte_accum) {
size_t wte_size = (size_t)m->cfg.vocab_size * m->cfg.n_embd;
/* v2: 64 字节对齐分配 — SIMD AVX-512 需要 64 字节对齐 */
#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));
}
/* Allocate norm weight gradients for each layer */
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));
}
/* BUG #50 FIX: Allocate Adam state for LayerNorm weights */
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); /* ensure per-thread buffers exist for parallel batch */
}
/* ---- per-thread resource management ---- */
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; /* attn_q,o,k,v,mlp_gate,down,up */
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++; }
/* 分配 n_bl_max (上限7) 个指针槽, 与 thr_res_free 的遍历上限一致,
* 避免 qkv_merged 或 act_type 导致 n_bl<7 时越界 free 野指针。 */
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));
}
/* === Ponder 循环思考缓冲 === */
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); /* 分配 n_layer 槽, 用前 rec_iters 个 */
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);
/* Ponder 缓冲释放 */
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;
}
/* Sum all per-thread gradient pools into the real grad_accum (which
* model_batch_begin already zeroed). Called once after the parallel region,
* on the master thread (outside the parallel region). */
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]) {
/* v2: OpenMP 并行梯度合并 — 大矩阵 (sz > 4096) 才开 */
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];
}
}
/* v2: wte/wpe 梯度合并并行 — 32768×512 = 1670万, 最大瓶颈 */
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];
/* Zero norm weight gradients */
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));
}
}
/* Zero embedding and norm gradients */
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));
}
/* Zero per-thread gradient pools so they start fresh each step.
* (Real grad_accum above is what the optimizer consumes; these pools are
* accumulated into it by thr_grad_reduce and must not carry across steps.) */
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));
}
}
}
}
/* ========================================================================
* Unified Sliding-Window Training (端到端: 训练 = 推理路径)
* ========================================================================
* 目的: 让训练和推理用完全相同的前向路径, 消除 train/infer 不一致.
*
* 实现策略:
* - forward: 用 model_stateful_begin + 逐 token forward (与推理完全一致)
* 但 act 用 per-thread 的 g_thr[tid].acts[l] (支持 batch 并行 + backward)
* - backward: 用 trans_layer_backward_sliding (attention 部分用 sliding 版本)
*
* 与原 model_forward/model_batch_backward 的区别:
* 1. attention 窗口: {0..sink} ∪ {pos-w+1..pos} 两段式 (训练原版是单段跳过 sinks)
* 2. KV cache 索引: pos % n_ctx 环形 (训练原版是 pos 线性)
* 3. 位置编码: wpe[pos % n_ctx] (训练原版是 wpe[pos])
* 4. 逐 token forward, 不做 prefill 加速 (与推理逐 token 一致)
*
* 性能影响: 比 model_forward 慢 (不能 prefill 复用), 但保证 train=infer.
* 单步预计: ~30-40s/step (2 核, vs 原 model_forward 13-17s/step)
* ===================================================================== */
/* ========================================================================
* PonderNet 循环思考 — 训练/推理前向反向 (见 lal_ponder.h 头部设计说明)
* ======================================================================== */
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;
/* A. 逐层停机: 层 0..L-2, 每层一个停机概率, 状态入缓存 */
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);
}
/* B. 末块: R≥2 块内循环 (权重共享, 迭代级 PonderNet), R=1 单遍 */
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) {
/* 保存迭代0 K/V — 该 token 对外暴露的 K/V (与 context prefill 单遍语义一致) */
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 {
/* 迭代结束恢复迭代0 K/V: 后续 token 只应 attend 首次通过的 K/V */
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; /* 末步是 remainder (强制停机) */
ponder_dist_fill(pb, halts, n_param);
/* 混合读出: out = Σ c_s·s_s (Σc = 1), 末步 c = 剩余质量 */
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));
/* 训练日志指标 (ste_train.c 读取) */
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; /* 读出梯度 (final norm 反向之后) */
float *V = r->g_pre4; /* 滚动梯度缓冲 (norm 反向已用完, 空闲) */
/* g_l = <G, s_l>: 任务 loss 对停机权重的显式梯度来源 (状态 detach) */
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;
}
/* 停机单元参数梯度 (dpre 已链到 pre-activation) */
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];
}
/* 主干链式反传 (c 视作常数, 与 detach 一致):
* V = c_s·G + 上方链式梯度 → 层/迭代反向就地更新 V */
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;
/* 恢复该迭代自身 K/V 到 cache 当前位 (attention backward 需重算分数) */
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);
}
}
/* V 现在是 x_0 (embedding 输出) 的梯度 → 交给既有 wte/wpe 尾部逻辑 */
memcpy(r->gh, V, n * sizeof(float));
}
/* 推理侧 ponder 前向 (stateful, 单线程): 混合读出 + 早退 + 思考深度统计 */
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; /* 块内早退: cache 已含迭代0 K/V, 无物理层被跳过 */
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) {
/* 层间早退: 被跳过的物理层用当前状态填 K/V (已收敛近似), 保证后续 token attention 完整 */
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; /* 预测 tokens[t+1] from context tokens[0..t] */
/* 训练时也用 stateful 机制: 开始前 reset KV cache (与推理 model_stateful_begin 一致).
* 但不用 g_sctx 的 act (那是推理用的全局 act), 而用 per-thread g_thr[tid].acts. */
if (!m->k_cache) model_kv_cache_alloc(m);
/* [加速] prefill 复用: 同一样本的多个 pred_pos (1, 1+stride, ...) 共享 context.
* 训练循环里 pred_pos 严格递增, 之前 prefill 到 last_prefill_to 的 KV cache
* 仍然有效, 只需 forward last_prefill_to..t 这一段.
* 当样本切换 (tokens 指针不同) 时, 清 cache 重新 prefill 全程.
* 复用率: n_preds=3 时 ~2/3 prefill 计算量节省. */
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) {
/* 样本切换: 清 KV cache */
need_clear = 1;
prefill_from = 0;
last_prefill_to = -1;
} else if (t > last_prefill_to) {
/* 同一样本, pred_pos 递增: 复用 [0, last_prefill_to] 的 KV cache */
prefill_from = last_prefill_to + 1;
if (prefill_from > t) prefill_from = t;
} else {
/* t <= last_prefill_to: 不可能 (pred_pos 递增), 安全处理 */
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();
}
/* 逐 token forward (与推理 model_stateful_forward_sliding 完全一致, 但用 per-thread act).
* [加速] 中间 token (p < t) 用 kv_only 快速路径: 只算 K/V 存 cache, 跳过 attn_o/MLP.
* 最后一个 token (p == t) 走完整 trans_layer_forward_sliding (含 attn_o/MLP).
* backward 只对最后一个 token 做, 所以中间 token 的 act 不需要完整保存.
* 节省 ~60% 计算量 (attn_o + MLP 占层 forward 的大头).
* [加速] prefill_from > 0 时跳过已缓存的 token. */
float *x = g_thr[tid].x;
for (int p = prefill_from; p <= t; p++) {
/* embedding + position (与推理一致: wpe[pos % ctx]) */
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) {
/* 中间 token: kv_only 快速路径 (只存 K/V, 不算 attn_o/MLP) */
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) {
/* 最后一个 token: PonderNet 循环思考前向 (逐层停机 + 块内循环 + 混合读出) */
ponder_train_forward(m, tid, cache_pos, p, window, n_sinks, ctx);
} else {
/* 最后一个 token: 完整 forward (act 供 backward 使用) */
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;
/* 最终 norm + logits (与推理一致) */
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);
/* 同步写一份到 m->final_ln (兼容诊断代码) */
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;
/* CE gradient (与 model_batch_backward 一致) */
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;
/* Gradient clipping */
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; }
/* Backprop through final norm */
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));
/* 关键: 只对最后一个 token (位置 t) 做 backward.
* 与原 model_batch_backward 一致 —— attention_backward_sliding 把 cached K/V
* 当常量, 只算当前 token 的 Q/K/V 梯度. 不需要回传到前面所有 token. */
if (g_ponder_cfg.enable && m->ponder_ready) {
/* PonderNet 循环思考反传: c_s 加权链式 + 停机单元显式梯度 */
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);
}
/* wte 梯度: 只给最后一个 token (与原版一致) */
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];
}
}
/* wpe 梯度: 给所有位置 (与原版一致, 平均) */
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) {
/* Apply accumulated gradients with Adam, averaged by batch_size.
* This does ONE optimizer step for the entire batch. */
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);
/* === LAL-aware Adam: group-wise second moment ===
* Standard Adam normalizes per-parameter: update = m / sqrt(v_per_param)
* This erases CORE/BINARY gradient differences because large CORE
* gradients produce large v, reducing the effective update.
*
* LAL-aware Adam shares v within each group (CORE, BINARY):
* v_core = EMA(mean(||grad_core||^2)) -- shared across ALL CORE params
* v_bin = EMA(mean(||grad_bin||^2)) -- shared across ALL BINARY params
* update = m[i] / sqrt(v_group) -- group-normalized
*
* This preserves relative gradient magnitudes: if CORE has 3x
* larger gradients than BINARY, the update is 3x larger too.
* Combined with g_core_lr_multiplier, CORE truly learns faster.
*
* PRUNE neurons: weight decay toward 0 + freeze if small. */
if (g_use_lal_adam && bl->logic_mask) {
/* Step 1: Compute group-wise gradient energy */
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; /* per-param average within this neuron */
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;
/* Step 2: EMA update of group v (persist across steps) */
/* Store in first CORE/BINARY neuron's v_adam[0] as proxy.
* This is safe because v_adam is per-param and we only
* read v_adam[0] of the first neuron in each group. */
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;
/* Step 3: Update weights using group-wise normalization
* [加速] OpenMP 并行: 每行 j 独立更新, 无数据依赖 */
#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) {
/* PRUNE: weight decay toward zero */
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;
/* BUG #22 FIX: use the right group's sqrt_v for each neuron.
* Previously hardcoded bin_sqrt_v for BOTH groups, which
* artificially amplified CORE updates (CORE has 10x larger
* gradients from alpha=2 vs beta=0.2 + sqrt(807/201) normalization,
* so core_v >> bin_v; using bin_sqrt_v as denominator makes
* CORE effective lr explode on top of g_core_lr_multiplier=3.0).
*
* With this fix, each group is normalized by its own
* second moment — relative gradient magnitudes within a
* group are preserved, and cross-group scaling is left
* to g_core_lr_multiplier alone. */
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;
/* GROUP-WISE v: use sqrt_v, not per-param v */
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;
}
/* Skip standard Adam loop — already done above */
goto layer_done;
}
/* [加速] OpenMP 并行: 标准 Adam 路径也并行 */
#pragma omp parallel for schedule(static)
for (int j = 0; j < out; j++) {
if (bl->logic_mask && bl->logic_mask[j] == 2) continue;
/* CORE neurons get boosted learning rate for faster differentiation */
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) {
/* Average gradient over batch */
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;
/* Adam moment updates */
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;
/* Bias-corrected update */
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 {
/* SGD: w -= lr * avg_grad */
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];
}
/* Update bias */
bl->bias[j] -= lr_j * bl->bias_grad_accum[j] * inv_batch;
}
layer_done:
/* BUG #54 FIX 方案I + v10: 定期检查 W_v effective rank + decay 0.999
*
* 根因: 正反馈循环让 W_v 退化为 rank-1
* v8 step100 rank=300 (好), step200 rank=5 (退化)
*
* 方案I: 每 50 步检查 W_v 的 effective rank,
* 如果 rank 太低 (Frobenius/max_row 比值 < 5), 用 Xavier 重新初始化.
*
* v10: decay 0.99→0.999 (0.999^200=0.819 vs 0.99^200=0.134)
* 数值 rank 从 5→509, 但 S[0] 仍主导 (eff_rank 5-8)
* 下一步需 orthogonal regularization 来 cap S[0]
*
* 近似 rank: ||W||_F / ||W||_max_row
* 满秩时 ≈ sqrt(out), rank-1 时 ≈ 1
*/
if (b == 0 && m->cfg.qkv_merged) {
int n = m->cfg.n_embd;
int in = bl->in_dim;
/* 只检查 W_v 部分 (rows 2*n 到 3*n) */
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);
}
/* v10: W_v weight decay 0.999 + noise
* v13l: Skip when g_skip_wv — W_v not in forward path, no need to regularize */
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; /* v10: gentle decay, 0.999^200=0.819 */
wf[i] += 0.001f * ((float)rand() / RAND_MAX * 2.0f - 1.0f); /* noise */
}
}
} /* end !g_skip_wv */
/* v11+v13l: Orthogonal regularization on W_v
* Loss += lambda * ||W_v^T @ W_v - I||^2_F
* Gradient: dW_v = 4 * lambda * W_v @ (W_v^T @ W_v - I)
*
* v13l enhancements:
* - Increased lambda 0.02→0.05 for stronger rank promotion
* - Added diagonal variance penalty: encourages uniform singular values
* (high effective rank). When all diag(G) entries are equal,
* all singular values are equal → maximum effective rank.
* - Skip when g_skip_wv (W_v not in forward path)
*
* Effect: pulls all singular values toward 1.
* - Caps S[0] (currently 30-72) down toward 1
* - Boosts S[1:] (currently 1-2.5) up toward 1
* - SVD: if W = U S V^T, then W^T W = V S^2 V^T
* Gradient W @ (W^T W - I) = U S V^T V (S^2 - I) V^T = U S (S^2-I) V^T
* So dW_v moves S[i] toward: S[i] - 4*lambda*S[i]*(S[i]^2-1)
* S[0]>1 → decrease, S[i]<1 → increase. Perfect!
*
* Compute: G = W_v^T @ W_v (n x n, only n=512)
* G -= I
* dW_v = 4 * lambda * W_v @ G
* Cost: 2 * n^2 * n = 2 * 512^3 ≈ 268M FLOPs per layer (negligible vs training) */
if (!g_skip_wv) {
float lambda_ortho = 0.05f; /* v13l: increased 0.02→0.05 for stronger rank promotion */
/* Allocate G on stack: n x n = 512*512 = 262144 floats = 1MB */
/* Use static to avoid stack overflow */
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;
}
/* Step 1: G = W_v^T @ W_v (original strided version) */
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;
}
}
/* Step 2: G -= I */
/* v13l: compute effective rank (participation ratio) before I subtraction
* eff_rank = (trace(G))^2 / trace(G^2) = (sum s_i^2)^2 / sum(s_i^4)
* Full rank → n, rank-1 → 1. Monitor this to track rank improvement. */
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]; /* Frobenius of G = trace(G^2) for symmetric */
}
}
float eff_rank = (tr_G * tr_G) / (tr_G2 + 1e-12f);
for (int i = 0; i < n; i++)
G[i * n + i] -= 1.0f;
/* Step 3: dW_v = 4 * lambda * W_v @ G, apply directly to w_float */
/* W_v[k][i] -= 4 * lambda * sum_j W_v[k][j] * G[j][i] */
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;
}
}
/* Log orthogonal regularization stats every 50 steps */
if (g_opt_step % 50 == 49) {
/* Recompute Frobenius of (W^T W - I) for monitoring */
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);
}
}
/* v11b: Orthogonal regularization on W_o (attn_o)
* 和 W_v 同样的正则化, 防止 W_o rank-1 退化
* W_o 是 b==1, 独立的 n×n 矩阵 (不是 QKV merged) */
if (b == 1) {
float lambda_ortho = 0.05f; /* v13l: increased 0.02→0.05 */
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;
/* G = W_o^T @ W_o (W_o shape [n, in=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_o];
dot += wf_row[i] * wf_row[j];
}
Go[i * n + j] = dot;
Go[j * n + i] = dot;
}
}
/* G -= I */
for (int i = 0; i < n; i++) Go[i * n + i] -= 1.0f;
/* dW_o = 4 * lambda * W_o @ G */
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);
}
}
}
/* v12: Orthogonal regularization on W_o (attn output projection)
* Same formula as W_v: Loss += lambda * ||W_o^T @ W_o - I||^2_F
* Gradient: dW_o = 4 * lambda * W_o @ (W_o^T @ W_o - I)
*
* v11 SVD showed W_o eff_rank=5-11 (severely rank-deficient).
* This causes layer collapse: different inputs project to same
* low-dimensional subspace → cosine(火,水)→1.0 after attention.
*
* W_o is the entire BinLayer (b==1), shape [n_embd, n_embd].
* Simpler than W_v (no QKV merge offset needed). */
if (b == 1) {
int n = m->cfg.n_embd;
int in = bl->in_dim;
float lambda_ortho_o = 0.05f; /* v13l: increased 0.02→0.05 */
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;
}
/* Go = W_o^T @ W_o (original) */
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;
}
}
/* Go -= I */
for (int i = 0; i < n; i++)
Go[i * n + i] -= 1.0f;
/* dW_o = 4 * lambda * W_o @ Go, apply directly */
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;
}
}
/* Log W_o orthogonal stats every 50 steps */
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);
}
}
/* Weight clipping + repack: per-neuron based on logic_mask.
* CORE (float): ±2.0 — needs room for precise differentiation.
* BINARY (sign): ±1.0 — must stay near ±1 for sign function.
* PRUNE: already skipped in update loop above. */
if (!g_use_pure_float) {
for (int j = 0; j < out; j++) {
float clip_val = 1.0f; /* BINARY default */
if (bl->logic_mask && bl->logic_mask[j] == 0)
clip_val = 2.0f; /* CORE: allow larger float weights */
/* PRUNE (mask==2) already skipped, but clip anyway for safety */
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
}
}
}
/* === CRITICAL FIX: Update token embeddings (wte) with Adam ===
* Without this, embeddings are frozen and the model cannot learn
* concept boundaries. This is the #1 fix for LAL whitebox training. */
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;
/* [加速] OpenMP 并行: 每个 vocab token 独立更新 */
#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];
/* BUG #52 FIX (v2 - gentler): Track if this token had any gradient this step.
* Tokens not in training data keep random init → high logit → sampled → garbage.
* Apply MILD weight decay (×0.9999) only to unused tokens with large norm.
* Previous v1 (×0.999) was too strong, shrunk all wte → CORE diff collapsed. */
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;
/* v17c: v_wte floor - prevent Adam cold-start amplification.
* When a token first receives gradient (from logic_reg or C3),
* va[i] = (1-beta2)*g^2 is tiny -> vh ~ sqrt(0.001)*|g| ~ 0.032*|g|,
* update = lr*mh/vh ~ lr/0.032 ~ 31*lr (amplified 30x).
* This crashed boundary 78->14 at step 100 in v16/v17.
* Floor vh to 1e-4 caps amplification at ~10x, logic_reg/C3 stay safe.
* Fix verified: v17c step 100 logic_reg trigger, boundary stable at 77. */
if (vh < 1e-4f) vh = 1e-4f;
w[i] -= lr * g_wte_lr_scale * mh / vh; /* v16: wte 对齐泵降速 */
}
}
/* BUG #52 v2: Only decay unused tokens with large norm (threshold-based) */
if (!has_grad) {
/* Compute norm, only decay if above average to avoid shrinking all embeddings */
float norm_sq = 0.0f;
for (int i = 0; i < n; i++) norm_sq += w[i] * w[i];
float norm = sqrtf(norm_sq);
/* Only decay if norm > 0.5 (above typical init scale 1/sqrt(n)≈0.042) */
if (norm > 0.5f) {
float decay = 0.9999f; /* Much gentler than v1's 0.999 */
for (int i = 0; i < n; i++) {
w[i] *= decay;
}
}
}
}
}
/* === Update position embeddings (wpe) with Adam + norm clipping ===
* Without this, position embeddings are random noise → model has no
* position awareness → attention collapses all positions → same output.
*
* BUG FIX (v17): wpe norm explosion — wpe[0] reached 11.37 (22.7x wte norm 0.50).
* When wpe dominates wte, token identity is drowned out → model loses semantic
* information → generation produces position-driven garbage, not content-driven.
* Fix: (1) reduced LR (0.3x) slows wpe growth; (2) hard norm clip at 1.0 caps it.
* wpe is added to wte: x = wte[tok] + wpe[pos], so wpe norm should be ≤ wte norm
* to avoid drowning token signal. Cap at 1.0 (2x wte norm) gives learning room. */
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; /* v17: reduced LR to slow wpe growth */
float wpe_max_norm = 1.0f; /* v17: hard cap — 2x typical wte norm */
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; /* v17(远程): wpe 降速 lr*0.3 — 与 v16 对齐泵修复同源 */
}
/* v17: Norm clipping — prevent wpe from dominating wte */
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;
}
}
}
/* === Update LayerNorm weights with proper Adam ===
* Now using correct gradients from layer_norm_backward (grad_w/grad_b).
* Previously these were stuck at init (w=1.0, b=0.0) because
* layer_norm_backward didn't compute grad_w, causing all inputs
* to produce identical final_ln. */
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) {
/* BUG #50 FIX: Use Adam for LayerNorm weights (was SGD+clip, caused norm_w→0) */
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; /* v13c: full LR for LayerNorm weights — Adam handles scaling */
for (int i = 0; i < n; i++) {
/* norm1_w — v13c: enable Adam training with reduced LR + clipping
* Previous BUG #50: SGD with large gradients caused norm_w→0.
* Fix: Adam naturally normalizes gradient scale; 0.1x LR adds safety margin. */
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;
/* v13g: revert to [0.5, 2.0] clip — v13f [0.95, 1.05] killed
* core_diff (2.36→1.90). LN weight growth is BENEFICIAL:
* it amplifies important feature dimensions, aiding concept
* differentiation despite slightly higher cosine similarity. */
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;
}
/* norm1_b */
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;
}
/* norm2_w — v13c: enable Adam training with reduced LR + clipping */
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;
/* v13g: revert to [0.5, 2.0] clip */
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;
}
/* norm2_b */
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;
}
}
}
}
/* ln_f weights with Adam */
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;
/* v13g: revert to [0.5, 2.0] clip */
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;
}
}
/* === Sync w_core and wbits from updated w_float ===
* [加速] 这个 repack 已经在 weight clipping 后做过了 (line 4394),
* 除非 logic_mask 被 100 步重分配改了, 否则不需要再 repack.
* 删除冗余 repack 节省 ~1s/step (40-70 次 repack × O(out×in)). */
/* Sync w_core and wbits from updated w_float
* (恢复: 删除后 Windows 产生 NaN, 可能是 clipping 后状态不一致) */
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]);
}
}
/* Increment Adam step once per batch */
if (g_use_adam) {
/* Ponder 停机单元 Adam 更新 (在 g_opt_step 自增前, bias-correction 用同一步数) */
ponder_apply(m, lr, batch_size, g_opt_step + 1);
g_opt_step++;
}
/* [已禁用] CORE/BINARY/PRUNE 动态重分配.
* 原来每 100 步按 w_float 范数重算 mask, 但训练早期 (step 100) 权重还没分化,
* 重分配会把已学到的概念结构打乱 (boundary 74→16, opp_sim 0.26→0.84).
* 现在 mask 只在 model_load 时算一次, 训练中固定不变.
* 如需调整比例, 修改 g_logic_core_ratio / g_logic_prune_ratio 的初始值. */
}
void model_stateful_begin(Model *m) {
/* Ensure KV cache is allocated */
if (!m->k_cache) model_kv_cache_alloc(m);
/* Reset KV cache to zero */
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);
}
/* Allocate stateful context buffers */
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(); /* v16 */
/* C3 概念图驱动长上下文记忆: 仅当概念图已加载(g_runtime_cg!=NULL)时分配缓冲.
* 与 --concept-graph 一体 —— 不加载图则退化普通滑动窗口. */
{
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;
}
}
/* Use the GLOBAL attention window/sink (g_attn_window / g_attn_sink) so
* inference matches training exactly. The ModelConfig.sliding_window field
* defaults to 9996 and is NOT synced from --attn-window, so relying on it
* here caused a train/infer window mismatch -> garbled generation.
* Single source of truth: the global flags (set by --attn-window/--attn-sink
* and default 1024/64). */
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);
}
/* ─── Stateful Inference: Reset KV Cache ────────────────────────── */
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(); /* v16: 新一轮生成 — 清空信使缓存 */
/* C3 概念图驱动长上下文记忆: 新一轮生成时清零聚合 */
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));
}
}
/* ─── Stateful Forward with Sliding Window ──────────────────────── */
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;
/* Must match training: use GLOBAL g_attn_window / g_attn_sink, not
* ModelConfig.sliding_window (which defaults to 9996 and is not synced
* from --attn-window). See model_stateful_begin() for the same fix. */
int window = g_attn_window > 0 ? g_attn_window : ctx;
int n_sinks = g_attn_sink;
/* Circular buffer: no need to shift. Just wrap around. */
int pos = g_sctx.kv_pos; /* logical position in cache */
int abs_pos = g_sctx.total_pos; /* absolute position in sequence */
int pe_pos = (m->cfg.attn_type == ATTN_LEARNED) ? (abs_pos % ctx) : abs_pos;
float *x = g_sctx.x;
/* Embedding lookup + position encoding */
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];
}
/* C3 概念图驱动长上下文记忆: 把被滑动窗口挤出的中间段 token 按"概念归属"
* 聚合进概念状态槽. 概念归属由来概念图给出: anchor = neighbor[eject*K+0]
* (该 token 在 wte 几何里最近的概念 token), slot = anchor % LCTX_SLOTS.
* 同一份概念图既引导生成(graph_concept_bias)又驱动长上下文记忆 — 一体. */
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]; /* 最近概念 token */
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;
}
}
}
/* Forward through layers with sliding window attention */
if (g_ponder_cfg.enable && m->ponder_ready) {
/* PonderNet 循环思考推理: 混合读出 + 早退 + 思考深度统计 */
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);
}
/* Final norm + logits (tied embeddings) */
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;
/* Logits: raw dot product (tied embeddings).
* Cosine normalization removed — it compressed logit range too much,
* making sampling unable to distinguish good tokens from noise.
* Repetition penalty in generation handles mode collapse instead.
*
* BUG #42 FIX: was a scalar loop (1 mul-add per iteration).
* Same computation as compute_full_logits and model_forward_float_logits,
* which both use 8-way unrolled loops for SIMD vectorization.
* The scalar version was 4-8x slower on vocab=32768. Now matches. */
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; /* v16 */
}
/* Advance circular buffer pointer */
g_sctx.kv_pos = (g_sctx.kv_pos + 1) % ctx;
g_sctx.total_pos++;
return g_sctx.logits;
}
/* ─── Configure Sliding Window at Runtime ───────────────────────── */
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);
}
/* ========================================================================
* ========================================================================
* Concept-Aware Attention (基于「理解(概念-边界) + 推理(关系演化)」框架)
* ========================================================================
* ========================================================================
* 四层设计实现:
* Layer 1: 基于概念边界的语义片段切分 + segment-messenger
* Layer 2: 关系强度门控(概念边界预筛选)
* Layer 3: 异构多头算力分配(不同头不同访问域)
* Layer 4: 推理侧 KV-Cache 概念复用(含信使 cache)
*
* 设计原点:
* - 注意力的计算开销来自"两两概念对的关系匹配"
* - 优化原则:保留真实需要建立关系的概念对的完整 Q-K 匹配
* - 对于边界隔离、本就弱关系的概念对,要么过滤,要么走信使间接通信
*
* 本优化只改造理解阶段(Attention)的信息交互通路,不改动 FFN 推理演化逻辑。
* ======================================================================== */
/* 全局概念感知注意力配置(默认关闭,需显式开启) */
ConceptAttnConfig g_concept_attn_cfg = {0};
/* C3 概念图驱动的长上下文记忆配置(仅当概念图已加载时启用) */
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);
}
/* v16: 注意力残差配额 (v13j 防塌缩设计, 默认 0.15). LAL_ATTN_RES_SCALE 可调 —
* 配额太低时注意力架构变化在输出端"隐形" */
float g_attn_res_scale = 0.15f;
/* ConceptAttnStats 结构 + g_ca_stats + concept_attn_stats_reset 已上移到
* attention_forward_concept_ctx 之前 (line 2868+), 因为 ctx 函数需要统计.
* 此处保留此注释作为指针. */
/* === Bug Fix 2: 门控分数运行统计 (自适应分位数阈值) ===
* 维护一个滑动窗口记录最近的门控分数,计算 P25 分位数作为阈值
* g_gate_ring: 循环缓冲区, g_gate_head: 写入位置, g_gate_n_samples: 已有样本数 */
#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; /* P25 分位数 (初始回退到默认阈值) */
/* 记录门控分数到滑动窗口,并更新 P25 分位数 */
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++;
/* 每 64 个样本重新计算一次 P25 (避免频繁排序) */
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));
/* 简单插入排序 (n <= 512, 复杂度可接受) */
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;
}
/* P25 = 第 25 百分位 */
int idx = (int)(0.25f * (n - 1));
g_gate_p25 = tmp[idx];
}
}
/* 全局信使缓存(每层一个,按 layer_idx 索引) */
MessengerCache *g_messenger_caches = NULL;
static int g_messenger_caches_n_layer = 0;
/* ─── Layer 1: Messenger Cache Management ─────────────────────── */
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;
}
/* ─── Layer 1: Segment Messenger Generation ──────────────────────
* 在每个 segment 内部,基于本片段全部 V,聚合生成少量信使向量。
* 信使是本片段全部概念与关系状态的压缩载体。
*
* 聚合策略:均匀分桶 + 均值池化
* - 将片段内 V[0..seg_len-1] 均匀分成 num_messengers 个桶
* - 每个桶内做均值池化,得到一个信使向量
* - 信使的 K = 信使的 V(自关联,简化)
*
* 语义意义:远方片段的整体语义,由信使代为表达。
* 普通token通过信使间接获得远方概念集合的状态。
*/
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;
/* === Bug Fix 1: 信使去中心化 ===
* 原始实现: 信使 = 桶内 V 的均值 → 携带公共模式,4个信使几乎相同
* 修复: 先计算片段内全局 V 均值,信使 = 桶均值 - 全局均值
* 只保留偏差信息(本桶的"特色"而非"共识")
* 同时对信使做范数钳制,防止越训越大 */
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;
}
/* 计算片段内全局 V 均值 */
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;
}
/* 范数钳制上限: 嵌入维度的 sqrt(n_embd) 量级 */
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; /* 信使 K = V(自关联简化) */
}
/* 范数钳制: 防止信使范数失控膨胀 */
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);
}
/* ─── Layer 2: Concept Boundary Gate (关系强度门控) ───────────────
* 给定 token-i(Q侧)、token-j(K侧),利用距离先验 + 粗粒度相似度
* 快速预判:如果预判两个概念边界隔离,潜在关系极弱,
* 直接把该位置置 -inf,不参与完整内积计算。
*
* 软门控(保留回退通路,避免硬切断长距离指代):
* sim_coarse = <Q_i, K_j> / (||Q_i|| * ||K_j|| + eps)
* dist_prior = exp(-distance / tau) // tau = segment_len
* gate_score = sim_coarse + gate_distance_prior * dist_prior * 0.5
* if gate_score < gate_threshold:
* 以 (1 - gate_fallback_prob) 概率屏蔽
* 以 gate_fallback_prob 概率保留(回退通路)
*
* 返回:1 = 保留(参与完整 QK 计算),0 = 屏蔽(置 -inf)
*/
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; /* 门控禁用,全部保留 */
/* 计算粗粒度余弦相似度(用前 1/4 维度做快速预判,省算力) */
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; /* 距离近的 token 有先验加分 */
}
/* === Bug Fix 2: 自适应分位数阈值 ===
* 原始实现: 固定阈值 0.1,但训练中分数分布整体漂移(52% > 0.5)
* 导致门控要么全开要么全关,无法稳定兑现"概念边界隔离"
* 修复: 使用运行统计的分位数作为阈值
* g_gate_score_p25 = 观察到的分数分布的 25th percentile
* 屏蔽最低 25% 的分数对,而非用一个死阈值
*
* 原理: 不管训练如何移动分数的绝对值,分位数始终代表
* "当前分布中关系最弱的 25%"——这才是"边界隔离"的语义 */
float effective_threshold = cfg->gate_threshold; /* 默认回退 */
/* 使用运行统计的分位数(如果可用) */
if (g_gate_n_samples > 50) {
/* 有足够样本时,用 P25 分位数作为阈值 */
effective_threshold = g_gate_p25;
}
/* 软门控判定 */
if (gate_score < effective_threshold) {
/* 概率回退通路:用 hash(distance, sim) 做确定性伪随机,
* 避免引入 rand() 影响可复现性 */
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; /* 屏蔽:概念边界隔离,不参与完整 QK 计算 */
}
return 1; /* 保留:可能存在有效关系 */
}
/* ─── Layer 3: Heterogeneous Head Access Configuration ───────────
* 不同类型关系本身就有不同的"概念交互范围",不需要统一全序列扫描。
* - 头A:局部语法关系(主谓宾、修饰):强局部性,适合小窗口。
* - 头B:指代、实体绑定:偶尔需要长距离跳跃。
* - 头C:因果、时序关系:中等范围依赖。
*/
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;
/* 自动分配:local = n_head/2, messenger = n_head/4, global = 剩余 */
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; /* Bug Fix 3: 不再减半, 保持对称 */
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);
}
/* ─── Layer 4: Concept-Aware Attention Forward (主入口) ───────────
* 概念感知注意力前向传播。整合四层优化:
* 1. 切成语义片段(segment_len)
* 2. 片段内部:完整 QKV,充分做片段内概念理解
* 3. 生成本片段信使:聚合本片段全部概念-关系状态
* 4. 本片段普通 token:只和【局部窗口 + 本片段信使 + 邻近片段信使】做匹配
* 5. 关系门控:过滤边界隔离的概念对(Layer 2)
* 6. 异构多头:不同头不同访问域(Layer 3)
* 7. KV-Cache:历史 K/V 直接复用,信使也进 cache(Layer 4)
*
* 数学复杂度(设片段长度 L,每个片段信使数目 S,S << L):
* - 片段内部:O(n L d)
* - 信使交互:O((n/L * S)^2 d),该项很小
* - 普通token与信使:O(n * S * d),远小于 O(n^2 d)
*/
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) {
/* 主开关关闭 → 回退到 sliding window attention (端到端统一) */
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; /* 概念注意力内部用 seq_pos 直接索引 cache,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;
/* Layer 4: 写入 KV-Cache(与标准 attention_forward 一致) */
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));
/* Layer 1: 片段切分 + 信使生成
* 当前 token 属于片段 seg_idx = seq_pos / segment_len
* 当一个片段的最后一个 token 处理完时,生成本片段的信使 */
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;
/* Layer 1: 当片段填满时(actual_seg_len >= seg_len)或这是该片段最后一个 token
* 时,生成/更新该片段的信使。
* 修复:短样本(对话数据平均 11 token)尾部若累积 ≥ min_seg_len 也强制封口,
* 否则信使机制永远空转、概念注意力在训练时收不到梯度。*/
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]) {
/* 从 v_cache 取本片段的 V,生成信使 */
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;
/* 探针: 信使间相似度 + 信使范数 (审查建议的核心验证项)
* 去同质化目标: 信使间余弦 < 0.2 说明信使携带的是"差异"而非"共识均值"
* 范数钳制目标: 信使范数应被 MSG_NORM_CAP=4.0 约束, 不被注意力按范数主导 */
{
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++;
}
}
}
/* 构建当前 token 的注意力候选集:
* - 局部窗口:[max(0, seq_pos - window), seq_pos]
* - 本片段信使(如果当前片段已完成)
* - 邻近片段信使(前 messenger_neighbors 个已完成的片段)
* - 全局头:可以访问全部已生成信使
*
* 注意:候选集大小受限于 scratch buffer(10240) */
int n_attend = 0;
int pos_list[10240];
int is_messenger[10240]; /* 标记该位置是信使还是普通 token */
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;
/* 1. 局部窗口(因果:只看 seq_pos 之前 + 自己) */
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; /* -1 = 普通 token, >=0 = 信使索引 */
n_attend++;
}
/* 2. 本片段信使 + 邻近片段信使(仅 MESSENGER/GLOBAL 头) */
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;
/* 邻近片段:seg_idx - n_neighbor .. seg_idx - 1(已完成的) */
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 = -1, messenger_idx = s * num_messengers + 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;
}
/* v16 探针: 候选集统计。
* 修正(指标口径 bug):candidates 只累计【普通 token 候选】(窗口内的真实 token),
* 信使成本单独计入 msg_candidates。否则短样本下「窗口截断到 seq_pos+1 + 信使」
* 会让 n_attend 超过 full_equiv,导致"候选精简"显示为负,误导为机制失效。
* 概念注意力的精简收益来自「用少量信使替代大量历史 token」,普通 token 候选
* 应 ≤ 窗口(截断到 seq_pos+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) {
/* 信使 K */
int msg_idx = is_messenger[i];
K_jh = mc->messenger_k + (size_t)msg_idx * n_embd + h * head_dim;
} else {
/* 普通 token K(从 KV cache) */
int j = pos_list[i];
int phys_j = j % n_ctx;
K_jh = k_cache + (size_t)phys_j * n_embd + h * head_dim;
}
/* Layer 2: 关系强度门控(仅对普通 token,信使总是保留) */
if (is_messenger[i] < 0 && cfg->gate_enable) {
int j = pos_list[i];
int distance = seq_pos - j;
g_ca_stats.gate_pairs++;
/* Bug Fix 2: 记录门控分数到滑动窗口以计算自适应分位数阈值 */
{
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;
}
/* Softmax */
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;
/* v16 探针: 信使注意力质量 */
for (int i = 0; i < n_attend; i++)
if (is_messenger[i] >= 0) g_ca_stats.msg_mass += attn_w[i];
/* 加权求和 V */
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];
}
}
}
/* ─── Layer 4: Concept-Aware Attention Backward ──────────────────
* 概念感知注意力反向传播。计算当前 token 的 Q/K/V 梯度。
* 缓存的 K/V(位置 0..seq_pos-1)视为常量(与 attention_backward 一致)。
* 信使视为常量(不回传梯度到信使生成路径,简化实现)。
*/
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) {
/* 主开关关闭 → 回退到 sliding window attention backward (端到端统一) */
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; /* -1 = 普通 token, >=0 = 信使索引 */
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;
}
/* 重算 scores + softmax(K 在 cache 中) */
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;
/* g_w[i] = <g_out, V_i> */
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];
}
/* g_scores[i] = w[i] * (g_w[i] - <g_w, w>) */
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);
}
/* g_Q[d] += sum_i g_scores[i] * K_i[d] * scale */
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];
}
/* 当前 token 的 K/V 梯度(只在 seq_pos 在候选集中时) */
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) {
/* g_K_cur[d] += g_scores[self_idx] * Q[d] * scale */
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];
/* g_V_cur[d] += w[self_idx] * g_out[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];
}
}
}
/* ─── Integration Guide: trans_layer_forward_concept ─────────────
* trans_layer_forward 的实现包含比例缩放、残差归一化等复杂逻辑,
* 完整复制易引入 bug。推荐集成方式:
*
* 在现有 trans_layer_forward() 中,将 attention_forward 调用替换为:
*
* if (g_concept_attn_cfg.enable && g_messenger_caches) {
* attention_forward_concept(act->attn_out, qkv_ptr,
* n, cfg->n_head, abs_pos,
* tl->kv_k, tl->kv_v, cfg->n_ctx,
* &g_concept_attn_cfg,
* &g_messenger_caches[layer_idx]);
* } else {
* attention_forward(act->attn_out, qkv_ptr, n, cfg->n_head,
* abs_pos, tl->kv_k, tl->kv_v);
* }
*
* 反向传播同理:将 attention_backward 替换为 attention_backward_concept。
*
* 通过 model_set_concept_attn() 在运行时配置,无需修改模型结构。
* ──────────────────────────────────────────────────────────────────── */
/* ─── Global Messenger Cache Management (per-layer) ──────────────
* 在 model_load 时分配,model_free 时释放。
* 每层一个 MessengerCache,按 layer_idx 索引。
*/
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; /* +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]);
}
/* ─── Configure Concept-Aware Attention at Runtime ─────────────── */
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");
}
}
/* v16: 概念注意力探针 — 输出聚合统计并重置计数器 */
void concept_attn_probe_print(void) {
ConceptAttnStats *s = &g_ca_stats;
/* 诊断: 若两个版本都没走过前向 (forwards==0 && forwards_ctx==0), 打印根因.
* 修复 (2026-08-17): 旧探针只看 forwards (简单版), 但训练 forward 实际
* 走的是 attention_forward_concept_ctx (长上下文记忆版), 导致 fwd=0 假警报
* 团队持续误以为概念注意力没参与前向. 现在两个计数都看. */
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;
/* 审查判定: 信使间余弦 < 0.2 才算去同质化达标 (机制潜力挖完的判据) */
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);
/* ctx 版本专属统计: 概念槽命中情况 */
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();
}