如何结合有效性校验函数对抗训练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
相关产品推荐
相关产品推荐

