ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

Transformer的架构实现

Transformer的架构实现 一.整体架构二.分层实现2.1嵌入层和位置编码层代码如下:class PositionalEncoding(nn.Module): # 初始化预计算位置编码表 def __init__(self, d_modelDIM_MODEL, max_lenSEQ_LEN): super(PositionalEncoding, self).__init__() # 定义位置编码表初始全0 pe torch.zeros(max_len, d_model) # 遍历每个位置 for pos in range(max_len): # 遍历每个2i值计算这一对正余弦值 for _2i in range(0, d_model, 2): # 带公式计算 pe[pos, _2i] math.sin( pos / (10000 ** (_2i / d_model)) ) pe[pos, _2i1] math.cos( pos / (10000 ** (_2i / d_model)) ) # 注册缓冲 self.register_buffer(pe, pe) # 前向传播传入词向量 (N, L, E)返回叠加位置编码后的向量 def forward(self, x): # 从位置编码表中截取 L 个位置的编码向量形状 (L, E) pe_vectors self.pe[0 : x.shape[1]] return x pe_vectors # 广播后相加 class PositionalEncoding(nn.Module): # 初始化预计算位置编码表 def __init__(self, d_modelDIM_MODEL, max_lenSEQ_LEN): super(PositionalEncoding, self).__init__() # 定义位置编码表初始全0 pe torch.zeros(max_len, d_model) # 遍历每个位置 for pos in range(max_len): # 遍历每个2i值计算这一对正余弦值 for _2i in range(0, d_model, 2): # 带公式计算 pe[pos, _2i] math.sin( pos / (10000 ** (_2i / d_model)) ) pe[pos, _2i1] math.cos( pos / (10000 ** (_2i / d_model)) ) # 注册缓冲 self.register_buffer(pe, pe) # 前向传播传入词向量 (N, L, E)返回叠加位置编码后的向量 def forward(self, x): # 从位置编码表中截取 L 个位置的编码向量形状 (L, E) pe_vectors self.pe[0 : x.shape[1]] return x pe_vectors # 广播后相加下面进行逐行解释1)初始化需要传入的参数:d_model代表每个位置的向量宽度;max_len代表每句话的最大长度。2)定义位置编码表初始为全0pe torch.zeros(max_len, d_model)假设原来的词向量是(N,L,128),那么我们初始化得到的位置编码表就是(L,128)。下面我们将对位置编码表每个位置进行赋值。3)编码表赋值# 遍历每个位置 for pos in range(max_len): # 遍历每个2i值计算这一对正余弦值 for _2i in range(0, d_model, 2): # 带公式计算 pe[pos, _2i] math.sin( pos / (10000 ** (_2i / d_model)) ) pe[pos, _2i1] math.cos( pos / (10000 ** (_2i / d_model)) )Transformer 模型完全摒弃了 RNN 结构意味着它不再按顺序处理序列而是可以并行处理所有位置的信息。尽管这带来了显著的计算效率提升却也引发了一个问题Transformer 无法像 RNN 那样天然地捕捉词语之间的顺序关系。为了解决这一问题Transformer 引入了一个关键机制——位置编码Positional Encoding。该机制为每个词引入一个表示其位置信息的向量并将其与对应的词向量相加作为模型输入的一部分。这样一来模型在处理每个词时既能获取词义信息也能感知其在句子中的位置从而具备对基本语序的理解能力。最后的位置编码之后的向量如图所示。4前向传播得到位置编码表之后我们需要将它与我们词嵌入的向量进行相加。已知我们需要L个位置的编码向量所以我们直接在编码表里面取然后将其与词嵌入的向量进行广播相加。2.2自定义模型代码如下:class TranslationModel(nn.Module): def __init__(self, input_vocab_size, target_vocab_size, input_padding_idx, target_padding_idx): super(TranslationModel, self).__init__() # 词嵌入层针对src和tgt分别定义 self.src_embedding nn.Embedding(num_embeddingsinput_vocab_size, embedding_dimDIM_MODEL, padding_idxinput_padding_idx) self.tgt_embedding nn.Embedding(num_embeddingstarget_vocab_size, embedding_dimDIM_MODEL, padding_idxtarget_padding_idx) # 位置编码层 self.positional_encoding PositionalEncoding() # Transformer层 #d_model:每个token的向量宽度 #nhead:注意力切成几路 #num_encoder_layers:Encoder 叠几层 #num_decoder_layers:Decoder 叠几层 #batch_first:张量是 (N, 长度, E) self.transformer nn.Transformer( d_modelDIM_MODEL, nheadNUM_HEADS, num_encoder_layersNUM_ENCODER_LAYERS, num_decoder_layersNUM_DECODER_LAYERS, batch_firstTrue, ) # 关闭 nested tensor 快速路径避免 RTX 50 上不稳定的 fused kernel self.transformer.encoder.enable_nested_tensor False self.transformer.encoder.use_nested_tensor False # 线性层 self.fc nn.Linear(in_featuresDIM_MODEL, out_featurestarget_vocab_size) # 前向传播 def forward(self, src, tgt, src_padding_mask, tgt_mask): # 编码包含词嵌入和位置编码 memory self.encode(src, src_padding_mask) # 解码包含线性层输出 output self.decode(tgt, memory, tgt_mask, memory_padding_masksrc_padding_mask) return output # 编码方法得到 编码器的输出 memory def encode(self, src, src_padding_mask): # 1. 词嵌入输入src形状 (N, S)输出形状 (N, S, E)N是一次有几个句子S是原句补齐后有多少个tokenE是每个Token的向量宽度 embed self.src_embedding(src) # 2. 叠加位置编码 input self.positional_encoding(embed) # 3. Transformer编码 memory self.transformer.encoder(srcinput, src_key_padding_masksrc_padding_mask) return memory # 返回形状 (N, S, E) # 解码方法 def decode(self, tgt, memory, tgt_mask, memory_padding_mask): # 1. 词嵌入输入tgt形状 (N, T)输出形状 (N, T, E) embed self.tgt_embedding(tgt) # 2. 叠加位置编码 input self.positional_encoding(embed) # 3. Transformer解码输出形状 (N, T, E) #tgt / inputDecoder 要生成的那一侧例如 sos i love you 的向量。 #memoryencode(src) 的结果整句源序列的表示。 #tgt_mask下三角可见、上三角挡住训练时不能偷看后面的词。 #memory_key_padding_mask一般就是 src_padding_mask告诉 Decoder「源句的 pad 别当信息用」。 output self.transformer.decoder( tgtinput, memorymemory, tgt_masktgt_mask, memory_key_padding_maskmemory_padding_mask, ) # 4. 线性层整合输出 output self.fc(output) return output # 返回形状 (N, T, target_vocab_size)下面进行代码的逐行解释:1)input_vocab_size源语言词表大小。用途建源语言 Embedding表有多少行nn.Embedding(num_embeddingsinput_vocab_size, embedding_dimDIM_MODEL, ...)target_vocab_size目标语言词表大小。用在两处目标语言 EmbeddingDecoder 输入的英文词 → 向量最后的线性层隐向量 → 英文词表上每个词的分数self.fc nn.Linear(in_featuresDIM_MODEL, out_featurestarget_vocab_size)input_padding_idx源语言 PAD 的 id。中文句子补齐时填的那个 token 编号一般是zh_tokenizer.pad_id词表里PAD排在最前经常是 0。target_padding_idx目标语言 PAD 的 id英文补齐用的 PAD 编号en_tokenizer.pad_id。2)对词嵌入层进行初始化因为解码器和编码器都需要词嵌入层所以我们分别要进行初始化。其中分别需要传入的参数是中文词表文件的行数词表大小以及PAD的id。3)对位置编码层进行初始化4)对Transformer进行初始化self.transformer nn.Transformer( d_modelDIM_MODEL, nheadNUM_HEADS, num_encoder_layersNUM_ENCODER_LAYERS, num_decoder_layersNUM_DECODER_LAYERS, batch_firstTrue, )其中需要传入的参数如下:d_model代表每个token的向量维度nhead:代表注意力需要切成几路。num_encoder_layers代表编码器里面的多头注意力机制和全连接的层数num_decoder_layers代表的是解码器的层数。batch_first代表张量的第 0 维是 batch。这样 Encoder / Decoder 吃的和吐的都是(N,L,E)。5)对线性层进行初始化self.fc nn.Linear(in_featuresDIM_MODEL, out_featurestarget_vocab_size)其中DIM_MODEL代表输入的多宽target_vocab_size代表输出有多宽6)前向传播# 前向传播 def forward(self, src, tgt, src_padding_mask, tgt_mask): # 编码包含词嵌入和位置编码 memory self.encode(src, src_padding_mask) # 解码包含线性层输出 output self.decode(tgt, memory, tgt_mask, memory_padding_masksrc_padding_mask) return output其中:src:形状(N, S)的每个格子都是int表示中文词表里的编号。tgt:形状NT)的每个格子都是int表示中文词表里面的编号。src_padding_mask:形状(N, S)代表哪些位置是PADTrue表示这个位置是填充。tgt_mask:英文侧不许看后面的词形状(T, T)下三角能看、上三角挡住。7)编码方法# 编码方法得到 编码器的输出 memory def encode(self, src, src_padding_mask): # 1. 词嵌入输入src形状 (N, S)输出形状 (N, S, E)N是一次有几个句子S是原句补齐后有多少个tokenE是每个Token的向量宽度 embed self.src_embedding(src) # 2. 叠加位置编码 input self.positional_encoding(embed) # 3. Transformer编码 memory self.transformer.encoder(srcinput, src_key_padding_masksrc_padding_mask) return memory # 返回形状 (N, S, E)第一步是词嵌入由原来的(N,S)变为后来的(N,S,E)第二步叠加位置编码第三步使用transformer编码。最后返回(N,S,E)的memory用于后面解码使用。8)解码方法# 解码方法 def decode(self, tgt, memory, tgt_mask, memory_padding_mask): # 1. 词嵌入输入tgt形状 (N, T)输出形状 (N, T, E) embed self.tgt_embedding(tgt) # 2. 叠加位置编码 input self.positional_encoding(embed) # 3. Transformer解码输出形状 (N, T, E) #tgt / inputDecoder 要生成的那一侧例如 sos i love you 的向量。 #memoryencode(src) 的结果整句源序列的表示。 #tgt_mask下三角可见、上三角挡住训练时不能偷看后面的词。 #memory_key_padding_mask一般就是 src_padding_mask告诉 Decoder「源句的 pad 别当信息用」。 output self.transformer.decoder( tgtinput, memorymemory, tgt_masktgt_mask, memory_key_padding_maskmemory_padding_mask, ) # 4. 线性层整合输出 output self.fc(output) return output # 返回形状 (N, T, target_vocab_size)其中词嵌入和位置编码的叠加跟编码器相同。而解码部分需要传入的是输入记忆单元目标掩码记忆掩码这四样东西。其实记忆掩码就是src_padding_mask用来告诉解码器哪些位置是PAD的。2.3解码器和编码器的具体原理其实我们已经掌握了他们的api如何使用但还需要了解具体的工作原理。2.3.1编码器下面先看编码器层内部:每个编码器内部都包括N个层其中每一层由自注意力子层和前馈神经网络组成。1自注意力子层:自注意力机制Self-Attention是 Transformer 编码器的核心结构之一它的作用是在序列内部建立各位置之间的依赖关系使模型能够为每个位置生成融合全局信息的表示。自注意力的计算过程如下:1.生成Q,KV向量。自注意力机制的第一步是将输入序列中的每个位置表示映射为三个不同的向量分别是 查询Query、键Key 和 值Value。Query表示当前词的用于发起注意力匹配的向量Key表示序列中每个位置的内容标识用于与 Query 进行匹配Value表示该位置携带的信息用于加权汇总得到新的表示。具体是怎么计算的呢?2.计算位置相关性完成 Query、Key、Value 向量的生成后模型会使用每个位置的 Query 向量与所有位置的 Key 向量进行相关性评分。评分函数采用向量点积形式。由于在高维空间中点积的数值可能过大会影响 softmax 的稳定性因此在实际计算中对结果进行了缩放。最终的评分函数为是key向量的维度用于缩放点积的幅度。这个分数越大表示第 i 个位置越应该关注第 j 个位置的信息。对于整个序列可以通过矩阵运算一次性计算所有位置之间的评分计算公式如下图所示其中dk是key向量的维度用于缩放点积的幅度。这个分数越大表示第 i 个位置越应该关注第 j 个位置的信息。对于整个序列可以通过矩阵运算一次性计算所有位置之间的评分计算公式如下图所示3.计算注意力权重在得到每个位置与所有位置之间的相关性评分后模型会使用 softmax 函数进行归一化确保每个位置对所有位置的关注程度之和为 1从而形成一个有效的加权分布。对于整个序列模型要做的是对之前得到的注意力评分矩阵的每一行进行softmax归一化。4.加权汇总生成输出最后模型会根据注意力权重对所有位置的 Value 向量进行加权求和得到每个位置融合全局信息后的新表示。综上整个自注意力机制的完整公式如下:2)多头注意力计算过程Transformer 引入了多头注意力机制Multi-Head Attention。其核心思想是通过多组独立的 Query、Key、Value 投影让不同注意力头分别专注于不同的语义关系最后将各头的输出拼接融合。接着进行多头注意力合并多个输出矩阵按维度拼接再乘以W0得到最终的多头注意力的输出。3前馈神经网络层前馈神经网络Feed-Forward Network简称 FFN是 Transformer 编码器中每个子层的重要组成部分紧接在多头注意力子层之后。它通过对每个位置的表示进行逐位置、非线性的特征变换进一步提升模型对复杂语义的建模能力。个标准的 FFN 子层包含两个线性变换和一个非线性激活函数中间通常使用 ReLU激活。其计算公式如下计算过程如下图:4) 残差连接和层归一化在 Transformer 的每个编码器层中每个子层包括自注意力子层和前馈神经网络子层其输出都要经过残差连接Residual Connection和层归一化Layer Normalization处理。这两者是深层神经网络中常用的结构用于缓解模型训练中的梯度消失、收敛困难等问题对于Transformer能够堆叠多层至关重要。11残差链接残差连接Residual Connection也称“跳跃连接”或“捷径连接”最初在计算机视觉领域被提出用于缓解深层神经网络中的梯度消失问题。其核心思想是将子层的输入直接与其输出相加形成一条跨越子层的“捷径”其数学形式为2) 层归一化每个子层在残差连接之后都会进行层归一化Layer Normalization简称 LayerNorm。它的主要作用是规范输入序列中每个token的特征分布某个token的表示可能在不同维度上有较大数值差异提升模型训练的稳定性。该操作会将每个token的向量调整为均值为 0、方差为 1 的规范分布。2.3.2 解码器Transformer 解码器的主要功能是根据编码器的输出逐步生成目标序列中的每一个词。其生成方式采用自回归机制autoregressive每一步的输入由此前已生成的所有词组成模型将输出一个与当前输入长度相同的序列表示。我们只取最后一个位置的输出作为当前步的预测结果。这一过程会不断重复直到生成特殊的结束标记 eos表示序列生成完成。每个Decoder Layer都包含三个子层分别是Masked自注意力子层、编码器-解码器注意力子层Encoder-Decoder Attention和前馈神经网络子层Feed-Forward Network。1)Masked 自注意力子层该子层的主要作用是建模目标序列中当前位置与前文之间的依赖关系为当前词的生成提供上下文语义支持。由于 Transformer 不具备像 RNN 那样的隐藏状态传递机制无法在序列生成过程中保留上下文信息因此在生成每一个词时必须将此前已生成的所有词作为输入通过自注意力机制重新建模上下文关系以预测下一个词。此外从结构上看Transformer 编解码器都具有一个典型特性输入多少个词就输出多少个表示。需要注意的是在推理阶段我们只使用解码器最后一个位置的输出作为当前步的预测结果如下图所示如果训练阶段也完全按照推理流程进行就必须将每个目标序列拆分成多个训练样本每个样本输入一段前文只预测一个词。如下图所示这种方式虽然逻辑合理但训练效率极低完全无法利用 Transformer 并行计算的优势。为提升效率Transformer 采用了并行训练策略一次性输入完整目标序列同时预测每个位置的词。如下图所示但如果不加限制这种方式会让模型在预测每个位置时“看到”后面的词即提前访问未来信息破坏生成任务的因果结构。为解决这个问题解码器在自注意力机制中引入了遮盖机制Mask。该机制会在计算注意力时阻止模型访问当前位置之后的词只允许它依赖自身及前文的信息。这样即使在并行训练时模型也只能像逐词生成一样“看见”它应该看到的内容。Mask 机制的实现非常简单只需将注意力得分矩阵中当前位置对其后续位置的评分设置为 −∞这样在经过 softmax 运算后这些位置的权重会趋近于 0。最终在加权求和时来自未来位置的信息几乎不会参与计算从而实现了“当前词只能看到它前面的词”的约束。2) 编码器-解码器注意力子层该子层的主要作用是建模当前解码位置与源语言序列中各位置之间的依赖关系帮助模型在生成目标词时有效地参考输入内容相当于Seq2Seq模型中的注意力机制。编码器-解码器注意力的核心机制与前面讲过的自注意力机制完全一致区别仅在于Query 来自解码器当前的输入表示即当前生成状态Key和Value 来自编码器的输出表示即整个源序列的上下文。也就是说当前生成位置使用自己的Query去“询问”编码器输出中的哪些位置最相关。注意力机制会根据 Query 与所有 Key 的相似度为每个源位置分配一个权重然后用这些权重对 Value 进行加权求和得到当前生成词所需的上下文信息。三.训练模型def train_one_epoch(model, device, dataloader, loss_fn, optimizer): model.train() total_loss 0 # 按批次迭代训练集数据 for inputs, targets in tqdm(dataloader, desc训练): inputs, targets inputs.to(device), targets.to(device) # 0. 准备数据和掩码 # 从targets中分离解码器真正的input和target decoder_inputs targets[:, :-1] # 去掉eos decoder_targets targets[:, 1:] # 去掉sos # 定义src的填充掩码 src_padding_mask (inputs model.src_embedding.padding_idx) # 定义tgt的掩码 tgt_mask model.transformer.generate_square_subsequent_mask(decoder_inputs.shape[1]).to(device) # 1. 前向传播 # RTX 50 系 fused attention / TF32 GEMM 可能触发 illegal instruction with sdpa_kernel([SDPBackend.MATH]): decoder_outputs model(inputs, decoder_inputs, src_padding_mask, tgt_mask) # 2. 计算损失 loss loss_fn(decoder_outputs.transpose(1, 2), decoder_targets) # 3. 反向传播计算梯度 loss.backward() # 4. 更新参数 optimizer.step() # 5. 梯度清零 optimizer.zero_grad() # 累加损失 total_loss loss.item() return total_loss / len(dataloader)重要步骤如下:传入模型参数计算损失反向传播计算梯度更新参数梯度清零。训练总流程:def train(): # 1. 定义设备 device torch.device(cuda if torch.cuda.is_available() else cpu) print(f使用设备: {device}) if device.type cuda: configure_cuda() print(fGPU: {torch.cuda.get_device_name(0)}) # 2. 构建数据集加载器 train_loader get_dataloader() # 3. 创建分词器 zh_tokenizer ZHTokenizer.from_vocab(MODELS_DIR/ZH_VOCAB_FILE) en_tokenizer ENTokenizer.from_vocab(MODELS_DIR/EN_VOCAB_FILE) # 4. 定义模型 model TranslationModel(input_vocab_sizezh_tokenizer.vocab_size, target_vocab_sizeen_tokenizer.vocab_size, input_padding_idxzh_tokenizer.pad_id, target_padding_idxen_tokenizer.pad_id) model.to(device) # 5. 定义损失函数 loss_fn nn.CrossEntropyLoss(ignore_indexen_tokenizer.pad_id) # 6. 定义优化器 optimizer optim.Adam(model.parameters(), lrLEARNING_RATE, fusedFalse, foreachFalse) # 定义一个写入器 writer SummaryWriter(log_dirLOGS_DIR / time.strftime(%Y-%m-%d_%H-%M-%S)) # 7. 开始训练 min_loss float(inf) for epoch in range(EPOCHS): print(fEpoch {epoch1}/{EPOCHS}: ) train_loss train_one_epoch(modelmodel, devicedevice, dataloadertrain_loader, loss_fnloss_fn, optimizeroptimizer) print(fTrain Loss: {train_loss}) writer.add_scalar(Loss/train, train_loss, epoch) # 判断是否需要保存模型 if train_loss min_loss: state_dict model.state_dict() torch.save(state_dict, MODELS_DIR/BEST_MODEL) min_loss train_loss print(模型保存成功) writer.close()1.定义设备2.构建数据加载器3.创建分词器中文分词器是继承于自定义分词器定义方法如下:class BaseTokenizer(): # 定义类属性 unk_token UNK_TOKEN pad_token PAD_TOKEN start_token START_TOKEN end_token END_TOKEN # 初始化基于词表构建分词器对象 def __init__(self, vocab_list): self.vocab_size len(vocab_list) self.id2word vocab_list self.word2id {word:id for id, word in enumerate(vocab_list)} # 定义特殊token及其id # self.unk_token UNK_TOKEN self.unk_id self.word2id[self.unk_token] self.pad_id self.word2id[self.pad_token] self.start_id self.word2id[self.start_token] self.end_id self.word2id[self.end_token] # 定义工厂方法从文件中加载词表创建一个分词器对象 classmethod def from_vocab(cls, vocab_file_path): # 打开文件获取词表 with open( vocab_file_path, r, encodingutf-8) as f: vocab_list [line.strip() for line in f.readlines()] return cls(vocab_list) # 根据语料句子创建词表并保存为文件 classmethod def build_vocab(cls, vocab_file_path, sentences): vocab_set set() for sentence in tqdm(sentences, desc构建词表): vocab_set.update(cls.tokenize(sentence)) # 增加特殊token填充词未登录词 id2word [cls.pad_token, cls.unk_token, cls.start_token, cls.end_token] list(vocab_set) print(f词表大小{len(id2word)}) # 保存词表 with open(vocab_file_path, w, encodingutf-8) as f: f.write(\n.join(id2word)) # 分词方法类方法 classmethod def tokenize(cls, text): pass # 编码方法id化传入是否添加开始结束标记的判断 def encode(self, text, markFalse): tokens self.tokenize(text) # 如果是目标序列文本就加入sos和eos特殊标记 if mark: tokens [self.start_token] tokens [self.end_token] ids [ self.word2id.get(token, self.unk_id) for token in tokens ] return ids# 中文分词器 class ZHTokenizer(BaseTokenizer): classmethod def tokenize(cls, text): return list(text)# 英文分词器 class ENTokenizer(BaseTokenizer): tokenizer TreebankWordTokenizer() detokenizer TreebankWordDetokenizer() classmethod def tokenize(cls, text): return cls.tokenizer.tokenize(text) # 解码方法传入id列表得到完整的英文句子 def decode(self, ids): tokens [ self.id2word[id] for id in ids ] return self.detokenizer.detokenize(tokens)4.定义模型model TranslationModel(input_vocab_sizezh_tokenizer.vocab_size, target_vocab_sizeen_tokenizer.vocab_size, input_padding_idxzh_tokenizer.pad_id, target_padding_idxen_tokenizer.pad_id) model.to(device)5.定义损失函数# 5. 定义损失函数 loss_fn nn.CrossEntropyLoss(ignore_indexen_tokenizer.pad_id)6.定义优化器# 6. 定义优化器 optimizer optim.Adam(model.parameters(), lrLEARNING_RATE, fusedFalse, foreachFalse)7.定义写入器# 定义一个写入器 writer SummaryWriter(log_dirLOGS_DIR / time.strftime(%Y-%m-%d_%H-%M-%S))8.开始训练min_loss float(inf) for epoch in range(EPOCHS): print(fEpoch {epoch1}/{EPOCHS}: ) train_loss train_one_epoch(modelmodel, devicedevice, dataloadertrain_loader, loss_fnloss_fn, optimizeroptimizer) print(fTrain Loss: {train_loss}) writer.add_scalar(Loss/train, train_loss, epoch) # 判断是否需要保存模型 if train_loss min_loss: state_dict model.state_dict() torch.save(state_dict, MODELS_DIR/BEST_MODEL) min_loss train_loss print(模型保存成功) writer.close()四. 预测# 调用预测逻辑的应用程序 def run_predict(): # 1. 定义设备 device torch.device(cuda if torch.cuda.is_available() else cpu) # 2. 创建分词器 zh_tokenizer ZHTokenizer.from_vocab(MODELS_DIR/ZH_VOCAB_FILE) en_tokenizer ENTokenizer.from_vocab(MODELS_DIR/EN_VOCAB_FILE) print(词表加载成功) # 3. 加载模型 model TranslationModel(input_vocab_sizezh_tokenizer.vocab_size, input_padding_idxzh_tokenizer.pad_id, target_vocab_sizeen_tokenizer.vocab_size, target_padding_idxen_tokenizer.pad_id).to(device) state_dict torch.load(MODELS_DIR/BEST_MODEL) model.load_state_dict(state_dict) print(模型加载成功) print(欢迎使用中英翻译模型输入 quit 或者 q 退出...) # 核心是一个死循环 while True: # 等待用户输入 user_input_text input(中文) # 如果是 quit 或者 q就退出 if user_input_text in [q, quit]: print(欢迎下次使用) break # 如果输入为空字符继续输入 if user_input_text.strip() : print(请输入有效内容...) continue # 调用预测逻辑得到预测结果 result predict(user_input_text, modelmodel, input_tokenizerzh_tokenizer, target_tokenizeren_tokenizer, devicedevice) print(英文, result) if __name__ __main__: run_predict()当用户输入之后系统会调用predict函数函数内部如下:# 预测函数传入一串文本得到接下来最可能的5个词 def predict(text, model, input_tokenizer, target_tokenizer, device): # 1. 处理文本得到模型输入 # 1.1 分词、id化 ids input_tokenizer.encode(text) # 1.2 转换为Tensor形状 (N1, L) inputs torch.tensor([ids]).to(device) # 2. 调用核心预测逻辑 result predict_batch(model, inputs, target_tokenizer, device) # 3. 解码得到目标译文句子 sentence target_tokenizer.decode(result[0]) return sentence接着我们看predict_batchdef predict_batch(model, inputs, tokenizer, device): model.eval() with torch.no_grad(): # Transformer 的 encode 需要 padding mask只返回 memory (N, S, E) src_padding_mask (inputs model.src_embedding.padding_idx) with sdpa_kernel([SDPBackend.MATH]): memory model.encode(inputs, src_padding_mask) # 解码从 SOS 开始之后每步把「已经生成的整段」一起送进去 batch_size inputs.shape[0] decoder_input torch.full((batch_size, 1), tokenizer.start_id, devicedevice) generated [] is_finished torch.zeros(batch_size, dtypetorch.bool, devicedevice) for _ in range(SEQ_LEN): tgt_mask model.transformer.generate_square_subsequent_mask( decoder_input.shape[1] ).to(device) with sdpa_kernel([SDPBackend.MATH]): decoder_output model.decode( decoder_input, memory, tgt_mask, memory_padding_masksrc_padding_mask ) # 只取最后一个时间步的词表分数贪心选下一个词 (N, 1) next_token_ids torch.argmax(decoder_output[:, -1:, :], dim-1) generated.append(next_token_ids) # Transformer 没有逐步 hidden要把新词拼到已生成序列后面 decoder_input torch.cat([decoder_input, next_token_ids], dim1) is_finished | (next_token_ids.squeeze(1) tokenizer.end_id) if bool(is_finished.all()): break五.结语代码存放在:​​​​​​​gzh2219/transformer-
返回列表