
1. 为什么看懂Transformer源码不是“大神专利”而是每个想真正用好它的人都该跨过的门槛我带过不少刚从学校出来的实习生也辅导过不少转行做AI工程的职场人。他们有个共同点能调用torch.nn.TransformerEncoderLayer能跑通Hugging Face的pipeline甚至能微调BERT——但只要模型输出结果异常或者想改一个注意力计算的mask逻辑就立刻卡住。不是不会查文档而是文档里写的“attn_maskis applied before softmax”和代码里那一行attn_weights attn_weights.masked_fill(attn_mask, float(-inf))之间隔着一层看不见的玻璃。这层玻璃就是源码。很多人误以为“看懂源码”等于“从零手写一个Transformer”。其实完全不是。PyTorch官方实现的torch.nn.Transformer模块核心逻辑就集中在不到300行Python代码里不含注释和空行。它不追求极致性能不堆砌CUDA内核而是用最清晰、最符合论文原意的方式组织结构。它就像一本用Python写成的《Attention Is All You Need》教科书——每一行都在翻译公式每一个变量名都在呼应论文图2。你不需要成为编译器专家也不需要精通GPU调度只需要理解词嵌入怎么变成向量、多头注意力怎么并行计算、前馈网络为什么是两层线性加激活——这些全在源码里明明白白写着。关键词“Transformer”“PyTorch”“源代码”“词嵌入”“多头注意力”不是孤立的标签它们是这条理解路径上的五个路标。词嵌入是起点把文字变成数字多头注意力是心脏决定信息如何流动PyTorch是工具让数学表达可执行源代码是地图告诉你每一步脚踩在哪块砖上而Transformer是整座建筑的名字。这篇解读不讲宏观架构图不列公式推导就打开torch/nn/modules/transformer.py和torch/nn/functional.py这两个文件一行一行带你走完从输入文本到最终输出的完整数据流。你会发现所谓“源码”不过是把论文里的方框箭头翻译成了x self.norm1(x self._sa_block(x, src_mask, src_key_padding_mask))这样一句可读、可调试、可修改的Python。提示本文所有代码片段均来自PyTorch 2.3.0官方源码torch2.3.0路径为torch/nn/modules/transformer.py主模块和torch/nn/functional.py底层函数。请确保你的环境版本一致避免因API变更导致理解偏差。文中所有变量命名、参数顺序、默认值均严格对应源码不做任何“简化版”或“教学版”改写。2. 从nn.Transformer入口开始解构一个标准Encoder-Decoder模型的初始化逻辑当你写下model nn.Transformer(d_model512, nhead8, num_encoder_layers6)时PyTorch做的第一件事不是构建计算图而是校验参数的数学合理性。这一步藏在__init__方法的开头却常被忽略——它直接决定了后续所有张量运算能否成立。2.1 参数校验为什么d_model必须能被nhead整除源码中第一段关键逻辑是if d_model % nhead ! 0: raise ValueError(fembed_dim {d_model} not divisible by num_heads {nhead})这行检查背后是多头注意力机制的数学硬约束。d_model是整个模型的隐藏层维度即每个词向量的总长度nhead是头的数量。每个头要独立处理一部分特征所以必须将d_model平均分配给nhead个头每个头分得d_k d_v d_model // nhead维。如果不能整除比如d_model512, nhead3那么512//3≈170.666无法分配整数维的向量。这不是PyTorch的“任性”而是Vaswani论文中d_k d_v d_model / h这一定义的必然要求。我曾见过有人强行绕过此检查把d_model设为513、nhead设为3结果在_scaled_dot_product_attention函数里q.size(-1)即d_k变成了非整数直接触发RuntimeError: expected scalar type Float but found Long——因为PyTorch张量维度必须是整数。2.2 模块组装Encoder与Decoder的“骨架”是如何搭起来的nn.Transformer的主体结构非常清晰它内部持有self.encoder和self.decoder两个子模块而这两个子模块又分别由多个TransformerEncoderLayer或TransformerDecoderLayer堆叠而成。源码中关键的一行是self.encoder TransformerEncoder(encoder_layer, num_encoder_layers, normencoder_norm) self.decoder TransformerDecoder(decoder_layer, num_decoder_layers, normdecoder_norm)这里没有魔法。encoder_layer是一个单层Encoder的模板num_encoder_layers6意味着用这个模板复制6份并按顺序连接。这种设计体现了PyTorch的“组合优于继承”哲学你不需要为6层Encoder写一个新类只需定义好1层的行为然后用nn.Sequential或自定义容器将其堆叠。TransformerEncoder类本身就是一个轻量级包装器其核心逻辑只有一行def forward(self, src, maskNone, src_key_padding_maskNone): output src for mod in self.layers: output mod(output, src_maskmask, src_key_padding_masksrc_key_padding_mask) if self.norm is not None: output self.norm(output) return output注意for mod in self.layers:这一循环。它明确告诉你Transformer的深度就是这个for循环的迭代次数。每一层的输出都作为下一层的输入。而mod(output, ...)调用的正是TransformerEncoderLayer的forward方法。这意味着理解整个Encoder等价于彻底吃透TransformerEncoderLayer这一单层的全部逻辑。我们接下来就聚焦于此。2.3TransformerEncoderLayer一个“标准单元”的四步流水线TransformerEncoderLayer是整个Transformer架构的原子单位。它的forward方法完美复现了论文图1中Encoder Block的四个核心组件多头自注意力Multi-head Self-Attention、Add Norm、前馈网络Feed-Forward Network、再次Add Norm。源码将其组织为一条清晰的四步流水线def forward(self, src, src_maskNone, src_key_padding_maskNone): # Step 1: Multi-head self-attention src2 self.self_attn(src, src, src, attn_masksrc_mask, key_padding_masksrc_key_padding_mask)[0] # Step 2: Add Norm (residual connection layer norm) src self.norm1(src src2) # Step 3: Feed-forward network src2 self.linear2(self.dropout(self.activation(self.linear1(src)))) # Step 4: Add Norm again src self.norm2(src src2) return src这段代码的精妙之处在于它把论文中复杂的并行计算拆解成了程序员最熟悉的“变量赋值函数调用”序列。src2是注意力层的输出src是原始输入src self.norm1(src src2)这一行同时完成了残差连接src src2和层归一化self.norm1(...)前馈网络则被展开为linear1 - activation - dropout - linear2的线性变换链。这里没有任何隐藏的控制流没有异步调度就是纯粹的数据流。你可以在PyCharm里在这一行打上断点运行时亲眼看到src的shape从(seq_len, batch_size, d_model)经过self_attn后src2的shape保持完全一致——这是Transformer“恒等映射”特性的直接体现也是它能稳定训练的基石。注意self_attn是一个MultiheadAttention实例它本身也是一个nn.Module。这意味着src2 self.self_attn(...)[0]这行代码会进一步调用MultiheadAttention.forward()。我们将在下一节深入这个核心组件。现在请牢牢抓住这个四步流水线的节奏Attention → Norm → FFN → Norm。这是所有Transformer变体BERT、GPT、ViT共享的DNA。3.MultiheadAttention揭开多头自注意力机制的“黑箱”看清每一行代码对应的数学含义如果说TransformerEncoderLayer是骨架那么MultiheadAttention就是心脏。它的forward方法是整个Transformer源码中数学密度最高的部分。但别被“黑箱”吓住——它只是把论文公式(1)到(3)逐字翻译成了Python和PyTorch张量操作。我们来一行一行把它“翻译”回人类语言。3.1 输入预处理Q/K/V矩阵的生成与维度重塑MultiheadAttention.forward()的开头是三组线性变换q, k, v F.linear(query, self.in_proj_weight, self.in_proj_bias).chunk(3, dim-1)这行代码是理解多头注意力的钥匙。F.linear是PyTorch的底层线性层query是输入张量shape为(L, N, E)即seq_len, batch_size, embed_dim。self.in_proj_weight是一个巨大的权重矩阵其shape为(3*E, E)self.in_proj_bias是(3*E,)的偏置向量。F.linear(...)的输出是一个(L, N, 3*E)的张量然后.chunk(3, dim-1)将其在最后一个维度dim-1即embed_dim维度上切成三等份得到q,k,v三个张量每个都是(L, N, E)。这对应着论文中的公式Q XWQ, K XWK, V XWV但这里有一个关键细节PyTorch没有为Q/K/V分别定义三个独立的线性层而是用一个大的权重矩阵一次性计算再切分。这是为了提升GPU内存访问效率——一次大矩阵乘法比三次小矩阵乘法对GPU更友好。self.in_proj_weight的前E行对应WsupQ/sup中间E行对应WsupK/sup最后E行对应WsupV/sup。你可以把它想象成一个“三合一”的投影仪一束光query照进去同时投射出三幅不同的影子q,k,v。紧接着代码对这三个张量进行维度重塑为多头并行做准备q q.contiguous().view(q.shape[0], q.shape[1], self.num_heads, self.head_dim).transpose(0, 2) k k.contiguous().view(k.shape[0], k.shape[1], self.num_heads, self.head_dim).transpose(0, 2) v v.contiguous().view(v.shape[0], v.shape[1], self.num_heads, self.head_dim).transpose(0, 2)view(...)将(L, N, E)reshape为(L, N, H, D_h)其中H是头数D_h E // H是每个头的维度。transpose(0, 2)则交换第0维L和第2维H得到(H, N, L, D_h)。这个变换的意义是把“序列长度×批次×总维度”的张量变成“头数×批次×序列长度×头维度”。这样每个头的计算就可以在H这个维度上完全并行互不干扰。q[0]就是第一个头的Queryq[1]是第二个头的Query……q[H-1]是最后一个头的Query。这就是“多头”的物理实现。3.2 核心计算缩放点积注意力Scaled Dot-Product Attention的完整实现多头注意力的核心是_scaled_dot_product_attention这个函数。它封装了论文公式(1)的全部逻辑计算QKT缩放应用masksoftmax再乘以V。源码如下已简化注释def _scaled_dot_product_attention(query, key, value, attn_maskNone, dropout_p0.0, is_causalFalse): # Step 1: Compute Q K^T B, Nt, E query.shape # (batch, target_seq_len, head_dim) query query / math.sqrt(E) # Scale by sqrt(d_k) # Step 2: Compute attention scores attn torch.bmm(query, key.transpose(-2, -1)) # (B, Nt, Ns) # Step 3: Apply attention mask (if provided) if attn_mask is not None: attn attn_mask # Step 4: Apply softmax to get attention weights attn torch.softmax(attn, dim-1) # Step 5: Apply dropout (optional) if dropout_p 0.0: attn torch.dropout(attn, dropout_p, trainTrue) # Step 6: Compute weighted sum of values output torch.bmm(attn, value) # (B, Nt, E) return output, attn这里有几个极易误解的点必须澄清缩放的位置query query / math.sqrt(E)发生在bmm之前。这是为了防止QKsupT/sup的数值过大导致softmax梯度消失。E在这里是head_dim即d_k不是d_model。很多初学者误以为是除以d_model这是错误的。mask的加法而非乘法attn attn_mask。attn_mask通常是一个全0或全-inf的张量。加-inf会使对应位置的softmax输出趋近于0从而屏蔽掉那些位置的注意力。这是一种数值稳定的实现方式比用masked_fill更高效。bmm的维度torch.bmm是batch matrix multiplication要求输入是(B, N, M)和(B, M, P)。query是(B, Nt, E)key.transpose(-2, -1)是(B, E, Ns)所以输出attn是(B, Nt, Ns)即每个目标位置对每个源位置的注意力分数。这正是注意力权重矩阵的形状。我曾经在一个长文本生成任务中发现模型总是忽略开头的几个词。调试时打印出attn矩阵发现attn_mask的形状是(1, 1, 512, 512)而attn是(batch, Nt, Ns)。由于维度不匹配mask根本没有生效根源在于attn_mask的构造方式错误。正确的做法是对于因果掩码causal mask应使用torch.triu(torch.full((seq_len, seq_len), float(-inf)), diagonal1)然后unsqueeze(0).unsqueeze(0)扩展到(1, 1, seq_len, seq_len)再传入。源码的健壮性恰恰要求你理解每一维的物理意义。3.3 多头融合如何把H个头的输出拼接回一个向量经过_scaled_dot_product_attention我们得到了H个头各自的输出每个都是(H, N, L, D_h)。下一步是把它们“缝合”回一个(L, N, E)的张量。源码用了一行极其优雅的代码完成attn_output attn_output.transpose(0, 2).contiguous().view(L, N, E)让我们逆向拆解attn_output初始shape是(H, N, L, D_h)头数×批次×序列长×头维。transpose(0, 2)交换头数和序列长维度得到(L, N, H, D_h)。contiguous()确保内存连续为view做准备。view(L, N, E)将最后两个维度H和D_h合并为E H * D_h得到(L, N, E)。这行代码就是论文中公式(2)Concat(head_1, ..., head_h)WsupO/sup的前半部分。“Concat”操作在PyTorch里就是view的reshape而WsupO/sup则由self.out_proj这个线性层完成attn_output self.out_proj(attn_output)self.out_proj的权重矩阵shape是(E, E)它把拼接后的E维向量再次投影回E维。这个投影层至关重要——它让不同头学到的特征能够相互“交流”和“混合”而不是简单地并列堆叠。没有它多头注意力就退化成了H个独立的单头注意力失去了“多头”带来的表征能力提升。提示MultiheadAttention的forward方法末尾还有一行return attn_output, attn_weights。attn_weights是(H, N, L, L)的张量记录了每个头在每个位置对所有位置的注意力权重。这是调试和可视化注意力模式的黄金数据。你可以用torchvision.utils.make_grid把它画成热力图直观看到模型到底在“看”哪里。4. 词嵌入与位置编码Transformer的“输入端”如何将文字转化为可计算的向量Transformer不吃文字只吃数字。所以从原始文本到nn.Transformer的src输入中间必须经过两道关键工序词嵌入Word Embedding和位置编码Positional Encoding。PyTorch本身不提供完整的文本预处理管道但它为这两步提供了最基础、最灵活的构建块。理解它们是读懂整个数据流的起点。4.1 词嵌入nn.Embedding——一个查表器的朴素智慧nn.Embedding是PyTorch中最简单的模块之一但它承载着NLP最核心的思想用稠密向量表示稀疏符号。它的源码几乎就是一行class Embedding(Module): def forward(self, input: Tensor) - Tensor: return F.embedding(input, self.weight, self.padding_idx, self.max_norm, self.norm_type, self.scale_grad_by_freq, self.sparse)F.embedding是底层C实现但它的行为可以用一句话概括把输入张量inputshape为(N,)或(N, L)的整数索引当作“地址”去self.weightshape为(V, D)的词表矩阵里查找对应的行向量并返回这些向量组成的张量。例如假设你的词表大小V10000嵌入维度D512那么self.weight就是一个10000×512的矩阵。input torch.tensor([2, 5, 10])F.embedding(input, weight)就会返回一个3×512的张量其中第0行是weight[2]第1行是weight[5]第2行是weight[10]。这就是“查表”。这里的关键洞察是nn.Embedding不关心你输入的整数是什么含义。它可以是词ID可以是字符ID甚至可以是图像patch的ID如ViT。它只是一个通用的“索引→向量”映射器。self.weight的初始值通常是随机初始化的然后在训练中通过反向传播不断更新让语义相近的词如“猫”和“狗”在向量空间中距离更近。我见过一个常见误区有人试图用nn.Linear替代nn.Embedding。这是行不通的因为Linear的输入是浮点向量而Embedding的输入是离散的整数索引。Linear无法处理input中可能出现的padding_idx填充符也无法保证input中的每个值都在[0, V)范围内。Embedding的健壮性正在于它对输入索引的严格校验和边界处理。4.2 位置编码正弦波的魔力与可学习编码的务实选择词嵌入解决了“是什么”的问题但没解决“在哪里”的问题。Transformer没有RNN那样的时序记忆也没有CNN那样的局部感受野所以必须显式地告诉模型“这个词在句子中排第几位”。这就是位置编码Positional Encoding的使命。PyTorch官方实现提供了两种方案固定正弦位置编码nn.Transformer.generate_square_subsequent_mask配合手动实现和可学习位置编码nn.Embedding。后者是更主流、更实用的选择。源码中nn.Transformer并没有内置位置编码模块。它把选择权交给了用户。最常见的做法是像词嵌入一样用一个nn.Embedding来学习位置self.pos_embedding nn.Embedding(max_len, d_model)max_len是最大序列长度如512d_model是模型维度如512。self.pos_embedding(torch.arange(max_len))会生成一个max_len × d_model的矩阵每一行代表一个位置的编码向量。这个矩阵在训练开始时是随机初始化的然后和词嵌入一起通过反向传播学习最优的位置表示。为什么不用论文中那个著名的正弦函数论文公式(4)给出的PE(pos, 2i) sin(pos / 10000^(2i/d_model))其设计初衷是让模型能外推到训练时没见过的更长序列。但实践中绝大多数任务如机器翻译、文本分类的序列长度是固定的或有明确上限的。一个可学习的nn.Embedding其灵活性和拟合能力远超手工设计的正弦函数。它能自动学习到任务特定的位置模式比如在问答任务中“答案”往往出现在句末模型就会学到一个强指向末尾的位置编码。当然如果你真想复现正弦编码PyTorch社区有大量现成的实现。核心逻辑就是用torch.arange生成位置索引pos用torch.arange生成维度索引i然后套用sin/cos公式计算。但请记住可学习的位置编码是工业界的事实标准正弦编码更多是教学和研究场景的展示。在阅读源码时看到pos_embedding你首先应该想到的是nn.Embedding而不是一堆三角函数。4.3 输入端的完整数据流从文本到src的七步旅程现在让我们把词嵌入和位置编码串联起来还原一个典型的Transformer Encoder输入流程。假设你有一批文本已经过tokenizer处理得到了input_idsshape为(batch_size, seq_len)的整数张量词嵌入查找src self.word_embedding(input_ids)→(B, L, E)位置嵌入查找pos self.pos_embedding(torch.arange(L).to(input_ids.device))→(L, E)广播相加src src pos.unsqueeze(0)→(B, L, E)。pos.unsqueeze(0)将其变为(1, L, E)利用PyTorch的广播机制自动加到每个batch的序列上。Dropout可选src self.dropout(src)。这是标准的正则化手段防止过拟合。转置以适配nn.Transformer接口PyTorch的nn.Transformer期望输入是(seq_len, batch_size, embed_dim)即L, B, E。所以需要src src.transpose(0, 1)→(L, B, E)。生成注意力掩码可选对于Encoder通常不需要因果掩码但可能需要src_key_padding_mask来屏蔽填充符padding。这通常是一个(B, L)的布尔张量True表示该位置是padding。喂入模型output self.transformer_encoder(src, src_key_padding_masksrc_key_padding_mask)。这七步就是从一行文本到模型内部张量的完整旅程。每一步都对应着源码中一个明确的、可调试的操作。当你在forward函数里看到src self.word_embedding(src)时你就知道此刻文字已经正式变成了数字进入了可计算的世界。注意nn.Transformer的forward方法签名是forward(src, tgt, src_maskNone, tgt_maskNone, memory_maskNone, src_key_padding_maskNone, tgt_key_padding_maskNone, memory_key_padding_maskNone)。其中src是Encoder输入tgt是Decoder输入。对于纯Encoder任务如BERTtgt是不需要的你可以只传入src和相关的mask。不要被这个长长的参数列表吓住它们都是可选的且都有合理的默认值None。5. 实战调试如何在PyCharm中一步步跟踪Transformer的前向传播定位真实问题看懂源码的最高境界不是背下所有函数名而是能在模型出错时像侦探一样沿着数据流精准定位问题发生的那一行。PyCharm的调试器Debugger是你的最佳搭档。下面我以一个真实场景为例演示如何用调试器“走进”Transformer的内部。5.1 场景设定一个诡异的nan值从何而来假设你正在训练一个文本分类模型一切正常直到某一轮loss突然变成nan。你怀疑是某个层的输出出现了nan。常规做法是打印每一层的输出但那样太慢。更好的办法是设置一个条件断点让程序在nan出现的瞬间停下。步骤1在MultiheadAttention.forward中设置条件断点打开torch/nn/modules/activation.py或你的PyTorch安装路径下的对应文件找到MultiheadAttention.forward方法。在attn_output self.out_proj(attn_output)这一行左侧的灰色区域点击设置一个断点。然后右键点击断点选择“More...”在弹出的对话框中勾选“Condition”输入条件torch.isnan(attn_output).any()这个条件的意思是“只有当attn_output张量中存在任何一个nan值时才触发断点”。这样程序会在nan首次产生时精确停在out_proj这行。步骤2启动调试观察变量状态运行你的训练脚本选择“Debug”模式。当断点触发时PyCharm会暂停执行。此时打开“Variables”面板你会看到当前作用域下的所有变量。重点关注attn_output它的值已经是nan证明问题出在out_proj之前。q,k,v检查它们的max(),min(),std()。如果q或k的值异常巨大如1e8那问题就在前面的线性变换。attn_weights检查它的sum(dim-1)是否为1.0softmax的性质。如果不是说明softmax的输入即QK^T可能有nan或inf。步骤3向上追溯定位根因如果q有nan那就继续在F.linear调用处设置断点检查query和self.in_proj_weight。query来自上一层的输出self.in_proj_weight是可学习参数。如果query正常而in_proj_weight有nan那就是参数更新出了问题如学习率过大梯度爆炸。我曾遇到一个案例nan的源头是src_key_padding_mask。这个mask本应是bool类型但被错误地转换成了float导致attn attn_mask时-inf被加到了一个很大的正数上产生了nan。调试器让你一眼就能看到attn_mask.dtype是torch.float32而不是预期的torch.bool。这种细节只靠print是很难发现的。5.2 可视化注意力用attn_weights画出模型的“视线”MultiheadAttention.forward的返回值中attn_weights是每个头的注意力权重。这是理解模型行为的金钥匙。我们可以把它提取出来画成热力图。# 在你的模型forward中捕获注意力权重 def forward_with_attn(self, src, src_maskNone, src_key_padding_maskNone): # ... 原始forward逻辑 ... # 在调用self.self_attn时获取返回的attn_weights src2, attn_weights self.self_attn(src, src, src, attn_masksrc_mask, key_padding_masksrc_key_padding_mask) # ... 后续逻辑 ... return output, attn_weights # 返回attn_weights供外部使用然后在推理时model.eval() with torch.no_grad(): output, attn_weights model.forward_with_attn(src) # attn_weights shape: (H, N, L, L) # 取第一个样本、第一个头 attn_map attn_weights[0, 0].cpu().numpy() # (L, L) plt.imshow(attn_map, cmapviridis) plt.colorbar() plt.title(Attention Map (Head 0, Sample 0)) plt.show()你会看到一张L×L的热力图颜色越亮表示位置i对位置j的关注度越高。对于一个训练良好的模型你应该能看到清晰的对角线关注自己以及一些跨越短距离的亮斑关注邻近词。如果整张图都是均匀的灰色说明注意力机制没有学到有效的模式如果只有对角线亮其他地方全黑说明模型过于“自闭”没有建立长程依赖。5.3 修改源码一个安全、可逆的定制化实验有时你需要微调Transformer的行为比如改变注意力的缩放因子或者添加一个新的mask逻辑。直接修改PyTorch源码是危险的但你可以通过继承和重写来安全地实现。class CustomMultiheadAttention(nn.MultiheadAttention): def forward(self, query, key, value, key_padding_maskNone, need_weightsTrue, attn_maskNone): # 调用父类方法获取原始输出 attn_output, attn_output_weights super().forward( query, key, value, key_padding_mask, need_weights, attn_mask ) # 在这里添加你的定制逻辑 # 例如对attn_output_weights进行后处理 if attn_output_weights is not None: # 将注意力权重限制在[0.1, 0.9]之间防止过于集中或分散 attn_output_weights torch.clamp(attn_output_weights, 0.1, 0.9) attn_output_weights attn_output_weights / attn_output_weights.sum(dim-1, keepdimTrue) return attn_output, attn_output_weights # 在你的模型中使用 self.self_attn CustomMultiheadAttention(embed_dim512, num_heads8)这种方法的优势在于它完全兼容PyTorch的API不影响其他模块且易于测试和回滚。你不需要动torch/目录下的任何一行代码所有的定制都在你的项目代码里。这是工程实践中最推荐的“源码级”定制方式。提示在调试时善用PyCharm的“Evaluate Expression”功能快捷键AltF8。你可以随时输入q.mean().item()、k.std().item()来查看张量的统计信息而无需修改源码添加print语句。这能极大提升调试效率。6. 从源码到应用如何基于PyTorch官方实现快速搭建一个可落地的文本分类Pipeline理解源码的终极目的不是为了成为源码贡献者而是为了成为一个更强大、更自主的使用者。现在让我们把前面所有知识点整合成一个端到端的、可立即运行的文本分类Pipeline。这个Pipeline不依赖Hugging Face完全基于PyTorch原生API代码量控制在200行以内但涵盖了数据加载、模型定义、训练循环、评估和推理的全部环节。6.1 数据准备用torchtext构建一个极简的文本流水线我们使用经典的AG_NEWS数据集。torchtext提供了便捷的加载器from torchtext.datasets import AG_NEWS from torchtext.data.utils import get_tokenizer from torchtext.vocab import build_vocab_from_iterator from torchtext.data.functional import to_map_style_dataset # 1. 获取分词器 tokenizer get_tokenizer(basic_english) # 2. 构建词表 train_iter AG_NEWS(splittrain) vocab build_vocab_from_iterator( map(tokenizer, [label text for (label, text) in train_iter]), min_freq1, specials[unk, pad] ) vocab.set_default_index(vocab[unk]) # 3. 定义数值化函数 def yield_tokens(data_iter): for _, text in data_iter: yield tokenizer(text) # 4. 创建数据集 train_dataset to_map_style_dataset(AG_NEWS(splittrain)) test_dataset to_map_style_dataset(AG_NEWS(splittest)) # 5. 定义collate_fn负责batching def collate_batch(batch): label_list, text_list, offsets [], [], [0] for _label, _text in batch: