CS336 Lecture 7-8:Parallelism

本篇主题 这两节Lecture都是关于并行训练的,主要内容包括: 理解训练超大模型时系统层面的复杂性 掌握不同的 parallelization paradigms,以及为什么人们通常会同时使用多种并行方式 了解大规模训练任务通常是如何进行的 Lecture 7 - Parallelism 1 单块GPU无法满足SCALING的需求,必须使用多块GPU进行训练,所以我们需要Multi-GPU、Multi-Machine的并行训练方法,如下图所示: 首先讲了一些Basic Concepts: All Reduce:通信操作,所有参与的进程都将输入数据进行规约(如求和、最大值等)后,将结果分发给所有进程。常用于分布式训练中的梯度同步。 Broadcast:通信操作,数据从一个进程发送到所有其他进程。常用于分布式训练中的分发模型参数或者初始化数据。 All Gather:通信操作,所有参与的进程将各自的数据发送给所有其他进程,最终每个进程都获得所有数据的集合。常用于分布式训练中的收集模型输出或者中间结果。和All Reduce的区别在于,All Reduce会对数据进行规约操作(如求和),而All Gather只是简单地收集数据,不进行任何计算。 Reduce Scatter:通信操作,所有参与的进程将各自的数据发送给所有其他进程,并对数据进行规约操作(如求和),最终每个进程都获得规约后的结果的一部分。常用于分布式训练中的分布式梯度更新。 了解了初始知识之后,可以正式进入核心部分,不同的并行方式: Data Parallelism:最早的并行方式,模型复制到每个GPU上,每个GPU处理不同的数据batch,计算梯度后进行同步更新。优点是实现简单;缺点是通信开销大,尤其是模型参数较大时。 假设计算一个 SGD:我们会把 B 大小下的 batch 分给 M 个不同的机器,然后交换梯度去同步计算。这种情况对于 Compute scaling,每个 GPU 计算 B/M 个数据(不错!);对于 Communication overhead,每个 batch 需要转移两次梯度(发送和接受);对于 Memory scaling,完全没有,每个 GPU 都需要复制一次模型参数。 早期的data parallelism问题是只缓解了计算压力,内存压力没有缓解,需要每个GPU都复制一份模型参数,通信压力也很大。 Zero[HTTPS://arxiv.org/pdf/1910.02054]可以解决这个问题,核心的思想是将模型的state(参数和优化器状态)分布在不同的GPU上,每个GPU只存储模型的一部分,这样可以显著减少每个GPU的内存占用,同时通过通信操作来同步更新模型参数。 ZeRO分三个阶段: ZeRO-1:主要聚焦于 optimizer state sharding。把 optimizer state(first + second moments)分给每个 GPU,每个 GPU 都有 parameters + gradients,负责更新一部分 params。 每个 GPU 根据分配到的 batch 子集计算完整的梯度,此时只拥有局部梯度 利用 Reduce-Scatter 将所有 GPU 的局部梯度汇总并分散到每个设备上,现在每个 GPU 只持有全局梯度的一部分 每个 GPU 利用局部梯度和 optimizer state 更新该部分参数 利用 AllGather 将所有 GPU 的部分参数收集并分发给所有设备,确保每个 GPU 拥有完整的参数。 这种方法的通信开销并没有增加,而且 memory 减少了接近四倍(优化器状态是fp32,参数是fp16)。 ...

April 9, 2026 · 4 min

Pytorch-Tricks

写在前面:Pytorch 内容繁杂,实验中用到的技巧与机制若不及时记录,容易遗忘并导致重复造轮子。本文档用于记录 Pytorch 的相关技巧与机制,将持续更新,方便后续查阅。 动态计算图 缘起 在一次实验日志中,发现一个固定的 CNN 模块推理时间存在较大波动。排查后发现是数据流水线卡顿或 GPU 内存管理抖动所致,但也因此注意到了 Pytorch 的动态计算图机制。 Pytorch 在每次前向传播时都会重新构建计算图,这意味着每次前向传播都会产生额外开销,尤其在模型结构复杂或输入数据变化较大时,可能导致推理时间波动。 具体而言: 每次前向传播,即便模型结构未变,Pytorch 仍会依据当前输入数据和模型结构重建计算图。 若输入数据形状或模型结构发生变化,计算图会重新构建,进而引发性能波动。 可能导致性能波动的情况 1. 条件分支结构 模型中包含数据依赖的条件分支(如 if 语句)时,不同前向传播可能执行不同计算路径,导致计算图结构变化。 class BadEncoder(nn.Module): def forward(self, x): if x.sum() > 0: # 数据依赖的条件 x = self.heavy_conv(x) else: x = self.light_conv(x) return x 2. 变长序列 处理变长序列(如 RNN 或 Transformer)时,每次输入长度可能不同,导致计算图结构变化。 class BadRNN(nn.Module): def forward(self, x): for i in range(x.size(1)): # 依赖输入序列长度 x = self.rnn_cell(x[:, i, :], x) return x 3. 其他情况 动态形状 随机性操作 稀疏操作 解决方案 使用静态计算图:若模型结构固定,可借助 TorchScript 或 ONNX 转换为静态计算图,避免每次前向传播都重新构建。 Padding + Mask:对于变长序列,采用填充至统一长度,并利用 mask 标识有效部分,保持计算图结构一致。 模型训练/推理速度分析与优化 背景 修改模型与训练代码后发现速度明显下降。排查发现数据加载耗时过长,遂将 DataLoader 的在线预处理改为离线处理。但模型推理速度仍慢于基线,后续需进一步优化。 ...

April 7, 2026 · 4 min

CS336 Lecture 2: Resource Accounting

本讲主线 这节课围绕一个核心问题展开:如何精确估算训练成本。主要分为三块: 内存占用 FLOPs 估算 模型与训练配置 这和 Lecture 1 提到的“效率优先”直接相关。 例题:70B 模型训练时长 问题:使用 1024 张 H100,在 15T tokens 上训练一个 70B 模型,需要多久? 根据 Transformer 常用训练量估算公式($6 \times$ 参数量 $\times$ token 数): $$ \text{总 FLOPs} = 6 \times 70 \times 10^9 \times 15 \times 10^{12} = 6.3 \times 10^{24}\,\text{FLOPs} $$单卡 H100(BF16)理论性能约为 $2 \times 10^{15}$ FLOPs/s。若考虑大模型训练实际效率约为 30%,则单卡有效 FLOPs/s 为: $$ 0.3 \times 2 \times 10^{15} = 6 \times 10^{14}\,\text{FLOPs/s} $$总时长: $$ \text{总时间} = \frac{\text{总 FLOPs}}{\text{每秒有效 FLOPs} \times \text{GPU 数量}} = \frac{6.3 \times 10^{24}}{6 \times 10^{14} \times 1024} $$ 约等于 119 天。 ...

April 6, 2026 · 1 min

CS336 Lecture 3-4: Architectures, Hyperparameters, and MoE

本讲主线 老师推荐了 The Illustrated Transformer 作为辅助材料。 这节课基于 Tatsu H 对 2017-2025 年经典模型的统计,讨论两个问题: 不同模型在架构与超参数上有哪些共性与分歧? 这些演化背后的工程动机是什么? 架构与超参数观察 Pre-Norm vs Post-Norm 几乎所有现代 LLM(BERT 这类早期模型除外)都采用 Pre-Norm。核心原因是训练更稳定,梯度路径更顺畅。近年还出现了 double norm 等变体(如 Grok、Gemma 2),但尚未完全普及。 LayerNorm vs RMSNorm RMSNorm 不计算均值,只做均方根归一化,计算更简单。表面上看节省 FLOPs 不多,但在真实系统中,norm 的时间开销往往受 data movement 影响明显,因此 RMSNorm 的工程收益更像是在缓解内存墙,而不仅是减少算术量。 激活函数演化 从早期 ReLU,到 GPT 系列常见的 GeLU,再到近年广泛采用的 SwiGLU(门控激活变体)。门控结构通过逐元素调制,提升了表达能力与训练效率的折中表现。 串行层与并行层 传统 Transformer block 以串行堆叠为主。并行 block(如部分 GPT-J 设计)存在,但总体采用度不高。 位置编码路线 Sine/Absolute/Relative 都在不同阶段被主流模型采用。近年 RoPE(rotary position embeddings)几乎成为事实标准之一。 FFN 维度比(ffn dim : model dim) 传统经验常用 4 倍宽度;在 GLU 系列激活流行后,很多模型会选在 2.5-3.5 区间(常见约 8/3),以平衡参数规模与表达能力。 ...

April 6, 2026 · 1 min

CS336-写在所有之前

CS336 是一门很好的课,但是很多内容都比较基础,或者在网上已经有很多资源了。 2026/04/06:我学到了Lecture4,目前来看,CS336的精华内容还是在Lab上面,如果想要系统学习LLM的训练细节,建议直接看Lab的内容,Lecture部分可以作为辅助材料来理解一些概念和背景知识。 2026/04/09:Lecture4之后因为对并行训练的兴趣,直接跳到了Lecture7-8,感觉内容非常有价值,值得细看一下lecture视频或者讲义,精华在于不同并行方式的介绍和分析,在Lecture8里有一些简单的通信操作代码示例,帮助理解。

April 6, 2026 · 1 min

CS336 Lecture 1: Overview and Tokenization

课程导读 这门课的讲义基于 notebook 组织,不同章节通过函数模块化展开。 开场老师提出了一个很尖锐的问题:研究者正在与底层技术逐步脱节。他给了一个时间线: 八年前,研究者会自己补充数据并训练模型; 六年前,研究者还会下载模型后进行 fine-tune; 现在,很多人直接向闭源模型(GPT-4/Claude/Gemini)提问。 这虽然有些夸张,但确实反映了趋势:模型能力提升后,很多基础环节被“封装”了。老师强调: “Full understanding of this technology is necessary for fundamental research” 随着模型规模持续增大,训练中的计算压力正在从 Attention 逐步转向 FFN。 课程收获 Mechanics:系统如何工作(如 Transformer 结构、GPU 并行方式) Mindset:如何最大化利用硬件与内存,如何看待 scaling laws Intuition:哪些数据和建模决策更可能带来更好结果 老师也指出一个常见误解:只要堆算力,模型就会自动变好。他给出的观点是:准确率 = 效率 × 资源。 在数据和资源固定时,效率决定了上限。 老师将近年的 LLM 工作大致分成两类: 闭源路线:以 GPT 系列为代表,最早系统性拥抱 Scale,但细节封闭。 开放权重路线:以 Qwen 等 open-weight 模型为代表,拥抱 Scale 的同时提升开放度,但常见情况是只公开部分训练细节,数据与失败案例仍不完整。 Tokenization 为什么需要 tokenization?因为 LLM 在 token 序列上建模概率分布,我们需要把原始字符串编码为 token,并保证可逆解码。 🔗 Tokenization 可视化工具: tiktokenizer.vercel.app 常见 tokenization 方法: ...

April 3, 2026 · 1 min

博客网站优化日志

📅 2026-04-03 ✅ 已完成工作 环境搭建:Hugo+PaperMod 部署。 域名申请:ganko.asia 申请。 🚧 进行中(等待审核) 域名备案/审核:域名实名审核/备案状态,不能通过自定义域名访问。 📅 2026-04-04 头像发光效果,文章卡片悬浮效果。 📅 2026-04-07 域名审核通过,绑定成功,使用Nginx+Certbot配置HTTPS。 增加阅读的目录,直接跳转到文章不同部分。

April 3, 2026 · 1 min

胡思乱想

这里用来记录一些无论是产品角度的想法,还是算法角度的想法,或者是一些其他的想法。总之就是一些胡思乱想,这里不会引用任何论文,我也不会去查找任何资料,完全是凭空想象的。也就是说,这里记录的内容可能完全是错误的,甚至是荒谬的,但我觉得有趣,所以就记录下来. 神奇的特征空间 2026-4-1 在这个神奇的特征空间中,同一物体的特征不再是无规律的,而是可以通过运算得到的,例如: 一个朝向我可乐瓶的特征可以通过旋转得到一个背向我的可乐瓶的特征。 一半可乐瓶的特征应该是contact另一半可乐瓶的特征得到完整可乐瓶的特征。 具体来说,就是一个可旋转,拼接,以及不同角度拼接的特征空间,如果可以的话,这个特征空间应该是三维的(不过很明显这里的特征不包括-一瓶水甜还是苦这种单从图片无法推测的属性) VIT和CNN都有的惰性问题:依赖背景对主体进行识别 2026-4-18 在一些数据集中,模型可能会过度依赖背景信息来识别主体。例如,在一个包含大量猫的图片数据集中,模型可能会学会通过识别草地或室内环境来判断图片中是否有猫,而不是通过识别猫的特征。这种依赖背景的现象可能会导致模型在面对新的环境或背景时表现不佳,因为它没有真正学会识别主体的特征。 找到哪些受上下文影响较大的特征,让模型仅通过这些特征识别,让总体的识别减去这些敏感特征的识别,效果会不会变好。因为从最近研究的因果角度来说,背景信息某种程度上也会影响主体的识别,所以完全去除背景信息可能也不是最好的选择,而是找到一个平衡点,让模型减去完全使用背景信息盲猜的结果,在因果图上相当于切断了背景信息对主体识别的直接影响,但保留了背景信息通过主体再影响识别的间接影响。 补充:四月刷到了一篇关于解决VIT这个问题的论文,不过没看,不知道他用的什么方法(还是抽空看一下吧);现在是5月26,这篇论文看了,笔记在vit_lazy_aggregation但是还是有点不一样,和我想的有一点不一样,他还是在针对前景和背景这两个具体的块做,本质还是针对不同信息密度的区域做不同处理,有点像是任务导致的区别,分类任务本质就是要找到主体,而类分割任务要的是主要根据自身区域并一定程度上参考其他区域能给予的辅助信息来预测。(讲的好像不是很好,不过我暂时没心思做这么大的任务,说笼统一点就是,不分主体背景,模型就是根据那片区域本身和上下文区域的辅助信息来做预测,无论是分类任务还是类分割任务,这是一个完全符合人类逻辑的因果链条) VLM对一个图像中不同区域的识别问题 论文的方案想用VLM做标注,因为用原有数据集的位姿轨迹自监督效果并不好,主要是想让VLM去评估一个off-road场景下不同区域的Risk,方案开始是超像素分割->合并并区域标号->VLM评估每个区域的Risk。在用的时候发现,区域数目20-30(测的22)的时候,VLM会有自己编造的区域,会自己编造到30个区域。 原因可能感觉有很多: 一个是图像的分辨率低,为了省钱,图像分辨率我用的是960*640,可能VLM觉得这个分辨率太低了,无法识别出足够的细节。 第二个是超像素分割的问题,没有仔细调超像素的参,分割出很多长条状,形状极不规则的区域,增大了VLM的识别难度。 第三个就是老生常谈的VLM根本没办法真正理解图片。 最终没采用VLM标注的方案,因为试了每个区域一个一个送,发现标完所有近两万张图的花费太高了,不知道老师会不会报销。 相对数据会比绝对数据更有用吗 2026-5-7 这个idea已经总结成了一个确切的方案了,对比了一些现有的工作感觉值得做,打算做完现在这个和老师提一下,现在先自己做一些小范围的验证,反正有了ai之后做的快了很多。 突然想起来毕设的任务里有一个小的深度估计器,还有当时对输入输出的分析,当时对比好的结果和坏的结果,发现好的定位结果的文本描述中相对位置的描述特别多,例如“在左边”、“在右边”、“在上面”、“在下面”等等,而坏的定位结果的文本描述中相对位置的描述特别少,更多的是绝对位置的描述,例如“在某个车右边0.5米处”,好像答辩的时候也提过这个问题,觉得相对位置的描述可能更有用。这辈子还有机会做这个问题吗?感觉这个问题挺有意思的,可能也挺有用的,如果能证明相对位置的描述更有用的话,可能会对定位任务优化有用,或许已经有人做了呢?一个本科毕设能分析出啥好东西来。 嘴型特征作为分布中心,猜测文本和真实文本以他为中心 2026-5-15 这只是个猜测,这个点从一些up主一个人说话不出声,另一个人猜什么并和第三个人打电话沟通来的,我感觉,猜的文本很明显不是瞎猜的,是根据嘴型(主要),对方的性格,当时的情景等(不过要是一个人纯瞎说肯定没法猜了)。这个肯定可以用在从嘴型猜文本这个任务上,但是这个任务不知道有啥用(偷拍别人视频然后恢复人家说话来偷听?),我想的是之前做过语音驱动数字人(嘴型),在刚刚的猜测链条上加一个语音并且反过来,文本增强(如果能搞到这个分布的话增强更好)语音驱动数字人。还有一个,直接从嘴型猜文本很难,加上语音(语音本身直接就可以)相当于把之前以嘴型为中心的分布进行约束得到最终的文本,那如果用不同的语音约束呢。 这个想法真怪,但是这个以嘴型特征分布为中心,我确信应该是对的,只是不知道怎么用。 LLM+生成模型来增强数据集来达到提升很多不同任务模型的因果推断能力 2026-5-24 以图像分类任务为例,假设我们现在有一个模型,想提升它的因果推断能力.首先我们需要知道,原本的错误的因果链是什么,例如在图像分类任务中,模型可能过度依赖背景信息来识别主体,那么这个错误的因果链就是背景信息→主体识别。我们想要提升模型的因果推断能力,就需要打破这个错误的因果链,让模型学会真正识别主体的特征,而不是依赖背景信息。我们不通过理论的方法来打破这个因果链,依然通过数据的方式。动机是,对于人类来说,其实我们也有很多错误的因果认知,不过人类可以在经由他人指出错误的因果认知之后,直接修改这个因果链,直接把错误的因果链改成正确的因果链,而不需要通过大量的数据来重新学习这个因果链。 但是对于模型来说,没有直接修改因果链这个说法,除非模型本身显示构建了,所以我们只能通过数据来间接修改这个因果链。 大概的方案可能是,首先对输入数据进行全方位的结构化描述,然后根据预测结果进行分类,分为正确的预测结果和错误的预测结果两类,因为之前的描述很结构化,可以使用自动化工具来构建出模型的认知因果链,例如如果观测到,模型对一个猫的预测正确的描述中,大量的背景描述相似,并且错误预测中背景描述和正确预测的背景描述差异较大,那么我们就可以推断出模型可能过度依赖背景信息来识别主体。 接下来我们就可以针对这个错误的因果链,构建一个新的数据集,这个数据集中的图像会有更多的背景变化,例如同一只猫在不同的背景下出现,这样模型就需要学会真正识别猫的特征,而不是依赖背景信息来识别猫了。通过这样的方式,我们就可以提升模型的因果推断能力,让模型学会真正识别主体的特征,而不是依赖背景信息来识别主体了。 这个方案不会提出任何具体的模型结构和针对任何任务,而是一个通用的方案,可以应用于很多不同的任务,例如图像分类、目标检测、语义分割等等。通过这个方案,我们可以提升很多不同任务模型的因果推断能力。但是有个前提是,我们需要能够生成新的输入数据和确切的标签数据,对于分类任务来说很简单,例如我们可以生成同一只猫在不同背景下的图像,并且标签仍然是猫。对于其他任务来说可能会更复杂一些,例如目标检测,我们需要生成同一只猫在不同背景下的图像,并且标签需要包含猫的位置和大小等信息,这可能需要一些更复杂的生成策略实现。 人类直觉和AI的联系 2026-6-5 很明显,上面idea很多都很直觉,我发现很多现有的领域/任务/方法都和人类有一些联系。这也让我重新审视一些之前习以为常的做法,例如特征金字塔实际上做的是,人看整体、看局部(他是同时看的,那如果是看了整体之后再看局部会更好还是更差) 放大图片但不增加分辨率 在我看不清一个图片的时候,我会放大,但是实际上图片的分辨率并没有提高,这实际上是人类视觉系统与计算机视觉在处理有限信息时的区别。已经有很多工作有这种思想了,主要在于怎么“放大”,“放大”的本质是:用固定的感知资源去“扫描”一个更小的空间范围,从而提高该区域的信息采样密度。 空间注意力机制:模型学出一个“关注度热图”(如Transformer中的注意力权重),然后让后续层只对高权重区域进行精细计算,其他区域粗略处理,模拟人类聚焦的思想。 多尺度推理 + 空间插值: 模拟放大并重新看的思想,将图像(或特征图)的一个局部区域通过双线性插值等操作“物理放大”到原来尺寸,再送入模型进行二次判断。虽然插值不增加信息,但它让后续的卷积核能够在该区域上用更多的参数去扫描,相当于人类眼球在局部区域做更多扫视。 递归/循环放大 (Recursive Zoom):模型反复对自己生成的“放大预测”进行再分析,逐轮修正。例如,先用低分辨率图像得到一个粗糙的语义图,然后将其作为先验,引导对原始低分辨率图像的第二次“注意力扫描”。 特征金字塔 (Feature Pyramid Networks):模型从原始图像中提取多尺度的特征图(例如原图的 1/4, 1/8, 1/16 大小)。高层特征感受野大,负责看整体;低层特征分辨率高,负责看细节。在进行推理时,融合高低层信息,相当于同时对全图(宏观)和局部(微观)进行“放大观察”。

March 27, 2026 · 1 min

NanoGPT

所有之前:这篇是关于NanoGPT项目的核心代码的总结和理解,我会持续把主流技术加入到这个项目中,并对核心代码进行解释。完整项目代码需要查询NanoGPT. 模型架构 Base 最基础的模型架构是一个Decoder-only的Transformer模型,代码实现参考CodeByMyself,在CodeByMyself的基础上,使用了更现代化的Rope位置编码和SwiGLU激活函数, 并且加入了Moe用于和FFN对比,可在配置文件中选择是否使用Moe。部分代码如下: # swiglu class ffn_swiglu(nn.Module): def __init__(self, d_model, d_hidden, dropout=0.1): super().__init__() #第一个 Linear 输出 2 * d_hidden self.w_gate_up = nn.Linear(d_model, 2 * d_hidden) #同时生成 gate 和 up self.w_down = nn.Linear(d_hidden, d_model) self.dropout = nn.Dropout(dropout) def forward(self, x): # x: [B, L, d_model] gate_up = self.w_gate_up(x) # [B, L, 2 * d_hidden] gate, up = gate_up.chunk(2, dim=-1) # each: [B, L, d_hidden] swiglu_out = F.silu(gate) * up # [B, L, d_hidden] out = self.w_down(self.dropout(swiglu_out)) # [B, L, d_model] return out # moe class moe(nn.Module): def __init__(self,n_expert,d_model,top_k=2,dropout=0.1): super().__init__() self.d_model = d_model self.n_expert = n_expert self.top_k = top_k self.dropout = dropout self.gate = nn.Linear(d_model,n_expert,bias=False) self.softmax = nn.Softmax(dim=-1) # self.experts = nn.ModuleList([ffn(self.d_model,self.d_model*4,dropout) for _ in range(n_expert)]) relu激活的expert self.experts = nn.ModuleList([ffn_swiglu(self.d_model,self.d_model*2,dropout) for _ in range(n_expert)]) def forward(self,x): b,t,d = x.shape assert d == self.d_model,f"输入维度和moe设置维度不匹配" x_flat = x.view(-1,d) N = x_flat.shape[0] gate_logits = self.gate(x_flat) topk_weights,topk_indices = torch.topk(gate_logits,self.top_k,dim=-1) topk_weights= self.softmax(topk_weights) out = torch.zeros_like(x_flat) for i,expert in enumerate(self.experts): mask = (topk_indices == i) if not mask.any(): continue token_indices,expert_pos = torch.where(mask) select_x = x_flat[token_indices] expert_out = expert(select_x) weights = topk_weights[token_indices, expert_pos] out.index_add_(0, token_indices, expert_out * weights.unsqueeze(1)) return out.view(b,t,d) def precompute_freqs_cis(dim: int, end: int, theta: float = 10000.0): """ 预计算旋转角度的复数表示(cos + i*sin) """ freqs = 1.0 / (theta ** (torch.arange(0, dim, 2)[: (dim // 2)].float() / dim)) t = torch.arange(end, device=freqs.device) freqs = torch.outer(t, freqs).float() # [end, dim//2] freqs_cos = freqs.cos() freqs_sin = freqs.sin() return freqs_cos, freqs_sin # 分开存储更便于后续操作,不使用torch.complex类型,因为部分系统不支持 SwiGLU激活函数的实现是通过将线性层的输出分成两部分,一部分作为gate,另一部分作为up,然后使用SILU函数对gate进行激活,并与up相乘得到最终的输出。 ...

November 24, 2025 · 2 min

CodeByMySelf-Diffusion

所有之前:CodeByMySelf 系列之 Diffusion Models 的实现记录,变体不会在这篇里更新。 Diffusion Models 原理 具体的一些公式推导请看李宏毅老师的视频: 李宏毅Diffusion 基本概念 (以图像生成为例)首先Diffusion的基本运作过程是,首先Sample出一个全是噪声的图片,然后经过一步一步的Denoise最终得到图片,Denoise的步数是事先确定好的。 在Diffusion中,深度学习模型的唯一作用就是预测噪声,其他过程均由数学公式推理得到 进一步说,不同情况的噪声需要的处理肯定是不一样的,一个Denoise模型肯定无法做到处理不同的情况的噪声,因此Denoise模型需要一个“噪声严重程度”作为输入,在Diffusion中称为时间步。 Denoise内部的工作是,接受带Noise的图片,对噪声进行预测,然后使用带Noise的图片减去预测噪声那么就可以得到。 我们清楚Noise predictor模型的输入输出之后,该怎么训练这个模型呢?已知模型的输入-输出对是Noise图片加上“噪声程度”-Noise数据对,我们需要人为创造这个数据对出来,具体来说,从image list中拿一张图片,然后从GasuionN中采样出一张纯噪声的图片,加了一定的步数之后,就得到了很多数据对(Noise图片、时间步–噪声),这个过程称为Diffusion Process。 现在,已经可以从一个纯噪声得到一张图片,但是我们不能根据噪声随机生成图片,所以我们需要把文字考虑进来,让模型考虑进文字之后再预测噪声。 完整算法 首先是Training的过程,从immage list中取出一张干净的图,之后从1-T中随机取一个值作为时间步,然后从Gaussion中取样出一个noise,之后就是一步一步的梯度更新训练模型。这种方法和我们的想法有些不一样,我们想像中noise是一步一步加进去的,但是在实际中是一下子直接加进去的。 具体的加噪公式推导过程如下: 然后是推理过程,也就是Sampling这个过程。首先,采样得到一个纯噪声图片X_T,然后进行T次公式中的循环,$\epsilon_\theta$表示预测噪声的模型。 但是为什么要加这个z噪声:每次都要重新采样出一个噪声,但是只有当t>1的时候.这个z的使用和生成语句中一个问题类似:为什么每次总要SAMPLE而不是直接取Mean(概率最大)?在The curious case of nerual text degeneration论文中有一个分析,当直接取Mean时候,会一直重复相同的回答和出现跳帧现象,并且他们分析得到人在写文章的时候并不会一定会选概率最大的词。 简单代码实现 Diffusion已经有了许多变体,下面是一个简单的实现,变体不会在这篇里更新。 import torch from torch import nn import math import torch.nn.functional as F import numpy as np import time class WeightedLoss(nn.Module): def __init__(self): super(WeightedLoss,self).__init__() def forward(self,pred,target,weighted=1.0): loss = self._loss(pred,target) weighted_loss = (loss * weighted).mean() # mean by batch return weighted_loss class L1Loss(WeightedLoss): def _loss(self,pred,target): return torch.abs(pred - target) class L2Loss(WeightedLoss): def _loss(self,pred,target): return (pred - target) ** 2 Losses = { 'l1':L1Loss, 'l2':L2Loss } class SinusoidalPosEmb(nn.Module): def __init__(self,dim): super(SinusoidalPosEmb,self).__init__() self.dim = dim def forward(self,t): ''' 扩散模型时间步的位置编码公式和Transformer中的位置编码公式是类似的。 ''' device = t.device half_dim = self.dim // 2 emb = math.log(10000) / (half_dim - 1) emb = torch.exp(torch.arange(half_dim,device=device) * -emb) emb = t[:,None] * emb[None,:] emb = torch.cat((emb.sin(),emb.cos()),dim=-1) return emb def extract(a,t,x_shape): ''' a: (T,) t: (B,) x_shape: (B,...) ''' batch_size = t.shape[0] out = a.gather(-1,t) # (B,) return out.view(batch_size, *((1,) * (len(x_shape) - 1))) # (B,1,1,1...) class MLP(nn.Module): ''' 这是一个简单的多层感知机(MLP)模型,结合了时间步的位置编码,用于处理状态和动作的输入(强化学习) ''' def __init__(self,state_dim,action_dim,hidden_dim,device,t_dim): super(MLP,self).__init__() self.device = device self.t_dim = t_dim self.a_tim = action_dim self.time_mlp = nn.Sequential( SinusoidalPosEmb(t_dim), nn.Linear(t_dim,t_dim*2), nn.Mish(), nn.Linear(t_dim*2,t_dim) ) input_dim = state_dim + action_dim + t_dim self.mid_layer = nn.Sequential( nn.Linear(input_dim,hidden_dim), nn.Mish(), nn.Linear(hidden_dim,hidden_dim), nn.Mish(), nn.Linear(hidden_dim,hidden_dim), nn.Mish(), ) self.f_layer = nn.Linear(hidden_dim,action_dim) def init_weights(self): ''' 初始化模型权重,好的权重初始化可以帮助模型更快收敛,提升训练效果。 ''' for m in self.modules(): if isinstance(m,nn.Linear): nn.init.xavier_normal_(m.weight) if m.bias is not None: nn.init.zeros_(m.bias) def forward(self,state,action,t): t_emb = self.time_mlp(t) # (batch,t_dim) x = torch.cat((state,action,t_emb),dim=-1) # (batch,state_dim+action_dim+t_dim) x = self.mid_layer(x) # (batch,hidden_dim) action_delta = self.f_layer(x) # (batch,action_dim) return action_delta class Diffusion(nn.Module): def __init__(self,loss_type,beta_schedule='linear',clip_denoise=True,predict_epsilon=True,**kwargs): ''' loss_type: 损失函数类型 beta_schedule: beta(时间)的离散方式,常见的有'linear'(线性)和'cosine'(余弦)。 clip_denoise: 是否在去噪过程中裁剪输出,以防止数值过大。 ''' super(Diffusion,self).__init__() self.device = kwargs.get('device','cpu') self.t_dim = kwargs.get('t_dim',16) self.state_dim = kwargs.get('state_dim',10) self.action_dim = kwargs.get('action_dim',2) self.hidden_dim = kwargs.get('hidden_dim',256) self.loss_type = loss_type self.beta_schedule = beta_schedule self.clip_denoise = clip_denoise self.T = kwargs.get('timesteps',1000) self.model = MLP(self.state_dim,self.action_dim,self.hidden_dim,self.device,self.t_dim).to(self.device) self.kwargs = kwargs self.model.init_weights() if beta_schedule == 'linear': betas = torch.linspace(1e-4,0.02,self.kwargs.get('timesteps', 1000),dtype=torch.float32) alphas = 1. - betas alphas_cumprod = torch.cumprod(alphas,dim=0) # [1,2,3] -> [1,1*2,1*2*3] alphas_cumprod_prev = torch.cat((torch.tensor([1.],dtype=torch.float32),alphas_cumprod[:-1]),dim=0) self.register_buffer('betas',betas) self.register_buffer('alphas',alphas) self.register_buffer('alphas_cumprod',alphas_cumprod) self.register_buffer('alphas_cumprod_prev',alphas_cumprod_prev) self.register_buffer('sqrt_alphas_cumprod',torch.sqrt(alphas_cumprod)) # (前向过程) self.register_buffer('sqrt_alphas_cumprob',torch.sqrt(alphas_cumprod)) self.register_buffer('sqrt_one_minus_alphas_cumprod',torch.sqrt(1. - alphas_cumprod)) # (反向过程) posterior_variance = betas * (1. - alphas_cumprod_prev) / (1. - alphas_cumprod) self.register_buffer('posterior_variance',posterior_variance) # 用于从xt求得x0 self.register_buffer('sqrt_recip_alphas_cumprod',torch.sqrt(1. / alphas_cumprod)) self.register_buffer('sqrt_recipm1_alphas_cumprod',torch.sqrt(1. / alphas_cumprod - 1)) self.register_buffer('posterior_mean_coef1',betas *torch.sqrt(alphas_cumprod_prev) / (1. - alphas_cumprod)) self.register_buffer('posterior_mean_coef2', (1. - alphas_cumprod_prev) * torch.sqrt(alphas) / (1. - alphas_cumprod)) self.loss_func = Losses[loss_type]() def forward(self,state,*args,**kwargs): return self.sample(state,*args,**kwargs) def sample(self,state,*args,**kwargs): batch_size = state.shape[0] shape = [batch_size,self.action_dim] action = self.p_sample_loop(state,shape,*args,**kwargs) return action.clamp(-1.,1.) def q_posterior(self,x0,xt,t): ''' 计算后验分布q(x_{t-1}|x_t,x_0)的均值和方差 ''' posterior_mean = (extract(self.posterior_mean_coef1,t,x0.shape) * x0 + extract(self.posterior_mean_coef2,t,x0.shape) * xt) posterior_variance = extract(self.posterior_variance,t,x0.shape) posterior_log_variance = torch.log(posterior_variance.clamp(min=1e-20)) return posterior_mean,posterior_variance,posterior_log_variance def p_sample_loop(self,state,shape,*args,**kwargs): device = self.kwargs.get('device','cpu') batch_size = shape[0] x = torch.randn(shape,device=device,requires_grad=False) # 标准DDPM噪声不需要梯度 for i in reversed(range(0,self.kwargs.get('timesteps',1000))): t = torch.full((batch_size,),i,dtype=torch.long,device=device) x = self.p_sample(x,t,state) return x def predict_x0_from_noise(self,x,t,noise): ''' 根据公式计算x0 ''' return (extract(self.sqrt_recip_alphas_cumprod,t,x.shape) * x - extract(self.sqrt_recipm1_alphas_cumprod,t,x.shape) * noise) def p_mean_variance(self,x,t,state): pred_noise = self.model(state,x,t) x_0 = self.predict_x0_from_noise(x,t,pred_noise) x_0.clamp_(-1.,1.) if self.clip_denoise else x_0 model_mean,posterior_variance,posterior_log_variance = self.q_posterior(x_0,x,t) return model_mean,posterior_log_variance def p_sample(self,x,t,state): ''' x: 当前的噪声状态xt t: 当前的时间步t state: 状态 ''' model_mean,model_log_variance = self.p_mean_variance(x,t,state) noise = torch.randn_like(x) if t.sum() > 0 else 0. # 如果t=0,则不添加噪声 return model_mean + torch.exp(0.5 * model_log_variance) * noise def q_sample(self,x0,t,noise): sample = (extract(self.sqrt_alphas_cumprod,t,x0.shape) * x0 + extract(self.sqrt_one_minus_alphas_cumprod,t,x0.shape) * noise) return sample def p_losses(self,x0,state,t,weights): noise = torch.randn_like(x0) x_noisy = self.q_sample(x0,t,noise) x_recon = self.model(state,x_noisy,t) loss = self.loss_func(x_recon,noise,weights) return loss def loss(self,x,state,weights=1.0): batch_size = x.shape[0] t = torch.randint(0,self.T,(batch_size,),device=self.device).long() return self.p_losses(x,state,t,weights) if __name__ == '__main__': x = torch.randn(256,2) state = torch.randn(256,11) model = Diffusion(loss_type='l2',beta_schedule='linear',clip_denoise=True,predict_epsilon=True, device='cpu',t_dim=16,state_dim=11,action_dim=2,hidden_dim=256,timesteps=100) result = model(state) print(result) loss = model.loss(x,state) print(f"loss: {loss.item()}")

September 24, 2025 · 3 min