| """ |
| 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 |
|
|
| |
| 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") |
| |
| |
| self.feedback_manager = FeedbackManager(feedback_dir, data_dir) |
| |
| |
| self._active_model: Optional[BaselineModel] = None |
| |
| |
| 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 |
| """ |
| |
| new_model = BaselineModel() |
| |
| |
| 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) |
| |
| |
| 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"] |
| |
| |
| X, y = new_model.load_data_from_directory(self.data_dir, balance_data=True) |
| |
| |
| print("🚀 Training new model version...") |
| accuracy = new_model.train(X, y) |
| |
| if accuracy is None or accuracy == False: |
| return {"error": "Training failed", "success": False} |
| |
| |
| model_info = new_model.get_model_info() |
| |
| |
| history = self._load_history() |
| version_num = len(history.get("versions", [])) + 1 |
| version_id = f"v{version_num}" |
| |
| |
| 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 |
| } |
| |
| |
| for v in history.get("versions", []): |
| v["is_active"] = False |
| |
| |
| history["versions"].append(version_entry) |
| history["active_version"] = version_id |
| history["total_retrains"] = history.get("total_retrains", 0) + 1 |
| |
| self._save_history(history) |
| |
| |
| if include_feedback and feedback_count > 0: |
| self.feedback_manager.clear_verified_data() |
| |
| |
| self.feedback_manager.mark_retrain_complete() |
| |
| |
| 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() |
| |
| |
| 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"} |
| |
| |
| 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"} |
| |
| |
| shutil.copy2(model_path, self.active_model_file) |
| |
| |
| for v in history["versions"]: |
| v["is_active"] = (v["version"] == version_id) |
| history["active_version"] = version_id |
| |
| self._save_history(history) |
| |
| |
| 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 |
|
|
|
|