You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何重塑输入数据适配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

关键修改说明

  1. GRU维度适配:开启batch_first=True,让输入输出维度顺序贴合你的批量数据格式。
  2. 新增token嵌入层:解决“42维token转256维GRU输入”的维度匹配问题,保证自回归生成的循环逻辑成立。
  3. 分阶段逻辑:
    • 训练时用teacher forcing:直接喂入目标序列的前127个token,一次性完成GRU计算,提升训练效率与稳定性。
    • 推理时用自回归生成:从初始输入开始,每一步生成的token作为下一轮输入,循环128次得到完整序列。
  4. 初始输入处理:通过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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.19 06:54:55