/* runtime/lal_ponder.h — PonderNet 式循环思考 (adaptive halting) * * 设计原点(Banino et al., 2021, "PonderNet: Learning to Ponder" 的机制思想, * 结合本仓库 三值STE/滑窗KV-cache 的工程约束重新推导): * - 不同 token 需要的计算深度不同: 简单 token 浅层即可预测, 困难 token 需要更深层计算。 * - 固定深度前向把算力平均分配给所有 token, 既浪费(简单 token)又可能不足(困难 token)。 * - PonderNet: 每个"计算步"输出停机概率 p, 构成停机分布, 读出用期望加权, * 训练用 任务loss + 停机分布NLL(几何先验) + 不停机惩罚 三项联合塑形。 * * 双重混合模式 (可独立开关, 供消融实验): * A. 逐层停机 (LAL_PONDER_LAYER=1): 层 0..L-2 每层一个停机单元, 读出状态 = 各层期望加权和。 * B. 块内循环 (LAL_PONDER_REC=N≥2): 末层 block 以共享权重循环迭代最多 N 次 * (Universal Transformer 风格), 每次迭代一个停机概率(共享单元), 迭代内 PonderNet 混合。 * 两者可同时开启: 步序列 = 层0..层L-2 + 末块迭代0..R-1, 最后一步取剩余质量(remainder)。 * * 关键工程决策 (与代码库约束的适配): * 1. 读出混合 (read-out mixing): 各层照常完整前向 (KV cache 语义不变), 停机只影响 * 最终 norm+logits 使用的状态: out = Σ_l c_l·s_l。数学上等价于 * out = x_0 + Σ_j reach_j·Δ_j (每层残差增量按"到达概率"缩放), 见 PONDER_NOTES.md。 * 2. 停机单元输入 detach (stop-grad): 任务梯度不通过 c_l 流回主干 (避免 O(L²) 反传), * 改为显式辅助梯度注入停机单元参数 — 保留"任务 loss 教会模型何时停机"的自适应性。 * 3. 停机单元用全精度浮点小线性层 (n_embd→1) + 独立 Adam, 不做三值量化 * (停机概率需要连续细粒度, 三值会退化成硬开关)。 * 4. 块内循环的 KV cache: 每个 token 对外只暴露"第一次迭代"的 K/V (与 context * prefill 的单遍语义一致); 后续迭代内部 attention 使用自身 K/V, 迭代结束即恢复。 * 5. 推理早退: 累计停机概率 Q ≥ threshold 时停止计算更深层的 block, 剩余质量直接 * 加在当前状态上 (假设已收敛近似), 被跳过的物理层用 kv_only 填充 K/V 保证 * 后续 token 的 attention 完整。输出平均思考深度统计。 * * 损失 (PonderNet 式): * L = L_CE + β·L_AL + γ·L_P * L_AL = -Σ_l prior_l·log(P_l + eps) 停机分布 vs 几何先验 的交叉熵 * L_P = Σ_l (1 - Q_l)² 每步未停机剩余质量的平方惩罚 * prior_l = λ(1-λ)^l, prior 末步取余 保证 Σprior = 1 * P_l = p_l·Π_{ij} prior_l * dL_P/dp_j = -(2/(1-p_j))·Σ_{l≥j} Rm[l+1]², Rm[l+1] = 步 l 后的剩余质量 * dL_CE/dp_j = Rm[j]·g_j - tailG_j/(1-p_j), g_j = (G=读出梯度), * tailG_j = Σ_{l>j} c_l·g_l * d(pre-act)_j = dp_j · p_j(1-p_j) * * 环境变量: * LAL_PONDER=0 总开关 (0=完全回退标准前向, 行为与改造前逐位一致) * LAL_PONDER_LAYER=0/1 逐层停机 (默认 1) * LAL_PONDER_REC=N 末块迭代次数, 1=关闭 (默认 2) * LAL_PONDER_LAMBDA=x 几何先验 λ (默认 0.5) * LAL_PONDER_BETA=x L_AL 权重 β (默认 0.01) * LAL_PONDER_GAMMA=x L_P 权重 γ (默认 0.01) * LAL_PONDER_LR=x 停机单元学习率相对 CE lr 的倍率 (默认 1.0) * LAL_PONDER_THRESHOLD=x 推理早退累计概率阈值 (默认 0.99) * LAL_PONDER_MIN_LAYER=k 推理最早可在第 k 步早退 (默认 1) * * This is a single-header library. In exactly ONE .c translation unit define * LAL_PONDER_IMPLEMENTATION before including to get the definitions. */ #ifndef LAL_PONDER_H #define LAL_PONDER_H #include #include #include #include #define LAL_PONDER_MAX_STEPS 40 /* (n_layer-1) + rec_iters 上限: 31+8 覆盖所有阶段 */ #define LAL_PONDER_MAX_REC 8 /* rec_iters 硬上限 */ #define LAL_PONDER_EPS 1e-6f /* 停机概率钳制 */ #define LAL_PONDER_LOG_EPS 1e-12f /* log 数值下限 */ /* ------------------------------------------------------------------ */ /* 配置 */ /* ------------------------------------------------------------------ */ typedef struct { int enable; /* 总开关 */ int layer_halt; /* A. 逐层停机 */ int rec_iters; /* B. 末块迭代次数 (1=关闭循环) */ float lambda; /* 几何先验 λ */ float beta; /* L_AL 权重 */ float gamma; /* L_P 权重 */ float lr_scale; /* 停机单元 lr = CE lr × lr_scale */ float threshold; /* 推理早退阈值 (累计停机概率) */ int infer_min_layer; /* 推理最早早退步 */ } PonderConfig; extern PonderConfig g_ponder_cfg; /* 最近一次训练前向的 ponder 指标 (ste_train.c 日志读取, 末次预测点值) */ extern float g_ponder_last_al; /* L_AL */ extern float g_ponder_last_p; /* L_P */ extern float g_ponder_last_mean; /* E[停机步] */ /* ------------------------------------------------------------------ */ /* 停机单元 (全精度浮点小线性层 n_embd → 1, 独立 Adam) */ /* ------------------------------------------------------------------ */ typedef struct { int in_dim; float *w; /* [in_dim] */ float b; float *grad_w; /* [in_dim] */ float grad_b; float *m_w, *v_w; /* Adam 一阶/二阶矩 */ float m_b, v_b; } PonderLayer; void ponder_layer_alloc(PonderLayer *u, int in_dim); void ponder_layer_init(PonderLayer *u); /* w=0, b=logit(lambda) → 初始分布=几何先验 */ void ponder_layer_free(PonderLayer *u); /* ------------------------------------------------------------------ */ /* 停机分布缓冲 (per-thread / per-inference 上下文) */ /* ------------------------------------------------------------------ */ typedef struct { int n_steps; /* 总步数 (含末步 remainder) */ int n_param; /* 有停机参数的步数 (= n_steps - 1) */ float p[LAL_PONDER_MAX_STEPS]; /* 停机概率 (参数步) */ float c[LAL_PONDER_MAX_STEPS]; /* 停机分布 P_l (贡献权重) */ float Rm[LAL_PONDER_MAX_STEPS]; /* Rm[l] = 步 l 之前的剩余质量 (Rm[0]=1) */ float prior[LAL_PONDER_MAX_STEPS]; /* 几何先验 */ float gdot[LAL_PONDER_MAX_STEPS]; /* 反向时填: g_l = */ float loss_al; float loss_p; float mean_step; /* E[停机步] (0-indexed) */ int active; } PonderBuf; /* 由各参数步的停机概率填充分布 + 计算 L_AL / L_P / E[步] */ void ponder_dist_fill(PonderBuf *pb, const float *halts, int n_param); /* 停机概率: p = σ(w·x + b) (x 为调用方 detach 后的状态) */ float ponder_halt(const PonderLayer *u, const float *x, int n); /* 反向: 依据 gdot 计算各参数步 d(pre-act)_j (调用方负责累加进单元参数梯度) */ void ponder_grad(PonderBuf *pb, float *dpre); /* ------------------------------------------------------------------ */ /* 推理侧统计 (思考深度) */ /* ------------------------------------------------------------------ */ typedef struct { long n_tokens; /* 统计的 token 数 */ double depth_sum; /* Σ E[步] (0-indexed) */ int early_exits; /* 触发早退的 token 数 */ int last_depth; /* 最近 token 的 round(E[步]) */ float last_c[LAL_PONDER_MAX_STEPS]; /* 最近 token 的停机分布 (demo/可视化用) */ int last_n_steps; } PonderStats; extern PonderStats g_ponder_stats; void ponder_stats_reset(void); void ponder_stats_record(const PonderBuf *pb, int early_exit); void ponder_stats_print(const char *tag); /* ------------------------------------------------------------------ */ /* 配置解析 (env) — 启动时调用一次 */ /* ------------------------------------------------------------------ */ void ponder_config_from_env(void); #ifdef LAL_PONDER_IMPLEMENTATION PonderConfig g_ponder_cfg = {0}; float g_ponder_last_al = 0.0f; float g_ponder_last_p = 0.0f; float g_ponder_last_mean = 0.0f; PonderStats g_ponder_stats = {0}; static float ponder_sigmoidf(float z) { if (z > 30.0f) z = 30.0f; if (z < -30.0f) z = -30.0f; return 1.0f / (1.0f + expf(-z)); } void ponder_layer_alloc(PonderLayer *u, int in_dim) { memset(u, 0, sizeof(*u)); u->in_dim = in_dim; u->w = calloc(in_dim, sizeof(float)); u->grad_w = calloc(in_dim, sizeof(float)); u->m_w = calloc(in_dim, sizeof(float)); u->v_w = calloc(in_dim, sizeof(float)); u->b = 0.0f; u->grad_b = 0.0f; u->m_b = 0.0f; u->v_b = 0.0f; } void ponder_layer_init(PonderLayer *u) { /* w=0, b=logit(λ): 初始每步停机概率 = λ → 初始停机分布恰为几何先验, * L_AL 初始梯度最小, 让任务 loss 从第一步起主导塑形停机行为 */ for (int i = 0; i < u->in_dim; i++) u->w[i] = 0.0f; float lam = g_ponder_cfg.lambda; if (lam < 0.01f) lam = 0.01f; if (lam > 0.99f) lam = 0.99f; u->b = logf(lam / (1.0f - lam)); for (int i = 0; i < u->in_dim; i++) u->grad_w[i] = 0.0f; u->grad_b = 0.0f; for (int i = 0; i < u->in_dim; i++) { u->m_w[i] = 0.0f; u->v_w[i] = 0.0f; } u->m_b = 0.0f; u->v_b = 0.0f; } void ponder_layer_free(PonderLayer *u) { free(u->w); free(u->grad_w); free(u->m_w); free(u->v_w); memset(u, 0, sizeof(*u)); } /* 停机概率: p = σ(w·x + b), x 为调用方 detach 后的状态 */ float ponder_halt(const PonderLayer *u, const float *x, int n) { float z = u->b; for (int i = 0; i < n; i++) z += u->w[i] * x[i]; float p = ponder_sigmoidf(z); if (p < LAL_PONDER_EPS) p = LAL_PONDER_EPS; if (p > 1.0f - LAL_PONDER_EPS) p = 1.0f - LAL_PONDER_EPS; return p; } void ponder_dist_fill(PonderBuf *pb, const float *halts, int n_param) { /* n_param 个参数步 + 1 个 remainder 步 */ int N = n_param + 1; if (N > LAL_PONDER_MAX_STEPS) N = LAL_PONDER_MAX_STEPS; n_param = N - 1; pb->n_steps = N; pb->n_param = n_param; pb->active = 1; float lam = g_ponder_cfg.lambda; /* 几何先验: prior_l = λ(1-λ)^l, 末步取余保证归一 */ { float acc = 0.0f; for (int l = 0; l < n_param; l++) { pb->prior[l] = lam * powf(1.0f - lam, (float)l); acc += pb->prior[l]; } pb->prior[n_param] = 1.0f - acc; if (pb->prior[n_param] < LAL_PONDER_LOG_EPS) pb->prior[n_param] = LAL_PONDER_LOG_EPS; } /* 前向分布: c_l = Rm[l]·p_l, 末步 c = 余量 */ float Rm = 1.0f, Q = 0.0f; pb->Rm[0] = 1.0f; for (int l = 0; l < n_param; l++) { float p = halts[l]; if (p < LAL_PONDER_EPS) p = LAL_PONDER_EPS; if (p > 1.0f - LAL_PONDER_EPS) p = 1.0f - LAL_PONDER_EPS; pb->p[l] = p; pb->c[l] = Rm * p; Q += pb->c[l]; Rm = 1.0f - Q; if (Rm < 0.0f) Rm = 0.0f; pb->Rm[l + 1] = Rm; } pb->p[n_param] = 1.0f; /* 末步强制停机 (remainder 语义) */ pb->c[n_param] = Rm; pb->Rm[N] = 0.0f; /* 末步之后剩余为 0 (防越界读) */ /* L_AL = -Σ prior_l·log(c_l) (含末步 remainder 项) */ float al = 0.0f; for (int l = 0; l < N; l++) { float cl = pb->c[l]; if (cl < LAL_PONDER_LOG_EPS) cl = LAL_PONDER_LOG_EPS; al -= pb->prior[l] * logf(cl); } pb->loss_al = al; /* L_P = Σ_{l=0}^{N-2} (1-Q_l)² = Σ Rm[l+1]² (末步后余量恒 0 不计) */ float lp = 0.0f; for (int l = 1; l < N; l++) { float r = pb->Rm[l]; lp += r * r; } pb->loss_p = lp; /* E[停机步] = Σ l·c_l (0-indexed) */ float ms = 0.0f; for (int l = 0; l < N; l++) ms += (float)l * pb->c[l]; pb->mean_step = ms; } void ponder_grad(PonderBuf *pb, float *dpre) { int n_param = pb->n_param; int N = pb->n_steps; /* --- 任务梯度 (exact, detach 状态): dp_task_j = Rm[j]·g_j - tailG_j/(1-p_j) --- */ /* --- L_AL 梯度: dp_al_j = -prior_j/p_j + tailprior_j/(1-p_j) --- */ /* --- L_P 梯度 (后缀和): dp_p_j = -(2/(1-p_j))·TailSq_j, * TailSq_j = Σ_{l≥j}^{N-2} Rm[l+1]² ---------------------------------------- */ /* 后缀和 tailG_j = Σ_{l>j} c_l·g_l */ float tailG[LAL_PONDER_MAX_STEPS]; float tailPrior[LAL_PONDER_MAX_STEPS]; float tailSq[LAL_PONDER_MAX_STEPS]; tailG[n_param] = 0.0f; tailPrior[n_param] = 0.0f; tailSq[n_param] = 0.0f; /* TailSq_j = Σ_{l=j}^{N-2} Rm[l+1]²: 递推 TailSq_j = TailSq_{j+1} + Rm[j+1]² */ for (int j = n_param - 1; j >= 0; j--) { tailG[j] = tailG[j + 1] + pb->c[j + 1] * pb->gdot[j + 1]; tailPrior[j] = tailPrior[j + 1] + pb->prior[j + 1]; tailSq[j] = tailSq[j + 1] + pb->Rm[j + 1] * pb->Rm[j + 1]; } for (int j = 0; j < n_param; j++) { float p = pb->p[j]; float one_minus_p = 1.0f - p; if (one_minus_p < LAL_PONDER_EPS) one_minus_p = LAL_PONDER_EPS; float dp_task = pb->Rm[j] * pb->gdot[j] - tailG[j] / one_minus_p; float dp_al = -pb->prior[j] / p + tailPrior[j] / one_minus_p; float dp_p = -2.0f * tailSq[j] / one_minus_p; float dp = dp_task + g_ponder_cfg.beta * dp_al + g_ponder_cfg.gamma * dp_p; /* 数值安全钳制 (三值 STE 主干对停机单元的梯度扰动敏感) */ if (dp > 100.0f) dp = 100.0f; if (dp < -100.0f) dp = -100.0f; /* 链到 pre-activation: σ' = p(1-p) */ dpre[j] = dp * p * one_minus_p; } } /* ---------------- 统计 ---------------- */ void ponder_stats_reset(void) { g_ponder_stats.n_tokens = 0; g_ponder_stats.depth_sum = 0.0; g_ponder_stats.early_exits = 0; g_ponder_stats.last_depth = 0; g_ponder_stats.last_n_steps = 0; } void ponder_stats_record(const PonderBuf *pb, int early_exit) { g_ponder_stats.n_tokens++; g_ponder_stats.depth_sum += pb->mean_step; if (early_exit) g_ponder_stats.early_exits++; g_ponder_stats.last_depth = (int)(pb->mean_step + 0.5f); g_ponder_stats.last_n_steps = pb->n_steps; int n = pb->n_steps; if (n > LAL_PONDER_MAX_STEPS) n = LAL_PONDER_MAX_STEPS; for (int l = 0; l < n; l++) g_ponder_stats.last_c[l] = pb->c[l]; } void ponder_stats_print(const char *tag) { if (!g_ponder_cfg.enable || g_ponder_stats.n_tokens == 0) return; float avg = (float)(g_ponder_stats.depth_sum / (double)g_ponder_stats.n_tokens); printf("[PONDER%s] tokens=%ld avg_depth=%.2f early_exit=%ld/%ld last_dist=[", tag ? tag : "", g_ponder_stats.n_tokens, avg, g_ponder_stats.early_exits, g_ponder_stats.n_tokens); int n = g_ponder_stats.last_n_steps; if (n > 8) n = 8; for (int l = 0; l < n; l++) printf("%.3f%s", g_ponder_stats.last_c[l], l + 1 < n ? " " : ""); if (g_ponder_stats.last_n_steps > 8) printf("..."); printf("]\n"); } /* ---------------- 配置解析 ---------------- */ static float ponder_env_float(const char *key, float dflt) { const char *v = getenv(key); if (!v || !v[0]) return dflt; float x = (float)atof(v); return x; } static int ponder_env_int(const char *key, int dflt) { const char *v = getenv(key); if (!v || !v[0]) return dflt; return atoi(v); } void ponder_config_from_env(void) { PonderConfig *c = &g_ponder_cfg; c->enable = ponder_env_int("LAL_PONDER", 1); c->layer_halt = ponder_env_int("LAL_PONDER_LAYER", 1); c->rec_iters = ponder_env_int("LAL_PONDER_REC", 2); if (c->rec_iters < 1) c->rec_iters = 1; if (c->rec_iters > LAL_PONDER_MAX_REC) c->rec_iters = LAL_PONDER_MAX_REC; c->lambda = ponder_env_float("LAL_PONDER_LAMBDA", 0.5f); if (c->lambda < 0.01f) c->lambda = 0.01f; if (c->lambda > 0.99f) c->lambda = 0.99f; c->beta = ponder_env_float("LAL_PONDER_BETA", 0.01f); c->gamma = ponder_env_float("LAL_PONDER_GAMMA", 0.01f); c->lr_scale = ponder_env_float("LAL_PONDER_LR", 1.0f); if (c->lr_scale < 0.0f) c->lr_scale = 0.0f; c->threshold = ponder_env_float("LAL_PONDER_THRESHOLD", 0.99f); if (c->threshold < 0.5f) c->threshold = 0.5f; if (c->threshold > 1.0f) c->threshold = 1.0f; c->infer_min_layer = ponder_env_int("LAL_PONDER_MIN_LAYER", 1); if (c->infer_min_layer < 0) c->infer_min_layer = 0; /* 两者全关 = 等效关闭 */ if (c->enable && !c->layer_halt && c->rec_iters <= 1) { c->enable = 0; } printf("[PONDER] enable=%d layer_halt=%d rec_iters=%d lambda=%.3f beta=%.4f gamma=%.4f " "lr_scale=%.2f threshold=%.2f min_layer=%d\n", c->enable, c->layer_halt, c->rec_iters, c->lambda, c->beta, c->gamma, c->lr_scale, c->threshold, c->infer_min_layer); if (!c->enable) { printf("[PONDER] 已停用 — 前向/反向/推理与改造前完全一致 (baseline)\n"); } else { int nL_est = 0; /* 由调用方补充打印实际步数 */ (void)nL_est; printf("[PONDER] 模式: %s%s%s\n", c->layer_halt ? "逐层停机" : "", (c->layer_halt && c->rec_iters > 1) ? " + " : "", c->rec_iters > 1 ? "块内循环(末块)" : "单遍读出"); } } #endif /* LAL_PONDER_IMPLEMENTATION */ #endif /* LAL_PONDER_H */