Integrated FinBERT model.
This commit is contained in:
parent
df759d0ead
commit
7f8d6a8131
@ -12,6 +12,9 @@ builder.Services.AddSwaggerGen();
|
||||
builder.Services.AddSingleton<IHistoricalMarketDataCache, PostgresHistoricalMarketDataCache>();
|
||||
builder.Services.AddSingleton<MarketDataService>();
|
||||
builder.Services.AddSingleton<IUserInterestStore, PostgresUserInterestStore>();
|
||||
builder.Services.AddSingleton<FinBertScoringService>();
|
||||
builder.Services.AddSingleton<INewsSentimentStore, PostgresNewsSentimentStore>();
|
||||
builder.Services.AddSingleton<NewsSentimentService>();
|
||||
|
||||
builder.Services.AddCors(options =>
|
||||
{
|
||||
|
||||
@ -13,6 +13,7 @@
|
||||
|
||||
|
||||
<ItemGroup>
|
||||
<PackageReference Include="Microsoft.ML.OnnxRuntime" Version="1.22.0" />
|
||||
<PackageReference Include="Npgsql" Version="10.0.2" />
|
||||
<PackageReference Include="Swashbuckle.AspNetCore" Version="10.1.5" />
|
||||
</ItemGroup>
|
||||
|
||||
46
Controller/NewsSentimentController.cs
Normal file
46
Controller/NewsSentimentController.cs
Normal 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() });
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
14
DatabaseHandler/NewsSentimentRecord.cs
Normal file
14
DatabaseHandler/NewsSentimentRecord.cs
Normal 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;
|
||||
}
|
||||
}
|
||||
123
DatabaseHandler/PostgresNewsSentimentStore.cs
Normal file
123
DatabaseHandler/PostgresNewsSentimentStore.cs
Normal 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;
|
||||
}
|
||||
}
|
||||
}
|
||||
19
Interface/INewsSentimentStore.cs
Normal file
19
Interface/INewsSentimentStore.cs
Normal 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);
|
||||
}
|
||||
}
|
||||
220
ServiceHandler/FinBertScoringService.cs
Normal file
220
ServiceHandler/FinBertScoringService.cs
Normal 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();
|
||||
}
|
||||
}
|
||||
}
|
||||
101
ServiceHandler/NewsSentimentService.cs
Normal file
101
ServiceHandler/NewsSentimentService.cs
Normal 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);
|
||||
}
|
||||
}
|
||||
}
|
||||
Loading…
Reference in New Issue
Block a user