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

求助:C#中是否有可复现Keras CTC解码功能的类库?

.NET环境下复现TensorFlow/Keras CTC解码方案

下面提供两种可行的实现方式,对应你Python代码中的CTC贪婪解码逻辑:


方案一:基于TensorFlow.NET直接实现

TensorFlow.NET是TensorFlow的官方.NET绑定库,可以直接复用类似Keras的CTC解码API,兼容性更好。

1. 前置准备

安装TensorFlow.NET和NumSharp的NuGet包:

Install-Package TensorFlow.NET
Install-Package NumSharp

2. 解码函数实现

using TensorFlow;
using NumSharp;
using System.Collections.Generic;
using System.Linq;

public List<string> DecodeBatchPredictions(NDArray pred, int maxLength, Dictionary<int, string> numToChar)
{
    // 构造input_len,对应Python中np.ones(pred.shape[0]) * pred.shape[1]
    int batchSize = pred.Shape[0];
    int seqLength = pred.Shape[1];
    NDArray inputLen = np.ones(batchSize) * seqLength;

    // 执行CTC贪婪解码
    using var tfInputLen = TFNDArray.From(inputLen);
    using var tfPred = TFNDArray.From(pred);
    var ctcDecodeResult = TF.OperationDecoder.CtcDecode(tfPred, tfInputLen, greedy: true);
    
    // 提取解码结果并截断到maxLength
    var decodedTensor = ctcDecodeResult[0][0];
    var decodedArray = decodedTensor.ToArray<int>();
    var outputText = new List<string>();

    for (int i = 0; i < batchSize; i++)
    {
        // 按批次提取序列
        var seq = decodedArray.Skip(i * seqLength).Take(maxLength).ToList();
        // 转换为文本
        var text = string.Join("", seq.Select(num => numToChar.TryGetValue(num, out var c) ? c : string.Empty));
        outputText.Add(text);
    }

    return outputText;
}

注意事项

  • 提前将Python中用于映射的num_to_char字典转换为.NET的Dictionary<int, string>
  • 确保TensorFlow.NET版本与训练模型的TensorFlow主版本一致(比如都是2.x)
  • Keras保存的模型可以用KerasModel.LoadModel方法直接加载到TensorFlow.NET中

方案二:基于ONNX Runtime手动实现CTC解码

如果不想依赖TensorFlow生态,可以将Keras模型导出为ONNX格式,用ONNX Runtime推理后手动实现贪婪解码逻辑,适合跨平台场景。

1. Python端导出ONNX模型

import tensorflow as tf
from tensorflow.keras.models import load_model
import tf2onnx

# 加载训练好的Keras模型
model = load_model("your_trained_model.h5")
# 定义输入签名(根据你的模型输入维度调整)
input_spec = (tf.TensorSpec((None, 100, 64), tf.float32, name="input"),)
# 导出为ONNX格式
tf2onnx.convert.from_keras(model, input_signature=input_spec, output_path="model.onnx")

2. .NET端解码实现

先安装ONNX Runtime NuGet包:

Install-Package Microsoft.ML.OnnxRuntime

解码函数代码:

using Microsoft.ML.OnnxRuntime;
using Microsoft.ML.OnnxRuntime.Tensors;
using System.Collections.Generic;
using System.Linq;

public List<string> DecodeBatchWithOnnx(Tensor<float> predTensor, int maxLength, Dictionary<int, string> numToChar, int blankIndex = 0)
{
    var outputText = new List<string>();
    int batchSize = predTensor.Dimensions[0];
    int seqLength = predTensor.Dimensions[1];

    foreach (int batchIdx in Enumerable.Range(0, batchSize))
    {
        var currentChars = new List<int>();
        int prevChar = -1;

        for (int timeStep = 0; timeStep < seqLength; timeStep++)
        {
            // 找到当前时间步概率最大的类别索引
            int maxIdx = 0;
            float maxProb = predTensor[batchIdx, timeStep, 0];
            for (int cls = 1; cls < predTensor.Dimensions[2]; cls++)
            {
                if (predTensor[batchIdx, timeStep, cls] > maxProb)
                {
                    maxProb = predTensor[batchIdx, timeStep, cls];
                    maxIdx = cls;
                }
            }

            // 跳过空白符和连续重复字符(CTC核心规则)
            if (maxIdx != blankIndex && maxIdx != prevChar)
            {
                currentChars.Add(maxIdx);
                prevChar = maxIdx;
            }
        }

        // 截断到指定长度并转换为文本
        var truncatedSeq = currentChars.Take(maxLength).ToList();
        var text = string.Join("", truncatedSeq.Select(num => numToChar[num]));
        outputText.Add(text);
    }

    return outputText;
}

注意事项

  • blankIndex要和训练时CTC层设置的空白符索引一致(通常为0)
  • ONNX模型推理时,需要将输入转换为符合模型要求的Tensor<float>格式
  • 导出ONNX时要确保输入维度与模型实际输入匹配

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.02 15:41:46