<?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>Transformer on Ganko Space</title><link>https://ganko.asia/tags/transformer/</link><description>Recent content in Transformer on Ganko Space</description><generator>Hugo</generator><language>zh-cn</language><lastBuildDate>Fri, 03 Apr 2026 00:00:00 +0000</lastBuildDate><atom:link href="https://ganko.asia/tags/transformer/index.xml" rel="self" type="application/rss+xml"/><item><title>CS336 Lecture 1: Overview and Tokenization</title><link>https://ganko.asia/posts/cs336-lecture1/</link><pubDate>Fri, 03 Apr 2026 00:00:00 +0000</pubDate><guid>https://ganko.asia/posts/cs336-lecture1/</guid><description>&lt;h2 id="课程导读">课程导读&lt;/h2>
&lt;p>这门课的讲义基于 notebook 组织，不同章节通过函数模块化展开。&lt;/p>
&lt;p>开场老师提出了一个很尖锐的问题：&lt;strong>研究者正在与底层技术逐步脱节&lt;/strong>。他给了一个时间线：&lt;/p>
&lt;ul>
&lt;li>八年前，研究者会自己补充数据并训练模型；&lt;/li>
&lt;li>六年前，研究者还会下载模型后进行 fine-tune；&lt;/li>
&lt;li>现在，很多人直接向闭源模型（GPT-4/Claude/Gemini）提问。&lt;/li>
&lt;/ul>
&lt;p>这虽然有些夸张，但确实反映了趋势：模型能力提升后，很多基础环节被“封装”了。老师强调：&lt;/p>
&lt;blockquote>
&lt;p>&lt;strong>&amp;ldquo;Full understanding of this technology is necessary for fundamental research&amp;rdquo;&lt;/strong>&lt;/p>&lt;/blockquote>
&lt;p>随着模型规模持续增大，训练中的计算压力正在从 Attention 逐步转向 FFN。&lt;/p>
&lt;p align="center">
&lt;img src="https://ganko.asia/images/计算压力转换.png" alt="计算压力转换" />
&lt;/p>
&lt;h3 id="课程收获">课程收获&lt;/h3>
&lt;ul>
&lt;li>&lt;strong>Mechanics&lt;/strong>：系统如何工作（如 Transformer 结构、GPU 并行方式）&lt;/li>
&lt;li>&lt;strong>Mindset&lt;/strong>：如何最大化利用硬件与内存，如何看待 scaling laws&lt;/li>
&lt;li>&lt;strong>Intuition&lt;/strong>：哪些数据和建模决策更可能带来更好结果&lt;/li>
&lt;/ul>
&lt;p>老师也指出一个常见误解：只要堆算力，模型就会自动变好。他给出的观点是：&lt;strong>准确率 = 效率 × 资源&lt;/strong>。&lt;/p>
&lt;p>在数据和资源固定时，效率决定了上限。&lt;/p>
&lt;p>老师将近年的 LLM 工作大致分成两类：&lt;/p>
&lt;ol>
&lt;li>&lt;strong>闭源路线&lt;/strong>：以 GPT 系列为代表，最早系统性拥抱 Scale，但细节封闭。&lt;/li>
&lt;li>&lt;strong>开放权重路线&lt;/strong>：以 Qwen 等 open-weight 模型为代表，拥抱 Scale 的同时提升开放度，但常见情况是只公开部分训练细节，数据与失败案例仍不完整。&lt;/li>
&lt;/ol>
&lt;hr>
&lt;h2 id="tokenization">Tokenization&lt;/h2>
&lt;p>为什么需要 tokenization？因为 LLM 在 token 序列上建模概率分布，我们需要把原始字符串编码为 token，并保证可逆解码。&lt;/p>
&lt;p align="center">
&lt;img src="https://ganko.asia/images/tokenization.png" alt="tokenization" />
&lt;/p>
&lt;p>🔗 &lt;strong>Tokenization 可视化工具&lt;/strong>: &lt;a href="https://tiktokenizer.vercel.app/?encoder=gpt2">tiktokenizer.vercel.app&lt;/a>&lt;/p>
&lt;p>常见 tokenization 方法：&lt;/p></description></item><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><item><title>CodeByMySelf-Transformer</title><link>https://ganko.asia/posts/codebymyself-transformer/</link><pubDate>Wed, 24 Sep 2025 00:00:00 +0000</pubDate><guid>https://ganko.asia/posts/codebymyself-transformer/</guid><description>&lt;blockquote>
&lt;p>所有之前：CodeByMySelf 系列之 Transformer 的实现记录，只有标题没有内容的是TODO，针对经典算法的变体和改进算法在其他文章。&lt;/p>&lt;/blockquote>
&lt;h2 id="transformer">Transformer&lt;/h2>
&lt;p>下面是一个最基本的Transformer实现，主要包含以下几个部分：&lt;/p>
&lt;ul>
&lt;li>基本模块
&lt;ul>
&lt;li>多头注意力模块&lt;/li>
&lt;li>Token和位置嵌入模块&lt;/li>
&lt;li>LayerNorm&lt;/li>
&lt;li>前馈神经网络FFN&lt;/li>
&lt;/ul>
&lt;/li>
&lt;li>$EncoderLayer -> Encoder$&lt;/li>
&lt;li>$DecoderLayer -> Decoder$&lt;/li>
&lt;li>$Encoder+Decoder -> Transformer$&lt;/li>
&lt;/ul>
&lt;h3 id="多头注意力mha">多头注意力MHA&lt;/h3>
&lt;p>多头注意力机制的公式如下：
&lt;/p>
$$ \text{MultiHead}(Q, K, V) = \text{Concat}(\text{head}_1, \ldots, \text{head}_h)W^O $$&lt;p>其中，每一个头的计算分别为：
&lt;/p>
$$ \text{Attention}(Q, K, V) = \text{softmax}\left(\frac{QK^\top}{\sqrt{d_k}}\right)V $$&lt;p>对于图像数据和序列数据，多头注意力机制的实现是不同，但是核心原理是相同的。对于序列数据，每个位置对应一个token；对于图像数据，每个位置对应一个patch块或者一个像素块，如果是$patch$块，需要展开到一维向量。这里实现的为序列数据的多头注意力机制,初始$shape$为$(batch,time,d_model)$&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>&lt;/span>&lt;span style="display:flex;">&lt;span>&lt;span style="color:#f92672">import&lt;/span> torch
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>&lt;span style="color:#f92672">from&lt;/span> torch &lt;span style="color:#f92672">import&lt;/span> nn
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>&lt;span style="color:#f92672">import&lt;/span> torch.functional &lt;span style="color:#66d9ef">as&lt;/span> F
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>&lt;span style="color:#f92672">import&lt;/span> math
&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">MultiHeadAttention&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,d_model,n_head):
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> super(MultiHeadAttention, self)&lt;span style="color:#f92672">.&lt;/span>__init__()
&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">assert&lt;/span> d_model &lt;span style="color:#f92672">%&lt;/span> n_head &lt;span style="color:#f92672">==&lt;/span> &lt;span style="color:#ae81ff">0&lt;/span>
&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_head &lt;span style="color:#f92672">=&lt;/span> n_head
&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>w_k &lt;span style="color:#f92672">=&lt;/span> nn&lt;span style="color:#f92672">.&lt;/span>Linear(d_model,d_model)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> self&lt;span style="color:#f92672">.&lt;/span>w_q &lt;span style="color:#f92672">=&lt;/span> nn&lt;span style="color:#f92672">.&lt;/span>Linear(d_model,d_model)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> self&lt;span style="color:#f92672">.&lt;/span>w_v &lt;span style="color:#f92672">=&lt;/span> nn&lt;span style="color:#f92672">.&lt;/span>Linear(d_model,d_model)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> self&lt;span style="color:#f92672">.&lt;/span>w_conbine &lt;span style="color:#f92672">=&lt;/span> nn&lt;span style="color:#f92672">.&lt;/span>Linear(d_model,d_model,bias&lt;span style="color:#f92672">=&lt;/span>&lt;span style="color:#66d9ef">False&lt;/span>) &lt;span style="color:#75715e"># 对应W^O&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>&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,q,k,v,mask&lt;span style="color:#f92672">=&lt;/span>&lt;span style="color:#66d9ef">None&lt;/span>):
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> batch,time,dimension &lt;span style="color:#f92672">=&lt;/span> q&lt;span style="color:#f92672">.&lt;/span>shape() &lt;span style="color:#75715e"># q,k,v shape: (batch,time,d_model)&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> n_d &lt;span style="color:#f92672">=&lt;/span> self&lt;span style="color:#f92672">.&lt;/span>d_model &lt;span style="color:#f92672">//&lt;/span> self&lt;span style="color:#f92672">.&lt;/span>n_head
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> q,k,v &lt;span style="color:#f92672">=&lt;/span> self&lt;span style="color:#f92672">.&lt;/span>w_q(q),self&lt;span style="color:#f92672">.&lt;/span>w_k(k),self&lt;span style="color:#f92672">.&lt;/span>w_v(v)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> print(q&lt;span style="color:#f92672">.&lt;/span>shape, batch, time, self&lt;span style="color:#f92672">.&lt;/span>n_head, n_d)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> q &lt;span style="color:#f92672">=&lt;/span> q&lt;span style="color:#f92672">.&lt;/span>view(batch,time,self&lt;span style="color:#f92672">.&lt;/span>n_head,n_d)&lt;span style="color:#f92672">.&lt;/span>transpose(&lt;span style="color:#ae81ff">1&lt;/span>,&lt;span style="color:#ae81ff">2&lt;/span>) &lt;span style="color:#75715e"># self.n_head * n_d = d_model&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> k &lt;span style="color:#f92672">=&lt;/span> k&lt;span style="color:#f92672">.&lt;/span>view(batch,time,self&lt;span style="color:#f92672">.&lt;/span>n_head,n_d)&lt;span style="color:#f92672">.&lt;/span>transpose(&lt;span style="color:#ae81ff">1&lt;/span>,&lt;span style="color:#ae81ff">2&lt;/span>) &lt;span style="color:#75715e"># shape: (batch,n_head,time,n_d)&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> v &lt;span style="color:#f92672">=&lt;/span> v&lt;span style="color:#f92672">.&lt;/span>view(batch,time,self&lt;span style="color:#f92672">.&lt;/span>n_head,n_d)&lt;span style="color:#f92672">.&lt;/span>transpose(&lt;span style="color:#ae81ff">1&lt;/span>,&lt;span style="color:#ae81ff">2&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:#75715e"># compute attention by the formula&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> score &lt;span style="color:#f92672">=&lt;/span> q &lt;span style="color:#f92672">@&lt;/span> k&lt;span style="color:#f92672">.&lt;/span>transpose(&lt;span style="color:#f92672">-&lt;/span>&lt;span style="color:#ae81ff">2&lt;/span>,&lt;span style="color:#f92672">-&lt;/span>&lt;span style="color:#ae81ff">1&lt;/span>) &lt;span style="color:#f92672">/&lt;/span> math&lt;span style="color:#f92672">.&lt;/span>sqrt(n_d) &lt;span style="color:#75715e"># (batch,n_head,time,time)&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#66d9ef">if&lt;/span> mask &lt;span style="color:#f92672">is&lt;/span> &lt;span style="color:#f92672">not&lt;/span> &lt;span style="color:#66d9ef">None&lt;/span>:
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> score &lt;span style="color:#f92672">=&lt;/span> score&lt;span style="color:#f92672">.&lt;/span>masked_fill(mask&lt;span style="color:#f92672">==&lt;/span>&lt;span style="color:#ae81ff">0&lt;/span>,float(&lt;span style="color:#e6db74">&amp;#39;-inf&amp;#39;&lt;/span>))
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> attn &lt;span style="color:#f92672">=&lt;/span> self&lt;span style="color:#f92672">.&lt;/span>softmax(score) &lt;span style="color:#f92672">@&lt;/span> v &lt;span style="color:#75715e"># (batch,n_head,time,n_d)&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> attn &lt;span style="color:#f92672">=&lt;/span> attn&lt;span style="color:#f92672">.&lt;/span>transpose(&lt;span style="color:#ae81ff">1&lt;/span>,&lt;span style="color:#ae81ff">2&lt;/span>)&lt;span style="color:#f92672">.&lt;/span>contiguous()&lt;span style="color:#f92672">.&lt;/span>view(batch,time,dimension) &lt;span style="color:#75715e"># (batch,time,dimension)&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> attn &lt;span style="color:#f92672">=&lt;/span> self&lt;span style="color:#f92672">.&lt;/span>w_conbine(attn) &lt;span style="color:#75715e"># (batch,time,dimension)&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#66d9ef">return&lt;/span> attn
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>d_model &lt;span style="color:#f92672">=&lt;/span> &lt;span style="color:#ae81ff">512&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>n_head &lt;span style="color:#f92672">=&lt;/span> &lt;span style="color:#ae81ff">8&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>mha &lt;span style="color:#f92672">=&lt;/span> MultiHeadAttention(d_model,n_head)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>x &lt;span style="color:#f92672">=&lt;/span> torch&lt;span style="color:#f92672">.&lt;/span>randn(&lt;span style="color:#ae81ff">2&lt;/span>,&lt;span style="color:#ae81ff">10&lt;/span>,d_model)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>y &lt;span style="color:#f92672">=&lt;/span> mha(x,x,x) &lt;span style="color:#75715e"># self-multi-head-attention&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>print(&lt;span style="color:#e6db74">f&lt;/span>&lt;span style="color:#e6db74">&amp;#34;shape: &lt;/span>&lt;span style="color:#e6db74">{&lt;/span>y&lt;span style="color:#f92672">.&lt;/span>shape&lt;span style="color:#e6db74">}&lt;/span>&lt;span style="color:#ae81ff">\n&lt;/span>&lt;span style="color:#e6db74">输出:&lt;/span>&lt;span style="color:#e6db74">{&lt;/span>y&lt;span style="color:#e6db74">}&lt;/span>&lt;span style="color:#e6db74">&amp;#34;&lt;/span>) &lt;span style="color:#75715e"># torch.Size([2, 10, 512])&lt;/span>
&lt;/span>&lt;/span>&lt;/code>&lt;/pre>&lt;/div>&lt;h3 id="token和位置嵌入模块">Token和位置嵌入模块&lt;/h3>
&lt;p>Token Embedding 是将离散的词（token）转换为连续的向量表示的过程。通常通过查找表（如 nn.Embedding）将每个词的索引映射为一个高维向量，使模型能够处理和学习词之间的语义关系。
公式表示： &lt;/p></description></item></channel></rss>