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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.24 16:45:08