如何重塑输入数据适配PyTorch GRU,实现分子SELFIES字符串生成
问题分析与解决方案
首先明确几个核心概念,帮你理清GRU的维度逻辑:
- PyTorch的GRU默认输入格式是
[seq_len, batch_size, input_size],设置batch_first=True后会切换为你预期的[batch_size, seq_len, input_size],更适配批量数据处理。 - 你的需求属于自回归序列生成:用初始256维分子编码启动GRU,每一步生成的token(转换为256维embedding后)作为下一轮GRU的输入,循环128次得到完整的SELFIES序列。
你的现有代码仅实现了单步GRU计算,未做循环生成逻辑,因此只能得到单个token。以下是修正后的完整方案:
修正后的DecoderNet代码
import torch import torch.nn as nn class DecoderNet(nn.Module): def __init__(self, input_size, hidden_size, num_layers, output_size, output_len): super(DecoderNet, self).__init__() self.input_size = input_size # 256 self.hidden_size = hidden_size # 256 self.num_layers = num_layers # 1 self.output_size = output_size # 42 self.output_len = output_len # 128 # 开启batch_first,适配[batch, seq_len, input]格式 self.gru = nn.GRU(input_size=input_size, hidden_size=hidden_size, num_layers=num_layers, batch_first=True) # 将GRU输出映射为token概率 self.fc = nn.Linear(hidden_size, output_size) # 将token索引转换为256维embedding,作为下一轮GRU输入 self.token_embedding = nn.Embedding(output_size, input_size) self.softmax = nn.Softmax(dim=2) def forward(self, x, h=None, target_seq=None): """ x: 初始输入,形状[batch_size, input_size] -> [64,256] h: 初始隐藏状态,可选,默认自动初始化 target_seq: 训练时的目标序列,形状[batch_size, output_len],用于teacher forcing 返回:生成的序列概率,形状[batch_size, output_len, output_size] """ batch_size = x.size(0) # 初始化隐藏状态(GRU的隐藏状态格式固定为[num_layers, batch_size, hidden_size]) if h is None: h = self.init_hidden(batch_size) # 初始输入转为[batch, 1, input_size],适配GRU单步输入要求 current_input = x.unsqueeze(1) # [64,1,256] if target_seq is not None: # 训练阶段:使用teacher forcing,直接喂入目标序列的前output_len-1个token target_embeds = self.token_embedding(target_seq[:, :-1]) # [64,127,256] # 拼接初始输入与目标embeds,得到完整输入序列 full_input = torch.cat([current_input, target_embeds], dim=1) # [64,128,256] # 一次性完成GRU计算,提升训练效率 out, h = self.gru(full_input, h) outputs = self.fc(out) # [64,128,42] else: # 推理阶段:自回归生成,循环output_len次 outputs = [] for _ in range(self.output_len): out, h = self.gru(current_input, h) # out: [64,1,256] token_logits = self.fc(out) # [64,1,42] outputs.append(token_logits) # 选择下一个token(这里用argmax,也可改用随机采样提升多样性) next_token = torch.argmax(token_logits, dim=2) # [64,1] # 将token转为embedding,作为下一轮输入 current_input = self.token_embedding(next_token) # [64,1,256] # 拼接所有步骤的输出 outputs = torch.cat(outputs, dim=1) # [64,128,42] outputs = self.softmax(outputs) return outputs, h def init_hidden(self, batch_size): # 隐藏状态与GRU参数同设备 h0 = torch.zeros(self.num_layers, batch_size, self.hidden_size, device=self.gru.weight.device) return h0
关键修改说明
- GRU维度适配:开启
batch_first=True,让输入输出维度顺序贴合你的批量数据格式。 - 新增token嵌入层:解决“42维token转256维GRU输入”的维度匹配问题,保证自回归生成的循环逻辑成立。
- 分阶段逻辑:
- 训练时用teacher forcing:直接喂入目标序列的前127个token,一次性完成GRU计算,提升训练效率与稳定性。
- 推理时用自回归生成:从初始输入开始,每一步生成的token作为下一轮输入,循环128次得到完整序列。
- 初始输入处理:通过
unsqueeze(1)将[64,256]转为[64,1,256],符合GRU的单步输入要求。
使用示例
# 初始化模型 decoder = DecoderNet(input_size=256, hidden_size=256, num_layers=1, output_size=42, output_len=128) # 模拟输入:64个256维分子编码向量 batch_input = torch.randn(64, 256) # 推理模式生成序列 gen_seq_probs, _ = decoder(batch_input) print(gen_seq_probs.shape) # 输出: torch.Size([64, 128, 42]) # 训练模式(假设已有目标序列) target_seq = torch.randint(0, 42, (64, 128)) # 模拟128长度的目标token序列 train_output, _ = decoder(batch_input, target_seq=target_seq)
内容的提问来源于stack exchange,提问作者Hubert Rybka
相关产品推荐
相关产品推荐

