""" 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