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

如何使用自定义Cauchy-Schwarz散度损失函数训练Keras模型并解决报错

问题解决方案

错误根本原因

  • Keras自定义损失函数接收的输入p1(对应y_true标签)、p2(对应y_pred预测值)都是TensorFlow张量,而非numpy数组。TensorFlow构建静态计算图时,batch维度的大小p1.shape[0]是动态值,静态阶段返回None,直接传给range()就会触发类型错误。
  • 损失函数中使用的scipy.stats.gaussian_kde、Python原生sum、math.sqrt等都是numpy生态的操作,完全脱离TensorFlow计算图,无法自动求导,就算解决了shape问题也没法正常反向传播训练模型。

解决方法

方案1:改写为纯TensorFlow实现(推荐)

全部使用TensorFlow原生算子实现损失逻辑,保证计算图完整、支持自动求导。针对多分类one-hot标签场景,可直接基于类别概率向量计算CS散度,无需额外做核密度估计,示例代码如下:

import tensorflow as tf

def cs_divergence(y_true, y_pred):
    # 对每个样本的类别维度计算对应值
    numerator = tf.reduce_sum(y_true * y_pred, axis=-1)
    denominator = tf.sqrt(tf.reduce_sum(tf.square(y_true), axis=-1) * tf.reduce_sum(tf.square(y_pred), axis=-1))
    # 数值裁剪,防止除0、log输入为0导致的NaN问题
    denominator = tf.clip_by_value(denominator, 1e-7, 1e7)
    cos_similarity = tf.clip_by_value(numerator / denominator, 1e-7, 1.0 - 1e-7)
    return -tf.math.log(cos_similarity)

# 模型编译部分保持逻辑不变,调整SGD调用适配TF2.4写法
sgd = tf.keras.optimizers.SGD(learning_rate=0.0001, decay=1e-6, momentum=0.9, nesterov=True) 
model.compile(optimizer=sgd,
              loss=cs_divergence, 
              metrics=['accuracy'])

方案2:用tf.py_function包裹原有逻辑(不推荐)

如果业务必须使用scipy的核密度估计逻辑,可以用tf.py_function将Python函数包裹进计算图,但会大幅降低训练速度,且无法保证梯度传递正常:

import tensorflow as tf
from math import sqrt
from math import log
from scipy.stats import gaussian_kde

def cs_divergence_wrapper(p1, p2):
    p1_np = p1.numpy()
    p2_np = p2.numpy()
    r = range(0, p1_np.shape[0])
    p1_kernel = gaussian_kde(p1_np)
    p2_kernel = gaussian_kde(p2_np)
    p1_computed = p1_kernel(r)
    p2_computed = p2_kernel(r)
    numerator = sum(p1_computed * p2_computed)
    denominator = sqrt(sum(p1_computed ** 2) * sum(p2_computed**2))
    return -log(numerator/denominator)

def cs_divergence(p1, p2):
    loss = tf.py_function(func=cs_divergence_wrapper, inp=[p1, p2], Tout=tf.float32)
    loss.set_shape(())
    return loss

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.06 09:36:03