<?xml version="1.0" encoding="utf-8"?>
<feed xmlns="http://www.w3.org/2005/Atom">
  <author>
    <name>karmarkar</name>
  </author>
  <generator uri="https://hexo.io/">Hexo</generator>
  <id>https://karmarkar.top/</id>
  <link href="https://karmarkar.top/" rel="alternate"/>
  <link href="https://karmarkar.top/atom.xml" rel="self"/>
  <rights>All rights reserved 2026, karmarkar</rights>
  <subtitle>思考与代码</subtitle>
  <title>karmarkar 的博客</title>
  <updated>2026-09-02T06:42:19.002Z</updated>
  <entry>
    <author>
      <name>karmarkar</name>
    </author>
    <category term="人工智能" scheme="https://karmarkar.top/categories/%E4%BA%BA%E5%B7%A5%E6%99%BA%E8%83%BD/"/>
    <category term="深度学习" scheme="https://karmarkar.top/tags/%E6%B7%B1%E5%BA%A6%E5%AD%A6%E4%B9%A0/"/>
    <category term="Transformer" scheme="https://karmarkar.top/tags/Transformer/"/>
    <category term="机器学习" scheme="https://karmarkar.top/tags/%E6%9C%BA%E5%99%A8%E5%AD%A6%E4%B9%A0/"/>
    <content>
      <![CDATA[<h1>注意力机制</h1><blockquote><p>一切来源于一句朴素的直觉：<strong>&quot;注意力放在哪里，计算就倾向于哪里。&quot;</strong></p></blockquote><p>如果你接触过深度学习，多半听过一个说法：Transformer 时代的基石是注意力机制（Attention Mechanism）。无论是 GPT、BERT，还是机器翻译、图像生成，几乎所有大模型都离不开它。这篇文章不堆砌术语，而是从&quot;为什么要注意力&quot;讲起，一步步拆解它的动机、数学和实现。</p><h2 id="一、为什么要注意力：RNN-的瓶颈">一、为什么要注意力：RNN 的瓶颈</h2><p>在注意力出现之前，处理序列数据的主力是循环神经网络（RNN/LSTM）。RNN 把上一时刻的隐藏状态 <code>h_&#123;t-1&#125;</code> 传入下一时刻，像一条流水线一样逐词处理句子：</p><figure class="highlight 1c"><table><tr><td class="gutter"><pre><span class="line">1</span><br></pre></td><td class="code"><pre><code class="hljs 1c"><span class="hljs-string">&quot;我 喜欢 深度学习&quot;</span>  →  逐词读入 → 最终隐藏状态<br></code></pre></td></tr></table></figure><p>这种方式有两个致命的弱点：</p><ol><li><strong>长距离依赖难</strong>：信息要经过很多步才能从序列开头传到结尾，中间经过非线性变换，早期的信息很容易&quot;遗忘&quot;。</li><li><strong>无法并行</strong>：t 时刻必须等 t-1 时刻算完，GPU 的并行能力被浪费。</li></ol><p>注意力机制正是来解这两个问题的：<strong>它让序列中任意两个位置之间可以直接建立联系，路径长度是 1，且计算天然可并行。</strong></p><h2 id="二、核心思想：Query-Key-Value">二、核心思想：Query / Key / Value</h2><p>注意力机制的直觉，用一个&quot;查资料&quot;的场景就能说清：</p><blockquote><p>你（<strong>Query</strong>，问题）要在一堆资料里找答案。每份资料上有个标题（<strong>Key</strong>，索引），以及正文（<strong>Value</strong>，内容）。你先扫一遍标题，找出与自己问题最相关的几份，然后主要阅读这几份的正文。</p></blockquote><p>映射到神经网络里：</p><ul><li><strong>Q（Query）</strong>：你当前在关注什么（比如要预测的词）。</li><li><strong>K（Key）</strong>：序列里每个位置&quot;是什么&quot;（用来被匹配）。</li><li><strong>V（Value）</strong>：序列里每个位置&quot;携带什么信息&quot;（真正被取用的内容）。</li></ul><p>注意力要做的就一件事：<strong>用 Q 去匹配每一个 K，算出一个&quot;相关度分数&quot;，再按分数加权求和所有的 V。</strong></p><p>分数越高的位置，在输出里占的比重越大——这就是&quot;注意力&quot;。</p><h2 id="三、缩放点积注意力（Scaled-Dot-Product-Attention）">三、缩放点积注意力（Scaled Dot-Product Attention）</h2><p>现代 Transformer 用的是最常用的形式，公式如下：</p><section><eqn><span class="katex-display"><span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML" display="block"><semantics><mrow><mtext>Attention</mtext><mo stretchy="false">(</mo><mi>Q</mi><mo separator="true">,</mo><mi>K</mi><mo separator="true">,</mo><mi>V</mi><mo stretchy="false">)</mo><mo>=</mo><mtext>softmax</mtext><mrow><mo fence="true">(</mo><mfrac><mrow><mi>Q</mi><msup><mi>K</mi><mi>T</mi></msup></mrow><msqrt><msub><mi>d</mi><mi>k</mi></msub></msqrt></mfrac><mo fence="true">)</mo></mrow><mi>V</mi></mrow><annotation encoding="application/x-tex">\text{Attention}(Q, K, V) = \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right) V</annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="katex-base"><span class="katex-strut" style="height:1em;vertical-align:-0.25em;"></span><span class="mord text"><span class="mord">Attention</span></span><span class="mopen">(</span><span class="mord mathnormal">Q</span><span class="mpunct">,</span><span class="mspace" style="margin-right:0.1667em;"></span><span class="mord mathnormal" style="margin-right:0.0715em;">K</span><span class="mpunct">,</span><span class="mspace" style="margin-right:0.1667em;"></span><span class="mord mathnormal" style="margin-right:0.2222em;">V</span><span class="mclose">)</span><span class="mspace" style="margin-right:0.2778em;"></span><span class="mrel">=</span><span class="mspace" style="margin-right:0.2778em;"></span></span><span class="katex-base"><span class="katex-strut" style="height:2.4684em;vertical-align:-0.95em;"></span><span class="mord text"><span class="mord">softmax</span></span><span class="mspace" style="margin-right:0.1667em;"></span><span class="minner"><span class="mopen delimcenter" style="top:0em;"><span class="delimsizing size3">(</span></span><span class="mord"><span class="mopen nulldelimiter"></span><span class="mfrac"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:1.5183em;"><span style="top:-2.2528em;"><span class="pstrut" style="height:3em;"></span><span class="mord"><span class="mord sqrt"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.8572em;"><span class="svg-align" style="top:-3em;"><span class="pstrut" style="height:3em;"></span><span class="mord" style="padding-left:0.833em;"><span class="mord"><span class="mord mathnormal">d</span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.3361em;"><span style="top:-2.55em;margin-left:0em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="katex-sizing reset-size6 size3 mtight"><span class="mord mathnormal mtight" style="margin-right:0.0315em;">k</span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.15em;"><span></span></span></span></span></span></span></span></span><span style="top:-2.8172em;"><span class="pstrut" style="height:3em;"></span><span class="hide-tail" style="min-width:0.853em;height:1.08em;"><svg xmlns="http://www.w3.org/2000/svg" width="400em" height="1.08em" viewBox="0 0 400000 1080" preserveAspectRatio="xMinYMin slice"><path d="M95,702c-2.7,0,-7.17,-2.7,-13.5,-8c-5.8,-5.3,-9.5,-10,-9.5,-14c0,-2,0.3,-3.3,1,-4c1.3,-2.7,23.83,-20.7,67.5,-54c44.2,-33.3,65.8,-50.3,66.5,-51c1.3,-1.3,3,-2,5,-2c4.7,0,8.7,3.3,12,10s173,378,173,378c0.7,0,35.3,-71,104,-213c68.7,-142,137.5,-285,206.5,-429c69,-144,104.5,-217.7,106.5,-221l0 -0c5.3,-9.3,12,-14,20,-14H400000v40H845.2724s-225.272,467,-225.272,467s-235,486,-235,486c-2.7,4.7,-9,7,-19,7c-6,0,-10,-1,-12,-3s-194,-422,-194,-422s-65,47,-65,47zM834 80h400000v40h-400000z"/></svg></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.1828em;"><span></span></span></span></span></span></span></span><span style="top:-3.23em;"><span class="pstrut" style="height:3em;"></span><span class="frac-line" style="border-bottom-width:0.04em;"></span></span><span style="top:-3.677em;"><span class="pstrut" style="height:3em;"></span><span class="mord"><span class="mord mathnormal">Q</span><span class="mord"><span class="mord mathnormal" style="margin-right:0.0715em;">K</span><span class="msupsub"><span class="vlist-t"><span class="vlist-r"><span class="vlist" style="height:0.8413em;"><span style="top:-3.063em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="katex-sizing reset-size6 size3 mtight"><span class="mord mathnormal mtight" style="margin-right:0.1389em;">T</span></span></span></span></span></span></span></span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.93em;"><span></span></span></span></span></span><span class="mclose nulldelimiter"></span></span><span class="mclose delimcenter" style="top:0em;"><span class="delimsizing size3">)</span></span></span><span class="mspace" style="margin-right:0.1667em;"></span><span class="mord mathnormal" style="margin-right:0.2222em;">V</span></span></span></span></span></eqn></section><p>拆成四步，一步步看：</p><h3 id="第-1-步：算相关度">第 1 步：算相关度</h3><p>用 Q 与所有 K 做点积，得到一个分数矩阵：</p><section><eqn><span class="katex-display"><span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML" display="block"><semantics><mrow><mtext>score</mtext><mo stretchy="false">(</mo><mi>q</mi><mo separator="true">,</mo><msub><mi>k</mi><mi>i</mi></msub><mo stretchy="false">)</mo><mo>=</mo><mi>q</mi><mo>⋅</mo><msub><mi>k</mi><mi>i</mi></msub></mrow><annotation encoding="application/x-tex">\text{score}(q, k_i) = q \cdot k_i</annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="katex-base"><span class="katex-strut" style="height:1em;vertical-align:-0.25em;"></span><span class="mord text"><span class="mord">score</span></span><span class="mopen">(</span><span class="mord mathnormal" style="margin-right:0.0359em;">q</span><span class="mpunct">,</span><span class="mspace" style="margin-right:0.1667em;"></span><span class="mord"><span class="mord mathnormal" style="margin-right:0.0315em;">k</span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.3117em;"><span style="top:-2.55em;margin-left:-0.0315em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="katex-sizing reset-size6 size3 mtight"><span class="mord mathnormal mtight">i</span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.15em;"><span></span></span></span></span></span></span><span class="mclose">)</span><span class="mspace" style="margin-right:0.2778em;"></span><span class="mrel">=</span><span class="mspace" style="margin-right:0.2778em;"></span></span><span class="katex-base"><span class="katex-strut" style="height:0.6389em;vertical-align:-0.1944em;"></span><span class="mord mathnormal" style="margin-right:0.0359em;">q</span><span class="mspace" style="margin-right:0.2222em;"></span><span class="mbin">⋅</span><span class="mspace" style="margin-right:0.2222em;"></span></span><span class="katex-base"><span class="katex-strut" style="height:0.8444em;vertical-align:-0.15em;"></span><span class="mord"><span class="mord mathnormal" style="margin-right:0.0315em;">k</span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.3117em;"><span style="top:-2.55em;margin-left:-0.0315em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="katex-sizing reset-size6 size3 mtight"><span class="mord mathnormal mtight">i</span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.15em;"><span></span></span></span></span></span></span></span></span></span></span></eqn></section><p>点积越大，说明两者方向越接近、越相关。</p><h3 id="第-2-步：缩放-1-√d-k">第 2 步：缩放 <code>1/√d_k</code></h3><p>这就是&quot;缩放&quot;二字的由来。为什么要除？</p><blockquote><p>当向量维度 <code>d_k</code> 很大时，点积的值会变得很大，Softmax 的梯度会趋近于 0（处于饱和区），训练变得困难。除以 <code>√d_k</code> 可以让点积的方差保持稳定，让 Softmax 始终落在梯度合理的区间。</p></blockquote><h3 id="第-3-步：Softmax-归一化">第 3 步：Softmax 归一化</h3><p>对每一行做 Softmax，把分数变成&quot;和为 1 的概率分布&quot;：</p><section><eqn><span class="katex-display"><span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML" display="block"><semantics><mrow><msub><mi>α</mi><mi>i</mi></msub><mo>=</mo><mfrac><mrow><mi>exp</mi><mo>⁡</mo><mo stretchy="false">(</mo><msub><mtext>score</mtext><mi>i</mi></msub><mo stretchy="false">)</mo></mrow><mrow><munder><mo>∑</mo><mi>j</mi></munder><mi>exp</mi><mo>⁡</mo><mo stretchy="false">(</mo><msub><mtext>score</mtext><mi>j</mi></msub><mo stretchy="false">)</mo></mrow></mfrac></mrow><annotation encoding="application/x-tex">\alpha_i = \frac{\exp(\text{score}_i)}{\sum_j \exp(\text{score}_j)}</annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="katex-base"><span class="katex-strut" style="height:0.5806em;vertical-align:-0.15em;"></span><span class="mord"><span class="mord mathnormal" style="margin-right:0.0037em;">α</span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.3117em;"><span style="top:-2.55em;margin-left:-0.0037em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="katex-sizing reset-size6 size3 mtight"><span class="mord mathnormal mtight">i</span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.15em;"><span></span></span></span></span></span></span><span class="mspace" style="margin-right:0.2778em;"></span><span class="mrel">=</span><span class="mspace" style="margin-right:0.2778em;"></span></span><span class="katex-base"><span class="katex-strut" style="height:2.5488em;vertical-align:-1.1218em;"></span><span class="mord"><span class="mopen nulldelimiter"></span><span class="mfrac"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:1.427em;"><span style="top:-2.314em;"><span class="pstrut" style="height:3em;"></span><span class="mord"><span class="mop"><span class="mop op-symbol small-op" style="position:relative;top:0em;">∑</span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.162em;"><span style="top:-2.4003em;margin-left:0em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="katex-sizing reset-size6 size3 mtight"><span class="mord mathnormal mtight" style="margin-right:0.0572em;">j</span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.4358em;"><span></span></span></span></span></span></span><span class="mspace" style="margin-right:0.1667em;"></span><span class="mop">exp</span><span class="mopen">(</span><span class="mord"><span class="mord text"><span class="mord">score</span></span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.3117em;"><span style="top:-2.55em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="katex-sizing reset-size6 size3 mtight"><span class="mord mathnormal mtight" style="margin-right:0.0572em;">j</span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.2861em;"><span></span></span></span></span></span></span><span class="mclose">)</span></span></span><span style="top:-3.23em;"><span class="pstrut" style="height:3em;"></span><span class="frac-line" style="border-bottom-width:0.04em;"></span></span><span style="top:-3.677em;"><span class="pstrut" style="height:3em;"></span><span class="mord"><span class="mop">exp</span><span class="mopen">(</span><span class="mord"><span class="mord text"><span class="mord">score</span></span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.3117em;"><span style="top:-2.55em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="katex-sizing reset-size6 size3 mtight"><span class="mord mathnormal mtight">i</span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.15em;"><span></span></span></span></span></span></span><span class="mclose">)</span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:1.1218em;"><span></span></span></span></span></span><span class="mclose nulldelimiter"></span></span></span></span></span></span></eqn></section><p>分数最高的位置拿到最大的权重，这就是&quot;注意力集中&quot;的体现。</p><h3 id="第-4-步：加权求和">第 4 步：加权求和</h3><p>用归一化后的权重对 V 加权求和，得到最终的输出向量：</p><section><eqn><span class="katex-display"><span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML" display="block"><semantics><mrow><mtext>output</mtext><mo>=</mo><munder><mo>∑</mo><mi>i</mi></munder><msub><mi>α</mi><mi>i</mi></msub><msub><mi>v</mi><mi>i</mi></msub></mrow><annotation encoding="application/x-tex">\text{output} = \sum_i \alpha_i v_i</annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="katex-base"><span class="katex-strut" style="height:0.8095em;vertical-align:-0.1944em;"></span><span class="mord text"><span class="mord">output</span></span><span class="mspace" style="margin-right:0.2778em;"></span><span class="mrel">=</span><span class="mspace" style="margin-right:0.2778em;"></span></span><span class="katex-base"><span class="katex-strut" style="height:2.3277em;vertical-align:-1.2777em;"></span><span class="mop op-limits"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:1.05em;"><span style="top:-1.8723em;margin-left:0em;"><span class="pstrut" style="height:3.05em;"></span><span class="katex-sizing reset-size6 size3 mtight"><span class="mord mathnormal mtight">i</span></span></span><span style="top:-3.05em;"><span class="pstrut" style="height:3.05em;"></span><span><span class="mop op-symbol large-op">∑</span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:1.2777em;"><span></span></span></span></span></span><span class="mspace" style="margin-right:0.1667em;"></span><span class="mord"><span class="mord mathnormal" style="margin-right:0.0037em;">α</span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.3117em;"><span style="top:-2.55em;margin-left:-0.0037em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="katex-sizing reset-size6 size3 mtight"><span class="mord mathnormal mtight">i</span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.15em;"><span></span></span></span></span></span></span><span class="mord"><span class="mord mathnormal" style="margin-right:0.0359em;">v</span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.3117em;"><span style="top:-2.55em;margin-left:-0.0359em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="katex-sizing reset-size6 size3 mtight"><span class="mord mathnormal mtight">i</span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.15em;"><span></span></span></span></span></span></span></span></span></span></span></eqn></section><h2 id="四、自注意力（Self-Attention）：关注自己">四、自注意力（Self-Attention）：关注自己</h2><p>如果 Q、K、V 全部来自<strong>同一个序列本身</strong>，就称为自注意力。此时每个词都能看到句子里的其他所有词，并决定自己该&quot;多关注谁&quot;。</p><p>以 &quot;<strong>The animal didn't cross the street because it was too tired</strong>&quot; 为例，编码 &quot;it&quot; 时，自注意力会把较大的权重分给 &quot;animal&quot;，从而知道这里的 &quot;it&quot; 指的是动物，而不是别的名词。</p><p>自注意力让每个 token 的表示都<strong>融合了全局上下文</strong>，这正是&quot;长距离依赖&quot;问题的终极解法——任意两词之间只需一步就能互相看见。</p><h2 id="五、多头注意力（Multi-Head-Attention）">五、多头注意力（Multi-Head Attention）</h2><p>只算一次注意力，相当于所有人只从一个角度去&quot;看&quot;问题。多头注意力则是把 Q、K、V 切分成多个子空间，各自独立算注意力，最后拼接起来：</p><section><eqn><span class="katex-display"><span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML" display="block"><semantics><mrow><mtext>MultiHead</mtext><mo stretchy="false">(</mo><mi>Q</mi><mo separator="true">,</mo><mi>K</mi><mo separator="true">,</mo><mi>V</mi><mo stretchy="false">)</mo><mo>=</mo><mtext>Concat</mtext><mo stretchy="false">(</mo><msub><mtext>head</mtext><mn>1</mn></msub><mo separator="true">,</mo><mo>…</mo><mo separator="true">,</mo><msub><mtext>head</mtext><mi>h</mi></msub><mo stretchy="false">)</mo><msup><mi>W</mi><mi>O</mi></msup></mrow><annotation encoding="application/x-tex">\text{MultiHead}(Q, K, V) = \text{Concat}(\text{head}_1, \ldots, \text{head}_h) W^O</annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="katex-base"><span class="katex-strut" style="height:1em;vertical-align:-0.25em;"></span><span class="mord text"><span class="mord">MultiHead</span></span><span class="mopen">(</span><span class="mord mathnormal">Q</span><span class="mpunct">,</span><span class="mspace" style="margin-right:0.1667em;"></span><span class="mord mathnormal" style="margin-right:0.0715em;">K</span><span class="mpunct">,</span><span class="mspace" style="margin-right:0.1667em;"></span><span class="mord mathnormal" style="margin-right:0.2222em;">V</span><span class="mclose">)</span><span class="mspace" style="margin-right:0.2778em;"></span><span class="mrel">=</span><span class="mspace" style="margin-right:0.2778em;"></span></span><span class="katex-base"><span class="katex-strut" style="height:1.1413em;vertical-align:-0.25em;"></span><span class="mord text"><span class="mord">Concat</span></span><span class="mopen">(</span><span class="mord"><span class="mord text"><span class="mord">head</span></span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.3011em;"><span style="top:-2.55em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="katex-sizing reset-size6 size3 mtight"><span class="mord mtight">1</span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.15em;"><span></span></span></span></span></span></span><span class="mpunct">,</span><span class="mspace" style="margin-right:0.1667em;"></span><span class="minner">…</span><span class="mspace" style="margin-right:0.1667em;"></span><span class="mpunct">,</span><span class="mspace" style="margin-right:0.1667em;"></span><span class="mord"><span class="mord text"><span class="mord">head</span></span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.3361em;"><span style="top:-2.55em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="katex-sizing reset-size6 size3 mtight"><span class="mord mathnormal mtight">h</span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.15em;"><span></span></span></span></span></span></span><span class="mclose">)</span><span class="mord"><span class="mord mathnormal" style="margin-right:0.1389em;">W</span><span class="msupsub"><span class="vlist-t"><span class="vlist-r"><span class="vlist" style="height:0.8913em;"><span style="top:-3.113em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="katex-sizing reset-size6 size3 mtight"><span class="mord mathnormal mtight" style="margin-right:0.0278em;">O</span></span></span></span></span></span></span></span></span></span></span></span></eqn></section><p>其中每个 head：</p><section><eqn><span class="katex-display"><span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML" display="block"><semantics><mrow><msub><mtext>head</mtext><mi>i</mi></msub><mo>=</mo><mtext>Attention</mtext><mo stretchy="false">(</mo><mi>Q</mi><msubsup><mi>W</mi><mi>i</mi><mi>Q</mi></msubsup><mo separator="true">,</mo><mi>K</mi><msubsup><mi>W</mi><mi>i</mi><mi>K</mi></msubsup><mo separator="true">,</mo><mi>V</mi><msubsup><mi>W</mi><mi>i</mi><mi>V</mi></msubsup><mo stretchy="false">)</mo></mrow><annotation encoding="application/x-tex">\text{head}_i = \text{Attention}(QW_i^Q, KW_i^K, VW_i^V)</annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="katex-base"><span class="katex-strut" style="height:0.8444em;vertical-align:-0.15em;"></span><span class="mord"><span class="mord text"><span class="mord">head</span></span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.3117em;"><span style="top:-2.55em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="katex-sizing reset-size6 size3 mtight"><span class="mord mathnormal mtight">i</span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.15em;"><span></span></span></span></span></span></span><span class="mspace" style="margin-right:0.2778em;"></span><span class="mrel">=</span><span class="mspace" style="margin-right:0.2778em;"></span></span><span class="katex-base"><span class="katex-strut" style="height:1.2361em;vertical-align:-0.2769em;"></span><span class="mord text"><span class="mord">Attention</span></span><span class="mopen">(</span><span class="mord mathnormal">Q</span><span class="mord"><span class="mord mathnormal" style="margin-right:0.1389em;">W</span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.9592em;"><span style="top:-2.4231em;margin-left:-0.1389em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="katex-sizing reset-size6 size3 mtight"><span class="mord mathnormal mtight">i</span></span></span><span style="top:-3.1809em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="katex-sizing reset-size6 size3 mtight"><span class="mord mathnormal mtight">Q</span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.2769em;"><span></span></span></span></span></span></span><span class="mpunct">,</span><span class="mspace" style="margin-right:0.1667em;"></span><span class="mord mathnormal" style="margin-right:0.0715em;">K</span><span class="mord"><span class="mord mathnormal" style="margin-right:0.1389em;">W</span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.8913em;"><span style="top:-2.453em;margin-left:-0.1389em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="katex-sizing reset-size6 size3 mtight"><span class="mord mathnormal mtight">i</span></span></span><span style="top:-3.113em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="katex-sizing reset-size6 size3 mtight"><span class="mord mathnormal mtight" style="margin-right:0.0715em;">K</span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.247em;"><span></span></span></span></span></span></span><span class="mpunct">,</span><span class="mspace" style="margin-right:0.1667em;"></span><span class="mord mathnormal" style="margin-right:0.2222em;">V</span><span class="mord"><span class="mord mathnormal" style="margin-right:0.1389em;">W</span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.8913em;"><span style="top:-2.453em;margin-left:-0.1389em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="katex-sizing reset-size6 size3 mtight"><span class="mord mathnormal mtight">i</span></span></span><span style="top:-3.113em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="katex-sizing reset-size6 size3 mtight"><span class="mord mathnormal mtight" style="margin-right:0.2222em;">V</span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.247em;"><span></span></span></span></span></span></span><span class="mclose">)</span></span></span></span></span></eqn></section><blockquote><p>不同 head 能学到不同的关注模式：有的 head 关注语法关系，有的关注相邻词，有的关注长距离指代。<strong>多个视角叠加，表达能力更强。</strong></p></blockquote><h2 id="六、别忘了位置：位置编码">六、别忘了位置：位置编码</h2><p>注意力公式本身是&quot;无序&quot;的——它对序列位置不敏感，把顺序打乱结果不变。但语言是有顺序的，&quot;你打了我&quot;和&quot;我打了你&quot;完全不同。</p><p>所以 Transformer 会给每个 token 加上<strong>位置编码（Positional Encoding）</strong>，把位置信息注入输入向量，让模型能区分&quot;第 1 个词&quot;和&quot;第 5 个词&quot;。</p><h2 id="七、PyTorch-实现：从零写一个多头注意力">七、PyTorch 实现：从零写一个多头注意力</h2><p>理论讲完，来点能跑的。下面是一个完整的 PyTorch 实现，包含单头注意力与多头注意力：</p><figure class="highlight python"><table><tr><td class="gutter"><pre><span class="line">1</span><br><span class="line">2</span><br><span class="line">3</span><br><span class="line">4</span><br><span class="line">5</span><br><span class="line">6</span><br><span class="line">7</span><br><span class="line">8</span><br><span class="line">9</span><br><span class="line">10</span><br><span class="line">11</span><br><span class="line">12</span><br><span class="line">13</span><br><span class="line">14</span><br><span class="line">15</span><br><span class="line">16</span><br><span class="line">17</span><br><span class="line">18</span><br><span class="line">19</span><br><span class="line">20</span><br><span class="line">21</span><br><span class="line">22</span><br><span class="line">23</span><br><span class="line">24</span><br><span class="line">25</span><br><span class="line">26</span><br><span class="line">27</span><br><span class="line">28</span><br><span class="line">29</span><br><span class="line">30</span><br><span class="line">31</span><br><span class="line">32</span><br><span class="line">33</span><br><span class="line">34</span><br><span class="line">35</span><br><span class="line">36</span><br><span class="line">37</span><br><span class="line">38</span><br><span class="line">39</span><br><span class="line">40</span><br><span class="line">41</span><br><span class="line">42</span><br><span class="line">43</span><br><span class="line">44</span><br><span class="line">45</span><br><span class="line">46</span><br><span class="line">47</span><br><span class="line">48</span><br><span class="line">49</span><br><span class="line">50</span><br><span class="line">51</span><br><span class="line">52</span><br><span class="line">53</span><br><span class="line">54</span><br><span class="line">55</span><br><span class="line">56</span><br><span class="line">57</span><br><span class="line">58</span><br><span class="line">59</span><br><span class="line">60</span><br><span class="line">61</span><br><span class="line">62</span><br><span class="line">63</span><br><span class="line">64</span><br><span class="line">65</span><br><span class="line">66</span><br><span class="line">67</span><br></pre></td><td class="code"><pre><code class="hljs python"><span class="hljs-keyword">import</span> torch<br><span class="hljs-keyword">import</span> torch.nn <span class="hljs-keyword">as</span> nn<br><span class="hljs-keyword">import</span> torch.nn.functional <span class="hljs-keyword">as</span> F<br><br><br><span class="hljs-keyword">class</span> <span class="hljs-title class_">ScaledDotProductAttention</span>(nn.Module):<br>    <span class="hljs-string">&quot;&quot;&quot;缩放点积注意力&quot;&quot;&quot;</span><br><br>    <span class="hljs-keyword">def</span> <span class="hljs-title function_">__init__</span>(<span class="hljs-params">self, d_k</span>):<br>        <span class="hljs-built_in">super</span>().__init__()<br>        self.d_k = d_k  <span class="hljs-comment"># 每个 head 的维度，用于缩放</span><br><br>    <span class="hljs-keyword">def</span> <span class="hljs-title function_">forward</span>(<span class="hljs-params">self, q, k, v, mask=<span class="hljs-literal">None</span></span>):<br>        <span class="hljs-comment"># q, k: (B, heads, L, d_k)  v: (B, heads, L, d_v)</span><br>        scores = torch.matmul(q, k.transpose(-<span class="hljs-number">2</span>, -<span class="hljs-number">1</span>))  <span class="hljs-comment"># (B, heads, L, L)</span><br>        scores = scores / (self.d_k ** <span class="hljs-number">0.5</span>)            <span class="hljs-comment"># 缩放，稳定梯度</span><br><br>        <span class="hljs-keyword">if</span> mask <span class="hljs-keyword">is</span> <span class="hljs-keyword">not</span> <span class="hljs-literal">None</span>:<br>            scores = scores.masked_fill(mask == <span class="hljs-number">0</span>, <span class="hljs-built_in">float</span>(<span class="hljs-string">&#x27;-inf&#x27;</span>))<br><br>        weights = F.softmax(scores, dim=-<span class="hljs-number">1</span>)            <span class="hljs-comment"># 归一化成权重</span><br>        out = torch.matmul(weights, v)                 <span class="hljs-comment"># 加权求和</span><br>        <span class="hljs-keyword">return</span> out<br><br><br><span class="hljs-keyword">class</span> <span class="hljs-title class_">MultiHeadAttention</span>(nn.Module):<br>    <span class="hljs-string">&quot;&quot;&quot;多头注意力：切分 d_model，并行算注意力后拼接&quot;&quot;&quot;</span><br><br>    <span class="hljs-keyword">def</span> <span class="hljs-title function_">__init__</span>(<span class="hljs-params">self, d_model, n_heads</span>):<br>        <span class="hljs-built_in">super</span>().__init__()<br>        <span class="hljs-keyword">assert</span> d_model % n_heads == <span class="hljs-number">0</span>, <span class="hljs-string">&quot;d_model 必须能被 n_heads 整除&quot;</span><br>        self.d_model = d_model<br>        self.n_heads = n_heads<br>        self.d_k = d_model // n_heads<br><br>        self.w_q = nn.Linear(d_model, d_model, bias=<span class="hljs-literal">False</span>)<br>        self.w_k = nn.Linear(d_model, d_model, bias=<span class="hljs-literal">False</span>)<br>        self.w_v = nn.Linear(d_model, d_model, bias=<span class="hljs-literal">False</span>)<br>        self.w_o = nn.Linear(d_model, d_model, bias=<span class="hljs-literal">False</span>)<br><br>        self.attention = ScaledDotProductAttention(self.d_k)<br><br>    <span class="hljs-keyword">def</span> <span class="hljs-title function_">_split_heads</span>(<span class="hljs-params">self, x</span>):<br>        <span class="hljs-comment"># (B, L, d_model) -&gt; (B, n_heads, L, d_k)</span><br>        B, L, _ = x.size()<br>        <span class="hljs-keyword">return</span> x.view(B, L, self.n_heads, self.d_k).transpose(<span class="hljs-number">1</span>, <span class="hljs-number">2</span>)<br><br>    <span class="hljs-keyword">def</span> <span class="hljs-title function_">forward</span>(<span class="hljs-params">self, q, k, v</span>):<br>        <span class="hljs-comment"># q/k/v: (B, L, d_model)</span><br>        q, k, v = self.w_q(q), self.w_k(k), self.w_v(v)<br>        q, k, v = self._split_heads(q), self._split_heads(k), self._split_heads(v)<br><br>        out = self.attention(q, k, v)           <span class="hljs-comment"># (B, n_heads, L, d_k)</span><br>        out = out.transpose(<span class="hljs-number">1</span>, <span class="hljs-number">2</span>).contiguous()  <span class="hljs-comment"># 拼回头维度</span><br>        B, L, _, _ = out.size()<br>        out = out.view(B, L, self.d_model)      <span class="hljs-comment"># (B, L, d_model)</span><br>        <span class="hljs-keyword">return</span> self.w_o(out)<br><br><br><span class="hljs-comment"># 快速验证</span><br><span class="hljs-keyword">if</span> __name__ == <span class="hljs-string">&quot;__main__&quot;</span>:<br>    torch.manual_seed(<span class="hljs-number">42</span>)<br>    d_model, n_heads, L = <span class="hljs-number">64</span>, <span class="hljs-number">8</span>, <span class="hljs-number">10</span><br>    x = torch.randn(<span class="hljs-number">2</span>, L, d_model)             <span class="hljs-comment"># 模拟 (batch, 序列长度, 特征维度)</span><br>    mha = MultiHeadAttention(d_model, n_heads)<br>    y = mha(x, x, x)                           <span class="hljs-comment"># 自注意力：Q=K=V=x</span><br>    <span class="hljs-built_in">print</span>(<span class="hljs-string">&quot;输入:&quot;</span>, x.shape, <span class="hljs-string">&quot;-&gt; 输出:&quot;</span>, y.shape)<br></code></pre></td></tr></table></figure><p>运行结果应该是 <code>输入: torch.Size([2, 10, 64]) -&gt; 输出: torch.Size([2, 10, 64])</code>，维度保持不变，方便直接接入 Transformer 的残差连接。</p><p>代码里值得注意的两个点：</p><ul><li><strong><code>scores / d_k**0.5</code></strong> 就是公式里的&quot;缩放&quot;，别小看这一行，少了它深层模型很容易训不动。</li><li><strong><code>mask</code></strong>：解码器里做自回归时需要挡住&quot;未来的词&quot;（把分数设为 <code>-inf</code>，Softmax 后权重为 0），这是 GPT 等模型&quot;只允许看过去&quot;的关键。</li></ul><h2 id="八、注意力机制的应用">八、注意力机制的应用</h2><ul><li><strong>机器翻译</strong>：Attention 最早在 NMT（神经机器翻译）里大放异彩，解决了 RNN 翻译长句时的信息丢失。</li><li><strong>Transformer</strong>：纯注意力架构，&quot;Attention is All You Need&quot;。</li><li><strong>BERT / GPT</strong>：预训练语言模型全部建立在 Transformer 之上，自注意力是其核心组件。</li><li><strong>计算机视觉</strong>：Vision Transformer（ViT）把图像切成 patch 当序列处理；目标检测里的 DETR 也用注意力建模目标关系。</li><li><strong>扩散模型 / 多模态</strong>：Stable Diffusion 用交叉注意力把文本条件注入图像生成过程。</li></ul><h2 id="九、总结">九、总结</h2><p>最后用一张图收个尾，把整个流程串起来：</p><figure class="highlight css"><table><tr><td class="gutter"><pre><span class="line">1</span><br><span class="line">2</span><br><span class="line">3</span><br></pre></td><td class="code"><pre><code class="hljs css"><span class="hljs-selector-tag">Q</span> ─┐<br>K ─┤→  <span class="hljs-selector-tag">Q</span>·Kᵀ →  /√d_k →  softmax →  × V →  输出<br>V ─┘<br></code></pre></td></tr></table></figure><p>注意力机制用一个统一的模式解决了序列建模的两大痛点：</p><ol><li><strong>任意两位置一步可达</strong>，长距离依赖不再是问题；</li><li><strong>计算可并行</strong>，大模型才训得动。</li></ol><p>理解了 Q/K/V、缩放、Softmax、多头这些零件，你再去读任何一篇 Transformer 相关的论文，都会觉得熟悉——因为它们的核心，都是注意力。</p><blockquote><p>相关主题可以继续阅读：Transformer 架构详解、位置编码的直觉、预训练语言模型 BERT/GPT。</p></blockquote>]]>
    </content>
    <id>https://karmarkar.top/posts/141d1667/</id>
    <link href="https://karmarkar.top/posts/141d1667/"/>
    <published>2026-09-01T08:10:00.000Z</published>
    <summary>注意力机制（Attention Mechanism）是深度学习最重要的思想之一。本文从 RNN 的瓶颈讲起，用直觉和公式拆解 Q/K/V、缩放点积、Softmax 与多头注意力，并给出可运行的 PyTorch 实现。</summary>
    <title>注意力机制</title>
    <updated>2026-09-02T06:42:19.002Z</updated>
  </entry>
</feed>
