构建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
相关产品推荐
相关产品推荐

