upload runtime/lal_concept_attn.h (commit 9ef903f)
Browse files- src/runtime/lal_concept_attn.h +332 -0
src/runtime/lal_concept_attn.h
ADDED
|
@@ -0,0 +1,332 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
/* runtime/lal_concept_attn.h — Concept-Aware Attention
|
| 2 |
+
*
|
| 3 |
+
* 基于「理解(概念-边界) + 推理(关系演化)」框架优化的注意力机制。
|
| 4 |
+
*
|
| 5 |
+
* 设计原点(从框架公理推导,非抄工程现成方案):
|
| 6 |
+
* - Attention 的本职:发现概念之间真实的关系,给关系分配权重,聚合V信息。
|
| 7 |
+
* - 算力浪费根源:标准自注意力强制对全部两两概念对执行 Q-K 匹配,
|
| 8 |
+
* 大量 token-pair 概念边界本身互相隔离,本来就几乎不会产生有效关系,
|
| 9 |
+
* 却依然完整计算 QK^T。
|
| 10 |
+
* - 优化目标:不破坏"识别关系强弱、绑定实体"这个理解能力,
|
| 11 |
+
* 砍掉大量无效概念对的匹配计算。
|
| 12 |
+
*
|
| 13 |
+
* 四层设计:
|
| 14 |
+
* Layer 1: 基于概念边界,区分「需要精细匹配」和「可以间接通信」的概念集
|
| 15 |
+
* - 片段内部:完整多头注意力
|
| 16 |
+
* - 片段之间:通过 segment-messenger 间接通信
|
| 17 |
+
* Layer 2: 关系强度门控,过滤本就边界隔离的概念对
|
| 18 |
+
* - 距离先验 + 粗粒度相似度快速预判
|
| 19 |
+
* - 软门控(保留极小概率回退通路,避免切断长距离指代)
|
| 20 |
+
* Layer 3: 多头的差异化算力分配(异构多头)
|
| 21 |
+
* - 局部语法头:只在局部窗口做匹配
|
| 22 |
+
* - 指代/因果头:可以访问窗口 + 信使
|
| 23 |
+
* Layer 4: 推理侧约束:区分训练阶段和推理 KV-Cache 的概念复用
|
| 24 |
+
* - 历史概念的 K/V 已编码完成,直接复用
|
| 25 |
+
* - 信使 token 也进入 KV-cache,保证间接长距离通路复用
|
| 26 |
+
*
|
| 27 |
+
* 与现有方案的本质区别:
|
| 28 |
+
* - Mistral 滑动窗口:单纯位置窗口,没有信使;长距离依赖能力弱。
|
| 29 |
+
* - Longformer global-token:全局固定几个 token 读取全部位置;
|
| 30 |
+
* 不是每个片段动态聚合本片段语义状态的信使。
|
| 31 |
+
* - Performer 线性注意力:纯数学核近似,不关心"概念边界、关系强弱"。
|
| 32 |
+
*
|
| 33 |
+
* 本优化只改造理解阶段(Attention)的信息交互通路,不改动 FFN 推理演化逻辑。
|
| 34 |
+
* 只要真实关系的匹配通路保留,上层推理能力就不会被破坏。
|
| 35 |
+
*
|
| 36 |
+
* Build: 与 lal_runtime.c 一起编译(#include "lal_concept_attn.h")
|
| 37 |
+
*/
|
| 38 |
+
#ifndef LAL_CONCEPT_ATTN_H
|
| 39 |
+
#define LAL_CONCEPT_ATTN_H
|
| 40 |
+
|
| 41 |
+
#include <stdint.h>
|
| 42 |
+
#include <stddef.h>
|
| 43 |
+
|
| 44 |
+
#ifdef __cplusplus
|
| 45 |
+
extern "C" {
|
| 46 |
+
#endif
|
| 47 |
+
|
| 48 |
+
/* ========================================================================
|
| 49 |
+
* Concept-Aware Attention Configuration
|
| 50 |
+
* ======================================================================== */
|
| 51 |
+
|
| 52 |
+
/* 头的访问域类型(Layer 3: 异构多头) */
|
| 53 |
+
typedef enum {
|
| 54 |
+
HEAD_LOCAL = 0, /* 局部语法关系(主谓宾、修饰):强局部性,小窗口 */
|
| 55 |
+
HEAD_MESSENGER = 1, /* 指代、实体绑定、因果时序:窗口 + 信使 */
|
| 56 |
+
HEAD_GLOBAL = 2, /* 全局兜底(少数头,访问全部历史 + 信使) */
|
| 57 |
+
} HeadAccessType;
|
| 58 |
+
|
| 59 |
+
/* 概念感知注意力配置 */
|
| 60 |
+
typedef struct {
|
| 61 |
+
/* === Layer 1: 片段 + 信使 === */
|
| 62 |
+
int segment_len; /* 语义片段长度 L(token 数)。0 = 禁用片段化 */
|
| 63 |
+
int num_messengers; /* 每个片段的信使数目 S (S << L)。默认 4 */
|
| 64 |
+
int messenger_neighbors; /* 邻近片段信使数(普通 token 可见的邻居片段数)。默认 2 */
|
| 65 |
+
int min_seg_len; /* 尾部封口最小片段长度:序列末尾累积 ≥ 此值时也触发信使生成。
|
| 66 |
+
默认 16。修复短样本(对话数据平均 11 token)信使永远不激活的 bug */
|
| 67 |
+
|
| 68 |
+
/* === Layer 2: 关系强度门控 === */
|
| 69 |
+
int gate_enable; /* 1 = 启用关系预筛选门控,0 = 禁用 */
|
| 70 |
+
int gate_window; /* 门控生效的局部窗口大小(窗口内才做门控) */
|
| 71 |
+
float gate_threshold; /* 概念边界相似度阈值,低于此值判定为边界隔离 */
|
| 72 |
+
float gate_fallback_prob; /* 回退概率(极小,避免硬切断长距离指代)。默认 0.01 */
|
| 73 |
+
int gate_distance_prior; /* 距离先验权重(>0 启用)。距离越远,门控越严 */
|
| 74 |
+
|
| 75 |
+
/* === Layer 3: 异构多头 === */
|
| 76 |
+
int hetero_enable; /* 1 = 启用异构多头,0 = 所有头同等访问 */
|
| 77 |
+
int n_local_heads; /* 局部语法头数量 */
|
| 78 |
+
int n_messenger_heads; /* 指代/因果头数量(访问窗口+信使) */
|
| 79 |
+
/* 剩余头 = HEAD_GLOBAL */
|
| 80 |
+
|
| 81 |
+
/* === Layer 4: 推理 KV-Cache === */
|
| 82 |
+
int cache_messenger; /* 1 = 信使也进入 KV-cache(保证间接长距离通路复用) */
|
| 83 |
+
|
| 84 |
+
/* === 主开关 === */
|
| 85 |
+
int enable; /* 1 = 启用概念感知注意力,0 = 回退到标准 attention_forward */
|
| 86 |
+
} ConceptAttnConfig;
|
| 87 |
+
|
| 88 |
+
/* 默认配置(唯一路线底座:CORE/BINARY/PRUNE + 浮点 + 概念注意力)
|
| 89 |
+
*
|
| 90 |
+
* 这是项目唯一保留的注意力路线——不再是"可选开关"。
|
| 91 |
+
* 四层优化(片段信使 / 关系门控 / 异构多�� / KV 复用)默认全开,
|
| 92 |
+
* 并按审查意见固化:去中心化门控、去中心化信使、窗口对称(128)、自适应阈值。
|
| 93 |
+
* 环境变量 LAL_CONCEPT_ATTN=0 仅作为调试逃生口(强制关,回退标准 attention)。 */
|
| 94 |
+
static inline ConceptAttnConfig concept_attn_default_config(void) {
|
| 95 |
+
ConceptAttnConfig c;
|
| 96 |
+
c.segment_len = 32; /* v21: 32 token 一个片段.
|
| 97 |
+
原 64 在对话数据(avg~11 token, 短样本多)下填不满,
|
| 98 |
+
信使机制空转、概念注意力收不到梯度 (探针实证).
|
| 99 |
+
改 32 后信使候选/质量/统计片段全部恢复(见 seg_len=32 测试).
|
| 100 |
+
注: segment_len 仍可由 LAL_CA_SEG_LEN 环境变量覆盖. */
|
| 101 |
+
c.num_messengers = 4; /* 每片段 4 个信使 */
|
| 102 |
+
c.min_seg_len = 4; /* 尾部封口最小片段长度:≥4 token 的短样本尾部也生成信使,
|
| 103 |
+
避免对话数据(avg 11 token)信使永远空转、概念注意力收不到梯度 */
|
| 104 |
+
c.messenger_neighbors = 2; /* 看 2 个邻近片段的信使 */
|
| 105 |
+
c.gate_enable = 1;
|
| 106 |
+
c.gate_window = 128; /* 窗口对称:概念路径与标准路径一致 (审查 Bug Fix 3: 32→128) */
|
| 107 |
+
/* 阈值语义:固定 0.1 在塌缩表征下失效(审查 Bug 2)。
|
| 108 |
+
* 运行时用门控分数分布的 P25 分位数自适应覆盖(概念注意力前向内维护 g_ca_quantile)。
|
| 109 |
+
* 此处的 0.1 仅作初始/兜底值,不再作为稳定工作点。 */
|
| 110 |
+
c.gate_threshold = 0.1f;
|
| 111 |
+
c.gate_fallback_prob = 0.01f; /* 1% 回退概率,避免硬切断长距离指代 */
|
| 112 |
+
c.gate_distance_prior = 1; /* 启用距离先验(去中心化后作为相对偏差项) */
|
| 113 |
+
c.hetero_enable = 1; /* 异构多头:局部/信使/全局分层算力 */
|
| 114 |
+
c.n_local_heads = -1; /* -1 = 自动:n_head 的一半 */
|
| 115 |
+
c.n_messenger_heads = -1; /* -1 = 自动:n_head 的 1/4 */
|
| 116 |
+
c.cache_messenger = 1; /* 信使进 KV-cache,保证间接长距离通路复用 */
|
| 117 |
+
c.enable = 1; /* 默认开启——概念注意力是基座路线,非可选 */
|
| 118 |
+
return c;
|
| 119 |
+
}
|
| 120 |
+
|
| 121 |
+
/* ========================================================================
|
| 122 |
+
* Messenger Cache (Layer 1 + Layer 4)
|
| 123 |
+
* ========================================================================
|
| 124 |
+
* 每层一个 MessengerCache,存储已聚合的片段信使 K/V。
|
| 125 |
+
* 信使是本片段全部概念与关系状态的压缩载体。
|
| 126 |
+
*
|
| 127 |
+
* 内存布局:
|
| 128 |
+
* messenger_k: [segment_capacity * num_messengers * n_embd] floats
|
| 129 |
+
* messenger_v: 同上
|
| 130 |
+
* segment_filled[segment_capacity]: 该片段的信使是否已生成
|
| 131 |
+
*/
|
| 132 |
+
typedef struct {
|
| 133 |
+
float *messenger_k; /* [segment_capacity * num_messengers * n_embd] */
|
| 134 |
+
float *messenger_v; /* 同上 */
|
| 135 |
+
uint8_t *segment_filled; /* [segment_capacity] 该片段信使是否已生成 */
|
| 136 |
+
int segment_capacity; /* 最大片段数(= n_ctx / segment_len + 1) */
|
| 137 |
+
int num_messengers; /* 每片段信使数 */
|
| 138 |
+
int n_embd; /* 嵌入维度 */
|
| 139 |
+
int n_filled; /* 已生成的片段数 */
|
| 140 |
+
} MessengerCache;
|
| 141 |
+
|
| 142 |
+
/* 分配信使缓存。在 model_load 后调用。 */
|
| 143 |
+
void messenger_cache_alloc(MessengerCache *mc, int segment_capacity,
|
| 144 |
+
int num_messengers, int n_embd);
|
| 145 |
+
|
| 146 |
+
/* 释放信使缓存 */
|
| 147 |
+
void messenger_cache_free(MessengerCache *mc);
|
| 148 |
+
|
| 149 |
+
/* 重置信使缓存(推理新会话开始时调用) */
|
| 150 |
+
void messenger_cache_reset(MessengerCache *mc);
|
| 151 |
+
|
| 152 |
+
/* ========================================================================
|
| 153 |
+
* Layer 1: Segment Messenger Generation
|
| 154 |
+
* ========================================================================
|
| 155 |
+
* 在每个 segment 内部,基于本片段全部 V,聚合生成少量信使向量。
|
| 156 |
+
* 信使是本片段全部概念与关系状态的压缩载体。
|
| 157 |
+
*
|
| 158 |
+
* 聚合策略:对片段内 V 做均匀分桶 + 均值池化
|
| 159 |
+
* - 将片段内 V[0..seg_len-1] 均匀分成 num_messengers 个桶
|
| 160 |
+
* - 每个桶内做均值池化,得到一个信使向量
|
| 161 |
+
* - 信使的 K = 信使的 V(自关联,简化)
|
| 162 |
+
*
|
| 163 |
+
* 数学复杂度:O(seg_len * n_embd),远小于注意力本身
|
| 164 |
+
*
|
| 165 |
+
* 语义意义:远方片段的整体语义,由信使代为表达。
|
| 166 |
+
* 普通token通过信使间接获得远方概念集合的状态,
|
| 167 |
+
* 而不是挨个访问每一个远方概念。
|
| 168 |
+
*
|
| 169 |
+
* 参数:
|
| 170 |
+
* v_seg: [seg_len * n_embd] — 本片段的 V 缓存
|
| 171 |
+
* seg_len: 片段长度
|
| 172 |
+
* out_k: [num_messengers * n_embd] — 输出信使 K
|
| 173 |
+
* out_v: [num_messengers * n_embd] — 输出信使 V
|
| 174 |
+
*/
|
| 175 |
+
void generate_segment_messengers(const float *v_seg, int seg_len, int n_embd,
|
| 176 |
+
int num_messengers,
|
| 177 |
+
float *out_k, float *out_v);
|
| 178 |
+
|
| 179 |
+
/* ========================================================================
|
| 180 |
+
* Layer 2: Concept Boundary Gate (关系强度门控)
|
| 181 |
+
* ========================================================================
|
| 182 |
+
* 给定 token-i(Q侧)、token-j(K侧),利用距离先验 + 粗粒度相似度
|
| 183 |
+
* 快速预判:如果预判两个概念边界隔离,潜在关系极弱,
|
| 184 |
+
* 直接把该位置置 -inf,不参与完整内积计算。
|
| 185 |
+
*
|
| 186 |
+
* 注意:不是简单按位置距离硬截断;是"概念边界是否可能产生关系"的软判断。
|
| 187 |
+
* 避免错误切断长距离指代这种真实强关系。
|
| 188 |
+
*
|
| 189 |
+
* 门控公式(软门控,保留回退通路):
|
| 190 |
+
* sim_coarse = <Q_i, K_j> / (||Q_i|| * ||K_j|| + eps) // 粗粒度余弦相似度
|
| 191 |
+
* dist_prior = exp(-distance / tau) // 距离先验
|
| 192 |
+
* gate_score = sim_coarse * (1 + gate_distance_prior * dist_prior)
|
| 193 |
+
* if gate_score < gate_threshold:
|
| 194 |
+
* 以 (1 - gate_fallback_prob) 概率置 -inf
|
| 195 |
+
* 以 gate_fallback_prob 概率保留(回退通路)
|
| 196 |
+
*
|
| 197 |
+
* 参数:
|
| 198 |
+
* q_i: [head_dim] — 当前 token 的 Q(某头)
|
| 199 |
+
* k_j: [head_dim] — 候选 token 的 K(某头)
|
| 200 |
+
* distance: i 和 j 的位置距离
|
| 201 |
+
* cfg: 概念感知注意力配置
|
| 202 |
+
* 返回:1 = 保留(参与完整 QK 计算),0 = 屏蔽(置 -inf)
|
| 203 |
+
*/
|
| 204 |
+
int concept_boundary_gate(const float *q_i, const float *k_j,
|
| 205 |
+
int head_dim, int distance,
|
| 206 |
+
const ConceptAttnConfig *cfg);
|
| 207 |
+
|
| 208 |
+
/* ========================================================================
|
| 209 |
+
* Layer 3: Heterogeneous Head Access Configuration
|
| 210 |
+
* ========================================================================
|
| 211 |
+
* 不同类型关系本身就有不同的"概念交互范围",不需要统一全序列扫描。
|
| 212 |
+
* - 头A:局部语法关系(主谓宾、修饰):强局部性,适合小窗口。
|
| 213 |
+
* - 头B:指代、实体绑定:偶尔需要长距离跳跃。
|
| 214 |
+
* - 头C:因果、时序关系:中等范围依赖。
|
| 215 |
+
*
|
| 216 |
+
* 配置:根据 n_head 自动分配头的访问域
|
| 217 |
+
* - 前 n_local_heads 个头 = HEAD_LOCAL
|
| 218 |
+
* - 接下来 n_messenger_heads 个头 = HEAD_MESSENGER
|
| 219 |
+
* - 剩余 = HEAD_GLOBAL
|
| 220 |
+
*/
|
| 221 |
+
HeadAccessType get_head_access_type(int head_idx, int n_head,
|
| 222 |
+
const ConceptAttnConfig *cfg);
|
| 223 |
+
|
| 224 |
+
/* 获取某头的实际访问窗口大小 */
|
| 225 |
+
int get_head_window(int head_idx, int n_head, int base_window,
|
| 226 |
+
const ConceptAttnConfig *cfg);
|
| 227 |
+
|
| 228 |
+
/* 某头是否可以访问信使 */
|
| 229 |
+
int head_can_access_messenger(int head_idx, int n_head,
|
| 230 |
+
const ConceptAttnConfig *cfg);
|
| 231 |
+
|
| 232 |
+
/* ========================================================================
|
| 233 |
+
* Layer 4: Concept-Aware Attention Forward (主入口)
|
| 234 |
+
* ========================================================================
|
| 235 |
+
* 概念感知注意力前向传播。整合四层优化:
|
| 236 |
+
*
|
| 237 |
+
* 1. 切成语义片段(segment_len)
|
| 238 |
+
* 2. 片段内部:完整 QKV,充分做片段内概念理解
|
| 239 |
+
* 3. 生成本片段信使:聚合本片段全部概念-关系状态
|
| 240 |
+
* 4. 本片段普通 token:只和【局部窗口 + 本片段信使 + 邻近片段信使】做匹配
|
| 241 |
+
* 5. 关系门控:过滤边界隔离的概念对(Layer 2)
|
| 242 |
+
* 6. 异构多头:不同头不同访问域(Layer 3)
|
| 243 |
+
* 7. KV-Cache:历史 K/V 直接复用,信使也进 cache(Layer 4)
|
| 244 |
+
*
|
| 245 |
+
* 参数:
|
| 246 |
+
* attn_out: [n_embd] — 输出
|
| 247 |
+
* qkv: [3 * n_embd] — Q | K | V 拼接,单 token
|
| 248 |
+
* n_embd, n_head: — 维度与头数
|
| 249 |
+
* seq_pos: — 当前序列位置
|
| 250 |
+
* k_cache, v_cache: — [n_ctx * n_embd] KV 缓存
|
| 251 |
+
* n_ctx: — 上下文长度
|
| 252 |
+
* cfg: — 概念感知注意力配置
|
| 253 |
+
* mc: — 信使缓存(可为 NULL,则不使用信使)
|
| 254 |
+
*
|
| 255 |
+
* 因果性:seq_pos 只关注 0..seq_pos(含)。
|
| 256 |
+
* 多头:head_dim = n_embd / n_head(必须整除)。
|
| 257 |
+
*/
|
| 258 |
+
void attention_forward_concept(float *attn_out, const float *qkv,
|
| 259 |
+
int n_embd, int n_head, int seq_pos,
|
| 260 |
+
float *k_cache, float *v_cache,
|
| 261 |
+
int n_ctx,
|
| 262 |
+
const ConceptAttnConfig *cfg,
|
| 263 |
+
MessengerCache *mc);
|
| 264 |
+
|
| 265 |
+
/* ========================================================================
|
| 266 |
+
* Layer 4: Concept-Aware Attention Backward
|
| 267 |
+
* ========================================================================
|
| 268 |
+
* 概念感知注意力反向传播。计算当前 token 的 Q/K/V 梯度。
|
| 269 |
+
* 缓存的 K/V(位置 0..seq_pos-1)视为常量(与 attention_backward 一致)。
|
| 270 |
+
*
|
| 271 |
+
* 参数:
|
| 272 |
+
* grad_qkv: [3 * n_embd] — 输出梯度(Q|K|V)
|
| 273 |
+
* grad_attn_out: [n_embd] — 来自上层的 attn_out 梯度
|
| 274 |
+
* qkv: [3 * n_embd] — 前向时用的 Q|K|V
|
| 275 |
+
* n_embd, n_head, seq_pos
|
| 276 |
+
* k_cache, v_cache, n_ctx
|
| 277 |
+
* cfg, mc
|
| 278 |
+
*/
|
| 279 |
+
void attention_backward_concept(float *grad_qkv, const float *grad_attn_out,
|
| 280 |
+
const float *qkv, int n_embd, int n_head,
|
| 281 |
+
int seq_pos,
|
| 282 |
+
const float *k_cache, const float *v_cache,
|
| 283 |
+
int n_ctx,
|
| 284 |
+
const ConceptAttnConfig *cfg,
|
| 285 |
+
MessengerCache *mc);
|
| 286 |
+
|
| 287 |
+
/* ========================================================================
|
| 288 |
+
* Transformer Layer Forward Integration
|
| 289 |
+
* ========================================================================
|
| 290 |
+
* trans_layer_forward 的实现包含比例缩放、残差归一化等复杂逻辑,
|
| 291 |
+
* 完整复制易引入 bug。推荐集成方式:
|
| 292 |
+
*
|
| 293 |
+
* 在现有 trans_layer_forward() 中,将 attention_forward 调用替换为:
|
| 294 |
+
* if (g_concept_attn_cfg.enable && g_messenger_caches) {
|
| 295 |
+
* attention_forward_concept(act->attn_out, qkv_ptr,
|
| 296 |
+
* n, cfg->n_head, abs_pos,
|
| 297 |
+
* tl->kv_k, tl->kv_v, cfg->n_ctx,
|
| 298 |
+
* &g_concept_attn_cfg,
|
| 299 |
+
* &g_messenger_caches[layer_idx]);
|
| 300 |
+
* } else {
|
| 301 |
+
* attention_forward(act->attn_out, qkv_ptr, n, cfg->n_head,
|
| 302 |
+
* abs_pos, tl->kv_k, tl->kv_v);
|
| 303 |
+
* }
|
| 304 |
+
*
|
| 305 |
+
* 反向传播同理:将 attention_backward 替换为 attention_backward_concept。
|
| 306 |
+
* 通过 model_set_concept_attn() 在运行时配置,无需修改模型结构。
|
| 307 |
+
*/
|
| 308 |
+
|
| 309 |
+
/* ========================================================================
|
| 310 |
+
* Global Config (便于全局开关,无需改 ModelConfig)
|
| 311 |
+
* ========================================================================
|
| 312 |
+
* 全局概念感知注意力配置。设为 enable=1 即启用。
|
| 313 |
+
* 优先级:ModelConfig 中的 concept_cfg > 全局 g_concept_attn_cfg。
|
| 314 |
+
*/
|
| 315 |
+
extern ConceptAttnConfig g_concept_attn_cfg;
|
| 316 |
+
|
| 317 |
+
/* 全局信使缓存(每层一个,按 layer_idx 索引)。
|
| 318 |
+
* 在 model_load 时分配,model_free 时释放。 */
|
| 319 |
+
extern MessengerCache *g_messenger_caches; /* [n_layer] */
|
| 320 |
+
|
| 321 |
+
/* Model-dependent functions (declared in lal_runtime.h after Model is defined):
|
| 322 |
+
* void model_messenger_caches_alloc(Model *m, const ConceptAttnConfig *cfg);
|
| 323 |
+
* void model_messenger_caches_free(void);
|
| 324 |
+
* void model_messenger_caches_reset(void);
|
| 325 |
+
* void model_set_concept_attn(Model *m, const ConceptAttnConfig *cfg);
|
| 326 |
+
*/
|
| 327 |
+
|
| 328 |
+
#ifdef __cplusplus
|
| 329 |
+
}
|
| 330 |
+
#endif
|
| 331 |
+
|
| 332 |
+
#endif /* LAL_CONCEPT_ATTN_H */
|