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

Keras自定义损失函数结合sample_weight使用的相关疑问

关于Keras自定义损失使用sample_weight的解答

结论

你当前的实现方式完全可以直接通过model.fit的sample_weight参数传入样本权重,不需要修改损失类结构,也不需要额外将权重作为网络输入传入。

原理说明

  • Keras的样本加权逻辑是在Loss基类的__call__方法中自动处理的,你自定义损失时实现的call方法只需要负责输出每个样本的独立损失即可,不需要主动接收和处理sample_weight参数。
  • 官方文档说明的「返回对应批次每个样本损失数组的损失函数自动支持样本加权」完全适用于你的场景:只要call方法返回的张量形状为(batch_size,)或者可被压缩为该形状的(batch_size, 1),框架就会自动将每个样本的损失和对应权重相乘,再完成后续的损失聚合操作。
  • 你之前看到的「将权重作为网络输入传入」的方案,仅适用于需要在损失计算逻辑内部对权重做特殊自定义处理的场景,普通的样本加权需求用内置的sample_weight机制即可满足。

代码校验说明

你当前实现的NegLogLikMixedGaussian损失类完全符合要求:

  • call方法返回的-tf.reduce_mean(log_likelihood, axis=-1)形状为(batch_size,),满足单样本损失的输出要求。
  • 你写的model.compile和model.fit调用方式都是正确的,直接传入sample_weight参数即可生效。

测试注意事项

你测试时使用的np.ones(len(y_train)) / len(dh.y_train_scaled)权重,和无权重的训练梯度方向完全一致,仅损失的绝对数值会缩小为原来的1/样本数,如果需要损失数值和无权重场景完全对齐,测试时将权重设置为全1即可。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.28 08:45:06