<?xml version="1.0" encoding="utf-8"?><feed xmlns="http://www.w3.org/2005/Atom" ><generator uri="https://jekyllrb.com/" version="3.10.0">Jekyll</generator><link href="https://ke-albert.github.io/feed.xml" rel="self" type="application/atom+xml" /><link href="https://ke-albert.github.io/" rel="alternate" type="text/html" /><updated>2026-07-26T14:11:23+08:00</updated><id>https://ke-albert.github.io/feed.xml</id><title type="html">欢迎来到Ke的站点</title><subtitle>Ke Xu&apos;s Personal Blog</subtitle><entry><title type="html">Jekyll Theme Wuk</title><link href="https://ke-albert.github.io/2026/07/24/jekyll-theme-WuK/" rel="alternate" type="text/html" title="Jekyll Theme Wuk" /><published>2026-07-24T00:00:00+08:00</published><updated>2026-07-24T00:00:00+08:00</updated><id>https://ke-albert.github.io/2026/07/24/jekyll-theme-WuK</id><content type="html" xml:base="https://ke-albert.github.io/2026/07/24/jekyll-theme-WuK/"><![CDATA[]]></content><author><name></name></author><summary type="html"><![CDATA[]]></summary></entry><entry><title type="html">transformer</title><link href="https://ke-albert.github.io/2025/11/21/transformer/" rel="alternate" type="text/html" title="transformer" /><published>2025-11-21T21:43:31+08:00</published><updated>2025-11-21T21:43:31+08:00</updated><id>https://ke-albert.github.io/2025/11/21/transformer</id><content type="html" xml:base="https://ke-albert.github.io/2025/11/21/transformer/"><![CDATA[<h1 id="第一种方式">第一种方式</h1>
<p>transformer的输入输出不改变维度大小。q,k,v单独构建，单独计算。</p>
<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><table class="rouge-table"><tbody><tr><td class="rouge-gutter gl"><pre class="lineno">1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
</pre></td><td class="rouge-code"><pre><span class="k">def</span> <span class="nf">forward</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span><span class="n">q</span><span class="p">,</span><span class="n">k</span><span class="p">,</span><span class="n">v</span><span class="p">,</span><span class="n">mask</span><span class="o">=</span><span class="bp">None</span><span class="p">):</span>
        <span class="n">batch</span><span class="p">,</span><span class="n">time</span><span class="p">,</span><span class="n">dimension</span><span class="o">=</span><span class="n">q</span><span class="p">.</span><span class="n">shape</span><span class="c1">#128,32,512,
</span>        <span class="n">n_d</span><span class="o">=</span><span class="bp">self</span><span class="p">.</span><span class="n">emb_dim</span><span class="o">//</span><span class="bp">self</span><span class="p">.</span><span class="n">n_head</span><span class="c1"># 512//8==64
</span>        <span class="n">q</span><span class="p">,</span><span class="n">k</span><span class="p">,</span><span class="n">v</span><span class="o">=</span><span class="bp">self</span><span class="p">.</span><span class="n">w_q</span><span class="p">(</span><span class="n">q</span><span class="p">),</span><span class="bp">self</span><span class="p">.</span><span class="n">w_k</span><span class="p">(</span><span class="n">k</span><span class="p">),</span><span class="bp">self</span><span class="p">.</span><span class="n">w_v</span><span class="p">(</span><span class="n">v</span><span class="p">)</span>
        <span class="n">q</span><span class="o">=</span><span class="n">q</span><span class="p">.</span><span class="n">view</span><span class="p">(</span><span class="n">batch</span><span class="p">,</span><span class="n">time</span><span class="p">,</span><span class="bp">self</span><span class="p">.</span><span class="n">n_head</span><span class="p">,</span><span class="n">n_d</span><span class="p">).</span><span class="n">permute</span><span class="p">(</span><span class="mi">0</span><span class="p">,</span><span class="mi">2</span><span class="p">,</span><span class="mi">1</span><span class="p">,</span><span class="mi">3</span><span class="p">)</span><span class="c1">#128,8,32,64
</span>        <span class="n">k</span><span class="o">=</span><span class="n">k</span><span class="p">.</span><span class="n">view</span><span class="p">(</span><span class="n">batch</span><span class="p">,</span><span class="n">time</span><span class="p">,</span><span class="bp">self</span><span class="p">.</span><span class="n">n_head</span><span class="p">,</span><span class="n">n_d</span><span class="p">).</span><span class="n">permute</span><span class="p">(</span><span class="mi">0</span><span class="p">,</span><span class="mi">2</span><span class="p">,</span><span class="mi">1</span><span class="p">,</span><span class="mi">3</span><span class="p">)</span>
        <span class="n">v</span><span class="o">=</span><span class="n">v</span><span class="p">.</span><span class="n">view</span><span class="p">(</span><span class="n">batch</span><span class="p">,</span><span class="n">time</span><span class="p">,</span><span class="bp">self</span><span class="p">.</span><span class="n">n_head</span><span class="p">,</span><span class="n">n_d</span><span class="p">).</span><span class="n">permute</span><span class="p">(</span><span class="mi">0</span><span class="p">,</span><span class="mi">2</span><span class="p">,</span><span class="mi">1</span><span class="p">,</span><span class="mi">3</span><span class="p">)</span>
        <span class="n">attn</span><span class="o">=</span><span class="n">q</span><span class="o">@</span><span class="n">k</span><span class="p">.</span><span class="n">transpose</span><span class="p">(</span><span class="o">-</span><span class="mi">2</span><span class="p">,</span><span class="o">-</span><span class="mi">1</span><span class="p">)</span><span class="o">/</span><span class="n">math</span><span class="p">.</span><span class="n">sqrt</span><span class="p">(</span><span class="n">n_d</span><span class="p">)</span><span class="c1">#128,8,32,32
</span>        <span class="k">if</span> <span class="n">mask</span> <span class="ow">is</span> <span class="ow">not</span> <span class="bp">None</span><span class="p">:</span>
            <span class="n">mask</span><span class="o">=</span><span class="n">mask</span><span class="p">.</span><span class="n">unsqueeze</span><span class="p">(</span><span class="mi">1</span><span class="p">).</span><span class="n">unsqueeze</span><span class="p">(</span><span class="mi">2</span><span class="p">)</span><span class="c1">#1,1,32,32
</span>            <span class="n">attn</span><span class="o">=</span><span class="n">attn</span><span class="p">.</span><span class="n">masked_fill</span><span class="p">(</span><span class="n">mask</span><span class="p">[:,:,:</span><span class="n">time</span><span class="p">,:</span><span class="n">time</span><span class="p">]</span><span class="o">==</span><span class="mi">0</span><span class="p">,</span><span class="nb">float</span><span class="p">(</span><span class="s">'-inf'</span><span class="p">))</span>
        <span class="n">attn</span><span class="o">=</span><span class="bp">self</span><span class="p">.</span><span class="n">softmax</span><span class="p">(</span><span class="n">attn</span><span class="p">)</span><span class="o">@</span><span class="n">v</span><span class="c1">#128,8,32,64
</span>        <span class="n">attn</span><span class="o">=</span><span class="n">attn</span><span class="p">.</span><span class="n">permute</span><span class="p">(</span><span class="mi">0</span><span class="p">,</span><span class="mi">2</span><span class="p">,</span><span class="mi">1</span><span class="p">,</span><span class="mi">3</span><span class="p">).</span><span class="n">contiguous</span><span class="p">().</span><span class="n">view</span><span class="p">(</span><span class="n">batch</span><span class="p">,</span><span class="n">time</span><span class="p">,</span><span class="n">dimension</span><span class="p">)</span>
        <span class="n">out</span><span class="o">=</span><span class="bp">self</span><span class="p">.</span><span class="n">w_combine</span><span class="p">(</span><span class="n">attn</span><span class="p">)</span>
        <span class="k">return</span> <span class="n">out</span>
</pre></td></tr></tbody></table></code></pre></div></div>
<h1 id="第二种方式">第二种方式</h1>
<p>q,k,v一起都在一个大的矩阵中，一起计算。</p>
<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><table class="rouge-table"><tbody><tr><td class="rouge-gutter gl"><pre class="lineno">1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
</pre></td><td class="rouge-code"><pre><span class="k">def</span> <span class="nf">forward</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span><span class="n">x</span><span class="p">,</span><span class="n">mask</span><span class="o">=</span><span class="bp">None</span><span class="p">):</span>
        <span class="n">B</span><span class="p">,</span><span class="n">T</span><span class="p">,</span><span class="n">C</span><span class="o">=</span><span class="n">x</span><span class="p">.</span><span class="n">shape</span>
        <span class="n">qkv</span><span class="o">=</span><span class="bp">self</span><span class="p">.</span><span class="n">qkv</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
        <span class="n">qkv</span><span class="o">=</span><span class="n">qkv</span><span class="p">.</span><span class="n">reshape</span><span class="p">(</span><span class="n">B</span><span class="p">,</span><span class="n">T</span><span class="p">,</span><span class="mi">3</span><span class="p">,</span><span class="bp">self</span><span class="p">.</span><span class="n">num_heads</span><span class="p">,</span><span class="bp">self</span><span class="p">.</span><span class="n">head_dim</span><span class="p">)</span>
        <span class="n">qkv</span><span class="o">=</span><span class="n">qkv</span><span class="p">.</span><span class="n">permute</span><span class="p">(</span><span class="mi">2</span><span class="p">,</span><span class="mi">0</span><span class="p">,</span><span class="mi">3</span><span class="p">,</span><span class="mi">1</span><span class="p">,</span><span class="mi">4</span><span class="p">)</span><span class="c1">#(3,B,num_heads,T,head_dim)
</span>        <span class="n">q</span><span class="p">,</span><span class="n">k</span><span class="p">,</span><span class="n">v</span><span class="o">=</span><span class="n">qkv</span><span class="p">.</span><span class="n">unbind</span><span class="p">(</span><span class="mi">0</span><span class="p">)</span><span class="c1">#(B,num_heads,T,head_dim)
</span>        <span class="n">attn</span><span class="o">=</span><span class="p">(</span><span class="n">q</span><span class="o">@</span><span class="n">k</span><span class="p">.</span><span class="n">transpose</span><span class="p">(</span><span class="o">-</span><span class="mi">2</span><span class="p">,</span><span class="o">-</span><span class="mi">1</span><span class="p">))</span><span class="o">*</span><span class="bp">self</span><span class="p">.</span><span class="n">scale</span><span class="c1">#(B,num_heads,T,T)
</span>        <span class="k">if</span> <span class="n">mask</span> <span class="ow">is</span> <span class="ow">not</span> <span class="bp">None</span><span class="p">:</span>
            <span class="n">mask</span><span class="o">=</span><span class="n">mask</span><span class="p">.</span><span class="n">unsqueeze</span><span class="p">(</span><span class="mi">1</span><span class="p">).</span><span class="n">unsqueeze</span><span class="p">(</span><span class="mi">2</span><span class="p">)</span>
            <span class="n">attn</span><span class="o">=</span><span class="n">attn</span><span class="p">.</span><span class="n">masked_fill</span><span class="p">(</span><span class="n">mask</span><span class="p">[:,:,:</span><span class="n">T</span><span class="p">,:</span><span class="n">T</span><span class="p">]</span><span class="o">==</span><span class="mi">0</span><span class="p">,</span><span class="nb">float</span><span class="p">(</span><span class="s">'-inf'</span><span class="p">))</span>
        <span class="n">attn</span><span class="o">=</span><span class="bp">self</span><span class="p">.</span><span class="n">softmax</span><span class="p">(</span><span class="n">attn</span><span class="p">)</span>
        <span class="n">out</span><span class="o">=</span><span class="p">(</span><span class="n">attn</span><span class="o">@</span><span class="n">v</span><span class="p">)</span><span class="c1">#(B,num_heads,T,head_dim)
</span>        <span class="n">out</span><span class="o">=</span><span class="n">out</span><span class="p">.</span><span class="n">transpose</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span><span class="mi">2</span><span class="p">)</span><span class="c1">#(B,T,num_heads,head_dim)
</span>        <span class="n">out</span><span class="o">=</span><span class="n">out</span><span class="p">.</span><span class="n">reshape</span><span class="p">(</span><span class="n">B</span><span class="p">,</span><span class="n">T</span><span class="p">,</span><span class="o">-</span><span class="mi">1</span><span class="p">)</span><span class="c1">#(B,T,C)
</span>        <span class="n">out</span><span class="o">=</span><span class="bp">self</span><span class="p">.</span><span class="n">proj</span><span class="p">(</span><span class="n">out</span><span class="p">)</span>
        <span class="k">return</span> <span class="n">out</span>
</pre></td></tr></tbody></table></code></pre></div></div>
<h1 id="第三种方式">第三种方式</h1>
<p>q,k,v单独构建，单独计算。但是变化维度的矩阵分为了正向和反向变化，用来改变q,k,v的维度。</p>
<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><table class="rouge-table"><tbody><tr><td class="rouge-gutter gl"><pre class="lineno">1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
</pre></td><td class="rouge-code"><pre><span class="k">def</span> <span class="nf">forward</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span><span class="n">x</span><span class="p">,</span><span class="n">mask</span><span class="o">=</span><span class="bp">None</span><span class="p">):</span>
        <span class="c1"># queries，keys，values的形状:
</span>        <span class="c1"># (batch_size，查询或者“键－值”对的个数，num_hiddens)
</span>        <span class="c1"># valid_lens　的形状:
</span>        <span class="c1"># (batch_size，)或(batch_size，查询的个数)
</span>        <span class="c1"># 经过变换后，输出的queries，keys，values　的形状:
</span>        <span class="c1"># (batch_size*num_heads，查询或者“键－值”对的个数，
</span>        <span class="c1"># num_hiddens/num_heads)
</span>        <span class="n">queries</span><span class="o">=</span><span class="n">transpose_qkv</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">q</span><span class="p">(</span><span class="n">x</span><span class="p">),</span><span class="bp">self</span><span class="p">.</span><span class="n">num_heads</span><span class="p">)</span>
        <span class="n">keys</span><span class="o">=</span><span class="n">transpose_qkv</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">k</span><span class="p">(</span><span class="n">x</span><span class="p">),</span><span class="bp">self</span><span class="p">.</span><span class="n">num_heads</span><span class="p">)</span>
        <span class="n">values</span><span class="o">=</span><span class="n">transpose_qkv</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">v</span><span class="p">(</span><span class="n">x</span><span class="p">),</span><span class="bp">self</span><span class="p">.</span><span class="n">num_heads</span><span class="p">)</span>

        <span class="n">attn</span><span class="o">=</span><span class="n">queries</span><span class="o">@</span><span class="n">keys</span><span class="p">.</span><span class="n">transpose</span><span class="p">(</span><span class="o">-</span><span class="mi">2</span><span class="p">,</span><span class="o">-</span><span class="mi">1</span><span class="p">)</span>
        <span class="k">if</span> <span class="n">mask</span> <span class="ow">is</span> <span class="ow">not</span> <span class="bp">None</span><span class="p">:</span>
            <span class="n">mask</span><span class="o">=</span><span class="n">mask</span><span class="p">.</span><span class="n">unsqueeze</span><span class="p">(</span><span class="mi">1</span><span class="p">)</span>
            <span class="n">attn</span><span class="o">=</span><span class="n">attn</span><span class="p">.</span><span class="n">masked_fill</span><span class="p">(</span><span class="n">mask</span><span class="p">[:,:</span><span class="n">queries</span><span class="p">.</span><span class="n">shape</span><span class="p">[</span><span class="o">-</span><span class="mi">1</span><span class="p">],:</span><span class="n">queries</span><span class="p">.</span><span class="n">shape</span><span class="p">[</span><span class="o">-</span><span class="mi">1</span><span class="p">]]</span><span class="o">==</span><span class="mi">0</span><span class="p">,</span><span class="nb">float</span><span class="p">(</span><span class="s">'-inf'</span><span class="p">))</span>
        <span class="n">attn</span><span class="o">=</span><span class="bp">self</span><span class="p">.</span><span class="n">softmax</span><span class="p">(</span><span class="n">attn</span><span class="p">)</span>
        <span class="n">y</span><span class="o">=</span><span class="n">attn</span><span class="o">@</span><span class="n">values</span>
        <span class="n">y</span><span class="o">=</span><span class="n">transpose_out</span><span class="p">(</span><span class="n">y</span><span class="p">,</span><span class="bp">self</span><span class="p">.</span><span class="n">num_heads</span><span class="p">)</span>
        <span class="n">y</span><span class="o">=</span><span class="bp">self</span><span class="p">.</span><span class="n">o</span><span class="p">(</span><span class="n">y</span><span class="p">)</span>
        <span class="k">return</span> <span class="n">y</span>
</pre></td></tr></tbody></table></code></pre></div></div>
<h1 id="第四种方式">第四种方式</h1>
<p>几乎和第二种一模一样。</p>
<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><table class="rouge-table"><tbody><tr><td class="rouge-gutter gl"><pre class="lineno">1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
</pre></td><td class="rouge-code"><pre><span class="k">def</span> <span class="nf">forward</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span><span class="n">x</span><span class="p">,</span><span class="n">mask</span><span class="o">=</span><span class="bp">None</span><span class="p">):</span>
        <span class="n">B</span><span class="p">,</span><span class="n">T</span><span class="p">,</span><span class="n">C</span><span class="o">=</span><span class="n">x</span><span class="p">.</span><span class="n">shape</span> <span class="c1"># B:batch_size,T:seq_len,C:embed_dim
</span>        <span class="n">qkv</span><span class="o">=</span><span class="bp">self</span><span class="p">.</span><span class="n">qkv_layer</span><span class="p">(</span><span class="n">x</span><span class="p">)</span><span class="c1">#B,T,3*embed_dim
</span>        <span class="n">q</span><span class="p">,</span><span class="n">k</span><span class="p">,</span><span class="n">v</span><span class="o">=</span><span class="n">qkv</span><span class="p">.</span><span class="n">split</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">embed_dim</span><span class="p">,</span><span class="n">dim</span><span class="o">=-</span><span class="mi">1</span><span class="p">)</span><span class="c1">#B,T,embed_dim
</span>        <span class="n">q</span><span class="o">=</span><span class="n">q</span><span class="p">.</span><span class="n">view</span><span class="p">(</span><span class="n">B</span><span class="p">,</span><span class="n">T</span><span class="p">,</span><span class="bp">self</span><span class="p">.</span><span class="n">num_heads</span><span class="p">,</span><span class="bp">self</span><span class="p">.</span><span class="n">head_dim</span><span class="p">).</span><span class="n">transpose</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span><span class="mi">2</span><span class="p">)</span><span class="c1">#B,num_heads,T,head_dim
</span>        <span class="n">k</span><span class="o">=</span><span class="n">k</span><span class="p">.</span><span class="n">view</span><span class="p">(</span><span class="n">B</span><span class="p">,</span><span class="n">T</span><span class="p">,</span><span class="bp">self</span><span class="p">.</span><span class="n">num_heads</span><span class="p">,</span><span class="bp">self</span><span class="p">.</span><span class="n">head_dim</span><span class="p">).</span><span class="n">transpose</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span><span class="mi">2</span><span class="p">)</span><span class="c1">#B,num_heads,T,head_dim
</span>        <span class="n">v</span><span class="o">=</span><span class="n">v</span><span class="p">.</span><span class="n">view</span><span class="p">(</span><span class="n">B</span><span class="p">,</span><span class="n">T</span><span class="p">,</span><span class="bp">self</span><span class="p">.</span><span class="n">num_heads</span><span class="p">,</span><span class="bp">self</span><span class="p">.</span><span class="n">head_dim</span><span class="p">).</span><span class="n">transpose</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span><span class="mi">2</span><span class="p">)</span><span class="c1">#B,num_heads,T,head_dim
</span>        <span class="n">att</span><span class="o">=</span><span class="n">q</span><span class="o">@</span><span class="n">k</span><span class="p">.</span><span class="n">transpose</span><span class="p">(</span><span class="o">-</span><span class="mi">2</span><span class="p">,</span><span class="o">-</span><span class="mi">1</span><span class="p">)</span><span class="o">*</span><span class="p">(</span><span class="mi">1</span><span class="o">/</span><span class="n">math</span><span class="p">.</span><span class="n">sqrt</span><span class="p">(</span><span class="n">k</span><span class="p">.</span><span class="n">shape</span><span class="p">[</span><span class="o">-</span><span class="mi">1</span><span class="p">]))</span><span class="c1">#B,num_heads,T,T
</span>        <span class="k">if</span> <span class="n">mask</span> <span class="ow">is</span> <span class="ow">not</span> <span class="bp">None</span><span class="p">:</span>
            <span class="n">mask</span><span class="o">=</span><span class="n">mask</span><span class="p">.</span><span class="n">unsqueeze</span><span class="p">(</span><span class="mi">1</span><span class="p">).</span><span class="n">unsqueeze</span><span class="p">(</span><span class="mi">2</span><span class="p">)</span>
            <span class="n">att</span><span class="o">=</span><span class="n">att</span><span class="p">.</span><span class="n">masked_fill</span><span class="p">(</span><span class="n">mask</span><span class="p">[:,:,:</span><span class="n">T</span><span class="p">,:</span><span class="n">T</span><span class="p">]</span><span class="o">==</span><span class="mi">0</span><span class="p">,</span><span class="nb">float</span><span class="p">(</span><span class="s">'-inf'</span><span class="p">))</span>
        <span class="n">att</span><span class="o">=</span><span class="bp">self</span><span class="p">.</span><span class="n">softmax</span><span class="p">(</span><span class="n">att</span><span class="p">)</span>
        <span class="n">y</span><span class="o">=</span><span class="n">att</span><span class="o">@</span><span class="n">v</span><span class="c1">#B,num_heads,T,head_dim
</span>        <span class="n">y</span><span class="o">=</span><span class="n">y</span><span class="p">.</span><span class="n">transpose</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span><span class="mi">2</span><span class="p">).</span><span class="n">contiguous</span><span class="p">().</span><span class="n">view</span><span class="p">(</span><span class="n">B</span><span class="p">,</span><span class="n">T</span><span class="p">,</span><span class="n">C</span><span class="p">)</span>
        <span class="n">y</span><span class="o">=</span><span class="bp">self</span><span class="p">.</span><span class="n">proj_layer</span><span class="p">(</span><span class="n">y</span><span class="p">)</span>
        <span class="k">return</span> <span class="n">y</span>
</pre></td></tr></tbody></table></code></pre></div></div>

<h1 id="vision-transformer">Vision Transformer</h1>]]></content><author><name></name></author><category term="人工智能" /><summary type="html"><![CDATA[第一种方式 transformer的输入输出不改变维度大小。q,k,v单独构建，单独计算。 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 def forward(self,q,k,v,mask=None): batch,time,dimension=q.shape#128,32,512, n_d=self.emb_dim//self.n_head# 512//8==64 q,k,v=self.w_q(q),self.w_k(k),self.w_v(v) q=q.view(batch,time,self.n_head,n_d).permute(0,2,1,3)#128,8,32,64 k=k.view(batch,time,self.n_head,n_d).permute(0,2,1,3) v=v.view(batch,time,self.n_head,n_d).permute(0,2,1,3) attn=q@k.transpose(-2,-1)/math.sqrt(n_d)#128,8,32,32 if mask is not None: mask=mask.unsqueeze(1).unsqueeze(2)#1,1,32,32 attn=attn.masked_fill(mask[:,:,:time,:time]==0,float('-inf')) attn=self.softmax(attn)@v#128,8,32,64 attn=attn.permute(0,2,1,3).contiguous().view(batch,time,dimension) out=self.w_combine(attn) return out 第二种方式 q,k,v一起都在一个大的矩阵中，一起计算。 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 def forward(self,x,mask=None): B,T,C=x.shape qkv=self.qkv(x) qkv=qkv.reshape(B,T,3,self.num_heads,self.head_dim) qkv=qkv.permute(2,0,3,1,4)#(3,B,num_heads,T,head_dim) q,k,v=qkv.unbind(0)#(B,num_heads,T,head_dim) attn=(q@k.transpose(-2,-1))*self.scale#(B,num_heads,T,T) if mask is not None: mask=mask.unsqueeze(1).unsqueeze(2) attn=attn.masked_fill(mask[:,:,:T,:T]==0,float('-inf')) attn=self.softmax(attn) out=(attn@v)#(B,num_heads,T,head_dim) out=out.transpose(1,2)#(B,T,num_heads,head_dim) out=out.reshape(B,T,-1)#(B,T,C) out=self.proj(out) return out 第三种方式 q,k,v单独构建，单独计算。但是变化维度的矩阵分为了正向和反向变化，用来改变q,k,v的维度。 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 def forward(self,x,mask=None): # queries，keys，values的形状: # (batch_size，查询或者“键－值”对的个数，num_hiddens) # valid_lens　的形状: # (batch_size，)或(batch_size，查询的个数) # 经过变换后，输出的queries，keys，values　的形状: # (batch_size*num_heads，查询或者“键－值”对的个数， # num_hiddens/num_heads) queries=transpose_qkv(self.q(x),self.num_heads) keys=transpose_qkv(self.k(x),self.num_heads) values=transpose_qkv(self.v(x),self.num_heads) attn=queries@keys.transpose(-2,-1) if mask is not None: mask=mask.unsqueeze(1) attn=attn.masked_fill(mask[:,:queries.shape[-1],:queries.shape[-1]]==0,float('-inf')) attn=self.softmax(attn) y=attn@values y=transpose_out(y,self.num_heads) y=self.o(y) return y 第四种方式 几乎和第二种一模一样。 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 def forward(self,x,mask=None): B,T,C=x.shape # B:batch_size,T:seq_len,C:embed_dim qkv=self.qkv_layer(x)#B,T,3*embed_dim q,k,v=qkv.split(self.embed_dim,dim=-1)#B,T,embed_dim q=q.view(B,T,self.num_heads,self.head_dim).transpose(1,2)#B,num_heads,T,head_dim k=k.view(B,T,self.num_heads,self.head_dim).transpose(1,2)#B,num_heads,T,head_dim v=v.view(B,T,self.num_heads,self.head_dim).transpose(1,2)#B,num_heads,T,head_dim att=q@k.transpose(-2,-1)*(1/math.sqrt(k.shape[-1]))#B,num_heads,T,T if mask is not None: mask=mask.unsqueeze(1).unsqueeze(2) att=att.masked_fill(mask[:,:,:T,:T]==0,float('-inf')) att=self.softmax(att) y=att@v#B,num_heads,T,head_dim y=y.transpose(1,2).contiguous().view(B,T,C) y=self.proj_layer(y) return y Vision Transformer]]></summary></entry><entry><title type="html">zero-to-hero</title><link href="https://ke-albert.github.io/2025/10/19/zero-to-hero/" rel="alternate" type="text/html" title="zero-to-hero" /><published>2025-10-19T15:43:31+08:00</published><updated>2025-10-19T15:43:31+08:00</updated><id>https://ke-albert.github.io/2025/10/19/zero-to-hero</id><content type="html" xml:base="https://ke-albert.github.io/2025/10/19/zero-to-hero/"><![CDATA[<div class="post-toc" id="post-toc">
  <div class="post-toc-title">目录</div>
</div>

<p>这篇文章记录了从手动实现自动微分到构建 GPT 训练逻辑的完整学习路径，重点包括微分、二元模型、N-gram、BatchNorm、Transformer 与 Tokenization 的核心思想与代码实现。</p>

<h1 id="part-1-micrograd">Part 1-Micrograd</h1>
<p>要自动计算梯度，需要一个类来实现，这个类能够表示标量和张量，需要计算变量的梯度时，能够调用这个类（实例）的相关计算梯度的方法。可以简单地创建一个Value类:</p>
<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><table class="rouge-table"><tbody><tr><td class="rouge-gutter gl"><pre class="lineno">1
2
3
4
5
6
7
</pre></td><td class="rouge-code"><pre><span class="k">class</span> <span class="nc">Value</span><span class="p">:</span>
    <span class="k">def</span> <span class="nf">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">data</span><span class="p">,</span> <span class="n">_children</span><span class="o">=</span><span class="p">(),</span> <span class="n">_op</span><span class="o">=</span><span class="s">''</span><span class="p">):</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">data</span> <span class="o">=</span> <span class="n">data</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">grad</span> <span class="o">=</span> <span class="mi">0</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">_backward</span> <span class="o">=</span> <span class="k">lambda</span><span class="p">:</span> <span class="bp">None</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">_prev</span> <span class="o">=</span> <span class="nb">set</span><span class="p">(</span><span class="n">_children</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">_op</span> <span class="o">=</span> <span class="n">_op</span>
</pre></td></tr></tbody></table></code></pre></div></div>

<p>这样，梯度和求导函数都包括在了这个类里面，还记录了它依赖于哪些变量即由哪些变量和操作计算得来。以tanh为例：</p>
<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><table class="rouge-table"><tbody><tr><td class="rouge-gutter gl"><pre class="lineno">1
2
3
4
5
6
7
8
9
10
11
12
</pre></td><td class="rouge-code"><pre><span class="k">class</span> <span class="nc">Value</span><span class="p">:</span>

    <span class="p">...</span>

    <span class="k">def</span> <span class="nf">tanh</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
        <span class="n">x</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">data</span>
        <span class="n">t</span> <span class="o">=</span> <span class="p">(</span><span class="n">math</span><span class="p">.</span><span class="n">exp</span><span class="p">(</span><span class="mi">2</span><span class="o">*</span><span class="n">x</span><span class="p">)</span> <span class="o">-</span> <span class="mi">1</span><span class="p">)</span><span class="o">/</span><span class="p">(</span><span class="n">math</span><span class="p">.</span><span class="n">exp</span><span class="p">(</span><span class="mi">2</span><span class="o">*</span><span class="n">x</span><span class="p">)</span> <span class="o">+</span> <span class="mi">1</span><span class="p">)</span>
        <span class="n">out</span> <span class="o">=</span> <span class="n">Value</span><span class="p">(</span><span class="n">t</span><span class="p">,</span> <span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="p">),</span> <span class="s">'tanh'</span><span class="p">)</span>
        <span class="k">def</span> <span class="nf">_backward</span><span class="p">():</span>
            <span class="bp">self</span><span class="p">.</span><span class="n">grad</span> <span class="o">+=</span> <span class="p">(</span><span class="mi">1</span> <span class="o">-</span> <span class="n">t</span><span class="o">**</span><span class="mi">2</span><span class="p">)</span> <span class="o">*</span> <span class="n">out</span><span class="p">.</span><span class="n">grad</span>
        <span class="n">out</span><span class="p">.</span><span class="n">_backward</span> <span class="o">=</span> <span class="n">_backward</span>
    <span class="k">return</span> <span class="n">out</span>
</pre></td></tr></tbody></table></code></pre></div></div>
<p>在这里，out为计算的tanh值，它是一个Value对象并最终返回out，例如a=Value(1),b=a.tanh()，在计算tanh时，也会创建一个<code class="language-plaintext highlighter-rouge">_backward()</code>函数，它被赋值给了新创建的out变量，它记录了当前变量需要计算梯度时的公式，即这样的逻辑<code class="language-plaintext highlighter-rouge">还是以a,b为例，当需要从b反向传播计算a的梯度时，将b的梯度设为1，调用b的反向传播函数，它是用以计算a的梯度。反向传播函数不是同级对等的关系，即a的梯度需要调用b的反向传播函数计算得到，而不是a的反向传播函数，如果a的反向传播函数不为None，那么它则是计算在a之前的变量的梯度，而不是a本身的梯度，以y=x**2为例，要计算x关于y的导数，是通过y对x求导，而不是对x本身求导，y对其自身的导数为1</code>，这样当需要调用反向传播函数时，就可以计算得到当前变量的梯度值，这里的梯度是<code class="language-plaintext highlighter-rouge">+=</code>而不是直接赋值，这是因为同一变量可能会在多个不同的地方使用到，正确的计算方式就是将这些不同地方的但是是同一变量的梯度进行累加。这也解释了在<code class="language-plaintext highlighter-rouge">PyTorch</code>中，每一轮训练开始在进行反向传播之前需要将变量的梯度清0,为的就是不让上一轮的梯度继续与本轮的梯度累积。</p>

<p>当需要进行反向传播时，我们需要一个搜索算法来将所有的与最终作为反向传播开始的变量相关的变量都找出来，例如a=Value(1),b=a.tanh(),c=b**2，当需要从c反向传播计算a的梯度时，需要先将c的梯度设为1，然后调用c的反向传播函数，它会计算b的梯度，而a的梯度需要调用b的反向传播函数计算得到，所以需要一个搜索算法来将所有的与c相关的变量都找出来，即a,b,c。</p>
<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><table class="rouge-table"><tbody><tr><td class="rouge-gutter gl"><pre class="lineno">1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
</pre></td><td class="rouge-code"><pre><span class="k">class</span> <span class="nc">Value</span><span class="p">:</span>
     <span class="p">...</span>
     <span class="k">def</span> <span class="nf">backward</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
        <span class="n">topo</span><span class="o">=</span><span class="p">[]</span>
        <span class="n">visited</span><span class="o">=</span><span class="nb">set</span><span class="p">()</span>
        <span class="k">def</span> <span class="nf">build_topo</span><span class="p">(</span><span class="n">v</span><span class="p">):</span>
            <span class="k">if</span> <span class="n">v</span> <span class="ow">not</span> <span class="ow">in</span> <span class="n">visited</span><span class="p">:</span>
                <span class="n">visited</span><span class="p">.</span><span class="n">append</span><span class="p">(</span><span class="n">v</span><span class="p">)</span>
                <span class="k">for</span> <span class="n">child</span> <span class="ow">in</span> <span class="n">v</span><span class="p">.</span><span class="n">_prev</span><span class="p">:</span>
                    <span class="n">build_topo</span><span class="p">(</span><span class="n">child</span><span class="p">)</span>
                <span class="n">topo</span><span class="p">.</span><span class="n">append</span><span class="p">(</span><span class="n">v</span><span class="p">)</span>
        <span class="n">build_topo</span><span class="p">(</span><span class="bp">self</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">grad</span><span class="o">=</span><span class="mf">1.0</span>
        <span class="k">for</span> <span class="n">v</span> <span class="ow">in</span> <span class="nb">reversed</span><span class="p">(</span><span class="n">topo</span><span class="p">):</span>
            <span class="n">v</span><span class="p">.</span><span class="n">_backward</span><span class="p">()</span>
</pre></td></tr></tbody></table></code></pre></div></div>
<p>这是一个拓扑排序的过程，它将所有的与最终作为反向传播开始的变量相关的变量都找出来，然后从后往前调用每个变量的反向传播函数，这样就可以计算得到所有变量的梯度值。使用的是深度优先算法，这里用了两个变量来记录拓扑排序和已访问变量节点，理论上来说可以只使用<code class="language-plaintext highlighter-rouge">topo</code>一个变量，这里用了<code class="language-plaintext highlighter-rouge">visited|set</code>是因为集合查找比列表查找更快，列表的查找时间复杂度是O(log(n)),而集合则是O(1)。
有了能够进行自动微分反向传播的类，我们就可以按照Neuron-&gt;Layer-&gt;MLP的顺序构建一个简单的多层感知机，并进行训练。它不支持矩阵并行计算这样的复杂操作，底层实现逻辑还是通过循环来获取单一的元素进行运算的。</p>

<h1 id="part-2-bigram">Part 2-Bigram</h1>
<p>从统计模型角度来看二元语法模型，字符级别的预测例如句子<code class="language-plaintext highlighter-rouge">我喜欢你</code>,我们根据单个字符进行预测，首先在句子的首位添加起止符，比如可以使用<code class="language-plaintext highlighter-rouge">.我喜欢你.</code>这样的表达，<code class="language-plaintext highlighter-rouge">.</code>表示句子的开始和结束，起止符号可以相同，也可以采用不同的表示，这里采样相同的表示。将字符转换成数字显然更加符合计算机的习惯，也是更方便计算。所以对于’abc…z’这样的字符，再加上<code class="language-plaintext highlighter-rouge">.</code>这样的起止符，在这个简单的模型中一共有27个独特的字符，设计两个字典方便从字符到数字之间的互相转换。可以统计在使用的数据集中，二元字符对出现的频率，这应该是一个[27,27]的二维矩阵，因为每个字符都可以作为二元模型中的第一个字符，也可以作为二元模型中的第二个字符，总之，我们可以经过统计得到这样的以每个字符开头的二元对的频率，并将它友好地显示出来。从可视化中可以明显看出，一些字符对的出现频率很高，一些则一次都没有，这即揭示了一些统计规律（也可能是当前数据集不够具有代表性），但不管怎样，我们可以根据当前的数据集得到符合当前数据集的统计规律，比如<code class="language-plaintext highlighter-rouge">wq</code>字符对在当前数据集中出现的频率为0。接着，我们可以独立地对每一行求和，并计算各自的概率即softmax运算。这样就得到了基于当前字符作为第一个字符，其二元模型中第二个字符出现的各个近似概率（总之如果数据集足够具有代表性，也即数据集足够大，大数定律告诉我们当然可以近似地认为这就是这个字符对的概率）。这就是简单的二元语法的统计模型，接下来我们就可以进行采样了，生成一系列<code class="language-plaintext highlighter-rouge">预测</code>的合理的人名。</p>

<p>简单的统计模型，换个角度我们也可以使用神经网络的方式来实现，因为当到达三元、四元、n元模型时，简单的统计模型就显得捉襟见肘了，单单想一想各种组合的可能性，就是呈指数级的爆炸增长。而神经网络可以在简单和复杂中都能有效地进行表示（建模）和计算。在概率与数理统计中，我们学习过最大似然估计，简而言之，最大似然估计构建了一个这样的模型，这个模型符合使得当前数据集中n元模型出现的概率达到最大–即利用已知的样本结果信息，反推最具有可能（最大概率）导致这些样本结果出现的模型参数值，模型的复杂问题一切都包含在了什么模型最能导致当前的样本结果了。<code class="language-plaintext highlighter-rouge">∏𝑖=1𝑛𝑝𝜃^(𝑥𝑖)，这就是最大似然函数。对于连续型随机变量，有相同的结论。</code>。深度学习中我们一般都是最小化损失，换个思路最大化似然函数-&gt;最小化负的最大似然函数-&gt;最小化数据集中的所有数据的平均负的最大似然函数，按照惯例为了计算方便，我们将其对数化，这样乘法就变成了加法<code class="language-plaintext highlighter-rouge">log(a*b*c)=loga+logb+logc</code>，于是这样的逻辑思路就是<code class="language-plaintext highlighter-rouge">利用最大似然函数的思想我们构建了一个对模型损失的评估，通过反向传播算法我们不断根据当前的损失去更新模型的参数。</code>不过当遇到从没有出现过的字符对时，这个log值就为<code class="language-plaintext highlighter-rouge">-inf</code>，可以给整体各自加上一个较小的值，比如1，这样最小频率为1，这个方式称为模型平滑。为了训练神经网络，需要构建训练样本标签对，例如<code class="language-plaintext highlighter-rouge">.emma.</code>这样的单词，划分成训练样本对为<code class="language-plaintext highlighter-rouge">xtr=[.,e,m,m]</code>和<code class="language-plaintext highlighter-rouge">ytr=[e,m,m,a,.]</code>，在输入网络前还需要将其转换成数字的形式<code class="language-plaintext highlighter-rouge">[0,5,13,13,1]</code>和<code class="language-plaintext highlighter-rouge">[5,13,13,1,0]</code>。转换成数字的形式之后，可以看到时间序列xtr中当前的词去预测下一个词，在下一个时间步中被预测的下一个词又作为当前的词继续去预测它的下一个词，因为要使用神经网络，y=X@W+b这样的形式，可以把xtr-ytr变为二维即<code class="language-plaintext highlighter-rouge">shape=[time_step,1]</code>的形式，但是这样矩阵相乘时，就要求W的形状是<code class="language-plaintext highlighter-rouge">shape=[1,hidden_representation]</code>的大小，即<code class="language-plaintext highlighter-rouge">[[0],[5],[13],[13],[1]]@[[n1,...,n_h]]</code>，简单的矩阵乘法知识就可以知道，它的结果仅仅是<code class="language-plaintext highlighter-rouge">C=A@B</code>中，C的每一行的结果都是A的每一行乘以B的每一列。而使用<code class="language-plaintext highlighter-rouge">one_hot</code>编码，将<code class="language-plaintext highlighter-rouge">x_tr-y_tr</code>根据其字典大小的长度进行编码，可以让<code class="language-plaintext highlighter-rouge">x_tr-y_tr</code>的形状变成<code class="language-plaintext highlighter-rouge">shape=[time_step,27]</code>，这样W的形状就可以是<code class="language-plaintext highlighter-rouge">shape=[27,hidden_representation]</code>，可以进行更复杂的线性组合，获得更复杂的表示，独热码只是encodeing的其中一种最简单的方式，本质就是给当前的字符一个在高维空间上的表示方法，还可以通过一些词袋模型计算更复杂的编码表示方式比如<code class="language-plaintext highlighter-rouge">word2world</code>。在这里，简单将隐藏表示维度设为27，那么<code class="language-plaintext highlighter-rouge">logits=xenc@W</code>的结果被赋予了<code class="language-plaintext highlighter-rouge">logits</code>的称呼，这也即与统计模型中的出现次数对应一个级别，但是这里的<code class="language-plaintext highlighter-rouge">次数</code>有正有负，因为W是随机初始化的。对<code class="language-plaintext highlighter-rouge">counts=logits.exp()</code>将其全部转为正数，然后计算其softmax结果，就得到了统计模型中每一行相同的频率（概率）表示形式，有了概率表示就可以计算最大似然函数，在未训练模型时，我们还可以估计平均最大似然函数正常的取值范围，因为在没有训练之前，没有任何理由可以认为什么组合的出现概率最高，所以它们出现的概率是相等的，平等看待每一个组合，这样就计算出一个最大似然函数的值，可以用来评估初始化权重矩阵是否合理。从负的最大似然函数开始进行反向传播，使用梯度下降算法训练网络。最后同样地可以得到与统计模型相似地结果。</p>

<p>W权重矩阵和二元模型概率分布及其正则化。W的形状是<code class="language-plaintext highlighter-rouge">shape=[27,27]</code>，而我们使用<code class="language-plaintext highlighter-rouge">one_hot</code>编码，根据矩阵相乘的结果，每一行即每一个时间步的输入，都被映射到了一个27维的向量空间中，这个向量空间中的每个维度，都对应了一个字符，而这个字符的出现概率，就是这个向量空间中这个维度的数值。所以W的每一行，就是对应了一个字符作为第一个字符，其二元模型中第二个字符出现的概率分布。通过矩阵相乘获取了与统计模型相同的表示结果。而在统计模型中为了避免log取值出现负无穷，我们通过为最小值增加一点值作为模型平滑的结果，这个量加得越大，最终计算的概率越平滑，各个组合出现的概率越相近，同样地在W权重矩阵中，如果W的各个元素其初始值都被设置为0，那么取得的结果就是一个均值，各个组合的概率相同。这就引出了正则化，通过L2正则化，我们在实现梯度下降更新W参数的同时，也在尽可能让W变小，我们可以实现相同的平滑效果，这就是正则化和模型平滑的一个共通解释，很有趣。</p>

<h1 id="part-3-n-gram模型">Part 3-N-gram模型</h1>
<p>现在，可以利用神经网络模型，将二元模型扩展到N元模型，以N=3为例。可以构建训练验证测试数据集。现在我们设置embedding的大小为10,即<code class="language-plaintext highlighter-rouge">C.shape=[27,10]</code>,它代表了每个字符被映射到了一个10维的向量空间中。通过索引<code class="language-plaintext highlighter-rouge">C[X]</code>,可以获取到输入序列X中每个时间步下每个字符对应的10维向量表示，而<code class="language-plaintext highlighter-rouge">C[X].shape=[N,3,10]</code>，接下来就是相似地构建<code class="language-plaintext highlighter-rouge">W1.shape=[30,200],W2.shape=[200,27]</code>权重矩阵,同样地训练模型。在计算损失函数时，我们的操作一般是将取得的logits结果进行取指数(e)运算，全部变为正值，再使用softmax函数计算各自的概率，最后取log值求平均计算最大似然函数，现在，这些操作全部可以简化为<code class="language-plaintext highlighter-rouge">F.cross_entropy(logits,Y)</code>这一个表达式中。这就是N元模型，相比于二元模型没有任何特别的，都可以使用神经网络模型一步一步构建得到,但是一些地方也有变化，比如网络深度，隐藏层数，网络宽度，embedding大小，学习率变化，这些都是构建网络过程中的超参数。
此外，无论如何构建多么复杂的网络，网络的损失都不会变为0，一个较为直观的解释是，每个N元模型的开始都是以<code class="language-plaintext highlighter-rouge">&lt;S&gt;</code>作为第一个字符,所以通过起止符去预测下一个元素时，我们需要这种变化，如果每一次通过起止符去预测下一个元素的值都是不变的，这本来就自相矛盾。比如<code class="language-plaintext highlighter-rouge">&lt;S&gt;你好&lt;E&gt;</code>，<code class="language-plaintext highlighter-rouge">&lt;S&gt;我喜欢你&lt;E&gt;</code>。</p>

<h1 id="part-4-batchnorm">Part 4-BatchNorm</h1>
<p>在早期的深度学习实践中，对于权重矩阵的初始化需要十分精准，避免在张量传播过程中，其值域范围的均值和方差出现极端的不稳定情况。</p>
<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><table class="rouge-table"><tbody><tr><td class="rouge-gutter gl"><pre class="lineno">1
2
3
4
5
6
</pre></td><td class="rouge-code"><pre><span class="n">g</span><span class="o">=</span><span class="n">torch</span><span class="p">.</span><span class="n">Generator</span><span class="p">().</span><span class="n">manual_seed</span><span class="p">(</span><span class="mi">2147483647</span><span class="p">)</span>
<span class="n">C</span><span class="o">=</span><span class="n">torch</span><span class="p">.</span><span class="n">randn</span><span class="p">((</span><span class="n">vocab_size</span><span class="p">,</span><span class="n">n_embd</span><span class="p">),</span><span class="n">generator</span><span class="o">=</span><span class="n">g</span><span class="p">)</span>
<span class="n">W1</span><span class="o">=</span><span class="n">torch</span><span class="p">.</span><span class="n">randn</span><span class="p">((</span><span class="n">n_embd</span><span class="o">*</span><span class="n">block_size</span><span class="p">,</span><span class="n">n_hidden</span><span class="p">),</span><span class="n">generator</span><span class="o">=</span><span class="n">g</span><span class="p">)</span><span class="o">*</span><span class="p">(</span><span class="mi">5</span><span class="o">/</span><span class="mi">3</span><span class="p">)</span><span class="o">/</span><span class="p">((</span><span class="n">n_embd</span><span class="o">*</span><span class="n">block_size</span><span class="p">)</span><span class="o">**</span><span class="mf">0.5</span><span class="p">)</span> <span class="c1">#*0.2
</span><span class="n">b1</span><span class="o">=</span><span class="n">torch</span><span class="p">.</span><span class="n">randn</span><span class="p">(</span><span class="n">n_hidden</span><span class="p">,</span><span class="n">generator</span><span class="o">=</span><span class="n">g</span><span class="p">)</span><span class="o">*</span><span class="mf">0.01</span>
<span class="n">W2</span><span class="o">=</span><span class="n">torch</span><span class="p">.</span><span class="n">randn</span><span class="p">((</span><span class="n">n_hidden</span><span class="p">,</span><span class="n">vocab_size</span><span class="p">),</span><span class="n">generator</span><span class="o">=</span><span class="n">g</span><span class="p">)</span><span class="o">*</span><span class="mf">0.01</span>
<span class="n">b2</span><span class="o">=</span><span class="n">torch</span><span class="p">.</span><span class="n">randn</span><span class="p">(</span><span class="n">vocab_size</span><span class="p">,</span><span class="n">generator</span><span class="o">=</span><span class="n">g</span><span class="p">)</span><span class="o">*</span><span class="mf">0.01</span>
</pre></td></tr></tbody></table></code></pre></div></div>
<p>在该代码示例中，假如初始化时没有乘以系数，那么在第一次前向传播过程中，各个层的所得的均值和方差就会出现极端的不稳定情况，这会导致梯度消失或梯度爆炸的问题，从而影响模型的训练效果。以W1为例，<code class="language-plaintext highlighter-rouge">h=tanh(X@W1+b1)</code>,tanh函数的图像如下所示，它也是属于sigmoid函数簇
<img src="/assets/images/2025/tanh.png" alt="tahn" />
它的导数为<code class="language-plaintext highlighter-rouge">1-h**2</code>,链式规则为<code class="language-plaintext highlighter-rouge">self.grad=(1-h**2)*out.grad</code>从公式上可以看出，当h接近-1或1时，导数接近0，这会导致梯度消失的问题。而当h接近0时，导数接近1，它会传递梯度值。而如果没有合理的初始化W1权重矩阵，在第一次前向传播过程中，h的取值会非常大或非常小，这会导致梯度消失或梯度爆炸的问题。所以，在初始化W1权重矩阵时，需要乘以一个系数，比如<code class="language-plaintext highlighter-rouge">(5/3)/((n_embd*block_size)**0.5)</code>，这是根据tanh函数的性质推导出来的一个系数，它可以确保在第一次前向传播过程中，h的均值和方差不会出现较大的波动，从而避免梯度消失的问题。
当然我们可以凭借直觉和反复实现观察其内部的数值，来给定一个较为还不错的系数。而系统性的初始化方法，比如Xavier初始化，kaiming初始化，则在工程上系统性地对初始化进行了优化，避免了手动给定系数的过程，同时也确保了模型的训练效果。
当然这只是网络第一次训练中的第一个batch过程，为训练开了一个比较好的头。为了在整个训练过程中都保持较好的稳定状态，提出了一系列的归一化方法。比如batch normalization，layer normalization，instance normalization等。这些方法的基本思想都是在每个batch或每个样本中，对输入的特征进行归一化，从而避免梯度消失或梯度爆炸的问题。我们可以这样想象，有一个X=[x1,x2,…,xn]它有n个特征，其中每个特征的数值尺度是不一样的，范围从1-10000变化，那么进行梯度下降时，它们的更新速度或者说走的step的尺度也是不一样的，有的走得快，有的走得很慢。比较好的处理方法就是把这些特征的尺度都进行归一化处理，让它们都在同一个尺度下面进行训练。这就是归一化。
怎么在代码中手动实现一个batch normalization呢？原理很简单，我们可以在每一个batch中，计算当前batch的均值和方差，再对当前batch中的每个样本，减去均值后再除以方差，从而实现归一化。<code class="language-plaintext highlighter-rouge">hpreact=bngain*(hpreact-hpreact.mean(0,keepdim=True))/hpreact.std(0,keepdim=True)+bnbias</code>，为了让网络能够调整分布，使得一些神经元激活一些不敏感，引入bngain/gamma和beta/bnbias两个可学习的参数，这是因为我们不希望一直让网络强制保持标准的高斯分布，只是在前期没有任何知识的情况下无法假设，只能保持公平性，不让网络对任何一个对象有偏爱，但当随着网络训练的过程，网络能够偏爱一些胜过另一些，也就是神经元敏感与不敏感。
现在我们引入了bn，但这也引入了一个问题，我们的训练过程与数据出现了耦合。隐藏状态、激活值除了依赖于输入X、函数，还依赖随机选取的batch形成的组合数据，比如一些batch中的样本的某个特征的取值范围很大，而另一些样本的取值范围很小，这种随机抖动的作用，反而可以作为一种正则化，引入一点熵让模型难以过拟合。
不过，这也引入了一个问题，就是在预测时，如何在模型中已有的计算批量状态下的bn中进行适配？预测时我们输入的是一个样本，而不是一个batch，但是现在有一段代码是计算bn的。这里有两种实现方式，一种是固定训练集中的bn值，当训练完成后，重新计算整个训练集的bn值，当预测时使用这个固定的bn值。另一种实现方式则是采用动态更新的策略，动态评估bn值。省略了重新计算整个训练集的步骤。它的实现方式如下,基本思想就是记录一个running的状态，然后根据是否是训练过程来决定要不要计算当前批次的均值和方差，如果不是训练过程，则直接将记录的running状态的均值和方差赋值给实际要使用的均值和方差变量，再进行归一化处理。并且，如果是训练过程，就会在最后根据momentum这个动量的大小决定要更新到running状态上的分量，一般而言momentum取0.1代表取当前计算得到的批次的均值和方差的0.1，将其加入到running状态值中，这样就实现了根据每一个批次的均值和方差，对最终的均值和方差的更新。</p>
<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><table class="rouge-table"><tbody><tr><td class="rouge-gutter gl"><pre class="lineno">1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
</pre></td><td class="rouge-code"><pre><span class="k">class</span> <span class="nc">BatchNormld</span><span class="p">:</span>
  <span class="k">def</span> <span class="nf">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span><span class="n">dim</span><span class="p">,</span><span class="n">eps</span><span class="o">=</span><span class="mf">1e-5</span><span class="p">,</span><span class="n">momentum</span><span class="o">=</span><span class="mf">0.1</span><span class="p">):</span>
    <span class="bp">self</span><span class="p">.</span><span class="n">eps</span><span class="o">=</span><span class="n">eps</span>
    <span class="bp">self</span><span class="p">.</span><span class="n">momentum</span><span class="o">=</span><span class="n">momentum</span>
    <span class="bp">self</span><span class="p">.</span><span class="n">training</span><span class="o">=</span><span class="bp">True</span>
    <span class="c1">#parameters trained with backprop
</span>    <span class="bp">self</span><span class="p">.</span><span class="n">gamma</span><span class="o">=</span><span class="n">torch</span><span class="p">.</span><span class="n">ones</span><span class="p">(</span><span class="n">dim</span><span class="p">)</span>
    <span class="bp">self</span><span class="p">.</span><span class="n">beta</span><span class="o">=</span><span class="n">torch</span><span class="p">.</span><span class="n">zeros</span><span class="p">(</span><span class="n">dim</span><span class="p">)</span>
    <span class="c1">#buffers (trained with a running 'momentum update')
</span>    <span class="bp">self</span><span class="p">.</span><span class="n">running_mean</span><span class="o">=</span><span class="n">torch</span><span class="p">.</span><span class="n">zeros</span><span class="p">(</span><span class="n">dim</span><span class="p">)</span>
    <span class="bp">self</span><span class="p">.</span><span class="n">running_var</span><span class="o">=</span><span class="n">torch</span><span class="p">.</span><span class="n">ones</span><span class="p">(</span><span class="n">dim</span><span class="p">)</span>

  <span class="k">def</span> <span class="nf">__call__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span><span class="n">x</span><span class="p">):</span>
    <span class="c1">#calculate the forward pass
</span>    <span class="k">if</span> <span class="bp">self</span><span class="p">.</span><span class="n">training</span><span class="p">:</span>
      <span class="n">xmean</span><span class="o">=</span><span class="n">x</span><span class="p">.</span><span class="n">mean</span><span class="p">(</span><span class="mi">0</span><span class="p">,</span><span class="n">keepdim</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span>
      <span class="n">xvar</span><span class="o">=</span><span class="n">x</span><span class="p">.</span><span class="n">var</span><span class="p">(</span><span class="mi">0</span><span class="p">,</span><span class="n">keepdim</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span>
    <span class="k">else</span><span class="p">:</span>
      <span class="n">xmean</span><span class="o">=</span><span class="bp">self</span><span class="p">.</span><span class="n">running_mean</span>
      <span class="n">xvar</span><span class="o">=</span><span class="bp">self</span><span class="p">.</span><span class="n">running_var</span>
    <span class="n">xhat</span><span class="o">=</span><span class="p">(</span><span class="n">x</span><span class="o">-</span><span class="n">xmean</span><span class="p">)</span><span class="o">/</span><span class="n">torch</span><span class="p">.</span><span class="n">sqrt</span><span class="p">(</span><span class="n">xvar</span><span class="o">+</span><span class="bp">self</span><span class="p">.</span><span class="n">eps</span><span class="p">)</span>
    <span class="bp">self</span><span class="p">.</span><span class="n">out</span><span class="o">=</span><span class="bp">self</span><span class="p">.</span><span class="n">gamma</span><span class="o">*</span><span class="n">xhat</span><span class="o">+</span><span class="bp">self</span><span class="p">.</span><span class="n">beta</span>
    <span class="c1">#update the buffers
</span>    <span class="k">if</span> <span class="bp">self</span><span class="p">.</span><span class="n">training</span><span class="p">:</span>
      <span class="k">with</span> <span class="n">torch</span><span class="p">.</span><span class="n">no_grad</span><span class="p">():</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">running_mean</span><span class="o">=</span><span class="p">(</span><span class="mi">1</span><span class="o">-</span><span class="bp">self</span><span class="p">.</span><span class="n">momentum</span><span class="p">)</span><span class="o">*</span><span class="bp">self</span><span class="p">.</span><span class="n">running_mean</span><span class="o">+</span><span class="bp">self</span><span class="p">.</span><span class="n">momentum</span><span class="o">*</span><span class="n">xmean</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">running_var</span><span class="o">=</span><span class="p">(</span><span class="mi">1</span><span class="o">-</span><span class="bp">self</span><span class="p">.</span><span class="n">momentum</span><span class="p">)</span><span class="o">*</span><span class="bp">self</span><span class="p">.</span><span class="n">running_var</span><span class="o">+</span><span class="bp">self</span><span class="p">.</span><span class="n">momentum</span><span class="o">*</span><span class="n">xvar</span>
    <span class="k">return</span> <span class="bp">self</span><span class="p">.</span><span class="n">out</span>

  <span class="k">def</span> <span class="nf">parameters</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
    <span class="k">return</span> <span class="p">[</span><span class="bp">self</span><span class="p">.</span><span class="n">gamma</span><span class="p">,</span><span class="bp">self</span><span class="p">.</span><span class="n">beta</span><span class="p">]</span>
</pre></td></tr></tbody></table></code></pre></div></div>
<h1 id="part-5-backprob">Part 5-Backprob</h1>
<p>针对简单的多层感知机，其梯度反向传播较为简单，在计算时需要从矩阵的角度考察衡量，要注意是否需要沿着某一个维度/轴进行求和，因为在前向传播过程中会隐含地出现广播这个操作，所以很容易忽略掉。如果原来的变量是矩阵形式，那么其反向传播时的梯度也是矩阵形式。简而言之，变量正向和反向传播的形状都是不变的。</p>

<h1 id="part-6-building-a-wavenet">Part 6-Building a WaveNet</h1>
<p>这里面基本是按照第三部分的内容，不过将一些混乱的结构进行了整理，将Embedding的过程单独整理成了一个类。添加了展平层，此外对bn层进行了扩展，支持不同的维度。</p>

<h1 id="part-7-lets-build-gpt-from-scratch">Part 7-Let’s build GPT: from scratch</h1>
<p>通过一个简单的二元模型，该二元模型做的仅仅是将输入的token的整数索引映射到一个向量中，该向量的大小也是token的词典大小，表示的是从该token映射到下一个可能token的Logits值。它的参数大小是[vocab_size,vocab_size]，<code class="language-plaintext highlighter-rouge">logits=self.token_embedding_table(idx)</code>，举例该例子主要是为了实现一个简单的生成器，根据当前的输入生成下一个预测的输出。可以看到，我们只需要拿取计算的最后一个时间步的logits，然后通过softmax计算得到概率，根据该概率采样得到下一个预测的词的idx_next，然后将这个id_next添加到idx中，作为输入，在这个例子中由于是二元模型，所以只会用到前后两个词，但是作为通用的生成器，这里可以简单地修改，就可以利用规定的上下文大小比如8个词的上下文，去预测下一个词,比如我们可以更改block_size的大小，就可以只关注于最新的需要查看的上下文大小的词，使用这些最新的词去预测下一个词，而计算注意力这些额外的计算，都放在了forward函数中，如果我们注释掉这一行，那么并且现在forward函数中没有实施注意力相关的计算，就退回了最初的二元模型。此时在没有实施注意力代码的时候，你会怎么考虑控制上下文大小的注意力计算呢？我最初的想法很直接，注意力计算要考虑上下文大小，那就直接在注意力计算的过程中进行控制，但是这样又引入了新的问题，就是控制所选上下文的窗口的移动，比如现在的上下文窗口大小是8，现在我的输入有32个，那就要控制找到新的8个词的内容，这就会要求我们额外地在forward函数中添加额外的控制。但是我们不想因为这个改变通用的写法，所以换个思路，把这个窗口大小的控制逻辑放在了生成器中，直接截取idx的最新的上下文窗口大小的内容，这样，在forward函数中只需要按照输入计算注意力就可以了。</p>
<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><table class="rouge-table"><tbody><tr><td class="rouge-gutter gl"><pre class="lineno">1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
</pre></td><td class="rouge-code"><pre><span class="k">def</span> <span class="nf">forward</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">idx</span><span class="p">,</span> <span class="n">targets</span><span class="o">=</span><span class="bp">None</span><span class="p">):</span>

        <span class="c1"># idx and targets are both (B,T) tensor of integers
</span>        <span class="n">logits</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">token_embedding_table</span><span class="p">(</span><span class="n">idx</span><span class="p">)</span> <span class="c1"># (B,T,C)
</span>
        <span class="k">if</span> <span class="n">targets</span> <span class="ow">is</span> <span class="bp">None</span><span class="p">:</span>
            <span class="n">loss</span> <span class="o">=</span> <span class="bp">None</span>
        <span class="k">else</span><span class="p">:</span>
            <span class="n">B</span><span class="p">,</span> <span class="n">T</span><span class="p">,</span> <span class="n">C</span> <span class="o">=</span> <span class="n">logits</span><span class="p">.</span><span class="n">shape</span>
            <span class="n">logits</span> <span class="o">=</span> <span class="n">logits</span><span class="p">.</span><span class="n">view</span><span class="p">(</span><span class="n">B</span><span class="o">*</span><span class="n">T</span><span class="p">,</span> <span class="n">C</span><span class="p">)</span>
            <span class="n">targets</span> <span class="o">=</span> <span class="n">targets</span><span class="p">.</span><span class="n">view</span><span class="p">(</span><span class="n">B</span><span class="o">*</span><span class="n">T</span><span class="p">)</span>
            <span class="n">loss</span> <span class="o">=</span> <span class="n">F</span><span class="p">.</span><span class="n">cross_entropy</span><span class="p">(</span><span class="n">logits</span><span class="p">,</span> <span class="n">targets</span><span class="p">)</span>

        <span class="k">return</span> <span class="n">logits</span><span class="p">,</span> <span class="n">loss</span>

<span class="k">def</span> <span class="nf">generate</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">idx</span><span class="p">,</span> <span class="n">max_new_tokens</span><span class="p">):</span>
    <span class="c1"># idx is (B, T) array of indices in the current context
</span>    <span class="k">for</span> <span class="n">_</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="n">max_new_tokens</span><span class="p">):</span>
        <span class="c1"># get the predictions
</span>        <span class="n">idx_cond</span><span class="o">=</span><span class="n">idx</span><span class="p">[:,</span><span class="o">-</span><span class="n">block_size</span><span class="p">:]</span>
        <span class="n">logits</span><span class="p">,</span> <span class="n">loss</span> <span class="o">=</span> <span class="bp">self</span><span class="p">(</span><span class="n">idx_cond</span><span class="p">)</span>
        <span class="c1"># focus only on the last time step
</span>        <span class="n">logits</span> <span class="o">=</span> <span class="n">logits</span><span class="p">[:,</span> <span class="o">-</span><span class="mi">1</span><span class="p">,</span> <span class="p">:]</span> <span class="c1"># becomes (B, C)
</span>        <span class="c1"># apply softmax to get probabilities
</span>        <span class="n">probs</span> <span class="o">=</span> <span class="n">F</span><span class="p">.</span><span class="n">softmax</span><span class="p">(</span><span class="n">logits</span><span class="p">,</span> <span class="n">dim</span><span class="o">=-</span><span class="mi">1</span><span class="p">)</span> <span class="c1"># (B, C)
</span>        <span class="c1"># sample from the distribution
</span>        <span class="n">idx_next</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">multinomial</span><span class="p">(</span><span class="n">probs</span><span class="p">,</span> <span class="n">num_samples</span><span class="o">=</span><span class="mi">1</span><span class="p">)</span> <span class="c1"># (B, 1)
</span>        <span class="c1"># append sampled index to the running sequence
</span>        <span class="n">idx</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">cat</span><span class="p">((</span><span class="n">idx</span><span class="p">,</span> <span class="n">idx_next</span><span class="p">),</span> <span class="n">dim</span><span class="o">=</span><span class="mi">1</span><span class="p">)</span> <span class="c1"># (B, T+1)
</span>    <span class="k">return</span> <span class="n">idx</span>
</pre></td></tr></tbody></table></code></pre></div></div>
<p>注意力的数学本质，就是求得加权后的新值，理解的意思就是W是权重矩阵，而X是待加权的变量，通过W@X，就可以根据W中的权重来重新调整X中每一个元素的值，该值参考了其它值，在该元素对应位置生成的新的加权后的元素值，其是由与该位置有关的原来的元素值，以及其它相关值通过加权运算得到的。</p>
<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><table class="rouge-table"><tbody><tr><td class="rouge-gutter gl"><pre class="lineno">1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
</pre></td><td class="rouge-code"><pre><span class="n">torch</span><span class="p">.</span><span class="n">manual_seed</span><span class="p">(</span><span class="mi">1337</span><span class="p">)</span>
<span class="n">B</span><span class="p">,</span><span class="n">T</span><span class="p">,</span><span class="n">C</span> <span class="o">=</span> <span class="mi">4</span><span class="p">,</span><span class="mi">8</span><span class="p">,</span><span class="mi">32</span> <span class="c1"># batch, time, channels
</span><span class="n">x</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">randn</span><span class="p">(</span><span class="n">B</span><span class="p">,</span><span class="n">T</span><span class="p">,</span><span class="n">C</span><span class="p">)</span>

<span class="c1"># let's see a single Head perform self-attention
</span><span class="n">head_size</span> <span class="o">=</span> <span class="mi">16</span>
<span class="n">key</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">(</span><span class="n">C</span><span class="p">,</span> <span class="n">head_size</span><span class="p">,</span> <span class="n">bias</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>
<span class="n">query</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">(</span><span class="n">C</span><span class="p">,</span> <span class="n">head_size</span><span class="p">,</span> <span class="n">bias</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>
<span class="n">value</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">(</span><span class="n">C</span><span class="p">,</span> <span class="n">head_size</span><span class="p">,</span> <span class="n">bias</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>
<span class="n">k</span> <span class="o">=</span> <span class="n">key</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>   <span class="c1"># (B, T, 16)
</span><span class="n">q</span> <span class="o">=</span> <span class="n">query</span><span class="p">(</span><span class="n">x</span><span class="p">)</span> <span class="c1"># (B, T, 16)
</span><span class="n">wei</span> <span class="o">=</span>  <span class="n">q</span> <span class="o">@</span> <span class="n">k</span><span class="p">.</span><span class="n">transpose</span><span class="p">(</span><span class="o">-</span><span class="mi">2</span><span class="p">,</span> <span class="o">-</span><span class="mi">1</span><span class="p">)</span> <span class="c1"># (B, T, 16) @ (B, 16, T) ---&gt; (B, T, T)
</span>
<span class="n">tril</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">tril</span><span class="p">(</span><span class="n">torch</span><span class="p">.</span><span class="n">ones</span><span class="p">(</span><span class="n">T</span><span class="p">,</span> <span class="n">T</span><span class="p">))</span>
<span class="c1">#wei = torch.zeros((T,T))
</span><span class="n">wei</span> <span class="o">=</span> <span class="n">wei</span><span class="p">.</span><span class="n">masked_fill</span><span class="p">(</span><span class="n">tril</span> <span class="o">==</span> <span class="mi">0</span><span class="p">,</span> <span class="nb">float</span><span class="p">(</span><span class="s">'-inf'</span><span class="p">))</span>
<span class="n">wei</span> <span class="o">=</span> <span class="n">F</span><span class="p">.</span><span class="n">softmax</span><span class="p">(</span><span class="n">wei</span><span class="p">,</span> <span class="n">dim</span><span class="o">=-</span><span class="mi">1</span><span class="p">)</span>

<span class="n">v</span> <span class="o">=</span> <span class="n">value</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
<span class="n">out</span> <span class="o">=</span> <span class="n">wei</span> <span class="o">@</span> <span class="n">v</span>
</pre></td></tr></tbody></table></code></pre></div></div>
<p>在这个例子中，key和query都是从原x中通过矩阵投影而来，通过这两个投影计算得到wei即权重矩阵，同样地我们也再一次对x投影得到v值，那么v的加权后的新值，就是<code class="language-plaintext highlighter-rouge">out=wei@v</code>，这就是注意力的计算过程。不过此时的形状还是[B,T,head_size]，还要使用一个投影矩阵将其投影回原来的大小[B,T,C]，我们需要记住的是，经过注意力加权后，新的变量的形状应该与原来的变量形状保持一致，因为我们仅仅是做一个加权处理，如果最终改变了形状就不正确了。同样地需要注意，在生成式语言模型中，就是根据上下文去预测下一个词时，按照正常逻辑，我们应该只能看到截止到当前时间步及其以前时间步的内容，计算它们之间的注意力值，这里我们使用掩膜将上三角的值换成了<code class="language-plaintext highlighter-rouge">-inf</code>,这就符合了[生成]这个词所表达的含义。
需要注意的是，</p>
<ol>
  <li>在计算注意力时没有空间相对位置的关系，想想<code class="language-plaintext highlighter-rouge">a=a_1*w_1+a_2*w_2+a_3*w_3</code>等同于<code class="language-plaintext highlighter-rouge">a=a_3*w_3+a_1*w_1+a_2*w_2</code>，所以我们需要自己额外添加一个位置编码进去。</li>
  <li>注意力可以看作一种交流机制，可以被视为有向图中的节点，这些节点相互观察，并通过来自所有指向它们的节点的加权和来聚合信息，其中权重取决于数据。</li>
  <li>在batch维度跨越的样本之间是无法交流的，因为矩阵乘法只发生在最后两个维度。</li>
  <li>在编码器中，只需要删除进行掩膜的那行代码，就从解码器变成了编码器，从历史的原因分析，最初transformer架构的提出是针对语言翻译的，所以使用编码器去了解翻译对象的全部信息，再使用解码器也就是对注意力权重矩阵增加了掩膜的部分，去生成翻译的目标语言。单独的解码器架构也就是自回归模型，常用于大语言模型。</li>
  <li>自注意力，就是说除了query是来自于x外，key,value也都来自于x的投影变换，如果query，value来自于其它数据，就叫做cross-attention。</li>
  <li>“缩放”注意力机制额外将wei除以1/√(头大小)。这样一来，当输入Q、K具有单位方差时，wei也会具有单位方差，并且Softmax会保持分散状态，不会过度饱和。</li>
</ol>

<h1 id="part-8-tokenization">Part 8-Tokenization</h1>
<p>LLM的许多奇怪的现象都和tokenization有关。例如，大语言模型不能很好的拼写单词，在编写代码方面表现很差，这些都与tokenization有关。
在本例中针对字符级别，将其转为utf-8编码，然后转为整数。</p>
<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><table class="rouge-table"><tbody><tr><td class="rouge-gutter gl"><pre class="lineno">1
2
3
4
</pre></td><td class="rouge-code"><pre><span class="k">print</span><span class="p">(</span><span class="s">'你好'</span><span class="p">.</span><span class="n">encode</span><span class="p">(</span><span class="s">'utf-8'</span><span class="p">))</span>
<span class="sa">b</span><span class="s">'</span><span class="se">\xe4\xbd\xa0\xe5\xa5\xbd</span><span class="s">'</span>
<span class="k">print</span><span class="p">(</span><span class="nb">list</span><span class="p">(</span><span class="nb">map</span><span class="p">(</span><span class="nb">int</span><span class="p">,</span><span class="s">'你好'</span><span class="p">.</span><span class="n">encode</span><span class="p">(</span><span class="s">'utf-8'</span><span class="p">))))</span>
<span class="p">[</span><span class="mi">228</span><span class="p">,</span> <span class="mi">189</span><span class="p">,</span> <span class="mi">160</span><span class="p">,</span> <span class="mi">229</span><span class="p">,</span> <span class="mi">165</span><span class="p">,</span> <span class="mi">189</span><span class="p">]</span>
</pre></td></tr></tbody></table></code></pre></div></div>
<p>这样我们便得到了用整数表示的一些列token。仅仅一个<code class="language-plaintext highlighter-rouge">你好</code>的token表示就是6个整数，每个整数都在0到255之间。而具有这样规律的整数在大量文本中会重复出现，这就是语义的表达。而大预言模型只是接收这些整数作为输入，然后根据其内部的参数，进行预测下一个token出现的概率。我们可以将[228, 189, 160, 229, 165, 189]这样的具有语义信息的整数进行融合，创造出一个新的整数来代表该序列，就是使用256来代表<code class="language-plaintext highlighter-rouge">你好</code>，重复这样的操作我们就可以得到扩充了的，压缩了信息的token库。但是token库的大小也不能无限大，例如使用10000这个整数来表达<code class="language-plaintext highlighter-rouge">针对简单的多层感知机，其梯度反向传播较为简单</code>这样一句话，已经压缩了许多信息，无法获取其内部的详细语义信息，所以token库的大小要适合就行，也就对应了我们需要融合的次数，每次我们融合都选择出现频率最高的一对整数，并将其赋予新的值，只要我们在合适的合并次数下停止，就能够得到较为合理的，既能合理高效表示语义信息，又不至于压缩得太厉害的程度。</p>
<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><table class="rouge-table"><tbody><tr><td class="rouge-gutter gl"><pre class="lineno">1
2
3
4
5
</pre></td><td class="rouge-code"><pre><span class="k">def</span> <span class="nf">get_stats</span><span class="p">(</span><span class="n">ids</span><span class="p">):</span>
    <span class="n">counts</span> <span class="o">=</span> <span class="p">{}</span>
    <span class="k">for</span> <span class="n">pair</span> <span class="ow">in</span> <span class="nb">zip</span><span class="p">(</span><span class="n">ids</span><span class="p">,</span> <span class="n">ids</span><span class="p">[</span><span class="mi">1</span><span class="p">:]):</span> <span class="c1"># Pythonic way to iterate consecutive elements
</span>        <span class="n">counts</span><span class="p">[</span><span class="n">pair</span><span class="p">]</span> <span class="o">=</span> <span class="n">counts</span><span class="p">.</span><span class="n">get</span><span class="p">(</span><span class="n">pair</span><span class="p">,</span> <span class="mi">0</span><span class="p">)</span> <span class="o">+</span> <span class="mi">1</span>
    <span class="k">return</span> <span class="n">counts</span>
</pre></td></tr></tbody></table></code></pre></div></div>
<p>我们使用<code class="language-plaintext highlighter-rouge">get_stats</code>函数，对输入得tokens进行统计，得到所有相邻的token对出现的次数并返回。</p>
<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><table class="rouge-table"><tbody><tr><td class="rouge-gutter gl"><pre class="lineno">1
2
3
4
5
6
7
8
9
10
11
12
13
</pre></td><td class="rouge-code"><pre><span class="k">def</span> <span class="nf">merge</span><span class="p">(</span><span class="n">ids</span><span class="p">,</span> <span class="n">pair</span><span class="p">,</span> <span class="n">idx</span><span class="p">):</span>
  <span class="c1"># in the list of ints (ids), replace all consecutive occurences of pair with the new token idx
</span>  <span class="n">newids</span> <span class="o">=</span> <span class="p">[]</span>
  <span class="n">i</span> <span class="o">=</span> <span class="mi">0</span>
  <span class="k">while</span> <span class="n">i</span> <span class="o">&lt;</span> <span class="nb">len</span><span class="p">(</span><span class="n">ids</span><span class="p">):</span>
    <span class="c1"># if we are not at the very last position AND the pair matches, replace it
</span>    <span class="k">if</span> <span class="n">i</span> <span class="o">&lt;</span> <span class="nb">len</span><span class="p">(</span><span class="n">ids</span><span class="p">)</span> <span class="o">-</span> <span class="mi">1</span> <span class="ow">and</span> <span class="n">ids</span><span class="p">[</span><span class="n">i</span><span class="p">]</span> <span class="o">==</span> <span class="n">pair</span><span class="p">[</span><span class="mi">0</span><span class="p">]</span> <span class="ow">and</span> <span class="n">ids</span><span class="p">[</span><span class="n">i</span><span class="o">+</span><span class="mi">1</span><span class="p">]</span> <span class="o">==</span> <span class="n">pair</span><span class="p">[</span><span class="mi">1</span><span class="p">]:</span>
      <span class="n">newids</span><span class="p">.</span><span class="n">append</span><span class="p">(</span><span class="n">idx</span><span class="p">)</span>
      <span class="n">i</span> <span class="o">+=</span> <span class="mi">2</span>
    <span class="k">else</span><span class="p">:</span>
      <span class="n">newids</span><span class="p">.</span><span class="n">append</span><span class="p">(</span><span class="n">ids</span><span class="p">[</span><span class="n">i</span><span class="p">])</span>
      <span class="n">i</span> <span class="o">+=</span> <span class="mi">1</span>
  <span class="k">return</span> <span class="n">newids</span>
</pre></td></tr></tbody></table></code></pre></div></div>
<p>然后我们可以使用<code class="language-plaintext highlighter-rouge">merge</code>函数，融合指定的token pair，并赋值为新的token idx，返回的是融合后的token序列，原来的token pair被token idx替代了。所以，如果原token长度为n，其中有m个token pair，那么融合后token序列的长度为n-m。下面这段代码将整个融合的流程串联了起来，首先获取token pair对，然后对获取的token pair对进行排序，获得出现频率最高的pair，之后为这个pair赋值一个新的token整数值，来代替原来的pair，实现信息的高效、压缩表示，并用一个字典记录被替换的pair和它的新的整数token值的表示之间的映射关系，方便后续解码的时候反推回原来的表示。</p>
<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><table class="rouge-table"><tbody><tr><td class="rouge-gutter gl"><pre class="lineno">1
2
3
4
5
6
7
8
9
10
11
12
</pre></td><td class="rouge-code"><pre><span class="n">vocab_size</span> <span class="o">=</span> <span class="mi">276</span> <span class="c1"># the desired final vocabulary size
</span><span class="n">num_merges</span> <span class="o">=</span> <span class="n">vocab_size</span> <span class="o">-</span> <span class="mi">256</span>
<span class="n">ids</span> <span class="o">=</span> <span class="nb">list</span><span class="p">(</span><span class="n">tokens</span><span class="p">)</span> <span class="c1"># copy so we don't destroy the original list
</span>
<span class="n">merges</span> <span class="o">=</span> <span class="p">{}</span> <span class="c1"># (int, int) -&gt; int
</span><span class="k">for</span> <span class="n">i</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="n">num_merges</span><span class="p">):</span>
  <span class="n">stats</span> <span class="o">=</span> <span class="n">get_stats</span><span class="p">(</span><span class="n">ids</span><span class="p">)</span>
  <span class="n">pair</span> <span class="o">=</span> <span class="nb">max</span><span class="p">(</span><span class="n">stats</span><span class="p">,</span> <span class="n">key</span><span class="o">=</span><span class="n">stats</span><span class="p">.</span><span class="n">get</span><span class="p">)</span>
  <span class="n">idx</span> <span class="o">=</span> <span class="mi">256</span> <span class="o">+</span> <span class="n">i</span>
  <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">"merging </span><span class="si">{</span><span class="n">pair</span><span class="si">}</span><span class="s"> into a new token </span><span class="si">{</span><span class="n">idx</span><span class="si">}</span><span class="s">"</span><span class="p">)</span>
  <span class="n">ids</span> <span class="o">=</span> <span class="n">merge</span><span class="p">(</span><span class="n">ids</span><span class="p">,</span> <span class="n">pair</span><span class="p">,</span> <span class="n">idx</span><span class="p">)</span>
  <span class="n">merges</span><span class="p">[</span><span class="n">pair</span><span class="p">]</span> <span class="o">=</span> <span class="n">idx</span>
</pre></td></tr></tbody></table></code></pre></div></div>
<p>我们通过这个循环，融合了20个token pair，得到了一个新的token库，其中包含了256个原始token，以及20个新的token，每个新的token都是由两个原始token融合而来的。并且可以比较融合后的token长度和原始token长度得到一个压缩率，用来表示压缩的程度。Tokenizer是一个独立于LLM的模块，有其自己的训练集，使用上面的方法即Byte pair encoding(BPE)算法，它将在文本和一些列的tokens之间转换，输入到LLM中的只有转换后的tokens。
<img src="/assets/images/2025/tokenizer.png" alt="tokenizer" /></p>
<h2 id="decoding">Decoding</h2>
<p>现在我们已经实现了融合token pair的操作，就是我们已经根据训练集训练好了一个tokenizer，现在我们需要应用到实际中。但是怎么解码呢？即怎么从一段tokens中，恢复出原来的文本？很简单的直观的方法就是根据融合的链式替换，一步步将新的token idx替换为原来的token pair，直到恢复出原来的文本。不过使用字节级的形式，可以很取巧地快速解码，因为字节的表示可以相加，就像字符一样进行拼接。<code class="language-plaintext highlighter-rouge">bytes([66])+bytes([67])==b'BC'</code></p>
<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><table class="rouge-table"><tbody><tr><td class="rouge-gutter gl"><pre class="lineno">1
2
3
4
5
6
7
8
9
</pre></td><td class="rouge-code"><pre><span class="n">vocab</span> <span class="o">=</span> <span class="p">{</span><span class="n">idx</span><span class="p">:</span> <span class="nb">bytes</span><span class="p">([</span><span class="n">idx</span><span class="p">])</span> <span class="k">for</span> <span class="n">idx</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="mi">256</span><span class="p">)}</span>
<span class="k">for</span> <span class="p">(</span><span class="n">p0</span><span class="p">,</span> <span class="n">p1</span><span class="p">),</span> <span class="n">idx</span> <span class="ow">in</span> <span class="n">merges</span><span class="p">.</span><span class="n">items</span><span class="p">():</span>
    <span class="n">vocab</span><span class="p">[</span><span class="n">idx</span><span class="p">]</span> <span class="o">=</span> <span class="n">vocab</span><span class="p">[</span><span class="n">p0</span><span class="p">]</span> <span class="o">+</span> <span class="n">vocab</span><span class="p">[</span><span class="n">p1</span><span class="p">]</span>

<span class="k">def</span> <span class="nf">decode</span><span class="p">(</span><span class="n">ids</span><span class="p">):</span>
  <span class="c1"># given ids (list of integers), return Python string
</span>  <span class="n">tokens</span> <span class="o">=</span> <span class="sa">b</span><span class="s">""</span><span class="p">.</span><span class="n">join</span><span class="p">(</span><span class="n">vocab</span><span class="p">[</span><span class="n">idx</span><span class="p">]</span> <span class="k">for</span> <span class="n">idx</span> <span class="ow">in</span> <span class="n">ids</span><span class="p">)</span>
  <span class="n">text</span> <span class="o">=</span> <span class="n">tokens</span><span class="p">.</span><span class="n">decode</span><span class="p">(</span><span class="s">"utf-8"</span><span class="p">,</span> <span class="n">errors</span><span class="o">=</span><span class="s">"replace"</span><span class="p">)</span>
  <span class="k">return</span> <span class="n">text</span>
</pre></td></tr></tbody></table></code></pre></div></div>
<p>我们首先创建一个字典vocab，包含了256个原始token的字节表示，然后根据merges字典，将融合后的token pair的字节表示，赋值为原来的token pair的字节表示的拼接。最后我们可以使用decode函数，将一段tokens恢复为原来的文本。这样就不用一步一步去替换新的token idx为原来的token pair，而是直接将所有的token idx对应的字节表示拼接起来，再解码为文本。</p>
<h2 id="encoding">Encoding</h2>
<p>和解码相对的就是我们怎么编码文本为tokens。编码的过程和解码的过程是相反的，我们从文本开始，将其编码为utf-8的字节序列，并转为list对象，此刻就成为了整数表示形式，然后我们反复更新这个tokens序列也就是编码它，除非它的长度小于2，即它无法再融合，或者最小的token pair都已经不在merges字典中了，就停止编码并返回此刻的tokens。为什么不从最大的pair对开始融合呢，因为最大的pair对也是从较小的融合而来的，在不融合出它的子集之前，在tokens中找不到找这个较大的pair对。</p>
<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><table class="rouge-table"><tbody><tr><td class="rouge-gutter gl"><pre class="lineno">1
2
3
4
5
6
7
8
9
10
11
12
13
</pre></td><td class="rouge-code"><pre><span class="k">def</span> <span class="nf">encode</span><span class="p">(</span><span class="n">text</span><span class="p">):</span>
  <span class="c1"># given a string, return list of integers (the tokens)
</span>  <span class="n">tokens</span> <span class="o">=</span> <span class="nb">list</span><span class="p">(</span><span class="n">text</span><span class="p">.</span><span class="n">encode</span><span class="p">(</span><span class="s">"utf-8"</span><span class="p">))</span>
  <span class="k">while</span> <span class="nb">len</span><span class="p">(</span><span class="n">tokens</span><span class="p">)</span> <span class="o">&gt;=</span> <span class="mi">2</span><span class="p">:</span>
    <span class="n">stats</span> <span class="o">=</span> <span class="n">get_stats</span><span class="p">(</span><span class="n">tokens</span><span class="p">)</span>
    <span class="n">pair</span> <span class="o">=</span> <span class="nb">min</span><span class="p">(</span><span class="n">stats</span><span class="p">,</span> <span class="n">key</span><span class="o">=</span><span class="k">lambda</span> <span class="n">p</span><span class="p">:</span> <span class="n">merges</span><span class="p">.</span><span class="n">get</span><span class="p">(</span><span class="n">p</span><span class="p">,</span> <span class="nb">float</span><span class="p">(</span><span class="s">"inf"</span><span class="p">)))</span>
    <span class="k">if</span> <span class="n">pair</span> <span class="ow">not</span> <span class="ow">in</span> <span class="n">merges</span><span class="p">:</span>
      <span class="k">break</span> <span class="c1"># nothing else can be merged
</span>    <span class="n">idx</span> <span class="o">=</span> <span class="n">merges</span><span class="p">[</span><span class="n">pair</span><span class="p">]</span>
    <span class="n">tokens</span> <span class="o">=</span> <span class="n">merge</span><span class="p">(</span><span class="n">tokens</span><span class="p">,</span> <span class="n">pair</span><span class="p">,</span> <span class="n">idx</span><span class="p">)</span>
  <span class="k">return</span> <span class="n">tokens</span>

<span class="k">print</span><span class="p">(</span><span class="n">encode</span><span class="p">(</span><span class="s">""</span><span class="p">))</span>
</pre></td></tr></tbody></table></code></pre></div></div>
<p>所以，merges这个变量记录了所有的融合操作，可以用于编码阶段，而vocab变量记录了所有的token的字节表示，可以用于解码阶段。</p>
<h2 id="正则表达式">正则表达式</h2>
<p>有了编解码，我们就可以利用tokenizer，将文本转换为tokens，然后输入到LLM中进行处理，再从LLM中输出tokens，最后利用tokenizer将tokens转换为文本。
但是，训练这样一个简单的tokenizer无法解决一些特殊的问题，比如在英文句子中，单词之间有空格，而tokenizer是基于字节级的，所以空格会被编码为一个token，在训练集中就会出现很多这样的模式，例如<code class="language-plaintext highlighter-rouge">dog</code>这个单词，在文本中会出现许多代表相同意思的但是稍微有些不同字节表示的形式<code class="language-plaintext highlighter-rouge">'dog.',' dog','dog!','dog?'</code>，所以BPE会对每一个这样的不同模式的<code class="language-plaintext highlighter-rouge">dog</code>构建一个token，获得了只是稍有不同的<code class="language-plaintext highlighter-rouge">dog</code>这样的token，就是BPE把一些不该编码的部分也加入了进来，比如单词和标点符号。
所以必须要有一种人工的干预，强制不让一些符合规则的字符合并在一起。所以基本上我们可以构建这样的一个正则表达式<code class="language-plaintext highlighter-rouge">"""'s|'t|'re|'ve|'m|'ll|'d| ?\p{L}+| ?\p{N}+| ?[^\s\p{L}\p{N}]+|\s+(?!\S)|\s+"""</code>，用来分离那些不该被合并的字符。它的实现逻辑就是我们不去找那些烦扰的字符，而是去匹配我们合理的关心的模式，把这些模式提取出来并形成一个列表，那么没有提取出来的就是我们不关心的东西，所以我们只提取出了关切的部分。在<code class="language-plaintext highlighter-rouge">hello world how are your</code>这个例子中，经过正则匹配我们可以得到一个列表<code class="language-plaintext highlighter-rouge">['hello',' world',' how',' are',' you']</code>，我们在训练tokenizer之前，首先做的就是人工的划分单词，就是<code class="language-plaintext highlighter-rouge">' are' ' you'</code>中<code class="language-plaintext highlighter-rouge">e</code>不会和<code class="language-plaintext highlighter-rouge"> you</code>中的第一个空格融合。接下来所有的融合都只会各自发生在列表中的各个元素中，不会跨元素进行融合，融合之后的结果再进行拼接就得到了我们训练好的tokenizer。</p>
<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><table class="rouge-table"><tbody><tr><td class="rouge-gutter gl"><pre class="lineno">1
2
3
4
5
6
7
8
9
10
</pre></td><td class="rouge-code"><pre><span class="k">for</span> <span class="n">i</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="mi">101</span><span class="p">):</span>
    <span class="k">if</span> <span class="n">i</span> <span class="o">%</span> <span class="mi">3</span> <span class="o">==</span> <span class="mi">0</span> <span class="ow">and</span> <span class="n">i</span> <span class="o">%</span> <span class="mi">5</span> <span class="o">==</span> <span class="mi">0</span><span class="p">:</span>
        <span class="k">print</span><span class="p">(</span><span class="s">"FizzBuzz"</span><span class="p">)</span>
    <span class="k">elif</span> <span class="n">i</span> <span class="o">%</span> <span class="mi">3</span> <span class="o">==</span> <span class="mi">0</span><span class="p">:</span>
        <span class="k">print</span><span class="p">(</span><span class="s">"Fizz"</span><span class="p">)</span>
    <span class="k">elif</span> <span class="n">i</span> <span class="o">%</span> <span class="mi">5</span> <span class="o">==</span> <span class="mi">0</span><span class="p">:</span>
        <span class="k">print</span><span class="p">(</span><span class="s">"Buzz"</span><span class="p">)</span>
    <span class="k">else</span><span class="p">:</span>
        <span class="k">print</span><span class="p">(</span><span class="n">i</span><span class="p">)</span>
<span class="p">[</span><span class="s">'</span><span class="se">\n</span><span class="s">'</span><span class="p">,</span> <span class="s">'for'</span><span class="p">,</span> <span class="s">' i'</span><span class="p">,</span> <span class="s">' in'</span><span class="p">,</span> <span class="s">' range'</span><span class="p">,</span> <span class="s">'('</span><span class="p">,</span> <span class="s">'1'</span><span class="p">,</span> <span class="s">','</span><span class="p">,</span> <span class="s">' 101'</span><span class="p">,</span> <span class="s">'):'</span><span class="p">,</span> <span class="s">'</span><span class="se">\n</span><span class="s">   '</span><span class="p">,</span> <span class="s">' if'</span><span class="p">,</span> <span class="s">' i'</span><span class="p">,</span> <span class="s">' %'</span><span class="p">,</span> <span class="s">' 3'</span><span class="p">,</span> <span class="s">' =='</span><span class="p">,</span> <span class="s">' 0'</span><span class="p">,</span> <span class="s">' and'</span><span class="p">,</span> <span class="s">' i'</span><span class="p">,</span> <span class="s">' %'</span><span class="p">,</span> <span class="s">' 5'</span><span class="p">,</span> <span class="s">' =='</span><span class="p">,</span> <span class="s">' 0'</span><span class="p">,</span> <span class="s">':'</span><span class="p">,</span> <span class="s">'</span><span class="se">\n</span><span class="s">       '</span><span class="p">,</span> <span class="s">' print'</span><span class="p">,</span> <span class="s">'("'</span><span class="p">,</span> <span class="s">'FizzBuzz'</span><span class="p">,</span> <span class="s">'")'</span><span class="p">,</span> <span class="s">'</span><span class="se">\n</span><span class="s">   '</span><span class="p">,</span> <span class="s">' elif'</span><span class="p">,</span> <span class="s">' i'</span><span class="p">,</span> <span class="s">' %'</span><span class="p">,</span> <span class="s">' 3'</span><span class="p">,</span> <span class="s">' =='</span><span class="p">,</span> <span class="s">' 0'</span><span class="p">,</span> <span class="s">':'</span><span class="p">,</span> <span class="s">'</span><span class="se">\n</span><span class="s">       '</span><span class="p">,</span> <span class="s">' print'</span><span class="p">,</span> <span class="s">'("'</span><span class="p">,</span> <span class="s">'Fizz'</span><span class="p">,</span> <span class="s">'")'</span><span class="p">,</span> <span class="s">'</span><span class="se">\n</span><span class="s">   '</span><span class="p">,</span> <span class="s">' elif'</span><span class="p">,</span> <span class="s">' i'</span><span class="p">,</span> <span class="s">' %'</span><span class="p">,</span> <span class="s">' 5'</span><span class="p">,</span> <span class="s">' =='</span><span class="p">,</span> <span class="s">' 0'</span><span class="p">,</span> <span class="s">':'</span><span class="p">,</span> <span class="s">'</span><span class="se">\n</span><span class="s">       '</span><span class="p">,</span> <span class="s">' print'</span><span class="p">,</span> <span class="s">'("'</span><span class="p">,</span> <span class="s">'Buzz'</span><span class="p">,</span> <span class="s">'")'</span><span class="p">,</span> <span class="s">'</span><span class="se">\n</span><span class="s">   '</span><span class="p">,</span> <span class="s">' else'</span><span class="p">,</span> <span class="s">':'</span><span class="p">,</span> <span class="s">'</span><span class="se">\n</span><span class="s">       '</span><span class="p">,</span> <span class="s">' print'</span><span class="p">,</span> <span class="s">'('</span><span class="p">,</span> <span class="s">'i'</span><span class="p">,</span> <span class="s">')'</span><span class="p">,</span> <span class="s">'</span><span class="se">\n</span><span class="s">'</span><span class="p">]</span>
</pre></td></tr></tbody></table></code></pre></div></div>
<p><a href="https://tiktokenizer.vercel.app">tiktokenizer</a>这个网站中可以查看大语言模型的tokenizer结果,在GPT-2中，这些空格都没有进行合并，所以GPT-2中的tokenizer训练不仅仅是简单地将BPE应用到每一个提取的单词中，而是额外添加了一些规则，不过在这里我们理解预处理的过程，使用简单的处理方式就可以了。
<img src="/assets/images/2025/token.png" alt="Token" />
在GPT-4o中，可以看到同样的代码，但是在tokenizer中，空格被合并了，对tokenizer的训练进行了优化。
<img src="/assets/images/2025/gpt-4o tokenizer.png" alt="gpt-4o tokenizer" />
tiktoken库是OpenAI官方的一个tokenizer库，它可以用来将文本转换为tokens，也可以将tokens转换为文本。它的使用方法和我们之前实现的tokenizer类似，但是它是基于OpenAI的模型训练的，所以它的tokenizer结果和OpenAI的模型训练结果是一致的。</p>
<h2 id="special-tokens">special tokens</h2>
<p>len(encoder) #256 raw  byte tokens. 5,000 merges. +1 special token。这个special token就是<code class="language-plaintext highlighter-rouge">&lt;|endoftext|&gt;</code>，它在OpenAI的模型中被用来表示一段文本的结束，也就是说这个特殊token前后的两段文本是独立的，它们没有任何关系。假设我们有一个很大的数据集，这些数据集都是独立的文本，从各个数据源获取的，我们当然希望这些文本之间不应该有语义上的关联，比如A文档是在叙述一段小说，而B文档则是关于科学的东西，所以需要这个标志告诉模型前后之间不相关，上一段的文档内容提供的信息不应该继承到下一段文本中来。当然，可以注册许多特殊的token，比如用来区分用户和系统提示信息，不希望将系统提示信息暴露给用户方。</p>
<h2 id="实现正则化版本的tokenizer">实现正则化版本的tokenizer</h2>
<p>这段代码就是一个正则化版本的tokenizer训练过程，它的基本用到的函数都没有改变，只是现在需要作用到每一个提取的单词中，而不是直接作用到整个文本中。获取每个独立单词中出现的所有字符对的统计信息，然后根据统计信息对每个单词进行合并操作，直到达到指定的合并次数。</p>
<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><table class="rouge-table"><tbody><tr><td class="rouge-gutter gl"><pre class="lineno">1
2
3
4
5
6
7
8
9
10
11
12
13
</pre></td><td class="rouge-code"><pre><span class="n">text_chunks</span><span class="o">=</span><span class="n">re</span><span class="p">.</span><span class="n">findall</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">compiled_pattern</span><span class="p">,</span><span class="n">text</span><span class="p">)</span>
<span class="n">ids</span><span class="o">=</span><span class="p">[</span><span class="nb">list</span><span class="p">(</span><span class="n">ch</span><span class="p">.</span><span class="n">encode</span><span class="p">(</span><span class="s">'utf-8'</span><span class="p">))</span> <span class="k">for</span> <span class="n">ch</span> <span class="ow">in</span> <span class="n">text_chunks</span><span class="p">]</span>

<span class="k">for</span> <span class="n">i</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="n">num_merges</span><span class="p">):</span>
  <span class="n">idx</span><span class="o">=</span><span class="n">i</span><span class="o">+</span><span class="mi">256</span>
  <span class="n">stats</span><span class="o">=</span><span class="p">{}</span>
  <span class="k">for</span> <span class="n">ch</span> <span class="ow">in</span> <span class="n">ids</span><span class="p">:</span>
    <span class="n">stats</span><span class="o">=</span><span class="bp">self</span><span class="p">.</span><span class="n">_get_stats</span><span class="p">(</span><span class="n">ch</span><span class="p">,</span><span class="n">stats</span><span class="p">)</span>
  <span class="n">pair</span><span class="o">=</span><span class="nb">max</span><span class="p">(</span><span class="n">stats</span><span class="p">,</span><span class="n">key</span><span class="o">=</span><span class="n">stats</span><span class="p">.</span><span class="n">get</span><span class="p">)</span>
  <span class="n">ids</span><span class="o">=</span><span class="p">[</span><span class="bp">self</span><span class="p">.</span><span class="n">_merge</span><span class="p">(</span><span class="n">ch_ids</span><span class="p">,</span><span class="n">pair</span><span class="p">,</span><span class="n">idx</span><span class="p">)</span> <span class="k">for</span> <span class="n">ch_ids</span> <span class="ow">in</span> <span class="n">ids</span><span class="p">]</span>
  <span class="bp">self</span><span class="p">.</span><span class="n">merges</span><span class="p">[</span><span class="n">pair</span><span class="p">]</span><span class="o">=</span><span class="n">idx</span>
  <span class="n">first</span><span class="p">,</span><span class="n">second</span><span class="o">=</span><span class="n">pair</span>
  <span class="bp">self</span><span class="p">.</span><span class="n">vocab</span><span class="p">[</span><span class="n">idx</span><span class="p">]</span><span class="o">=</span><span class="bp">self</span><span class="p">.</span><span class="n">vocab</span><span class="p">[</span><span class="n">first</span><span class="p">]</span><span class="o">+</span><span class="bp">self</span><span class="p">.</span><span class="n">vocab</span><span class="p">[</span><span class="n">second</span><span class="p">]</span>
</pre></td></tr></tbody></table></code></pre></div></div>
<h2 id="sentencepiece">sentencepiece</h2>
<p>sentencepiece是Google提出的一个tokenizer库，它是作用在码点上的。要理解码点，先要分清楚utf-8和unicode的区别。简单来说，unicode是一个字符集，它定义了每个字符的唯一编码号，这个编码号被称为码点，比如中文里面的<code class="language-plaintext highlighter-rouge">我</code>的码点就是<code class="language-plaintext highlighter-rouge">U+6211</code>，而utf-8是一个编码方案，计算机只能存储0和1，为了将码点存储下来，使用utf-8编码方案，<code class="language-plaintext highlighter-rouge">U+6211</code>的utf-8编码就是<code class="language-plaintext highlighter-rouge">0xE6 0x88 0x91</code>这3个字节，它将每个码点映射为一个或多个字节序列。sentencepiece和tiktoken的区别简而言之，就是它们使用BPE的层级不同，前者在unicode的码点下进行合并，而后者在更贴近计算机底层的utf-8层级使用。一个更贴近人类语言，一个更贴近计算机语言。它们面向的场景不同，tiktoken面向通用文本，跨语言，因为所有语言底层逻辑都是用二进制存储，也就是可以用utf-8表示。而sentencepiece面向多语言精细化分词，字符是人类语言的基本表意单元，从字符出发合并更符合语言的语义结构。</p>
<h2 id="vocab_size">vocab_size</h2>
<p>在transformer中，vocab_size是指模型中使用的词汇表大小，它包括了所有可能的输入和输出符号。它会出现在两个地方，一个是token_embedding_table(shape=(vocab_size, d_model))，这个层会将每个token映射为一个d_model维的向量，所以vocab_size就是token_embedding_table的第一维。除此之外还有一个是在transformer的末端LM_head层，会将d_model维的向量映射为vocab_size维的向量，所以LM_head的输出维度就是vocab_size。就是我们会为每一个token(vocab_size大小)在每一个时间步上生成一个预测概率，随着我们有着越来越多的token，就需要预测更多的概率，就是在最后一个线性层上进行更多的点积运算。
所以token_embedding_table会随着vocab_size的增加而增加，LM_head线性层会随着vocab_size的增加而增加，会有更多的计算；更多的token意味着有更多的参数，可能担心很多参数没有得到充分的训练，因为引入更多的token，就是均摊了其它token在数据集上出现的频率，总体上所有token都会出现更少的示例占比，所以token的频率降低可能意味着它们在前向后向传播过程中的机会不多；此外，更多的token意味着对数据集的压缩更大，在适当的情况下我们可以使用较少的token来表达更多的文本(比如原来用80个token来表示一段话，现在只需要用50个token)，但是如果vocab_size过多，也可能会导致一大段话被压缩成一个token，这样模型在思考处理一定量的字符时，时间就不会那么充裕，因为过多的信息被过多压缩了。</p>

<h1 id="part-9-reproduce-gpt-2124m">Part 9-Reproduce GPT-2(124M)</h1>
<p>考虑到项目代码比较长，在这里详细解释各个部分，不如直接深入源码。在这里只是提炼其中的核心并分析，更细节的部分可以直接参考源码，源码中会有详细注释。</p>
<h2 id="最初的样子">最初的样子</h2>
<p>项目最初有CausalSelfAttention、MLP、Block、GPTConfig、GPT一共5个类型。</p>
<ol>
  <li>CausalSelfAttention
key,query和value以一行代码的批量计算紧凑形式组织在了一起，使用register_buffer，表示生成的下三角矩阵张量是缓冲器即非可学习参数，会随模型保存，但不参与梯度更新。在前向传播时，会调整各个维度的位置，只需要记住，参与计算的最后两个维度一定是[T，head_size]。这里为了加速注意力计算，使用了pytorch提供的接口函数<code class="language-plaintext highlighter-rouge">F.scaled_dot_product_attention</code>。</li>
  <li>MLP
提供的是注意力计算后的FFN步骤。不过这里使用了gelu激活函数，且使用<code class="language-plaintext highlighter-rouge">tanh</code>近似GELU加快训练</li>
  <li>Blcok
将前两个模块组装到一起，由于是自然语言，所以这里使用的是层归一化。</li>
  <li>GPTConfig
这个类使用python的dataclass进行注册，方便参数的管理</li>
  <li>GPT
这个类从名称上就可以看出，是GPT模型的实现。其中重要的是其可以加载OpenAI官方GPT-2预训练好的权重参数。加载预训练权重时，需要获取二者的参数字典，即自定义实现的模型的参数字典和OpenAI官方GPT-2模型的参数字典。此外，OpenAI使用Conv1D,它的输入和输出的维度与普通的Linear层是相反的，所以需要转置后才能copy到自定义的模型中。
    <h2 id="实现forwardautoregressive">实现forward,autoregressive</h2>
    <p>前向传播分为三个部分，将输入idx(shape=(B,T))进行token_embedding，得到(shape=(B,T,n_embd))和位置编码(shape=(T,n_embd))，将这两个张量相加。以block为组计算transformer结构，最后使用层归一化和线性层得到每一个待预测的token的logits(shape=(B,T,vocab_size))。
实现自回归的生成器，首先使用tiktoken的gpt-2的tokenizer对输入的文本进行编码，之后就可以传入模型得到logits,套路和之前的第7部分的生成器内容是差不多的，不过这里是从前50最大概率中选取一个token作为下一个token。</p>
    <h2 id="实现损失计算">实现损失计算</h2>
    <p>在forward中添加了损失计算部分，如果传入的target不为空，则计算损失。</p>
    <h2 id="实现一个数据加载器">实现一个数据加载器</h2>
    <p>要训练的数据从哪里来呢？怎么组织好直接可以传入模型进行训练呢？这就需要一个数据加载器帮助我们完成这些工作。DataLoader不仅会加载数据到内存中缓存好，还可以分批次扔出数据用于当前批次数据作为输入进入模型，进行训练，它通过一个<code class="language-plaintext highlighter-rouge">current_position</code>记录当前应该是第几个批次的数据，之后通过<code class="language-plaintext highlighter-rouge">next_batch</code>去获取对应位置的数据然后传递给模型进行训练。不过该数据加载中存在一个未解决的小bug,就是最后一段数据不足以支撑B<em>T批次大小时，会重置当前位置为0，意味着末尾有一段数据永远不会用于训练。保持这样处理的好处是，每个批次的数据都是固定的，内部的统计量不会改变，如果采用循环读取操作（当末尾的数据填充不满B</em>T时，使用头部的一段数据进行填充）会导致每个epoch训练时，批次的数据其统计量变化，可能（我猜测）会影响模型训练效果，此外，不利用末尾这一小段数据，在整个样本比例中其实占比不大，影响可以忽略。
  之后简单地使用一个优化器，训练模型，更新权重。</p>
    <h2 id="权重共享">权重共享</h2>
    <p>权重共享有两个问题，一个是理解为什么要这样做？第二个就是分清权重矩阵的存储逻辑和运算逻辑。
先来看第一个问题，为什么要权重共享？</p>
  </li>
  <li>降低参数量，这个是最直接的收益。参数的收益量是vocab_size*num_embd。</li>
  <li>符合自回归预测的任务逻辑
  语言模型的核心任务是子回归预测，给定前序token，预测下一个token。过程的本质是‘嵌入-编码-解码’的闭环过程，权重共享让这个闭环更合理。如果不共享权重，lm_head会学习另外一套逆映射，lm_head的本质是在学习一套wte映射的逆操作，将经过嵌入并编码后的结果映射回token。使用共享权重，就像是用同一把钥匙锁门和开门。</li>
  <li>词嵌入层wte的权重矩阵中，每一行对应一个token的词嵌入向量，行与行之间的距离(比如欧式距离)体现了token的语义相似度，当lm_head共享这套权重时，本质上就是<code class="language-plaintext highlighter-rouge">每个token的预测得分=隐藏向量*对应token的嵌入向量</code>，点积越大，说明二者的语义越匹配，预测概率越高，这套共享机制让预测过程直接利用了嵌入层学到的语义信息，确保语义相似的token在预测时会得到更高的关联得分，让模型的预测更符合语义逻辑。本质上是因为我们对嵌入层的嵌入向量之间的相似性解释，在预测时也希望能够利用上这种相似性，从而体现语义上的连贯性。</li>
  <li>此外，无偏设计（没有偏置）是保证了不会破坏语义上的对称。
权重矩阵的运算逻辑和存储逻辑是相反的
  pytorch中，变量x和权重矩阵w相乘时，其实是<code class="language-plaintext highlighter-rouge">x@W.T</code>，也就是说构建权重矩阵时我们是按照运算的逻辑构建的即权重矩阵的shape=[input,output]，而存储逻辑的形状其实是[output,input]，所以对于wte权重矩阵，因为在由token转到词嵌入空间中我们没有使用投影运算（矩阵相乘），而是直接使用索引，索引到当前token整数对应的wte的行，取出该行的词嵌入向量，所以wte通过<code class="language-plaintext highlighter-rouge">nn.Embedding</code>构建，而lm_head需要对编码后的向量进行投影，所以需要使用矩阵运算，构建时使用<code class="language-plaintext highlighter-rouge">nn.Linear</code>，这样，它构建时需要使用[input,output]这样的形式，但是底层存储时的形状会是[output,input]，转置一下就符合矩阵运算需要的形状了。
    <h2 id="初始化权重">初始化权重</h2>
    <p>为了最真实地贯彻GPT-2的路线，初始化的设置也选择了最贴近GPT-2的初始化设置。</p>
    <h2 id="控制残差流动引起的方差变化">控制残差流动引起的方差变化</h2>
    <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><table class="rouge-table"><tbody><tr><td class="rouge-gutter gl"><pre class="lineno">1
2
3
</pre></td><td class="rouge-code"><pre><span class="n">x</span><span class="o">=</span><span class="n">torch</span><span class="p">.</span><span class="n">zero</span><span class="p">(</span><span class="mi">768</span><span class="p">)</span>
<span class="k">for</span> <span class="n">i</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="mi">100</span><span class="p">):</span>
  <span class="n">x</span><span class="o">+=</span><span class="n">torch</span><span class="p">.</span><span class="n">randn</span><span class="p">(</span><span class="mi">768</span><span class="p">)</span>
</pre></td></tr></tbody></table></code></pre></div>    </div>
    <p>通过以上的例子，模拟了一个残差流不断累加的过程，其方差会不断增加。为了控制方差，维持前后方差的一致性，需要对残差流动进行归一化处理。所以在每个transformer block中，有两次残差连接，所以针对残差连接处的计算，在初始化投影矩阵权重时，需要除以根号下的两倍总的transformer block数。
不过，这里有一个问题，就是这种方式确保的是最终输出的变量，其方差是不变的，但是在每个transformer block中，其方差都缩小了。就是这样一个意思，原本我们为了确保初始化时，矩阵乘法后，输出的方差是基本不变的，所以除以了根号下的输入维度数，比如768，这样每一层的输出其方差都维持在这个位置。但是呢，由于引入了残差连接，所以确保在最终输出的方差不变的情况下，又对残差处进一步除以了一个值，这就导致每一层其方差变小了。不再是原来的输出后其方差都维持在这个位置的说法了。</p>
    <h2 id="启用tf32训练">启用TF32训练</h2>
    <p>TF32是GPU内部的一种计算方式，默认情况下使用FP32，当激活了TF32后，在GPU内部计算矩阵乘法时(其它运算依然使用FP32)，会使用TF32进行计算，从而提高计算效率。输入输出依然是FP32，只是在计算时采用了TF32。矩阵乘法用 TF32 加速（FP32 输入→TF32 计算→FP32 输出），其他运算仍为 FP32。</p>
    <h2 id="启用bf16">启用BF16</h2>
    <p>有必要简单地区分下FP32、TF32、FP16和BF16，FP32有8位指数位，用来表示范围，23位精度位用来确保精度，TF32相对于FP32，在精度位上只使用了10位。FP16，在精度位上和TF32保持一致，都使用10位，但是它的指数位只有5位，而BF16，它的指数为和TF32一样，有8位，但是精度位进一步缩减，变为了7位。它的最佳实践文档可以参考<a href="https://docs.pytorch.org/tutorials/recipes/recipes/amp_recipe.html">AMP</a>。基本上来说，就是在中间变量，将FP32转为BF16，进行训练和反向传播，梯度更新时再转回FP32确保精度。所以结合TF32(仅仅针对矩阵乘法)，在没有启用BF16时，它的输入输出是FP32，仅仅是中间计算时用了TF32加速，现在输入输出变成了BF16。所以BF16是整体上将FP32迁移成了BF16，不仅仅包括矩阵乘法，而是其所有的中间变量的临时存储都是BF16。总结起来就是核心是计算用BF16，参数存储用FP32。所以这就是为什么叫做混合精度的原因，一些东西在pytroch中仍然保持FP32(权重矩阵)，而一些东西精度降低了，成了BF16(激活值、中间计算的临时变量等等)。但是，又来了，又是但是，仔细查看上面的最佳实践文档可以知道，其实也并不是所有层都转成了BF16，不过总而言之，言而总之，启用BF16结合TF32，确实加速了我们的训练，节省了内存。</p>
    <h2 id="启用torchcompile">启用torch.compile</h2>
    <p>torch.compile对于加速的作用主要来自于减少python的开销和GPU读取次数。减少python的开销，简单地比喻就是python解释器就像厨师做饭，按照菜谱，要一步一步做，而自动化厨房只需要把原材料给它，就会自动运行。对于减少python的开销，torch.compile可以看到你所要操作的整个流程，而python解释器是逐行运行，它并不知道接下来会发生什么。torch.compile不会以一种<code class="language-plaintext highlighter-rouge">eager</code>模式运行，会优化运行的过程。它会首先移除python解释器在前向传播中的作用，将整个神经网络编译成不涉及python解释器的单一对象，然后直接运行。而对于GPU读取次数，举个简单例子就是说进行各种乘除法运算时，如果针对的同一个变量(该变量需要经过多种复合运算得到)，在没有使用torch.compile时，GPU会一步一步地在内核中运算-存储到GPU内存中-读取GPU内存中的内容-再次运算，而使用了torch.compile后，它就像知道了所有这些操作的最终结果都是为了得到那一个变量，就会直接在内核中进行连续的操作运算，而不用反复地读取存储了，这就是GPU读取次数的意思。</p>
    <h2 id="转换到flash-attentino">转换到Flash Attentino</h2>
    <p>简而言之，pytorch内部实现了这个机制，使得计算注意力分数的速度有了提升。</p>
    <h2 id="fit-nice-number">fit nice number</h2>
    <p>计算机科学中，由于二进制的原因，许多对于数字的设定都偏爱2的次幂，这些神奇的数字。所以对于一些设定的参数，比如批次、时间步长等等，优化成2的倍数，能发挥更好的计算机性能。</p>
    <h2 id="梯度裁剪">梯度裁剪</h2>
    <p>就是防止出现梯度爆炸、梯度更新反复横跳等现象。在反向传播和梯度更新之间执行，一般常用的是收集所有参数的梯度计算其L2范数，之后设定阈值进行裁剪。</p>
    <h2 id="余弦衰减学习率调度器">余弦衰减学习率调度器</h2>
    <p>需要注意，学习率调度器和优化器是两个不同的东西，一个是针对学习率，一个是在学习率固定下来后怎么去计算梯度更新的策略。在优化器的参数中，去更新刷新后的学习率。余弦学习率调度器本质上就是根据当前的训练epoch次数，去计算应该使用什么大小的学习率。
<img src="/assets/images/2025/cos-learning.png" alt="cos-learning" /></p>
    <h2 id="权重衰减和梯度累积">权重衰减和梯度累积</h2>
    <p>在GPT-2，GPT-3训练过程中采用了变化的batch size大小的策略进行训练。基于这样的观察和解释，在模型训练的初期，基本上是在学习忽略那些不常出现在训练集中的token，学习非常简单的偏差和类似的东西，每一个样本都在告诉模型，使用这些token，不使用那些Token,来自每一个样本的梯度实际上是高度相关的，在优化的初始阶段，它们看起来都大致相同，因为它们都在告诉模型这些token出现，那些token不出现。所以在训练初期，没有必要使用很大的batch size。只有当跨越过初期阶段，使用大批量样本才会展现出统计上的意义，去学习更深层次的东西，比如“语境歧义”（如 “苹果” 是水果还是公司）。怎么理解这段话呢？
简而言之，打个比方把模型训练比作 “老师教学生学语文”，batch size 比作 “一次布置的作业量”：
初期（学拼音、常用字）：学生的核心任务是 “记住常用字怎么写、怎么读”—— 所有作业（样本）都在重复 “听写常用字”，学生的错误（梯度）都集中在 “生僻字不会写”“常用字写错笔画” 上（梯度高度相关）。这时候布置 10 道题（小 batch）和 100 道题（大 batch）的效果一样 —— 学生都是在纠正相同的错误，100 道题只会让学生更累（浪费时间），不会更快掌握。
后期（学阅读理解、写作）：学生的任务是 “理解语境、掌握多义词、组织逻辑”—— 作业题（样本）五花八门：有的考 “‘打’在‘打球’和‘打电话’中的不同含义”，有的考 “议论文的论点论据”，有的考 “散文的情感表达”（梯度多样性高）。这时候布置 100 道题（大 batch）比 10 道题（小 batch）效果好 —— 学生能接触更多场景，避免 “只懂一道题，换题就错”（降低梯度噪声），学到的规律更全面（统计上的通用逻辑）。
所以，训练初期的核心是 “学简单规律”，梯度同质化，小 batch 足够用，且效率更高；
训练后期的核心是 “学复杂规律”，梯度异质化，大 batch 能覆盖更全的样本分布，让模型学到统计上的通用规律；
但在本项目的实现中，跳过了这个步骤，因为它会把问题复杂化，比如怎么动态地处理batch size带来的数据变化。
真实地遵循GPT-3的训练策略，其中提到了使用0.1的权重衰减，就是L2正则化。但是不是对所有的参数都进行正则化，主要是针对矩阵乘法部分的参数，比如线性层的权重矩阵。为此需要构建一个<code class="language-plaintext highlighter-rouge">configure_optimizer</code>函数，来确哪些需要权重衰减，哪些不需要。其中优化器使用到了<code class="language-plaintext highlighter-rouge">fused</code>这个选项，就是把许多计算需要用到许多核的情况，融合成了一个核，简而言之就是加快计算速度。
  在有限的资源，比如一个GPU中，如何使用0.5M大小的总token数(=B*T)进行训练呢？如果直接指定对应的batch size，那么GPU的内存会溢出。这个时候就要使用梯度累积的策略了。它允许我们采用串行的方式，模拟任何批次大小的数据。代价就是运行时间变成，处理多个序列然后把这些梯度加起来。所以基本策略就是使用一个小batch size，多次前向传播-反向传播，但是不更新梯度，直到重复<code class="language-plaintext highlighter-rouge">B/min_batch_size</code>次，再更新梯度。需要注意的是，我们使用了loss_accum来记录最终的0.5M这个batch size大小下的损失，它是被detach掉的。</p>
    <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><table class="rouge-table"><tbody><tr><td class="rouge-gutter gl"><pre class="lineno">1
2
3
4
5
6
7
8
</pre></td><td class="rouge-code"><pre><span class="k">for</span> <span class="n">micro_step</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="n">grad_accum_steps</span><span class="p">):</span>
     <span class="n">x</span><span class="p">,</span><span class="n">y</span><span class="o">=</span><span class="n">train_loader</span><span class="p">.</span><span class="n">next_batch</span><span class="p">()</span>
     <span class="n">x</span><span class="p">,</span><span class="n">y</span><span class="o">=</span><span class="n">x</span><span class="p">.</span><span class="n">to</span><span class="p">(</span><span class="n">device</span><span class="p">),</span><span class="n">y</span><span class="p">.</span><span class="n">to</span><span class="p">(</span><span class="n">device</span><span class="p">)</span>
     <span class="k">with</span> <span class="n">torch</span><span class="p">.</span><span class="n">autocast</span><span class="p">(</span><span class="n">device_type</span><span class="o">=</span><span class="n">device</span><span class="p">,</span><span class="n">dtype</span><span class="o">=</span><span class="n">torch</span><span class="p">.</span><span class="n">bfloat16</span><span class="p">):</span>
         <span class="n">logits</span><span class="p">,</span><span class="n">loss</span><span class="o">=</span><span class="n">model</span><span class="p">(</span><span class="n">x</span><span class="p">,</span><span class="n">y</span><span class="p">)</span>
     <span class="n">loss</span><span class="o">=</span><span class="n">loss</span><span class="o">/</span><span class="n">grad_accum_steps</span>
     <span class="n">loss_accum</span><span class="o">+=</span><span class="n">loss</span><span class="p">.</span><span class="n">detach</span><span class="p">()</span>
     <span class="n">loss</span><span class="p">.</span><span class="n">backward</span><span class="p">()</span>
</pre></td></tr></tbody></table></code></pre></div>    </div>
    <h2 id="distributeddataparalleddp">Distributeddataparalle(DDP)</h2>
    <p>怎么利用多GPU进行训练。这里的难点在于我们需要想像有8个并行运行的程序在运行相同的代码，它们的区别就是<code class="language-plaintext highlighter-rouge">ddp_rank</code>。因为我们有了多GPU，原本代码中的一些参数值就得考虑到平均后的正确数值是多少了，比如原本我们在单个GPU中使用mini_batch，需要前向-反向传播<code class="language-plaintext highlighter-rouge">B/mini_batch</code>次才能进行一次梯度更新，现在考虑到使用多GPU，这个前向-反向传播次数进一步被平均了，所以需要<code class="language-plaintext highlighter-rouge">B/(mini_batch*num_GPU)次就可以了</code>。我们需要让<code class="language-plaintext highlighter-rouge">DataLoader</code>根据不同的GPU加载属于自己的那份数据，而不是每个GPU都加载同一段的数据。</p>
  </li>
</ol>

<p>目前总结下代码中的一个容易误导的地方，就是现在是按照总的优化器更新次数来训练模型，而不是按照epoch次数训练模型。按照传统epoch次数的理解，举例100个epoch，就需要每个epoch都要完整地遍历一次训练数据集。而现在，采用的是总的优化器更新次数，每次都遍历一定量的token数，比如设定总的优化器更新次数为max_steps，那么它等同于<code class="language-plaintext highlighter-rouge">max_steps/(total_token/batch_size_token)</code>个epoch。</p>]]></content><author><name></name></author><category term="人工智能" /><summary type="html"><![CDATA[目录 这篇文章记录了从手动实现自动微分到构建 GPT 训练逻辑的完整学习路径，重点包括微分、二元模型、N-gram、BatchNorm、Transformer 与 Tokenization 的核心思想与代码实现。 Part 1-Micrograd 要自动计算梯度，需要一个类来实现，这个类能够表示标量和张量，需要计算变量的梯度时，能够调用这个类（实例）的相关计算梯度的方法。可以简单地创建一个Value类: 1 2 3 4 5 6 7 class Value: def __init__(self, data, _children=(), _op=''): self.data = data self.grad = 0 self._backward = lambda: None self._prev = set(_children) self._op = _op 这样，梯度和求导函数都包括在了这个类里面，还记录了它依赖于哪些变量即由哪些变量和操作计算得来。以tanh为例： 1 2 3 4 5 6 7 8 9 10 11 12 class Value: ... def tanh(self): x = self.data t = (math.exp(2*x) - 1)/(math.exp(2*x) + 1) out = Value(t, (self, ), 'tanh') def _backward(): self.grad += (1 - t**2) * out.grad out._backward = _backward return out 在这里，out为计算的tanh值，它是一个Value对象并最终返回out，例如a=Value(1),b=a.tanh()，在计算tanh时，也会创建一个_backward()函数，它被赋值给了新创建的out变量，它记录了当前变量需要计算梯度时的公式，即这样的逻辑还是以a,b为例，当需要从b反向传播计算a的梯度时，将b的梯度设为1，调用b的反向传播函数，它是用以计算a的梯度。反向传播函数不是同级对等的关系，即a的梯度需要调用b的反向传播函数计算得到，而不是a的反向传播函数，如果a的反向传播函数不为None，那么它则是计算在a之前的变量的梯度，而不是a本身的梯度，以y=x**2为例，要计算x关于y的导数，是通过y对x求导，而不是对x本身求导，y对其自身的导数为1，这样当需要调用反向传播函数时，就可以计算得到当前变量的梯度值，这里的梯度是+=而不是直接赋值，这是因为同一变量可能会在多个不同的地方使用到，正确的计算方式就是将这些不同地方的但是是同一变量的梯度进行累加。这也解释了在PyTorch中，每一轮训练开始在进行反向传播之前需要将变量的梯度清0,为的就是不让上一轮的梯度继续与本轮的梯度累积。 当需要进行反向传播时，我们需要一个搜索算法来将所有的与最终作为反向传播开始的变量相关的变量都找出来，例如a=Value(1),b=a.tanh(),c=b**2，当需要从c反向传播计算a的梯度时，需要先将c的梯度设为1，然后调用c的反向传播函数，它会计算b的梯度，而a的梯度需要调用b的反向传播函数计算得到，所以需要一个搜索算法来将所有的与c相关的变量都找出来，即a,b,c。 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 class Value: ... def backward(self): topo=[] visited=set() def build_topo(v): if v not in visited: visited.append(v) for child in v._prev: build_topo(child) topo.append(v) build_topo(self) self.grad=1.0 for v in reversed(topo): v._backward() 这是一个拓扑排序的过程，它将所有的与最终作为反向传播开始的变量相关的变量都找出来，然后从后往前调用每个变量的反向传播函数，这样就可以计算得到所有变量的梯度值。使用的是深度优先算法，这里用了两个变量来记录拓扑排序和已访问变量节点，理论上来说可以只使用topo一个变量，这里用了visited|set是因为集合查找比列表查找更快，列表的查找时间复杂度是O(log(n)),而集合则是O(1)。 有了能够进行自动微分反向传播的类，我们就可以按照Neuron-&gt;Layer-&gt;MLP的顺序构建一个简单的多层感知机，并进行训练。它不支持矩阵并行计算这样的复杂操作，底层实现逻辑还是通过循环来获取单一的元素进行运算的。 Part 2-Bigram 从统计模型角度来看二元语法模型，字符级别的预测例如句子我喜欢你,我们根据单个字符进行预测，首先在句子的首位添加起止符，比如可以使用.我喜欢你.这样的表达，.表示句子的开始和结束，起止符号可以相同，也可以采用不同的表示，这里采样相同的表示。将字符转换成数字显然更加符合计算机的习惯，也是更方便计算。所以对于’abc…z’这样的字符，再加上.这样的起止符，在这个简单的模型中一共有27个独特的字符，设计两个字典方便从字符到数字之间的互相转换。可以统计在使用的数据集中，二元字符对出现的频率，这应该是一个[27,27]的二维矩阵，因为每个字符都可以作为二元模型中的第一个字符，也可以作为二元模型中的第二个字符，总之，我们可以经过统计得到这样的以每个字符开头的二元对的频率，并将它友好地显示出来。从可视化中可以明显看出，一些字符对的出现频率很高，一些则一次都没有，这即揭示了一些统计规律（也可能是当前数据集不够具有代表性），但不管怎样，我们可以根据当前的数据集得到符合当前数据集的统计规律，比如wq字符对在当前数据集中出现的频率为0。接着，我们可以独立地对每一行求和，并计算各自的概率即softmax运算。这样就得到了基于当前字符作为第一个字符，其二元模型中第二个字符出现的各个近似概率（总之如果数据集足够具有代表性，也即数据集足够大，大数定律告诉我们当然可以近似地认为这就是这个字符对的概率）。这就是简单的二元语法的统计模型，接下来我们就可以进行采样了，生成一系列预测的合理的人名。 简单的统计模型，换个角度我们也可以使用神经网络的方式来实现，因为当到达三元、四元、n元模型时，简单的统计模型就显得捉襟见肘了，单单想一想各种组合的可能性，就是呈指数级的爆炸增长。而神经网络可以在简单和复杂中都能有效地进行表示（建模）和计算。在概率与数理统计中，我们学习过最大似然估计，简而言之，最大似然估计构建了一个这样的模型，这个模型符合使得当前数据集中n元模型出现的概率达到最大–即利用已知的样本结果信息，反推最具有可能（最大概率）导致这些样本结果出现的模型参数值，模型的复杂问题一切都包含在了什么模型最能导致当前的样本结果了。∏𝑖=1𝑛𝑝𝜃^(𝑥𝑖)，这就是最大似然函数。对于连续型随机变量，有相同的结论。。深度学习中我们一般都是最小化损失，换个思路最大化似然函数-&gt;最小化负的最大似然函数-&gt;最小化数据集中的所有数据的平均负的最大似然函数，按照惯例为了计算方便，我们将其对数化，这样乘法就变成了加法log(a*b*c)=loga+logb+logc，于是这样的逻辑思路就是利用最大似然函数的思想我们构建了一个对模型损失的评估，通过反向传播算法我们不断根据当前的损失去更新模型的参数。不过当遇到从没有出现过的字符对时，这个log值就为-inf，可以给整体各自加上一个较小的值，比如1，这样最小频率为1，这个方式称为模型平滑。为了训练神经网络，需要构建训练样本标签对，例如.emma.这样的单词，划分成训练样本对为xtr=[.,e,m,m]和ytr=[e,m,m,a,.]，在输入网络前还需要将其转换成数字的形式[0,5,13,13,1]和[5,13,13,1,0]。转换成数字的形式之后，可以看到时间序列xtr中当前的词去预测下一个词，在下一个时间步中被预测的下一个词又作为当前的词继续去预测它的下一个词，因为要使用神经网络，y=X@W+b这样的形式，可以把xtr-ytr变为二维即shape=[time_step,1]的形式，但是这样矩阵相乘时，就要求W的形状是shape=[1,hidden_representation]的大小，即[[0],[5],[13],[13],[1]]@[[n1,...,n_h]]，简单的矩阵乘法知识就可以知道，它的结果仅仅是C=A@B中，C的每一行的结果都是A的每一行乘以B的每一列。而使用one_hot编码，将x_tr-y_tr根据其字典大小的长度进行编码，可以让x_tr-y_tr的形状变成shape=[time_step,27]，这样W的形状就可以是shape=[27,hidden_representation]，可以进行更复杂的线性组合，获得更复杂的表示，独热码只是encodeing的其中一种最简单的方式，本质就是给当前的字符一个在高维空间上的表示方法，还可以通过一些词袋模型计算更复杂的编码表示方式比如word2world。在这里，简单将隐藏表示维度设为27，那么logits=xenc@W的结果被赋予了logits的称呼，这也即与统计模型中的出现次数对应一个级别，但是这里的次数有正有负，因为W是随机初始化的。对counts=logits.exp()将其全部转为正数，然后计算其softmax结果，就得到了统计模型中每一行相同的频率（概率）表示形式，有了概率表示就可以计算最大似然函数，在未训练模型时，我们还可以估计平均最大似然函数正常的取值范围，因为在没有训练之前，没有任何理由可以认为什么组合的出现概率最高，所以它们出现的概率是相等的，平等看待每一个组合，这样就计算出一个最大似然函数的值，可以用来评估初始化权重矩阵是否合理。从负的最大似然函数开始进行反向传播，使用梯度下降算法训练网络。最后同样地可以得到与统计模型相似地结果。 W权重矩阵和二元模型概率分布及其正则化。W的形状是shape=[27,27]，而我们使用one_hot编码，根据矩阵相乘的结果，每一行即每一个时间步的输入，都被映射到了一个27维的向量空间中，这个向量空间中的每个维度，都对应了一个字符，而这个字符的出现概率，就是这个向量空间中这个维度的数值。所以W的每一行，就是对应了一个字符作为第一个字符，其二元模型中第二个字符出现的概率分布。通过矩阵相乘获取了与统计模型相同的表示结果。而在统计模型中为了避免log取值出现负无穷，我们通过为最小值增加一点值作为模型平滑的结果，这个量加得越大，最终计算的概率越平滑，各个组合出现的概率越相近，同样地在W权重矩阵中，如果W的各个元素其初始值都被设置为0，那么取得的结果就是一个均值，各个组合的概率相同。这就引出了正则化，通过L2正则化，我们在实现梯度下降更新W参数的同时，也在尽可能让W变小，我们可以实现相同的平滑效果，这就是正则化和模型平滑的一个共通解释，很有趣。 Part 3-N-gram模型 现在，可以利用神经网络模型，将二元模型扩展到N元模型，以N=3为例。可以构建训练验证测试数据集。现在我们设置embedding的大小为10,即C.shape=[27,10],它代表了每个字符被映射到了一个10维的向量空间中。通过索引C[X],可以获取到输入序列X中每个时间步下每个字符对应的10维向量表示，而C[X].shape=[N,3,10]，接下来就是相似地构建W1.shape=[30,200],W2.shape=[200,27]权重矩阵,同样地训练模型。在计算损失函数时，我们的操作一般是将取得的logits结果进行取指数(e)运算，全部变为正值，再使用softmax函数计算各自的概率，最后取log值求平均计算最大似然函数，现在，这些操作全部可以简化为F.cross_entropy(logits,Y)这一个表达式中。这就是N元模型，相比于二元模型没有任何特别的，都可以使用神经网络模型一步一步构建得到,但是一些地方也有变化，比如网络深度，隐藏层数，网络宽度，embedding大小，学习率变化，这些都是构建网络过程中的超参数。 此外，无论如何构建多么复杂的网络，网络的损失都不会变为0，一个较为直观的解释是，每个N元模型的开始都是以&lt;S&gt;作为第一个字符,所以通过起止符去预测下一个元素时，我们需要这种变化，如果每一次通过起止符去预测下一个元素的值都是不变的，这本来就自相矛盾。比如&lt;S&gt;你好&lt;E&gt;，&lt;S&gt;我喜欢你&lt;E&gt;。 Part 4-BatchNorm 在早期的深度学习实践中，对于权重矩阵的初始化需要十分精准，避免在张量传播过程中，其值域范围的均值和方差出现极端的不稳定情况。 1 2 3 4 5 6 g=torch.Generator().manual_seed(2147483647) C=torch.randn((vocab_size,n_embd),generator=g) W1=torch.randn((n_embd*block_size,n_hidden),generator=g)*(5/3)/((n_embd*block_size)**0.5) #*0.2 b1=torch.randn(n_hidden,generator=g)*0.01 W2=torch.randn((n_hidden,vocab_size),generator=g)*0.01 b2=torch.randn(vocab_size,generator=g)*0.01 在该代码示例中，假如初始化时没有乘以系数，那么在第一次前向传播过程中，各个层的所得的均值和方差就会出现极端的不稳定情况，这会导致梯度消失或梯度爆炸的问题，从而影响模型的训练效果。以W1为例，h=tanh(X@W1+b1),tanh函数的图像如下所示，它也是属于sigmoid函数簇 它的导数为1-h**2,链式规则为self.grad=(1-h**2)*out.grad从公式上可以看出，当h接近-1或1时，导数接近0，这会导致梯度消失的问题。而当h接近0时，导数接近1，它会传递梯度值。而如果没有合理的初始化W1权重矩阵，在第一次前向传播过程中，h的取值会非常大或非常小，这会导致梯度消失或梯度爆炸的问题。所以，在初始化W1权重矩阵时，需要乘以一个系数，比如(5/3)/((n_embd*block_size)**0.5)，这是根据tanh函数的性质推导出来的一个系数，它可以确保在第一次前向传播过程中，h的均值和方差不会出现较大的波动，从而避免梯度消失的问题。 当然我们可以凭借直觉和反复实现观察其内部的数值，来给定一个较为还不错的系数。而系统性的初始化方法，比如Xavier初始化，kaiming初始化，则在工程上系统性地对初始化进行了优化，避免了手动给定系数的过程，同时也确保了模型的训练效果。 当然这只是网络第一次训练中的第一个batch过程，为训练开了一个比较好的头。为了在整个训练过程中都保持较好的稳定状态，提出了一系列的归一化方法。比如batch normalization，layer normalization，instance normalization等。这些方法的基本思想都是在每个batch或每个样本中，对输入的特征进行归一化，从而避免梯度消失或梯度爆炸的问题。我们可以这样想象，有一个X=[x1,x2,…,xn]它有n个特征，其中每个特征的数值尺度是不一样的，范围从1-10000变化，那么进行梯度下降时，它们的更新速度或者说走的step的尺度也是不一样的，有的走得快，有的走得很慢。比较好的处理方法就是把这些特征的尺度都进行归一化处理，让它们都在同一个尺度下面进行训练。这就是归一化。 怎么在代码中手动实现一个batch normalization呢？原理很简单，我们可以在每一个batch中，计算当前batch的均值和方差，再对当前batch中的每个样本，减去均值后再除以方差，从而实现归一化。hpreact=bngain*(hpreact-hpreact.mean(0,keepdim=True))/hpreact.std(0,keepdim=True)+bnbias，为了让网络能够调整分布，使得一些神经元激活一些不敏感，引入bngain/gamma和beta/bnbias两个可学习的参数，这是因为我们不希望一直让网络强制保持标准的高斯分布，只是在前期没有任何知识的情况下无法假设，只能保持公平性，不让网络对任何一个对象有偏爱，但当随着网络训练的过程，网络能够偏爱一些胜过另一些，也就是神经元敏感与不敏感。 现在我们引入了bn，但这也引入了一个问题，我们的训练过程与数据出现了耦合。隐藏状态、激活值除了依赖于输入X、函数，还依赖随机选取的batch形成的组合数据，比如一些batch中的样本的某个特征的取值范围很大，而另一些样本的取值范围很小，这种随机抖动的作用，反而可以作为一种正则化，引入一点熵让模型难以过拟合。 不过，这也引入了一个问题，就是在预测时，如何在模型中已有的计算批量状态下的bn中进行适配？预测时我们输入的是一个样本，而不是一个batch，但是现在有一段代码是计算bn的。这里有两种实现方式，一种是固定训练集中的bn值，当训练完成后，重新计算整个训练集的bn值，当预测时使用这个固定的bn值。另一种实现方式则是采用动态更新的策略，动态评估bn值。省略了重新计算整个训练集的步骤。它的实现方式如下,基本思想就是记录一个running的状态，然后根据是否是训练过程来决定要不要计算当前批次的均值和方差，如果不是训练过程，则直接将记录的running状态的均值和方差赋值给实际要使用的均值和方差变量，再进行归一化处理。并且，如果是训练过程，就会在最后根据momentum这个动量的大小决定要更新到running状态上的分量，一般而言momentum取0.1代表取当前计算得到的批次的均值和方差的0.1，将其加入到running状态值中，这样就实现了根据每一个批次的均值和方差，对最终的均值和方差的更新。 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 class BatchNormld: def __init__(self,dim,eps=1e-5,momentum=0.1): self.eps=eps self.momentum=momentum self.training=True #parameters trained with backprop self.gamma=torch.ones(dim) self.beta=torch.zeros(dim) #buffers (trained with a running 'momentum update') self.running_mean=torch.zeros(dim) self.running_var=torch.ones(dim) def __call__(self,x): #calculate the forward pass if self.training: xmean=x.mean(0,keepdim=True) xvar=x.var(0,keepdim=True) else: xmean=self.running_mean xvar=self.running_var xhat=(x-xmean)/torch.sqrt(xvar+self.eps) self.out=self.gamma*xhat+self.beta #update the buffers if self.training: with torch.no_grad(): self.running_mean=(1-self.momentum)*self.running_mean+self.momentum*xmean self.running_var=(1-self.momentum)*self.running_var+self.momentum*xvar return self.out def parameters(self): return [self.gamma,self.beta] Part 5-Backprob 针对简单的多层感知机，其梯度反向传播较为简单，在计算时需要从矩阵的角度考察衡量，要注意是否需要沿着某一个维度/轴进行求和，因为在前向传播过程中会隐含地出现广播这个操作，所以很容易忽略掉。如果原来的变量是矩阵形式，那么其反向传播时的梯度也是矩阵形式。简而言之，变量正向和反向传播的形状都是不变的。 Part 6-Building a WaveNet 这里面基本是按照第三部分的内容，不过将一些混乱的结构进行了整理，将Embedding的过程单独整理成了一个类。添加了展平层，此外对bn层进行了扩展，支持不同的维度。 Part 7-Let’s build GPT: from scratch 通过一个简单的二元模型，该二元模型做的仅仅是将输入的token的整数索引映射到一个向量中，该向量的大小也是token的词典大小，表示的是从该token映射到下一个可能token的Logits值。它的参数大小是[vocab_size,vocab_size]，logits=self.token_embedding_table(idx)，举例该例子主要是为了实现一个简单的生成器，根据当前的输入生成下一个预测的输出。可以看到，我们只需要拿取计算的最后一个时间步的logits，然后通过softmax计算得到概率，根据该概率采样得到下一个预测的词的idx_next，然后将这个id_next添加到idx中，作为输入，在这个例子中由于是二元模型，所以只会用到前后两个词，但是作为通用的生成器，这里可以简单地修改，就可以利用规定的上下文大小比如8个词的上下文，去预测下一个词,比如我们可以更改block_size的大小，就可以只关注于最新的需要查看的上下文大小的词，使用这些最新的词去预测下一个词，而计算注意力这些额外的计算，都放在了forward函数中，如果我们注释掉这一行，那么并且现在forward函数中没有实施注意力相关的计算，就退回了最初的二元模型。此时在没有实施注意力代码的时候，你会怎么考虑控制上下文大小的注意力计算呢？我最初的想法很直接，注意力计算要考虑上下文大小，那就直接在注意力计算的过程中进行控制，但是这样又引入了新的问题，就是控制所选上下文的窗口的移动，比如现在的上下文窗口大小是8，现在我的输入有32个，那就要控制找到新的8个词的内容，这就会要求我们额外地在forward函数中添加额外的控制。但是我们不想因为这个改变通用的写法，所以换个思路，把这个窗口大小的控制逻辑放在了生成器中，直接截取idx的最新的上下文窗口大小的内容，这样，在forward函数中只需要按照输入计算注意力就可以了。 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 def forward(self, idx, targets=None): # idx and targets are both (B,T) tensor of integers logits = self.token_embedding_table(idx) # (B,T,C) if targets is None: loss = None else: B, T, C = logits.shape logits = logits.view(B*T, C) targets = targets.view(B*T) loss = F.cross_entropy(logits, targets) return logits, loss def generate(self, idx, max_new_tokens): # idx is (B, T) array of indices in the current context for _ in range(max_new_tokens): # get the predictions idx_cond=idx[:,-block_size:] logits, loss = self(idx_cond) # focus only on the last time step logits = logits[:, -1, :] # becomes (B, C) # apply softmax to get probabilities probs = F.softmax(logits, dim=-1) # (B, C) # sample from the distribution idx_next = torch.multinomial(probs, num_samples=1) # (B, 1) # append sampled index to the running sequence idx = torch.cat((idx, idx_next), dim=1) # (B, T+1) return idx 注意力的数学本质，就是求得加权后的新值，理解的意思就是W是权重矩阵，而X是待加权的变量，通过W@X，就可以根据W中的权重来重新调整X中每一个元素的值，该值参考了其它值，在该元素对应位置生成的新的加权后的元素值，其是由与该位置有关的原来的元素值，以及其它相关值通过加权运算得到的。 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 torch.manual_seed(1337) B,T,C = 4,8,32 # batch, time, channels x = torch.randn(B,T,C) # let's see a single Head perform self-attention head_size = 16 key = nn.Linear(C, head_size, bias=False) query = nn.Linear(C, head_size, bias=False) value = nn.Linear(C, head_size, bias=False) k = key(x) # (B, T, 16) q = query(x) # (B, T, 16) wei = q @ k.transpose(-2, -1) # (B, T, 16) @ (B, 16, T) ---&gt; (B, T, T) tril = torch.tril(torch.ones(T, T)) #wei = torch.zeros((T,T)) wei = wei.masked_fill(tril == 0, float('-inf')) wei = F.softmax(wei, dim=-1) v = value(x) out = wei @ v 在这个例子中，key和query都是从原x中通过矩阵投影而来，通过这两个投影计算得到wei即权重矩阵，同样地我们也再一次对x投影得到v值，那么v的加权后的新值，就是out=wei@v，这就是注意力的计算过程。不过此时的形状还是[B,T,head_size]，还要使用一个投影矩阵将其投影回原来的大小[B,T,C]，我们需要记住的是，经过注意力加权后，新的变量的形状应该与原来的变量形状保持一致，因为我们仅仅是做一个加权处理，如果最终改变了形状就不正确了。同样地需要注意，在生成式语言模型中，就是根据上下文去预测下一个词时，按照正常逻辑，我们应该只能看到截止到当前时间步及其以前时间步的内容，计算它们之间的注意力值，这里我们使用掩膜将上三角的值换成了-inf,这就符合了[生成]这个词所表达的含义。 需要注意的是， 在计算注意力时没有空间相对位置的关系，想想a=a_1*w_1+a_2*w_2+a_3*w_3等同于a=a_3*w_3+a_1*w_1+a_2*w_2，所以我们需要自己额外添加一个位置编码进去。 注意力可以看作一种交流机制，可以被视为有向图中的节点，这些节点相互观察，并通过来自所有指向它们的节点的加权和来聚合信息，其中权重取决于数据。 在batch维度跨越的样本之间是无法交流的，因为矩阵乘法只发生在最后两个维度。 在编码器中，只需要删除进行掩膜的那行代码，就从解码器变成了编码器，从历史的原因分析，最初transformer架构的提出是针对语言翻译的，所以使用编码器去了解翻译对象的全部信息，再使用解码器也就是对注意力权重矩阵增加了掩膜的部分，去生成翻译的目标语言。单独的解码器架构也就是自回归模型，常用于大语言模型。 自注意力，就是说除了query是来自于x外，key,value也都来自于x的投影变换，如果query，value来自于其它数据，就叫做cross-attention。 “缩放”注意力机制额外将wei除以1/√(头大小)。这样一来，当输入Q、K具有单位方差时，wei也会具有单位方差，并且Softmax会保持分散状态，不会过度饱和。 Part 8-Tokenization LLM的许多奇怪的现象都和tokenization有关。例如，大语言模型不能很好的拼写单词，在编写代码方面表现很差，这些都与tokenization有关。 在本例中针对字符级别，将其转为utf-8编码，然后转为整数。 1 2 3 4 print('你好'.encode('utf-8')) b'\xe4\xbd\xa0\xe5\xa5\xbd' print(list(map(int,'你好'.encode('utf-8')))) [228, 189, 160, 229, 165, 189] 这样我们便得到了用整数表示的一些列token。仅仅一个你好的token表示就是6个整数，每个整数都在0到255之间。而具有这样规律的整数在大量文本中会重复出现，这就是语义的表达。而大预言模型只是接收这些整数作为输入，然后根据其内部的参数，进行预测下一个token出现的概率。我们可以将[228, 189, 160, 229, 165, 189]这样的具有语义信息的整数进行融合，创造出一个新的整数来代表该序列，就是使用256来代表你好，重复这样的操作我们就可以得到扩充了的，压缩了信息的token库。但是token库的大小也不能无限大，例如使用10000这个整数来表达针对简单的多层感知机，其梯度反向传播较为简单这样一句话，已经压缩了许多信息，无法获取其内部的详细语义信息，所以token库的大小要适合就行，也就对应了我们需要融合的次数，每次我们融合都选择出现频率最高的一对整数，并将其赋予新的值，只要我们在合适的合并次数下停止，就能够得到较为合理的，既能合理高效表示语义信息，又不至于压缩得太厉害的程度。 1 2 3 4 5 def get_stats(ids): counts = {} for pair in zip(ids, ids[1:]): # Pythonic way to iterate consecutive elements counts[pair] = counts.get(pair, 0) + 1 return counts 我们使用get_stats函数，对输入得tokens进行统计，得到所有相邻的token对出现的次数并返回。 1 2 3 4 5 6 7 8 9 10 11 12 13 def merge(ids, pair, idx): # in the list of ints (ids), replace all consecutive occurences of pair with the new token idx newids = [] i = 0 while i &lt; len(ids): # if we are not at the very last position AND the pair matches, replace it if i &lt; len(ids) - 1 and ids[i] == pair[0] and ids[i+1] == pair[1]: newids.append(idx) i += 2 else: newids.append(ids[i]) i += 1 return newids 然后我们可以使用merge函数，融合指定的token pair，并赋值为新的token idx，返回的是融合后的token序列，原来的token pair被token idx替代了。所以，如果原token长度为n，其中有m个token pair，那么融合后token序列的长度为n-m。下面这段代码将整个融合的流程串联了起来，首先获取token pair对，然后对获取的token pair对进行排序，获得出现频率最高的pair，之后为这个pair赋值一个新的token整数值，来代替原来的pair，实现信息的高效、压缩表示，并用一个字典记录被替换的pair和它的新的整数token值的表示之间的映射关系，方便后续解码的时候反推回原来的表示。 1 2 3 4 5 6 7 8 9 10 11 12 vocab_size = 276 # the desired final vocabulary size num_merges = vocab_size - 256 ids = list(tokens) # copy so we don't destroy the original list merges = {} # (int, int) -&gt; int for i in range(num_merges): stats = get_stats(ids) pair = max(stats, key=stats.get) idx = 256 + i print(f"merging {pair} into a new token {idx}") ids = merge(ids, pair, idx) merges[pair] = idx 我们通过这个循环，融合了20个token pair，得到了一个新的token库，其中包含了256个原始token，以及20个新的token，每个新的token都是由两个原始token融合而来的。并且可以比较融合后的token长度和原始token长度得到一个压缩率，用来表示压缩的程度。Tokenizer是一个独立于LLM的模块，有其自己的训练集，使用上面的方法即Byte pair encoding(BPE)算法，它将在文本和一些列的tokens之间转换，输入到LLM中的只有转换后的tokens。 Decoding 现在我们已经实现了融合token pair的操作，就是我们已经根据训练集训练好了一个tokenizer，现在我们需要应用到实际中。但是怎么解码呢？即怎么从一段tokens中，恢复出原来的文本？很简单的直观的方法就是根据融合的链式替换，一步步将新的token idx替换为原来的token pair，直到恢复出原来的文本。不过使用字节级的形式，可以很取巧地快速解码，因为字节的表示可以相加，就像字符一样进行拼接。bytes([66])+bytes([67])==b'BC' 1 2 3 4 5 6 7 8 9 vocab = {idx: bytes([idx]) for idx in range(256)} for (p0, p1), idx in merges.items(): vocab[idx] = vocab[p0] + vocab[p1] def decode(ids): # given ids (list of integers), return Python string tokens = b"".join(vocab[idx] for idx in ids) text = tokens.decode("utf-8", errors="replace") return text 我们首先创建一个字典vocab，包含了256个原始token的字节表示，然后根据merges字典，将融合后的token pair的字节表示，赋值为原来的token pair的字节表示的拼接。最后我们可以使用decode函数，将一段tokens恢复为原来的文本。这样就不用一步一步去替换新的token idx为原来的token pair，而是直接将所有的token idx对应的字节表示拼接起来，再解码为文本。 Encoding 和解码相对的就是我们怎么编码文本为tokens。编码的过程和解码的过程是相反的，我们从文本开始，将其编码为utf-8的字节序列，并转为list对象，此刻就成为了整数表示形式，然后我们反复更新这个tokens序列也就是编码它，除非它的长度小于2，即它无法再融合，或者最小的token pair都已经不在merges字典中了，就停止编码并返回此刻的tokens。为什么不从最大的pair对开始融合呢，因为最大的pair对也是从较小的融合而来的，在不融合出它的子集之前，在tokens中找不到找这个较大的pair对。 1 2 3 4 5 6 7 8 9 10 11 12 13 def encode(text): # given a string, return list of integers (the tokens) tokens = list(text.encode("utf-8")) while len(tokens) &gt;= 2: stats = get_stats(tokens) pair = min(stats, key=lambda p: merges.get(p, float("inf"))) if pair not in merges: break # nothing else can be merged idx = merges[pair] tokens = merge(tokens, pair, idx) return tokens print(encode("")) 所以，merges这个变量记录了所有的融合操作，可以用于编码阶段，而vocab变量记录了所有的token的字节表示，可以用于解码阶段。 正则表达式 有了编解码，我们就可以利用tokenizer，将文本转换为tokens，然后输入到LLM中进行处理，再从LLM中输出tokens，最后利用tokenizer将tokens转换为文本。 但是，训练这样一个简单的tokenizer无法解决一些特殊的问题，比如在英文句子中，单词之间有空格，而tokenizer是基于字节级的，所以空格会被编码为一个token，在训练集中就会出现很多这样的模式，例如dog这个单词，在文本中会出现许多代表相同意思的但是稍微有些不同字节表示的形式'dog.',' dog','dog!','dog?'，所以BPE会对每一个这样的不同模式的dog构建一个token，获得了只是稍有不同的dog这样的token，就是BPE把一些不该编码的部分也加入了进来，比如单词和标点符号。 所以必须要有一种人工的干预，强制不让一些符合规则的字符合并在一起。所以基本上我们可以构建这样的一个正则表达式"""'s|'t|'re|'ve|'m|'ll|'d| ?\p{L}+| ?\p{N}+| ?[^\s\p{L}\p{N}]+|\s+(?!\S)|\s+"""，用来分离那些不该被合并的字符。它的实现逻辑就是我们不去找那些烦扰的字符，而是去匹配我们合理的关心的模式，把这些模式提取出来并形成一个列表，那么没有提取出来的就是我们不关心的东西，所以我们只提取出了关切的部分。在hello world how are your这个例子中，经过正则匹配我们可以得到一个列表['hello',' world',' how',' are',' you']，我们在训练tokenizer之前，首先做的就是人工的划分单词，就是' are' ' you'中e不会和 you中的第一个空格融合。接下来所有的融合都只会各自发生在列表中的各个元素中，不会跨元素进行融合，融合之后的结果再进行拼接就得到了我们训练好的tokenizer。 1 2 3 4 5 6 7 8 9 10 for i in range(1, 101): if i % 3 == 0 and i % 5 == 0: print("FizzBuzz") elif i % 3 == 0: print("Fizz") elif i % 5 == 0: print("Buzz") else: print(i) ['\n', 'for', ' i', ' in', ' range', '(', '1', ',', ' 101', '):', '\n ', ' if', ' i', ' %', ' 3', ' ==', ' 0', ' and', ' i', ' %', ' 5', ' ==', ' 0', ':', '\n ', ' print', '("', 'FizzBuzz', '")', '\n ', ' elif', ' i', ' %', ' 3', ' ==', ' 0', ':', '\n ', ' print', '("', 'Fizz', '")', '\n ', ' elif', ' i', ' %', ' 5', ' ==', ' 0', ':', '\n ', ' print', '("', 'Buzz', '")', '\n ', ' else', ':', '\n ', ' print', '(', 'i', ')', '\n'] tiktokenizer这个网站中可以查看大语言模型的tokenizer结果,在GPT-2中，这些空格都没有进行合并，所以GPT-2中的tokenizer训练不仅仅是简单地将BPE应用到每一个提取的单词中，而是额外添加了一些规则，不过在这里我们理解预处理的过程，使用简单的处理方式就可以了。 在GPT-4o中，可以看到同样的代码，但是在tokenizer中，空格被合并了，对tokenizer的训练进行了优化。 tiktoken库是OpenAI官方的一个tokenizer库，它可以用来将文本转换为tokens，也可以将tokens转换为文本。它的使用方法和我们之前实现的tokenizer类似，但是它是基于OpenAI的模型训练的，所以它的tokenizer结果和OpenAI的模型训练结果是一致的。 special tokens len(encoder) #256 raw byte tokens. 5,000 merges. +1 special token。这个special token就是&lt;|endoftext|&gt;，它在OpenAI的模型中被用来表示一段文本的结束，也就是说这个特殊token前后的两段文本是独立的，它们没有任何关系。假设我们有一个很大的数据集，这些数据集都是独立的文本，从各个数据源获取的，我们当然希望这些文本之间不应该有语义上的关联，比如A文档是在叙述一段小说，而B文档则是关于科学的东西，所以需要这个标志告诉模型前后之间不相关，上一段的文档内容提供的信息不应该继承到下一段文本中来。当然，可以注册许多特殊的token，比如用来区分用户和系统提示信息，不希望将系统提示信息暴露给用户方。 实现正则化版本的tokenizer 这段代码就是一个正则化版本的tokenizer训练过程，它的基本用到的函数都没有改变，只是现在需要作用到每一个提取的单词中，而不是直接作用到整个文本中。获取每个独立单词中出现的所有字符对的统计信息，然后根据统计信息对每个单词进行合并操作，直到达到指定的合并次数。 1 2 3 4 5 6 7 8 9 10 11 12 13 text_chunks=re.findall(self.compiled_pattern,text) ids=[list(ch.encode('utf-8')) for ch in text_chunks] for i in range(num_merges): idx=i+256 stats={} for ch in ids: stats=self._get_stats(ch,stats) pair=max(stats,key=stats.get) ids=[self._merge(ch_ids,pair,idx) for ch_ids in ids] self.merges[pair]=idx first,second=pair self.vocab[idx]=self.vocab[first]+self.vocab[second] sentencepiece sentencepiece是Google提出的一个tokenizer库，它是作用在码点上的。要理解码点，先要分清楚utf-8和unicode的区别。简单来说，unicode是一个字符集，它定义了每个字符的唯一编码号，这个编码号被称为码点，比如中文里面的我的码点就是U+6211，而utf-8是一个编码方案，计算机只能存储0和1，为了将码点存储下来，使用utf-8编码方案，U+6211的utf-8编码就是0xE6 0x88 0x91这3个字节，它将每个码点映射为一个或多个字节序列。sentencepiece和tiktoken的区别简而言之，就是它们使用BPE的层级不同，前者在unicode的码点下进行合并，而后者在更贴近计算机底层的utf-8层级使用。一个更贴近人类语言，一个更贴近计算机语言。它们面向的场景不同，tiktoken面向通用文本，跨语言，因为所有语言底层逻辑都是用二进制存储，也就是可以用utf-8表示。而sentencepiece面向多语言精细化分词，字符是人类语言的基本表意单元，从字符出发合并更符合语言的语义结构。 vocab_size 在transformer中，vocab_size是指模型中使用的词汇表大小，它包括了所有可能的输入和输出符号。它会出现在两个地方，一个是token_embedding_table(shape=(vocab_size, d_model))，这个层会将每个token映射为一个d_model维的向量，所以vocab_size就是token_embedding_table的第一维。除此之外还有一个是在transformer的末端LM_head层，会将d_model维的向量映射为vocab_size维的向量，所以LM_head的输出维度就是vocab_size。就是我们会为每一个token(vocab_size大小)在每一个时间步上生成一个预测概率，随着我们有着越来越多的token，就需要预测更多的概率，就是在最后一个线性层上进行更多的点积运算。 所以token_embedding_table会随着vocab_size的增加而增加，LM_head线性层会随着vocab_size的增加而增加，会有更多的计算；更多的token意味着有更多的参数，可能担心很多参数没有得到充分的训练，因为引入更多的token，就是均摊了其它token在数据集上出现的频率，总体上所有token都会出现更少的示例占比，所以token的频率降低可能意味着它们在前向后向传播过程中的机会不多；此外，更多的token意味着对数据集的压缩更大，在适当的情况下我们可以使用较少的token来表达更多的文本(比如原来用80个token来表示一段话，现在只需要用50个token)，但是如果vocab_size过多，也可能会导致一大段话被压缩成一个token，这样模型在思考处理一定量的字符时，时间就不会那么充裕，因为过多的信息被过多压缩了。 Part 9-Reproduce GPT-2(124M) 考虑到项目代码比较长，在这里详细解释各个部分，不如直接深入源码。在这里只是提炼其中的核心并分析，更细节的部分可以直接参考源码，源码中会有详细注释。 最初的样子 项目最初有CausalSelfAttention、MLP、Block、GPTConfig、GPT一共5个类型。 CausalSelfAttention key,query和value以一行代码的批量计算紧凑形式组织在了一起，使用register_buffer，表示生成的下三角矩阵张量是缓冲器即非可学习参数，会随模型保存，但不参与梯度更新。在前向传播时，会调整各个维度的位置，只需要记住，参与计算的最后两个维度一定是[T，head_size]。这里为了加速注意力计算，使用了pytorch提供的接口函数F.scaled_dot_product_attention。 MLP 提供的是注意力计算后的FFN步骤。不过这里使用了gelu激活函数，且使用tanh近似GELU加快训练 Blcok 将前两个模块组装到一起，由于是自然语言，所以这里使用的是层归一化。 GPTConfig 这个类使用python的dataclass进行注册，方便参数的管理 GPT 这个类从名称上就可以看出，是GPT模型的实现。其中重要的是其可以加载OpenAI官方GPT-2预训练好的权重参数。加载预训练权重时，需要获取二者的参数字典，即自定义实现的模型的参数字典和OpenAI官方GPT-2模型的参数字典。此外，OpenAI使用Conv1D,它的输入和输出的维度与普通的Linear层是相反的，所以需要转置后才能copy到自定义的模型中。 实现forward,autoregressive 前向传播分为三个部分，将输入idx(shape=(B,T))进行token_embedding，得到(shape=(B,T,n_embd))和位置编码(shape=(T,n_embd))，将这两个张量相加。以block为组计算transformer结构，最后使用层归一化和线性层得到每一个待预测的token的logits(shape=(B,T,vocab_size))。 实现自回归的生成器，首先使用tiktoken的gpt-2的tokenizer对输入的文本进行编码，之后就可以传入模型得到logits,套路和之前的第7部分的生成器内容是差不多的，不过这里是从前50最大概率中选取一个token作为下一个token。 实现损失计算 在forward中添加了损失计算部分，如果传入的target不为空，则计算损失。 实现一个数据加载器 要训练的数据从哪里来呢？怎么组织好直接可以传入模型进行训练呢？这就需要一个数据加载器帮助我们完成这些工作。DataLoader不仅会加载数据到内存中缓存好，还可以分批次扔出数据用于当前批次数据作为输入进入模型，进行训练，它通过一个current_position记录当前应该是第几个批次的数据，之后通过next_batch去获取对应位置的数据然后传递给模型进行训练。不过该数据加载中存在一个未解决的小bug,就是最后一段数据不足以支撑BT批次大小时，会重置当前位置为0，意味着末尾有一段数据永远不会用于训练。保持这样处理的好处是，每个批次的数据都是固定的，内部的统计量不会改变，如果采用循环读取操作（当末尾的数据填充不满BT时，使用头部的一段数据进行填充）会导致每个epoch训练时，批次的数据其统计量变化，可能（我猜测）会影响模型训练效果，此外，不利用末尾这一小段数据，在整个样本比例中其实占比不大，影响可以忽略。 之后简单地使用一个优化器，训练模型，更新权重。 权重共享 权重共享有两个问题，一个是理解为什么要这样做？第二个就是分清权重矩阵的存储逻辑和运算逻辑。 先来看第一个问题，为什么要权重共享？ 降低参数量，这个是最直接的收益。参数的收益量是vocab_size*num_embd。 符合自回归预测的任务逻辑 语言模型的核心任务是子回归预测，给定前序token，预测下一个token。过程的本质是‘嵌入-编码-解码’的闭环过程，权重共享让这个闭环更合理。如果不共享权重，lm_head会学习另外一套逆映射，lm_head的本质是在学习一套wte映射的逆操作，将经过嵌入并编码后的结果映射回token。使用共享权重，就像是用同一把钥匙锁门和开门。 词嵌入层wte的权重矩阵中，每一行对应一个token的词嵌入向量，行与行之间的距离(比如欧式距离)体现了token的语义相似度，当lm_head共享这套权重时，本质上就是每个token的预测得分=隐藏向量*对应token的嵌入向量，点积越大，说明二者的语义越匹配，预测概率越高，这套共享机制让预测过程直接利用了嵌入层学到的语义信息，确保语义相似的token在预测时会得到更高的关联得分，让模型的预测更符合语义逻辑。本质上是因为我们对嵌入层的嵌入向量之间的相似性解释，在预测时也希望能够利用上这种相似性，从而体现语义上的连贯性。 此外，无偏设计（没有偏置）是保证了不会破坏语义上的对称。 权重矩阵的运算逻辑和存储逻辑是相反的 pytorch中，变量x和权重矩阵w相乘时，其实是x@W.T，也就是说构建权重矩阵时我们是按照运算的逻辑构建的即权重矩阵的shape=[input,output]，而存储逻辑的形状其实是[output,input]，所以对于wte权重矩阵，因为在由token转到词嵌入空间中我们没有使用投影运算（矩阵相乘），而是直接使用索引，索引到当前token整数对应的wte的行，取出该行的词嵌入向量，所以wte通过nn.Embedding构建，而lm_head需要对编码后的向量进行投影，所以需要使用矩阵运算，构建时使用nn.Linear，这样，它构建时需要使用[input,output]这样的形式，但是底层存储时的形状会是[output,input]，转置一下就符合矩阵运算需要的形状了。 初始化权重 为了最真实地贯彻GPT-2的路线，初始化的设置也选择了最贴近GPT-2的初始化设置。 控制残差流动引起的方差变化 1 2 3 x=torch.zero(768) for i in range(100): x+=torch.randn(768) 通过以上的例子，模拟了一个残差流不断累加的过程，其方差会不断增加。为了控制方差，维持前后方差的一致性，需要对残差流动进行归一化处理。所以在每个transformer block中，有两次残差连接，所以针对残差连接处的计算，在初始化投影矩阵权重时，需要除以根号下的两倍总的transformer block数。 不过，这里有一个问题，就是这种方式确保的是最终输出的变量，其方差是不变的，但是在每个transformer block中，其方差都缩小了。就是这样一个意思，原本我们为了确保初始化时，矩阵乘法后，输出的方差是基本不变的，所以除以了根号下的输入维度数，比如768，这样每一层的输出其方差都维持在这个位置。但是呢，由于引入了残差连接，所以确保在最终输出的方差不变的情况下，又对残差处进一步除以了一个值，这就导致每一层其方差变小了。不再是原来的输出后其方差都维持在这个位置的说法了。 启用TF32训练 TF32是GPU内部的一种计算方式，默认情况下使用FP32，当激活了TF32后，在GPU内部计算矩阵乘法时(其它运算依然使用FP32)，会使用TF32进行计算，从而提高计算效率。输入输出依然是FP32，只是在计算时采用了TF32。矩阵乘法用 TF32 加速（FP32 输入→TF32 计算→FP32 输出），其他运算仍为 FP32。 启用BF16 有必要简单地区分下FP32、TF32、FP16和BF16，FP32有8位指数位，用来表示范围，23位精度位用来确保精度，TF32相对于FP32，在精度位上只使用了10位。FP16，在精度位上和TF32保持一致，都使用10位，但是它的指数位只有5位，而BF16，它的指数为和TF32一样，有8位，但是精度位进一步缩减，变为了7位。它的最佳实践文档可以参考AMP。基本上来说，就是在中间变量，将FP32转为BF16，进行训练和反向传播，梯度更新时再转回FP32确保精度。所以结合TF32(仅仅针对矩阵乘法)，在没有启用BF16时，它的输入输出是FP32，仅仅是中间计算时用了TF32加速，现在输入输出变成了BF16。所以BF16是整体上将FP32迁移成了BF16，不仅仅包括矩阵乘法，而是其所有的中间变量的临时存储都是BF16。总结起来就是核心是计算用BF16，参数存储用FP32。所以这就是为什么叫做混合精度的原因，一些东西在pytroch中仍然保持FP32(权重矩阵)，而一些东西精度降低了，成了BF16(激活值、中间计算的临时变量等等)。但是，又来了，又是但是，仔细查看上面的最佳实践文档可以知道，其实也并不是所有层都转成了BF16，不过总而言之，言而总之，启用BF16结合TF32，确实加速了我们的训练，节省了内存。 启用torch.compile torch.compile对于加速的作用主要来自于减少python的开销和GPU读取次数。减少python的开销，简单地比喻就是python解释器就像厨师做饭，按照菜谱，要一步一步做，而自动化厨房只需要把原材料给它，就会自动运行。对于减少python的开销，torch.compile可以看到你所要操作的整个流程，而python解释器是逐行运行，它并不知道接下来会发生什么。torch.compile不会以一种eager模式运行，会优化运行的过程。它会首先移除python解释器在前向传播中的作用，将整个神经网络编译成不涉及python解释器的单一对象，然后直接运行。而对于GPU读取次数，举个简单例子就是说进行各种乘除法运算时，如果针对的同一个变量(该变量需要经过多种复合运算得到)，在没有使用torch.compile时，GPU会一步一步地在内核中运算-存储到GPU内存中-读取GPU内存中的内容-再次运算，而使用了torch.compile后，它就像知道了所有这些操作的最终结果都是为了得到那一个变量，就会直接在内核中进行连续的操作运算，而不用反复地读取存储了，这就是GPU读取次数的意思。 转换到Flash Attentino 简而言之，pytorch内部实现了这个机制，使得计算注意力分数的速度有了提升。 fit nice number 计算机科学中，由于二进制的原因，许多对于数字的设定都偏爱2的次幂，这些神奇的数字。所以对于一些设定的参数，比如批次、时间步长等等，优化成2的倍数，能发挥更好的计算机性能。 梯度裁剪 就是防止出现梯度爆炸、梯度更新反复横跳等现象。在反向传播和梯度更新之间执行，一般常用的是收集所有参数的梯度计算其L2范数，之后设定阈值进行裁剪。 余弦衰减学习率调度器 需要注意，学习率调度器和优化器是两个不同的东西，一个是针对学习率，一个是在学习率固定下来后怎么去计算梯度更新的策略。在优化器的参数中，去更新刷新后的学习率。余弦学习率调度器本质上就是根据当前的训练epoch次数，去计算应该使用什么大小的学习率。 权重衰减和梯度累积 在GPT-2，GPT-3训练过程中采用了变化的batch size大小的策略进行训练。基于这样的观察和解释，在模型训练的初期，基本上是在学习忽略那些不常出现在训练集中的token，学习非常简单的偏差和类似的东西，每一个样本都在告诉模型，使用这些token，不使用那些Token,来自每一个样本的梯度实际上是高度相关的，在优化的初始阶段，它们看起来都大致相同，因为它们都在告诉模型这些token出现，那些token不出现。所以在训练初期，没有必要使用很大的batch size。只有当跨越过初期阶段，使用大批量样本才会展现出统计上的意义，去学习更深层次的东西，比如“语境歧义”（如 “苹果” 是水果还是公司）。怎么理解这段话呢？ 简而言之，打个比方把模型训练比作 “老师教学生学语文”，batch size 比作 “一次布置的作业量”： 初期（学拼音、常用字）：学生的核心任务是 “记住常用字怎么写、怎么读”—— 所有作业（样本）都在重复 “听写常用字”，学生的错误（梯度）都集中在 “生僻字不会写”“常用字写错笔画” 上（梯度高度相关）。这时候布置 10 道题（小 batch）和 100 道题（大 batch）的效果一样 —— 学生都是在纠正相同的错误，100 道题只会让学生更累（浪费时间），不会更快掌握。 后期（学阅读理解、写作）：学生的任务是 “理解语境、掌握多义词、组织逻辑”—— 作业题（样本）五花八门：有的考 “‘打’在‘打球’和‘打电话’中的不同含义”，有的考 “议论文的论点论据”，有的考 “散文的情感表达”（梯度多样性高）。这时候布置 100 道题（大 batch）比 10 道题（小 batch）效果好 —— 学生能接触更多场景，避免 “只懂一道题，换题就错”（降低梯度噪声），学到的规律更全面（统计上的通用逻辑）。 所以，训练初期的核心是 “学简单规律”，梯度同质化，小 batch 足够用，且效率更高； 训练后期的核心是 “学复杂规律”，梯度异质化，大 batch 能覆盖更全的样本分布，让模型学到统计上的通用规律； 但在本项目的实现中，跳过了这个步骤，因为它会把问题复杂化，比如怎么动态地处理batch size带来的数据变化。 真实地遵循GPT-3的训练策略，其中提到了使用0.1的权重衰减，就是L2正则化。但是不是对所有的参数都进行正则化，主要是针对矩阵乘法部分的参数，比如线性层的权重矩阵。为此需要构建一个configure_optimizer函数，来确哪些需要权重衰减，哪些不需要。其中优化器使用到了fused这个选项，就是把许多计算需要用到许多核的情况，融合成了一个核，简而言之就是加快计算速度。 在有限的资源，比如一个GPU中，如何使用0.5M大小的总token数(=B*T)进行训练呢？如果直接指定对应的batch size，那么GPU的内存会溢出。这个时候就要使用梯度累积的策略了。它允许我们采用串行的方式，模拟任何批次大小的数据。代价就是运行时间变成，处理多个序列然后把这些梯度加起来。所以基本策略就是使用一个小batch size，多次前向传播-反向传播，但是不更新梯度，直到重复B/min_batch_size次，再更新梯度。需要注意的是，我们使用了loss_accum来记录最终的0.5M这个batch size大小下的损失，它是被detach掉的。 1 2 3 4 5 6 7 8 for micro_step in range(grad_accum_steps): x,y=train_loader.next_batch() x,y=x.to(device),y.to(device) with torch.autocast(device_type=device,dtype=torch.bfloat16): logits,loss=model(x,y) loss=loss/grad_accum_steps loss_accum+=loss.detach() loss.backward() Distributeddataparalle(DDP) 怎么利用多GPU进行训练。这里的难点在于我们需要想像有8个并行运行的程序在运行相同的代码，它们的区别就是ddp_rank。因为我们有了多GPU，原本代码中的一些参数值就得考虑到平均后的正确数值是多少了，比如原本我们在单个GPU中使用mini_batch，需要前向-反向传播B/mini_batch次才能进行一次梯度更新，现在考虑到使用多GPU，这个前向-反向传播次数进一步被平均了，所以需要B/(mini_batch*num_GPU)次就可以了。我们需要让DataLoader根据不同的GPU加载属于自己的那份数据，而不是每个GPU都加载同一段的数据。 目前总结下代码中的一个容易误导的地方，就是现在是按照总的优化器更新次数来训练模型，而不是按照epoch次数训练模型。按照传统epoch次数的理解，举例100个epoch，就需要每个epoch都要完整地遍历一次训练数据集。而现在，采用的是总的优化器更新次数，每次都遍历一定量的token数，比如设定总的优化器更新次数为max_steps，那么它等同于max_steps/(total_token/batch_size_token)个epoch。]]></summary></entry><entry><title type="html">YOLO-目标检测</title><link href="https://ke-albert.github.io/2025/10/05/YOLO/" rel="alternate" type="text/html" title="YOLO-目标检测" /><published>2025-10-05T15:43:31+08:00</published><updated>2025-10-05T15:43:31+08:00</updated><id>https://ke-albert.github.io/2025/10/05/YOLO</id><content type="html" xml:base="https://ke-albert.github.io/2025/10/05/YOLO/"><![CDATA[<h1 id="yolo-目标检测">YOLO-目标检测</h1>

<p>TLTR
1-学习YOLO的过程
2-Andrew Ng深度学习课程中关于YOLO的讲解
3-霹雳吧啦Wz的讲解，https://www.bilibili.com/video/BV1yi4y1g7ro/</p>

<h2 id="学习yolo的过程">学习YOLO的过程</h2>

<p>我是从25年7月份正式开始学习YOLO目标检测算法的，在这之前听了吴恩达深度学习课程中的关于YOLO的讲解，但是当时没有听懂，所以相当于重新开始。学习的是霹雳吧啦Wz的讲解，首先介绍了目标检测算法的两个方案：two-stage和one-stage。two-stage的方案包括先验框预测和分类器，代表是Faster R-CNN，而one-stage的方案则是直接预测目标框和分类，代表是YOLO，以及DETR等结合了注意力机制的模型，不过我现在主要是学习YOLO，DETR没有接触过。学习了YOLOv1-v3，其中YOLO-V3还有一个SPP的版本，这也是我跟随学习代码实现的版本，目前已经完成了YOLOV3-SPP的源码学习。当然，在这中间我还穿插着复习过Andrew Ng深度学习中涉及到的YOLO知识， <code class="language-plaintext highlighter-rouge">[还发现了别人总结好的网页](http://www.ai-start.com/dl2017/)</code></p>

<h2 id="andrew-ng-深度学习课程中关于yolo的讲解">Andrew Ng 深度学习课程中关于YOLO的讲解</h2>
<p>课程中，Andrew 首先将目标定位分解为了两个子问题：定位和分类。分类问题很简单，判断输入的图片中的物体是什么类型，将分类和定位结合起来，除了要判断物体的类型，还要判断物体的位置，这两个问题都是适用于单一目标，即每个图片中只有一个目标。进一步提出目标检测，图中存在着多个目标，我们都需要将它们分类和定位，这是一个多目标问题。
接着从分类的输出中，扩充定位信息，普通的分类输出是一个经过softmax的向量，对应物体的类别数量，需要定位我们只需要在输出中额外增加4个值，用来表示物体的边框信息。有了输出，我们需要定位监督学习的目标标签，目标标签是一个向量，包括一个Pc用来表示是否含有对象（可以理解为置信度），4个表示边框的值，物体类别数量的值，在物体类别中还可以使用一个值来表示，这个值从1-n变化。
Andrew 将4个边框值的表示扩充到了特征点检测问题，广义上神经网络可以输出图片上特征点的坐标，在目标检测中输出4个值表示边框信息，需要检测特征点可以设置任意需要统一输出特征点的数量，然后制作目标标签进行训练。</p>
<h3 id="基于滑动窗口的目标检测算法">基于滑动窗口的目标检测算法</h3>
<p>对图片进行裁剪，以识别汽车为例，只保留包含汽车的区域，其它区域裁剪掉，然后对裁剪后的图片进行分类，判断是否为汽车，输出y=0|1，这样就训练出了一个分类网络，对于多类别的同理输出y=0|1|…|n。网络训练好之后就可以基于滑动窗口来实现目标检测了，其中的定位问题，在滑动的时候就已经内含在其中了，即只要在当前窗口中识别出了一个类别，那么这个窗口的位置就是该类别的位置，我们可以记录窗口的横向和纵向滑动距离。滑动窗口的实现很灵活，首先可以选择是否以要重叠，即下一次的滑动是否与本次的部分区域重叠，从图上左上角向右下滑动，通过步距进行控制。还可以控制滑动窗口的大小。总之，这些人为控制的操作都是为了有效识别出物体，比如有些物体可能会存在两次滑动窗口之间。
这也就引出了滑动窗口的问题：计算成本太高。如果用小步幅，无法准确定位图中的对象，如果用大步幅（包括多个目标），粗糙间隔会影响性能。
为了提高计算效率，将滑动窗口使用卷积来实现。你可能会想，之前的滑动窗口不就是基于卷积实现的吗？再仔细分析下，最初的滑动窗口首先通过卷积判断当前窗口中有没有待分类的物体，至于滑动窗口的移动和卷积没有一点关系，可以通过两层循环实现这个逻辑。而这里的滑动窗口的卷积实现，是指将滑动窗口的移动步骤也以卷积的方式内含实现，在卷积的过程中就相当于移动了滑动窗口，过程不同但是在结果上是等价的，以至于我们可以理解为就像移动了窗口，但这是利用了卷积计算原理和它的高效操作实现的，因为卷积窗口也是需要滑动的，不过这个滑动是在pytorh等深度学习框架中实现了的，借助了GPU的高性能计算，就是说这个滑动是优化过的，提出该方法的论文是Sermanet, Pierre, et al. “OverFeat: Integrated Recognition, Localization and Detection using Convolutional Networks.” Eprint Arxiv (2013)。其实就是将一维的全连接层替换成了多维的表示，不过这个多维表示也就表明它可以在其它维度上增加数量，从而将输出的特征层1x1xN变成mxmxn。该卷积操作的原理就是我们不需要把图像分割成子集，分别执行前向传播，而是将图像整体输入给卷积网络计算，其中有许多区域可以共享计算，最终得到输出层。换个角度就是感受野的解释，我们通过控制步距，窗口大小（这里的窗口大小不是卷积层的窗口大小，而是滑动窗口的大小，也就是在还没有使用卷积实现滑动窗口的时候，滑动窗口的大小，而这时我们可以去调整卷积层的大小，固定后，即卷积层大小固定，滑动窗口大小固定后，可以使用卷积的形式实现滑动窗口），来实现不同的下采样，在最后的输出结果中我们就可以得到感受野，最后输出的特征层大小就是相当于我们滑动了多少次窗口，它的位置映射回原图上就是目标框看的区域内容。
但是该算法仍然存在不能完美定位的问题，比如目标跨过多个区域时，一个窗口定位只能定位到部分区域，还有些目标更适合用长方形的框来定位。</p>
<h3 id="yolo">YOLO</h3>
<p>YOLO算法其实也是根据感受野的解释反推在原图上的目标框位置，不过它输出了边框参数，所以在边框的形状上可以很灵活的变化学习。将原图划分成SxS个小区域，这SxS个小区域在原图上经过卷积网络后得到的输出层大小就是SxS(pixel)，所以YOLO预先定义输出层大小，通过控制输出层的尺寸来控制细粒度，以更好地识别目标。对于这SxS个输出，每个位置都包括一些参数(Pc,boxes-info,class)，通过构建对应的目标标签进行训练这样的一个目标检测网络。
这只是网络的原理，还需要结合额外的人为控制才能更好地训练网络，使用预测的边框。交并比、非极大值抑制、Pc阈值、Anchor-box。
Andrew 介绍的YOLO算法令人恍然大悟，但是对于其中的细节没有深入探究，代码部分的实战只是简单的实现，没有关于训练的部分内容，比如数据集处理、损失函数计算、正负样本选取这些。为了识别同一区域内的多个目标，增加了anchor box后，怎么确定哪个anchor box预测的是正确的，怎么根据anchor box来计算损失，这些问题没有解答。</p>
<h2 id="yolov3-spp源码实战">YOLOV3-SPP源码实战</h2>
<p>从最初的YOLOV1开始，就已经有了很多的版本，比如YOLOV2、YOLOV3、YOLOV4等，每个版本都有自己的改进，比如SPP、FPN、PAN等。而YOLOV3-SPP是在YOLOV3的基础上，增加了SPP层，用来提取不同尺度的特征，这也是我跟随学习代码实现的版本。
学习路线首先是定义网络结构，为了让网络具有扩展性，将网络结构存储在了.cfg文件中，通过读取.cfg文件来定义网络结构，这也是YOLO系列的一个特点。有了网络结构，需要解析网络结构，将其转为在内存中规则的数据结构，再根据这个数据结构来搭建网络模型。网络模型搭建好后，训练时需要导入数据集，就需要自己制作数据集入口，包括一系列的检查文件路径、预处理、缓存操作等，都定义在了一个类里面，通过继承pytorch的<code class="language-plaintext highlighter-rouge">Dataset</code>类，传入dataloader中。数据集加载后就是训练，每个批次经过网络输出，得到预测值，需要计算预测值与关联标签的损失，再反向传播，反向传播时又定义了调度器，学习率衰减规则，打印日志等模块内容，等等这些构成了YOLOV3-SPP算法的全部过程。在训练时，除了损失计算部分是与YOLO强关联外，其它的内容都是可以迁移通用的，比如调度器、日志打印模块。</p>
<h3 id="cfg网络定义">.cfg网络定义</h3>
<p>.cfg网络定义内容其实就是一些规则结构的字符串，描述了网络的层结构。并搭配了对应的解析器，用来解析.cfg文件，将其转为在内存中规则的数据结构。它主要有convolutional层，shortcut层，池化层、route层、upsample层、YOLO层。其中shortcut层和route层需要特别注意，不像其它层需要做具体的工作（比如进行大量计算），这两个层之所以单独独立成一个层（在编号中占据一个位置），shortcut层负责残差连接，所以它的关键字<code class="language-plaintext highlighter-rouge">from</code>经常等于-3，以当前shortcut层为索引0，向其上层索引3个，route层有两个作用，拼接多个层的输出和将当前的指针（代表网络当前输出所处的位置）退回到某一层（在SPP的输入多分支时有用），这是因为网络定义的顺序是线性的，即使平行结构的SPP，也需要按照线性顺序排序。YOLO层不是预测头所处的层，它是预测头的下一层（紧连着预测头），它的作用是初始化anchor模板，定义特征层的网格，判断训练还是预测，从而输出不同的结果。</p>
<h3 id="自定义数据集">自定义数据集</h3>
<p>数据集集中了对图片和标签载入的操作预处理。包括设置批量大小，设置预处理输出图片大小、数据缓存、是否进行数据增强等操作。</p>
<h3 id="损失计算">损失计算</h3>
<p>正样本计算目标框损失、分类损失和置信度损失，负样本只计算置信度损失。在<code class="language-plaintext highlighter-rouge">compute_loss</code>函数中，根据预测值和关联标签，计算损失。首先根据预测值和关联标签，筛选出正样本，也即在<code class="language-plaintext highlighter-rouge">build_target</code>函数中，筛选出正样本的类别标签（用于计算类别损失）、正样本对应的gt box信息、含有正样本的图像索引、anchor模板索引、所处的哪一个grid位置信息、anchor模板的大小信息，遍历每一个YOLO输出层，计算这三个损失。其中，置信度损失计算时首先会创建一个tobj，初始值都为0，之后对于筛选出来的正样本，根据正样本计算的iou值动态（每个批次正样本会变）去得到一个标签置信度并在tobj对应的正样本位置填上该值，其它未变的即为负样本的默认值0，然后使用预测的值和tobj进行二值交叉熵计算，得到置信度损失。可以注意到，这里并没有按照正负样本按照一定比例选取，而是计算了全部的负样本，不过使用了Focal loss、调整置信度权重等方式有效地缓解了正负样本不平衡的问题。</p>

<ul>
  <li>build_target函数解析
    <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><table class="rouge-table"><tbody><tr><td class="rouge-gutter gl"><pre class="lineno">1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
</pre></td><td class="rouge-code"><pre>  <span class="k">def</span> <span class="nf">build_targets</span><span class="p">(</span><span class="n">p</span><span class="p">,</span> <span class="n">targets</span><span class="p">,</span> <span class="n">model</span><span class="p">):</span>
  <span class="c1"># Build targets for compute_loss(), input targets(image_idx,class,x,y,w,h)
</span>  <span class="c1"># p: predictions [batch_size, num_anchors, grid_h, grid_w, num_params]
</span>  <span class="c1">#选出正样本
</span>  <span class="n">nt</span> <span class="o">=</span> <span class="n">targets</span><span class="p">.</span><span class="n">shape</span><span class="p">[</span><span class="mi">0</span><span class="p">]</span>
  <span class="n">tcls</span><span class="p">,</span> <span class="n">tbox</span><span class="p">,</span> <span class="n">indices</span><span class="p">,</span> <span class="n">anch</span> <span class="o">=</span> <span class="p">[],</span> <span class="p">[],</span> <span class="p">[],</span> <span class="p">[]</span>
  <span class="c1"># gain：用于将归一化坐标转换到“网格空间”的缩放因子（初始为全1，后续按层更新）
</span>  <span class="n">gain</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">ones</span><span class="p">(</span><span class="mi">6</span><span class="p">,</span> <span class="n">device</span><span class="o">=</span><span class="n">targets</span><span class="p">.</span><span class="n">device</span><span class="p">).</span><span class="nb">long</span><span class="p">()</span>  <span class="c1"># normalized to gridspace gain
</span>
  <span class="n">multi_gpu</span> <span class="o">=</span> <span class="nb">type</span><span class="p">(</span><span class="n">model</span><span class="p">)</span> <span class="ow">in</span> <span class="p">(</span><span class="n">nn</span><span class="p">.</span><span class="n">parallel</span><span class="p">.</span><span class="n">DataParallel</span><span class="p">,</span> <span class="n">nn</span><span class="p">.</span><span class="n">parallel</span><span class="p">.</span><span class="n">DistributedDataParallel</span><span class="p">)</span>
  <span class="k">for</span> <span class="n">i</span><span class="p">,</span> <span class="n">j</span> <span class="ow">in</span> <span class="nb">enumerate</span><span class="p">(</span><span class="n">model</span><span class="p">.</span><span class="n">yolo_layers</span><span class="p">):</span>  <span class="c1"># j: [89, 101, 113]
</span>      <span class="c1"># 获取该yolo predictor对应的anchors
</span>      <span class="c1"># 注意anchor_vec是anchors缩放到对应特征层上的尺度
</span>      <span class="n">anchors</span> <span class="o">=</span> <span class="n">model</span><span class="p">.</span><span class="n">module</span><span class="p">.</span><span class="n">module_list</span><span class="p">[</span><span class="n">j</span><span class="p">].</span><span class="n">anchor_vec</span> <span class="k">if</span> <span class="n">multi_gpu</span> <span class="k">else</span> <span class="n">model</span><span class="p">.</span><span class="n">module_list</span><span class="p">[</span><span class="n">j</span><span class="p">].</span><span class="n">anchor_vec</span>
      <span class="c1"># p[i].shape: [batch_size, 3, grid_h, grid_w, num_params]
</span>      <span class="n">gain</span><span class="p">[</span><span class="mi">2</span><span class="p">:]</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">tensor</span><span class="p">(</span><span class="n">p</span><span class="p">[</span><span class="n">i</span><span class="p">].</span><span class="n">shape</span><span class="p">)[[</span><span class="mi">3</span><span class="p">,</span> <span class="mi">2</span><span class="p">,</span> <span class="mi">3</span><span class="p">,</span> <span class="mi">2</span><span class="p">]]</span>  <span class="c1"># xyxy gain
</span>      <span class="n">na</span> <span class="o">=</span> <span class="n">anchors</span><span class="p">.</span><span class="n">shape</span><span class="p">[</span><span class="mi">0</span><span class="p">]</span>  <span class="c1"># number of anchors
</span>      <span class="c1"># [3] -&gt; [3, 1] -&gt; [3, nt]
</span>      <span class="n">at</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">arange</span><span class="p">(</span><span class="n">na</span><span class="p">).</span><span class="n">view</span><span class="p">(</span><span class="n">na</span><span class="p">,</span> <span class="mi">1</span><span class="p">).</span><span class="n">repeat</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="n">nt</span><span class="p">).</span><span class="n">to</span><span class="p">(</span><span class="s">'cuda'</span><span class="p">)</span>  <span class="c1"># anchor tensor, same as .repeat_interleave(nt)
</span>
      <span class="c1"># Match targets to anchors
</span>      <span class="n">a</span><span class="p">,</span> <span class="n">t</span><span class="p">,</span> <span class="n">offsets</span> <span class="o">=</span> <span class="p">[],</span> <span class="n">targets</span> <span class="o">*</span> <span class="n">gain</span><span class="p">,</span> <span class="mi">0</span>
      <span class="k">if</span> <span class="n">nt</span><span class="p">:</span>  <span class="c1"># 如果存在target的话
</span>          <span class="c1"># 通过计算anchor模板与所有target的wh_iou来匹配正样本
</span>          <span class="c1"># j: [3, nt] , iou_t = 0.20
</span>          <span class="n">j</span> <span class="o">=</span> <span class="p">(</span><span class="n">wh_iou</span><span class="p">(</span><span class="n">anchors</span><span class="p">,</span> <span class="n">t</span><span class="p">[:,</span> <span class="mi">4</span><span class="p">:</span><span class="mi">6</span><span class="p">])</span> <span class="o">&gt;</span> <span class="n">model</span><span class="p">.</span><span class="n">hyp</span><span class="p">[</span><span class="s">'iou_t'</span><span class="p">]).</span><span class="n">to</span><span class="p">(</span><span class="s">'cuda'</span><span class="p">)</span>  <span class="c1"># iou(3,n) = wh_iou(anchors(3,2), gwh(n,2))
</span>          <span class="c1"># t.repeat(na, 1, 1): [nt, 6] -&gt; [3, nt, 6]
</span>          <span class="c1"># 获取正样本对应的anchor模板与target信息
</span>          <span class="n">a</span><span class="p">,</span> <span class="n">t</span> <span class="o">=</span> <span class="n">at</span><span class="p">[</span><span class="n">j</span><span class="p">],</span> <span class="n">t</span><span class="p">.</span><span class="n">repeat</span><span class="p">(</span><span class="n">na</span><span class="p">,</span> <span class="mi">1</span><span class="p">,</span> <span class="mi">1</span><span class="p">)[</span><span class="n">j</span><span class="p">]</span>  <span class="c1"># filter
</span>
      <span class="c1"># Define
</span>      <span class="c1"># long等于to(torch.int64), 数值向下取整
</span>      <span class="n">b</span><span class="p">,</span> <span class="n">c</span> <span class="o">=</span> <span class="n">t</span><span class="p">[:,</span> <span class="p">:</span><span class="mi">2</span><span class="p">].</span><span class="nb">long</span><span class="p">().</span><span class="n">T</span>  <span class="c1"># image_idx, class
</span>      <span class="n">gxy</span> <span class="o">=</span> <span class="n">t</span><span class="p">[:,</span> <span class="mi">2</span><span class="p">:</span><span class="mi">4</span><span class="p">]</span>  <span class="c1"># grid xy
</span>      <span class="n">gwh</span> <span class="o">=</span> <span class="n">t</span><span class="p">[:,</span> <span class="mi">4</span><span class="p">:</span><span class="mi">6</span><span class="p">]</span>  <span class="c1"># grid wh 
</span>      <span class="n">gij</span> <span class="o">=</span> <span class="p">(</span><span class="n">gxy</span> <span class="o">-</span> <span class="n">offsets</span><span class="p">).</span><span class="nb">long</span><span class="p">()</span>  <span class="c1"># 匹配targets所在的grid cell左上角坐标
</span>      <span class="n">gi</span><span class="p">,</span> <span class="n">gj</span> <span class="o">=</span> <span class="n">gij</span><span class="p">.</span><span class="n">T</span>  <span class="c1"># grid xy indices
</span>
      <span class="c1"># Append
</span>      <span class="c1"># gain[3]: grid_h, gain[2]: grid_w
</span>      <span class="c1"># image_idx, anchor_idx, grid indices(y, x)
</span>      <span class="n">indices</span><span class="p">.</span><span class="n">append</span><span class="p">((</span><span class="n">b</span><span class="p">,</span> <span class="n">a</span><span class="p">,</span> <span class="n">gj</span><span class="p">.</span><span class="n">clamp_</span><span class="p">(</span><span class="mi">0</span><span class="p">,</span> <span class="n">gain</span><span class="p">[</span><span class="mi">3</span><span class="p">]</span><span class="o">-</span><span class="mi">1</span><span class="p">),</span> <span class="n">gi</span><span class="p">.</span><span class="n">clamp_</span><span class="p">(</span><span class="mi">0</span><span class="p">,</span> <span class="n">gain</span><span class="p">[</span><span class="mi">2</span><span class="p">]</span><span class="o">-</span><span class="mi">1</span><span class="p">)))</span>
      <span class="n">tbox</span><span class="p">.</span><span class="n">append</span><span class="p">(</span><span class="n">torch</span><span class="p">.</span><span class="n">cat</span><span class="p">((</span><span class="n">gxy</span> <span class="o">-</span> <span class="n">gij</span><span class="p">,</span> <span class="n">gwh</span><span class="p">),</span> <span class="mi">1</span><span class="p">))</span>  <span class="c1"># gt box相对anchor的x,y偏移量以及w,h
</span>      <span class="n">anch</span><span class="p">.</span><span class="n">append</span><span class="p">(</span><span class="n">anchors</span><span class="p">[</span><span class="n">a</span><span class="p">])</span>  <span class="c1"># anchors
</span>      <span class="n">tcls</span><span class="p">.</span><span class="n">append</span><span class="p">(</span><span class="n">c</span><span class="p">)</span>  <span class="c1"># class
</span>      <span class="k">if</span> <span class="n">c</span><span class="p">.</span><span class="n">shape</span><span class="p">[</span><span class="mi">0</span><span class="p">]:</span>  <span class="c1"># if any targets
</span>          <span class="c1"># 目标的标签数值不能大于给定的目标类别数
</span>          <span class="k">assert</span> <span class="n">c</span><span class="p">.</span><span class="nb">max</span><span class="p">()</span> <span class="o">&lt;</span> <span class="n">model</span><span class="p">.</span><span class="n">nc</span><span class="p">,</span> <span class="s">'Model accepts %g classes labeled from 0-%g, however you labelled a class %g. '</span> \
                                     <span class="s">'See https://github.com/ultralytics/yolov3/wiki/Train-Custom-Data'</span> <span class="o">%</span> <span class="p">(</span>
                                         <span class="n">model</span><span class="p">.</span><span class="n">nc</span><span class="p">,</span> <span class="n">model</span><span class="p">.</span><span class="n">nc</span> <span class="o">-</span> <span class="mi">1</span><span class="p">,</span> <span class="n">c</span><span class="p">.</span><span class="nb">max</span><span class="p">())</span>

  <span class="k">return</span> <span class="n">tcls</span><span class="p">,</span> <span class="n">tbox</span><span class="p">,</span> <span class="n">indices</span><span class="p">,</span> <span class="n">anch</span>
</pre></td></tr></tbody></table></code></pre></div>    </div>
    <p>该函数接收预测值(P,shapes=[3,Batch,anchor_num,p_h,p_w,5+num_classes]),真值标签(targets,shapes=[num_target,6(img_idx,class_idx,x,y,w,h)],xywh是中心坐标和宽高的归一化值)和模型(model)。该函数的目标是根据模板anchor和真值标签，筛选出正样本对应的anchor和target，再结合预测值计算偏移量，返回正样本对应的类别信息、box偏移量信息、图片索引信息和anchor信息。
  首先获取真值的个数num_target，每个batch可能有不同数量的目标，这个个数是动态变化的。然后创建<code class="language-plaintext highlighter-rouge">tcls, tbox, indices, anch = [], [], [], []</code>，以备后续添加筛选出正样本对应的信息。<code class="language-plaintext highlighter-rouge">gain</code>(每个1对应着targets中这些位置[img_idx,class_idx,x,y,w,h])值用于将归一化坐标转换到“网格空间”的缩放因子（初始为全1，后续按层更新），用于将targets的坐标转换到与预测值相同的网格空间，因为targets最初是采用归一化的形式的。
  <code class="language-plaintext highlighter-rouge">anchors</code>获取到当前yolo predictor对应的anchor模板（就是原本anchor根据步距缩放到当前predictor的特征层上的尺度），形状为[3,2]，表示3个anchor模板，每个模板有2个参数（宽高）。gain参数更新x,y,w,h的缩放因子，用于将targets的坐标转换到与预测值相同的网格空间。<code class="language-plaintext highlighter-rouge">na</code>是anchor模板的个数，一般是3，通过展开维度得到<code class="language-plaintext highlighter-rouge">at</code>，它的形状是[na,num_target]，表示的意思是每个anchor模板对应当前batch中的真实目标数，因为此时我们不知道正样本它的anchor模板与目标对是哪些，有可能一个anchor模板会对应上多个目标(每个网格规定了有3个anchor，这三个anchor都源自于最初的那3个anchor模板)，有可能这些anchor模板(3个)与其中的某个目标都不符合正样本的条件。
  准备好以上的信息后，就可以根据anchor模板与targets的wh_iou来筛选出正样本了。<code class="language-plaintext highlighter-rouge">t</code>是将targes和gain缩放系数相乘得到的缩放到当前预测特征图尺寸上的预测结果，它的归一化坐标变成了当前特征图上的坐标。如果当前batch中存在目标的话，计算anchor模板和t的IoU值，这里的代码是<code class="language-plaintext highlighter-rouge">j = (wh_iou(anchors, t[:, 4:6]) &gt; model.hyp['iou_t'])</code>,它采用了不严格的计算方式，通过把所有的anchor和目标，把它们的中心点平移到坐标原点，再去按照它们的高宽计算IoU值，其实是一种简化的计算路径。通过与预先设定的IoU阈值比较可以得到一个布尔值掩膜<code class="language-plaintext highlighter-rouge">j</code>，它的形状和<code class="language-plaintext highlighter-rouge">at</code>是一样的，都是[na,num_target]，表示每个anchor模板与每个目标的IoU值是否大于阈值。通过这个掩膜<code class="language-plaintext highlighter-rouge">j</code>，就可以筛选出正样本对应的anchor模板和target信息，分别存储在<code class="language-plaintext highlighter-rouge">a</code>和<code class="language-plaintext highlighter-rouge">t</code>中,它们的形状一个是[None,]，一个是[None,6]，这里的None表示的是不清楚匹配到几个正样本，通过debug自定义训练样本此时的None等于22，而num_target的值是13，表示当前batch中有13个目标，但是筛选出来了22个正样本，所以这就验证了多个anchor也会匹配到同一个目标的情况，他们都是正样本。<code class="language-plaintext highlighter-rouge">a</code>中元素的值是anchor模板的索引例如[0,0,0,1,1,1,2,2,2]这样的，它对应着<code class="language-plaintext highlighter-rouge">t</code>中对应元素也就是此刻目标匹配的样本模板，比如a[0]正好对应t[0]。这样我们就得到了正样本(anchor模板-目标对)。
  <code class="language-plaintext highlighter-rouge">b</code>和<code class="language-plaintext highlighter-rouge">c</code>它们的形状都是[None,]，表示的是正样本对应的图片索引和类别索引。<code class="language-plaintext highlighter-rouge">gxy</code>和<code class="language-plaintext highlighter-rouge">gwh</code>则是正样本的中心坐标和宽高(缩放到当前特征图尺寸)，<code class="language-plaintext highlighter-rouge">gij</code>通过向下取整得到的是正样本所在的网格cell的左上角坐标，<code class="language-plaintext highlighter-rouge">gi, gj</code>则是<code class="language-plaintext highlighter-rouge">gij</code>的元素，分别表示网格的y轴和x轴索引。
  indices将(b,a,gj,gi)包裹到一个元组中并加入到当前列表中。记录了正样本对应的图片索引、anchor模板索引、x轴和y轴的网格索引，用来定位正样本在特征图上的所属网格位置。tbox将gt box相对于当前grid cell的x,y偏移量和宽高加入到当前列表中。anch将当前正样本对应的anchor模板加入到当前列表中。tcls将当前正样本对应的类别索引加入到当前列表中。
  以上都是在一个预测特征图尺寸下的正样本筛选和信息记录，不同的预测特征图尺寸下的正样本筛选和信息记录是独立的，互不干扰。总结下我发现最重要的代码部分就是</p>
  </li>
  <li>compute_loss函数解析
    <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><table class="rouge-table"><tbody><tr><td class="rouge-gutter gl"><pre class="lineno">1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
</pre></td><td class="rouge-code"><pre>  <span class="k">def</span> <span class="nf">compute_loss</span><span class="p">(</span><span class="n">p</span><span class="p">,</span> <span class="n">targets</span><span class="p">,</span> <span class="n">model</span><span class="p">):</span>  <span class="c1"># predictions, targets, model
</span>  <span class="n">device</span> <span class="o">=</span> <span class="n">p</span><span class="p">[</span><span class="mi">0</span><span class="p">].</span><span class="n">device</span>
  <span class="n">lcls</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">zeros</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="n">device</span><span class="o">=</span><span class="n">device</span><span class="p">)</span>  <span class="c1"># Tensor(0)
</span>  <span class="n">lbox</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">zeros</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="n">device</span><span class="o">=</span><span class="n">device</span><span class="p">)</span>  <span class="c1"># Tensor(0)
</span>  <span class="n">lobj</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">zeros</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="n">device</span><span class="o">=</span><span class="n">device</span><span class="p">)</span>  <span class="c1"># Tensor(0)
</span>  <span class="n">tcls</span><span class="p">,</span> <span class="n">tbox</span><span class="p">,</span> <span class="n">indices</span><span class="p">,</span> <span class="n">anchors</span> <span class="o">=</span> <span class="n">build_targets</span><span class="p">(</span><span class="n">p</span><span class="p">,</span> <span class="n">targets</span><span class="p">,</span> <span class="n">model</span><span class="p">)</span>  <span class="c1"># targets
</span>  <span class="n">h</span> <span class="o">=</span> <span class="n">model</span><span class="p">.</span><span class="n">hyp</span>  <span class="c1"># hyperparameters
</span>  <span class="n">red</span> <span class="o">=</span> <span class="s">'mean'</span>  <span class="c1"># Loss reduction (sum or mean)
</span>
  <span class="c1"># Define criteria
</span>  <span class="n">BCEcls</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">BCEWithLogitsLoss</span><span class="p">(</span><span class="n">pos_weight</span><span class="o">=</span><span class="n">torch</span><span class="p">.</span><span class="n">tensor</span><span class="p">([</span><span class="n">h</span><span class="p">[</span><span class="s">'cls_pw'</span><span class="p">]],</span> <span class="n">device</span><span class="o">=</span><span class="n">device</span><span class="p">),</span> <span class="n">reduction</span><span class="o">=</span><span class="n">red</span><span class="p">)</span>
  <span class="n">BCEobj</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">BCEWithLogitsLoss</span><span class="p">(</span><span class="n">pos_weight</span><span class="o">=</span><span class="n">torch</span><span class="p">.</span><span class="n">tensor</span><span class="p">([</span><span class="n">h</span><span class="p">[</span><span class="s">'obj_pw'</span><span class="p">]],</span> <span class="n">device</span><span class="o">=</span><span class="n">device</span><span class="p">),</span> <span class="n">reduction</span><span class="o">=</span><span class="n">red</span><span class="p">)</span>

  <span class="c1"># class label smoothing https://arxiv.org/pdf/1902.04103.pdf eqn 3
</span>  <span class="n">cp</span><span class="p">,</span> <span class="n">cn</span> <span class="o">=</span> <span class="n">smooth_BCE</span><span class="p">(</span><span class="n">eps</span><span class="o">=</span><span class="mf">0.0</span><span class="p">)</span>

  <span class="c1"># focal loss
</span>  <span class="n">g</span> <span class="o">=</span> <span class="n">h</span><span class="p">[</span><span class="s">'fl_gamma'</span><span class="p">]</span>  <span class="c1"># focal loss gamma
</span>  <span class="k">if</span> <span class="n">g</span> <span class="o">&gt;</span> <span class="mi">0</span><span class="p">:</span>
      <span class="n">BCEcls</span><span class="p">,</span> <span class="n">BCEobj</span> <span class="o">=</span> <span class="n">FocalLoss</span><span class="p">(</span><span class="n">BCEcls</span><span class="p">,</span> <span class="n">g</span><span class="p">),</span> <span class="n">FocalLoss</span><span class="p">(</span><span class="n">BCEobj</span><span class="p">,</span> <span class="n">g</span><span class="p">)</span>

  <span class="c1"># per output
</span>  <span class="k">for</span> <span class="n">i</span><span class="p">,</span> <span class="n">pi</span> <span class="ow">in</span> <span class="nb">enumerate</span><span class="p">(</span><span class="n">p</span><span class="p">):</span>  <span class="c1"># layer index, layer predictions
</span>      <span class="n">b</span><span class="p">,</span> <span class="n">a</span><span class="p">,</span> <span class="n">gj</span><span class="p">,</span> <span class="n">gi</span> <span class="o">=</span> <span class="n">indices</span><span class="p">[</span><span class="n">i</span><span class="p">]</span>  <span class="c1"># image_idx, anchor_idx, grid_y, grid_x
</span>      <span class="n">tobj</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">zeros_like</span><span class="p">(</span><span class="n">pi</span><span class="p">[...,</span> <span class="mi">0</span><span class="p">],</span> <span class="n">device</span><span class="o">=</span><span class="n">device</span><span class="p">)</span>  <span class="c1"># target obj
</span>
      <span class="n">nb</span> <span class="o">=</span> <span class="n">b</span><span class="p">.</span><span class="n">shape</span><span class="p">[</span><span class="mi">0</span><span class="p">]</span>  <span class="c1"># number of positive samples
</span>      <span class="k">if</span> <span class="n">nb</span><span class="p">:</span>
          <span class="c1"># 对应匹配到正样本的预测信息
</span>          <span class="n">ps</span> <span class="o">=</span> <span class="n">pi</span><span class="p">[</span><span class="n">b</span><span class="p">,</span> <span class="n">a</span><span class="p">,</span> <span class="n">gj</span><span class="p">,</span> <span class="n">gi</span><span class="p">]</span>  <span class="c1"># prediction subset corresponding to targets
</span>
          <span class="c1"># GIoU
</span>          <span class="n">pxy</span> <span class="o">=</span> <span class="n">ps</span><span class="p">[:,</span> <span class="p">:</span><span class="mi">2</span><span class="p">].</span><span class="n">sigmoid</span><span class="p">()</span>
          <span class="n">pwh</span> <span class="o">=</span> <span class="n">ps</span><span class="p">[:,</span> <span class="mi">2</span><span class="p">:</span><span class="mi">4</span><span class="p">].</span><span class="n">exp</span><span class="p">().</span><span class="n">clamp</span><span class="p">(</span><span class="nb">max</span><span class="o">=</span><span class="mf">1E3</span><span class="p">)</span> <span class="o">*</span> <span class="n">anchors</span><span class="p">[</span><span class="n">i</span><span class="p">]</span>
          <span class="n">pbox</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">cat</span><span class="p">((</span><span class="n">pxy</span><span class="p">,</span> <span class="n">pwh</span><span class="p">),</span> <span class="mi">1</span><span class="p">)</span>  <span class="c1"># predicted box
</span>          <span class="n">giou</span> <span class="o">=</span> <span class="n">bbox_iou</span><span class="p">(</span><span class="n">pbox</span><span class="p">.</span><span class="n">t</span><span class="p">(),</span> <span class="n">tbox</span><span class="p">[</span><span class="n">i</span><span class="p">],</span> <span class="n">x1y1x2y2</span><span class="o">=</span><span class="bp">False</span><span class="p">,</span> <span class="n">GIoU</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span>  <span class="c1"># giou(prediction, target)
</span>          <span class="n">lbox</span> <span class="o">+=</span> <span class="p">(</span><span class="mf">1.0</span> <span class="o">-</span> <span class="n">giou</span><span class="p">).</span><span class="n">mean</span><span class="p">()</span>  <span class="c1"># giou loss
</span>
          <span class="c1"># Obj
</span>          <span class="n">tobj</span><span class="p">[</span><span class="n">b</span><span class="p">,</span> <span class="n">a</span><span class="p">,</span> <span class="n">gj</span><span class="p">,</span> <span class="n">gi</span><span class="p">]</span> <span class="o">=</span> <span class="p">(</span><span class="mf">1.0</span> <span class="o">-</span> <span class="n">model</span><span class="p">.</span><span class="n">gr</span><span class="p">)</span> <span class="o">+</span> <span class="n">model</span><span class="p">.</span><span class="n">gr</span> <span class="o">*</span> <span class="n">giou</span><span class="p">.</span><span class="n">detach</span><span class="p">().</span><span class="n">clamp</span><span class="p">(</span><span class="mi">0</span><span class="p">).</span><span class="nb">type</span><span class="p">(</span><span class="n">tobj</span><span class="p">.</span><span class="n">dtype</span><span class="p">)</span>  <span class="c1"># giou ratio
</span>
          <span class="c1"># Class
</span>          <span class="k">if</span> <span class="n">model</span><span class="p">.</span><span class="n">nc</span> <span class="o">&gt;</span> <span class="mi">1</span><span class="p">:</span>  <span class="c1"># cls loss (only if multiple classes)
</span>              <span class="n">t</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">full_like</span><span class="p">(</span><span class="n">ps</span><span class="p">[:,</span> <span class="mi">5</span><span class="p">:],</span> <span class="n">cn</span><span class="p">,</span> <span class="n">device</span><span class="o">=</span><span class="n">device</span><span class="p">)</span>  <span class="c1"># targets
</span>              <span class="n">t</span><span class="p">[</span><span class="nb">range</span><span class="p">(</span><span class="n">nb</span><span class="p">),</span> <span class="n">tcls</span><span class="p">[</span><span class="n">i</span><span class="p">]]</span> <span class="o">=</span> <span class="n">cp</span>
              <span class="n">lcls</span> <span class="o">+=</span> <span class="n">BCEcls</span><span class="p">(</span><span class="n">ps</span><span class="p">[:,</span> <span class="mi">5</span><span class="p">:],</span> <span class="n">t</span><span class="p">)</span>  <span class="c1"># BCE
</span>
          <span class="c1"># Append targets to text file
</span>          <span class="c1"># with open('targets.txt', 'a') as file:
</span>          <span class="c1">#     [file.write('%11.5g ' * 4 % tuple(x) + '\n') for x in torch.cat((txy[i], twh[i]), 1)]
</span>
      <span class="n">lobj</span> <span class="o">+=</span> <span class="n">BCEobj</span><span class="p">(</span><span class="n">pi</span><span class="p">[...,</span> <span class="mi">4</span><span class="p">],</span> <span class="n">tobj</span><span class="p">)</span>  <span class="c1"># obj loss
</span>
  <span class="c1"># 乘上每种损失的对应权重
</span>  <span class="n">lbox</span> <span class="o">*=</span> <span class="n">h</span><span class="p">[</span><span class="s">'giou'</span><span class="p">]</span>
  <span class="n">lobj</span> <span class="o">*=</span> <span class="n">h</span><span class="p">[</span><span class="s">'obj'</span><span class="p">]</span>
  <span class="n">lcls</span> <span class="o">*=</span> <span class="n">h</span><span class="p">[</span><span class="s">'cls'</span><span class="p">]</span>

  <span class="c1"># loss = lbox + lobj + lcls
</span>  <span class="k">return</span> <span class="p">{</span><span class="s">"box_loss"</span><span class="p">:</span> <span class="n">lbox</span><span class="p">,</span>
          <span class="s">"obj_loss"</span><span class="p">:</span> <span class="n">lobj</span><span class="p">,</span>
          <span class="s">"class_loss"</span><span class="p">:</span> <span class="n">lcls</span><span class="p">}</span>
</pre></td></tr></tbody></table></code></pre></div>    </div>
    <p>该函数就是计算目标框损失、置信度损失和类别损失。其中负样本只会计算置信度损失。目标框损失的计算方式是计算预测的目标框参数和正样本的目标框参数的gIoU值。置信度损失的计算方式是首先根据每一个预测特征图的大小创建一个tobj变量，它的形状是[B,num_anchor,h,w],为属于正样本的位置赋值<code class="language-plaintext highlighter-rouge">tobj[b, a, gj, gi] = (1.0 - model.gr) + model.gr * giou.detach().clamp(0).type(tobj.dtype)</code>,其它位置不变保持0值，然后计算它的交叉熵损失<code class="language-plaintext highlighter-rouge">lobj += BCEobj(pi[..., 4], tobj)</code>。对于类别损失，当类别大于1时才进行计算，首先构建一个<code class="language-plaintext highlighter-rouge">t</code>变量，它的形状是[None,80]，因为这里None表示根据实际的正样本数决定，而80则是当前使用coco数据集有80个类别，它的值默认赋值cn就是负样本的值。之后通过<code class="language-plaintext highlighter-rouge">t[range(nb),tcls[i]]=cp</code>将正样本对应的那个类别的值赋值为cp，这样就构建好了类别标签，通过<code class="language-plaintext highlighter-rouge">lcls+=BCEcls(ps[:,5:],t)</code>计算类别的交叉熵损失。最后以字典形式返回这三个损失，在返回之前还要乘上每种损失对应的权重。</p>
  </li>
</ul>

<h2 id="ultralytics源码解析">Ultralytics源码解析</h2>
<h3 id="损失计算-1">损失计算</h3>
<div class="language-plaintext highlighter-rouge"><div class="highlight"><pre class="highlight"><code><table class="rouge-table"><tbody><tr><td class="rouge-gutter gl"><pre class="lineno">1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
</pre></td><td class="rouge-code"><pre>在Ultralytics的YOLO11目标检测网络中，损失计算使用的是v8DetectionLoss类。YOLOv8的损失计算包括了类别损失(原来的置信度损失现在合并到了类别损失中，由类别损失来执行YOLOv3中类别和置信度的责任)、DFL损失、Bbox损失，DFL的核心思想是模型不去直接预测边界框的偏移量，而是预测偏移量在离散区间上的分布。再通过分布加权求和得到最终的偏移量，其目标就是为了更精准地预测边界框。比如，它会将偏移量等分成16等份，然后去预测偏移量落在这16等份离散区间上的分布值。 ```python
def __call__(self, preds: Any, batch: Dict[str, torch.Tensor]) -&gt; Tuple[torch.Tensor, torch.Tensor]:
    """Calculate the sum of the loss for box, cls and dfl multiplied by batch size."""
    loss = torch.zeros(3, device=self.device)  # box, cls, dfl
    feats = preds[1] if isinstance(preds, tuple) else preds
    pred_distri, pred_scores = torch.cat([xi.view(feats[0].shape[0], self.no, -1) for xi in feats], 2).split(
        (self.reg_max * 4, self.nc), 1
    )

    pred_scores = pred_scores.permute(0, 2, 1).contiguous()
    pred_distri = pred_distri.permute(0, 2, 1).contiguous()

    dtype = pred_scores.dtype
    batch_size = pred_scores.shape[0]
    imgsz = torch.tensor(feats[0].shape[2:], device=self.device, dtype=dtype) * self.stride[0]  # image size (h,w)
    anchor_points, stride_tensor = make_anchors(feats, self.stride, 0.5)

    # Targets
    targets = torch.cat((batch["batch_idx"].view(-1, 1), batch["cls"].view(-1, 1), batch["bboxes"]), 1)
    targets = self.preprocess(targets.to(self.device), batch_size, scale_tensor=imgsz[[1, 0, 1, 0]])
    gt_labels, gt_bboxes = targets.split((1, 4), 2)  # cls, xyxy
    mask_gt = gt_bboxes.sum(2, keepdim=True).gt_(0.0)

    # Pboxes
    pred_bboxes = self.bbox_decode(anchor_points, pred_distri)  # xyxy, (b, h*w, 4)
    # dfl_conf = pred_distri.view(batch_size, -1, 4, self.reg_max).detach().softmax(-1)
    # dfl_conf = (dfl_conf.amax(-1).mean(-1) + dfl_conf.amax(-1).amin(-1)) / 2

    _, target_bboxes, target_scores, fg_mask, _ = self.assigner(
        # pred_scores.detach().sigmoid() * 0.8 + dfl_conf.unsqueeze(-1) * 0.2,
        pred_scores.detach().sigmoid(),
        (pred_bboxes.detach() * stride_tensor).type(gt_bboxes.dtype),
        anchor_points * stride_tensor,
        gt_labels,
        gt_bboxes,
        mask_gt,
    )

    target_scores_sum = max(target_scores.sum(), 1)

    # Cls loss
    # loss[1] = self.varifocal_loss(pred_scores, target_scores, target_labels) / target_scores_sum  # VFL way
    loss[1] = self.bce(pred_scores, target_scores.to(dtype)).sum() / target_scores_sum  # BCE

    # Bbox loss
    if fg_mask.sum():
        target_bboxes /= stride_tensor
        loss[0], loss[2] = self.bbox_loss(
            pred_distri, pred_bboxes, anchor_points, target_bboxes, target_scores, target_scores_sum, fg_mask
        )

    loss[0] *= self.hyp.box  # box gain
    loss[1] *= self.hyp.cls  # cls gain
    loss[2] *= self.hyp.dfl  # dfl gain

    return loss * batch_size, loss.detach()  # loss(box, cls, dfl) ``` 1. 初始化损失张量
计算损失时，首先会创建一个loss变量分别存储「边界框损失、分类 + 置信度损失、DFL 损失」。 2. 提取模型预测特征    feats是模型的预测输出，它是一个list有3个元素，分别是三个检测头的预测输出。feates=[ (B, nc+reg_max*4, H1, W1), (B, nc+reg_max*4, H2, W2), (B, nc+reg_max*4, H3, W3) ]，其中B是批量大小，h和w是特征图的高度和宽度，4是目标框的参数数量，reg_max是DFL的离散区间数量。 3. 特征图展平+分离坐标分布和分类分数    pred_distri, pred_scores = torch.cat([xi.view(feats[0].shape[0], self.no, -1) for xi in feats], 2).split(
(self.reg_max * 4, self.nc), 1 )
将预测图的高宽展平为一个维度，并将不同尺度的特征图的预测结果按展平的维度拼接起来，然后将坐标分布和分类分数分离出来。pred_distri的形状为(B,坐标分布通道数,总锚点数)，pred_scores的形状为(B, 类别数, 总锚点数)。 4. 调整通道维度顺序，适配后续计算    将pred_distri和pred_scores的总锚点数和坐标分布通道数/类别数的维度位置进行交换，使其更合理，按照批次维度-锚点维度-坐标分布通道数/类别数维度进行索引。 5. 计算输入图像尺寸+生成锚点    这一步和YOLOv3SPP中的YOLOLayer那一步的目的是一样的，都是生成anchor。不过，虽说都是生成anchor,但是YOLOv8是free anchor的，即每一个网格都只有一个坐标框的相关参数预测，而YOLOv3SPP中每个网格有多个坐标框的相关参数预测（YOLOv3中有一个anchor_num的维度，YOLOv8中没有该维度）。    首先会根据当前第一个特征图的最后两个维度(H,W)乘以对应特征图的stride，复原回原图的尺寸大小。比如(64,64)x8-&gt;(512,512)。然后会调用make_anchors(feats,self.stride,0.5)来生成所有尺度的anchor中心点坐标(0.5为步距)和对应的步长张量。    sy,sx=torch.meshgrid(sy,sx,indexing='ij')，其中sy代表行，sx代表列，因为计算机中原点是在左上角，往右是x轴，往下是y轴，所以sx是列索引，sy是行索引。而我们常用的索引方式是先说行，再说列，第几行第几列，所以sy,sx这样的索引方式是符合我们的习惯的。但是在数学上，是按照(x,y)的方式进行索引的，就是先第几列，再第几行。所以，meshgrid对每一个网格生成它是第几行的索引坐标集合和它是第几列的坐标集合，所以sy代表每一个网格是第几行，它的形状是[h,w]，第一个维度的元素是[h_i,h_i,...,h_i]，代表这一行w个网络它是第几行，同理，sx的形状是[h,w]，第一个维度的元素是[w_1,w_2,...,w_w]，代表这一行w个网格它是第几列。torch.stack((sx,sy),-1).view(-1,2)是按照(x,y)的索引方式在列的维度[h,w]进行拼接，再调整形状为[锚点的数量,2]，这样就得到所有锚点它的中心点(x,y)坐标索引了。stride_tensor.append(torch.full((h * w, 1), stride, dtype=dtype, device=device))为每一个锚点生成该特征尺度下相同的stride值。    最后输出的锚点和步距张量的形状都是二维的，其中anchor的形状是[总的锚点数量,2(x,y)]，stride的形状是[总的锚点数量，1(stride_value)]。 6. 处理目标标签    现在我们需要对真值标签进行处理。首先将batch_idx(代表目标所在的图片索引批次号)，cls(类别),bbox(xywh)沿着第一维进行拼接，所以我们得到的targets的维度是二维的，形状是[num_target,6]。    然后在process函数中传入targets，batch_size和scale_tensor。如果当前batch中没有目标，则直接返回一个全为0的张量，形状是(batch_size,0,ne-1)，而有目标的话，会先创建一个out张量默认值为0，形状是(batch_size,counts.max(),ne-1)，这里因为用了batch_size维度，所以原来targets中第一维的第一个元素batch_idx，就不用再放入了，这里的counts是来自于targets第0维度的所有目标的所属图片批次的索引，它表示每一个图片批次中有多少个目标存在。比如第一个批次图片索引中，有17个目标，第二个批次图片索引中有2个目标等等。而counts.max()则是找到所有这些图片中，在一张图片上存在的最大目标的数量。并以此为准，创建了counts.max()这个维度，out[j,:n]=targets[matches,1:]就是说将每张图片中存在的目标的除了批次号的元素(包括类别和目标框参数)都进行赋值，所有小于counts.max()的批次号，它只会填充前n个，其它的则默认为0，而遇到counts.max()所在的批次号时，就会全部填满。执行完这个后，就将原来的targets由目标数量的索引形式转为了以批次进行索引的形式。再将out中的目标框由(xywh)的归一化表示转为(xyxy)的表示形式，并乘以scale_tensor还原回最初输入网络的时候图片的大小。    所以process的作用就是将原来的targets的索引方式由目标数量的索引形式转为了以批次进行索引的形式，并将xywh表示转为xyxy表示，然后恢复成原图大小。    gt_labels和gt_bboxes分别来自于处理后的targets的第2维度的第一个元素和其余元素。因为在process中第1个维度采用了counts.max()，所以其它图片中目标数小于最大图片的目标数时，会有默认的0值出现，所以需要mask_gt=gt_bboxes.sum(2,keepdim=True).gt_(0.)来得到一个有效的掩膜，标记有效目标。 7. 解码预测边界框参数    因为预测的边界框参数是偏移量的分布，所以需要对其进行解码，从分布得到真实偏移量。pred_bboxes=self.bbox_decode(anchor_points,pred_distri)，其中anchor_points是所有锚点的中心点坐标形状是[总的锚点数量,2(x,y)]，pred_distri是预测的边界框参数分布形状是[批量大小,总的锚点数量，坐标分布通道数]。YOLOv8默认使用dfl，pred_dist = pred_dist.view(b, a, 4, c // 4).softmax(3).matmul(self.proj.type(pred_dist.dtype))，它将4个方向(ltrb)单独成立了一个维度，然后按照这个维度计算softmax，得到概率分布值，相加为1。然后通过矩阵相乘计算每个离散区间的加权和，权重就是刚才计算得到的概率分布值。在这里,self.proj分成了16个离散区间，其值从[0-15]变化，经过这样的计算,pred_dist的形状是[b,a,4]，就得到了(ltrb)四个方向的偏移量，dist2bbox(pred_dist, anchor_points, xywh=False)的作用就是将(ltrb)四个方向的偏移量转为(xyxy)坐标的形式。    所以，返回的pred_bboxes的形状是从[b,a,reg_max*4]变为[b,a,4]，它的第2维度的每个元素代表一个锚点的(xyxy)坐标表示。 8.  正负样本匹配
和YOLOv3SPP中一样，也要进行正负样本匹配。
```python
@torch.no_grad()
def forward(self, pd_scores, pd_bboxes, anc_points, gt_labels, gt_bboxes, mask_gt):
    """
    Compute the task-aligned assignment.

    Args:
        pd_scores (torch.Tensor): Predicted classification scores with shape (bs, num_total_anchors, num_classes).
        pd_bboxes (torch.Tensor): Predicted bounding boxes with shape (bs, num_total_anchors, 4).
        anc_points (torch.Tensor): Anchor points with shape (num_total_anchors, 2).
        gt_labels (torch.Tensor): Ground truth labels with shape (bs, n_max_boxes, 1).
        gt_bboxes (torch.Tensor): Ground truth boxes with shape (bs, n_max_boxes, 4).
        mask_gt (torch.Tensor): Mask for valid ground truth boxes with shape (bs, n_max_boxes, 1).

    Returns:
        target_labels (torch.Tensor): Target labels with shape (bs, num_total_anchors).
        target_bboxes (torch.Tensor): Target bounding boxes with shape (bs, num_total_anchors, 4).
        target_scores (torch.Tensor): Target scores with shape (bs, num_total_anchors, num_classes).
        fg_mask (torch.Tensor): Foreground mask with shape (bs, num_total_anchors).
        target_gt_idx (torch.Tensor): Target ground truth indices with shape (bs, num_total_anchors).

    References:
        https://github.com/Nioolek/PPYOLOE_pytorch/blob/master/ppyoloe/assigner/tal_assigner.py
    """
    self.bs = pd_scores.shape[0]
    self.n_max_boxes = gt_bboxes.shape[1]
    device = gt_bboxes.device

    if self.n_max_boxes == 0:
        return (
            torch.full_like(pd_scores[..., 0], self.num_classes),
            torch.zeros_like(pd_bboxes),
            torch.zeros_like(pd_scores),
            torch.zeros_like(pd_scores[..., 0]),
            torch.zeros_like(pd_scores[..., 0]),
        )

    try:
        return self._forward(pd_scores, pd_bboxes, anc_points, gt_labels, gt_bboxes, mask_gt)
    except torch.cuda.OutOfMemoryError:
        # Move tensors to CPU, compute, then move back to original device
        LOGGER.warning("CUDA OutOfMemoryError in TaskAlignedAssigner, using CPU")
        cpu_tensors = [t.cpu() for t in (pd_scores, pd_bboxes, anc_points, gt_labels, gt_bboxes, mask_gt)]
        result = self._forward(*cpu_tensors)
        return tuple(t.to(device) for t in result) ```
需要注意的是，传入的pd_scores经过了sigmoid()激活，而pd_bboxes和anc_points都还原回了原图大小下尺寸表示，所以涉及到的这些计算都是在原图尺寸下的计算。如果当前所有批次中，没有目标(counts.max()==0)，则直接返回全0的target_labels、target_bboxes、target_scores、fg_mask、target_gt_idx。否则，调用_forward()方法进行正负样本匹配。
```python
    def _forward(self, pd_scores, pd_bboxes, anc_points, gt_labels, gt_bboxes, mask_gt):
    """
    Compute the task-aligned assignment.

    Args:
        pd_scores (torch.Tensor): Predicted classification scores with shape (bs, num_total_anchors, num_classes).
        pd_bboxes (torch.Tensor): Predicted bounding boxes with shape (bs, num_total_anchors, 4).
        anc_points (torch.Tensor): Anchor points with shape (num_total_anchors, 2).
        gt_labels (torch.Tensor): Ground truth labels with shape (bs, n_max_boxes, 1).
        gt_bboxes (torch.Tensor): Ground truth boxes with shape (bs, n_max_boxes, 4).
        mask_gt (torch.Tensor): Mask for valid ground truth boxes with shape (bs, n_max_boxes, 1).

    Returns:
        target_labels (torch.Tensor): Target labels with shape (bs, num_total_anchors).
        target_bboxes (torch.Tensor): Target bounding boxes with shape (bs, num_total_anchors, 4).
        target_scores (torch.Tensor): Target scores with shape (bs, num_total_anchors, num_classes).
        fg_mask (torch.Tensor): Foreground mask with shape (bs, num_total_anchors).
        target_gt_idx (torch.Tensor): Target ground truth indices with shape (bs, num_total_anchors).
    """
    mask_pos, align_metric, overlaps = self.get_pos_mask(
        pd_scores, pd_bboxes, gt_labels, gt_bboxes, anc_points, mask_gt
    )

    target_gt_idx, fg_mask, mask_pos = self.select_highest_overlaps(mask_pos, overlaps, self.n_max_boxes)

    # Assigned target
    target_labels, target_bboxes, target_scores = self.get_targets(gt_labels, gt_bboxes, target_gt_idx, fg_mask)

    # Normalize
    align_metric *= mask_pos
    pos_align_metrics = align_metric.amax(dim=-1, keepdim=True)  # b, max_num_obj
    pos_overlaps = (overlaps * mask_pos).amax(dim=-1, keepdim=True)  # b, max_num_obj
    norm_align_metric = (align_metric * pos_overlaps / (pos_align_metrics + self.eps)).amax(-2).unsqueeze(-1)
    target_scores = target_scores * norm_align_metric

    return target_labels, target_bboxes, target_scores, fg_mask.bool(), target_gt_idx ```
self.get_pos_mask的作用是为每一个真值box获取正样本掩膜
1. 调用select_candidates_in_gts方法，选取落在在真值框内的正锚点中心点mask_in_gts。
   1. xy_centers[None]的作用是给锚点中心增加维度，shape从[n_anchors,2]变为了[1,n_anchors,2]，再通过广播与lt([bs*counts.max(),1,2])匹配，得到[bs*counts.max(),n_anchors,2]其含义是每个真实框与每个锚点的中心x1差、中心y1差。
   2. rb-xy_centers[None]同理，得到每个真实框与每个锚点的中心x2差、中心y2差。
   3. torch.cat(...,dim=2)，拼接后坐标维度变成4，shape[bs*n_boxes,n_anchors,4]这4个值分别是x_center-x1,y_center-y1,x2-x_center,y2-y_center。
   4. 最后view恢复批次和真实框的维度，最终bbox_deltas的shape为[bs,n_boxes,n_anchors,4]。
   5. return bbox_deltas.amin(3).gt_(eps)。测分数
      1. 最终判断锚点是否在真实框内，amin(3)表示对第3维度(4个偏移差)取最小值，每个批次、真实框、锚点对应一个值，代表4个偏移差中最小的那个。
      2. gt_(eps)，判断这个最小值是否大小eps极小值，避免数值误差
      3. 最终返回shape为[bs,n_boxes(counts.max()),n_anchors]的布尔张量，True表示该锚点的4个偏移差均为正，即锚点中心在真实框的内部：x1&lt;x_center&lt;x2且y1&lt;y_center&lt;y2，False表示锚点在真实框外部或者边界处。
   6. 总结下来，该方法就是判断锚点是否是在真实框的内部，而里面的n_boxes其实是之前构建真实框时用的counts.max()，所有xyxy为默认0值的当然最终其布尔值为False。
2. 调用get_box_metrics方法，基于预测和真值框，获取align_metric,overlaps。
   1. mask_gt是有效的gt框掩膜，shape为[bs,n_max_boxes,n_anchors]，True表示该真实框-anchor对是有效的，False表示无效的gt-anchor对。
   2. overlaps是用来记录gt与anchor的重叠度(IOU)，shape为[bs,n_max_boxes,n_anchors]，记录真实框-anchor对的IOU值。
   3. bbox_scores是用来记录每个真实框-anchor对的类别预测分数，shape为[bs,n_max_boxes,n_anchors]，记录每个真实框-anchor对的类别预测分数。
   4. 构建类别索引，提取对应的预测分数
      1. ind=torch.zeros([2, self.bs, self.n_max_boxes], dtype=torch.long)
      2. ind[0]=torch.arange(self.bs)[:,None].expand(-1,self.n_max_boxes),shape是[bs,n_max_boxes]，其值是[[0,...,0],...,[bs-1,...,bs-1]]这样的。
      3. ind[1]=gt_labels.long().squeeze(-1)，shape是[bs,n_max_boxes]
      4. bbox_scores[mask_gt]=pd_scores[ind[0], :, ind[1]][mask_gt]，在有效的gt-anchor对处，赋值为pd_scores中对应的值，shape是[bs,n_max_boxes,n_anchors]
   5. 计算GT和anchor的IOU，pd_boxes = pd_bboxes.unsqueeze(1).expand(-1, self.n_max_boxes, -1, -1)[mask_gt]，shape是[N,4]，表示N个有效gt-anchor对的预测框坐标。
      1. gt_boxes = gt_bboxes.unsqueeze(2).expand(-1, -1, na, -1)[mask_gt]，shape是[N,4]，表示N个有效gt-anchor对的真实框坐标
      2. overlaps[mask_gt]=self.iou_calculation(gt_boxes, pd_boxes)，计算这些有效gt-anchor对的CIOU值
   6. align_metric = bbox_scores.pow(self.alpha) * overlaps.pow(self.beta)，计算对齐度量，将类别预测分数和IOU分别加权后相乘，综合反映“预测框与GT的匹配程度”，即分数越高、IOU越大，匹配度越高。
   7. 输出align_metric(用于正负样本分配或损失加权)和overlaps(原始的IOU值)，它们的shape都是[bs,n_max_boxes,n_anchors]
   8. 总结下，该方法的核心目的就是为每个真是目标框找到最匹配的预测框，给后续的正负样本划分或损失计算提供依据。
3. mask_topk = self.select_topk_candidates(align_metric, topk_mask=mask_gt.expand(-1, -1, self.topk).bool())，获取前topk个度量metric的掩膜
4. mask_pos = mask_topk * mask_in_gts * mask_gt，获取最终的正样本掩膜
5. 返回mask_pos, align_metric, overlaps，它们的形状都是[bs,n_max_boxes,n_anchors]
self.select_highest_overlaps的作用是当多个真实框对应同一个锚框时，根据overlaps选择IOU最大的那个真实框，返回其索引、前景掩膜、正样本掩膜。
self.get_targets的作用是根据选择的真实框索引，获取对应的目标标签、目标框、目标分数。
最终，输出target_bboxes(正样本的目标框),target_scores(正样本=类别概率，负样本=0),fg_mask(正样本掩码，True=正样本，False=负样本)
</pre></td></tr></tbody></table></code></pre></div></div>

<h2 id="多标签分类与多分类">多标签分类与多分类</h2>
<div class="language-plaintext highlighter-rouge"><div class="highlight"><pre class="highlight"><code><table class="rouge-table"><tbody><tr><td class="rouge-gutter gl"><pre class="lineno">1
2
3
4
5
6
7
8
9
</pre></td><td class="rouge-code"><pre>我一直没有搞明白BCELoss和CrossEntropyLoss它们有什么区别。现在终于动了动我的小脑筋，克服了拖延症，细细分析下它们的区别，以及延申到YOLO目标检测中使用多标签分类，这也是把它们放在这里叙述的原因。

首先我们看BCELoss。顾名思义，就是二值交叉熵损失，它用于二分类任务，计算预测值与真实标签之间的差异。它的公式如下： L(y, \hat{y}) = - \left[ y \cdot \log(\hat{y}) + (1 - y) \cdot \log(1 - \hat{y}) \right] y的取值范围是{0,1}，当y取1时，公式简化为`L = -\log(\hat{y})`，同样地当y取0时，公式简化为`L = -\log(1 - \hat{y})`。 对于N个样本，我们取平均值 `L = -\frac{1}{N} \sum_{i=1}^{N} \left[ y_i \cdot \log(\hat{y}_i) + (1 - y_i) \cdot \log(1 - \hat{y}_i) \right]` 并且，实际使用时常常讲sigmoid函数和BCE损失计算进行合并，BCEWithLogitsLoss，以求数值稳定。 此外，我们还可以使用带权重的变体形式，以平衡二分类下正负样本的类别不平衡问题。 `L = -\frac{1}{N} \sum_{i=1}^{N} \left[ \omega \cdot y_i \cdot \log(\hat{y}_i) + (1 - \omega) \cdot (1 - y_i) \cdot \log(1 - \hat{y}_i) \right]`
接下来我们看多类别交叉熵损失。它是针对多分类任务（输出≥3 类，且样本 “互斥唯一”，即一个样本只能属于一类）设计的损失函数，本质是基于多项式分布（单个样本有 K 种输出可能，且概率和为 1）的交叉熵计算。它的激活函数是softmax，将模型输出的原始分数转换为概率分布，并且概率之和为1。
当只有一个样本时，它的形式是
`L(y, \hat{y}) = - \sum_{k=1}^{K} y_k \cdot \log(\hat{y}_k)`
而因为采用独特编码，所以当有多个样本时，公式可以简化成
`L = -\frac{1}{N} \sum_{i=1}^{N} \log(\hat{y}_{i, y_i})`
不要因为有多类别，就混淆了多标签分类(multi label classification)和多分类(multi-class classification)。多标签分类是指每个样本可以属于多个类别，而多分类是指每个样本只能属于一个类别。例如，一个图片可以同时包含猫、狗和鸟，这是一个多标签分类问题。而一个图片只能是猫、狗或鸟中的一种，这是一个多分类问题。在YOLO目标检测中计算目标的类别损失时，其实就是一个多标签分类问题，使用的是BCELoss，对每一个类别单独计算它的二值交叉熵损失。对于80类的检测任务，YOLO不会让模型输出“80个概率和为1”的结果，而是让模型输出80个独立的[0,1]值——每个值表示“该框属于这个类别的概率”，彼此独立、互不影响（比如一个框可以同时输出“人0.95”、“车：0.02”、“猫：0.01”，无需求和为1）它的计算逻辑是对每一个类别单独计算BCE损失，再求平均。
</pre></td></tr></tbody></table></code></pre></div></div>

<h2 id="yolo中类别损失计算时怎么为不同类别设置不同的权重值">YOLO中类别损失计算时怎么为不同类别设置不同的权重值</h2>
<div class="language-plaintext highlighter-rouge"><div class="highlight"><pre class="highlight"><code><table class="rouge-table"><tbody><tr><td class="rouge-gutter gl"><pre class="lineno">1
</pre></td><td class="rouge-code"><pre>在YOLO中，类别损失计算时，为了平衡不同类别之间的差异，通常会为每个类别设置不同的权重值。这是因为不同类别的样本数量可能不均衡，某些类别可能出现的频率更高，而某些类别可能出现的频率更低。为了使模型更关注出现频率较低的类别，我们可以为这些类别设置较高的权重值。在二分类问题中，常常出现正负样本比例不均衡的问题，这时就需要为正负样本设置不同的权重值来调整它们的损失值。在YOLO中的类别损失计算是一个多标签分类问题，它们是单独计算每一个类别的二值交叉熵损失的，根据在nn.BCEWithLogits中的posi_weight这个参数，可以为每一个类别设置不同的权重值（针对该类别下的二值交叉熵损失计算的正样本赋予该权重）。比如数据集中有三个类别，人、猫和狗，发现人的样本数量特别少，无论是人相对于猫和狗，又或者相对于在人的二分类下的正负样本，人的样本数量都很少，就可以设置pos_weight=[3,1,1]，为人的正样本在计算BCE时设置3倍权重值。以平衡类别不平衡的问题。YOLO11 在多标签场景下，输出层会为每个类别生成独立的logits，给每个类别配置对应的pos_weight实现逐类损失加权。pos_weight的计算依据是该类别自身的正负样本比例（人类正样本数/人类负样本数），而非跨类别对比（人类样本数/猫样本数）。整个任务是多标签（同时识别人和猫狗），但每个类内部是独立二分类（人：有 / 无；猫：有 / 无），pos_weight 只针对 “单类内部” 的正负失衡，而非整个任务的正负失衡。
</pre></td></tr></tbody></table></code></pre></div></div>]]></content><author><name></name></author><category term="人工智能" /><summary type="html"><![CDATA[YOLO-目标检测 TLTR 1-学习YOLO的过程 2-Andrew Ng深度学习课程中关于YOLO的讲解 3-霹雳吧啦Wz的讲解，https://www.bilibili.com/video/BV1yi4y1g7ro/ 学习YOLO的过程 我是从25年7月份正式开始学习YOLO目标检测算法的，在这之前听了吴恩达深度学习课程中的关于YOLO的讲解，但是当时没有听懂，所以相当于重新开始。学习的是霹雳吧啦Wz的讲解，首先介绍了目标检测算法的两个方案：two-stage和one-stage。two-stage的方案包括先验框预测和分类器，代表是Faster R-CNN，而one-stage的方案则是直接预测目标框和分类，代表是YOLO，以及DETR等结合了注意力机制的模型，不过我现在主要是学习YOLO，DETR没有接触过。学习了YOLOv1-v3，其中YOLO-V3还有一个SPP的版本，这也是我跟随学习代码实现的版本，目前已经完成了YOLOV3-SPP的源码学习。当然，在这中间我还穿插着复习过Andrew Ng深度学习中涉及到的YOLO知识， [还发现了别人总结好的网页](http://www.ai-start.com/dl2017/) Andrew Ng 深度学习课程中关于YOLO的讲解 课程中，Andrew 首先将目标定位分解为了两个子问题：定位和分类。分类问题很简单，判断输入的图片中的物体是什么类型，将分类和定位结合起来，除了要判断物体的类型，还要判断物体的位置，这两个问题都是适用于单一目标，即每个图片中只有一个目标。进一步提出目标检测，图中存在着多个目标，我们都需要将它们分类和定位，这是一个多目标问题。 接着从分类的输出中，扩充定位信息，普通的分类输出是一个经过softmax的向量，对应物体的类别数量，需要定位我们只需要在输出中额外增加4个值，用来表示物体的边框信息。有了输出，我们需要定位监督学习的目标标签，目标标签是一个向量，包括一个Pc用来表示是否含有对象（可以理解为置信度），4个表示边框的值，物体类别数量的值，在物体类别中还可以使用一个值来表示，这个值从1-n变化。 Andrew 将4个边框值的表示扩充到了特征点检测问题，广义上神经网络可以输出图片上特征点的坐标，在目标检测中输出4个值表示边框信息，需要检测特征点可以设置任意需要统一输出特征点的数量，然后制作目标标签进行训练。 基于滑动窗口的目标检测算法 对图片进行裁剪，以识别汽车为例，只保留包含汽车的区域，其它区域裁剪掉，然后对裁剪后的图片进行分类，判断是否为汽车，输出y=0|1，这样就训练出了一个分类网络，对于多类别的同理输出y=0|1|…|n。网络训练好之后就可以基于滑动窗口来实现目标检测了，其中的定位问题，在滑动的时候就已经内含在其中了，即只要在当前窗口中识别出了一个类别，那么这个窗口的位置就是该类别的位置，我们可以记录窗口的横向和纵向滑动距离。滑动窗口的实现很灵活，首先可以选择是否以要重叠，即下一次的滑动是否与本次的部分区域重叠，从图上左上角向右下滑动，通过步距进行控制。还可以控制滑动窗口的大小。总之，这些人为控制的操作都是为了有效识别出物体，比如有些物体可能会存在两次滑动窗口之间。 这也就引出了滑动窗口的问题：计算成本太高。如果用小步幅，无法准确定位图中的对象，如果用大步幅（包括多个目标），粗糙间隔会影响性能。 为了提高计算效率，将滑动窗口使用卷积来实现。你可能会想，之前的滑动窗口不就是基于卷积实现的吗？再仔细分析下，最初的滑动窗口首先通过卷积判断当前窗口中有没有待分类的物体，至于滑动窗口的移动和卷积没有一点关系，可以通过两层循环实现这个逻辑。而这里的滑动窗口的卷积实现，是指将滑动窗口的移动步骤也以卷积的方式内含实现，在卷积的过程中就相当于移动了滑动窗口，过程不同但是在结果上是等价的，以至于我们可以理解为就像移动了窗口，但这是利用了卷积计算原理和它的高效操作实现的，因为卷积窗口也是需要滑动的，不过这个滑动是在pytorh等深度学习框架中实现了的，借助了GPU的高性能计算，就是说这个滑动是优化过的，提出该方法的论文是Sermanet, Pierre, et al. “OverFeat: Integrated Recognition, Localization and Detection using Convolutional Networks.” Eprint Arxiv (2013)。其实就是将一维的全连接层替换成了多维的表示，不过这个多维表示也就表明它可以在其它维度上增加数量，从而将输出的特征层1x1xN变成mxmxn。该卷积操作的原理就是我们不需要把图像分割成子集，分别执行前向传播，而是将图像整体输入给卷积网络计算，其中有许多区域可以共享计算，最终得到输出层。换个角度就是感受野的解释，我们通过控制步距，窗口大小（这里的窗口大小不是卷积层的窗口大小，而是滑动窗口的大小，也就是在还没有使用卷积实现滑动窗口的时候，滑动窗口的大小，而这时我们可以去调整卷积层的大小，固定后，即卷积层大小固定，滑动窗口大小固定后，可以使用卷积的形式实现滑动窗口），来实现不同的下采样，在最后的输出结果中我们就可以得到感受野，最后输出的特征层大小就是相当于我们滑动了多少次窗口，它的位置映射回原图上就是目标框看的区域内容。 但是该算法仍然存在不能完美定位的问题，比如目标跨过多个区域时，一个窗口定位只能定位到部分区域，还有些目标更适合用长方形的框来定位。 YOLO YOLO算法其实也是根据感受野的解释反推在原图上的目标框位置，不过它输出了边框参数，所以在边框的形状上可以很灵活的变化学习。将原图划分成SxS个小区域，这SxS个小区域在原图上经过卷积网络后得到的输出层大小就是SxS(pixel)，所以YOLO预先定义输出层大小，通过控制输出层的尺寸来控制细粒度，以更好地识别目标。对于这SxS个输出，每个位置都包括一些参数(Pc,boxes-info,class)，通过构建对应的目标标签进行训练这样的一个目标检测网络。 这只是网络的原理，还需要结合额外的人为控制才能更好地训练网络，使用预测的边框。交并比、非极大值抑制、Pc阈值、Anchor-box。 Andrew 介绍的YOLO算法令人恍然大悟，但是对于其中的细节没有深入探究，代码部分的实战只是简单的实现，没有关于训练的部分内容，比如数据集处理、损失函数计算、正负样本选取这些。为了识别同一区域内的多个目标，增加了anchor box后，怎么确定哪个anchor box预测的是正确的，怎么根据anchor box来计算损失，这些问题没有解答。 YOLOV3-SPP源码实战 从最初的YOLOV1开始，就已经有了很多的版本，比如YOLOV2、YOLOV3、YOLOV4等，每个版本都有自己的改进，比如SPP、FPN、PAN等。而YOLOV3-SPP是在YOLOV3的基础上，增加了SPP层，用来提取不同尺度的特征，这也是我跟随学习代码实现的版本。 学习路线首先是定义网络结构，为了让网络具有扩展性，将网络结构存储在了.cfg文件中，通过读取.cfg文件来定义网络结构，这也是YOLO系列的一个特点。有了网络结构，需要解析网络结构，将其转为在内存中规则的数据结构，再根据这个数据结构来搭建网络模型。网络模型搭建好后，训练时需要导入数据集，就需要自己制作数据集入口，包括一系列的检查文件路径、预处理、缓存操作等，都定义在了一个类里面，通过继承pytorch的Dataset类，传入dataloader中。数据集加载后就是训练，每个批次经过网络输出，得到预测值，需要计算预测值与关联标签的损失，再反向传播，反向传播时又定义了调度器，学习率衰减规则，打印日志等模块内容，等等这些构成了YOLOV3-SPP算法的全部过程。在训练时，除了损失计算部分是与YOLO强关联外，其它的内容都是可以迁移通用的，比如调度器、日志打印模块。 .cfg网络定义 .cfg网络定义内容其实就是一些规则结构的字符串，描述了网络的层结构。并搭配了对应的解析器，用来解析.cfg文件，将其转为在内存中规则的数据结构。它主要有convolutional层，shortcut层，池化层、route层、upsample层、YOLO层。其中shortcut层和route层需要特别注意，不像其它层需要做具体的工作（比如进行大量计算），这两个层之所以单独独立成一个层（在编号中占据一个位置），shortcut层负责残差连接，所以它的关键字from经常等于-3，以当前shortcut层为索引0，向其上层索引3个，route层有两个作用，拼接多个层的输出和将当前的指针（代表网络当前输出所处的位置）退回到某一层（在SPP的输入多分支时有用），这是因为网络定义的顺序是线性的，即使平行结构的SPP，也需要按照线性顺序排序。YOLO层不是预测头所处的层，它是预测头的下一层（紧连着预测头），它的作用是初始化anchor模板，定义特征层的网格，判断训练还是预测，从而输出不同的结果。 自定义数据集 数据集集中了对图片和标签载入的操作预处理。包括设置批量大小，设置预处理输出图片大小、数据缓存、是否进行数据增强等操作。 损失计算 正样本计算目标框损失、分类损失和置信度损失，负样本只计算置信度损失。在compute_loss函数中，根据预测值和关联标签，计算损失。首先根据预测值和关联标签，筛选出正样本，也即在build_target函数中，筛选出正样本的类别标签（用于计算类别损失）、正样本对应的gt box信息、含有正样本的图像索引、anchor模板索引、所处的哪一个grid位置信息、anchor模板的大小信息，遍历每一个YOLO输出层，计算这三个损失。其中，置信度损失计算时首先会创建一个tobj，初始值都为0，之后对于筛选出来的正样本，根据正样本计算的iou值动态（每个批次正样本会变）去得到一个标签置信度并在tobj对应的正样本位置填上该值，其它未变的即为负样本的默认值0，然后使用预测的值和tobj进行二值交叉熵计算，得到置信度损失。可以注意到，这里并没有按照正负样本按照一定比例选取，而是计算了全部的负样本，不过使用了Focal loss、调整置信度权重等方式有效地缓解了正负样本不平衡的问题。 build_target函数解析 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 def build_targets(p, targets, model): # Build targets for compute_loss(), input targets(image_idx,class,x,y,w,h) # p: predictions [batch_size, num_anchors, grid_h, grid_w, num_params] #选出正样本 nt = targets.shape[0] tcls, tbox, indices, anch = [], [], [], [] # gain：用于将归一化坐标转换到“网格空间”的缩放因子（初始为全1，后续按层更新） gain = torch.ones(6, device=targets.device).long() # normalized to gridspace gain multi_gpu = type(model) in (nn.parallel.DataParallel, nn.parallel.DistributedDataParallel) for i, j in enumerate(model.yolo_layers): # j: [89, 101, 113] # 获取该yolo predictor对应的anchors # 注意anchor_vec是anchors缩放到对应特征层上的尺度 anchors = model.module.module_list[j].anchor_vec if multi_gpu else model.module_list[j].anchor_vec # p[i].shape: [batch_size, 3, grid_h, grid_w, num_params] gain[2:] = torch.tensor(p[i].shape)[[3, 2, 3, 2]] # xyxy gain na = anchors.shape[0] # number of anchors # [3] -&gt; [3, 1] -&gt; [3, nt] at = torch.arange(na).view(na, 1).repeat(1, nt).to('cuda') # anchor tensor, same as .repeat_interleave(nt) # Match targets to anchors a, t, offsets = [], targets * gain, 0 if nt: # 如果存在target的话 # 通过计算anchor模板与所有target的wh_iou来匹配正样本 # j: [3, nt] , iou_t = 0.20 j = (wh_iou(anchors, t[:, 4:6]) &gt; model.hyp['iou_t']).to('cuda') # iou(3,n) = wh_iou(anchors(3,2), gwh(n,2)) # t.repeat(na, 1, 1): [nt, 6] -&gt; [3, nt, 6] # 获取正样本对应的anchor模板与target信息 a, t = at[j], t.repeat(na, 1, 1)[j] # filter # Define # long等于to(torch.int64), 数值向下取整 b, c = t[:, :2].long().T # image_idx, class gxy = t[:, 2:4] # grid xy gwh = t[:, 4:6] # grid wh gij = (gxy - offsets).long() # 匹配targets所在的grid cell左上角坐标 gi, gj = gij.T # grid xy indices # Append # gain[3]: grid_h, gain[2]: grid_w # image_idx, anchor_idx, grid indices(y, x) indices.append((b, a, gj.clamp_(0, gain[3]-1), gi.clamp_(0, gain[2]-1))) tbox.append(torch.cat((gxy - gij, gwh), 1)) # gt box相对anchor的x,y偏移量以及w,h anch.append(anchors[a]) # anchors tcls.append(c) # class if c.shape[0]: # if any targets # 目标的标签数值不能大于给定的目标类别数 assert c.max() &lt; model.nc, 'Model accepts %g classes labeled from 0-%g, however you labelled a class %g. ' \ 'See https://github.com/ultralytics/yolov3/wiki/Train-Custom-Data' % ( model.nc, model.nc - 1, c.max()) return tcls, tbox, indices, anch 该函数接收预测值(P,shapes=[3,Batch,anchor_num,p_h,p_w,5+num_classes]),真值标签(targets,shapes=[num_target,6(img_idx,class_idx,x,y,w,h)],xywh是中心坐标和宽高的归一化值)和模型(model)。该函数的目标是根据模板anchor和真值标签，筛选出正样本对应的anchor和target，再结合预测值计算偏移量，返回正样本对应的类别信息、box偏移量信息、图片索引信息和anchor信息。 首先获取真值的个数num_target，每个batch可能有不同数量的目标，这个个数是动态变化的。然后创建tcls, tbox, indices, anch = [], [], [], []，以备后续添加筛选出正样本对应的信息。gain(每个1对应着targets中这些位置[img_idx,class_idx,x,y,w,h])值用于将归一化坐标转换到“网格空间”的缩放因子（初始为全1，后续按层更新），用于将targets的坐标转换到与预测值相同的网格空间，因为targets最初是采用归一化的形式的。 anchors获取到当前yolo predictor对应的anchor模板（就是原本anchor根据步距缩放到当前predictor的特征层上的尺度），形状为[3,2]，表示3个anchor模板，每个模板有2个参数（宽高）。gain参数更新x,y,w,h的缩放因子，用于将targets的坐标转换到与预测值相同的网格空间。na是anchor模板的个数，一般是3，通过展开维度得到at，它的形状是[na,num_target]，表示的意思是每个anchor模板对应当前batch中的真实目标数，因为此时我们不知道正样本它的anchor模板与目标对是哪些，有可能一个anchor模板会对应上多个目标(每个网格规定了有3个anchor，这三个anchor都源自于最初的那3个anchor模板)，有可能这些anchor模板(3个)与其中的某个目标都不符合正样本的条件。 准备好以上的信息后，就可以根据anchor模板与targets的wh_iou来筛选出正样本了。t是将targes和gain缩放系数相乘得到的缩放到当前预测特征图尺寸上的预测结果，它的归一化坐标变成了当前特征图上的坐标。如果当前batch中存在目标的话，计算anchor模板和t的IoU值，这里的代码是j = (wh_iou(anchors, t[:, 4:6]) &gt; model.hyp['iou_t']),它采用了不严格的计算方式，通过把所有的anchor和目标，把它们的中心点平移到坐标原点，再去按照它们的高宽计算IoU值，其实是一种简化的计算路径。通过与预先设定的IoU阈值比较可以得到一个布尔值掩膜j，它的形状和at是一样的，都是[na,num_target]，表示每个anchor模板与每个目标的IoU值是否大于阈值。通过这个掩膜j，就可以筛选出正样本对应的anchor模板和target信息，分别存储在a和t中,它们的形状一个是[None,]，一个是[None,6]，这里的None表示的是不清楚匹配到几个正样本，通过debug自定义训练样本此时的None等于22，而num_target的值是13，表示当前batch中有13个目标，但是筛选出来了22个正样本，所以这就验证了多个anchor也会匹配到同一个目标的情况，他们都是正样本。a中元素的值是anchor模板的索引例如[0,0,0,1,1,1,2,2,2]这样的，它对应着t中对应元素也就是此刻目标匹配的样本模板，比如a[0]正好对应t[0]。这样我们就得到了正样本(anchor模板-目标对)。 b和c它们的形状都是[None,]，表示的是正样本对应的图片索引和类别索引。gxy和gwh则是正样本的中心坐标和宽高(缩放到当前特征图尺寸)，gij通过向下取整得到的是正样本所在的网格cell的左上角坐标，gi, gj则是gij的元素，分别表示网格的y轴和x轴索引。 indices将(b,a,gj,gi)包裹到一个元组中并加入到当前列表中。记录了正样本对应的图片索引、anchor模板索引、x轴和y轴的网格索引，用来定位正样本在特征图上的所属网格位置。tbox将gt box相对于当前grid cell的x,y偏移量和宽高加入到当前列表中。anch将当前正样本对应的anchor模板加入到当前列表中。tcls将当前正样本对应的类别索引加入到当前列表中。 以上都是在一个预测特征图尺寸下的正样本筛选和信息记录，不同的预测特征图尺寸下的正样本筛选和信息记录是独立的，互不干扰。总结下我发现最重要的代码部分就是 compute_loss函数解析 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 def compute_loss(p, targets, model): # predictions, targets, model device = p[0].device lcls = torch.zeros(1, device=device) # Tensor(0) lbox = torch.zeros(1, device=device) # Tensor(0) lobj = torch.zeros(1, device=device) # Tensor(0) tcls, tbox, indices, anchors = build_targets(p, targets, model) # targets h = model.hyp # hyperparameters red = 'mean' # Loss reduction (sum or mean) # Define criteria BCEcls = nn.BCEWithLogitsLoss(pos_weight=torch.tensor([h['cls_pw']], device=device), reduction=red) BCEobj = nn.BCEWithLogitsLoss(pos_weight=torch.tensor([h['obj_pw']], device=device), reduction=red) # class label smoothing https://arxiv.org/pdf/1902.04103.pdf eqn 3 cp, cn = smooth_BCE(eps=0.0) # focal loss g = h['fl_gamma'] # focal loss gamma if g &gt; 0: BCEcls, BCEobj = FocalLoss(BCEcls, g), FocalLoss(BCEobj, g) # per output for i, pi in enumerate(p): # layer index, layer predictions b, a, gj, gi = indices[i] # image_idx, anchor_idx, grid_y, grid_x tobj = torch.zeros_like(pi[..., 0], device=device) # target obj nb = b.shape[0] # number of positive samples if nb: # 对应匹配到正样本的预测信息 ps = pi[b, a, gj, gi] # prediction subset corresponding to targets # GIoU pxy = ps[:, :2].sigmoid() pwh = ps[:, 2:4].exp().clamp(max=1E3) * anchors[i] pbox = torch.cat((pxy, pwh), 1) # predicted box giou = bbox_iou(pbox.t(), tbox[i], x1y1x2y2=False, GIoU=True) # giou(prediction, target) lbox += (1.0 - giou).mean() # giou loss # Obj tobj[b, a, gj, gi] = (1.0 - model.gr) + model.gr * giou.detach().clamp(0).type(tobj.dtype) # giou ratio # Class if model.nc &gt; 1: # cls loss (only if multiple classes) t = torch.full_like(ps[:, 5:], cn, device=device) # targets t[range(nb), tcls[i]] = cp lcls += BCEcls(ps[:, 5:], t) # BCE # Append targets to text file # with open('targets.txt', 'a') as file: # [file.write('%11.5g ' * 4 % tuple(x) + '\n') for x in torch.cat((txy[i], twh[i]), 1)] lobj += BCEobj(pi[..., 4], tobj) # obj loss # 乘上每种损失的对应权重 lbox *= h['giou'] lobj *= h['obj'] lcls *= h['cls'] # loss = lbox + lobj + lcls return {"box_loss": lbox, "obj_loss": lobj, "class_loss": lcls} 该函数就是计算目标框损失、置信度损失和类别损失。其中负样本只会计算置信度损失。目标框损失的计算方式是计算预测的目标框参数和正样本的目标框参数的gIoU值。置信度损失的计算方式是首先根据每一个预测特征图的大小创建一个tobj变量，它的形状是[B,num_anchor,h,w],为属于正样本的位置赋值tobj[b, a, gj, gi] = (1.0 - model.gr) + model.gr * giou.detach().clamp(0).type(tobj.dtype),其它位置不变保持0值，然后计算它的交叉熵损失lobj += BCEobj(pi[..., 4], tobj)。对于类别损失，当类别大于1时才进行计算，首先构建一个t变量，它的形状是[None,80]，因为这里None表示根据实际的正样本数决定，而80则是当前使用coco数据集有80个类别，它的值默认赋值cn就是负样本的值。之后通过t[range(nb),tcls[i]]=cp将正样本对应的那个类别的值赋值为cp，这样就构建好了类别标签，通过lcls+=BCEcls(ps[:,5:],t)计算类别的交叉熵损失。最后以字典形式返回这三个损失，在返回之前还要乘上每种损失对应的权重。 Ultralytics源码解析 损失计算 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 在Ultralytics的YOLO11目标检测网络中，损失计算使用的是v8DetectionLoss类。YOLOv8的损失计算包括了类别损失(原来的置信度损失现在合并到了类别损失中，由类别损失来执行YOLOv3中类别和置信度的责任)、DFL损失、Bbox损失，DFL的核心思想是模型不去直接预测边界框的偏移量，而是预测偏移量在离散区间上的分布。再通过分布加权求和得到最终的偏移量，其目标就是为了更精准地预测边界框。比如，它会将偏移量等分成16等份，然后去预测偏移量落在这16等份离散区间上的分布值。 ```python def __call__(self, preds: Any, batch: Dict[str, torch.Tensor]) -&gt; Tuple[torch.Tensor, torch.Tensor]: """Calculate the sum of the loss for box, cls and dfl multiplied by batch size.""" loss = torch.zeros(3, device=self.device) # box, cls, dfl feats = preds[1] if isinstance(preds, tuple) else preds pred_distri, pred_scores = torch.cat([xi.view(feats[0].shape[0], self.no, -1) for xi in feats], 2).split( (self.reg_max * 4, self.nc), 1 ) pred_scores = pred_scores.permute(0, 2, 1).contiguous() pred_distri = pred_distri.permute(0, 2, 1).contiguous() dtype = pred_scores.dtype batch_size = pred_scores.shape[0] imgsz = torch.tensor(feats[0].shape[2:], device=self.device, dtype=dtype) * self.stride[0] # image size (h,w) anchor_points, stride_tensor = make_anchors(feats, self.stride, 0.5) # Targets targets = torch.cat((batch["batch_idx"].view(-1, 1), batch["cls"].view(-1, 1), batch["bboxes"]), 1) targets = self.preprocess(targets.to(self.device), batch_size, scale_tensor=imgsz[[1, 0, 1, 0]]) gt_labels, gt_bboxes = targets.split((1, 4), 2) # cls, xyxy mask_gt = gt_bboxes.sum(2, keepdim=True).gt_(0.0) # Pboxes pred_bboxes = self.bbox_decode(anchor_points, pred_distri) # xyxy, (b, h*w, 4) # dfl_conf = pred_distri.view(batch_size, -1, 4, self.reg_max).detach().softmax(-1) # dfl_conf = (dfl_conf.amax(-1).mean(-1) + dfl_conf.amax(-1).amin(-1)) / 2 _, target_bboxes, target_scores, fg_mask, _ = self.assigner( # pred_scores.detach().sigmoid() * 0.8 + dfl_conf.unsqueeze(-1) * 0.2, pred_scores.detach().sigmoid(), (pred_bboxes.detach() * stride_tensor).type(gt_bboxes.dtype), anchor_points * stride_tensor, gt_labels, gt_bboxes, mask_gt, ) target_scores_sum = max(target_scores.sum(), 1) # Cls loss # loss[1] = self.varifocal_loss(pred_scores, target_scores, target_labels) / target_scores_sum # VFL way loss[1] = self.bce(pred_scores, target_scores.to(dtype)).sum() / target_scores_sum # BCE # Bbox loss if fg_mask.sum(): target_bboxes /= stride_tensor loss[0], loss[2] = self.bbox_loss( pred_distri, pred_bboxes, anchor_points, target_bboxes, target_scores, target_scores_sum, fg_mask ) loss[0] *= self.hyp.box # box gain loss[1] *= self.hyp.cls # cls gain loss[2] *= self.hyp.dfl # dfl gain return loss * batch_size, loss.detach() # loss(box, cls, dfl) ``` 1. 初始化损失张量 计算损失时，首先会创建一个loss变量分别存储「边界框损失、分类 + 置信度损失、DFL 损失」。 2. 提取模型预测特征 feats是模型的预测输出，它是一个list有3个元素，分别是三个检测头的预测输出。feates=[ (B, nc+reg_max*4, H1, W1), (B, nc+reg_max*4, H2, W2), (B, nc+reg_max*4, H3, W3) ]，其中B是批量大小，h和w是特征图的高度和宽度，4是目标框的参数数量，reg_max是DFL的离散区间数量。 3. 特征图展平+分离坐标分布和分类分数 pred_distri, pred_scores = torch.cat([xi.view(feats[0].shape[0], self.no, -1) for xi in feats], 2).split( (self.reg_max * 4, self.nc), 1 ) 将预测图的高宽展平为一个维度，并将不同尺度的特征图的预测结果按展平的维度拼接起来，然后将坐标分布和分类分数分离出来。pred_distri的形状为(B,坐标分布通道数,总锚点数)，pred_scores的形状为(B, 类别数, 总锚点数)。 4. 调整通道维度顺序，适配后续计算 将pred_distri和pred_scores的总锚点数和坐标分布通道数/类别数的维度位置进行交换，使其更合理，按照批次维度-锚点维度-坐标分布通道数/类别数维度进行索引。 5. 计算输入图像尺寸+生成锚点 这一步和YOLOv3SPP中的YOLOLayer那一步的目的是一样的，都是生成anchor。不过，虽说都是生成anchor,但是YOLOv8是free anchor的，即每一个网格都只有一个坐标框的相关参数预测，而YOLOv3SPP中每个网格有多个坐标框的相关参数预测（YOLOv3中有一个anchor_num的维度，YOLOv8中没有该维度）。 首先会根据当前第一个特征图的最后两个维度(H,W)乘以对应特征图的stride，复原回原图的尺寸大小。比如(64,64)x8-&gt;(512,512)。然后会调用make_anchors(feats,self.stride,0.5)来生成所有尺度的anchor中心点坐标(0.5为步距)和对应的步长张量。 sy,sx=torch.meshgrid(sy,sx,indexing='ij')，其中sy代表行，sx代表列，因为计算机中原点是在左上角，往右是x轴，往下是y轴，所以sx是列索引，sy是行索引。而我们常用的索引方式是先说行，再说列，第几行第几列，所以sy,sx这样的索引方式是符合我们的习惯的。但是在数学上，是按照(x,y)的方式进行索引的，就是先第几列，再第几行。所以，meshgrid对每一个网格生成它是第几行的索引坐标集合和它是第几列的坐标集合，所以sy代表每一个网格是第几行，它的形状是[h,w]，第一个维度的元素是[h_i,h_i,...,h_i]，代表这一行w个网络它是第几行，同理，sx的形状是[h,w]，第一个维度的元素是[w_1,w_2,...,w_w]，代表这一行w个网格它是第几列。torch.stack((sx,sy),-1).view(-1,2)是按照(x,y)的索引方式在列的维度[h,w]进行拼接，再调整形状为[锚点的数量,2]，这样就得到所有锚点它的中心点(x,y)坐标索引了。stride_tensor.append(torch.full((h * w, 1), stride, dtype=dtype, device=device))为每一个锚点生成该特征尺度下相同的stride值。 最后输出的锚点和步距张量的形状都是二维的，其中anchor的形状是[总的锚点数量,2(x,y)]，stride的形状是[总的锚点数量，1(stride_value)]。 6. 处理目标标签 现在我们需要对真值标签进行处理。首先将batch_idx(代表目标所在的图片索引批次号)，cls(类别),bbox(xywh)沿着第一维进行拼接，所以我们得到的targets的维度是二维的，形状是[num_target,6]。 然后在process函数中传入targets，batch_size和scale_tensor。如果当前batch中没有目标，则直接返回一个全为0的张量，形状是(batch_size,0,ne-1)，而有目标的话，会先创建一个out张量默认值为0，形状是(batch_size,counts.max(),ne-1)，这里因为用了batch_size维度，所以原来targets中第一维的第一个元素batch_idx，就不用再放入了，这里的counts是来自于targets第0维度的所有目标的所属图片批次的索引，它表示每一个图片批次中有多少个目标存在。比如第一个批次图片索引中，有17个目标，第二个批次图片索引中有2个目标等等。而counts.max()则是找到所有这些图片中，在一张图片上存在的最大目标的数量。并以此为准，创建了counts.max()这个维度，out[j,:n]=targets[matches,1:]就是说将每张图片中存在的目标的除了批次号的元素(包括类别和目标框参数)都进行赋值，所有小于counts.max()的批次号，它只会填充前n个，其它的则默认为0，而遇到counts.max()所在的批次号时，就会全部填满。执行完这个后，就将原来的targets由目标数量的索引形式转为了以批次进行索引的形式。再将out中的目标框由(xywh)的归一化表示转为(xyxy)的表示形式，并乘以scale_tensor还原回最初输入网络的时候图片的大小。 所以process的作用就是将原来的targets的索引方式由目标数量的索引形式转为了以批次进行索引的形式，并将xywh表示转为xyxy表示，然后恢复成原图大小。 gt_labels和gt_bboxes分别来自于处理后的targets的第2维度的第一个元素和其余元素。因为在process中第1个维度采用了counts.max()，所以其它图片中目标数小于最大图片的目标数时，会有默认的0值出现，所以需要mask_gt=gt_bboxes.sum(2,keepdim=True).gt_(0.)来得到一个有效的掩膜，标记有效目标。 7. 解码预测边界框参数 因为预测的边界框参数是偏移量的分布，所以需要对其进行解码，从分布得到真实偏移量。pred_bboxes=self.bbox_decode(anchor_points,pred_distri)，其中anchor_points是所有锚点的中心点坐标形状是[总的锚点数量,2(x,y)]，pred_distri是预测的边界框参数分布形状是[批量大小,总的锚点数量，坐标分布通道数]。YOLOv8默认使用dfl，pred_dist = pred_dist.view(b, a, 4, c // 4).softmax(3).matmul(self.proj.type(pred_dist.dtype))，它将4个方向(ltrb)单独成立了一个维度，然后按照这个维度计算softmax，得到概率分布值，相加为1。然后通过矩阵相乘计算每个离散区间的加权和，权重就是刚才计算得到的概率分布值。在这里,self.proj分成了16个离散区间，其值从[0-15]变化，经过这样的计算,pred_dist的形状是[b,a,4]，就得到了(ltrb)四个方向的偏移量，dist2bbox(pred_dist, anchor_points, xywh=False)的作用就是将(ltrb)四个方向的偏移量转为(xyxy)坐标的形式。 所以，返回的pred_bboxes的形状是从[b,a,reg_max*4]变为[b,a,4]，它的第2维度的每个元素代表一个锚点的(xyxy)坐标表示。 8. 正负样本匹配 和YOLOv3SPP中一样，也要进行正负样本匹配。 ```python @torch.no_grad() def forward(self, pd_scores, pd_bboxes, anc_points, gt_labels, gt_bboxes, mask_gt): """ Compute the task-aligned assignment. Args: pd_scores (torch.Tensor): Predicted classification scores with shape (bs, num_total_anchors, num_classes). pd_bboxes (torch.Tensor): Predicted bounding boxes with shape (bs, num_total_anchors, 4). anc_points (torch.Tensor): Anchor points with shape (num_total_anchors, 2). gt_labels (torch.Tensor): Ground truth labels with shape (bs, n_max_boxes, 1). gt_bboxes (torch.Tensor): Ground truth boxes with shape (bs, n_max_boxes, 4). mask_gt (torch.Tensor): Mask for valid ground truth boxes with shape (bs, n_max_boxes, 1). Returns: target_labels (torch.Tensor): Target labels with shape (bs, num_total_anchors). target_bboxes (torch.Tensor): Target bounding boxes with shape (bs, num_total_anchors, 4). target_scores (torch.Tensor): Target scores with shape (bs, num_total_anchors, num_classes). fg_mask (torch.Tensor): Foreground mask with shape (bs, num_total_anchors). target_gt_idx (torch.Tensor): Target ground truth indices with shape (bs, num_total_anchors). References: https://github.com/Nioolek/PPYOLOE_pytorch/blob/master/ppyoloe/assigner/tal_assigner.py """ self.bs = pd_scores.shape[0] self.n_max_boxes = gt_bboxes.shape[1] device = gt_bboxes.device if self.n_max_boxes == 0: return ( torch.full_like(pd_scores[..., 0], self.num_classes), torch.zeros_like(pd_bboxes), torch.zeros_like(pd_scores), torch.zeros_like(pd_scores[..., 0]), torch.zeros_like(pd_scores[..., 0]), ) try: return self._forward(pd_scores, pd_bboxes, anc_points, gt_labels, gt_bboxes, mask_gt) except torch.cuda.OutOfMemoryError: # Move tensors to CPU, compute, then move back to original device LOGGER.warning("CUDA OutOfMemoryError in TaskAlignedAssigner, using CPU") cpu_tensors = [t.cpu() for t in (pd_scores, pd_bboxes, anc_points, gt_labels, gt_bboxes, mask_gt)] result = self._forward(*cpu_tensors) return tuple(t.to(device) for t in result) ``` 需要注意的是，传入的pd_scores经过了sigmoid()激活，而pd_bboxes和anc_points都还原回了原图大小下尺寸表示，所以涉及到的这些计算都是在原图尺寸下的计算。如果当前所有批次中，没有目标(counts.max()==0)，则直接返回全0的target_labels、target_bboxes、target_scores、fg_mask、target_gt_idx。否则，调用_forward()方法进行正负样本匹配。 ```python def _forward(self, pd_scores, pd_bboxes, anc_points, gt_labels, gt_bboxes, mask_gt): """ Compute the task-aligned assignment. Args: pd_scores (torch.Tensor): Predicted classification scores with shape (bs, num_total_anchors, num_classes). pd_bboxes (torch.Tensor): Predicted bounding boxes with shape (bs, num_total_anchors, 4). anc_points (torch.Tensor): Anchor points with shape (num_total_anchors, 2). gt_labels (torch.Tensor): Ground truth labels with shape (bs, n_max_boxes, 1). gt_bboxes (torch.Tensor): Ground truth boxes with shape (bs, n_max_boxes, 4). mask_gt (torch.Tensor): Mask for valid ground truth boxes with shape (bs, n_max_boxes, 1). Returns: target_labels (torch.Tensor): Target labels with shape (bs, num_total_anchors). target_bboxes (torch.Tensor): Target bounding boxes with shape (bs, num_total_anchors, 4). target_scores (torch.Tensor): Target scores with shape (bs, num_total_anchors, num_classes). fg_mask (torch.Tensor): Foreground mask with shape (bs, num_total_anchors). target_gt_idx (torch.Tensor): Target ground truth indices with shape (bs, num_total_anchors). """ mask_pos, align_metric, overlaps = self.get_pos_mask( pd_scores, pd_bboxes, gt_labels, gt_bboxes, anc_points, mask_gt ) target_gt_idx, fg_mask, mask_pos = self.select_highest_overlaps(mask_pos, overlaps, self.n_max_boxes) # Assigned target target_labels, target_bboxes, target_scores = self.get_targets(gt_labels, gt_bboxes, target_gt_idx, fg_mask) # Normalize align_metric *= mask_pos pos_align_metrics = align_metric.amax(dim=-1, keepdim=True) # b, max_num_obj pos_overlaps = (overlaps * mask_pos).amax(dim=-1, keepdim=True) # b, max_num_obj norm_align_metric = (align_metric * pos_overlaps / (pos_align_metrics + self.eps)).amax(-2).unsqueeze(-1) target_scores = target_scores * norm_align_metric return target_labels, target_bboxes, target_scores, fg_mask.bool(), target_gt_idx ``` self.get_pos_mask的作用是为每一个真值box获取正样本掩膜 1. 调用select_candidates_in_gts方法，选取落在在真值框内的正锚点中心点mask_in_gts。 1. xy_centers[None]的作用是给锚点中心增加维度，shape从[n_anchors,2]变为了[1,n_anchors,2]，再通过广播与lt([bs*counts.max(),1,2])匹配，得到[bs*counts.max(),n_anchors,2]其含义是每个真实框与每个锚点的中心x1差、中心y1差。 2. rb-xy_centers[None]同理，得到每个真实框与每个锚点的中心x2差、中心y2差。 3. torch.cat(...,dim=2)，拼接后坐标维度变成4，shape[bs*n_boxes,n_anchors,4]这4个值分别是x_center-x1,y_center-y1,x2-x_center,y2-y_center。 4. 最后view恢复批次和真实框的维度，最终bbox_deltas的shape为[bs,n_boxes,n_anchors,4]。 5. return bbox_deltas.amin(3).gt_(eps)。测分数 1. 最终判断锚点是否在真实框内，amin(3)表示对第3维度(4个偏移差)取最小值，每个批次、真实框、锚点对应一个值，代表4个偏移差中最小的那个。 2. gt_(eps)，判断这个最小值是否大小eps极小值，避免数值误差 3. 最终返回shape为[bs,n_boxes(counts.max()),n_anchors]的布尔张量，True表示该锚点的4个偏移差均为正，即锚点中心在真实框的内部：x1&lt;x_center&lt;x2且y1&lt;y_center&lt;y2，False表示锚点在真实框外部或者边界处。 6. 总结下来，该方法就是判断锚点是否是在真实框的内部，而里面的n_boxes其实是之前构建真实框时用的counts.max()，所有xyxy为默认0值的当然最终其布尔值为False。 2. 调用get_box_metrics方法，基于预测和真值框，获取align_metric,overlaps。 1. mask_gt是有效的gt框掩膜，shape为[bs,n_max_boxes,n_anchors]，True表示该真实框-anchor对是有效的，False表示无效的gt-anchor对。 2. overlaps是用来记录gt与anchor的重叠度(IOU)，shape为[bs,n_max_boxes,n_anchors]，记录真实框-anchor对的IOU值。 3. bbox_scores是用来记录每个真实框-anchor对的类别预测分数，shape为[bs,n_max_boxes,n_anchors]，记录每个真实框-anchor对的类别预测分数。 4. 构建类别索引，提取对应的预测分数 1. ind=torch.zeros([2, self.bs, self.n_max_boxes], dtype=torch.long) 2. ind[0]=torch.arange(self.bs)[:,None].expand(-1,self.n_max_boxes),shape是[bs,n_max_boxes]，其值是[[0,...,0],...,[bs-1,...,bs-1]]这样的。 3. ind[1]=gt_labels.long().squeeze(-1)，shape是[bs,n_max_boxes] 4. bbox_scores[mask_gt]=pd_scores[ind[0], :, ind[1]][mask_gt]，在有效的gt-anchor对处，赋值为pd_scores中对应的值，shape是[bs,n_max_boxes,n_anchors] 5. 计算GT和anchor的IOU，pd_boxes = pd_bboxes.unsqueeze(1).expand(-1, self.n_max_boxes, -1, -1)[mask_gt]，shape是[N,4]，表示N个有效gt-anchor对的预测框坐标。 1. gt_boxes = gt_bboxes.unsqueeze(2).expand(-1, -1, na, -1)[mask_gt]，shape是[N,4]，表示N个有效gt-anchor对的真实框坐标 2. overlaps[mask_gt]=self.iou_calculation(gt_boxes, pd_boxes)，计算这些有效gt-anchor对的CIOU值 6. align_metric = bbox_scores.pow(self.alpha) * overlaps.pow(self.beta)，计算对齐度量，将类别预测分数和IOU分别加权后相乘，综合反映“预测框与GT的匹配程度”，即分数越高、IOU越大，匹配度越高。 7. 输出align_metric(用于正负样本分配或损失加权)和overlaps(原始的IOU值)，它们的shape都是[bs,n_max_boxes,n_anchors] 8. 总结下，该方法的核心目的就是为每个真是目标框找到最匹配的预测框，给后续的正负样本划分或损失计算提供依据。 3. mask_topk = self.select_topk_candidates(align_metric, topk_mask=mask_gt.expand(-1, -1, self.topk).bool())，获取前topk个度量metric的掩膜 4. mask_pos = mask_topk * mask_in_gts * mask_gt，获取最终的正样本掩膜 5. 返回mask_pos, align_metric, overlaps，它们的形状都是[bs,n_max_boxes,n_anchors] self.select_highest_overlaps的作用是当多个真实框对应同一个锚框时，根据overlaps选择IOU最大的那个真实框，返回其索引、前景掩膜、正样本掩膜。 self.get_targets的作用是根据选择的真实框索引，获取对应的目标标签、目标框、目标分数。 最终，输出target_bboxes(正样本的目标框),target_scores(正样本=类别概率，负样本=0),fg_mask(正样本掩码，True=正样本，False=负样本) 多标签分类与多分类 1 2 3 4 5 6 7 8 9 我一直没有搞明白BCELoss和CrossEntropyLoss它们有什么区别。现在终于动了动我的小脑筋，克服了拖延症，细细分析下它们的区别，以及延申到YOLO目标检测中使用多标签分类，这也是把它们放在这里叙述的原因。 首先我们看BCELoss。顾名思义，就是二值交叉熵损失，它用于二分类任务，计算预测值与真实标签之间的差异。它的公式如下： L(y, \hat{y}) = - \left[ y \cdot \log(\hat{y}) + (1 - y) \cdot \log(1 - \hat{y}) \right] y的取值范围是{0,1}，当y取1时，公式简化为`L = -\log(\hat{y})`，同样地当y取0时，公式简化为`L = -\log(1 - \hat{y})`。 对于N个样本，我们取平均值 `L = -\frac{1}{N} \sum_{i=1}^{N} \left[ y_i \cdot \log(\hat{y}_i) + (1 - y_i) \cdot \log(1 - \hat{y}_i) \right]` 并且，实际使用时常常讲sigmoid函数和BCE损失计算进行合并，BCEWithLogitsLoss，以求数值稳定。 此外，我们还可以使用带权重的变体形式，以平衡二分类下正负样本的类别不平衡问题。 `L = -\frac{1}{N} \sum_{i=1}^{N} \left[ \omega \cdot y_i \cdot \log(\hat{y}_i) + (1 - \omega) \cdot (1 - y_i) \cdot \log(1 - \hat{y}_i) \right]` 接下来我们看多类别交叉熵损失。它是针对多分类任务（输出≥3 类，且样本 “互斥唯一”，即一个样本只能属于一类）设计的损失函数，本质是基于多项式分布（单个样本有 K 种输出可能，且概率和为 1）的交叉熵计算。它的激活函数是softmax，将模型输出的原始分数转换为概率分布，并且概率之和为1。 当只有一个样本时，它的形式是 `L(y, \hat{y}) = - \sum_{k=1}^{K} y_k \cdot \log(\hat{y}_k)` 而因为采用独特编码，所以当有多个样本时，公式可以简化成 `L = -\frac{1}{N} \sum_{i=1}^{N} \log(\hat{y}_{i, y_i})` 不要因为有多类别，就混淆了多标签分类(multi label classification)和多分类(multi-class classification)。多标签分类是指每个样本可以属于多个类别，而多分类是指每个样本只能属于一个类别。例如，一个图片可以同时包含猫、狗和鸟，这是一个多标签分类问题。而一个图片只能是猫、狗或鸟中的一种，这是一个多分类问题。在YOLO目标检测中计算目标的类别损失时，其实就是一个多标签分类问题，使用的是BCELoss，对每一个类别单独计算它的二值交叉熵损失。对于80类的检测任务，YOLO不会让模型输出“80个概率和为1”的结果，而是让模型输出80个独立的[0,1]值——每个值表示“该框属于这个类别的概率”，彼此独立、互不影响（比如一个框可以同时输出“人0.95”、“车：0.02”、“猫：0.01”，无需求和为1）它的计算逻辑是对每一个类别单独计算BCE损失，再求平均。 YOLO中类别损失计算时怎么为不同类别设置不同的权重值 1 在YOLO中，类别损失计算时，为了平衡不同类别之间的差异，通常会为每个类别设置不同的权重值。这是因为不同类别的样本数量可能不均衡，某些类别可能出现的频率更高，而某些类别可能出现的频率更低。为了使模型更关注出现频率较低的类别，我们可以为这些类别设置较高的权重值。在二分类问题中，常常出现正负样本比例不均衡的问题，这时就需要为正负样本设置不同的权重值来调整它们的损失值。在YOLO中的类别损失计算是一个多标签分类问题，它们是单独计算每一个类别的二值交叉熵损失的，根据在nn.BCEWithLogits中的posi_weight这个参数，可以为每一个类别设置不同的权重值（针对该类别下的二值交叉熵损失计算的正样本赋予该权重）。比如数据集中有三个类别，人、猫和狗，发现人的样本数量特别少，无论是人相对于猫和狗，又或者相对于在人的二分类下的正负样本，人的样本数量都很少，就可以设置pos_weight=[3,1,1]，为人的正样本在计算BCE时设置3倍权重值。以平衡类别不平衡的问题。YOLO11 在多标签场景下，输出层会为每个类别生成独立的logits，给每个类别配置对应的pos_weight实现逐类损失加权。pos_weight的计算依据是该类别自身的正负样本比例（人类正样本数/人类负样本数），而非跨类别对比（人类样本数/猫样本数）。整个任务是多标签（同时识别人和猫狗），但每个类内部是独立二分类（人：有 / 无；猫：有 / 无），pos_weight 只针对 “单类内部” 的正负失衡，而非整个任务的正负失衡。]]></summary></entry><entry><title type="html">测试</title><link href="https://ke-albert.github.io/test/2025/09/20/test/" rel="alternate" type="text/html" title="测试" /><published>2025-09-20T20:20:00+08:00</published><updated>2025-09-20T20:20:00+08:00</updated><id>https://ke-albert.github.io/test/2025/09/20/test</id><content type="html" xml:base="https://ke-albert.github.io/test/2025/09/20/test/"><![CDATA[<p>这是一个测试文章</p>]]></content><author><name></name></author><category term="test" /><summary type="html"><![CDATA[这是一个测试文章]]></summary></entry><entry><title type="html">Welcome to Jekyll!</title><link href="https://ke-albert.github.io/jekyll/update/2025/09/20/welcome-to-jekyll/" rel="alternate" type="text/html" title="Welcome to Jekyll!" /><published>2025-09-20T15:43:31+08:00</published><updated>2025-09-20T15:43:31+08:00</updated><id>https://ke-albert.github.io/jekyll/update/2025/09/20/welcome-to-jekyll</id><content type="html" xml:base="https://ke-albert.github.io/jekyll/update/2025/09/20/welcome-to-jekyll/"><![CDATA[<p>You’ll find this post in your <code class="language-plaintext highlighter-rouge">_posts</code> directory. Go ahead and edit it and re-build the site to see your changes. You can rebuild the site in many different ways, but the most common way is to run <code class="language-plaintext highlighter-rouge">jekyll serve</code>, which launches a web server and auto-regenerates your site when a file is updated.</p>

<p>Jekyll requires blog post files to be named according to the following format:</p>

<p><code class="language-plaintext highlighter-rouge">YEAR-MONTH-DAY-title.MARKUP</code></p>

<p>Where <code class="language-plaintext highlighter-rouge">YEAR</code> is a four-digit number, <code class="language-plaintext highlighter-rouge">MONTH</code> and <code class="language-plaintext highlighter-rouge">DAY</code> are both two-digit numbers, and <code class="language-plaintext highlighter-rouge">MARKUP</code> is the file extension representing the format used in the file. After that, include the necessary front matter. Take a look at the source for this post to get an idea about how it works.</p>

<p>Jekyll also offers powerful support for code snippets:</p>

<figure class="highlight"><pre><code class="language-ruby" data-lang="ruby"><span class="k">def</span> <span class="nf">print_hi</span><span class="p">(</span><span class="nb">name</span><span class="p">)</span>
  <span class="nb">puts</span> <span class="s2">"Hi, </span><span class="si">#{</span><span class="nb">name</span><span class="si">}</span><span class="s2">"</span>
<span class="k">end</span>
<span class="n">print_hi</span><span class="p">(</span><span class="s1">'Tom'</span><span class="p">)</span>
<span class="c1">#=&gt; prints 'Hi, Tom' to STDOUT.</span></code></pre></figure>

<p>Check out the <a href="https://jekyllrb.com/docs/home">Jekyll docs</a> for more info on how to get the most out of Jekyll. File all bugs/feature requests at <a href="https://github.com/jekyll/jekyll">Jekyll’s GitHub repo</a>. If you have questions, you can ask them on <a href="https://talk.jekyllrb.com/">Jekyll Talk</a>.</p>]]></content><author><name></name></author><category term="jekyll" /><category term="update" /><summary type="html"><![CDATA[You’ll find this post in your _posts directory. Go ahead and edit it and re-build the site to see your changes. You can rebuild the site in many different ways, but the most common way is to run jekyll serve, which launches a web server and auto-regenerates your site when a file is updated. Jekyll requires blog post files to be named according to the following format: YEAR-MONTH-DAY-title.MARKUP Where YEAR is a four-digit number, MONTH and DAY are both two-digit numbers, and MARKUP is the file extension representing the format used in the file. After that, include the necessary front matter. Take a look at the source for this post to get an idea about how it works. Jekyll also offers powerful support for code snippets: def print_hi(name) puts "Hi, #{name}" end print_hi('Tom') #=&gt; prints 'Hi, Tom' to STDOUT. Check out the Jekyll docs for more info on how to get the most out of Jekyll. File all bugs/feature requests at Jekyll’s GitHub repo. If you have questions, you can ask them on Jekyll Talk.]]></summary></entry></feed>