Mohamed424's picture
added categorizations
62e9f26
Raw History Blame Contribute Delete
27.4 kB
import asyncio
from motor.motor_asyncio import AsyncIOMotorClient
from app.config import config
from typing import Optional, List, Dict, Any
from datetime import datetime
class DatabaseManager:
def __init__(self):
self.client = None
self.database = None
self.collection = None
self.is_connected = False
async def initialize(self):
"""Initialize MongoDB connection if available"""
try:
if config.ENABLE_MONGODB and config.MONGODB_URL:
self.client = AsyncIOMotorClient(config.MONGODB_URL)
self.database = self.client[config.DATABASE_NAME]
self.collection = self.database[config.COLLECTION_NAME]
self.is_connected = True
print("MongoDB configured successfully")
else:
print("MongoDB not configured - predictions will not be persisted")
except asyncio.TimeoutError:
print("MongoDB connection timeout - running without database")
self.is_connected = False
except Exception as e:
print(f"MongoDB configuration failed: {e}")
print("Running without database - predictions will not be persisted")
self.is_connected = False
async def save_prediction(self, prediction_data: Dict[str, Any]) -> bool:
"""Save prediction to database if available"""
if not self.is_connected:
return False
try:
await self.collection.insert_one({
**prediction_data,
"_id": prediction_data["id"]
})
return True
except Exception as e:
print(f"Failed to save prediction: {e}")
return False
async def get_prediction(self, prediction_id: str) -> Optional[Dict[str, Any]]:
"""Get prediction by ID"""
if not self.is_connected:
return None
try:
doc = await self.collection.find_one({"_id": prediction_id})
if doc:
doc["id"] = doc.pop("_id")
return doc
except Exception as e:
print(f"Failed to get prediction: {e}")
return None
async def get_predictions_paginated(
self,
page: int = 1,
limit: int = 10,
filters: Optional[Dict[str, Any]] = None,
sort_field: str = "created_at",
sort_direction: int = -1
) -> tuple:
"""Get paginated predictions with filtering"""
if not self.is_connected:
return [], 0
try:
filter_query = filters or {}
skip = (page - 1) * limit
# Get total count
total_count = await self.collection.count_documents(filter_query)
# Get paginated results
cursor = self.collection.find(filter_query).sort(sort_field, sort_direction).skip(skip).limit(limit)
predictions = []
async for doc in cursor:
doc_dict = dict(doc)
doc_dict["id"] = doc_dict.pop("_id")
predictions.append(doc_dict)
return predictions, total_count
except Exception as e:
print(f"Failed to get predictions: {e}")
return [], 0
async def delete_prediction(self, prediction_id: str) -> bool:
"""Delete prediction by ID"""
if not self.is_connected:
return False
try:
result = await self.collection.delete_one({"_id": prediction_id})
return result.deleted_count > 0
except Exception as e:
print(f"Failed to delete prediction: {e}")
return False
async def get_statistics(self, filter_query: Optional[Dict[str, Any]] = None, include_enhanced: bool = True) -> Dict[str, Any]:
"""
Get prediction statistics with optional enhanced metrics
Args:
filter_query: MongoDB filter for date ranges, etc.
include_enhanced: Whether to include enhanced metrics (trends, analytics, performance, health)
Returns:
Dictionary with statistics data
"""
if not self.is_connected:
return {
"total_predictions": 0,
"stroke_detected": 0,
"normal_cases": 0,
"stroke_detection_rate": 0,
"average_confidence": 0
}
try:
filter_query = filter_query or {}
# Basic statistics
total_predictions = await self.collection.count_documents(filter_query)
stroke_detected = await self.collection.count_documents({**filter_query, "has_stroke": True})
# Get average confidence
pipeline = [
{"$match": filter_query},
{"$group": {"_id": None, "avg_confidence": {"$avg": "$confidence_score"}}}
]
avg_confidence_result = await self.collection.aggregate(pipeline).to_list(1)
avg_confidence = avg_confidence_result[0]["avg_confidence"] if avg_confidence_result else 0
basic_stats = {
"total_predictions": total_predictions,
"stroke_detected": stroke_detected,
"normal_cases": total_predictions - stroke_detected,
"stroke_detection_rate": (stroke_detected / total_predictions * 100) if total_predictions > 0 else 0,
"average_confidence": avg_confidence if avg_confidence else 0
}
if not include_enhanced:
return basic_stats
# Enhanced statistics
trends = await self._calculate_trends(filter_query, basic_stats)
detailed_analytics = await self._calculate_detailed_analytics(filter_query)
performance_metrics = await self._calculate_performance_metrics(filter_query)
stroke_type_breakdown = await self._calculate_stroke_type_breakdown(filter_query)
stroke_type_accuracy = await self._calculate_stroke_type_accuracy(filter_query)
return {
**basic_stats,
"stroke_type_breakdown": stroke_type_breakdown,
"stroke_type_accuracy": stroke_type_accuracy,
"trends": trends,
"detailed_analytics": detailed_analytics,
"performance_metrics": performance_metrics
}
except Exception as e:
print(f"Failed to get statistics: {e}")
return {
"total_predictions": 0,
"stroke_detected": 0,
"normal_cases": 0,
"stroke_detection_rate": 0,
"average_confidence": 0
}
async def _calculate_trends(self, filter_query: Dict[str, Any], current_stats: Dict[str, Any]) -> Dict[str, float]:
"""Calculate trends by comparing current period to previous period"""
try:
# Extract date range from filter
if "created_at" not in filter_query:
# No date filter, can't calculate trends reliably
return {
"total_predictions": 0.0,
"stroke_detected": 0.0,
"normal_cases": 0.0,
"detection_rate": 0.0,
"average_confidence": 0.0
}
current_date_filter = filter_query["created_at"]
if "$gte" not in current_date_filter:
return {
"total_predictions": 0.0,
"stroke_detected": 0.0,
"normal_cases": 0.0,
"detection_rate": 0.0,
"average_confidence": 0.0
}
# Calculate previous period
start_date = current_date_filter["$gte"]
end_date = current_date_filter.get("$lte", datetime.now())
period_length = end_date - start_date
previous_start = start_date - period_length
previous_end = start_date
# Get previous period stats
previous_filter = {
**{k: v for k, v in filter_query.items() if k != "created_at"},
"created_at": {
"$gte": previous_start,
"$lte": previous_end
}
}
prev_total = await self.collection.count_documents(previous_filter)
prev_stroke = await self.collection.count_documents({**previous_filter, "has_stroke": True})
pipeline = [
{"$match": previous_filter},
{"$group": {"_id": None, "avg_confidence": {"$avg": "$confidence_score"}}}
]
prev_conf_result = await self.collection.aggregate(pipeline).to_list(1)
prev_confidence = prev_conf_result[0]["avg_confidence"] if prev_conf_result else 0
# Calculate percentage changes
def calc_change(current, previous):
if previous == 0:
return 100.0 if current > 0 else 0.0
return ((current - previous) / previous) * 100
return {
"total_predictions": round(calc_change(current_stats["total_predictions"], prev_total), 1),
"stroke_detected": round(calc_change(current_stats["stroke_detected"], prev_stroke), 1),
"normal_cases": round(calc_change(current_stats["normal_cases"], prev_total - prev_stroke), 1),
"detection_rate": round(calc_change(current_stats["stroke_detection_rate"],
(prev_stroke / prev_total * 100) if prev_total > 0 else 0), 1),
"average_confidence": round(calc_change(current_stats["average_confidence"], prev_confidence if prev_confidence else 0), 1)
}
except Exception as e:
print(f"Failed to calculate trends: {e}")
return {
"total_predictions": 0.0,
"stroke_detected": 0.0,
"normal_cases": 0.0,
"detection_rate": 0.0,
"average_confidence": 0.0
}
async def _calculate_detailed_analytics(self, filter_query: Dict[str, Any]) -> Dict[str, float]:
"""Calculate detailed analytics metrics using real ground truth when available"""
try:
# Get processing time statistics
pipeline = [
{"$match": filter_query},
{"$group": {
"_id": None,
"avg_processing_time": {"$avg": "$processing_time"},
"total": {"$sum": 1},
"stroke_count": {
"$sum": {"$cond": ["$has_stroke", 1, 0]}
}
}}
]
result = await self.collection.aggregate(pipeline).to_list(1)
if not result:
return self._default_detailed_analytics()
data = result[0]
avg_processing_time = data.get("avg_processing_time", 0)
# Check if we have ground truth data
ground_truth_filter = {**filter_query, "doctor_expectation": {"$ne": None}}
ground_truth_count = await self.collection.count_documents(ground_truth_filter)
if ground_truth_count > 0:
# Calculate real FP/FN rates using doctor expectations
confusion_pipeline = [
{"$match": ground_truth_filter},
{"$group": {
"_id": {
"predicted": "$has_stroke",
"actual": "$doctor_expectation"
},
"count": {"$sum": 1}
}}
]
confusion_result = await self.collection.aggregate(confusion_pipeline).to_list(10)
tp = tn = fp = fn = 0
for item in confusion_result:
predicted = item["_id"]["predicted"]
actual = item["_id"]["actual"]
count = item["count"]
if actual == True and predicted == True:
tp = count
elif actual == False and predicted == False:
tn = count
elif actual == False and predicted == True:
fp = count # False Positive
elif actual == True and predicted == False:
fn = count # False Negative
total_with_gt = tp + tn + fp + fn
fp_rate = (fp / total_with_gt * 100) if total_with_gt > 0 else 0
fn_rate = (fn / total_with_gt * 100) if total_with_gt > 0 else 0
detection_accuracy = ((tp + tn) / total_with_gt * 100) if total_with_gt > 0 else 0
else:
# Estimate FP/FN rates based on confidence distribution
low_confidence_threshold = 0.6
low_conf_pipeline = [
{"$match": {**filter_query, "confidence_score": {"$lt": low_confidence_threshold}}},
{"$count": "count"}
]
low_conf_result = await self.collection.aggregate(low_conf_pipeline).to_list(1)
low_conf_count = low_conf_result[0]["count"] if low_conf_result else 0
total = data.get("total", 1)
# Estimate: low confidence predictions are potential false positives/negatives
fp_rate = (low_conf_count / total * 100) if total > 0 else 0
fn_rate = fp_rate * 0.8 # Assume FN is slightly lower than FP
detection_accuracy = (data.get("stroke_count", 0) / total * 100) if total > 0 else 0
return {
"detection_accuracy": round(min(detection_accuracy, 100.0), 1),
"detection_accuracy_trend": 0.0, # Would need historical data
"average_processing_time": round(avg_processing_time, 1),
"processing_time_trend": 0.0, # Would need historical data
"false_positive_rate": round(min(fp_rate, 100.0), 1),
"false_positive_trend": 0.0,
"false_negative_rate": round(min(fn_rate, 100.0), 1),
"false_negative_trend": 0.0
}
except Exception as e:
print(f"Failed to calculate detailed analytics: {e}")
return self._default_detailed_analytics()
def _default_detailed_analytics(self) -> Dict[str, float]:
"""Return default detailed analytics"""
return {
"detection_accuracy": 0.0,
"detection_accuracy_trend": 0.0,
"average_processing_time": 0.0,
"processing_time_trend": 0.0,
"false_positive_rate": 0.0,
"false_positive_trend": 0.0,
"false_negative_rate": 0.0,
"false_negative_trend": 0.0
}
async def _calculate_performance_metrics(self, filter_query: Dict[str, Any]) -> Dict[str, float]:
"""
Calculate ML model performance metrics using real doctor expectations when available.
Falls back to estimates based on confidence scores if no ground truth.
"""
try:
# First, check if we have any cases with doctor expectations
ground_truth_filter = {**filter_query, "doctor_expectation": {"$ne": None}}
ground_truth_count = await self.collection.count_documents(ground_truth_filter)
if ground_truth_count > 0:
# Use real ground truth data
return await self._calculate_real_performance_metrics(ground_truth_filter)
else:
# Fall back to estimates
return await self._calculate_estimated_performance_metrics(filter_query)
except Exception as e:
print(f"Failed to calculate performance metrics: {e}")
return self._default_performance_metrics()
async def _calculate_real_performance_metrics(self, filter_query: Dict[str, Any]) -> Dict[str, float]:
"""Calculate real performance metrics using doctor expectations as ground truth"""
try:
pipeline = [
{"$match": filter_query},
{"$facet": {
"confusion_matrix": [
{"$group": {
"_id": {
"predicted": "$has_stroke",
"actual": "$doctor_expectation"
},
"count": {"$sum": 1}
}}
],
"avg_confidence": [
{"$group": {"_id": None, "avg": {"$avg": "$confidence_score"}}}
]
}}
]
result = await self.collection.aggregate(pipeline).to_list(1)
if not result or not result[0]["confusion_matrix"]:
return self._default_performance_metrics()
# Parse confusion matrix
tp = tn = fp = fn = 0
for item in result[0]["confusion_matrix"]:
predicted = item["_id"]["predicted"]
actual = item["_id"]["actual"]
count = item["count"]
if actual == True and predicted == True:
tp = count # True Positive
elif actual == False and predicted == False:
tn = count # True Negative
elif actual == False and predicted == True:
fp = count # False Positive
elif actual == True and predicted == False:
fn = count # False Negative
# Calculate metrics
sensitivity = (tp / (tp + fn)) if (tp + fn) > 0 else 0 # True Positive Rate
specificity = (tn / (tn + fp)) if (tn + fp) > 0 else 0 # True Negative Rate
ppv = (tp / (tp + fp)) if (tp + fp) > 0 else 0 # Positive Predictive Value / Precision
npv = (tn / (tn + fn)) if (tn + fn) > 0 else 0 # Negative Predictive Value
# F1 Score
precision = ppv
recall = sensitivity
f1 = (2 * precision * recall / (precision + recall)) if (precision + recall) > 0 else 0
# AUC-ROC estimate (would need ROC curve calculation for exact value)
avg_conf = result[0]["avg_confidence"][0]["avg"] if result[0]["avg_confidence"] else 0.5
auc_roc = min(0.5 + (avg_conf * 0.5), 0.99)
return {
"sensitivity": round(sensitivity * 100, 1),
"specificity": round(specificity * 100, 1),
"positive_predictive_value": round(ppv * 100, 1),
"negative_predictive_value": round(npv * 100, 1),
"f1_score": round(f1 * 100, 1),
"auc_roc": round(auc_roc, 3)
}
except Exception as e:
print(f"Failed to calculate real performance metrics: {e}")
return self._default_performance_metrics()
async def _calculate_estimated_performance_metrics(self, filter_query: Dict[str, Any]) -> Dict[str, float]:
"""Calculate estimated performance metrics based on confidence scores (fallback method)"""
try:
pipeline = [
{"$match": filter_query},
{"$facet": {
"total": [{"$count": "count"}],
"stroke_predictions": [
{"$match": {"has_stroke": True}},
{"$count": "count"}
],
"high_confidence_stroke": [
{"$match": {"has_stroke": True, "confidence_score": {"$gte": 0.85}}},
{"$count": "count"}
],
"high_confidence_normal": [
{"$match": {"has_stroke": False, "confidence_score": {"$gte": 0.85}}},
{"$count": "count"}
],
"avg_confidence": [
{"$group": {"_id": None, "avg": {"$avg": "$confidence_score"}}}
]
}}
]
result = await self.collection.aggregate(pipeline).to_list(1)
if not result:
return self._default_performance_metrics()
data = result[0]
total = data["total"][0]["count"] if data["total"] else 0
if total == 0:
return self._default_performance_metrics()
stroke_pred = data["stroke_predictions"][0]["count"] if data["stroke_predictions"] else 0
normal_pred = total - stroke_pred
high_conf_stroke = data["high_confidence_stroke"][0]["count"] if data["high_confidence_stroke"] else 0
high_conf_normal = data["high_confidence_normal"][0]["count"] if data["high_confidence_normal"] else 0
avg_conf = data["avg_confidence"][0]["avg"] if data["avg_confidence"] else 0.5
# Estimate performance metrics
estimated_tp = high_conf_stroke
estimated_tn = high_conf_normal
estimated_fp = stroke_pred - high_conf_stroke
estimated_fn = normal_pred - high_conf_normal
# Ensure non-negative values
estimated_fp = max(0, estimated_fp)
estimated_fn = max(0, estimated_fn)
# Calculate metrics
sensitivity = (estimated_tp / (estimated_tp + estimated_fn)) if (estimated_tp + estimated_fn) > 0 else 0
specificity = (estimated_tn / (estimated_tn + estimated_fp)) if (estimated_tn + estimated_fp) > 0 else 0
ppv = (estimated_tp / (estimated_tp + estimated_fp)) if (estimated_tp + estimated_fp) > 0 else 0
npv = (estimated_tn / (estimated_tn + estimated_fn)) if (estimated_tn + estimated_fn) > 0 else 0
# F1 Score
precision = ppv
recall = sensitivity
f1 = (2 * precision * recall / (precision + recall)) if (precision + recall) > 0 else 0
# AUC-ROC estimate
auc_roc = min(0.5 + (avg_conf * 0.5), 0.99)
return {
"sensitivity": round(min(sensitivity * 100, 100.0), 1),
"specificity": round(min(specificity * 100, 100.0), 1),
"positive_predictive_value": round(min(ppv * 100, 100.0), 1),
"negative_predictive_value": round(min(npv * 100, 100.0), 1),
"f1_score": round(min(f1 * 100, 100.0), 1),
"auc_roc": round(auc_roc, 3)
}
except Exception as e:
print(f"Failed to calculate estimated performance metrics: {e}")
return self._default_performance_metrics()
async def _calculate_stroke_type_breakdown(self, filter_query: Dict[str, Any]) -> Dict[str, Any]:
"""Calculate breakdown of cases by doctor-classified stroke type"""
try:
pipeline = [
{"$match": filter_query},
{"$group": {
"_id": "$stroke_type",
"count": {"$sum": 1}
}}
]
results = await self.collection.aggregate(pipeline).to_list(10)
normal_count = 0
ischemia_count = 0
bleeding_count = 0
unclassified_count = 0
for item in results:
stroke_type = item["_id"]
count = item["count"]
if stroke_type == "normal":
normal_count = count
elif stroke_type == "ischemia":
ischemia_count = count
elif stroke_type == "bleeding":
bleeding_count = count
else:
unclassified_count += count
classified_total = normal_count + ischemia_count + bleeding_count
return {
"normal_count": normal_count,
"ischemia_count": ischemia_count,
"bleeding_count": bleeding_count,
"unclassified_count": unclassified_count,
"normal_percentage": round((normal_count / classified_total * 100) if classified_total > 0 else 0, 1),
"ischemia_percentage": round((ischemia_count / classified_total * 100) if classified_total > 0 else 0, 1),
"bleeding_percentage": round((bleeding_count / classified_total * 100) if classified_total > 0 else 0, 1),
}
except Exception as e:
print(f"Failed to calculate stroke type breakdown: {e}")
return {
"normal_count": 0, "ischemia_count": 0, "bleeding_count": 0,
"unclassified_count": 0, "normal_percentage": 0.0,
"ischemia_percentage": 0.0, "bleeding_percentage": 0.0
}
async def _calculate_stroke_type_accuracy(self, filter_query: Dict[str, Any]) -> Dict[str, float]:
"""Calculate AI accuracy for each stroke type category using doctor classification as ground truth"""
try:
# Only consider cases where doctor provided a stroke_type
typed_filter = {**filter_query, "stroke_type": {"$ne": None}}
typed_count = await self.collection.count_documents(typed_filter)
if typed_count == 0:
return {"normal_accuracy": 0.0, "ischemia_accuracy": 0.0, "bleeding_accuracy": 0.0}
# For each type, check how often AI agreed with the doctor
accuracies = {}
for stype in ["normal", "ischemia", "bleeding"]:
type_filter = {**filter_query, "stroke_type": stype}
type_total = await self.collection.count_documents(type_filter)
if type_total == 0:
accuracies[f"{stype}_accuracy"] = 0.0
continue
if stype == "normal":
# Doctor says normal → AI should predict no stroke (has_stroke=False)
correct_count = await self.collection.count_documents({**type_filter, "has_stroke": False})
else:
# Doctor says ischemia/bleeding → AI should predict stroke (has_stroke=True)
correct_count = await self.collection.count_documents({**type_filter, "has_stroke": True})
accuracies[f"{stype}_accuracy"] = round((correct_count / type_total * 100), 1)
return accuracies
except Exception as e:
print(f"Failed to calculate stroke type accuracy: {e}")
return {"normal_accuracy": 0.0, "ischemia_accuracy": 0.0, "bleeding_accuracy": 0.0}
def _default_performance_metrics(self) -> Dict[str, float]:
"""Return default performance metrics"""
return {
"sensitivity": 0.0,
"specificity": 0.0,
"positive_predictive_value": 0.0,
"negative_predictive_value": 0.0,
"f1_score": 0.0,
"auc_roc": 0.5
}
def is_available(self) -> bool:
"""Check if database is available"""
return self.is_connected
# Global database manager instance
db_manager = DatabaseManager()