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

如何结合有效性校验函数对抗训练Pytorch LSTM分子生成模型

适配的损失函数

你当前场景下校验函数是不可导的黑盒逻辑,无法直接将梯度回传给LSTM生成模型,最适配的是基于REINFORCE算法的策略梯度损失,不需要额外训练GAN的判别器,直接把你现有的check_validity输出作为奖励信号引导生成器优化。
损失的核心逻辑为:
Loss = - 平均( 调整后奖励 * 生成对应分子序列的对数概率 )
其中有效分子奖励为1,无效为0,减去批次平均奖励作为基线可以降低梯度估计方差,加快收敛。

Pytorch训练循环集成

首先需要你的生成模型在生成序列的同时,返回对应序列的总对数概率(记录梯度路径,不能断开计算图),另外需要实现将模型输出的token id序列转为SMILES分子字符串的转换逻辑,完整训练循环示例如下:

import torch
import torch.optim as optim

# 初始化生成模型、优化器,可根据自己的需求调整学习率
G = YourLSTMGenerator()
optimizer = optim.Adam(G.parameters(), lr=1e-4)
num_epochs = 100
batch_size = 32

for epoch in range(num_epochs):
    G.train()
    optimizer.zero_grad()

    # 1. 生成批量分子序列,返回值generated_seqs形状为[batch_size, 序列长度],log_probs形状为[batch_size],对应每个序列的总对数概率
    generated_seqs, log_probs = G.generate(batch_size=batch_size)

    # 2. 批量校验分子有效性,获取奖励
    rewards = []
    for seq in generated_seqs:
        # token_to_smiles为你自己实现的、将token id序列转为SMILES字符串的函数
        mol_str = token_to_smiles(seq.cpu().tolist())
        rewards.append(check_validity(mol_str))
    rewards = torch.tensor(rewards, dtype=torch.float32, device=log_probs.device)

    # 3. 奖励减去基线,降低梯度方差
    adjusted_rewards = rewards - rewards.mean()

    # 4. 计算损失并反向传播更新参数
    loss = - (adjusted_rewards * log_probs).mean()
    loss.backward()
    # 梯度裁剪,避免LSTM梯度爆炸
    torch.nn.utils.clip_grad_norm_(G.parameters(), max_norm=5.0)
    optimizer.step()

    # 训练日志打印
    if (epoch + 1) % 10 == 0:
        valid_rate = rewards.mean().item()
        print(f"Epoch {epoch+1} | 损失: {loss.item():.4f} | 有效分子占比: {valid_rate:.4f}")

注意事项

  • 建议先使用公开的有效SMILES数据集对生成器做交叉熵预训练,让模型先学会基本的SMILES语法规则,再接入上述对抗训练逻辑微调,避免训练初期有效率为0导致无梯度信号无法收敛。
  • 若效果不达预期,可以优化奖励逻辑,不要仅返回0/1,比如对存在部分语法错误但接近有效的分子给出中间值奖励,收敛速度会更快。
  • 生成序列的log_probs必须保留完整的计算图,不能在生成过程中转numpy或者detach,否则梯度无法回传更新模型参数。

内容的提问来源于stack exchange,提问作者mac179

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.26 04:06:03