SoftTraderBackend/ServiceHandler/FinBertScoringService.cs
2026-03-21 09:16:59 -04:00

221 lines
5.3 KiB
C#

using Microsoft.ML.OnnxRuntime;
using Microsoft.ML.OnnxRuntime.Tensors;
using System.Globalization;
using System.Text;
namespace Backend.ServiceHandler
{
public sealed class FinBertScoringService : IDisposable
{
private readonly InferenceSession _session;
private readonly Dictionary<string, int> _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<long>(new[] { 1, seqLen });
var attentionMaskTensor = new DenseTensor<long>(new[] { 1, seqLen });
var tokenTypeIdsTensor = new DenseTensor<long>(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>
{
NamedOnnxValue.CreateFromTensor(_inputIdsName, inputIdsTensor),
NamedOnnxValue.CreateFromTensor(_attentionMaskName, attentionMaskTensor),
NamedOnnxValue.CreateFromTensor(_tokenTypeIdsName, tokenTypeIdsTensor)
};
using var results = _session.Run(inputs);
var outputTensor = results.First().AsTensor<float>();
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<int> Tokenize(string text)
{
var ids = new List<int> { 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<string> BasicTokenize(string text)
{
var tokens = new List<string>();
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<int> WordPieceTokenize(string token)
{
var ids = new List<int>();
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<string, int> LoadVocab(string vocabPath)
{
var vocab = new Dictionary<string, int>();
int index = 0;
foreach (string line in File.ReadLines(vocabPath))
{
vocab[line.TrimEnd()] = index++;
}
return vocab;
}
public void Dispose()
{
_session.Dispose();
}
}
}