using Microsoft.ML.OnnxRuntime; using Microsoft.ML.OnnxRuntime.Tensors; using System.Globalization; using System.Text; namespace Backend.ServiceHander { public sealed class FinBertScoringService : IDisposable { private readonly InferenceSession _session; private readonly Dictionary _vocab; private readonly string _inputIdsName; private readonly string _attentionMaskName; private readonly string _tokenTypeIdsName; private readonly string _outputName; private const int MaxLength = 512; private const int ClsId = 101; private const int SepId = 102; private const int UnkId = 100; // Label mapping from config.json: 0=positive, 1=negative, 2=neutral private static readonly string[] Labels = ["positive", "negative", "neutral"]; public FinBertScoringService(IWebHostEnvironment env) { string modelDir = Path.Combine(env.ContentRootPath, "AIModels", "finbert"); string modelPath = Path.Combine(modelDir, "model.onnx"); string vocabPath = Path.Combine(modelDir, "vocab.txt"); _vocab = LoadVocab(vocabPath); _session = new InferenceSession(modelPath); var inputNames = _session.InputMetadata.Keys.ToList(); _inputIdsName = inputNames.First(n => n.Contains("input_id", StringComparison.OrdinalIgnoreCase)); _attentionMaskName = inputNames.First(n => n.Contains("attention_mask", StringComparison.OrdinalIgnoreCase)); _tokenTypeIdsName = inputNames.First(n => n.Contains("token_type", StringComparison.OrdinalIgnoreCase)); _outputName = _session.OutputMetadata.Keys.First(); } public (double Score, string Label) Score(string text) { var inputIds = Tokenize(text); int seqLen = inputIds.Count; var inputIdsTensor = new DenseTensor(new[] { 1, seqLen }); var attentionMaskTensor = new DenseTensor(new[] { 1, seqLen }); var tokenTypeIdsTensor = new DenseTensor(new[] { 1, seqLen }); for (int i = 0; i < seqLen; i++) { inputIdsTensor[0, i] = inputIds[i]; attentionMaskTensor[0, i] = 1; tokenTypeIdsTensor[0, i] = 0; } var inputs = new List { NamedOnnxValue.CreateFromTensor(_inputIdsName, inputIdsTensor), NamedOnnxValue.CreateFromTensor(_attentionMaskName, attentionMaskTensor), NamedOnnxValue.CreateFromTensor(_tokenTypeIdsName, tokenTypeIdsTensor) }; using var results = _session.Run(inputs); var outputTensor = results.First().AsTensor(); float[] logits = outputTensor.ToArray(); var probs = Softmax(logits); // Score: positive * 100 + neutral * 50 + negative * 0 double score = Math.Round(probs[0] * 100 + probs[2] * 50, 2); int maxIdx = Array.IndexOf(probs, probs.Max()); string label = Labels[maxIdx]; return (score, label); } private List Tokenize(string text) { var ids = new List { ClsId }; text = NormalizeText(text); var basicTokens = BasicTokenize(text); foreach (var token in basicTokens) { ids.AddRange(WordPieceTokenize(token)); } ids.Add(SepId); if (ids.Count > MaxLength) { ids = ids.Take(MaxLength - 1).ToList(); ids.Add(SepId); } return ids; } private static string NormalizeText(string text) { text = text.ToLowerInvariant(); // Strip accents: NFD normalize then remove combining marks text = new string(text.Normalize(NormalizationForm.FormD) .Where(c => CharUnicodeInfo.GetUnicodeCategory(c) != UnicodeCategory.NonSpacingMark) .ToArray()); var sb = new StringBuilder(text.Length); foreach (char c in text) { if (char.IsControl(c)) sb.Append(' '); else sb.Append(c); } return sb.ToString(); } private static List BasicTokenize(string text) { var tokens = new List(); var current = new StringBuilder(); foreach (char c in text) { if (char.IsWhiteSpace(c)) { if (current.Length > 0) { tokens.Add(current.ToString()); current.Clear(); } } else if (char.IsPunctuation(c) || char.IsSymbol(c)) { if (current.Length > 0) { tokens.Add(current.ToString()); current.Clear(); } tokens.Add(c.ToString()); } else { current.Append(c); } } if (current.Length > 0) tokens.Add(current.ToString()); return tokens; } private List WordPieceTokenize(string token) { var ids = new List(); int start = 0; while (start < token.Length) { int end = token.Length; bool found = false; while (start < end) { string sub = start > 0 ? "##" + token[start..end] : token[start..end]; if (_vocab.TryGetValue(sub, out int id)) { ids.Add(id); found = true; start = end; break; } end--; } if (!found) { ids.Add(UnkId); start++; } } return ids; } private static float[] Softmax(float[] logits) { float max = logits.Max(); var exps = logits.Select(l => MathF.Exp(l - max)).ToArray(); float sum = exps.Sum(); return exps.Select(e => e / sum).ToArray(); } private static Dictionary LoadVocab(string vocabPath) { var vocab = new Dictionary(); int index = 0; foreach (string line in File.ReadLines(vocabPath)) { vocab[line.TrimEnd()] = index++; } return vocab; } public void Dispose() { _session.Dispose(); } } }