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

构建ICtCP色彩转换模型时损失异常飙升的问题排查

色彩转换模型在ICtCP空间训练时损失值异常飙升至inf/nan问题

我有一个包含20万+色块的数据集,采集自两种不同介质,正在基于此构建色彩转换模型。最初采用直接RGB-to-RGB的输入输出方式,神经网络表现尚可,但为了更好地处理亮度和色彩对比度的转换关系,尝试改用亮度-色度空间。

一开始使用CIELAB和YCbCr空间,但数据集是存储在对数容器中的HDR场景数据,这两种空间不适合HDR场景表示,转换结果不准确。于是改用Dolby的ICtCP空间(基于无界线性场景信息构建)。已完成数据集到该空间的转换,确认输出和数组结构正确,但将变量输入网络后,损失值先飙升至天文数字,随后变为inf或nan,无法定位问题所在。

已做排查

  • 使用colour-science库进行内部色彩转换
  • 测试过针对ICtCP空间的自定义损失函数以及TensorFlow内置的MSE损失,两者均产生极端损失值
  • 将RGB和ICtCP值输出到文本文件验证:RGB值处于0-1区间,ICtCP值中I分量为0-1,Ct和Cp分量为-1-1,无超出范围的值

色彩转换函数

#Davinci Wide Gamut Intermediate to Dolby ICtCP HDR opponent space
def DWG_TO_ITP(rgb_values):
    cs = colour.models.RGB_COLOURSPACE_DAVINCI_WIDE_GAMUT
    
    #DWG DI to XYZ Linear
    xyzLin = colour.RGB_to_XYZ(rgb_values, cs.whitepoint, cs.whitepoint, cs.matrix_RGB_to_XYZ, cctf_decoding=cs.cctf_decoding)
    
    #XYZ to ICtCp
    ictcp = colour.XYZ_to_ICtCp(xyzLin)
    
    return ictcp

# Dolby ICtCp HDR opponent space to Davinci Wide Gamut Intermediate
def ITP_TO_DWG(itp_values):

    cs = colour.models.RGB_COLOURSPACE_DAVINCI_WIDE_GAMUT
    
    #ICtCp to XYZ
    xyzLin = colour.ICtCp_to_XYZ(itp_values)
    
    #XYZ Linear to DWG DI
    dwg = colour.XYZ_to_RGB(xyzLin, cs.whitepoint, cs.whitepoint, cs.matrix_XYZ_to_RGB, cctf_encoding=cs.cctf_encoding)
    
    return dwg

自定义损失函数(当前未启用)

def ITP_loss(y_true, y_pred):
    
    # Split the ICtCp values into I, T, and P components
    I_1, T_1, P_1 = tf.split(y_true, 3, axis=-1)
    I_2, T_2, P_2 = tf.split(y_pred, 3, axis=-1)

    
    # Adjust the T components as in the original delta_E_ITP function
    T_1 = T_1 * 0.5
    T_2 = T_2 * 0.5

    # Compute the squared differences
    d_E_ITP = 720 * tf.sqrt(
        tf.square(I_2 - I_1) +
        tf.square(T_2 - T_1) +
        tf.square(P_2 - P_1)
    )
    
    # Return the mean error as the loss
    return tf.reduce_mean(d_E_ITP)

神经网络代码

def transform_nn(combined_rgb_values, output_callback, epochs=10000, batch_size=32):
    source_rgb = np.vstack([rgb_pair[0] for rgb_pair in combined_rgb_values])
    target_rgb = np.vstack([rgb_pair[1] for rgb_pair in combined_rgb_values])

    source_itp = DWG_TO_ITP(source_rgb)
    target_itp = DWG_TO_ITP(target_rgb)
    
    # Neural network base model with L2 regularization
    alpha = 0  # no penalty for now
    model = keras.Sequential([
        keras.layers.Input(shape=(3,)),
        keras.layers.Dense(128, activation = 'gelu', kernel_regularizer = keras.regularizers.L2(alpha)),
        keras.layers.Dense(64, activation = 'gelu', kernel_regularizer = keras.regularizers.L2(alpha)),
        keras.layers.Dense(32, activation = 'gelu', kernel_regularizer = keras.regularizers.L2(alpha)),
        keras.layers.Dense(3,)
    ])

    # Model optimization with Adam
    adam_optimizer = keras.optimizers.Adam(learning_rate=0.001)
    model.compile(
        optimizer= adam_optimizer,
        loss= "mean_squared_error",
        metrics=['mean_squared_error'])
    
    #normal
    early_stopping_norm = EarlyStopping(
        monitor = 'val_loss',
        patience = 30,
        verbose=1,
        restore_best_weights=True
    )
    
    # Train without early stopping
    history = model.fit(x=source_itp, y=target_itp,
                        epochs=epochs, batch_size=batch_size, 
                        verbose="auto", validation_split=0.3, 
                        callbacks=[early_stopping_norm])
    
    def interpolator(input_rgb):
        input_itp = DWG_TO_ITP(input_rgb)
        output_itp = model.predict(input_itp)
        output_rgb = ITP_TO_DWG(output_itp)
        return output_rgb
    
    return interpolator

损失情况(使用内置MSE损失)

注:自定义损失也会产生类似极端值,使用CIELAB或YCbCr实现时从未遇到此问题

Epoch 1/10000
  70/9078 [..............................] - ETA: 6s - loss: 19151210161612029119172287351962936121302040109299793920.0000 - mean_squared_error: 1915121016161202911917228735196293612130 150/9078 [..............................] - ETA: 6s - loss: 8937231408752302941104862160146934414914780835554000896.0000 - mean_squared_error: 89372314087523029411048621601469344149147 236/9078 [..............................] - ETA: 5s - loss: 8411239422438024050387858001836140461389620391983448064.0000 - mean_squared_error: 84112394224380240503878580018361404613896 322/9078 [>.............................] - ETA: 5s - loss: 55694874365583834449267799576553768559551931724848365789378071082067252634355826658906428848197067342214382161787617280.0000 407/9078 [>.............................] - ETA: 5s - loss: 9272320170949610945087897503565859983725183487173275717008470165482614622395441710684957926712521227412477496744314184602899 494/9078 [>.............................] - ETA: 5s - loss: inf - mean_squared_error: inf                                                                                               9078/9078 [==============================] - 7s 686us/step - loss: nan - mean_squared_error: nan - val_loss: nan - val_mean_squared_error: nan                                         
Epoch 2/10000
9078/9078 [==============================] - 6s 682us/step - loss: nan - mean_squared_error: nan - val_loss: nan - val_mean_squared_error: nan
Epoch 3/10000
8987/9078 [============================>.] - ETA: 0s - loss: nan - mean_squared_error: nan%

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.19 19:14:56