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

如何在Keras中实现CTC损失?基础序列模型改造遇训练输入问题

现有代码修复

你现有实现的错误主要有4处:

  • 基础网络的第一个Dense层未传入输入张量,导致计算图断裂
  • 构造模型时未将labels输入层加入模型输入列表,导致训练时无法传入标签数据
  • 标签输入层的dtype设置为float32,不符合CTC损失要求的整数类型输入
  • 模型输出指定错误,未将CTCLayer的返回值作为模型输出,损失逻辑未被加入计算图

修复后的完整代码如下:

import tensorflow as tf
from tensorflow import keras
from tensorflow.keras import layers, Model

class CTCLayer(layers.Layer):
    def __init__(self, name=None):
        super().__init__(name=name)
        self.loss_fn = keras.backend.ctc_batch_cost

    def call(self, y_true, y_pred):
        batch_len = tf.cast(tf.shape(y_true)[0], dtype="int64")
        input_length = tf.cast(tf.shape(y_pred)[1], dtype="int64")
        label_length = tf.cast(tf.shape(y_true)[1], dtype="int64")

        input_length = input_length * tf.ones(shape=(batch_len, 1), dtype="int64")
        label_length = label_length * tf.ones(shape=(batch_len, 1), dtype="int64")

        loss = self.loss_fn(y_true, y_pred, input_length, label_length)
        self.add_loss(loss)
        return y_pred

# 定义所有输入层
labels = layers.Input(shape=(None,), dtype="int64")
input_layer = layers.Input(shape=(50,20))

# 基础网络结构
layer = layers.Dense(123, activation = 'relu')(input_layer)
layer = layers.LSTM(128, return_sequences = True)(layer)
outputs = layers.Dense(20, activation='softmax')(layer)
ctc_output = CTCLayer()(labels, outputs)

# 构造模型时传入所有输入
model = Model(inputs=[input_layer, labels], outputs=ctc_output)
# 编译时不需要额外指定损失,因为损失已经在CTCLayer中添加
model.compile(optimizer=keras.optimizers.Adam())

训练时按照双输入的格式传参即可:

# train_x 是形状为 (样本数, 50, 20) 的训练数据
# train_y 是形状为 (样本数, 最大标签长度) 的整数标签
model.fit(x=[train_x, train_y], epochs=10, batch_size=32)

推理时单独提取预测分支即可,不需要传入标签:

infer_model = Model(inputs=input_layer, outputs=outputs)
# 直接传入输入数据得到预测结果
pred = infer_model.predict(test_x)

更高效简洁的实现方案(推荐)

不需要自定义层,直接使用Keras内置的CTC损失接口实现,逻辑更清晰,适配标准Keras训练流程,部署更方便。
完整实现代码如下:

import tensorflow as tf
from tensorflow import keras
from tensorflow.keras import layers, Model

# 定义基础模型(训练和推理共用)
input_layer = layers.Input(shape=(50,20))
layer = layers.Dense(123, activation = 'relu')(input_layer)
layer = layers.LSTM(128, return_sequences = True)(layer)
outputs = layers.Dense(20, activation='softmax')(layer)
model = Model(input_layer, outputs)

# 定义CTC损失,默认空白符索引为最后一位(即19,对应你的20分类),如果空白符为第0位可修改blank参数
def ctc_loss(y_true, y_pred):
    batch_len = tf.cast(tf.shape(y_true)[0], dtype="int64")
    input_length = tf.cast(tf.shape(y_pred)[1], dtype="int64")
    label_length = tf.cast(tf.shape(y_true)[1], dtype="int64")
    
    input_length = input_length * tf.ones(shape=(batch_len, 1), dtype="int64")
    label_length = label_length * tf.ones(shape=(batch_len, 1), dtype="int64")
    return keras.backend.ctc_batch_cost(y_true, y_pred, input_length, label_length)

# 编译模型,直接指定损失函数
model.compile(optimizer=keras.optimizers.Adam(), loss=ctc_loss)

训练时使用标准Keras训练流程即可,不需要调整输入格式:

model.fit(x=train_x, y=train_y, epochs=10, batch_size=32)

推理时直接使用当前模型即可,不需要额外构造推理分支:

pred = model.predict(test_x)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.02 09:45:02