TensorFlow CRNN-CTC自定义指标不启用eager执行如何正常生效
核心问题原因
你遇到的问题本质是三个:
- 静态shape判断
y_true.shape[1] is not None在图模式下永远不成立:图执行阶段标签占位符的静态维度是未知的,该分支会被直接剪枝,指标代码永远不执行。 decode_batch_predictions调用CTCGreedyDecoder时输入不符合要求:CTC解码器要求传入的序列长度参数为[batch_size]格式的一维张量,原实现传入了标量,导致静态shape校验报错,因此你才需要加上述判断规避报错。- 原有准确率计算用了Python侧的shape取值逻辑,图模式下静态维度未知时会返回None,无法正常执行。
修改方案
1. 重写decode_batch_predictions,适配CTC解码器输入要求
def decode_batch_predictions(y_pred, max_label_length): batch_size = tf.shape(y_pred)[0] seq_len = tf.shape(y_pred)[1] # 生成符合要求的一维序列长度张量 seq_lengths = tf.fill([batch_size], seq_len) # 转置为CTC解码器要求的[序列长度, batch_size, 分类数]格式 y_pred = tf.transpose(y_pred, perm=[1, 0, 2]) # 执行解码 decoded, _ = tf.keras.backend.ctc_decode(y_pred, input_length=seq_lengths, greedy=True) # 补全到固定最大长度 decoded = decoded[0][:, :max_label_length] pad_len = max_label_length - tf.shape(decoded)[1] decoded = tf.pad(decoded, [[0, 0], [0, pad_len]], constant_values=-1) return decoded
2. 重构准确率计算函数,全用TensorFlow动态操作
def calculate_accuracy(y_true, y_pred, metric, unknown_placeholder): y_pred = tf.cast(y_pred, y_true.dtype) # 替换-1为未知字符占位符 y_pred = tf.where(y_pred == -1, tf.cast(unknown_placeholder, y_true.dtype), y_pred) if metric == 'word': # 逐样本判断整词匹配,直接算平均准确率 correct_words = tf.reduce_all(tf.equal(y_true, y_pred), axis=1) return tf.reduce_mean(tf.cast(correct_words, tf.float32)) elif metric == 'char': # 逐字符判断匹配,直接算平均准确率 correct_chars = tf.equal(y_true, y_pred) return tf.reduce_mean(tf.cast(correct_chars, tf.float32)) return tf.constant(0.0, dtype=tf.float32)
3. 重写CTCLayer,删除静态shape判断,添加指标状态管理
class CTCLayer(tf.keras.layers.Layer): def __init__(self, max_label_length, unknown_placeholder, **kwargs): super().__init__(**kwargs) self.max_label_length = max_label_length self.unknown_placeholder = unknown_placeholder # 初始化指标实例,支持跨batch累积准确率 self.word_acc = tf.keras.metrics.Mean(name="word_accuracy") self.char_acc = tf.keras.metrics.Mean(name="char_accuracy") def call(self, *args): y_true, y_pred = args batch_length = tf.shape(y_true)[0] input_length = tf.shape(y_pred)[1] label_length = tf.shape(y_true)[1] input_length = tf.fill([batch_length, 1], input_length) label_length = tf.fill([batch_length, 1], label_length) loss = tf.keras.backend.ctc_batch_cost( y_true, y_pred, input_length, label_length ) # 直接计算指标,无需静态shape判断 predictions = decode_batch_predictions(y_pred, self.max_label_length) # 对齐标签和预测结果的长度,避免动态shape不匹配 predictions = predictions[:, :label_length] word_acc = calculate_accuracy(y_true, predictions, 'word', self.unknown_placeholder) char_acc = calculate_accuracy(y_true, predictions, 'char', self.unknown_placeholder) # 更新指标状态 self.word_acc.update_state(word_acc) self.char_acc.update_state(char_acc) # 注册到层指标 self.add_metric(self.word_acc.result(), name='word_accuracy') self.add_metric(self.char_acc.result(), name='char_accuracy') self.add_loss(loss) return y_pred def reset_states(self): # 每个epoch结束重置指标状态 self.word_acc.reset_states() self.char_acc.reset_states()
效果验证
修改完成后无需开启eager执行,图模式下即可正常输出word_accuracy和char_accuracy指标,也不会触发shape校验报错。
内容的提问来源于stack exchange,提问作者watch-this
相关产品推荐
相关产品推荐

