<?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/%E7%BB%8F%E5%85%B8%E7%AE%97%E6%B3%95/</link><description>Recent content in 经典算法 on Ganko Space</description><generator>Hugo</generator><language>zh-cn</language><lastBuildDate>Wed, 24 Sep 2025 00:00:00 +0000</lastBuildDate><atom:link href="https://ganko.asia/tags/%E7%BB%8F%E5%85%B8%E7%AE%97%E6%B3%95/index.xml" rel="self" type="application/rss+xml"/><item><title>CodeByMySelf-Diffusion</title><link>https://ganko.asia/posts/codebymyself-diffusion/</link><pubDate>Wed, 24 Sep 2025 00:00:00 +0000</pubDate><guid>https://ganko.asia/posts/codebymyself-diffusion/</guid><description>&lt;blockquote>
&lt;p>所有之前：CodeByMySelf 系列之 Diffusion Models 的实现记录，变体不会在这篇里更新。&lt;/p>&lt;/blockquote>
&lt;h2 id="diffusion-models">Diffusion Models&lt;/h2>
&lt;h3 id="原理">原理&lt;/h3>
&lt;p>具体的一些公式推导请看李宏毅老师的视频：
&lt;a href="https://www.bilibili.com/video/BV1mLbQeRExa?t=0.0">李宏毅Diffusion&lt;/a>&lt;/p>
&lt;h3 id="基本概念">基本概念&lt;/h3>
&lt;p>（以图像生成为例）首先Diffusion的基本运作过程是，首先Sample出一个全是噪声的图片，然后经过一步一步的Denoise最终得到图片，Denoise的步数是事先确定好的。&lt;/p>
&lt;p>&lt;strong>在Diffusion中，深度学习模型的唯一作用就是预测噪声，其他过程均由数学公式推理得到&lt;/strong>&lt;/p>
&lt;p>进一步说，不同情况的噪声需要的处理肯定是不一样的，一个Denoise模型肯定无法做到处理不同的情况的噪声，因此Denoise模型需要一个“噪声严重程度”作为输入，在Diffusion中称为时间步。&lt;/p>
&lt;p>Denoise内部的工作是，&lt;strong>接受带Noise的图片，对噪声进行预测，然后使用带Noise的图片减去预测噪声那么就可以得到。&lt;/strong>&lt;/p>
&lt;p align="center">
&lt;img src="https://ganko.asia/images/CodeByMyself/基本概念.png" alt="基本概念"/>
&lt;/p>
&lt;p>我们清楚Noise predictor模型的输入输出之后，该怎么训练这个模型呢？已知模型的输入-输出对是Noise图片加上“噪声程度”-Noise数据对，我们需要人为创造这个数据对出来，具体来说，从image list中拿一张图片，然后从GasuionN中采样出一张纯噪声的图片，加了一定的步数之后，就得到了很多数据对（Noise图片、时间步&amp;ndash;噪声），这个过程称为Diffusion Process。&lt;/p>
&lt;p>现在，已经可以从一个纯噪声得到一张图片，但是我们不能根据噪声随机生成图片，所以我们需要把文字考虑进来，让模型考虑进文字之后再预测噪声。&lt;/p>
&lt;p align="center">
&lt;img src="https://ganko.asia/images/CodeByMyself/考虑text.png" alt="考虑文字"/>
&lt;/p>
&lt;h3 id="完整算法">完整算法&lt;/h3>
&lt;p align="center">
&lt;img src="https://ganko.asia/images/CodeByMyself/pric.PNG" alt="完整算法"/>
&lt;/p>
&lt;p>首先是Training的过程，从immage list中取出一张干净的图，之后从1-T中随机取一个值作为时间步，然后从Gaussion中取样出一个noise，之后就是一步一步的梯度更新训练模型。这种方法和我们的想法有些不一样，我们想像中noise是一步一步加进去的，但是在实际中是一下子直接加进去的。&lt;/p>
&lt;p align="center">
&lt;img src="https://ganko.asia/images/CodeByMyself/addnoise.png" alt="加噪过程"/>
&lt;/p>
&lt;p>具体的加噪公式推导过程如下：&lt;/p>
&lt;p align="center">
&lt;img src="https://ganko.asia/images/CodeByMyself/addnoise_formula.png" alt="加噪公式"/>
&lt;/p>
&lt;p>然后是推理过程，也就是Sampling这个过程。首先，采样得到一个纯噪声图片X_T，然后进行T次公式中的循环，$\epsilon_\theta$表示预测噪声的模型。&lt;/p>
&lt;p>但是为什么要加这个z噪声：每次都要重新采样出一个噪声，但是只有当t&amp;gt;1的时候.这个z的使用和生成语句中一个问题类似：为什么每次总要SAMPLE而不是直接取Mean（概率最大）？在&lt;a href="https://arxiv.org/abs/1904.09751">The curious case of nerual text degeneration&lt;/a>论文中有一个分析，当直接取Mean时候，会一直重复相同的回答和出现跳帧现象，并且他们分析得到人在写文章的时候并不会一定会选概率最大的词。&lt;/p>
&lt;h3 id="简单代码实现">简单代码实现&lt;/h3>
&lt;p>Diffusion已经有了许多变体，下面是一个简单的实现，变体不会在这篇里更新。&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:#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> math
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>&lt;span style="color:#f92672">import&lt;/span> torch.nn.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> numpy &lt;span style="color:#66d9ef">as&lt;/span> np
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>&lt;span style="color:#f92672">import&lt;/span> time
&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">WeightedLoss&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):
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> super(WeightedLoss,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">def&lt;/span> &lt;span style="color:#a6e22e">forward&lt;/span>(self,pred,target,weighted&lt;span style="color:#f92672">=&lt;/span>&lt;span style="color:#ae81ff">1.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> loss &lt;span style="color:#f92672">=&lt;/span> self&lt;span style="color:#f92672">.&lt;/span>_loss(pred,target)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> weighted_loss &lt;span style="color:#f92672">=&lt;/span> (loss &lt;span style="color:#f92672">*&lt;/span> weighted)&lt;span style="color:#f92672">.&lt;/span>mean() &lt;span style="color:#75715e"># mean by batch&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#66d9ef">return&lt;/span> weighted_loss
&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">L1Loss&lt;/span>(WeightedLoss):
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#66d9ef">def&lt;/span> &lt;span style="color:#a6e22e">_loss&lt;/span>(self,pred,target):
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#66d9ef">return&lt;/span> torch&lt;span style="color:#f92672">.&lt;/span>abs(pred &lt;span style="color:#f92672">-&lt;/span> target)
&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">L2Loss&lt;/span>(WeightedLoss):
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#66d9ef">def&lt;/span> &lt;span style="color:#a6e22e">_loss&lt;/span>(self,pred,target):
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#66d9ef">return&lt;/span> (pred &lt;span style="color:#f92672">-&lt;/span> target) &lt;span style="color:#f92672">**&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>Losses &lt;span style="color:#f92672">=&lt;/span> {
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#e6db74">&amp;#39;l1&amp;#39;&lt;/span>:L1Loss,
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#e6db74">&amp;#39;l2&amp;#39;&lt;/span>:L2Loss
&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">SinusoidalPosEmb&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,dim):
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> super(SinusoidalPosEmb,self)&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>dim &lt;span style="color:#f92672">=&lt;/span> dim
&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,t):
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#e6db74">&amp;#39;&amp;#39;&amp;#39;
&lt;/span>&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>&lt;span style="color:#e6db74"> 扩散模型时间步的位置编码公式和Transformer中的位置编码公式是类似的。
&lt;/span>&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>&lt;span style="color:#e6db74"> &amp;#39;&amp;#39;&amp;#39;&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> device &lt;span style="color:#f92672">=&lt;/span> t&lt;span style="color:#f92672">.&lt;/span>device
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> half_dim &lt;span style="color:#f92672">=&lt;/span> self&lt;span style="color:#f92672">.&lt;/span>dim &lt;span style="color:#f92672">//&lt;/span> &lt;span style="color:#ae81ff">2&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> emb &lt;span style="color:#f92672">=&lt;/span> math&lt;span style="color:#f92672">.&lt;/span>log(&lt;span style="color:#ae81ff">10000&lt;/span>) &lt;span style="color:#f92672">/&lt;/span> (half_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> emb &lt;span style="color:#f92672">=&lt;/span> torch&lt;span style="color:#f92672">.&lt;/span>exp(torch&lt;span style="color:#f92672">.&lt;/span>arange(half_dim,device&lt;span style="color:#f92672">=&lt;/span>device) &lt;span style="color:#f92672">*&lt;/span> &lt;span style="color:#f92672">-&lt;/span>emb)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> emb &lt;span style="color:#f92672">=&lt;/span> t[:,&lt;span style="color:#66d9ef">None&lt;/span>] &lt;span style="color:#f92672">*&lt;/span> emb[&lt;span style="color:#66d9ef">None&lt;/span>,:]
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> emb &lt;span style="color:#f92672">=&lt;/span> torch&lt;span style="color:#f92672">.&lt;/span>cat((emb&lt;span style="color:#f92672">.&lt;/span>sin(),emb&lt;span style="color:#f92672">.&lt;/span>cos()),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:#66d9ef">return&lt;/span> emb
&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">extract&lt;/span>(a,t,x_shape):
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#e6db74">&amp;#39;&amp;#39;&amp;#39;
&lt;/span>&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>&lt;span style="color:#e6db74"> a: (T,) t: (B,) x_shape: (B,...)
&lt;/span>&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>&lt;span style="color:#e6db74"> &amp;#39;&amp;#39;&amp;#39;&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> batch_size &lt;span style="color:#f92672">=&lt;/span> t&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> out &lt;span style="color:#f92672">=&lt;/span> a&lt;span style="color:#f92672">.&lt;/span>gather(&lt;span style="color:#f92672">-&lt;/span>&lt;span style="color:#ae81ff">1&lt;/span>,t) &lt;span style="color:#75715e"># (B,)&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(batch_size, &lt;span style="color:#f92672">*&lt;/span>((&lt;span style="color:#ae81ff">1&lt;/span>,) &lt;span style="color:#f92672">*&lt;/span> (len(x_shape) &lt;span style="color:#f92672">-&lt;/span> &lt;span style="color:#ae81ff">1&lt;/span>))) &lt;span style="color:#75715e"># (B,1,1,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>&lt;/span>&lt;span style="display:flex;">&lt;span>&lt;span style="color:#66d9ef">class&lt;/span> &lt;span style="color:#a6e22e">MLP&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:#e6db74">&amp;#39;&amp;#39;&amp;#39;
&lt;/span>&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>&lt;span style="color:#e6db74"> 这是一个简单的多层感知机(MLP)模型，结合了时间步的位置编码，用于处理状态和动作的输入（强化学习）
&lt;/span>&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>&lt;span style="color:#e6db74"> &amp;#39;&amp;#39;&amp;#39;&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#66d9ef">def&lt;/span> __init__(self,state_dim,action_dim,hidden_dim,device,t_dim):
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> super(MLP,self)&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>device &lt;span style="color:#f92672">=&lt;/span> device
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> self&lt;span style="color:#f92672">.&lt;/span>t_dim &lt;span style="color:#f92672">=&lt;/span> t_dim
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> self&lt;span style="color:#f92672">.&lt;/span>a_tim &lt;span style="color:#f92672">=&lt;/span> action_dim
&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>time_mlp &lt;span style="color:#f92672">=&lt;/span> nn&lt;span style="color:#f92672">.&lt;/span>Sequential(
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> SinusoidalPosEmb(t_dim),
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> nn&lt;span style="color:#f92672">.&lt;/span>Linear(t_dim,t_dim&lt;span style="color:#f92672">*&lt;/span>&lt;span style="color:#ae81ff">2&lt;/span>),
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> nn&lt;span style="color:#f92672">.&lt;/span>Mish(),
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> nn&lt;span style="color:#f92672">.&lt;/span>Linear(t_dim&lt;span style="color:#f92672">*&lt;/span>&lt;span style="color:#ae81ff">2&lt;/span>,t_dim)
&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> input_dim &lt;span style="color:#f92672">=&lt;/span> state_dim &lt;span style="color:#f92672">+&lt;/span> action_dim &lt;span style="color:#f92672">+&lt;/span> t_dim
&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>mid_layer &lt;span style="color:#f92672">=&lt;/span> nn&lt;span style="color:#f92672">.&lt;/span>Sequential(
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> nn&lt;span style="color:#f92672">.&lt;/span>Linear(input_dim,hidden_dim),
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> nn&lt;span style="color:#f92672">.&lt;/span>Mish(),
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> nn&lt;span style="color:#f92672">.&lt;/span>Linear(hidden_dim,hidden_dim),
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> nn&lt;span style="color:#f92672">.&lt;/span>Mish(),
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> nn&lt;span style="color:#f92672">.&lt;/span>Linear(hidden_dim,hidden_dim),
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> nn&lt;span style="color:#f92672">.&lt;/span>Mish(),
&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>f_layer &lt;span style="color:#f92672">=&lt;/span> nn&lt;span style="color:#f92672">.&lt;/span>Linear(hidden_dim,action_dim)
&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">init_weights&lt;/span>(self):
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#e6db74">&amp;#39;&amp;#39;&amp;#39;
&lt;/span>&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>&lt;span style="color:#e6db74"> 初始化模型权重,好的权重初始化可以帮助模型更快收敛，提升训练效果。
&lt;/span>&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>&lt;span style="color:#e6db74"> &amp;#39;&amp;#39;&amp;#39;&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#66d9ef">for&lt;/span> m &lt;span style="color:#f92672">in&lt;/span> self&lt;span style="color:#f92672">.&lt;/span>modules():
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#66d9ef">if&lt;/span> isinstance(m,nn&lt;span style="color:#f92672">.&lt;/span>Linear):
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> nn&lt;span style="color:#f92672">.&lt;/span>init&lt;span style="color:#f92672">.&lt;/span>xavier_normal_(m&lt;span style="color:#f92672">.&lt;/span>weight)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#66d9ef">if&lt;/span> m&lt;span style="color:#f92672">.&lt;/span>bias &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> nn&lt;span style="color:#f92672">.&lt;/span>init&lt;span style="color:#f92672">.&lt;/span>zeros_(m&lt;span style="color:#f92672">.&lt;/span>bias)
&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,state,action,t):
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> t_emb &lt;span style="color:#f92672">=&lt;/span> self&lt;span style="color:#f92672">.&lt;/span>time_mlp(t) &lt;span style="color:#75715e"># (batch,t_dim)&lt;/span>
&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>cat((state,action,t_emb),dim&lt;span style="color:#f92672">=-&lt;/span>&lt;span style="color:#ae81ff">1&lt;/span>) &lt;span style="color:#75715e"># (batch,state_dim+action_dim+t_dim)&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> x &lt;span style="color:#f92672">=&lt;/span> self&lt;span style="color:#f92672">.&lt;/span>mid_layer(x) &lt;span style="color:#75715e"># (batch,hidden_dim)&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> action_delta &lt;span style="color:#f92672">=&lt;/span> self&lt;span style="color:#f92672">.&lt;/span>f_layer(x) &lt;span style="color:#75715e"># (batch,action_dim)&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#66d9ef">return&lt;/span> action_delta
&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">Diffusion&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,loss_type,beta_schedule&lt;span style="color:#f92672">=&lt;/span>&lt;span style="color:#e6db74">&amp;#39;linear&amp;#39;&lt;/span>,clip_denoise&lt;span style="color:#f92672">=&lt;/span>&lt;span style="color:#66d9ef">True&lt;/span>,predict_epsilon&lt;span style="color:#f92672">=&lt;/span>&lt;span style="color:#66d9ef">True&lt;/span>,&lt;span style="color:#f92672">**&lt;/span>kwargs):
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#e6db74">&amp;#39;&amp;#39;&amp;#39;
&lt;/span>&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>&lt;span style="color:#e6db74"> loss_type: 损失函数类型
&lt;/span>&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>&lt;span style="color:#e6db74"> beta_schedule: beta(时间)的离散方式，常见的有&amp;#39;linear&amp;#39;（线性）和&amp;#39;cosine&amp;#39;（余弦）。
&lt;/span>&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>&lt;span style="color:#e6db74"> clip_denoise: 是否在去噪过程中裁剪输出，以防止数值过大。
&lt;/span>&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>&lt;span style="color:#e6db74"> &amp;#39;&amp;#39;&amp;#39;&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> super(Diffusion,self)&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>device &lt;span style="color:#f92672">=&lt;/span> kwargs&lt;span style="color:#f92672">.&lt;/span>get(&lt;span style="color:#e6db74">&amp;#39;device&amp;#39;&lt;/span>,&lt;span style="color:#e6db74">&amp;#39;cpu&amp;#39;&lt;/span>)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> self&lt;span style="color:#f92672">.&lt;/span>t_dim &lt;span style="color:#f92672">=&lt;/span> kwargs&lt;span style="color:#f92672">.&lt;/span>get(&lt;span style="color:#e6db74">&amp;#39;t_dim&amp;#39;&lt;/span>,&lt;span style="color:#ae81ff">16&lt;/span>)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> self&lt;span style="color:#f92672">.&lt;/span>state_dim &lt;span style="color:#f92672">=&lt;/span> kwargs&lt;span style="color:#f92672">.&lt;/span>get(&lt;span style="color:#e6db74">&amp;#39;state_dim&amp;#39;&lt;/span>,&lt;span style="color:#ae81ff">10&lt;/span>)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> self&lt;span style="color:#f92672">.&lt;/span>action_dim &lt;span style="color:#f92672">=&lt;/span> kwargs&lt;span style="color:#f92672">.&lt;/span>get(&lt;span style="color:#e6db74">&amp;#39;action_dim&amp;#39;&lt;/span>,&lt;span style="color:#ae81ff">2&lt;/span>)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> self&lt;span style="color:#f92672">.&lt;/span>hidden_dim &lt;span style="color:#f92672">=&lt;/span> kwargs&lt;span style="color:#f92672">.&lt;/span>get(&lt;span style="color:#e6db74">&amp;#39;hidden_dim&amp;#39;&lt;/span>,&lt;span style="color:#ae81ff">256&lt;/span>)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> self&lt;span style="color:#f92672">.&lt;/span>loss_type &lt;span style="color:#f92672">=&lt;/span> loss_type
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> self&lt;span style="color:#f92672">.&lt;/span>beta_schedule &lt;span style="color:#f92672">=&lt;/span> beta_schedule
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> self&lt;span style="color:#f92672">.&lt;/span>clip_denoise &lt;span style="color:#f92672">=&lt;/span> clip_denoise
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> self&lt;span style="color:#f92672">.&lt;/span>T &lt;span style="color:#f92672">=&lt;/span> kwargs&lt;span style="color:#f92672">.&lt;/span>get(&lt;span style="color:#e6db74">&amp;#39;timesteps&amp;#39;&lt;/span>,&lt;span style="color:#ae81ff">1000&lt;/span>)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> self&lt;span style="color:#f92672">.&lt;/span>model &lt;span style="color:#f92672">=&lt;/span> MLP(self&lt;span style="color:#f92672">.&lt;/span>state_dim,self&lt;span style="color:#f92672">.&lt;/span>action_dim,self&lt;span style="color:#f92672">.&lt;/span>hidden_dim,self&lt;span style="color:#f92672">.&lt;/span>device,self&lt;span style="color:#f92672">.&lt;/span>t_dim)&lt;span style="color:#f92672">.&lt;/span>to(self&lt;span style="color:#f92672">.&lt;/span>device)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> self&lt;span style="color:#f92672">.&lt;/span>kwargs &lt;span style="color:#f92672">=&lt;/span> kwargs
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> self&lt;span style="color:#f92672">.&lt;/span>model&lt;span style="color:#f92672">.&lt;/span>init_weights()
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#66d9ef">if&lt;/span> beta_schedule &lt;span style="color:#f92672">==&lt;/span> &lt;span style="color:#e6db74">&amp;#39;linear&amp;#39;&lt;/span>:
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> betas &lt;span style="color:#f92672">=&lt;/span> torch&lt;span style="color:#f92672">.&lt;/span>linspace(&lt;span style="color:#ae81ff">1e-4&lt;/span>,&lt;span style="color:#ae81ff">0.02&lt;/span>,self&lt;span style="color:#f92672">.&lt;/span>kwargs&lt;span style="color:#f92672">.&lt;/span>get(&lt;span style="color:#e6db74">&amp;#39;timesteps&amp;#39;&lt;/span>, &lt;span style="color:#ae81ff">1000&lt;/span>),dtype&lt;span style="color:#f92672">=&lt;/span>torch&lt;span style="color:#f92672">.&lt;/span>float32)
&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> alphas &lt;span style="color:#f92672">=&lt;/span> &lt;span style="color:#ae81ff">1.&lt;/span> &lt;span style="color:#f92672">-&lt;/span> betas
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> alphas_cumprod &lt;span style="color:#f92672">=&lt;/span> torch&lt;span style="color:#f92672">.&lt;/span>cumprod(alphas,dim&lt;span style="color:#f92672">=&lt;/span>&lt;span style="color:#ae81ff">0&lt;/span>) &lt;span style="color:#75715e"># [1,2,3] -&amp;gt; [1,1*2,1*2*3]&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> alphas_cumprod_prev &lt;span style="color:#f92672">=&lt;/span> torch&lt;span style="color:#f92672">.&lt;/span>cat((torch&lt;span style="color:#f92672">.&lt;/span>tensor([&lt;span style="color:#ae81ff">1.&lt;/span>],dtype&lt;span style="color:#f92672">=&lt;/span>torch&lt;span style="color:#f92672">.&lt;/span>float32),alphas_cumprod[:&lt;span style="color:#f92672">-&lt;/span>&lt;span style="color:#ae81ff">1&lt;/span>]),dim&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>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> self&lt;span style="color:#f92672">.&lt;/span>register_buffer(&lt;span style="color:#e6db74">&amp;#39;betas&amp;#39;&lt;/span>,betas)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> self&lt;span style="color:#f92672">.&lt;/span>register_buffer(&lt;span style="color:#e6db74">&amp;#39;alphas&amp;#39;&lt;/span>,alphas)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> self&lt;span style="color:#f92672">.&lt;/span>register_buffer(&lt;span style="color:#e6db74">&amp;#39;alphas_cumprod&amp;#39;&lt;/span>,alphas_cumprod)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> self&lt;span style="color:#f92672">.&lt;/span>register_buffer(&lt;span style="color:#e6db74">&amp;#39;alphas_cumprod_prev&amp;#39;&lt;/span>,alphas_cumprod_prev)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> self&lt;span style="color:#f92672">.&lt;/span>register_buffer(&lt;span style="color:#e6db74">&amp;#39;sqrt_alphas_cumprod&amp;#39;&lt;/span>,torch&lt;span style="color:#f92672">.&lt;/span>sqrt(alphas_cumprod))
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#75715e"># （前向过程）&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> self&lt;span style="color:#f92672">.&lt;/span>register_buffer(&lt;span style="color:#e6db74">&amp;#39;sqrt_alphas_cumprob&amp;#39;&lt;/span>,torch&lt;span style="color:#f92672">.&lt;/span>sqrt(alphas_cumprod))
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> self&lt;span style="color:#f92672">.&lt;/span>register_buffer(&lt;span style="color:#e6db74">&amp;#39;sqrt_one_minus_alphas_cumprod&amp;#39;&lt;/span>,torch&lt;span style="color:#f92672">.&lt;/span>sqrt(&lt;span style="color:#ae81ff">1.&lt;/span> &lt;span style="color:#f92672">-&lt;/span> alphas_cumprod))
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#75715e"># （反向过程）&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> posterior_variance &lt;span style="color:#f92672">=&lt;/span> betas &lt;span style="color:#f92672">*&lt;/span> (&lt;span style="color:#ae81ff">1.&lt;/span> &lt;span style="color:#f92672">-&lt;/span> alphas_cumprod_prev) &lt;span style="color:#f92672">/&lt;/span> (&lt;span style="color:#ae81ff">1.&lt;/span> &lt;span style="color:#f92672">-&lt;/span> alphas_cumprod)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> self&lt;span style="color:#f92672">.&lt;/span>register_buffer(&lt;span style="color:#e6db74">&amp;#39;posterior_variance&amp;#39;&lt;/span>,posterior_variance)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#75715e"># 用于从xt求得x0&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> self&lt;span style="color:#f92672">.&lt;/span>register_buffer(&lt;span style="color:#e6db74">&amp;#39;sqrt_recip_alphas_cumprod&amp;#39;&lt;/span>,torch&lt;span style="color:#f92672">.&lt;/span>sqrt(&lt;span style="color:#ae81ff">1.&lt;/span> &lt;span style="color:#f92672">/&lt;/span> alphas_cumprod))
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> self&lt;span style="color:#f92672">.&lt;/span>register_buffer(&lt;span style="color:#e6db74">&amp;#39;sqrt_recipm1_alphas_cumprod&amp;#39;&lt;/span>,torch&lt;span style="color:#f92672">.&lt;/span>sqrt(&lt;span style="color:#ae81ff">1.&lt;/span> &lt;span style="color:#f92672">/&lt;/span> alphas_cumprod &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> self&lt;span style="color:#f92672">.&lt;/span>register_buffer(&lt;span style="color:#e6db74">&amp;#39;posterior_mean_coef1&amp;#39;&lt;/span>,betas &lt;span style="color:#f92672">*&lt;/span>torch&lt;span style="color:#f92672">.&lt;/span>sqrt(alphas_cumprod_prev) &lt;span style="color:#f92672">/&lt;/span> (&lt;span style="color:#ae81ff">1.&lt;/span> &lt;span style="color:#f92672">-&lt;/span> alphas_cumprod))
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> self&lt;span style="color:#f92672">.&lt;/span>register_buffer(&lt;span style="color:#e6db74">&amp;#39;posterior_mean_coef2&amp;#39;&lt;/span>, (&lt;span style="color:#ae81ff">1.&lt;/span> &lt;span style="color:#f92672">-&lt;/span> alphas_cumprod_prev) &lt;span style="color:#f92672">*&lt;/span> torch&lt;span style="color:#f92672">.&lt;/span>sqrt(alphas) &lt;span style="color:#f92672">/&lt;/span> (&lt;span style="color:#ae81ff">1.&lt;/span> &lt;span style="color:#f92672">-&lt;/span> alphas_cumprod))
&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>loss_func &lt;span style="color:#f92672">=&lt;/span> Losses[loss_type]()
&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,state,&lt;span style="color:#f92672">*&lt;/span>args,&lt;span style="color:#f92672">**&lt;/span>kwargs):
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#66d9ef">return&lt;/span> self&lt;span style="color:#f92672">.&lt;/span>sample(state,&lt;span style="color:#f92672">*&lt;/span>args,&lt;span style="color:#f92672">**&lt;/span>kwargs)
&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">sample&lt;/span>(self,state,&lt;span style="color:#f92672">*&lt;/span>args,&lt;span style="color:#f92672">**&lt;/span>kwargs):
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> batch_size &lt;span style="color:#f92672">=&lt;/span> state&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> shape &lt;span style="color:#f92672">=&lt;/span> [batch_size,self&lt;span style="color:#f92672">.&lt;/span>action_dim]
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> action &lt;span style="color:#f92672">=&lt;/span> self&lt;span style="color:#f92672">.&lt;/span>p_sample_loop(state,shape,&lt;span style="color:#f92672">*&lt;/span>args,&lt;span style="color:#f92672">**&lt;/span>kwargs)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#66d9ef">return&lt;/span> action&lt;span style="color:#f92672">.&lt;/span>clamp(&lt;span style="color:#f92672">-&lt;/span>&lt;span style="color:#ae81ff">1.&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">q_posterior&lt;/span>(self,x0,xt,t):
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#e6db74">&amp;#39;&amp;#39;&amp;#39;
&lt;/span>&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>&lt;span style="color:#e6db74"> 计算后验分布q(x_{t-1}|x_t,x_0)的均值和方差
&lt;/span>&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>&lt;span style="color:#e6db74"> &amp;#39;&amp;#39;&amp;#39;&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> posterior_mean &lt;span style="color:#f92672">=&lt;/span> (extract(self&lt;span style="color:#f92672">.&lt;/span>posterior_mean_coef1,t,x0&lt;span style="color:#f92672">.&lt;/span>shape) &lt;span style="color:#f92672">*&lt;/span> x0 &lt;span style="color:#f92672">+&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> extract(self&lt;span style="color:#f92672">.&lt;/span>posterior_mean_coef2,t,x0&lt;span style="color:#f92672">.&lt;/span>shape) &lt;span style="color:#f92672">*&lt;/span> xt)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> posterior_variance &lt;span style="color:#f92672">=&lt;/span> extract(self&lt;span style="color:#f92672">.&lt;/span>posterior_variance,t,x0&lt;span style="color:#f92672">.&lt;/span>shape)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> posterior_log_variance &lt;span style="color:#f92672">=&lt;/span> torch&lt;span style="color:#f92672">.&lt;/span>log(posterior_variance&lt;span style="color:#f92672">.&lt;/span>clamp(min&lt;span style="color:#f92672">=&lt;/span>&lt;span style="color:#ae81ff">1e-20&lt;/span>))
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#66d9ef">return&lt;/span> posterior_mean,posterior_variance,posterior_log_variance
&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">p_sample_loop&lt;/span>(self,state,shape,&lt;span style="color:#f92672">*&lt;/span>args,&lt;span style="color:#f92672">**&lt;/span>kwargs):
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> device &lt;span style="color:#f92672">=&lt;/span> self&lt;span style="color:#f92672">.&lt;/span>kwargs&lt;span style="color:#f92672">.&lt;/span>get(&lt;span style="color:#e6db74">&amp;#39;device&amp;#39;&lt;/span>,&lt;span style="color:#e6db74">&amp;#39;cpu&amp;#39;&lt;/span>)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> batch_size &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> x &lt;span style="color:#f92672">=&lt;/span> torch&lt;span style="color:#f92672">.&lt;/span>randn(shape,device&lt;span style="color:#f92672">=&lt;/span>device,requires_grad&lt;span style="color:#f92672">=&lt;/span>&lt;span style="color:#66d9ef">False&lt;/span>) &lt;span style="color:#75715e"># 标准DDPM噪声不需要梯度&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#66d9ef">for&lt;/span> i &lt;span style="color:#f92672">in&lt;/span> reversed(range(&lt;span style="color:#ae81ff">0&lt;/span>,self&lt;span style="color:#f92672">.&lt;/span>kwargs&lt;span style="color:#f92672">.&lt;/span>get(&lt;span style="color:#e6db74">&amp;#39;timesteps&amp;#39;&lt;/span>,&lt;span style="color:#ae81ff">1000&lt;/span>))):
&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>full((batch_size,),i,dtype&lt;span style="color:#f92672">=&lt;/span>torch&lt;span style="color:#f92672">.&lt;/span>long,device&lt;span style="color:#f92672">=&lt;/span>device)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> x &lt;span style="color:#f92672">=&lt;/span> self&lt;span style="color:#f92672">.&lt;/span>p_sample(x,t,state)
&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">return&lt;/span> x
&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">predict_x0_from_noise&lt;/span>(self,x,t,noise):
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#e6db74">&amp;#39;&amp;#39;&amp;#39;
&lt;/span>&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>&lt;span style="color:#e6db74"> 根据公式计算x0
&lt;/span>&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>&lt;span style="color:#e6db74"> &amp;#39;&amp;#39;&amp;#39;&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#66d9ef">return&lt;/span> (extract(self&lt;span style="color:#f92672">.&lt;/span>sqrt_recip_alphas_cumprod,t,x&lt;span style="color:#f92672">.&lt;/span>shape) &lt;span style="color:#f92672">*&lt;/span> x &lt;span style="color:#f92672">-&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> extract(self&lt;span style="color:#f92672">.&lt;/span>sqrt_recipm1_alphas_cumprod,t,x&lt;span style="color:#f92672">.&lt;/span>shape) &lt;span style="color:#f92672">*&lt;/span> noise)
&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">p_mean_variance&lt;/span>(self,x,t,state):
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> pred_noise &lt;span style="color:#f92672">=&lt;/span> self&lt;span style="color:#f92672">.&lt;/span>model(state,x,t)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> x_0 &lt;span style="color:#f92672">=&lt;/span> self&lt;span style="color:#f92672">.&lt;/span>predict_x0_from_noise(x,t,pred_noise)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> x_0&lt;span style="color:#f92672">.&lt;/span>clamp_(&lt;span style="color:#f92672">-&lt;/span>&lt;span style="color:#ae81ff">1.&lt;/span>,&lt;span style="color:#ae81ff">1.&lt;/span>) &lt;span style="color:#66d9ef">if&lt;/span> self&lt;span style="color:#f92672">.&lt;/span>clip_denoise &lt;span style="color:#66d9ef">else&lt;/span> x_0
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> model_mean,posterior_variance,posterior_log_variance &lt;span style="color:#f92672">=&lt;/span> self&lt;span style="color:#f92672">.&lt;/span>q_posterior(x_0,x,t)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#66d9ef">return&lt;/span> model_mean,posterior_log_variance
&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>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#66d9ef">def&lt;/span> &lt;span style="color:#a6e22e">p_sample&lt;/span>(self,x,t,state):
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#e6db74">&amp;#39;&amp;#39;&amp;#39;
&lt;/span>&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>&lt;span style="color:#e6db74"> x: 当前的噪声状态xt
&lt;/span>&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>&lt;span style="color:#e6db74"> t: 当前的时间步t
&lt;/span>&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>&lt;span style="color:#e6db74"> state: 状态
&lt;/span>&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>&lt;span style="color:#e6db74"> &amp;#39;&amp;#39;&amp;#39;&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> model_mean,model_log_variance &lt;span style="color:#f92672">=&lt;/span> self&lt;span style="color:#f92672">.&lt;/span>p_mean_variance(x,t,state)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> noise &lt;span style="color:#f92672">=&lt;/span> torch&lt;span style="color:#f92672">.&lt;/span>randn_like(x) &lt;span style="color:#66d9ef">if&lt;/span> t&lt;span style="color:#f92672">.&lt;/span>sum() &lt;span style="color:#f92672">&amp;gt;&lt;/span> &lt;span style="color:#ae81ff">0&lt;/span> &lt;span style="color:#66d9ef">else&lt;/span> &lt;span style="color:#ae81ff">0.&lt;/span> &lt;span style="color:#75715e"># 如果t=0,则不添加噪声&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#66d9ef">return&lt;/span> model_mean &lt;span style="color:#f92672">+&lt;/span> torch&lt;span style="color:#f92672">.&lt;/span>exp(&lt;span style="color:#ae81ff">0.5&lt;/span> &lt;span style="color:#f92672">*&lt;/span> model_log_variance) &lt;span style="color:#f92672">*&lt;/span> noise
&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">q_sample&lt;/span>(self,x0,t,noise):
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> sample &lt;span style="color:#f92672">=&lt;/span> (extract(self&lt;span style="color:#f92672">.&lt;/span>sqrt_alphas_cumprod,t,x0&lt;span style="color:#f92672">.&lt;/span>shape) &lt;span style="color:#f92672">*&lt;/span> x0 &lt;span style="color:#f92672">+&lt;/span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> extract(self&lt;span style="color:#f92672">.&lt;/span>sqrt_one_minus_alphas_cumprod,t,x0&lt;span style="color:#f92672">.&lt;/span>shape) &lt;span style="color:#f92672">*&lt;/span> noise)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#66d9ef">return&lt;/span> sample
&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">p_losses&lt;/span>(self,x0,state,t,weights):
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> noise &lt;span style="color:#f92672">=&lt;/span> torch&lt;span style="color:#f92672">.&lt;/span>randn_like(x0)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> x_noisy &lt;span style="color:#f92672">=&lt;/span> self&lt;span style="color:#f92672">.&lt;/span>q_sample(x0,t,noise)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> x_recon &lt;span style="color:#f92672">=&lt;/span> self&lt;span style="color:#f92672">.&lt;/span>model(state,x_noisy,t)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> loss &lt;span style="color:#f92672">=&lt;/span> self&lt;span style="color:#f92672">.&lt;/span>loss_func(x_recon,noise,weights)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#66d9ef">return&lt;/span> loss
&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">loss&lt;/span>(self,x,state,weights&lt;span style="color:#f92672">=&lt;/span>&lt;span style="color:#ae81ff">1.0&lt;/span>):
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> batch_size &lt;span style="color:#f92672">=&lt;/span> x&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> t &lt;span style="color:#f92672">=&lt;/span> torch&lt;span style="color:#f92672">.&lt;/span>randint(&lt;span style="color:#ae81ff">0&lt;/span>,self&lt;span style="color:#f92672">.&lt;/span>T,(batch_size,),device&lt;span style="color:#f92672">=&lt;/span>self&lt;span style="color:#f92672">.&lt;/span>device)&lt;span style="color:#f92672">.&lt;/span>long()
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> &lt;span style="color:#66d9ef">return&lt;/span> self&lt;span style="color:#f92672">.&lt;/span>p_losses(x,state,t,weights)
&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">if&lt;/span> __name__ &lt;span style="color:#f92672">==&lt;/span> &lt;span style="color:#e6db74">&amp;#39;__main__&amp;#39;&lt;/span>:
&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">256&lt;/span>,&lt;span style="color:#ae81ff">2&lt;/span>)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> state &lt;span style="color:#f92672">=&lt;/span> torch&lt;span style="color:#f92672">.&lt;/span>randn(&lt;span style="color:#ae81ff">256&lt;/span>,&lt;span style="color:#ae81ff">11&lt;/span>)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> model &lt;span style="color:#f92672">=&lt;/span> Diffusion(loss_type&lt;span style="color:#f92672">=&lt;/span>&lt;span style="color:#e6db74">&amp;#39;l2&amp;#39;&lt;/span>,beta_schedule&lt;span style="color:#f92672">=&lt;/span>&lt;span style="color:#e6db74">&amp;#39;linear&amp;#39;&lt;/span>,clip_denoise&lt;span style="color:#f92672">=&lt;/span>&lt;span style="color:#66d9ef">True&lt;/span>,predict_epsilon&lt;span style="color:#f92672">=&lt;/span>&lt;span style="color:#66d9ef">True&lt;/span>,
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> device&lt;span style="color:#f92672">=&lt;/span>&lt;span style="color:#e6db74">&amp;#39;cpu&amp;#39;&lt;/span>,t_dim&lt;span style="color:#f92672">=&lt;/span>&lt;span style="color:#ae81ff">16&lt;/span>,state_dim&lt;span style="color:#f92672">=&lt;/span>&lt;span style="color:#ae81ff">11&lt;/span>,action_dim&lt;span style="color:#f92672">=&lt;/span>&lt;span style="color:#ae81ff">2&lt;/span>,hidden_dim&lt;span style="color:#f92672">=&lt;/span>&lt;span style="color:#ae81ff">256&lt;/span>,timesteps&lt;span style="color:#f92672">=&lt;/span>&lt;span style="color:#ae81ff">100&lt;/span>)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span>
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> result &lt;span style="color:#f92672">=&lt;/span> model(state)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> print(result)
&lt;/span>&lt;/span>&lt;span style="display:flex;">&lt;span> loss &lt;span style="color:#f92672">=&lt;/span> model&lt;span style="color:#f92672">.&lt;/span>loss(x,state)
&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;loss: &lt;/span>&lt;span style="color:#e6db74">{&lt;/span>loss&lt;span style="color:#f92672">.&lt;/span>item()&lt;span style="color:#e6db74">}&lt;/span>&lt;span style="color:#e6db74">&amp;#34;&lt;/span>)
&lt;/span>&lt;/span>&lt;/code>&lt;/pre>&lt;/div></description></item><item><title>CodeByMySelf-RL</title><link>https://ganko.asia/posts/codebymyself-rl/</link><pubDate>Wed, 24 Sep 2025 00:00:00 +0000</pubDate><guid>https://ganko.asia/posts/codebymyself-rl/</guid><description>&lt;blockquote>
&lt;p>所有之前：CodeByMySelf 系列之 RL（深度强化学习）的实现记录，只有标题没有内容的是TODO。&lt;/p>&lt;/blockquote>
&lt;h2 id="rl强化学习深度强化学习">RL强化学习（深度强化学习）&lt;/h2>
&lt;blockquote>
&lt;p>TODO&lt;/p>&lt;/blockquote></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><item><title>CodeByMySelf-VAE</title><link>https://ganko.asia/posts/codebymyself-vae/</link><pubDate>Wed, 24 Sep 2025 00:00:00 +0000</pubDate><guid>https://ganko.asia/posts/codebymyself-vae/</guid><description>&lt;blockquote>
&lt;p>所有之前：CodeByMySelf 系列之 VAE（变分自编码器）的实现记录，只有标题没有内容的是TODO。&lt;/p>&lt;/blockquote>
&lt;h2 id="vae变分自编码器">VAE（变分自编码器）&lt;/h2>
&lt;p>VAE也是一个非常经典的生成模型，虽然现在主流已经是Diffusion模型，但是它的原理和思想还是非常有启发性的，尤其是对于理解diffusion中的一些概念（如潜在空间、重参数化技巧等）非常有帮助。&lt;/p>
&lt;blockquote>
&lt;p>TODO&lt;/p>&lt;/blockquote></description></item></channel></rss>