<?xml version="1.0" encoding="utf-8" standalone="yes"?><rss version="2.0" xmlns:atom="http://www.w3.org/2005/Atom" xmlns:content="http://purl.org/rss/1.0/modules/content/"><channel><title>持续更新 on Ganko Space</title><link>https://ganko.asia/tags/%E6%8C%81%E7%BB%AD%E6%9B%B4%E6%96%B0/</link><description>Recent content in 持续更新 on Ganko Space</description><generator>Hugo</generator><language>zh-cn</language><lastBuildDate>Mon, 24 Nov 2025 00:00:00 +0000</lastBuildDate><atom:link href="https://ganko.asia/tags/%E6%8C%81%E7%BB%AD%E6%9B%B4%E6%96%B0/index.xml" rel="self" type="application/rss+xml"/><item><title>NanoGPT</title><link>https://ganko.asia/posts/nanogpt/</link><pubDate>Mon, 24 Nov 2025 00:00:00 +0000</pubDate><guid>https://ganko.asia/posts/nanogpt/</guid><description>&lt;blockquote>
&lt;p>所有之前:这篇是关于NanoGPT项目的核心代码的总结和理解，我会持续把主流技术加入到这个项目中，并对核心代码进行解释。完整项目代码需要查询&lt;a href="https://github.com/JialeChe/NanoGPT.git">NanoGPT&lt;/a>.&lt;/p>&lt;/blockquote>
&lt;h2 id="模型架构">模型架构&lt;/h2>
&lt;h3 id="base">Base&lt;/h3>
&lt;p>最基础的模型架构是一个Decoder-only的Transformer模型，代码实现参考&lt;a href="https://ganko.asia/posts/codebymyself/">CodeByMyself&lt;/a>,在CodeByMyself的基础上，使用了更现代化的Rope位置编码和SwiGLU激活函数， 并且加入了Moe用于和FFN对比，可在配置文件中选择是否使用Moe。部分代码如下:&lt;/p>
&lt;div class="highlight">&lt;pre tabindex="0" style="color:#f8f8f2;background-color:#272822;-moz-tab-size:4;-o-tab-size:4;tab-size:4;">&lt;code class="language-python" data-lang="python">&lt;span style="display:flex;">&lt;span>&lt;span style="color:#75715e"># swiglu&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>&lt;span style="color:#66d9ef">class&lt;/span> &lt;span style="color:#a6e22e">ffn_swiglu&lt;/span>(nn&lt;span style="color:#f92672">.&lt;/span>Module):
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#66d9ef">def&lt;/span> __init__(self, d_model, d_hidden, dropout&lt;span style="color:#f92672">=&lt;/span>&lt;span style="color:#ae81ff">0.1&lt;/span>):
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> super()&lt;span style="color:#f92672">.&lt;/span>__init__()
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#75715e">#第一个 Linear 输出 2 * d_hidden&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> self&lt;span style="color:#f92672">.&lt;/span>w_gate_up &lt;span style="color:#f92672">=&lt;/span> nn&lt;span style="color:#f92672">.&lt;/span>Linear(d_model, &lt;span style="color:#ae81ff">2&lt;/span> &lt;span style="color:#f92672">*&lt;/span> d_hidden) &lt;span style="color:#75715e">#同时生成 gate 和 up&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> self&lt;span style="color:#f92672">.&lt;/span>w_down &lt;span style="color:#f92672">=&lt;/span> nn&lt;span style="color:#f92672">.&lt;/span>Linear(d_hidden, d_model)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> self&lt;span style="color:#f92672">.&lt;/span>dropout &lt;span style="color:#f92672">=&lt;/span> nn&lt;span style="color:#f92672">.&lt;/span>Dropout(dropout)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#66d9ef">def&lt;/span> &lt;span style="color:#a6e22e">forward&lt;/span>(self, x):
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#75715e"># x: [B, L, d_model]&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> gate_up &lt;span style="color:#f92672">=&lt;/span> self&lt;span style="color:#f92672">.&lt;/span>w_gate_up(x) &lt;span style="color:#75715e"># [B, L, 2 * d_hidden]&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> gate, up &lt;span style="color:#f92672">=&lt;/span> gate_up&lt;span style="color:#f92672">.&lt;/span>chunk(&lt;span style="color:#ae81ff">2&lt;/span>, dim&lt;span style="color:#f92672">=-&lt;/span>&lt;span style="color:#ae81ff">1&lt;/span>) &lt;span style="color:#75715e"># each: [B, L, d_hidden]&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> swiglu_out &lt;span style="color:#f92672">=&lt;/span> F&lt;span style="color:#f92672">.&lt;/span>silu(gate) &lt;span style="color:#f92672">*&lt;/span> up &lt;span style="color:#75715e"># [B, L, d_hidden]&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> out &lt;span style="color:#f92672">=&lt;/span> self&lt;span style="color:#f92672">.&lt;/span>w_down(self&lt;span style="color:#f92672">.&lt;/span>dropout(swiglu_out)) &lt;span style="color:#75715e"># [B, L, d_model]&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#66d9ef">return&lt;/span> out
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>&lt;span style="color:#75715e"># moe&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>&lt;span style="color:#66d9ef">class&lt;/span> &lt;span style="color:#a6e22e">moe&lt;/span>(nn&lt;span style="color:#f92672">.&lt;/span>Module):
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#66d9ef">def&lt;/span> __init__(self,n_expert,d_model,top_k&lt;span style="color:#f92672">=&lt;/span>&lt;span style="color:#ae81ff">2&lt;/span>,dropout&lt;span style="color:#f92672">=&lt;/span>&lt;span style="color:#ae81ff">0.1&lt;/span>):
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> super()&lt;span style="color:#f92672">.&lt;/span>__init__()
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> self&lt;span style="color:#f92672">.&lt;/span>d_model &lt;span style="color:#f92672">=&lt;/span> d_model
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> self&lt;span style="color:#f92672">.&lt;/span>n_expert &lt;span style="color:#f92672">=&lt;/span> n_expert
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> self&lt;span style="color:#f92672">.&lt;/span>top_k &lt;span style="color:#f92672">=&lt;/span> top_k
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> self&lt;span style="color:#f92672">.&lt;/span>dropout &lt;span style="color:#f92672">=&lt;/span> dropout
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> self&lt;span style="color:#f92672">.&lt;/span>gate &lt;span style="color:#f92672">=&lt;/span> nn&lt;span style="color:#f92672">.&lt;/span>Linear(d_model,n_expert,bias&lt;span style="color:#f92672">=&lt;/span>&lt;span style="color:#66d9ef">False&lt;/span>)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> self&lt;span style="color:#f92672">.&lt;/span>softmax &lt;span style="color:#f92672">=&lt;/span> nn&lt;span style="color:#f92672">.&lt;/span>Softmax(dim&lt;span style="color:#f92672">=-&lt;/span>&lt;span style="color:#ae81ff">1&lt;/span>)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#75715e"># self.experts = nn.ModuleList([ffn(self.d_model,self.d_model*4,dropout) for _ in range(n_expert)]) relu激活的expert&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> self&lt;span style="color:#f92672">.&lt;/span>experts &lt;span style="color:#f92672">=&lt;/span> nn&lt;span style="color:#f92672">.&lt;/span>ModuleList([ffn_swiglu(self&lt;span style="color:#f92672">.&lt;/span>d_model,self&lt;span style="color:#f92672">.&lt;/span>d_model&lt;span style="color:#f92672">*&lt;/span>&lt;span style="color:#ae81ff">2&lt;/span>,dropout) &lt;span style="color:#66d9ef">for&lt;/span> _ &lt;span style="color:#f92672">in&lt;/span> range(n_expert)])
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#66d9ef">def&lt;/span> &lt;span style="color:#a6e22e">forward&lt;/span>(self,x):
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> b,t,d &lt;span style="color:#f92672">=&lt;/span> x&lt;span style="color:#f92672">.&lt;/span>shape
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#66d9ef">assert&lt;/span> d &lt;span style="color:#f92672">==&lt;/span> self&lt;span style="color:#f92672">.&lt;/span>d_model,&lt;span style="color:#e6db74">f&lt;/span>&lt;span style="color:#e6db74">&amp;#34;输入维度和moe设置维度不匹配&amp;#34;&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> x_flat &lt;span style="color:#f92672">=&lt;/span> x&lt;span style="color:#f92672">.&lt;/span>view(&lt;span style="color:#f92672">-&lt;/span>&lt;span style="color:#ae81ff">1&lt;/span>,d)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> N &lt;span style="color:#f92672">=&lt;/span> x_flat&lt;span style="color:#f92672">.&lt;/span>shape[&lt;span style="color:#ae81ff">0&lt;/span>]
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> gate_logits &lt;span style="color:#f92672">=&lt;/span> self&lt;span style="color:#f92672">.&lt;/span>gate(x_flat)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> topk_weights,topk_indices &lt;span style="color:#f92672">=&lt;/span> torch&lt;span style="color:#f92672">.&lt;/span>topk(gate_logits,self&lt;span style="color:#f92672">.&lt;/span>top_k,dim&lt;span style="color:#f92672">=-&lt;/span>&lt;span style="color:#ae81ff">1&lt;/span>)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> topk_weights&lt;span style="color:#f92672">=&lt;/span> self&lt;span style="color:#f92672">.&lt;/span>softmax(topk_weights)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> out &lt;span style="color:#f92672">=&lt;/span> torch&lt;span style="color:#f92672">.&lt;/span>zeros_like(x_flat)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#66d9ef">for&lt;/span> i,expert &lt;span style="color:#f92672">in&lt;/span> enumerate(self&lt;span style="color:#f92672">.&lt;/span>experts):
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> mask &lt;span style="color:#f92672">=&lt;/span> (topk_indices &lt;span style="color:#f92672">==&lt;/span> i)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#66d9ef">if&lt;/span> &lt;span style="color:#f92672">not&lt;/span> mask&lt;span style="color:#f92672">.&lt;/span>any():
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#66d9ef">continue&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> token_indices,expert_pos &lt;span style="color:#f92672">=&lt;/span> torch&lt;span style="color:#f92672">.&lt;/span>where(mask)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> select_x &lt;span style="color:#f92672">=&lt;/span> x_flat[token_indices]
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> expert_out &lt;span style="color:#f92672">=&lt;/span> expert(select_x)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> weights &lt;span style="color:#f92672">=&lt;/span> topk_weights[token_indices, expert_pos]
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> out&lt;span style="color:#f92672">.&lt;/span>index_add_(&lt;span style="color:#ae81ff">0&lt;/span>, token_indices, expert_out &lt;span style="color:#f92672">*&lt;/span> weights&lt;span style="color:#f92672">.&lt;/span>unsqueeze(&lt;span style="color:#ae81ff">1&lt;/span>))
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#66d9ef">return&lt;/span> out&lt;span style="color:#f92672">.&lt;/span>view(b,t,d)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>&lt;span style="color:#66d9ef">def&lt;/span> &lt;span style="color:#a6e22e">precompute_freqs_cis&lt;/span>(dim: int, end: int, theta: float &lt;span style="color:#f92672">=&lt;/span> &lt;span style="color:#ae81ff">10000.0&lt;/span>):
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#e6db74">&amp;#34;&amp;#34;&amp;#34;
&lt;/span>&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>&lt;span style="color:#e6db74"> 预计算旋转角度的复数表示（cos + i*sin）
&lt;/span>&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>&lt;span style="color:#e6db74"> &amp;#34;&amp;#34;&amp;#34;&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> freqs &lt;span style="color:#f92672">=&lt;/span> &lt;span style="color:#ae81ff">1.0&lt;/span> &lt;span style="color:#f92672">/&lt;/span> (theta &lt;span style="color:#f92672">**&lt;/span> (torch&lt;span style="color:#f92672">.&lt;/span>arange(&lt;span style="color:#ae81ff">0&lt;/span>, dim, &lt;span style="color:#ae81ff">2&lt;/span>)[: (dim &lt;span style="color:#f92672">//&lt;/span> &lt;span style="color:#ae81ff">2&lt;/span>)]&lt;span style="color:#f92672">.&lt;/span>float() &lt;span style="color:#f92672">/&lt;/span> dim))
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> t &lt;span style="color:#f92672">=&lt;/span> torch&lt;span style="color:#f92672">.&lt;/span>arange(end, device&lt;span style="color:#f92672">=&lt;/span>freqs&lt;span style="color:#f92672">.&lt;/span>device)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> freqs &lt;span style="color:#f92672">=&lt;/span> torch&lt;span style="color:#f92672">.&lt;/span>outer(t, freqs)&lt;span style="color:#f92672">.&lt;/span>float() &lt;span style="color:#75715e"># [end, dim//2]&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> freqs_cos &lt;span style="color:#f92672">=&lt;/span> freqs&lt;span style="color:#f92672">.&lt;/span>cos()
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> freqs_sin &lt;span style="color:#f92672">=&lt;/span> freqs&lt;span style="color:#f92672">.&lt;/span>sin()
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#66d9ef">return&lt;/span> freqs_cos, freqs_sin &lt;span style="color:#75715e"># 分开存储更便于后续操作，不使用torch.complex类型，因为部分系统不支持&lt;/span>
&lt;/span>&lt;/span>&lt;/code>&lt;/pre>&lt;/div>&lt;p>SwiGLU激活函数的实现是通过将线性层的输出分成两部分，一部分作为gate，另一部分作为up，然后使用SILU函数对gate进行激活，并与up相乘得到最终的输出。&lt;/p></description></item></channel></rss>