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

基于ONNX的M2M100_418M模型C#翻译代码补全及目标语言设置

问题描述

已从Hugging Face下载M2M100_418M模型并转换为encoder.onnx和decoder.onnx文件,希望实现文本翻译功能。目前已使用BertTokenizer开源库和Microsoft.ML.OnnxRuntime编写了以下C#代码,现需补全代码以获取翻译结果,并了解如何在此配置中指定目标语言,实现将输入文本翻译为任意目标语言的需求。

用户原代码:

public class EncoderInput
{
    [VectorType(1, 256)]
    [ColumnName("input_ids")]
    public long[] InputIds { get; set; }

    [VectorType(1, 256)]
    [ColumnName("attention_mask")]
    public long[] AttentionMask { get; set; }
}

public class DecoderInput
{
    [VectorType(1, 256)]
    [ColumnName("encoder_attention_mask")]
    public long[] EncoderAttensionMask { get; set; }

    [VectorType(1, 256)]
    [ColumnName("input_ids")]
    public long[] InputIds { get; set; }

    [VectorType(1, 256)]
    [ColumnName("encoder_hidden_states")]
    public float[] EncoderHiddenState { get; set; }
}

var source = "I love you";
var tokenizer = new BertMultilingualTokenizer();
var encoded = tokenizer.Encode(256, source);

var encoderInput = new EncoderInput()
{
    InputIds = encoded.Select(t => t.InputIds).ToArray(),
    AttentionMask = encoded.Select(t => t.AttentionMask).ToArray(),
};

using var inputIdsOrtValue = OrtValue.CreateTensorValueFromMemory(encoderInput.InputIds,
      new long[] { 1, encoderInput.InputIds.Length });

using var attMaskOrtValue = OrtValue.CreateTensorValueFromMemory(encoderInput.AttentionMask,
      new long[] { 1, encoderInput.AttentionMask.Length });

var encoder = new InferenceSession("./onnx/m2m100_418M/encoder_model.onnx");
var decoder = new InferenceSession("./onnx/m2m100_418M/decoder_model.onnx");

var encoderInputs = new Dictionary<string, OrtValue>
{
    { "input_ids", inputIdsOrtValue },
    { "attention_mask", attMaskOrtValue },
};

using var runOptions = new RunOptions();
using var encoderOutput = encoder.Run(runOptions, encoderInputs, encoder.OutputNames);

var t = encoderOutput.First().GetTensorTypeAndShape();

var decoderInput = new DecoderInput()
{
    InputIds = encoded.Select(t => t.InputIds).ToArray(),
    EncoderAttensionMask = encoded.Select(t => t.AttentionMask).ToArray(),
    EncoderHiddenState = encoderOutput.First().GetTensorDataAsSpan<float>().ToArray()
};

using var hiddenState = OrtValue.CreateTensorValueFromMemory(decoderInput.EncoderHiddenState,
      new long[] { 1, decoderInput.EncoderHiddenState.Length / 1024, 1024 });

var decoderInputs = new Dictionary<string, OrtValue>
{
    { "encoder_attention_mask", attMaskOrtValue },
    { "input_ids", inputIdsOrtValue },
    { "encoder_hidden_states", hiddenState },
};

using var decoderOutput = decoder.Run(runOptions, decoderInputs, decoder.OutputNames);
解决方案

1. 核心逻辑说明

M2M100是自回归式多语言翻译模型,翻译流程需要:

  • 编码器处理源文本生成隐藏状态
  • 解码器从目标语言起始token(格式为<s>[目标语言代码]</s>,比如中文是<s>zh</s>)开始,逐token生成翻译结果,直到触发结束token</s>
  • 必须使用M2M100专用Tokenizer,通用BertTokenizer无法正确处理其多语言映射规则

2. 补全代码并实现目标语言指定

以下是完整的可运行代码,包含目标语言配置、自回归生成翻译结果的逻辑:

using System;
using System.Collections.Generic;
using System.Linq;
using Microsoft.ML.OnnxRuntime;
using HuggingFace.Tokenizers;

// 编码器输入结构
public class EncoderInput
{
    [VectorType(1, 256)]
    [ColumnName("input_ids")]
    public long[] InputIds { get; set; }

    [VectorType(1, 256)]
    [ColumnName("attention_mask")]
    public long[] AttentionMask { get; set; }
}

// 解码器输入结构(适配动态序列长度)
public class DecoderInput
{
    [VectorType(1, -1)]
    [ColumnName("input_ids")]
    public long[] InputIds { get; set; }

    [VectorType(1, 256)]
    [ColumnName("encoder_attention_mask")]
    public long[] EncoderAttentionMask { get; set; }

    [VectorType(1, 256, 1024)]
    public float[] EncoderHiddenStates { get; set; }
}

class M2M100Translator
{
    static void Main(string[] args)
    {
        // 配置参数
        string sourceText = "I love you";
        string sourceLang = "en"; // 源语言代码
        string targetLang = "zh"; // 目标语言代码,可替换为fr/es/ja等支持语言
        int maxSeqLen = 256;
        string encoderModelPath = "./onnx/m2m100_418M/encoder_model.onnx";
        string decoderModelPath = "./onnx/m2m100_418M/decoder_model.onnx";
        string tokenizerPath = "./path/to/m2m100_tokenizer"; // 需从HuggingFace下载对应tokenizer文件

        // 初始化M2M100专用Tokenizer
        var tokenizer = new Tokenizer(tokenizerPath);
        tokenizer.SetSourceLanguage(sourceLang);
        tokenizer.SetTargetLanguage(targetLang);

        // 1. 编码源文本并运行编码器
        var sourceEncoding = tokenizer.Encode(sourceText, maxSeqLen, truncation: true, padding: PaddingStrategy.MaxLength);
        var encoderInput = new EncoderInput
        {
            InputIds = sourceEncoding.Ids,
            AttentionMask = sourceEncoding.AttentionMask.Select(x => (long)x).ToArray()
        };

        using var inputIdsOrt = OrtValue.CreateTensorValueFromMemory(encoderInput.InputIds, new long[] { 1, maxSeqLen });
        using var attMaskOrt = OrtValue.CreateTensorValueFromMemory(encoderInput.AttentionMask, new long[] { 1, maxSeqLen });

        var encoderSession = new InferenceSession(encoderModelPath);
        var encoderInputs = new Dictionary<string, OrtValue>
        {
            { "input_ids", inputIdsOrt },
            { "attention_mask", attMaskOrt }
        };

        using var encoderOutputs = encoderSession.Run(null, encoderInputs, encoderSession.OutputNames);
        var encoderHiddenStates = encoderOutputs.First().GetTensorDataAsSpan<float>().ToArray();
        var encoderHiddenShape = encoderOutputs.First().GetTensorTypeAndShape().GetShape();

        // 2. 初始化解码器输入:目标语言起始token
        var targetStartToken = tokenizer.Encode($"<s>{targetLang}</s>", maxSeqLen).Ids.Take(1).ToArray();
        var decoderInputIds = new List<long>(targetStartToken);
        var endTokenId = tokenizer.GetTokenId("</s>");

        var decoderSession = new InferenceSession(decoderModelPath);
        string translationResult = string.Empty;

        // 3. 自回归生成翻译结果
        while (decoderInputIds.Count < maxSeqLen && decoderInputIds.Last() != endTokenId)
        {
            var decoderInput = new DecoderInput
            {
                InputIds = decoderInputIds.ToArray(),
                EncoderAttentionMask = encoderInput.AttentionMask,
                EncoderHiddenStates = encoderHiddenStates
            };

            // 转换为符合要求的OrtValue
            using var decoderInputIdsOrt = OrtValue.CreateTensorValueFromMemory(decoderInput.InputIds, new long[] { 1, decoderInput.InputIds.Length });
            using var encoderAttMaskOrt = OrtValue.CreateTensorValueFromMemory(decoderInput.EncoderAttentionMask, encoderHiddenShape.Take(2).ToArray());
            using var encoderHiddenOrt = OrtValue.CreateTensorValueFromMemory(decoderInput.EncoderHiddenStates, encoderHiddenShape);

            var decoderInputs = new Dictionary<string, OrtValue>
            {
                { "input_ids", decoderInputIdsOrt },
                { "encoder_attention_mask", encoderAttMaskOrt },
                { "encoder_hidden_states", encoderHiddenOrt }
            };

            // 运行解码器推理
            using var decoderOutputs = decoderSession.Run(null, decoderInputs, decoderSession.OutputNames);
            var logits = decoderOutputs.First().GetTensorDataAsSpan<float>().ToArray();
            
            // 取概率最大的token作为下一个生成结果
            int currentPos = decoderInputIds.Count - 1;
            var tokenLogits = logits.Skip(currentPos * tokenizer.GetVocabSize()).Take(tokenizer.GetVocabSize()).ToArray();
            long nextTokenId = Array.IndexOf(tokenLogits, tokenLogits.Max());

            decoderInputIds.Add(nextTokenId);
        }

        // 4. 解码token得到最终翻译文本
        translationResult = tokenizer.Decode(decoderInputIds.Skip(1).TakeWhile(id => id != endTokenId).ToArray());
        Console.WriteLine($"翻译结果:{translationResult}");
    }
}

关键注意事项

  • Tokenizer替换:必须使用M2M100专用Tokenizer,加载从HuggingFace下载的tokenizer配置文件,通用BertTokenizer无法适配多语言token规则。
  • 目标语言指定:通过tokenizer.SetTargetLanguage(targetLang)配置目标语言,同时解码器初始输入必须是对应目标语言的起始token。
  • 自回归生成:解码器需要循环生成token,直到达到最大长度或触发结束token</s>,无法一次性生成完整序列。
  • 张量维度匹配:确保编码器输出的隐藏状态维度与解码器输入要求一致,通常为[batch_size, sequence_length, hidden_size]。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.22 00:47:03