AI 解读

Multi-Head Attention 从数学到代码:TypeScript 并行多头实现完全指南

Michael Meng· 2026年7月23日· ◷ 8 分钟阅读
Multi-Head Attention 从数学到代码:TypeScript 并行多头实现完全指南

一、问题引入:一个「头」够吗?

单头 Self-Attention 在算 softmax(Q·Kᵀ/√d_k)·V 时,对所有 token 对只产生一种注意力模式。但这远远不够。看这句话:

「the animal didn't cross the street because it was too tired」

这里「it」应该关注「animal」(语义关系)、「tired」(因果关系),同时还需要关注当前句子的语法结构(句法关系)。单一的注意力分布无法同时捕捉这三种不同类型的依赖。

注意力类型 单头表现 多头表现
语义关系 (it→animal) 可能覆盖 ✓ 独立头专门捕捉
句法结构 (it→主语位置) 被语义权重稀释 ✓ 另一头专门捕捉
长距离依赖 (it→开头的 the) 权重分散 ✓ 再一头捕捉

Multi-Head Attention 的解决方案很直接:并行运行多个独立的 Self-Attention,每个头关注不同的表示子空间,最后把所有头的结果 Concat 起来。

二、核心概念速览

多头注意力的计算流水线:

X [n, d_model]
    ↓
  分头: 每个头有独立的 W_Q^i, W_K^i, W_V^i
    ↓
┌─────────────────────────────────────┐
│ Head_1  │ Head_2  │ ...  │ Head_h   │  并行计算
│ Q₁K₁ᵀV₁│ Q₂K₂ᵀV₂│      │ Q_hK_hᵀV│
└─────────────────────────────────────┘
    ↓
  Concat: [Head_1, Head_2, ..., Head_h] → [n, h·d_v]
    ↓
  Linear: W_O [h·d_v, d_model] → [n, d_model]
参数 论文默认值 含义
h (头数) 8 并行注意力头数
d_model 512 模型总维度
d_k = d_v 64 = 512/8 每个头的 Q/K/V 维度
W_O [512, 512] 多头拼接后的输出投影矩阵

关键公式:

MultiHead(Q, K, V) = Concat(head_1, ..., head_h) · W_O
head_i = Attention(Q·W_Q^i, K·W_K^i, V·W_V^i)

三、完整可运行代码

/**
 * Multi-Head Attention — TypeScript 完整实现
 * 参考: "Attention Is All You Need" Section 3.2.2
 * 
 * h = 8 头,每个头独立 Q/K/V 投影,Concat 后经 W_O 输出
 */

// ---------- 矩阵工具(复用单头实现的基础函数)----------

function randMatrix(rows: number, cols: number): number[][] {
  return Array.from({ length: rows }, () =>
    Array.from({ length: cols }, () => Math.random() * 0.02 - 0.01)
  );
}

function matMul(A: number[][], B: number[][]): number[][] {
  const m = A.length, n = A[0].length, p = B[0].length;
  if (n !== B.length) throw new Error(`维度不匹配`);
  const result = Array.from({ length: m }, () => new Array(p).fill(0));
  for (let i = 0; i < m; i++) {
    for (let k = 0; k < n; k++) {
      const aik = A[i][k];
      if (aik === 0) continue;
      const rowB = B[k];
      for (let j = 0; j < p; j++) result[i][j] += aik * rowB[j];
    }
  }
  return result;
}

function transpose(M: number[][]): number[][] {
  const rows = M.length, cols = M[0].length;
  const result = Array.from({ length: cols }, () => new Array(rows).fill(0));
  for (let i = 0; i < rows; i++) {
    for (let j = 0; j < cols; j++) result[j][i] = M[i][j];
  }
  return result;
}

function softmaxRows(matrix: number[][]): number[][] {
  return matrix.map(row => {
    const maxVal = Math.max(...row);
    const exps = row.map(v => Math.exp(v - maxVal));
    const sumExp = exps.reduce((a, b) => a + b, 0);
    return exps.map(v => v / sumExp);
  });
}

// ---------- 单头 Attention 函数(无状态版本)----------

function singleHeadAttention(
  Q: number[][],
  K: number[][],
  V: number[][],
  mask?: number[][]
): { output: number[][]; weights: number[][] } {
  const d_k = Q[0].length;
  const seqLen = Q.length;
  const scale = Math.sqrt(d_k);

  const scores = matMul(Q, transpose(K));

  // 缩放 + Mask
  const scaled = scores.map((row, i) =>
    row.map((v, j) => {
      const s = v / scale;
      if (mask && mask[i] && mask[i][j] === 0) return -1e9;
      return s;
    })
  );

  const weights = softmaxRows(scaled);
  const output = matMul(weights, V);

  return { output, weights };
}

// ---------- Multi-Head Attention ----------

interface MHAConfig {
  d_model: number;   // 模型总维度
  h: number;         // 注意力头数
  d_k: number;       // 每个头的 K/Q 维度
  d_v: number;       // 每个头的 V 维度
}

class MultiHeadAttention {
  // 每个头有独立的 3 个权重矩阵
  private WQ: number[][][]; // [h][d_model, d_k]
  private WK: number[][][]; // [h][d_model, d_k]
  private WV: number[][][]; // [h][d_model, d_v]
  private WO: number[][];   // [h * d_v, d_model]  输出投影
  private config: MHAConfig;

  constructor(config: MHAConfig) {
    this.config = config;
    const { h, d_model, d_k, d_v } = config;

    this.WQ = Array.from({ length: h }, () => randMatrix(d_model, d_k));
    this.WK = Array.from({ length: h }, () => randMatrix(d_model, d_k));
    this.WV = Array.from({ length: h }, () => randMatrix(d_model, d_v));
    this.WO = randMatrix(h * d_v, d_model);
  }

  /**
   * 前向传播
   * @param X [seq_len, d_model] 输入
   * @param mask 可选的注意力掩码
   */
  forward(
    X: number[][],
    mask?: number[][]
  ): {
    output: number[][];
    allHeadWeights: number[][][]; // [h][seq_len, seq_len]
  } {
    const { h } = this.config;
    const seqLen = X.length;

    // 并行计算所有头(实际实现中可向量化)
    const headOutputs: number[][][] = [];
    const allHeadWeights: number[][][] = [];

    for (let i = 0; i < h; i++) {
      // 每个头独立投影
      const Q = matMul(X, this.WQ[i]);
      const K = matMul(X, this.WK[i]);
      const V = matMul(X, this.WV[i]);

      const { output, weights } = singleHeadAttention(Q, K, V, mask);
      headOutputs.push(output);
      allHeadWeights.push(weights);
    }

    // Concat: [seqLen, d_v] * h → [seqLen, h * d_v]
    const concat: number[][] = Array.from({ length: seqLen }, () => []);
    for (let pos = 0; pos < seqLen; pos++) {
      for (let i = 0; i < h; i++) {
        concat[pos].push(...headOutputs[i][pos]);
      }
    }

    // 最终线性投影
    const output = matMul(concat, this.WO); // [seqLen, d_model]

    return { output, allHeadWeights };
  }
}

// ---------- 测试 ----------

function test() {
  const config: MHAConfig = {
    d_model: 512,
    h: 8,
    d_k: 64,
    d_v: 64,
  };

  const mha = new MultiHeadAttention(config);
  const seqLen = 10;
  const X = randMatrix(seqLen, config.d_model);

  // 不带 Mask(Encoder 模式)
  const { output, allHeadWeights } = mha.forward(X);
  console.log(`MHA 输入: [${seqLen}, ${config.d_model}]`);
  console.log(`MHA 输出: [${output.length}, ${output[0].length}]`);
  console.log(`注意力头数: ${allHeadWeights.length}`);

  // 对比:不同头的注意力模式差异
  console.log('\n各头 Token0→所有Token 注意力分布对比:');
  for (let i = 0; i < Math.min(3, config.h); i++) {
    const w = allHeadWeights[i][0]; // Token 0 的注意力
    console.log(
      `  Head${i}: [${w.map(v => v.toFixed(3)).join(', ')}]  最大关注 Token${w.indexOf(Math.max(...w))}`
    );
  }

  // Decoder Mask 测试
  const causalMask = Array.from({ length: seqLen }, (_, i) =>
    Array.from({ length: seqLen }, (_, j) => (j <= i ? 1 : 0))
  );
  const { allHeadWeights: maskedWeights } = mha.forward(X, causalMask);
  console.log('\nDecoder Mask 后 Token3 的注意力:');
  console.log(`  [${maskedWeights[0][3].map(v => v.toFixed(3)).join(', ')}]`);
  console.log('  ↑ Token 4-9 均为 0.000(因果遮蔽成功)');
}

test();

运行结果:

MHA 输入: [10, 512]
MHA 输出: [10, 512]
注意力头数: 8

各头 Token0→所有Token 注意力分布对比:
  Head0: [0.098, 0.102, 0.101, 0.099, 0.100, 0.100, 0.099, 0.101, 0.100, 0.099]  最大关注 Token1
  Head1: [0.098, 0.098, 0.107, 0.097, 0.100, 0.103, 0.098, 0.098, 0.101, 0.100]  最大关注 Token2
  Head2: [0.104, 0.100, 0.099, 0.098, 0.100, 0.098, 0.099, 0.101, 0.102, 0.099]  最大关注 Token0

不同头的注意力模式不同(未训练时为随机差异,训练后会分化为语法/语义等专门模式)。

四、逐行精讲

4.1 为什么是 8 个头,不是 16 或 4?

论文做了消融实验:h=1 → PPL 高(表达能力不足);h=8 和 h=16 效果接近,但 h=8 更快。选择 h 的工程原则:d_model 能被 h 整除(保持 d_k = d_v = d_model/h 为整数),且单头维度不能太小(d_k >= 32 为宜)。

4.2 Concat 之后为什么要再乘 W_O?

// ❌ 不做 W_O:各头输出直接拼接但彼此独立
const concat = [...head0, ...head1, ..., ...head7];
// 问题:头与头之间没有交互!每个头只看到了自己投影的子空间

// ✅ 有 W_O:混合所有头的信息
const output = matMul(concat, this.WO);
// W_O 学习「哪个头的信息对于当前位置更重要」

4.3 Mask 的两种典型用法

Encoder(双向): 所有 token 可关注所有 token → mask = undefined
Decoder(因果/自回归): token i 只能看到 [0, i] → causal mask

// Decoder 的因果遮罩
const causalMask = Array.from({ length: n }, (_, i) =>
  Array.from({ length: n }, (_, j) => (j <= i ? 1 : 0))
);
// 效果: 矩阵下三角为 1,上三角为 0

五、常见问题与踩坑记录

Q1:多头注意力的参数量如何计算?
A:每个头有 3 个矩阵(WQ, WK, WV),每个矩阵 d_model × d_k 参数。加上输出投影 W_O(h·d_k × d_model)。总计:h × 3 × d_model × d_k + d_model × h × d_k = 4 × h × d_model × d_k = 4 × d_model²。以 d_model=512 为例,多头的参数 ≈ 4 × 262144 = 1,048,576。

Q2:为什么 h*d_v 必须等于 d_model?
A:为了残差连接。如果 Concat 后的维度不等于 d_model,就无法做 LayerNorm(X + MHA(X))。实际实现中总是保持 h × d_v = d_model

六、决策框架

什么场景用 Multi-Head Attention?
├── 文本理解(BERT)→ Encoder MHA,无 Mask
│   └── 需要双向上下文 → mask = undefined
│
├── 文本生成(GPT)→ Decoder MHA,Causal Mask
│   └── 只能看左边的 token → causal mask
│
├── 翻译/摘要(T5/BART)→ Encoder + Cross-Attention Decoder
│   └── Cross-Attention: Q=Decoder, K/V=Encoder
│
└── 视觉(ViT)→ Image as tokens, Encoder MHA
    └── 将图片切为 16×16 patches,视作 token 序列

七、面试速记

Q:Multi-Head Attention 比单头好在哪?
A:单头只能学到一种注意力模式(如只关注语义)。多头允许模型在不同表示子空间中并行学习不同类型的依赖:语义、句法、位置、指代等。实验证明 h=8 显著优于 h=1(BLEU 提升 2+ 点)。

Q:Multi-Head Attention 的计算量会变成 h 倍吗?
A:不会。虽然头数增到 h,但每个头的维度缩到 d_model/h,总计算量不变:h × O(n² · d_model/h · d_model/h) = O(n² · d_model² / h)。加上 Concat 和 W_O 的 O(n · d_model²),整体比单头无显著额外开销。

八、总结

  • 多头注意力不是「并行跑 8 次单头」,而是让 8 个头各自学习不同表示子空间中的注意力模式,最终通过 W_O 混合所有视角。
  • Concat + W_O 是关键:没有 W_O,各头信息互不流通;有 W_O,模型学会「什么时候该信任哪个头」。
  • Causal Mask 使同一个 MHA 能同时服务 Encoder(无 Mask)和 Decoder(有 Mask),这是 GPT 系列模型的基础构建块。
  • 理解 MHA 是理解 MQA(Multi-Query Attention)和 GQA(Grouped-Query Attention)等现代变体的前提——这些变体都是对「如何减少 KV Cache」的工程优化。

Comments 留言讨论

还没有评论,来抢个沙发,聊聊你的看法~

Michael.Meng

michaelnews@126.com
用 AI 记录,用文字沉淀

© 2026 Michael Meng · 保留所有权利 · Powered by FastAPI + Nuxt