Skip to content

Commit 6703af5

Browse files
committed
feat(noise-ops-demo): 新增投影点积(attention 分数)模式,附分布原理文档
1 parent f590ddb commit 6703af5

8 files changed

Lines changed: 272 additions & 6 deletions

File tree

Lines changed: 81 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,81 @@
1+
# 注意力分数的分布:从初始化到训练后
2+
3+
注意力头里单个分数的计算:A、B 两个 D 维向量先各自做低维投影
4+
(q = W_Q A,k = W_K B,W_Q、W_K 为 H×D,H 为头维),再在 H 维空间做点积:
5+
6+
$$s = q \cdot k = (W_Q A)^\top (W_K B)$$
7+
8+
问题:s 的概率密度分布由什么决定?本文推导两个阶段的答案——
9+
**初始化时**由 (D, H, σ²) 完全决定(且形状只看 H);
10+
**训练后**由 M = W_QᵀW_K 的奇异值谱决定(H 只是有效秩上限)。
11+
12+
演示工具:`web-tools/noise-ops-demo` 面板二「投影点积(attention 分数)」模式;
13+
基础分布(乘积正态、Bessel-K 点积、卡方)的推导见其 README §1–§2。
14+
15+
## 设定
16+
17+
- A、B 分量 iid,N(0, σ²)(对一次前向传播而言是随机的输入)
18+
- W_Q、W_K 元素 iid,N(0, σ_w²),**初始化后固定**(模型视角:权重给定、输入随机)
19+
20+
## 1. 初始化时:精确分布
21+
22+
q 的每个分量是独立正态的加权和(W 固定,属"一方固定"情形):
23+
24+
$$q_i = \sum_{j=1}^D (W_Q)_{ij} A_j \;\sim\; \mathcal N(0,\; D\sigma_w^2\sigma^2),
25+
\qquad q_1,\dots,q_H \text{ 相互独立}$$
26+
27+
k 同理。于是 s 精确地是 **H 维双方随机点积**
28+
29+
$$s = \sum_{i=1}^H q_i k_i, \qquad
30+
f(s) = \frac{(|s|/2a)^{\nu} K_\nu(|s|/a)}{\sqrt\pi\,\Gamma(H/2)\,a},
31+
\quad \nu = \frac{H-1}{2},\; a = D\sigma_w^2\sigma^2$$
32+
33+
$$\mathbb E[s] = 0, \qquad \mathrm{Var}(s) = H\,(D\sigma_w^2\sigma^2)^2 = H D^2 \sigma_w^4 \sigma^4$$
34+
35+
两个直接推论:
36+
37+
- **形状只看 H**:H=1 是乘积正态(尖峰重尾),H=2 是 Laplace,H 增大由 CLT 趋于正态;
38+
D 只通过分量方差进入尺度。
39+
- **σ_w² = 1/D(标准初始化)时投影保持分量方差**:q_i、k_i 的方差仍是 σ²,
40+
D 从公式中消失,Var(s) = H σ⁴。这就是 softmax 前除以 √H 的来源——
41+
双方随机点积按 1/√(维数)缩放(对比:线性层初始化按 1/D_in,
42+
因为那里是"一方固定"情形,见 noise-ops-demo README §3)。
43+
44+
## 2. 两条缩放规则在 s 上串联
45+
46+
注意力分数恰好把两种基本缩放各用了一次:
47+
48+
| 步骤 | 情形 | 规则 | 效果 |
49+
|---|---|---|---|
50+
| q = W_Q A(线性投影) | 一方固定(W 给定,A 随机) | σ_w² = 1/D | 分量方差不变(σ² → σ²) |
51+
| s = q·k(点积) | 双方随机(q、k 都来自数据) | ÷ √H | 分数方差归一(Hσ⁴ → σ⁴) |
52+
53+
## 3. 训练后:奇异值谱决定一切
54+
55+
把两次投影合并:s = Aᵀ M B,其中 M = W_QᵀW_K 是 D×D 矩阵、rank(M) ≤ H。
56+
对 M 做 SVD(M = Σ σ_i u_i v_iᵀ),由于 u_i、v_i 各自正交,
57+
ξ_i = u_i·A/σ、η_i = v_i·B/σ 仍是独立标准正态,故
58+
59+
$$s = \sum_{i=1}^{r} \sigma_i\,\xi_i\eta_i, \qquad
60+
\mathbb E[s]=0,\quad \mathrm{Var}(s) = \sigma^4 \sum_i \sigma_i^2 = \sigma^4 \|M\|_F^2$$
61+
62+
即 s 是**加权的独立乘积和**,分布形状由奇异值谱 {σ_i} 决定:
63+
64+
- **谱平**(σ_i 全等、r = H):退化为初始化时的 H 维点积形状;
65+
- **谱集中**(有效秩 r_eff = (Σσ_i)²/Σσ_i² ≪ H):趋向少数乘积项之和,
66+
分布变尖峰重尾;r=1 极限就是乘积正态(单个 ξη)。
67+
- 所以"H 决定形状"只在谱平时成立。训练倾向于把谱学得集中
68+
(注意力投影的有效秩通常远小于 H——这也是 LoRA 类低秩微调的前提之一),
69+
因此训练后的分数分布一般比初始化时更尖、尾更重。
70+
71+
初始化时 σ_w²=1/D 的谱是"平"的典型样本:E‖M‖_F² = D²·H·σ_w⁴ = H(σ_w=1/√D),
72+
平均到 H 个方向各摊 ~1,与 §1 的 Var(s) = Hσ⁴ 一致。
73+
74+
## 4. 演示
75+
76+
`web-tools/noise-ops-demo` 面板二「投影点积(attention 分数)」模式:
77+
参数 D(输入维)、H(头维)、σ²(输入分量方差),σ_w² = 1/D 固定;
78+
W_Q、W_K 由种子生成后固定。演示的是 **scaled** 分数 s/√H(scaled dot-product
79+
attention 的做法):理论线为 H 维 Bessel-K 的 √H 缩放版,Var = σ⁴ 恒定。
80+
拖 H:形状从尖峰(H=1)→ Laplace(H=2)→ 正态(大 H),而尺度不动——
81+
直观感受"除以 √H"如何把方差稳住,以及"形状只看 H"的含义。

web-tools/noise-ops-demo/README.md

Lines changed: 6 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -8,7 +8,8 @@
88
1. **按元素运算**(标量视角):同一批样本对 x、y 逐元素做 x、x+y、x∘y、x² 四种运算,
99
四条精确理论密度同图对比,选中项叠加蒙特卡洛直方图。直观看:
1010
相加保持正态;按元素乘变成尖峰重尾的乘积正态;平方把质量压到正半轴、均值漂离 0。
11-
2. **求和:点积与长度**(向量视角):点积(双方随机 / 一方固定 ±1)与 ‖x‖² 三个模式,
11+
2. **求和:点积与长度**(向量视角):点积(双方随机 / 一方固定 ±1)、‖x‖²,
12+
以及投影点积(attention 分数:先各自 H 维投影再点积)四个模式,
1213
直方图 + 精确理论曲线。拖 D 从 1 到 4096 可看乘积正态 → Laplace → 正态的
1314
中心极限过程,以及 ‖x‖² 的长度集中(RMSNorm 的前提);
1415
两种点积模式的理论/样本方差对比即回答"σ²=1/D 能否让点积方差为 1"。
@@ -74,6 +75,10 @@ $$f(z) = \frac{1}{\sigma\sqrt{2\pi z}}\,e^{-z/(2\sigma^2)},\quad z>0$$
7475
| x·y(双方随机) | Bessel-K 型,ν=(D−1)/2 | 0 | D·σx²σy² | 否(D=1 乘积正态、D=2 Laplace;D 大由 CLT 趋近正态) |
7576
| x·v(一方固定,v_i=±1) | N(0, Dσ²) | 0 | Dσ² | ****(精确正态) |
7677
| ‖x‖²(长度平方) | σ²χ²(D) | Dσ² | 2Dσ⁴ | 否(只有正半轴;D 大趋近正态) |
78+
| (W_QA)·(W_KB)/√H(attention 分数,σ_w²=1/D) | H 维双方随机点积 ÷√H | 0 | σ⁴(与 D、H 无关) | 否(形状只看 H) |
79+
80+
投影点积即 attention 分数的初始化分布;训练后取决于 W_QᵀW_K 的奇异值谱,
81+
完整讨论见 `docs/attention-score-distribution.md`
7782

7883
**双方随机** z = Σ x_i y_i:特征函数 φ(t) = (1 + a²t²)^{-D/2}(a = σxσy),反演得
7984

web-tools/noise-ops-demo/css/style.css

Lines changed: 11 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -231,6 +231,17 @@ main {
231231
width: 180px;
232232
}
233233

234+
/* 面板参数行内的控件小组(如头维 H:标签+滑块+输入框),可整体隐藏 */
235+
.panel-controls .control-group {
236+
display: inline-flex;
237+
align-items: center;
238+
gap: 8px;
239+
}
240+
241+
.panel-controls .control-group input[type="range"] {
242+
width: 140px;
243+
}
244+
234245
.controls-hint {
235246
font-size: 12px;
236247
color: #8c959f;

web-tools/noise-ops-demo/index.html

Lines changed: 11 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -72,12 +72,18 @@ <h2>向量求和:点积与长度分布</h2>
7272
<label><input type="radio" name="sumMode" value="dotRandom" checked /> 点积 x·y(双方随机)</label>
7373
<label><input type="radio" name="sumMode" value="dotFixed" /> 点积 x·v(一方固定 ±1)</label>
7474
<label><input type="radio" name="sumMode" value="norm2" /> 长度平方 ‖x‖²</label>
75+
<label><input type="radio" name="sumMode" value="projDot" /> 投影点积(attention 分数)</label>
7576
</div>
7677
</div>
7778
<div class="panel-controls">
7879
<label for="sliderD">维数 D</label>
7980
<input id="sliderD" type="range" min="0" max="12" step="1" value="4" />
8081
<input id="inputD" type="number" min="1" max="8192" step="1" value="16" />
82+
<span class="control-group" id="hControl" style="display: none">
83+
<label for="sliderH">头维 H</label>
84+
<input id="sliderH" type="range" min="0" max="9" step="1" value="6" />
85+
<input id="inputH" type="number" min="1" max="512" step="1" value="64" />
86+
</span>
8187
<label>分量方差 σ²</label>
8288
<div class="preset-group">
8389
<button type="button" data-preset="1">1</button>
@@ -98,6 +104,11 @@ <h2>向量求和:点积与长度分布</h2>
98104
LLM 对应:线性层初始化 W ~ N(0, 1/D<sub>in</sub>) 属一方固定(输入对本次前向传播是给定的),输出方差保持 1;
99105
attention 的 q·k 双方皆随机、分量方差 ≈ 1,点积方差 = D<sub>k</sub>,故 softmax 前除以 √D<sub>k</sub>
100106
‖x‖² ~ σ²χ²(D):均值 Dσ²、相对涨落 √(2/D),高维时长度集中到 σ√D(RMSNorm 能工作的前提)。
107+
<br />
108+
投影点积(attention 分数):s = (W_QA)·(W_KB)/√H,W_Q、W_K 为 H×D 随机初始化矩阵(σ_w²=1/D,由种子生成后固定)。
109+
初始化后投影保持分量方差(q、k 分量方差仍是 σ²),点积方差 = H·σ⁴,除以 √H 后 <strong>Var = σ⁴,与 D、H 都无关</strong>——
110+
拖 H 只剩形状变化(H=1 尖峰、H=2 Laplace、大 H 正态)而尺度稳定。训练后取决于 M = W_QᵀW_K 的奇异值谱(H 只是有效秩上限),
111+
推导见 <code>docs/attention-score-distribution.md</code>
101112
</p>
102113
</div>
103114
</main>

web-tools/noise-ops-demo/js/app.js

Lines changed: 70 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -47,6 +47,9 @@
4747
// 面板二
4848
sliderD: $('sliderD'),
4949
inputD: $('inputD'),
50+
hControl: $('hControl'),
51+
sliderH: $('sliderH'),
52+
inputH: $('inputH'),
5053
inputSigmaSum: $('inputSigmaSum'),
5154
btnSampleSum: $('btnSampleSum'),
5255
sumStats: $('sumStats'),
@@ -60,11 +63,12 @@
6063
useSeed: true,
6164
seed: 42,
6265
elem: { sigma2: 1, N: 200000, op: 'product' },
63-
sum: { D: 16, sigma2preset: '1/D', sigma2: 1 / 16, mode: 'dotRandom' },
66+
sum: { D: 16, H: 64, sigma2preset: '1/D', sigma2: 1 / 16, mode: 'dotRandom' },
6467
};
6568

6669
var pairs = null; // 面板一样本对 {x, y};null = 未采样/已失效
6770
var sumCache = {}; // 面板二缓存:modeId -> { samples, M };参数变更即清空
71+
var projW = null; // projDot 的固定投影矩阵 {D, H, key, wQ, wK};随参数/种子失效
6872
var charts = {};
6973

7074
/** 随机种子(全局共享);salt 区分各条随机流,保证两面板样本无交集 */
@@ -226,8 +230,38 @@
226230
});
227231
}
228232

233+
/**
234+
* projDot 的投影矩阵:σ_w² = 1/D(标准初始化),由种子生成后固定;
235+
* D、H 或种子变化时重建。
236+
*/
237+
function getProjW() {
238+
var D = state.sum.D;
239+
var H = state.sum.H;
240+
var key = D + '|' + H + '|' + currentSeed(0);
241+
if (projW && projW.key === key) return projW;
242+
var g = S.makeRng(currentSeed(41));
243+
var sigmaW = 1 / Math.sqrt(D);
244+
projW = {
245+
key: key,
246+
wQ: S.makeProjection(g, H, D, sigmaW),
247+
wK: S.makeProjection(g, H, D, sigmaW),
248+
};
249+
return projW;
250+
}
251+
229252
/** 采指定模式并存入缓存;M 随 D 自适应(总采样量封顶 ~1.6e7 个分量) */
230253
function sampleSumMode(modeId) {
254+
if (modeId === 'projDot') {
255+
var D = state.sum.D;
256+
var H = state.sum.H;
257+
var sigma = Math.sqrt(state.sum.sigma2);
258+
var w = getProjW();
259+
// 每样本 2HD 次乘加:预算 ~3e8,M 随之自适应
260+
var M = Math.max(500, Math.min(20000, Math.floor(3e8 / (2 * H * D))));
261+
var samples = S.sampleProjDot(S.makeRng(currentSeed(43)), M, D, H, sigma, w.wQ, w.wK);
262+
sumCache.projDot = { samples: samples, M: M };
263+
return;
264+
}
231265
var D = state.sum.D;
232266
var sigma = Math.sqrt(state.sum.sigma2);
233267
var M = Math.max(2000, Math.min(60000, Math.floor(1.6e7 / (2 * D))));
@@ -245,10 +279,11 @@
245279
function renderSum() {
246280
var sigma = Math.sqrt(state.sum.sigma2);
247281
var D = state.sum.D;
282+
var H = state.sum.H; // 仅 projDot 模式使用
248283
var mode = T.SUM_MODES.filter(function (m) {
249284
return m.id === state.sum.mode;
250285
})[0];
251-
var r = mode.range(D, sigma);
286+
var r = mode.range(D, sigma, H);
252287
var lo = r[0];
253288
var hi = r[1];
254289

@@ -263,7 +298,7 @@
263298
animation: false,
264299
data: theoryLine(
265300
function (z) {
266-
return mode.pdf(z, D, sigma);
301+
return mode.pdf(z, D, sigma, H);
267302
},
268303
lo,
269304
hi,
@@ -318,14 +353,15 @@
318353
var text =
319354
'D=' +
320355
D +
356+
(mode.id === 'projDot' ? '、H=' + H : '') +
321357
'、σ²=' +
322358
fmt(state.sum.sigma2) +
323359
':<b>' +
324360
mode.label +
325361
'</b> 理论 均值 ' +
326-
fmt(mode.mean(D, sigma)) +
362+
fmt(mode.mean(D, sigma, H)) +
327363
'、方差 ' +
328-
fmt(mode.variance(D, sigma));
364+
fmt(mode.variance(D, sigma, H));
329365
if (pack) {
330366
var mv = S.sampleMeanVar(pack.samples);
331367
text +=
@@ -342,6 +378,11 @@
342378
el.sumStats.innerHTML = text;
343379
}
344380

381+
/** H 控件只在投影点积模式下显示 */
382+
function refreshHVisibility() {
383+
el.hControl.style.display = state.sum.mode === 'projDot' ? '' : 'none';
384+
}
385+
345386
// ---------- 样本失效 ----------
346387
function invalidateElement() {
347388
pairs = null;
@@ -351,6 +392,7 @@
351392

352393
function invalidateSum() {
353394
sumCache = {};
395+
projW = null;
354396
markNeedSample(el.btnSampleSum, true);
355397
renderSum();
356398
}
@@ -433,9 +475,31 @@
433475
refreshSigma2();
434476
invalidateSum();
435477
});
478+
// H:滑块走 2 的幂,输入框允许任意 1~512
479+
function onHChange(H) {
480+
state.sum.H = H;
481+
invalidateSum();
482+
}
483+
el.sliderH.addEventListener('input', function () {
484+
var H = Math.pow(2, +el.sliderH.value);
485+
el.inputH.value = H;
486+
onHChange(H);
487+
});
488+
el.inputH.addEventListener('change', function () {
489+
var H = Math.round(+el.inputH.value);
490+
if (!isFinite(H)) {
491+
el.inputH.value = state.sum.H;
492+
return;
493+
}
494+
H = Math.max(1, Math.min(512, H));
495+
el.inputH.value = H;
496+
el.sliderH.value = Math.max(0, Math.min(9, Math.round(Math.log2(H))));
497+
onHChange(H);
498+
});
436499
document.querySelectorAll('input[name="sumMode"]').forEach(function (r) {
437500
r.addEventListener('change', function () {
438501
state.sum.mode = r.value;
502+
refreshHVisibility();
439503
renderSum(); // 切换模式不失效:有缓存则显示直方图,无则纯理论线
440504
});
441505
});
@@ -468,6 +532,7 @@
468532
charts.sum = echarts.init($('chartSum'));
469533
bindControls();
470534
refreshSigma2();
535+
refreshHVisibility();
471536
markNeedSample(el.btnSampleElem, true);
472537
markNeedSample(el.btnSampleSum, true);
473538
renderElement();

web-tools/noise-ops-demo/js/sampler.js

Lines changed: 44 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -122,6 +122,48 @@
122122
return out;
123123
}
124124

125+
/**
126+
* 生成 H×D 投影矩阵(attention 的 W_Q / W_K):元素 iid N(0, σw²),按行存储。
127+
* 初始化后固定使用——模型视角:权重给定、输入随机。
128+
*/
129+
function makeProjection(gauss, H, D, sigmaW) {
130+
var W = new Float64Array(H * D);
131+
for (var i = 0; i < W.length; i++) W[i] = sigmaW * gauss();
132+
return W;
133+
}
134+
135+
/**
136+
* 投影点积(attention 分数)采样:s = (W_Q A)·(W_K B)/√H
137+
* (scaled dot-product attention 的做法:÷√H 后 Var = σ⁴,与 D、H 无关)。
138+
* W_Q、W_K 固定(由 makeProjection 事先生成),A、B 每个样本新采(分量 N(0, σ²))。
139+
* 返回 M 个分数样本。
140+
*/
141+
function sampleProjDot(gauss, M, D, H, sigma, wQ, wK) {
142+
var out = new Float64Array(M);
143+
var A = new Float64Array(D);
144+
var B = new Float64Array(D);
145+
for (var m = 0; m < M; m++) {
146+
var i, j;
147+
for (j = 0; j < D; j++) {
148+
A[j] = sigma * gauss();
149+
B[j] = sigma * gauss();
150+
}
151+
var s = 0;
152+
for (i = 0; i < H; i++) {
153+
var off = i * D;
154+
var qi = 0;
155+
var ki = 0;
156+
for (j = 0; j < D; j++) {
157+
qi += wQ[off + j] * A[j];
158+
ki += wK[off + j] * B[j];
159+
}
160+
s += qi * ki;
161+
}
162+
out[m] = s / Math.sqrt(H);
163+
}
164+
return out;
165+
}
166+
125167
/**
126168
* 直方图分箱:把 samples 投到 [lo, hi] 的 nBins 个等宽箱。
127169
* 返回 { centers, density, under, over },density = count/(N·binWidth),
@@ -173,6 +215,8 @@
173215
applyElementOp: applyElementOp,
174216
makeFixedVector: makeFixedVector,
175217
sampleSum: sampleSum,
218+
makeProjection: makeProjection,
219+
sampleProjDot: sampleProjDot,
176220
histogram: histogram,
177221
sampleMeanVar: sampleMeanVar,
178222
};

0 commit comments

Comments
 (0)