feat(sentiment): complete pipeline overhaul with ONNX priority + LoRA retraining
- Added 30 new sources (5 RSS + 25 Telegram) for previously ZERO-coverage assets - Fixed model loading priority: ONNX > LoRA v2 > PyTorch > Mock - ONNX FinBERT (pre-trained on 1.2M financial docs) now PRIMARY - best for real-world text - LoRA v2 models trained on 518 carefully labeled samples (balanced Bearish/Bullish/Neutral) - Emotion LoRA v2 trained with weighted loss (greed/fear 2x, joy 1.5x) - 30 new sources: STX, FET, XTZ, ENJ, ETC, TRX, ONG, DASH, LTC, ZIL, NEAR, APT, SUI, ICP - Early stopping (patience=3) on both LoRA trainings - Human-in-the-loop verification CLI tool created - Disk-conscious: save_total_limit=1, adapters 6-8MB each Pipeline now correctly classifies: - BTC breaks 100k → +0.54 Bullish ✅ - Major hack → -0.23 Bearish ✅ - HODL → +0.91 Bullish ✅ - Rug pull → -0.30 Bearish ✅ - SEC sues → -0.30 Bearish ✅ - ETF approval → +0.32 Bullish ✅ - Whale accumulation → +0.31 Bullish ✅ Models: ONNX FinBERT (PRIORITY 1) + LoRA v2 adapters (6-8MB each) Training data: 518 carefully labeled samples (190 real + 328 synthetic) Early stopping (patience=3) on both FinBERT and DistilRoBERTa LoRA Emotion LoRA v2: weighted loss (greed/fear 2x, joy 1.5x) + early stopping
This commit is contained in:
@@ -0,0 +1,309 @@
|
||||
"""
|
||||
Comprehensive tests for EventClassifier.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
import asyncio
|
||||
import numpy as np
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
from sentiment_engine.nlp.event_classification import (
|
||||
EventClassifier, ONNXEventModel, EventType
|
||||
)
|
||||
from sentiment_engine.schemas.processed import EventClassification
|
||||
|
||||
|
||||
class TestEventClassifier:
|
||||
"""Tests for EventClassifier"""
|
||||
|
||||
@pytest.fixture
|
||||
def classifier(self):
|
||||
return EventClassifier()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_initialize_loads_model(self, classifier):
|
||||
"""Should initialize and load model"""
|
||||
await classifier.initialize()
|
||||
# May use ONNX, PyTorch, or keyword fallback
|
||||
|
||||
def test_classify_listing_keywords(self, classifier):
|
||||
"""Should classify listing events"""
|
||||
text = "Binance will list new token ABC tomorrow"
|
||||
events = classifier._classify_sync(text, ["ABC"])
|
||||
|
||||
assert len(events) >= 1
|
||||
assert events[0].event_type == EventType.LISTING
|
||||
|
||||
def test_classify_hack_keywords(self, classifier):
|
||||
"""Should classify hack events"""
|
||||
text = "Exchange hacked, millions stolen in exploit"
|
||||
events = classifier._classify_sync(text, ["BTC"])
|
||||
|
||||
assert len(events) >= 1
|
||||
assert events[0].event_type == EventType.HACK
|
||||
|
||||
def test_classify_regulatory_keywords(self, classifier):
|
||||
"""Should classify regulatory events"""
|
||||
text = "SEC investigation into crypto exchange"
|
||||
events = classifier._classify_sync(text, ["BTC"])
|
||||
|
||||
assert len(events) >= 1
|
||||
assert events[0].event_type == EventType.REGULATORY
|
||||
|
||||
def test_classify_upgrade_keywords(self, classifier):
|
||||
"""Should classify upgrade events"""
|
||||
text = "Ethereum Dencun upgrade activates Proto-Danksharding"
|
||||
events = classifier._classify_sync(text, ["ETH"])
|
||||
|
||||
assert len(events) >= 1
|
||||
assert events[0].event_type == EventType.UPGRADE
|
||||
|
||||
def test_classify_partnership_keywords(self, classifier):
|
||||
"""Should classify partnership events"""
|
||||
text = "JPMorgan and Coinbase announce strategic partnership"
|
||||
events = classifier._classify_sync(text, ["BTC"])
|
||||
|
||||
assert len(events) >= 1
|
||||
assert events[0].event_type == EventType.PARTNERSHIP
|
||||
|
||||
def test_classify_earnings_keywords(self, classifier):
|
||||
"""Should classify earnings events"""
|
||||
text = "Coinbase Q2 earnings beat estimates. Revenue up 50%"
|
||||
events = classifier._classify_sync(text, ["COIN"])
|
||||
|
||||
assert len(events) >= 1
|
||||
assert events[0].event_type == EventType.EARNINGS
|
||||
|
||||
def test_classify_macro_keywords(self, classifier):
|
||||
"""Should classify macro events"""
|
||||
text = "Fed cuts rates 50bps. Bitcoin rallies on macro pivot"
|
||||
events = classifier._classify_sync(text, ["BTC"])
|
||||
|
||||
assert len(events) >= 1
|
||||
assert events[0].event_type == EventType.MACRO
|
||||
|
||||
def test_classify_liquidation_keywords(self, classifier):
|
||||
"""Should classify liquidation events"""
|
||||
text = "Massive liquidation cascade wipes out $500M in longs"
|
||||
events = classifier._classify_sync(text, ["BTC"])
|
||||
|
||||
assert len(events) >= 1
|
||||
assert events[0].event_type == EventType.LIQUIDATION
|
||||
|
||||
def test_classify_whale_keywords(self, classifier):
|
||||
"""Should classify whale events"""
|
||||
text = "Whale moves 10,000 BTC after 5 years dormancy"
|
||||
events = classifier._classify_sync(text, ["BTC"])
|
||||
|
||||
assert len(events) >= 1
|
||||
assert events[0].event_type == EventType.WHALE
|
||||
|
||||
def test_classify_manipulation_keywords(self, classifier):
|
||||
"""Should classify manipulation events"""
|
||||
text = "Pump and dump scheme detected on low cap token"
|
||||
events = classifier._classify_sync(text, ["SHITCOIN"])
|
||||
|
||||
assert len(events) >= 1
|
||||
assert events[0].event_type == EventType.MANIPULATION
|
||||
|
||||
def test_classify_delisting_keywords(self, classifier):
|
||||
"""Should classify delisting events"""
|
||||
text = "Binance delists privacy coins XMR and ZEC"
|
||||
events = classifier._classify_sync(text, ["XMR", "ZEC"])
|
||||
|
||||
assert len(events) >= 1
|
||||
assert events[0].event_type == EventType.DELISTING
|
||||
|
||||
def test_classify_governance_keywords(self, classifier):
|
||||
"""Should classify governance events"""
|
||||
text = "Arbitrum DAO proposal passes with 95% approval"
|
||||
events = classifier._classify_sync(text, ["ARB"])
|
||||
|
||||
assert len(events) >= 1
|
||||
assert events[0].event_type == EventType.GOVERNANCE
|
||||
|
||||
def test_multiple_events_detected(self, classifier):
|
||||
"""Should detect multiple events in one text"""
|
||||
text = "SEC approves ETF and Binance lists new token"
|
||||
events = classifier._classify_sync(text, ["BTC"])
|
||||
|
||||
event_types = [e.event_type for e in events]
|
||||
assert EventType.REGULATORY in event_types
|
||||
assert EventType.LISTING in event_types
|
||||
|
||||
def test_confidence_calculation(self, classifier):
|
||||
"""Confidence should increase with more keyword matches"""
|
||||
text1 = "listing"
|
||||
text2 = "listing listed debut launch trading starts"
|
||||
|
||||
events1 = classifier._classify_sync(text1, ["BTC"])
|
||||
events2 = classifier._classify_sync(text2, ["BTC"])
|
||||
|
||||
assert events2[0].confidence >= events1[0].confidence
|
||||
|
||||
def test_severity_estimation(self, classifier):
|
||||
"""Severity should be higher for stronger language"""
|
||||
text1 = "hack"
|
||||
text2 = "major hack massive exploit emergency"
|
||||
|
||||
events1 = classifier._classify_sync(text1, ["BTC"])
|
||||
events2 = classifier._classify_sync(text2, ["BTC"])
|
||||
|
||||
assert events2[0].severity >= events1[0].severity
|
||||
|
||||
def test_find_involved_assets(self, classifier):
|
||||
"""Should find mentioned assets in text"""
|
||||
text = "BTC and ETH both surge on news"
|
||||
assets = ["BTC", "ETH", "SOL"]
|
||||
|
||||
involved = classifier._find_involved_assets(text, assets, EventType.LISTING)
|
||||
|
||||
assert "BTC" in involved
|
||||
assert "ETH" in involved
|
||||
assert "SOL" not in involved
|
||||
|
||||
def test_market_wide_events(self, classifier):
|
||||
"""Should assign MARKET for macro/regulatory without specific assets"""
|
||||
text = "Fed cuts rates 50bps"
|
||||
|
||||
involved = classifier._find_involved_assets(text, [], EventType.MACRO)
|
||||
|
||||
assert involved == ["MARKET"]
|
||||
|
||||
|
||||
class TestONNXEventModel:
|
||||
"""Tests for ONNXEventModel wrapper"""
|
||||
|
||||
def test_init_loads_session(self):
|
||||
"""Should load ONNX session"""
|
||||
with patch('onnxruntime.InferenceSession') as mock_session:
|
||||
mock_session.return_value.get_inputs.return_value = [
|
||||
MagicMock(name="input_ids"),
|
||||
MagicMock(name="attention_mask"),
|
||||
MagicMock(name="token_type_ids")
|
||||
]
|
||||
mock_session.return_value.get_outputs.return_value = [MagicMock(name="logits")]
|
||||
|
||||
with patch('transformers.AutoTokenizer.from_pretrained'):
|
||||
model = ONNXEventModel("path", "tokenizer_path")
|
||||
assert model.session is not None
|
||||
|
||||
def test_predict_returns_probabilities(self):
|
||||
"""predict should return probabilities summing to 1"""
|
||||
with patch('onnxruntime.InferenceSession') as mock_session:
|
||||
mock_session.return_value.get_inputs.return_value = [
|
||||
MagicMock(name="input_ids"),
|
||||
MagicMock(name="attention_mask"),
|
||||
MagicMock(name="token_type_ids")
|
||||
]
|
||||
mock_session.return_value.get_outputs.return_value = [MagicMock(name="logits")]
|
||||
mock_session.return_value.run.return_value = [np.array([[1.0, 2.0, 0.5] + [0.0]*9])]
|
||||
|
||||
with patch('transformers.AutoTokenizer.from_pretrained'):
|
||||
model = ONNXEventModel("path", "tokenizer_path")
|
||||
probs = model.predict(
|
||||
np.ones((1, 10), dtype=np.int64),
|
||||
np.ones((1, 10), dtype=np.int64)
|
||||
)
|
||||
assert len(probs) == 12
|
||||
assert abs(probs.sum() - 1.0) < 0.001
|
||||
|
||||
|
||||
class TestEventClassifierONNX:
|
||||
"""Tests for EventClassifier with ONNX model"""
|
||||
|
||||
@pytest.fixture
|
||||
def classifier(self):
|
||||
return EventClassifier()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_classify_with_onnx(self, classifier):
|
||||
"""Should use ONNX model when available"""
|
||||
with patch('onnxruntime.InferenceSession') as mock_session:
|
||||
mock_session.return_value.get_inputs.return_value = [
|
||||
MagicMock(name="input_ids"),
|
||||
MagicMock(name="attention_mask"),
|
||||
MagicMock(name="token_type_ids")
|
||||
]
|
||||
mock_session.return_value.get_outputs.return_value = [MagicMock(name="logits")]
|
||||
mock_session.return_value.run.return_value = [np.array([[0.8] + [0.02]*11])]
|
||||
|
||||
with patch('transformers.AutoTokenizer.from_pretrained') as mock_tokenizer:
|
||||
mock_tokenizer.return_value.return_value = {
|
||||
"input_ids": np.ones((1, 10), dtype=np.int64),
|
||||
"attention_mask": np.ones((1, 10), dtype=np.int64),
|
||||
"token_type_ids": np.zeros((1, 10), dtype=np.int64)
|
||||
}
|
||||
|
||||
classifier._onnx_model = ONNXEventModel("path", "tokenizer_path")
|
||||
classifier._use_onnx = True
|
||||
|
||||
events = await classifier.classify("SEC approves ETF", ["BTC"])
|
||||
|
||||
assert len(events) >= 1
|
||||
|
||||
|
||||
class TestEventClassifierEdgeCases:
|
||||
"""Tests for edge cases in event classification"""
|
||||
|
||||
@pytest.fixture
|
||||
def classifier(self):
|
||||
return EventClassifier()
|
||||
|
||||
def test_empty_text(self, classifier):
|
||||
"""Should handle empty text"""
|
||||
events = classifier._classify_sync("", ["BTC"])
|
||||
assert events == []
|
||||
|
||||
def test_no_keywords(self, classifier):
|
||||
"""Should return empty when no keywords match"""
|
||||
text = "The weather is nice today"
|
||||
events = classifier._classify_sync(text, ["BTC"])
|
||||
assert events == []
|
||||
|
||||
def test_case_insensitive_matching(self, classifier):
|
||||
"""Should match keywords case insensitively"""
|
||||
text = "HACK EXPLOIT STOLEN"
|
||||
events = classifier._classify_sync(text, ["BTC"])
|
||||
assert len(events) >= 1
|
||||
assert events[0].event_type == EventType.HACK
|
||||
|
||||
def test_partial_keyword_matching(self, classifier):
|
||||
"""Should match partial keywords"""
|
||||
text = "hacking attempt detected"
|
||||
events = classifier._classify_sync(text, ["BTC"])
|
||||
# "hacking" contains "hack"
|
||||
assert len(events) >= 1
|
||||
|
||||
def test_asset_not_in_text(self, classifier):
|
||||
"""Should not include assets not in text"""
|
||||
text = "BTC surges"
|
||||
assets = ["BTC", "ETH", "SOL"]
|
||||
|
||||
involved = classifier._find_involved_assets(text, assets, EventType.LISTING)
|
||||
|
||||
assert "BTC" in involved
|
||||
assert "ETH" not in involved
|
||||
assert "SOL" not in involved
|
||||
|
||||
def test_overlapping_keywords(self, classifier):
|
||||
"""Should handle overlapping keyword categories"""
|
||||
text = "SEC hack investigation" # both regulatory and hack
|
||||
events = classifier._classify_sync(text, ["BTC"])
|
||||
|
||||
event_types = [e.event_type for e in events]
|
||||
# Should detect both or the stronger one
|
||||
assert len(events) >= 1
|
||||
|
||||
def test_confidence_threshold(self, classifier):
|
||||
"""Should filter low confidence events"""
|
||||
# Single weak keyword match
|
||||
text = "maybe listing soon"
|
||||
events = classifier._classify_sync(text, ["BTC"])
|
||||
# Confidence might be below threshold
|
||||
for e in events:
|
||||
assert e.confidence >= 0.3
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
pytest.main([__file__, "-v"])
|
||||
Reference in New Issue
Block a user