如何在AllenNLP自编码器中区分训练验证阶段实现正确校验
实现方法
你要的逻辑可以直接借助PyTorch模型自带的属性实现,AllenNLP的模型默认继承torch.nn.Module,内置了布尔类型的self.training属性,该属性会被AllenNLP的训练器自动维护:训练阶段自动设置为True,验证、测试阶段自动设置为False,无需你手动传参或者维护状态。
你只需要把你预想的逻辑中的training判断替换为self.training即可,修改后的代码如下:
embedded = self._embedder(text) if labels is not None: encoded = self._encoder(labels) # 直接用self.training区分训练/验证状态 if self.training: decoded = self._decoder(encoded) else: decoded = self._decoder(embedded) # compute loss / accuracy encoder_loss = MSE(embedded, encoded) reconstruction_loss = CDL(labels, decoded) else: decoded = self._decoder(embedded)
逻辑说明
这个实现完全符合你的校验要求:
- 训练阶段解码用标签编码器的输出,保证损失回传能同时优化标签编码器、解码器、输入嵌入模块三个部分
- 验证阶段解码用输入嵌入模块的输出,完全不会用到标签编码器的输出做解码,不会出现标签泄露的问题,同时还是会计算
encoder_loss,可以额外监控输入嵌入模块和标签编码器的拟合效果
如果你需要手动调用模型做推理,只需要先执行model.eval()即可自动切换到验证逻辑。
内容的提问来源于stack exchange,提问作者David Waterworth
相关产品推荐
相关产品推荐

