lalmodel-code / src /runtime /lal_ponder.h
gasschina's picture
feat: PonderNet 循环思考 (逐层停机 + 末块循环 双重混合)
1423b78
Raw
History Blame Contribute Delete
17.9 kB
/* 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·Π_{i<l}(1-p_i), 末步 P 取余 (强制分布归一)
*
* 梯度 (闭式推导, 见 PONDER_NOTES.md 附录):
* dL_AL/dp_j = -prior_j/p_j + tailprior_j/(1-p_j), tailprior_j = Σ_{l>j} 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, s_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 <math.h>
#include <stdio.h>
#include <stdlib.h>
#include <string.h>
#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 = <G, s_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 */