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

如何在自定义损失函数中兼容pixelClassificationLayer的Classes与ClassWeights

在自定义损失函数中集成Classes和ClassWeights参数

处理ClassWeights参数

要在自定义损失函数中加入类别权重,你需要修改损失函数定义,将classWeights作为额外参数传入,并在crossentropy函数中指定Weights选项来应用权重。

首先更新损失函数:

function loss = modelLoss(Y, T, classWeights)
    % 处理NaN掩码,忽略无效标签
    mask = ~isnan(T);
    T(isnan(T)) = 0; % 将NaN标签置为0,避免计算错误
    
    % 带类别权重的交叉熵计算
    loss = crossentropy(Y, T, ...
        Mask=mask, ...
        NormalizationFactor="mask-included", ...
        Weights=classWeights);
end

调用trainnet时,通过匿名函数将classWeights传递给损失函数:

% 假设已定义好classNames和classWeights
netTrained = trainnet(images, net, @(Y,T) modelLoss(Y,T,classWeights), options);

处理Classes参数

Classes参数的核心作用是建立标签与网络输出通道的映射关系,确保标签索引和网络输出的每个通道一一对应。这个逻辑不需要在损失函数中处理,而是要在数据预处理阶段完成:

  1. 如果原始标签是字符串形式,先转换为分类变量再转为数值索引:
% 假设T_strings是原始的字符串标签数组
T_categorical = categorical(T_strings, classNames); % 按classNames指定顺序映射类别
T = double(T_categorical); % 转换为1~N的数值索引(N为类别数)
  1. 确保网络输出层的通道数等于numel(classNames),这样Y的每个通道就对应classNames中的一个类别,损失函数就能正确匹配标签和输出。

关键说明

  • classWeights直接作用于损失计算,通过crossentropy的Weights参数实现类别平衡;
  • Classes主要是规范标签的索引规则,保证训练数据的标签和网络输出的维度、语义对应,属于数据层面的准备工作,与损失函数计算逻辑无关。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.15 06:13:20