Integrated FinBERT model.

This commit is contained in:
caspar 2026-03-19 14:43:28 -04:00
parent df759d0ead
commit 7f8d6a8131
8 changed files with 527 additions and 0 deletions

View File

@ -12,6 +12,9 @@ builder.Services.AddSwaggerGen();
builder.Services.AddSingleton<IHistoricalMarketDataCache, PostgresHistoricalMarketDataCache>(); builder.Services.AddSingleton<IHistoricalMarketDataCache, PostgresHistoricalMarketDataCache>();
builder.Services.AddSingleton<MarketDataService>(); builder.Services.AddSingleton<MarketDataService>();
builder.Services.AddSingleton<IUserInterestStore, PostgresUserInterestStore>(); builder.Services.AddSingleton<IUserInterestStore, PostgresUserInterestStore>();
builder.Services.AddSingleton<FinBertScoringService>();
builder.Services.AddSingleton<INewsSentimentStore, PostgresNewsSentimentStore>();
builder.Services.AddSingleton<NewsSentimentService>();
builder.Services.AddCors(options => builder.Services.AddCors(options =>
{ {

View File

@ -13,6 +13,7 @@
<ItemGroup> <ItemGroup>
<PackageReference Include="Microsoft.ML.OnnxRuntime" Version="1.22.0" />
<PackageReference Include="Npgsql" Version="10.0.2" /> <PackageReference Include="Npgsql" Version="10.0.2" />
<PackageReference Include="Swashbuckle.AspNetCore" Version="10.1.5" /> <PackageReference Include="Swashbuckle.AspNetCore" Version="10.1.5" />
</ItemGroup> </ItemGroup>

View File

@ -0,0 +1,46 @@
using Backend.ServiceHander;
using Microsoft.AspNetCore.Mvc;
namespace Backend.Controller
{
[ApiController]
[Route("api/news-sentiment")]
public sealed class NewsSentimentController : ControllerBase
{
private readonly NewsSentimentService _service;
public NewsSentimentController(NewsSentimentService service)
{
_service = service;
}
[HttpGet]
public async Task<IActionResult> Analyze(
[FromQuery] string keyword,
CancellationToken cancellationToken)
{
try
{
var results = await _service.AnalyzeAsync(keyword, cancellationToken);
return Ok(results.Select(r => new
{
source = r.Source,
keyword = r.Keyword,
publishedAt = r.PublishedAt,
title = r.Title,
score = r.Score,
sentimentLabel = r.SentimentLabel,
url = r.Url
}));
}
catch (ArgumentException ex)
{
return BadRequest(ex.Message);
}
catch (Exception ex)
{
return StatusCode(500, new { error = ex.ToString() });
}
}
}
}

View File

@ -0,0 +1,14 @@
namespace Backend.DatabaseHandler
{
public sealed class NewsSentimentRecord
{
public string Source { get; init; } = string.Empty;
public string Keyword { get; init; } = string.Empty;
public DateTimeOffset PublishedAt { get; init; }
public string Title { get; init; } = string.Empty;
public double Score { get; init; }
public string SentimentLabel { get; init; } = string.Empty;
public string TitleHash { get; init; } = string.Empty;
public string Url { get; init; } = string.Empty;
}
}

View File

@ -0,0 +1,123 @@
using Backend.Interface;
using Npgsql;
namespace Backend.DatabaseHandler
{
public sealed class PostgresNewsSentimentStore : INewsSentimentStore
{
private readonly string _connectionString;
public PostgresNewsSentimentStore(IConfiguration configuration)
{
_connectionString = configuration.GetConnectionString("MarketDataDb")
?? throw new InvalidOperationException("Missing connection string: MarketDataDb");
}
public async Task<List<NewsSentimentRecord>> GetByKeywordAsync(
string keyword,
CancellationToken cancellationToken = default)
{
var results = new List<NewsSentimentRecord>();
await using var conn = new NpgsqlConnection(_connectionString);
await conn.OpenAsync(cancellationToken);
const string sql = """
SELECT source, keyword, published_at, title, score,
sentiment_label, title_hash, url
FROM news_sentiment_cache
WHERE keyword = @keyword
ORDER BY published_at DESC;
""";
await using var cmd = new NpgsqlCommand(sql, conn);
cmd.Parameters.AddWithValue("keyword", keyword);
await using var reader = await cmd.ExecuteReaderAsync(cancellationToken);
while (await reader.ReadAsync(cancellationToken))
{
results.Add(new NewsSentimentRecord
{
Source = reader.GetString(0),
Keyword = reader.GetString(1),
PublishedAt = reader.GetFieldValue<DateTimeOffset>(2),
Title = reader.GetString(3),
Score = reader.GetDouble(4),
SentimentLabel = reader.GetString(5),
TitleHash = reader.GetString(6).TrimEnd(),
Url = reader.GetString(7)
});
}
return results;
}
public async Task InsertAsync(
IEnumerable<NewsSentimentRecord> records,
CancellationToken cancellationToken = default)
{
var recordList = records.ToList();
if (recordList.Count == 0)
return;
await using var conn = new NpgsqlConnection(_connectionString);
await conn.OpenAsync(cancellationToken);
await using var tx = await conn.BeginTransactionAsync(cancellationToken);
const string sql = """
INSERT INTO news_sentiment_cache
(source, keyword, published_at, title, score,
sentiment_label, title_hash, url)
VALUES
(@source, @keyword, @published_at, @title, @score,
@sentiment_label, @title_hash, @url)
ON CONFLICT (title_hash) DO NOTHING;
""";
foreach (var record in recordList)
{
await using var cmd = new NpgsqlCommand(sql, conn, tx);
cmd.Parameters.AddWithValue("source", record.Source);
cmd.Parameters.AddWithValue("keyword", record.Keyword);
cmd.Parameters.AddWithValue("published_at", record.PublishedAt);
cmd.Parameters.AddWithValue("title", record.Title);
cmd.Parameters.AddWithValue("score", record.Score);
cmd.Parameters.AddWithValue("sentiment_label", record.SentimentLabel);
cmd.Parameters.AddWithValue("title_hash", record.TitleHash);
cmd.Parameters.AddWithValue("url", record.Url);
await cmd.ExecuteNonQueryAsync(cancellationToken);
}
await tx.CommitAsync(cancellationToken);
}
public async Task<HashSet<string>> GetExistingHashesAsync(
string keyword,
CancellationToken cancellationToken = default)
{
var hashes = new HashSet<string>();
await using var conn = new NpgsqlConnection(_connectionString);
await conn.OpenAsync(cancellationToken);
const string sql = """
SELECT title_hash
FROM news_sentiment_cache
WHERE keyword = @keyword;
""";
await using var cmd = new NpgsqlCommand(sql, conn);
cmd.Parameters.AddWithValue("keyword", keyword);
await using var reader = await cmd.ExecuteReaderAsync(cancellationToken);
while (await reader.ReadAsync(cancellationToken))
{
hashes.Add(reader.GetString(0).TrimEnd());
}
return hashes;
}
}
}

View File

@ -0,0 +1,19 @@
using Backend.DatabaseHandler;
namespace Backend.Interface
{
public interface INewsSentimentStore
{
Task<List<NewsSentimentRecord>> GetByKeywordAsync(
string keyword,
CancellationToken cancellationToken = default);
Task InsertAsync(
IEnumerable<NewsSentimentRecord> records,
CancellationToken cancellationToken = default);
Task<HashSet<string>> GetExistingHashesAsync(
string keyword,
CancellationToken cancellationToken = default);
}
}

View File

@ -0,0 +1,220 @@
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();
}
}
}

View File

@ -0,0 +1,101 @@
using Backend.DatabaseHandler;
using Backend.Interface;
using System.Security.Cryptography;
using System.Text;
using System.Xml.Linq;
namespace Backend.ServiceHander
{
public sealed class NewsSentimentService
{
private readonly INewsSentimentStore _store;
private readonly FinBertScoringService _scorer;
private static readonly HttpClient Http = new()
{
DefaultRequestHeaders = { { "User-Agent", "SoftTrader/1.0" } }
};
public NewsSentimentService(INewsSentimentStore store, FinBertScoringService scorer)
{
_store = store;
_scorer = scorer;
}
public async Task<List<NewsSentimentRecord>> AnalyzeAsync(
string keyword,
CancellationToken cancellationToken = default)
{
keyword = keyword.Trim();
if (string.IsNullOrEmpty(keyword))
throw new ArgumentException("Keyword cannot be empty.");
var articles = await FetchRssAsync(keyword, cancellationToken);
var existingHashes = await _store.GetExistingHashesAsync(keyword, cancellationToken);
var newRecords = new List<NewsSentimentRecord>();
foreach (var (title, url, publishedAt) in articles)
{
string titleHash = ComputeHash(title);
if (existingHashes.Contains(titleHash))
continue;
var (score, label) = _scorer.Score(title);
newRecords.Add(new NewsSentimentRecord
{
Source = "google_rss",
Keyword = keyword,
PublishedAt = publishedAt,
Title = title,
Score = score,
SentimentLabel = label,
TitleHash = titleHash,
Url = url
});
}
if (newRecords.Count > 0)
await _store.InsertAsync(newRecords, cancellationToken);
return await _store.GetByKeywordAsync(keyword, cancellationToken);
}
private static async Task<List<(string Title, string Url, DateTimeOffset PublishedAt)>> FetchRssAsync(
string keyword,
CancellationToken cancellationToken)
{
string encodedKeyword = Uri.EscapeDataString(keyword);
string url = $"https://news.google.com/rss/search?q={encodedKeyword}&hl=en-US&gl=US&ceid=US:en";
string xml = await Http.GetStringAsync(url, cancellationToken);
var doc = XDocument.Parse(xml);
var items = new List<(string, string, DateTimeOffset)>();
foreach (var item in doc.Descendants("item"))
{
string? title = item.Element("title")?.Value;
string? link = item.Element("link")?.Value;
string? pubDate = item.Element("pubDate")?.Value;
if (string.IsNullOrWhiteSpace(title) || string.IsNullOrWhiteSpace(link))
continue;
DateTimeOffset published = DateTimeOffset.TryParse(pubDate, out var dt)
? dt
: DateTimeOffset.UtcNow;
items.Add((title, link, published));
}
return items;
}
private static string ComputeHash(string text)
{
byte[] bytes = SHA256.HashData(Encoding.UTF8.GetBytes(text));
return Convert.ToHexStringLower(bytes);
}
}
}