BBPlease / src /model_manager.py
hamza-ksr's picture
Upload folder using huggingface_hub
1bb6efc verified
Raw
History Blame Contribute Delete
9.7 kB
"""
Baby Cry AI - Model Manager
Handles model versioning, training, and hot-swapping for continuous learning
"""
import os
import sys
import json
import shutil
from datetime import datetime
from typing import Optional, Dict, List, Tuple
import pickle
# Add the src directory to Python path for imports to work from any location
SRC_DIR = os.path.dirname(os.path.abspath(__file__))
if SRC_DIR not in sys.path:
sys.path.insert(0, SRC_DIR)
from models.baseline_model import BaselineModel
from feedback_manager import FeedbackManager
class ModelManager:
"""Manages model versions and continuous learning retraining"""
def __init__(
self,
models_dir: str = "models",
data_dir: str = "data",
feedback_dir: str = "feedback_data"
):
"""
Initialize the ModelManager.
Args:
models_dir: Directory containing model files
data_dir: Directory containing training data
feedback_dir: Directory containing feedback data
"""
self.models_dir = models_dir
self.data_dir = data_dir
self.feedback_dir = feedback_dir
self.history_file = os.path.join(models_dir, "model_history.json")
self.active_model_file = os.path.join(models_dir, "baseline_model.pkl")
# Initialize feedback manager
self.feedback_manager = FeedbackManager(feedback_dir, data_dir)
# Current active model
self._active_model: Optional[BaselineModel] = None
# Ensure directories exist
os.makedirs(models_dir, exist_ok=True)
def _load_history(self) -> Dict:
"""Load model history from file"""
if os.path.exists(self.history_file):
try:
with open(self.history_file, 'r') as f:
return json.load(f)
except (json.JSONDecodeError, IOError):
pass
return {
"versions": [],
"active_version": None,
"total_retrains": 0
}
def _save_history(self, history: Dict):
"""Save model history to file"""
with open(self.history_file, 'w') as f:
json.dump(history, f, indent=2)
def get_active_model(self) -> BaselineModel:
"""Get the currently active model, loading if necessary"""
if self._active_model is None:
self._active_model = BaselineModel(self.active_model_file)
if not self._active_model.load_model():
print("⚠️ No active model found")
return self._active_model
def get_versions(self) -> List[Dict]:
"""Get list of all model versions"""
history = self._load_history()
return history.get("versions", [])
def get_active_version(self) -> Optional[str]:
"""Get the currently active model version"""
history = self._load_history()
return history.get("active_version")
def retrain_with_feedback(self, include_feedback: bool = True) -> Dict:
"""
Retrain the model with current data plus feedback.
Args:
include_feedback: Whether to include feedback data in training
Returns:
Dict with training results
"""
# Create new model instance for training
new_model = BaselineModel()
# Load existing training data
print("📂 Loading training data...")
X, y = new_model.load_data_from_directory(self.data_dir, balance_data=True)
if X is None or len(X) == 0:
return {"error": "No training data available", "success": False}
initial_count = len(X)
# Optionally merge feedback data first
feedback_count = 0
if include_feedback:
feedback_stats = self.feedback_manager.get_stats()
if feedback_stats["verified_count"] > 0:
print(f"📥 Merging {feedback_stats['verified_count']} feedback samples...")
merge_result = self.feedback_manager.merge_to_training_data()
feedback_count = merge_result["merged_count"]
# Reload data with merged feedback
X, y = new_model.load_data_from_directory(self.data_dir, balance_data=True)
# Train the model
print("🚀 Training new model version...")
accuracy = new_model.train(X, y)
if accuracy is None or accuracy == False:
return {"error": "Training failed", "success": False}
# Get model info
model_info = new_model.get_model_info()
# Create version entry
history = self._load_history()
version_num = len(history.get("versions", [])) + 1
version_id = f"v{version_num}"
# Save versioned model
version_model_path = os.path.join(self.models_dir, f"baseline_model_{version_id}.pkl")
shutil.copy2(self.active_model_file, version_model_path)
version_entry = {
"version": version_id,
"created_at": datetime.now().isoformat(),
"training_samples": len(X),
"feedback_samples_merged": feedback_count,
"accuracy": accuracy,
"model_path": version_model_path,
"is_active": True
}
# Mark previous active version as inactive
for v in history.get("versions", []):
v["is_active"] = False
# Add new version
history["versions"].append(version_entry)
history["active_version"] = version_id
history["total_retrains"] = history.get("total_retrains", 0) + 1
self._save_history(history)
# Clear feedback data after successful merge
if include_feedback and feedback_count > 0:
self.feedback_manager.clear_verified_data()
# Mark retrain complete in feedback manager
self.feedback_manager.mark_retrain_complete()
# Hot-swap to new model
self._active_model = new_model
return {
"success": True,
"version": version_id,
"training_samples": len(X),
"feedback_merged": feedback_count,
"accuracy": accuracy,
"model_info": model_info
}
def switch_version(self, version_id: str) -> Dict:
"""
Switch to a specific model version.
Args:
version_id: Version ID to switch to (e.g., "v1", "v2")
Returns:
Dict with switch result
"""
history = self._load_history()
# Find the version
target_version = None
for v in history.get("versions", []):
if v["version"] == version_id:
target_version = v
break
if target_version is None:
return {"success": False, "error": f"Version {version_id} not found"}
# Check if model file exists
model_path = target_version.get("model_path")
if not model_path or not os.path.exists(model_path):
return {"success": False, "error": f"Model file for {version_id} not found"}
# Copy versioned model to active model path
shutil.copy2(model_path, self.active_model_file)
# Update history
for v in history["versions"]:
v["is_active"] = (v["version"] == version_id)
history["active_version"] = version_id
self._save_history(history)
# Reload active model
self._active_model = BaselineModel(self.active_model_file)
self._active_model.load_model()
return {
"success": True,
"version": version_id,
"accuracy": target_version.get("accuracy"),
"created_at": target_version.get("created_at")
}
def get_model_summary(self) -> Dict:
"""Get comprehensive summary of model state"""
history = self._load_history()
feedback_stats = self.feedback_manager.get_stats()
active_model = self.get_active_model()
model_info = active_model.get_model_info() if active_model.is_trained else {}
return {
"active_version": history.get("active_version", "v0 (initial)"),
"total_versions": len(history.get("versions", [])),
"total_retrains": history.get("total_retrains", 0),
"model_status": "trained" if active_model.is_trained else "not_trained",
"model_info": model_info,
"feedback": {
"pending_samples": feedback_stats.get("verified_count", 0),
"samples_since_retrain": feedback_stats.get("samples_since_retrain", 0),
"retrain_threshold": feedback_stats.get("retrain_threshold", 50),
"progress_to_retrain": feedback_stats.get("progress_to_retrain", 0),
"ready_for_retrain": feedback_stats.get("ready_for_retrain", False)
}
}
def check_and_auto_retrain(self) -> Optional[Dict]:
"""
Check if automatic retrain should be triggered and perform it.
Returns:
Retrain result if triggered, None otherwise
"""
feedback_stats = self.feedback_manager.get_stats()
if feedback_stats.get("ready_for_retrain", False):
print("🔄 Auto-retrain triggered!")
return self.retrain_with_feedback()
return None