221 lines
5.3 KiB
C#
221 lines
5.3 KiB
C#
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<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();
|
|
}
|
|
}
|
|
}
|