前言

本篇文章记录 CS336 作业 Assignment 1: Basics 中的 Transformer Language Model Architecture 作业要求,仅供自己参考😄

Assignment 1https://github.com/stanford-cs336/assignment1-basics

referencehttps://chatgpt.com/

1. Transformer Language Model Architecture 作业要求

以下内容均翻译自 cs336_spring2025_assignment1_basics.pdf,请大家查看原文档获取更详细的内容

语言模型以 一批(batched)整数形式的 token ID 序列 作为输入(即形状为 (batch_size, sequence_length) 的 PyTorch 张量),并输出在词表上的 (批量化的)归一化概率分布(即形状为 (batch_size, sequence_length, vocab_size) 的 PyTorch 张量),该输出表示:对于序列中每一个输入 token,模型都会给出其 下一个 token 的预测分布

在训练语言模型时,我们使用 下一个词预测(next-word prediction)的目标,计算真实下一个 token 与模型预测分布之间的交叉熵损失(cross-entropy loss)。在推理阶段生成文本时,我们从 最后一个时间步(即序列的末尾)得到的下一个词分布中生成下一个 token(例如选择概率最大的 token,或从分布中进行采样),将该 token 追加到输入序列中,然后重复这一过程

在本次作业的这一部分中,你将 从零开始构建一个 Transformer 语言模型,我们会先给出模型的高层整体描述,然后逐步深入介绍各个具体组成模块

1.1 Transformer LM

给定一串 token ID 序列,Transformer 语言模型首先通过 输入嵌入(input embedding)将 token ID 转换为稠密向量表示;随后,将这些嵌入后的向量依次送入 num_layers 个 Transformer 模块;最后,使用一个可学习的线性投影层(也称为 “输出嵌入” 或 LM head)来生成对下一个 token 的预测 logits,Figure 1 给出了该过程的示意图

1.1.1 Token Embeddings

在模型的第一步,Transformer 会将 (批量化的)token ID 序列嵌入为一系列向量,这些向量中包含了 token 身份信息(在 Figure 1 中以红色模块表示)

更具体地说,给定一串 token ID,Transformer 语言模型使用一个 token embedding 层 来生成对应的向量序列。该嵌入层接收一个形状为 (batch_size, sequence_length) 的整数张量作为输入,并输出一个形状为 (batch_size, sequence_length, d_model) 的向量序列

1.1.2 Pre-norm Transformer Block

在完成 embedding 之后,激活值会依次经过多个 结构完全相同的神经网络层 进行处理,一个标准的 仅解码器(decode-only)Transformer 语言模型由 num_layers 个相同的层 组成,这些层通常被称为 Transformer 模块(Transformer blocks)

每个 Transformer 模块接收一个形状为 (batch_size, sequence_length, d_model) 的输入,并输出一个同样形状的张量。每个模块都会在整个序列范围内 聚合信息(通过自注意力机制),并通过 前馈网络(feed-forward layers) 对其进行非线性变换

1.2 Output Normalization and Embedding

在经过 num_layers 个 Transformer 模块之后,我们将取最终得到的激活值,并将其转换为 在整个词表上的概率分布

我们将实现一种 前归一化(pre-norm)Transformer 模块(详见 §1.5),该结构还要求在 最后一个 Transformer 模块之后 额外应用一次 层归一化(Layer Normalization),以确保输出的尺度是合适的

完成这一步归一化后,我们将使用一个标准的、可学习的 线性变换层,把 Transformer 模块的输出转换为对 下一个 token 的预测 logits(可参考 [Radford+ 2018] 的公式 2)

1.3 Remark: Batching, Einsum and Efficient Computation

在整个 Transformer 模型中,我们会反复对许多 “类似批次” 的输入执行相同的计算,下面是几个常见的例子:

  • 批次维度(elements of a batch):对批次中的每一个样本,应用完全相同的 Transformer 前向计算
  • 序列长度维度(sequence length):像 RMSNorm前馈网络(feed-forward) 这样的 “按位置(position-wise)” 操作,会在序列中的每一个位置上以完全相同的方式执行
  • 注意力头维度(attention heads):在 多头注意力(multi-head attention) 中,注意力计算会在多个注意力头之间进行批处理

为了充分利用 GPU 的计算能力,并让代码保持 简洁、易读、易理解,我们需要一种高效而直观的方式来表达这些批量化操作。PyTorch 中的许多算子都支持在张量前部携带额外的 “类批次” 维度,并在这些维度上高效地重复或广播操作

例如,假设我们正在执行一个 按位置的批处理操作,我们有一个形状为 (batch_size, sequence_length, d_model) 的数据张量 D,希望将其与一个形状为 (d_model, d_model) 的矩阵 A 相乘,在这种情况下,D @ A 会执行一次 批量矩阵乘法,其中 (batch_size, sequence_length) 这两个维度会被自动视为批次维度并进行并行处理,这正是 PyTorch 中高效的基本操作之一

基于上述原因,我们在实现函数时,应当假设输入张量可能会带有 额外的类批次维度,并尽量将这些维度放在 PyTorch 张量的 最前面。然后,为了让张量适配这种批处理形式,往往需要频繁地使用 viewreshapetranspose 等操作对维度进行重排,这种做法不仅繁琐,而且代码很快就就会变得难以理解,难以判断张量当前的真实形状与语义

一种更为优雅的做法是使用 einsum 记号法,即通过 torch.einsum,或者使用与框架无关的库,如 einopseinx。其中,两类核心操作是:

  • einsum:用于在任意维度上执行张量收缩(tensor contraction)
  • rearrange:用于对张量维度进行重排、拼接和拆分

事实上,机器学习中的绝大多数运算,本质上都可以看作是 维度重排与张量收缩 的组合,外加少量(通常是逐元素)非线性函数,这意味着,使用 einsum 记号法可以让你的代码在保持灵活性的同时,也更加简洁和易读

我们 强烈建议 在本课程中学习并使用 einsum 记号法,如果你之前没有接触过 einsum,建议从 einops 入手([docs]);如果你已经熟悉 einops,则可以进一步学习功能更通用的 einx[docs]),这两个库在课程提供的环境中都已经预先安装

在后续内容中,我们会给出一些如何使用 einsum 记号法的示例,这些示例将作为 einops 官方文档的补充,而你也应当首先阅读 einops 的文档来建立基础认知

Note:需要注意的是,尽管 einops 拥有广泛的社区支持和成熟度,einx 目前仍未经过充分的生产级测试。如果你在使用 einx 时遇到限制或 bug,完全可以回退到 einops + 较为直接的 PyTorch 操作 这一更稳妥的组合


Example (einstein_example1): Batched matrix multiplication with einops.einsum

import torch
from einops import rearrange, einsum

基础实现

Y = D @ A.T
  • 这种写法 很难直观看出 输入和输出的张量形状,以及这些维度各自代表什么含义
  • 例如:DA 可以有哪些形状?是否会出现一些 不符合直觉的广播或行为?

使用 einsum:自文档化且更健壮

Y = einsum(D, A, "batch sequence d_in, d_out d_in -> batch sequence d_out")
  • 这里通过显示标注维度名:
    • D 的形状是 (batch, sequence, d_in)
    • A 的形状是 (d_out, d_in)
    • 输出 Y 的形状是 (batch, sequence, d_out)
  • 维度语义一目了然,不容易写错,也更容易读懂

更通用的批处理版本

Y = einsum(D, A, "... d_in, d_out d_in -> ... d_out")
  • 在这个版本中:
    • D 可以有 任意数量的前置批次维度(用 ... 表示)
    • A 的形状仍然受限为 (d_out, d_in)
  • 这使得代码在不同场景下 更具复用性


Example (einstein_example2): Broadcasted operations with einops.rearrange

假设我们有一批图像,并希望为 每一张图像生成 10 个不同亮度(dimmed)版本,亮度由一个缩放因子控制:

images = torch.randn(64, 128, 128, 3) # (batch, height, width, channel)
dim_by = torch.linspace(start=0.0, end=1.0, steps=10)

  • images:一批图像,形状为 (batch, height, width, channel)
  • dim_by:长度为 10 的缩放系数,用于控制变暗程度

重排维度并相乘

dim_value = rearrange(dim_by, "dim_value -> 1 dim_value 1 1 1")
images_rearr = rearrange(images, "b height width channel -> b 1 height width channel")
dimmed_images = images_rearr * dim_value
  • dim_by 被重排为 (1, dim_value, 1, 1, 1),以便进行广播
  • images 被重排为 (batch, 1, height, width, channels)
  • 相乘后得到的 dimmed_images 形状为 batch, dim_value, height, width, channel

一步完成(使用 einsum)

dimmed_images = einsum(
images, dim_by,
"batch height width channel, dim_value -> batch dim_value height width channel"
)
  • 这一写法 同时完成了维度对齐与广播
  • 不需要显式 rearrange,语义依然非常清晰
  • 非常适合复杂的多维广播场景


Example (einstein_example3): Pixel mixing with einops.rearrange

假设我们有一批图像,表示为形状为 (batch, height, width, channel) 的张量,我们希望对 图像中的所有像素 施加一个 线性变换,但该变换需要对 每个通道(channel)独立进行,这个线性变换用一个矩阵 B 表示,其形状为 (height x width, height x width)

channels_last = torch.randn(64, 32, 32, 3)  # (batch, height, width, channel)
B = torch.randn(32*32, 32*32)

将图像张量重排,以便在所有像素上进行混合

下面展示的是一种 传统的 PyTorch 写法,需要反复使用 viewtranspose 来对维度进行重排:

channels_last_flat = channels_last.view(
    -1, channels_last.size(1) * channels_last.size(2), channels_last.size(3)
)

channels_first_flat = channels_last_flat.transpose(1, 2)

channels_first_flat_transformed = channels_first_flat @ B.T

channels_last_flat_transformed = channels_first_flat_transformed.transpose(1, 2)

channels_last_transformed = channels_last_flat_transformed.view(*channels_last.shape)

这种写法的问题在于:

  • 需要在代码前后额外添加注释,才能弄清楚输入与输出的形状
  • 写法繁琐、可读性差
  • 容易出错(bug-prone)

使用 einops: rearrange 替代繁琐的 view + transpose

height = width = 32

channels_first = rearrange(
    channels_last,
    "batch height width channel -> batch channel (height width)"
)

这一步明确表达了语义:将 (height, width) 合并为一个像素维度,并把 channel 提前

使用 einsum 执行线性变换

channels_first_transformed = einsum(
    channels_first, B,
    "batch channel pixel_in, pixel_out pixel_in -> batch channel pixel_out"
)

这里清晰地表达了:

  • pixel_inpixel_out 之间的线性映射
  • 批次维度与通道维度均被自动批处理

将结果重排回原始图像格式

channels_last_transformed = rearrange(
    channels_first_transformed,
    "batch channel (height width) -> batch height width channel",
    height=height, width=width
)

最终结果的形状恢复为 (batch, height, width, channel)

更激进的写法:一步完成(使用 einx.dot)

如果你愿意 “一步到位”,可以使用 einx.doteinops.einsum 的等价扩展):

height = width = 32

channels_last_transformed = einx.dot(
    "batch row_in col_in channel, (row_out col_out) (row_in col_in)"
    "-> batch row_out col_out channel",
    channels_last, B,
    col_in=width, col_out=width
)

这种写法将维度重排、张量收缩以及输出格式全部统一在一条表达式中完成,语义非常明确,但也更硬核


einsum 记号法 不仅能够处理任意数量的输入批处理维度,还具有一个关键优势:自文档化(self-documenting),在使用 einsum 记号法的代码中,输入张量和输出张量的相关形状一目了然,代码本身就清楚地说明了张量之间的关系

对于其余的张量,你可以考虑使用 Tensor 类型提示(type hints) 来进一步增强代码的可读性,例如使用 jaxtyping 库(该库并不局限于 JAX)

我们将在 作业 2 中更深入地讨论使用 einsum 记号法对性能的影响,但在当前阶段,你只需要知道一点:在绝大多数情况下,einsum 的表现几乎总是优于替代方案

1.3.3.1 Mathematical Notation and Memory Ordering

许多机器学习论文在数学记号中使用 行向量(row vectors),这种表示方式与 NumPyPyTorch 默认采用的 行优先(row-major)内存排列 非常契合,在使用行向量的情况下,一个线性变换可以写成:

y = x W ⊤ , (1) y = x W^{\top}, \tag{1} y=xW,(1)

其中,矩阵 W ∈ R d out × d in W \in \mathbb{R}^{d_{\text{out}} \times d_{\text{in}}} WRdout×din 采用行优先存储,而向量 x ∈ R 1 × d in x \in \mathbb{R}^{1 \times d_{\text{in}}} xR1×din 是一个行向量。

线性代数 中,更常见的做法是使用 列向量(column vectors),此时线性变换通常写为:

y = W x , (2) y = W x, \tag{2} y=Wx,(2)

其中 W ∈ R d out × d in W \in \mathbb{R}^{d_{\text{out}} \times d_{\text{in}}} WRdout×din,而 x ∈ R d in x \in \mathbb{R}^{d_{\text{in}}} xRdin 是列向量。

在本次作业中,我们将 在数学记号上统一使用列向量表示,因为这种方式在推导和理解数学公式时通常更直观、更容易跟随

不过你需要注意:如果你在代码中使用 普通的矩阵乘法记号(而不是 einsum),那么由于 PyTorch 采用行优先的内存布局,你实际上需要按照 行向量约定 来应用矩阵乘法,如果你是使用 einsum 来实现矩阵运算,那么这一差异基本不会成为问题

1.4 Basic Building Blocks: Linear and Embedding Modules

1.4.1 Parameter Initialization

要有效地训练神经网络,通常需要 精心设计模型参数的初始化方式—不良的初始化可能会导致诸如 梯度消失梯度爆炸 等不理想的训练行为,虽然 Pre-norm transformer 在初始化方面通常具有较强的鲁棒性,但初始化方式仍然会对 训练速度和收敛性 产生显著影响

鉴于本次作业内容已经较多,我们将把更深入的初始化细节留到 作业 3 中讨论,在这里,我们先给出一些在大多数情况下都表现良好的 近似初始化方案,目前请使用以下初始化规则:

  • 线性层权重(Linear weights) N ( μ = 0 , σ 2 = 2 d i n + d o u t ) \mathcal{N}\left(\mu=0,\sigma^{2}=\tfrac{2}{d_{\mathrm{in}}+d_{\mathrm{out}}}\right) N(μ=0,σ2=din+dout2),服从均值为 0、方差为 2 d i n + d o u t \tfrac{2}{d_{\mathrm{in}}+d_{\mathrm{out}}} din+dout2 的正态分布,并在区间 [ − 3 σ , 3 σ ] [-3\sigma, 3\sigma] [3σ,3σ] 上进行截断
  • 嵌入层(Embedding) N ( μ = 0 , σ 2 = 1 ) \mathcal{N}\left(\mu=0,\sigma^{2}=1\right) N(μ=0,σ2=1),服从均值为 0、方差为 1 的正态分布,并在区间 [ − 3 , 3 ] [-3, 3] [3,3] 上进行截断
  • RMSNorm:初始化为 1
1.4.2 Linear Module

线性层是 Transformer 以及神经网络中最基本、最核心的构建模块之一,首先,你需要实现一个自己的 Linear 类,该类继承自 torch.nn.Module,并执行如下线性变换:

y = W x . (3) y = W x. \tag{3} y=Wx.(3)

需要注意的是,按照大多数现代大语言模型(LLM)的做法,我们在这里不包含偏置项

Problem (linear): Implementing the linear module (1 point)

Deliverable:请实现一个 Linear 类,该类继承自 torch.nn.Module,并执行线性变换,你的实现应当遵循 PyTorch 内置 nn.Linear 模块的接口设计,但不包含偏置(bias)参数或偏置项,我们推荐使用如下接口:

def __init__(self, in_features, out_features, device=None, dtype=None)

用于构造一个线性变换模块,该函数应当接收以下参数:

  • in_features: int:输入的最终维度
  • out_features: int:输出的最终维度
  • device: torch.device | None = None:用于存放参数的设备
  • dtype: torch.dtype | None = None:参数的数据类型
def forward(self, x: torch.Tensor) -> torch.Tensor

将线性变换应用到输入张量上

实现时请务必注意以下几点

  • 继承自 nn.Module
  • 调用父类构造函数(super().__init__()
  • 构造并存储参数矩阵为 W W W(而不是 W ⊤ W^{\top} W,这是出于内存排列顺序的考虑,该参数应存放在一个 nn.Parameter
  • 不要 使用 nn.Linearnn.functional.linear

关于参数初始化,请使用上文给出的初始化设置,并结合 torch.nn.init.trunc_normal_ 来初始化权重参数

为了测试你的 Linear 模块,请先实现测试适配器 [adapters.run_linear],该适配器应当将给定的权重加载到你的 Linear 模块中,你可以使用 Module.load_state_dict 来完成这一操作

随后,运行以下命令进行测试:

uv run pytest -k test_linear
1.4.3 Embedding Module

如前所述,Transformer 的第一层是一个 嵌入层(embedding layer),它将整数形式的 token ID 映射到维度为 d_model 的向量空间中

在本次作业中,我们将实现一个 自定义的 Embedding ,该类继承自 torch.nn.Module(因此你 不应使用 nn.Embedding)。在前向传播中,forward 方法应当通过 索引操作,从一个形状为 (vocab_size, d_model) 的嵌入矩阵中,为每一个 token ID 选取对应的嵌入向量

具体来说,输入是一个形状为 (batch_size, sequence_length)torch.LongTensor 类型的 token ID 张量,输出则是对应的嵌入向量序列

Problem (embedding): Implement the embedding module (1 point)

Deliverable:请实现一个 Embedding 类,该类继承自 torch.nn.Module,并执行嵌入查找(embedding lookup),你的实现应当遵循 PyTorch 内置 nn.Embedding 模块的接口设计,我们推荐使用如下接口:

def __init__(self, num_embeddings, embedding_dim, device=None, dtype=None)

用于构造一个嵌入模块,该函数应当接收以下参数:

  • num_embeddings:词表大小(vocabulary size)
  • embedding_dim : int:嵌入向量的维度,即 d_model
  • device: torch.device | None = None:用于存放参数的设备
  • dtype: torch.dtype | None = None:参数的数据类型
def forward(self, token_ids: torch.Tensor) -> torch.Tensor

根据给定的 token ID,查找并返回对应的嵌入向量

实现时请务必注意以下几点

  • 继承自 nn.Module
  • 调用父类构造函数(super().__init__()
  • 将嵌入矩阵初始化并存储为一个 nn.Parameter
  • 嵌入矩阵的最后一个维度必须是 d_model
  • 不要 使用 nn.Embeddingnn.functional.embedding

关于参数初始化,同样请使用前文给出的初始化设置,并使用 torch.nn.init.trunc_normal_ 来初始化嵌入权重

为了测试你的实现,请先实现测试适配器 [adapters.run_embedding],随后,运行以下命令进行测试:

uv run pytest -k test_embedding

1.5 Pre-Norm Transformer Block

每一个 Transformer 模块包含两个子层:多头自注意力机制(multi-head self-attention)逐位置前馈网络(position-wise feed-forward network) [Vaswani+ 2017],具体细节请参照 1.1 小节

在最初的 Transformer 论文中,模型在这两个子层的输出外各自加一个 残差连接(residual connection),随后再进行 层归一化(layer normalization),这种结构通常被称为 后归一化(post-norm)Transformer,因为层归一化作用在子层的输出上

然而,后续的多项研究发现,将层归一化从 子层输出处 移动到 子层输入处,并在最后一个 Transformer 模块之后再额外加一次层归一化,可以显著提升 Transformer 在训练过程中的稳定性 [Nguyen+ 2019] [Xiong+ 2020],这种结构被称为 前归一化(pre-norm)Transformer,其示意图可参考 Figure 2

在 pre-norm 结构中,每个 Transformer 子层的输出仍然通过残差连接加回到子层的输入上 [Vaswani+ 2017],对 pre-norm 的一种直观理解是:从输入嵌入一直到 Transformer 的最终输出,存在一条 不经过任何归一化操作的 “干净残差通路(residual stream)”,这被认为有助于改善梯度传播

目前,pre-norm Transformer 已经成为语言模型中的标准结构(例如 GPT-3、LLaMA、PaLM 等),因此,在本次作业中我们也将实现这一变体,接下来,我们将依次讲解并实现 pre-norm Transformer 模块的各个组成部分

1.5.1 Root Mean Square Layer Normalization

在最初的 Transformer 实现中,[Vaswani+ 2017] 使用 层归一化(layer normalization)[Ba+ 2016] 来对激活值进行归一化。在本次作业中,我们将遵循 [Touvron+ 2023] 的做法,采用均方根层归一化(RMSNorm) 来替代标准的 LayerNorm,具体可以参考 [Zhang+ 2019] 中的公式 4

给定一个维度为 d model d_{\text{model}} dmodel 的激活向量 a ∈ R d model a \in \mathbb{R}^{d_{\text{model}}} aRdmodel,RMSNorm 对每个激活分量的缩放方式如下:

RMSNorm ( a i ) = a i RMS ( a ) g i , (4) \text{RMSNorm}(a_i) = \frac{a_i}{\text{RMS}(a)} g_i, \tag{4} RMSNorm(ai)=RMS(a)aigi,(4)

其中:

RMS ( a ) = 1 d model ∑ i = 1 d model a i 2 + ε . \text{RMS}(a) = \sqrt{\frac{1}{d_{\text{model}}} \sum_{i=1}^{d_{\text{model}}} a_i^2 + \varepsilon}. RMS(a)=dmodel1i=1dmodelai2+ε .

这里, g i g_i gi 是一个 可学习的 “增益(gain)” 参数,一共存在 d model d_{\text{model}} dmodel 个这样的参数,而 ε \varepsilon ε 是一个超参数,通常规定为 10 − 5 10^{-5} 105,用于数值稳定性

在对输入进行平方运算时,你应当 先将输入转换为 torch.float32,以防止数值溢出,总体来说,你的 forward 方法应当大致如下所示:

in_dtype = x.dtype
x = x.to(torch.float32)

# Your code here performing RMSNorm
...

result = ...

# Return the result in the original dtype
return result.to(in_dtype)
Problem (rmsnorm): Root Mean Square Layer Normalization (1 point)

Deliverable:请将 RMSNorm 实现为一个 torch.nn.Module,我们推荐使用如下接口:

def __init__(self, d_model: int, eps: float = 1e-5, device=None, dtype=None)

用于构造 RMSNorm 模块,该函数应当接收以下参数:

  • d_model: int:模型的隐藏维度
  • eps: float = 1e-5:用于数值稳定性的 ε \varepsilon ε 参数
  • device: torch.device | None = None:用于存放参数的设备
  • dtype: torch.dtype | None = None:参数的数据类型
def forward(self, x: torch.Tensor) -> torch.Tensor

对形状为 (batch_size, sequence_length, d_model) 的输入张量进行处理,并返回 形状相同 的输出张量

Note:请记得在执行归一化之前,先将输入提升为 torch.float32,并在计算完成后 再转换回原始数据类型,如前文所述。

为了测试你的实现,请先实现测试适配器 [adapters.run_rmsnorm],随后运行以下命令进行测试:

uv run pytest -k test_rmsnorm
1.5.2 Position-Wise Feed-Forward Network

在最初的 Transformer 论文中 [Vaswani+ 2017],Transformer 的前馈网络由 两层线性变换 组成,中间使用 ReLU 激活函数(即 ReLU ( x ) = max ⁡ ( 0 , x ) \text{ReLU}(x) = \max(0, x) ReLU(x)=max(0,x)),该前馈网络中间层的维度通常是输入维度的 4 倍

然而,与这一最初设计相比,现代语言模型通常引入了两个主要改动:使用新的激活函数和引入门控(gating)机制

具体来说,在本次作业中,我们将实现一种称为 SwiGLU 的激活函数形式,该形式被诸如 LLaMA3 [Grattafiori+ 2024] 和 Qwen 2.5 [Yang+ 2024] 等模型所采用,SwiGLU 将 SiLU(通常也称为 Swish)激活函数 与一种称为 门控线性单元(Gated Linear Unit, GLU) 的机制相结合

此外,按照大多数现代语言模型如 PaLM [Chowdhery+ 2022] 和 LLaMA [Touvron+ 2023] 的做法,我们还将 省略线性层中的偏置项(bias)

SiLU(或 Swish)激活函数 [Hendrycks+ 2016]
[Elfwing+ 2017] 定义如下:

SiLU ( x ) = x ⋅ σ ( x ) = x 1 + e − x . (5) \text{SiLU}(x) = x \cdot \sigma(x) = \frac{x}{1 + e^{-x}}. \tag{5} SiLU(x)=xσ(x)=1+exx.(5)

如 Figure 3 所示,SiLU 激活函数 在整体形态上与 ReLU 激活函数 相似,但在 零点附近是平滑的

门控线性单元(Gated Linear Units, GLU) 最早由 [Dauphin+ 2017] 提出,其定义为:将一个线性变换经过 sigmoid 函数后的结果,与另一个线性变换的结果进行逐元素相乘

GLU ( x , W 1 , W 2 ) = σ ( W 1 x ) ⊙ W 2 x , (6) \text{GLU}(x, W_1, W_2) = \sigma(W_1 x) \odot W_2 x, \tag{6} GLU(x,W1,W2)=σ(W1x)W2x,(6)

其中, ⊙ \odot 表示逐元素乘法,门控线性单元被认为可以 在保留非线性表达能力的同时,通过为梯度提供一条线性通路,从而缓解深层网络中的梯度消失问题

SiLU/Swish 激活函数与 GLU 机制结合起来就得到了 SwiGLU,这也是我们在前馈网络中将要使用的结构,其形式定义如下:

FFN ( x ) = SwiGLU ( x , W 1 , W 2 , W 3 ) = W 2 ( SiLU ( W 1 x ) ⊙ W 3 x ) , (7) \text{FFN}(x) = \text{SwiGLU}(x, W_1, W_2, W_3) = W_2\big(\text{SiLU}(W_1 x) \odot W_3 x\big), \tag{7} FFN(x)=SwiGLU(x,W1,W2,W3)=W2(SiLU(W1x)W3x),(7)

其中 x ∈ R d model x \in \mathbb{R}^{d_{\text{model}}} xRdmodel W 1 , W 3 ∈ R d ff × d model W_1, W_3 \in \mathbb{R}^{d_{\text{ff}} \times d_{\text{model}}} W1,W3Rdff×dmodel W 2 ∈ R d model × d ff W_2 \in \mathbb{R}^{d_{\text{model}} \times d_{\text{ff}}} W2Rdmodel×dff,并且通常取 d ff = 8 3 d model d_{\text{ff}} = \tfrac{8}{3} d_{\text{model}} dff=38dmodel

[Shazeer 2020] 首次提出将 SiLU/Swish 激活函数GLU 机制相结合,并通过实验表明,SwiGLU 在语言建模任务上优于诸如 ReLUSiLU(不带门控) 等基线方法,在本次作业的后续部分中,你也将对 SwiGLUSiLU 进行对比分析

尽管我们已经对这些组件给出了一些启发式的解释,相关论文中也提供了更多实验支持,但从经验主义的角度来看,仍然值得保持一种开放态度,Shazeer 在其论文中有一句如今广为流传的话:“我们并不解释这些架构为何有效;和其他一切一样,我们将它们的成功归因于神的仁慈”

Problem (positionwise_feedforward): Implement the position-wise feed-forward network (2 points)

Deliverable:请实现一个 SwiGLU 前馈网络,该网络由 SiLU 激活函数GLU(门控线性单元) 组成

Note:在这一具体实现中,为了提高数值稳定性,你可以在代码中直接使用 torch.sigmoid

在实现时,你应当将前馈网络的中间维度 d ff d_\text{ff} dff 设为大约 d ff ≈ 8 3 × d model d_{\text{ff}} \approx \frac{8}{3} \times d_{\text{model}} dff38×dmodel,同时需要确保 前馈网络内部层的维度是 64 的整数倍,以便更好地利用硬件性能

为了使用我们提供的测试用例验证你的实现,你需要先实现测试适配器 [adapters.run_swiglu],随后运行以下命令进行测试:

uv run pytest -k test_swiglu
1.5.3 Relative Positional Embeddings

为了向模型中注入位置信息,我们将实现 旋转位置编码(Rotary Position Embeddings, ROPE) [su+ 2021],也常简称为 RoPE

对于位于位置 i i i 的某个查询(query)token,其向量表示为:

q ( i ) = W q x ( i ) ∈ R d q^{(i)} = W_q x^{(i)} \in \mathbb{R}^d q(i)=Wqx(i)Rd

我们将对其施加一个成对的旋转矩阵 R i R^i Ri,从而得到:

q ′ ( i ) = R i q ( i ) = R i W q x ( i ) q'^{(i)} = R^i q^{(i)} = R^i W_q x^{(i)} q(i)=Riq(i)=RiWqx(i)

在这里, R i R^i Ri 会将嵌入向量中的元素对 ( q 2 k − 1 ( i ) , q 2 k ( i ) ) \big(q^{(i)}_{2k-1}, q^{(i)}_{2k}\big) (q2k1(i),q2k(i)) 视为二维向量,并按角度 θ i , k \theta_{i,k} θi,k 进行旋转,其中:

θ i , k = i Θ ( 2 k − 2 ) / d , k ∈ { 1 , … , d / 2 } \theta_{i,k} = \frac{i}{\Theta^{(2k-2)/d}}, \quad k \in \{1, \dots, d/2\} θi,k=Θ(2k2)/di,k{1,,d/2}

Θ \Theta Θ 是一个常数

因此,可以将 R i R^i Ri 看作一个大小为 d × d d\times d d×d块对角矩阵(block-diagonal matrix),其对角块为 R k i ( k = 1 , … , d / 2 ) R_k^i(k = 1, \dots, d/2) Rki(k=1,,d/2),每个块的形式为:

R k i = [ cos ⁡ ( θ i , k ) − sin ⁡ ( θ i , k ) sin ⁡ ( θ i , k ) cos ⁡ ( θ i , k ) ] . (8) R_k^i = \begin{bmatrix} \cos(\theta_{i,k}) & -\sin(\theta_{i,k}) \\ \sin(\theta_{i,k}) & \cos(\theta_{i,k}) \end{bmatrix}. \tag{8} Rki=[cos(θi,k)sin(θi,k)sin(θi,k)cos(θi,k)].(8)

由此,我们可以得到完整的旋转矩阵:

R i = [ R 1 i 0 0 ⋯ 0 0 R 2 i 0 ⋯ 0 0 0 R 3 i ⋯ 0 ⋮ ⋮ ⋮ ⋱ ⋮ 0 0 0 ⋯ R d / 2 i ] , (9) R^i = \begin{bmatrix} R_1^i & 0 & 0 & \cdots & 0 \\ 0 & R_2^i & 0 & \cdots & 0 \\ 0 & 0 & R_3^i & \cdots & 0 \\ \vdots & \vdots & \vdots & \ddots & \vdots \\ 0 & 0 & 0 & \cdots & R_{d/2}^i \end{bmatrix}, \tag{9} Ri= R1i0000R2i0000R3i0000Rd/2i ,(9)

其中的 0 表示 x × 2 x \times 2 x×2 的零矩阵

虽然理论上可以显式构造完整的 d × d d \times d d×d 旋转矩阵,但在实现时,一个好的方案应当 利用该矩阵的结构特性,以更高效的方式完成变换。由于我们只关心用一序列中 token 之间的 相对旋转关系,因此可以在不同层、不同 batch 之间 复用 预先计算好的 cos ⁡ ( θ i , k ) \cos(\theta_{i,k}) cos(θi,k) sin ⁡ ( θ i , k ) \sin(\theta_{i,k}) sin(θi,k)

如果你希望进一步优化实现,可以使用一个 由所有层共享的 RoPE 模块,并在初始化时创建一个二维的正弦和余弦预计算缓冲区(buffer),通过 self.register_buffer(persistent=False) 进行注册,而不是将其作为 nn.Parameter,因为我们 并不希望学习这些固定的正弦和余弦值

对键(key)向量 k ( i ) k^{(i)} k(i) 的处理过程与对查询向量 q ( i ) q^{(i)} q(i) 完全相同,同样通过对应的旋转矩阵 R i R^i Ri 进行旋转。需要注意的是,这一层本身不包含任何可学习参数

Problem (rope): Implement RoPE (2 points)

Deliverable:请实现一个 RotaryPositionalEmbedding 类,用于将 RoPE(旋转位置编码) 应用于输入张量,推荐使用如下接口:

def __init__(self, theta: float, d_k: int, max_seq_len: int, device=None)

用于构造 RoPE 模块,并在需要时创建缓冲区(buffers),该构造函数应当接收以下参数:

  • theta: float:RoPE 中使用的常数 Θ \Theta Θ
  • d_k: int:查询(query)与键(key)向量的维度
  • max_seq_len: int:可能输入的最大序列长度
  • device: torch.device | None = None:用于存放缓冲区的设备
def forward(self, x: torch.Tensor, token_positions: torch.Tensor) -> torch.Tensor

对形状为 (..., seq_len, d_k) 的输入张量进行处理,并返回 形状相同 的输出张量

请注意以下几点:

  • 你的实现应当 支持任意数量的批处理维度,即 xseq_len 之前可以有任意多个 batch 维度
  • 可以假设 token_positions 是一个形状为 (..., seq_len) 的张量,用于指定序列维度上各 token 的位置
  • 你应当使用 token_positions,在序列维度上 切片(slice) 你可能已经预计算好的 cossin 张量

为了测试你的实现,请完成适配器 [adapters.run_rope],并确保通过以下测试命令:

uv run pytest -k test_rope
1.5.4 Scaled Dot-Product Attention

我们现在将实现 缩放点积注意力(scaled dot-product attention),其定义最早见于 [Vaswani+ 2017] 的论文,作为一个预备步骤,注意力(Attention)操作的定义会用到 softmax,该操作将一个 未归一化的分数向量 转换为一个 归一化的概率分布

softmax ( v ) i = exp ⁡ ( v i ) ∑ j = 1 n exp ⁡ ( v j ) . (10) \text{softmax}(v)_i = \frac{\exp(v_i)}{\sum_{j=1}^{n} \exp(v_j)}. \tag{10} softmax(v)i=j=1nexp(vj)exp(vi).(10)

需要注意的是,当 v i v_i vi 的值很大时, exp ⁡ ( v i ) \exp(v_i) exp(vi) 可能会变成无穷大(此时会出现 ∞ / ∞ = NaN \infty / \infty = \text{NaN} ∞/∞=NaN 的问题),我们可以利用 softmax 的一个重要性质来避免这一数值问题:对所有输入同时加上一个常数 c c c,softmax 的结果保持不变

因此,在实际实现中我们可以利用这一性质来提升数值稳定性,通常的做法是从向量 o i o_i oi 的所有元素中减去其中的最大值,使新的最大元素变为 0。接下来,你将使用这一技巧来实现 softmax,以确保计算过程的数值稳定性

现在我们可以从数学角度来定义注意力(Attention)操作:

A t t e n t i o n ( Q , K , V ) = s o f t m a x ( Q ⊤ K d k ) V (11) \mathrm{Attention}(Q, K, V) = \mathrm{softmax}\left(\frac{Q^\top K}{\sqrt{d_k}}\right) V \tag{11} Attention(Q,K,V)=softmax(dk QK)V(11)

其中, Q ∈ R n × d k Q \in \mathbb{R}^{n \times d_k} QRn×dk K ∈ R m × d k K \in \mathbb{R}^{m \times d_k} KRm×dk V ∈ R m × d v V \in \mathbb{R}^{m \times d_v} VRm×dv,这里的 Q Q Q K K K V V V 都是该操作的输入,需要注意的是,它们并不是可学习的参数,如果你疑惑为什么这里不是使用 Q K ⊤ QK^\top QK,可以参考第 1.3.3.1 小节 的讨论

在某些情况下,对注意力操作的输出进行 掩码(masking) 是很有用的,一个掩码应当具有形状 M ∈ { True , False } n × m M \in \{\text{True}, \text{False}\}^{n \times m} M{True,False}n×m,其中布尔矩阵的每一行 i i i 表示第 i i i 个 query 可以关注哪些 key

按照约定(虽然直觉上有些容易混淆):

  • 在位置 ( i , j ) (i,j) (i,j) 处取值为 True,表示第 i i i 个 query 可以 关注第 j j j 个 key
  • 取值为 False,表示第 i i i 个 query 不可以 关注该 key

换句话说,只有在 ( i , j ) (i,j) (i,j) 处为 True 时,“信息” 才允许在 query i i i 与 key j j j 之间流动,例如,考虑一个 1 × 3 1\times 3 1×3 的掩码矩阵 [ True , True , False ] [ \text{True}, \text{True}, \text{False} ] [True,True,False],这表示单个 query 向量只会关注前两个 key

在计算层面上,使用掩码通常比在子序列上单独计算注意力要高效得多,实现方式是在 softmax 之前得注意力分数 Q ⊤ K d k \frac{Q^\top K}{\sqrt{d_k}} dk QK 中,对掩码矩阵中为 False 的位置上加上一个 − ∞ -\infty ,从而在 softmax 之后将这些位置的权重压到 0

Problem (softmax): Implement softmax (1 point)

Deliverable:编写一个函数,用于对一个张量应用 softmax 操作,你的函数应当接受两个参数:一个输入张量(tensor)和一个维度索引 i i i,并在输入张量的第 i i i 个维度上应用 softmax 运算。

输出张量应当与输入张量具有 相同的性质,但其第 i i i 个维度上的值将构成一个 归一化的概率分布。为避免数值稳定性问题,请使用如下技巧:在第 i i i 个维度上,对该维度的所有元素减去该维度上的最大值,再计算 softmax

为了测试你的实现,请完成适配器 [adapters.run_softmax],并确保通过以下测试命令:

uv run pytest -k test_softmax_matches_pytorch
Problem (scaled_dot_product_attention): Implement scaled dot-product attention (5 points)

Deliverable:实现缩放点积注意力(scaled dot-product attention)函数,你的实现需要支持如下输入形式:

  • Query 和 Key 的形状为 ( batch_size , … , seq_len , d_k ) \left(\text{batch\_size}, \ldots, \text{seq\_len}, \text{d\_k}\right) (batch_size,,seq_len,d_k)
  • Value 的形状为: ( batch_size , … , seq_len , d_v ) \left(\text{batch\_size}, \ldots, \text{seq\_len}, \text{d\_v}\right) (batch_size,,seq_len,d_v)

其中,“…” 表示任意数量的其他类似 batch 的维度,函数应当返回形状为 ( batch_size , … , d_v ) \left(\text{batch\_size}, \ldots, \text{d\_v}\right) (batch_size,,d_v) 的输出张量,关于 batch-like 维度的详细讨论可参考 第 1.3 小节

你的实现还需要支持一个 可选的、由用户提供的布尔掩码(mask),其形状为 ( seq_len , seq_len ) \left(\text{seq\_len},\text{seq\_len}\right) (seq_len,seq_len),对应掩码中值为 True 的位置,其对应的注意力概率在该维度上应当 共同归一化为 1,而掩码值为 False 的位置,其注意力概率应当为 0

为了验证你的实现是否正确,你需要在 [adapters.run_scaled_dot_product_attention] 中完成测试适配器,随后运行:

uv run pytest -k test_scaled_dot_product_attention

以测试你对 三阶输入张量 的实现

运行:

uv run pytest -k test_4d_scaled_dot_product_attention

以测试你对 四阶输入张量 的实现

1.5.5 Causal Multi-Head Self-Attention

我们将按照 [Vaswani+ 2017] 论文中第 3.2.2 节中所描述的方法实现多头注意力机制,回顾一下,从数学角度来看,多头注意力的计算形式定义如下:

MultiHead ( Q , K , V ) = Concat ( head 1 , … , head h ) (12) \text{MultiHead}(Q, K, V) = \text{Concat}(\text{head}_1, \ldots, \text{head}_h) \tag{12} MultiHead(Q,K,V)=Concat(head1,,headh)(12)

其中,每一个注意力头定义为:

head i = Attention ( Q i , K i , V i ) (13) \text{head}_i = \text{Attention}(Q_i, K_i, V_i) \tag{13} headi=Attention(Qi,Ki,Vi)(13)

这里, Q i , K i , V i Q_i, K_i, V_i Qi,Ki,Vi 分别是从 Q , K , V Q, K, V Q,K,V 中切分出的第 i i i 个子张量, i ∈ 1 , … , h i \in {1, \ldots, h} i1,,h,其维度大小分别为 d k d_k dk d v d_v dv,对应于 Q , K , V Q, K, V Q,K,V 的嵌入维度切分

其中的 Attention 操作即为 第 1.5.4 节 中定义的缩放点积注意力(scaled dot-product attention),基于这一点,我们可以将 多头自注意力 操作写为:

MultiHeadSelfAttention ( x ) = W O MultiHead ( W Q x , W K x , W V x ) (14) \text{MultiHeadSelfAttention}(x) = W_O \text{MultiHead}(W_Q x, W_K x, W_V x) \tag{14} MultiHeadSelfAttention(x)=WOMultiHead(WQx,WKx,WVx)(14)

在这里,可学习的参数包括: W Q ∈ R h d k × d model , W K ∈ R h d k × d model , W V ∈ R h d v × d model , W O ∈ R d model × h d v W_Q \in \mathbb{R}^{h d_k \times d_{\text{model}}}, W_K \in \mathbb{R}^{h d_k \times d_{\text{model}}}, W_V \in \mathbb{R}^{h d_v \times d_{\text{model}}},W_O \in \mathbb{R}^{d_{\text{model}} \times h d_v} WQRhdk×dmodel,WKRhdk×dmodel,WVRhdv×dmodel,WORdmodel×hdv

由于在多头注意力中, Q Q Q K K K V V V 会沿着输出维度被切分为多个注意力头,因此你可以将 W Q W_Q WQ W K W_K WK W V W_V WV 理解为在输出维度上为每一个注意力头分别设置的投影矩阵,当你完成这一实现后,你会发现计算 key、query 和 value 投影本质上需要 三次矩阵乘法

Note:作为拓展目标,你可以尝试将 key、query 和 value 的投影合并为 一个 权重矩阵,从而只需要进行一次矩阵乘法

Causal masking

你的实现应当防止模型在序列中 关注未来的 token,换句话说,假设模型接收到一个 token 序列 t 1 , … , t n t_1, \ldots, t_n t1,,tn,并且我们希望基于前缀 t 1 , … , t i t_1, \ldots, t_i t1,,ti(其中 i < n i<n i<n)来计算下一个词的预测,那么模型 不应该 访问(或者说注意到)位置 t i + 1 , … , t n t_{i+1}, \ldots, t_n ti+1,,tn 处的 token 表示。这是因为在推理阶段生成文本时,模型并不能访问这些未来 token,而这些 token 会泄露真实下一个词的身份,从而使语言模型的预训练目标变得毫无意义

对于一个输入 token 序列 t 1 , … , t n t_1, \ldots, t_n t1,,tn,我们当然可以通过对每一个前缀分别运行一次多头自注意力来 “朴素地” 阻止访问未来 token(即对序列中每个唯一前缀单独计算),然而,更高效的做法是使用 因果注意力掩码(causal attention masking),它允许序列中位置 i i i 的 token 只关注所有满足 j ≤ i j\le i ji 的位置

你可以使用 torch.triu 或基于广播的索引来构造这个掩码,并且应当利用这样一个事实:你在第 1.5.4 节中实现的缩放点积注意力已经支持注意力掩码

Applying RoPE.

RoPE 应当应用在 query(查询)向量和 key(键)向量上,而不应用于 value(值)向量,此外,在多头注意力中,注意力头的维度应当被视作一个批次维度来处理,因为每个注意力头都是 彼此独立地 计算注意力的,这意味着:对于每一个注意力头,都应当对其 query 和 key 向量应用完全相同的 RoPE 旋转操作

Problem (multihead_self_attention): Implement causal multi-head self-attention (5 points)

Deliverable:实现一个 因果多头自注意力(causal multi-head self-attention) 模块,形式为一个 torch.nn.Module,你的实现至少应当接收以下参数:

  • d_model: int:Transformer 块输入的特征维度
  • num_heads: int:多头自注意力中使用的注意力头数量

按照 [Vaswani+ 2017] 的设定,令 d k = d v = d model / h d_k = d_v = d_{\text{model}}/{h} dk=dv=dmodel/h,其中 h h h 为注意力头的数量

为了使用我们提供的测试来验证你的实现,你需要在 [adapters.run_multihead_self_attention] 中实现对应的测试适配器,随后运行:

uv run pytest -k test_multihead_self_attention

来测试你的实现是否正确

1.6 The Full Transformer LM

我们现在开始 组装 Transformer 块(此时回顾 Figure 2 会很有帮助),一个 Transformer 块包含两个 “子层(sublayers)”,一个用于 多头自注意力(multi-head self-attention),另一个用于 前馈网络(feed-forward network)。在每一个子层中,计算流程是相同的:首先执行 RMSNorm,然后进行该子层的主要计算(MHA/FF),最后再加上 残差连接(residual connection)

更具体地说,Transformer 块中 第一个子层(即注意力子层)应当实现如下更新过程:给定输入 x x x,输出 y y y 的计算方式为:

y = x + MultiHeadSelfAttention ( RMSNorm ( x ) ) . (15) y = x + \text{MultiHeadSelfAttention}(\text{RMSNorm}(x)). \tag{15} y=x+MultiHeadSelfAttention(RMSNorm(x)).(15)

Problem (transformer_block): Implement the Transformer block (3 points)

请按照 §1.5 中的描述并参考 Figure 2,实现一个 pre-norm Transformer 块,你的 Transformer 块至少应当接受以下参数:

  • d_model: int:Transformer 块输入的特征维度
  • num_heads: int:多头自注意力中使用的注意力头数量
  • d_ff: int:位置前馈网络中内部隐藏层的维度

为了测试你的实现,请在 [adapters.run_transformer_block] 中实现对应的测试适配器,然后运行:

uv run pytest -k test_transformer_block

以验证你的实现是否正确

Deliverable:一份能够通过所有提供测试的 Transformer 块实现代码

Problem (transformer_lm): Implementing the Transformer LM (3 points)

现在我们将所有模块组合在一起,整体流程如 Figure 1 中的高层结构示意所示,按照 §1.1.1 中对嵌入层(embedding)的描述,首先对输入进行嵌入处理,然后将结果送入 num_layers 个 Transformer 块中,最后再将输出传入三个输出层,从而得到在整个词表上的概率分布

现在是把所有组件整合在一起的时候了!请按照 §1.1 中的描述,并结合 Figure 1 所示的结构,实现一个 Transformer 语言模型。至少,你的实现需要支持前面所有 Transformer 块的构造参数,此外还应支持以下额外参数:

  • vocab_size: int:词表大小,用于确定词嵌入矩阵(token embedding matrix)的维度
  • context_length: int:最大上下文长度,用于确定位置嵌入矩阵(position embedding matrix)的维度
  • num_layers: int:使用的 Transformer 块的数量

为了使用我们提供的测试来验证你的实现,你首先需要在 [adapters.run_transformer_lm] 中实现测试适配器,然后运行:

uv run pytest -k test_transformer_lm

以测试你的实现

Deliverable:一个能够通过上述测试的 Transformer 语言模型模块

Problem (transformer_accounting): Transformer LM resource accounting (5 points)

Resource accounting.

理解 Transformer 各个组成部分在 计算量和内存 方面的消耗是非常有帮助的,接下来我们将通过几个步骤进行一次基础的 FLOPs(浮点运算次数)核算

由于 Transformer 中绝大多数 FLOPs 都来自矩阵乘法,因此我们的核心思路非常简单:

  1. 列出 Transformer 前向传播过程中涉及的所有矩阵乘法
  2. 将每一个矩阵乘法转换为所需的 FLOPs 数量

在第二步中,下面这个事实会非常有用:

Rule: 给定矩阵 A ∈ R m × n A \in \mathbb{R}^{m \times n} ARm×n B ∈ R n × p B \in \mathbb{R}^{n \times p} BRn×p,矩阵乘积 A B AB AB 需要 2 m n p 2mnp 2mnp 次 FLOPs

这是因为 ( A B ) [ i , j ] = A [ i , : ] ⋅ B [ : , j ] (AB)[i, j] = A[i, :] \cdot B[:, j] (AB)[i,j]=A[i,:]B[:,j] 这个点积需要 n n n 次加法和 n n n 次乘法,总共是 2 n 2n 2n 次 FLOPs,而矩阵 A B AB AB 一共有 m × p m\times p m×p 个元素,因此,总 FLOPs 数为 ( 2 n ) ( m p ) = 2 m n p (2n)(mp) = 2mnp (2n)(mp)=2mnp

在继续下一个问题之前,建议你先逐一检查自己实现的 Transformer blockTransformer 语言模型(Transformer LM) 中的每一个组件,列出其中所有涉及的矩阵乘法,以及它们各自对应的 FLOPs 开销

(a)考虑 GPT-2 XL,其配置如下:

  • vocab_size:50,257
  • context_length:1,024
  • num_layers:48
  • d_model:1,600
  • num_heads:25
  • d_ff:6,400

假设我们使用上述配置构建了模型:

  • 这个模型一共有多少个 可训练参数
  • 如果每个参数都使用 单精度浮点数(float32) 表示,仅加载该模型需要多少内存?

Deliverable:用一到两句话回答。

(b)请识别完成一次 GPT-2 XL 规模模型 前向传播所需的 所有矩阵乘法操作,假设输入序列长度等于 context_length,这些矩阵乘法总共需要多少 FLOPs

Deliverable:列出所有矩阵乘法(附简要说明),并给出所需 FLOPs 的总数。

(c)基于你在上一步中的分析,模型中 哪些部分消耗了最多的 FLOPs

Deliverable:用一到两句话回答。

(d)请将你的分析扩展到以下模型配置:

  • GPT-2 small:12 层,d_model = 768,12 个注意力头
  • GPT-2 medium:24 层,d_model = 1024,16 个注意力头
  • GPT-2 large:36 层,d_model = 1280,20 个注意力头

随着模型规模的增大:

  • Transformer LM 的哪些组成部分在总 FLOPs 中所占比例 增加
  • 哪些部分所占比例 减少

Deliverable:对每个模型给出各组件的 FLOPs 分解(以占总前向 FLOPs 的比例表示),并用一到两句话说明模型规模变化如何影响各组件的 FLOPs 占比。

(e)以 GPT-2 XL 为例,将 context_length 增加到 16,384

  • 单次前向传播的 总 FLOPs 会如何变化?
  • 模型各组件在 FLOPs 中的 相对贡献比例 将如何变化?

Deliverable:用一到两句话回答。

结语

这篇文章我们系统性地梳理了 CS336 Assignment 1 中 Transformer Language Model Architecture 的全部作业要求,从整体模型结构出发,逐步拆解了 embedding、pre-norm Transformer block、多头自注意力、SwiGLU 前馈网络、RoPE 位置编码,以及最终语言模型的组装方式与计算资源核算思路

更详细的内容大家可以查看官方提供的相关文档

下篇文章我们就来一起看看 Transformer Language Model Architecture 具体该如何实现,敬请期待🤗

参考

Logo

码道开发者社区,聚焦华为云码道 CodeArts 代码智能体,沉淀 Agent、Skill、鸿蒙开发实战内容,供开发者查阅资料、交流技术、分享工程实践

更多推荐