Files
sentiment-engine/sentiment_engine/tests/unit/test_integrity_onnx_integration.py

469 lines
19 KiB
Python
Raw Normal View History

"""Integrity tests for ONNX model integration and component coupling"""
import pytest
import numpy as np
from pathlib import Path
# Test that ONNX models exist and can be loaded
class TestONNXModelAvailability:
"""Verify ONNX models are present and loadable"""
MODEL_PATHS = {
"finbert": "models/onnx/finbert/model.onnx",
"bert-events": "models/onnx/bert-base-event/model.onnx",
"distilroberta-emotion": "models/onnx/distilroberta-emotion/model.onnx",
"minilm": "models/onnx/minilm-l6-v2/model.onnx",
}
LABEL_MAP_PATHS = {
"finbert": "models/onnx/finbert/label_map.json",
"bert-events": "models/onnx/bert-base-event/label_map.json",
"distilroberta-emotion": "models/onnx/distilroberta-emotion/label_map.json",
}
ID2LABEL_PATHS = {
"finbert": "models/onnx/finbert/id2label.json",
"bert-events": "models/onnx/bert-base-event/id2label.json",
"distilroberta-emotion": "models/onnx/distilroberta-emotion/id2label.json",
}
def test_all_onnx_models_exist(self):
"""All ONNX model files must exist"""
for name, path in self.MODEL_PATHS.items():
assert Path(path).exists(), f"ONNX model {name} not found at {path}"
assert Path(path).stat().st_size > 1024, f"ONNX model {name} appears empty"
def test_all_label_maps_exist(self):
"""All label maps must exist"""
for name, path in self.LABEL_MAP_PATHS.items():
assert Path(path).exists(), f"Label map {name} not found at {path}"
def test_all_id2label_maps_exist(self):
"""All id2label maps must exist"""
for name, path in self.ID2LABEL_PATHS.items():
assert Path(path).exists(), f"id2label map {name} not found at {path}"
def test_tokenizers_exist(self):
"""All tokenizers must exist alongside models"""
for name, path in self.MODEL_PATHS.items():
tokenizer_dir = Path(path).parent
assert (tokenizer_dir / "tokenizer.json").exists(), f"Tokenizer missing for {name}"
assert (tokenizer_dir / "tokenizer_config.json").exists(), f"Tokenizer config missing for {name}"
class TestONNXModelInference:
"""Verify ONNX models run inference correctly"""
@pytest.fixture(scope="class")
def onnx_session(self):
"""Create ONNX Runtime sessions for all models"""
import onnxruntime as ort
sessions = {}
for name, path in TestONNXModelAvailability.MODEL_PATHS.items():
if Path(path).exists():
sessions[name] = ort.InferenceSession(path, providers=['CPUExecutionProvider'])
return sessions
def test_finbert_inference_shape(self, onnx_session):
"""FinBERT sentiment model produces correct output shape"""
if "finbert" not in onnx_session:
pytest.skip("FinBERT ONNX model not available")
session = onnx_session["finbert"]
input_ids = np.ones((1, 64), dtype=np.int64)
attention_mask = np.ones((1, 64), dtype=np.int64)
token_type_ids = np.zeros((1, 64), dtype=np.int64) # Required by FinBERT
outputs = session.run(None, {
"input_ids": input_ids,
"attention_mask": attention_mask,
"token_type_ids": token_type_ids
})
assert len(outputs) >= 1
logits = outputs[0]
assert logits.shape == (1, 3), f"Expected (1, 3) for FinBERT, got {logits.shape}"
def test_bert_events_inference_shape(self, onnx_session):
"""BERT events model produces correct output shape (12 events)"""
if "bert-events" not in onnx_session:
pytest.skip("BERT events ONNX model not available")
session = onnx_session["bert-events"]
input_ids = np.ones((1, 64), dtype=np.int64)
attention_mask = np.ones((1, 64), dtype=np.int64)
token_type_ids = np.zeros((1, 64), dtype=np.int64) # Required by BERT
outputs = session.run(None, {
"input_ids": input_ids,
"attention_mask": attention_mask,
"token_type_ids": token_type_ids
})
assert len(outputs) >= 1
logits = outputs[0]
assert logits.shape == (1, 12), f"Expected (1, 12) for BERT events, got {logits.shape}"
def test_distilroberta_emotion_inference_shape(self, onnx_session):
"""DistilRoBERTa emotion model produces correct output shape (6 emotions)"""
if "distilroberta-emotion" not in onnx_session:
pytest.skip("DistilRoBERTa emotion ONNX model not available")
session = onnx_session["distilroberta-emotion"]
input_ids = np.ones((1, 64), dtype=np.int64)
attention_mask = np.ones((1, 64), dtype=np.int64)
# DistilRoBERTa does NOT use token_type_ids
outputs = session.run(None, {
"input_ids": input_ids,
"attention_mask": attention_mask
})
assert len(outputs) >= 1
logits = outputs[0]
assert logits.shape == (1, 6), f"Expected (1, 6) for DistilRoBERTa emotion, got {logits.shape}"
def test_minilm_embedding_shape(self, onnx_session):
"""MiniLM produces embeddings of correct shape (768 dims for MiniLM-L6-v2)"""
if "minilm" not in onnx_session:
pytest.skip("MiniLM ONNX model not available")
session = onnx_session["minilm"]
input_ids = np.ones((1, 64), dtype=np.int64)
attention_mask = np.ones((1, 64), dtype=np.int64)
outputs = session.run(None, {
"input_ids": input_ids,
"attention_mask": attention_mask
})
# MiniLM outputs: last_hidden_state, pooler_output
assert len(outputs) >= 1
last_hidden = outputs[0]
assert last_hidden.shape == (1, 64, 768), f"Expected (1, 64, 768) for MiniLM, got {last_hidden.shape}"
class TestSentimentEmotionONNXIntegration:
"""Test SentimentEmotionAnalyzer with ONNX models"""
@pytest.fixture
def analyzer(self):
from sentiment_engine.nlp.sentiment_emotion import SentimentEmotionAnalyzer
return SentimentEmotionAnalyzer()
@pytest.mark.asyncio
async def test_initialize_loads_onnx(self, analyzer):
"""Analyzer should initialize and load ONNX models when available"""
await analyzer.initialize()
# Should have loaded ONNX models (not mock)
assert hasattr(analyzer, '_model')
assert hasattr(analyzer, '_tokenizer')
@pytest.mark.asyncio
async def test_analyze_returns_sentiment_scores(self, analyzer):
"""Analyze returns proper SentimentScores for each asset"""
from sentiment_engine.schemas.processed import SentimentScores
await analyzer.initialize()
text = "Bitcoin surges to new all-time high!"
asset_mentions = [{"asset_id": "BTC", "span": (0, 3)}]
sentiment_results, emotion_results = await analyzer.analyze(text, asset_mentions)
assert "BTC" in sentiment_results
assert isinstance(sentiment_results["BTC"], SentimentScores)
assert -1.0 <= sentiment_results["BTC"].polarity <= 1.0
assert 0.0 <= sentiment_results["BTC"].confidence <= 1.0
@pytest.mark.asyncio
async def test_analyze_returns_emotion_scores(self, analyzer):
"""Analyze returns proper EmotionScores for each asset"""
from sentiment_engine.schemas.processed import EmotionScores
await analyzer.initialize()
text = "Bitcoin surges to new all-time high! To the moon!"
asset_mentions = [{"asset_id": "BTC", "span": (0, 3)}]
sentiment_results, emotion_results = await analyzer.analyze(text, asset_mentions)
assert "BTC" in emotion_results
assert isinstance(emotion_results["BTC"], EmotionScores)
assert 0.0 <= emotion_results["BTC"].joy <= 1.0
assert 0.0 <= emotion_results["BTC"].fear <= 1.0
assert 0.0 <= emotion_results["BTC"].anger <= 1.0
assert 0.0 <= emotion_results["BTC"].greed <= 1.0
assert 0.0 <= emotion_results["BTC"].sadness <= 1.0
class TestEventClassifierONNXIntegration:
"""Test EventClassifier with ONNX model"""
@pytest.fixture
def classifier(self):
from sentiment_engine.nlp.event_classification import EventClassifier
return EventClassifier()
@pytest.mark.asyncio
async def test_initialize_loads_onnx(self, classifier):
"""Classifier should initialize and load ONNX model when available"""
await classifier.initialize()
# Should have loaded ONNX model
assert hasattr(classifier, '_onnx_model')
@pytest.mark.asyncio
async def test_classify_returns_event_classifications(self, classifier):
"""Classify returns proper EventClassification objects"""
from sentiment_engine.schemas.processed import EventClassification, EventType
await classifier.initialize()
text = "SEC approves spot Bitcoin ETF for trading"
asset_mentions = ["BTC"]
events = await classifier.classify(text, asset_mentions)
assert len(events) >= 1
assert all(isinstance(e, EventClassification) for e in events)
assert all(hasattr(e, 'event_type') for e in events)
assert all(hasattr(e, 'confidence') for e in events)
assert all(hasattr(e, 'severity') for e in events)
@pytest.mark.asyncio
async def test_classify_regulatory_event(self, classifier):
"""Classify correctly identifies regulatory events (keyword fallback works)"""
from sentiment_engine.schemas.processed import EventType
await classifier.initialize()
text = "SEC files lawsuit against exchange for unregistered securities"
asset_mentions = ["BTC"]
events = await classifier.classify(text, asset_mentions)
regulatory_events = [e for e in events if e.event_type == EventType.REGULATORY]
assert len(regulatory_events) >= 1
assert regulatory_events[0].confidence > 0.3
class TestEntityExtractorIntegration:
"""Test EntityExtractor with spaCy NER if available"""
@pytest.fixture
def extractor(self):
from sentiment_engine.nlp.entity_extraction import EntityExtractor, AssetMapper
return EntityExtractor(AssetMapper())
@pytest.mark.asyncio
async def test_extract_all_returns_entities(self, extractor):
"""Extract_all returns proper EntityExtraction objects"""
from sentiment_engine.schemas.processed import EntityExtraction
await extractor.initialize()
text = "BTC and ETH are pumping. Vitalik buys more ETH."
entities = await extractor.extract_all(text)
assert len(entities) >= 2
assert all(isinstance(e, EntityExtraction) for e in entities)
asset_ids = [e.asset_id for e in entities]
assert "BTC" in asset_ids
assert "ETH" in asset_ids
@pytest.mark.asyncio
async def test_extract_all_deduplicates(self, extractor):
"""Extract_all deduplicates overlapping mentions"""
await extractor.initialize()
text = "BTC BTC BTC"
entities = await extractor.extract_all(text)
# Should only have one BTC mention after deduplication
btc_entities = [e for e in entities if e.asset_id == "BTC"]
assert len(btc_entities) == 1
class TestTemporalAnchorerIntegration:
"""Test TemporalAnchorer functionality"""
@pytest.fixture
def anchorer(self):
from sentiment_engine.nlp.temporal import TemporalAnchorer
return TemporalAnchorer()
def test_anchor_returns_temporal_anchor(self, anchorer):
"""Anchor returns proper TemporalAnchor object"""
from sentiment_engine.schemas.processed import TemporalAnchor
text = "Breaking: BTC crashes now!"
anchor = anchorer.anchor(text, None)
assert isinstance(anchor, TemporalAnchor)
assert anchor.time_horizon in ["immediate", "near", "medium", "long"]
assert isinstance(anchor.is_breaking, bool)
assert isinstance(anchor.is_scheduled, bool)
def test_anchor_detects_breaking(self, anchorer):
"""Anchor correctly detects breaking news"""
text = "Breaking news: BTC crashes!"
anchor = anchorer.anchor(text, None)
assert anchor.is_breaking is True
def test_anchor_detects_scheduled(self, anchorer):
"""Anchor correctly detects scheduled events"""
text = "Scheduled for 2024-01-15: ETH upgrade"
anchor = anchorer.anchor(text, None)
assert anchor.is_scheduled is True
class TestNLPProcessingPipelineIntegration:
"""End-to-end tests for the full NLP pipeline"""
@pytest.fixture
def pipeline(self):
from sentiment_engine.nlp.pipeline import NLPProcessingPipeline
return NLPProcessingPipeline()
@pytest.mark.asyncio
async def test_pipeline_initializes_all_components(self, pipeline):
"""Pipeline initializes all components successfully"""
await pipeline.initialize()
assert pipeline._initialized is True
assert pipeline.entity_extractor is not None
assert pipeline.sentiment_analyzer is not None
assert pipeline.event_classifier is not None
assert pipeline.temporal_anchorer is not None
assert pipeline.credibility_scorer is not None
@pytest.mark.asyncio
async def test_process_returns_processed_item(self, pipeline):
"""Process returns complete ProcessedItem with all fields"""
from sentiment_engine.schemas.payload import NormalizedPayload, SourceType
from sentiment_engine.schemas.processed import ProcessedItem
await pipeline.initialize()
# Use text with clear event keywords to trigger keyword fallback
payload = NormalizedPayload(
source_id="test_source",
source_type=SourceType.NEWS,
source_credibility_base=0.8,
ingest_ts=1700000000.0,
publish_ts=1700000000.0,
content_length=150,
raw_text="SEC approves spot Bitcoin ETF for trading today. Major regulatory decision.",
metadata={}
)
result = await pipeline.process(payload)
assert isinstance(result, ProcessedItem)
assert result.payload_id is not None
assert result.source_id == "test_source"
assert len(result.entities) >= 1
assert len(result.sentiment_per_asset) >= 1
assert len(result.emotions_per_asset) >= 1
assert len(result.events) >= 1, f"Expected at least 1 event, got {len(result.events)}"
assert result.temporal is not None
assert result.credibility is not None
assert result.processing_latency_ms > 0
class TestComponentConfigurationCoupling:
"""Test that component configurations are properly coupled"""
def test_onnx_model_paths_match_config(self):
"""ONNX model paths in code match expected locations"""
from sentiment_engine.nlp.sentiment_emotion import SentimentEmotionAnalyzer
from sentiment_engine.nlp.event_classification import EventClassifier
analyzer = SentimentEmotionAnalyzer()
classifier = EventClassifier()
# Check that the paths used in initialization match actual model locations
import os
finbert_onnx = Path("models/onnx/finbert/model.onnx")
emotion_onnx = Path("models/onnx/distilroberta-emotion/model.onnx")
event_onnx = Path("models/onnx/bert-base-event/model.onnx")
assert finbert_onnx.exists(), "FinBERT ONNX path mismatch"
assert emotion_onnx.exists(), "Emotion ONNX path mismatch"
assert event_onnx.exists(), "Event ONNX path mismatch"
def test_label_maps_match_model_outputs(self):
"""Label maps have correct number of labels matching model outputs"""
import json
# FinBERT: 3 labels
with open("models/onnx/finbert/label_map.json") as f:
labels = json.load(f)
assert len(labels) == 3, f"FinBERT label map has {len(labels)} labels, expected 3"
# BERT Events: 12 labels
with open("models/onnx/bert-base-event/label_map.json") as f:
labels = json.load(f)
assert len(labels) == 12, f"BERT Events label map has {len(labels)} labels, expected 12"
# DistilRoBERTa Emotion: 6 labels
with open("models/onnx/distilroberta-emotion/label_map.json") as f:
labels = json.load(f)
assert len(labels) == 6, f"Emotion label map has {len(labels)} labels, expected 6"
def test_id2label_maps_match_label_maps(self):
"""id2label maps match label maps"""
import json
for model in ["finbert", "bert-base-event", "distilroberta-emotion"]:
with open(f"models/onnx/{model}/label_map.json") as f:
label_map = json.load(f)
with open(f"models/onnx/{model}/id2label.json") as f:
id2label = json.load(f)
# Both should have same labels (order might differ)
assert set(label_map.values()) == set(id2label.values()), f"{model} label mismatch"
class TestFineTunedModelVersioning:
"""Verify fine-tuned models have correct version metadata"""
def test_finetuned_models_exist(self):
"""Fine-tuned PyTorch models exist"""
model_dirs = [
"models/finbert-crypto-sentiment",
"models/bert-crypto-events",
"models/distilroberta-crypto-emotion",
]
for dir_path in model_dirs:
assert Path(dir_path).exists(), f"Fine-tuned model dir {dir_path} missing"
assert (Path(dir_path) / "model.safetensors").exists(), f"model.safetensors missing in {dir_path}"
assert (Path(dir_path) / "config.json").exists(), f"config.json missing in {dir_path}"
def test_finetuned_model_configs_match_onnx(self):
"""Fine-tuned model configs are compatible with ONNX exports"""
import json
for model_dir, onnx_dir in [
("models/finbert-crypto-sentiment", "models/onnx/finbert"),
("models/bert-crypto-events", "models/onnx/bert-base-event"),
("models/distilroberta-crypto-emotion", "models/onnx/distilroberta-emotion"),
]:
with open(Path(model_dir) / "config.json") as f:
pt_config = json.load(f)
with open(Path(onnx_dir) / "config.json") as f:
onnx_config = json.load(f)
# Key configs should match
assert pt_config.get("num_labels") == onnx_config.get("num_labels"), \
f"num_labels mismatch for {model_dir}"
if __name__ == "__main__":
pytest.main([__file__, "-v"])