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

Keras实现Generalized Dice Loss时梯度为None问题求助

解决Keras中Generalized Dice Loss的梯度None问题

我之前也碰到过一模一样的问题,你的错误根源在于代码里依赖了稀疏张量的操作(tf.sparse_reduce_sum),这类操作在TensorFlow的自动求导机制中通常没有定义反向传播逻辑,再加上几个小细节没处理好,才导致了这个报错。下面是具体的排查和解决方法:

问题分析

  1. 稀疏张量操作无梯度:tf.sparse_reduce_sum是针对稀疏张量设计的,而稀疏张量的大部分操作在TensorFlow的自动求导体系中没有实现梯度逻辑,这是报错的核心原因。
  2. Logits未转概率:你的函数注释里说明prediction是logits,但Generalized Dice Loss需要基于概率图计算,直接用logits会导致数值不稳定,也可能间接影响梯度计算。
  3. 自定义One-Hot转换的潜在风险:如果labels_to_one_hot函数内部用到了tf.argmax、tf.round这类不可微分操作,也会触发梯度None的问题。

修改后的完整代码

import tensorflow as tf
from tensorflow.keras import backend as K

def generalised_dice_loss(prediction, ground_truth, weight_map=None, type_weight='Square'): 
    """ 
    Function to calculate the Generalised Dice Loss defined in Sudre, C. et. al. (2017) 
    Generalised Dice overlap as a deep learning loss function for highly unbalanced segmentations. DLMIA 2017 
    :param prediction: the logits 
    :param ground_truth: the segmentation ground truth (integer labels)
    :param weight_map: optional weight map for each pixel/voxel
    :param type_weight: type of weighting allowed between labels:
                        Square (square of inverse of volume), 
                        Simple (inverse of volume), 
                        Uniform (no weighting)
    :return: the computed loss
    """ 
    # Convert logits to normalized probabilities (critical for Dice calculation)
    prediction = tf.nn.softmax(prediction)
    prediction = tf.cast(prediction, tf.float32) 
    
    # Adjust ground truth shape to match prediction's spatial dimensions
    if len(ground_truth.shape) == len(prediction.shape): 
        ground_truth = ground_truth[..., -1] 
    # Convert integer ground truth to dense one-hot tensor (avoid sparse ops)
    one_hot = tf.one_hot(tf.cast(ground_truth, tf.int32), depth=tf.shape(prediction)[-1])
    
    if weight_map is not None: 
        n_classes = prediction.shape[-1]
        # Expand weight map to match number of classes and tile accordingly
        weight_map_nclasses = tf.tile(tf.expand_dims(weight_map, axis=-1), [1, 1, n_classes])
        # Calculate volumes using dense tensor summation (supports gradient)
        # Adjust axes based on your input shape: e.g., (batch, H, W, C) uses [0,1,2]
        ref_vol = tf.reduce_sum(weight_map_nclasses * one_hot, axis=[0, 1, 2])
        intersect = tf.reduce_sum(weight_map_nclasses * one_hot * prediction, axis=[0, 1, 2])
        seg_vol = tf.reduce_sum(weight_map_nclasses * prediction, axis=[0, 1, 2])
    else: 
        # No weight map: sum over spatial and batch dimensions
        ref_vol = tf.reduce_sum(one_hot, axis=[0, 1, 2])
        intersect = tf.reduce_sum(one_hot * prediction, axis=[0, 1, 2])
        seg_vol = tf.reduce_sum(prediction, axis=[0, 1, 2])
    
    # Compute class weights based on volume
    if type_weight == 'Square': 
        weights = tf.reciprocal(tf.square(ref_vol)) 
    elif type_weight == 'Simple': 
        weights = tf.reciprocal(ref_vol) 
    elif type_weight == 'Uniform': 
        weights = tf.ones_like(ref_vol) 
    else: 
        raise ValueError(f"The variable type_weight \"{type_weight}\" is not defined.")
    
    # Handle cases where reference volume is 0 (avoids infinite weights)
    new_weights = tf.where(tf.math.is_inf(weights), tf.zeros_like(weights), weights) 
    max_valid_weight = tf.reduce_max(new_weights)
    weights = tf.where(tf.math.is_inf(weights), tf.ones_like(weights) * max_valid_weight, weights) 
    
    # Calculate Generalised Dice Loss
    generalised_dice_numerator = 2 * tf.reduce_sum(tf.multiply(weights, intersect)) 
    # Add small epsilon to avoid division by zero
    generalised_dice_denominator = tf.reduce_sum(tf.multiply(weights, seg_vol + ref_vol)) + 1e-6 
    generalised_dice_score = generalised_dice_numerator / generalised_dice_denominator 
    return 1 - generalised_dice_score 

关键修改点说明

  1. 替换稀疏操作为稠密操作:把所有tf.sparse_reduce_sum换成tf.reduce_sum,用稠密张量完成所有计算,确保梯度可以正常传播。
  2. 添加Softmax转换:将输入的logits转换成概率分布,这是Dice Loss计算的标准要求,也保证了数值稳定性。
  3. 使用原生tf.one_hot:替代自定义的labels_to_one_hot,确保生成的是可微分的稠密张量,避免潜在的不可微分操作。
  4. 调整Weight Map处理:用tf.expand_dims和tf.tile更安全地扩展weight map的维度,避免形状不匹配的问题。
  5. 更新Inf值处理:使用TensorFlow推荐的tf.math.is_inf替代旧版tf.is_inf,提升代码兼容性。

验证梯度是否正常

你可以用以下代码快速验证修改后的损失函数是否能生成有效梯度:

# Create test tensors matching your model's input/output shape
batch_size, height, width, num_classes = 2, 32, 32, 3
logits = tf.random.normal((batch_size, height, width, num_classes))
ground_truth = tf.random.uniform((batch_size, height, width), maxval=num_classes, dtype=tf.int32)

# Compute loss and check gradients
loss = generalised_dice_loss(logits, ground_truth)
grads = K.gradients(loss, logits)

print("Gradient shape:", grads[0].shape if grads else "None")
# 如果输出梯度形状而不是None,说明梯度计算正常

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.29 08:58:10