基于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
相关产品推荐
相关产品推荐

