项目仓库
Gitee 主仓:https://gitee.com/pei-xiaoguang/kestrel-llm
GitHub 镜像:https://github.com/m13253246268-ship-it/kestrel-llm

相关文档
术语与数据口径 |
性能与基准 |
构建与复现 |
快速上手 |
架构总览

0. 一句话结论

明文实现 softmax 时,第一件事是减去最大值(exp(x − max(x)))——这是防溢出的标准技巧。

但在密文里,max 是最不该出现的操作:它不是多项式,要比较,要分支,而密文里"分支"根本不存在。

我们的替代方案是:用定标常数把 max-sub 换掉。

score 槽 = (q·k / √d)          ← 链内消息已隐式除以 F = 4096²
num      = EXP_Q(score / F)     ← 对全部 s 求 exp(此时还没有 mask)
den      = Σ_{s≤t} num          ← 掩码后用求和
inv      = recip(den)           ← 拟合的倒数多项式
p        = num · inv

关键在第二行:因为 score 在链内已经被隐式缩放了 4096² 倍,所以喂给 exp 的输入天然落在一个非常窄的区间里——窄到可以用一条 6 次多项式精确覆盖。


1. 明文侧的 softmax 为什么不能照搬

标准 softmax:

p_i = exp(s_i − max_j s_j) / Σ_j exp(s_j − max_k s_k)

减 max 的唯一目的是:防止 exp 溢出。exp(50) 就超过 float32 范围了。

但这条路径在密文里每一步都走不通:

需要的能力 明文里有 密文里有吗
求 max 一行代码 ❌ 需要比较 + 分支
逐元素减同一个数 广播 ⚠️ 可以做(同态加法/标量乘),但需要先把 max 求出来
exp 库函数 ❌ 只有多项式

第一行是死结。 所以我们必须换一个思路。


2. 替代方案:用定标把 exp 的输入"框"进窄区间

思路是绕开 max,改用已知的、固定的定标:

  1. 先测:在明文侧统计真实 score 的取值范围;
  2. 再框:让链内消息携带一个隐式的大定标(F = 4096²),于是 exp 的输入被压进一个预先已知的窄区间;
  3. 后拟合:在那个窄区间上拟合 exp 为一条 6 次多项式。

这个方案的代价非常明确:它依赖"score 分布不会跑出统计区间"。一旦真实推理时的 score 越界,拟合多项式在区间外的行为没有任何保证。

这就是"用定标换掉 max-sub"的代价:max-sub 是自适应的(永远安全),定标是预定的(区间外失效)。


3. 拟合的实际做法

规划脚本 _cfold_plan.py 里保留了这条流水线的完整语义(注释原文):

# 折叠语义与 t23_m3p.c C 段一致:
#   score 槽 = (q·k/√d) [链内消息已隐式除 F=4096²]
#   num = EXP_Q(score/F)(引擎对全 s 求 exp,再 causal mask)
#   den = Σ_{s≤t} num;inv = recip(den);p = num·inv

而 _attn_fit.py 负责实际拟合。它的做法值得逐条看:

3.1 exp:窄区间 + 无平移 + 6 次

lo, hi = -0.6, 0.5
z = u            # 直接对 S 拟合(无平移,系数小)
c = np.polyfit(z, np.exp(z), 6)

三个细节:

  • 区间 [-0.6, 0.5]:这是实测出来的 score 范围(脚本先打印 score range [min, max] 再定区间),不是猜的;
  • 不做平移:因为区间本来就窄且靠近 0,直接拟合的系数很小(|c|max 会打印出来核对)——平移反而会引入额外误差;
  • 6 次足够:窄区间上的 exp 极其平滑,deg=5/6 都能到很小的误差。

3.2 为什么"对全 s 求 exp 再 mask"

注意上面注释里那句:引擎对全 s 求 exp,再 causal mask。

顺序是 exp → mask,而不是 mask → exp。原因很实际:密文里没有条件分支,你没法"跳过"某些位置的计算,只能全算完再乘 0/1 掩码。好在"乘 0/1"在同态里是廉价的(标量乘法)。

而 causal mask 的作用是"位置 t 只能看到 ≤ t 的位置",于是在求和那一步:

den = Σ_{s≤t} num

求和范围就自然被掩码限制住了。

3.3 recip:倒数也要拟合

softmax 的分母要取倒数,而 1/x 同样不是多项式。所以倒数也要拟合——而且它的区间是可以推出来的:

# den = sum_s exp(S(h,s)) in [4*exp(lo), 4*exp(hi)]
dlo, dhi = 4*math.exp(lo), 4*math.exp(hi)

因为序列长度是 4(这一层处理 4 个 token),每项落在一个 2 次多项式在区间里的范围,4 项求和后区间就是 [4·exp(lo), 4·exp(hi)]。然后按这个区间标准化、拟合 deg=8。

_cfold_plan.py 里的实现更讲究一点——用 Chebyshev 拟合再转成普通多项式:

u = np.cos(np.pi * (np.arange(n) + 0.5) / n)      # Chebyshev 节点
p = Chebyshev.fit(u, 1.0/den, deg, domain=[-1,1]).convert(kind=Polynomial)

Chebyshev 节点在区间端点附近更密,能压住端点误差——这是数值拟合的标准手法,比等距 linspace 上的普通最小二乘更稳。

3.4 拟合的产物直接是 C 数组

这两个脚本直接打印出可粘贴进 C 源码的系数数组:

const double EXP_Q[7] = { ... };
const double RECIP_Q[9] = { ... };

也就是说:Python 侧负责拟合,C 侧只负责当常量和求值(Horner 法)。这个分工把"数值优化"与"密文运算"彻底隔开——密文侧一行拟合代码都没有(这就是本系列第 2 篇说的"两个世界")。

EXP_Q 的实际值(升幂,取自 _cfold_plan.py):

EXP_Q = [0.0013345331468413693, 0.0084199656748899163, 0.041687167024342969,
         0.16665574107235029, 0.49999829447852773, 1.0000003483277793, 1.0000000191319709]

注意前几个系数很小、最后两个接近 1.0——这是"在 [-0.6, 0.5] 上拟合 exp"应有的形状(exp(x) ≈ 1 + x + x²/2 + …,且区间关于 0 大致对称)。


4. 一个模型结构上的细节:GQA

拟合脚本里 NQ = 16, NKV = 8, HD = 128,而代码里映射 head 的方式是:

hk = h % NKV      # 16 个 query head 共享 8 个 KV head

这是 GQA(分组查询注意力):16 个 query head 对应 8 个 KV head,每个 KV head 被两个 query head 复用。

在密文里这件事的意义是:KV 侧的密文数量只有一半,而"多算的一份 K/V"在明文里是免费的(只是索引复用),在密文里是实打实的密文乘法。所以 GQA 对我们的性能是净收益——这不是我们选的,是模型结构给的便宜。


5. 安全边界(务请读完)

本文所述参数为机制验证级,远低于 HE 参数标准的 128-bit 水平,不得用于保护真实数据。
本文主张的是:密文下 attention 的实现取舍。
本文不主张:安全强度、性能优越性(本文不含与明文实现的性能对比)。

6. 这一篇的未解问题

  1. score 越界后没有兜底。定标方案假定 score 落在统计区间内。如果某个输入让它越界,拟合多项式会给出无意义的值——而且不会有任何报错。这是本文最担心的一处。可能的缓解:在明文侧统计时留更宽的余量,或在多项式上做区间外的钳制(但钳制本身也需要比较)。
  2. 区间余量取了多大,没有记录。[-0.6, 0.5] 是实测范围还是加了余量后的范围?这份数据我们没有归档。
  3. 掩码的实现细节我没有在密文侧逐行确认。本文描述的顺序(exp → mask → 求和)来自规划脚本的注释与仿真代码;C 侧的实际算子序列以 t23_m3p.c 为准,这一处需要读者(以及我们自己)对照源码复核。
  4. 拟合精度到整链精度的传递没有量化。exp 的 max_abs_err 和 recip 的 max_abs_err 脚本里都打印了,但"它贡献了整链 max|err| 的多少"没有做归因。

下一篇我们回到协议层:密文链的两条不变式——lay 与 boot 为什么这样交替,以及 112 与 2100 这两个素数数字是怎么分工的。


原文地址: https://www.cveoy.top/t/topic/qHCY 著作权归作者所有。请勿转载和采集!

免费AI点我,无需注册和登录