469 lines
19 KiB
Python
469 lines
19 KiB
Python
"""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"])
|