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

如何解决Keras计算两个Ragged Tensor的KL散度时的报错?

解决Ragged Tensor计算KL散度的报错问题

你定义了两个Ragged Tensor分别存储预测分布与真实分布,尝试用Keras内置的KL散度损失函数计算逐行散度并得到整体损失时,出现如下报错:

ValueError: TypeError: object of type 'RaggedTensor' has no len()

原代码重现

import tensorflow as tf

# Tensor 1
pred_score = tf.ragged.constant([
    [[-0.51760715], [-0.18927467], [-0.10698503]],
    [[-0.58782816], [-0.13076714], [-0.04999146], [-0.1772059], [-0.14299354]]
])
pred_score = tf.squeeze(pred_score, axis=-1)
pred_score_dist = tf.nn.softmax(pred_score, axis=-1)
print(pred_score_dist)
print(pred_score_dist.shape)

# Tensor 2
actual_score = tf.ragged.constant([
    [3.0, 2.0, 2.0], 
    [3.0, 3.0, 1.0, 1.0, 0.0]
])
actual_score_dist = tf.nn.softmax(actual_score, axis=-1)
print(actual_score_dist)
print(actual_score_dist.shape)

# 报错的损失计算代码
loss = tf.keras.losses.KLDivergence()
batch_loss = loss(actual_score_dist, pred_score_dist)

问题原因

Keras内置的KLDivergence损失函数并未原生支持Ragged Tensor,无法处理每行长度不一致的张量结构,因此抛出类型错误。

解决方案

手动实现KL散度的逐行计算,利用TensorFlow对Ragged Tensor的原生支持完成运算,同时添加数值稳定性处理避免log(0)错误:

import tensorflow as tf

# 保留原有的Ragged Tensor定义与分布计算
pred_score = tf.ragged.constant([
    [[-0.51760715], [-0.18927467], [-0.10698503]],
    [[-0.58782816], [-0.13076714], [-0.04999146], [-0.1772059], [-0.14299354]]
])
pred_score = tf.squeeze(pred_score, axis=-1)
pred_score_dist = tf.nn.softmax(pred_score, axis=-1)

actual_score = tf.ragged.constant([
    [3.0, 2.0, 2.0], 
    [3.0, 3.0, 1.0, 1.0, 0.0]
])
actual_score_dist = tf.nn.softmax(actual_score, axis=-1)

# 手动计算KL散度
epsilon = 1e-10  # 防止log(0)的数值稳定项
# 逐行计算KL散度:sum(y_true * log(y_true / y_pred))
kl_per_sample = tf.reduce_sum(
    actual_score_dist * tf.math.log((actual_score_dist + epsilon) / (pred_score_dist + epsilon)),
    axis=-1
)
# 计算整体批次损失(可根据需求选择mean或sum)
batch_loss = tf.reduce_mean(kl_per_sample)

print("逐样本KL散度:", kl_per_sample)
print("整体批次损失:", batch_loss)

说明

  • tf.reduce_sum在axis=-1时会自动对Ragged Tensor的每行独立求和,适配不同长度的行结构。
  • 添加epsilon是为了避免因分布值为0导致的log(0)数值错误,保证计算稳定性。
  • 最终的整体损失可根据业务需求选择tf.reduce_mean(平均损失)或tf.reduce_sum(总损失)。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.30 10:52:51