如何在PyTorch中为NNI的DARTS网络结构搜索实现MSE评估指标
代码问题修正与实现方案
你自行编写的MSE计算函数存在3个需要调整的点:
- 没必要对output做转置操作:默认回归任务中,模型输出
output和标签target的维度完全匹配,随意转置会触发PyTorch广播机制,导致计算出的误差不符合预期 - 未适配多输出回归场景:如果你的任务是多维度回归,现有实现仅除以batch_size,会比标准MSE结果大「输出维度数」倍
- 未将张量转为Python数值:原accuracy返回的是用
.item()提取的纯数值,你当前返回的是PyTorch张量,可能导致日志记录异常、额外占用显存
标准实现代码
推荐直接调用PyTorch官方实现的MSE接口,避免手动计算出错:
import torch import torch.nn.functional as F def accuracy_mse(output, target): mse_val = F.mse_loss(output, target) res = dict() res["acc_mse"] = mse_val.item() return res
如果需要手动实现对应逻辑,可参考下方写法,和官方接口计算结果完全一致:
def accuracy_mse(output, target): # 计算所有元素的平方误差平均 total_elements = target.numel() diff = torch.square(output - target).sum() / total_elements res = dict() res["acc_mse"] = diff.item() return res
DartsTrainer参数替换
将原有Trainer中的metrics参数替换为以下内容即可:
metrics=lambda output, target: accuracy_mse(output, target),
额外提示:回归任务需要同步把损失函数criterion替换为torch.nn.MSELoss(),和评估指标保持匹配
内容的提问来源于stack exchange,提问作者pittnerf
相关产品推荐
相关产品推荐

