BBPlease / src /feedback_manager.py
hamza-ks's picture
Relaunch deploy: RF v5 (72.6%), voc2vec fallback, telemetry removed, secrets docs removed
84e6d52
Raw
History Blame Contribute Delete
9.36 kB
"""
Baby Cry AI - Feedback Manager
Handles continuous learning: saving audio feedback, metadata tracking, and retrain triggers
"""
import os
import json
import shutil
import uuid
from datetime import datetime
from typing import Optional, Dict, List, Tuple
from collections import Counter
class FeedbackManager:
"""Manages user feedback for continuous learning"""
VALID_CATEGORIES = ['hunger', 'sleep', 'pain', 'discomfort', 'fussiness']
RETRAIN_THRESHOLD = 50 # Trigger retrain after this many new samples
def __init__(self, feedback_dir: str = "feedback_data", data_dir: str = "data"):
"""
Initialize the FeedbackManager.
Args:
feedback_dir: Directory to store feedback audio files
data_dir: Directory containing training data
"""
self.feedback_dir = feedback_dir
self.data_dir = data_dir
self.verified_dir = os.path.join(feedback_dir, "verified")
self.pending_dir = os.path.join(feedback_dir, "pending")
self.metadata_file = os.path.join(feedback_dir, "metadata.json")
# Ensure directories exist
self._ensure_directories()
def _ensure_directories(self):
"""Create necessary directory structure"""
# Create main feedback directory
os.makedirs(self.feedback_dir, exist_ok=True)
os.makedirs(self.pending_dir, exist_ok=True)
# Create verified directories for each category
for category in self.VALID_CATEGORIES:
os.makedirs(os.path.join(self.verified_dir, category), exist_ok=True)
def _load_metadata(self) -> Dict:
"""Load metadata from file"""
if os.path.exists(self.metadata_file):
try:
with open(self.metadata_file, 'r') as f:
return json.load(f)
except (json.JSONDecodeError, IOError):
pass
return {
"total_submissions": 0,
"verified_count": 0,
"pending_count": 0,
"last_retrain": None,
"samples_since_retrain": 0,
"category_counts": {cat: 0 for cat in self.VALID_CATEGORIES},
"submissions": []
}
def _save_metadata(self, metadata: Dict):
"""Save metadata to file"""
with open(self.metadata_file, 'w') as f:
json.dump(metadata, f, indent=2)
def save_feedback(
self,
audio_path: str,
predicted_label: str,
correct_label: str,
confidence: float,
is_correct: bool
) -> Tuple[str, bool]:
"""
Save audio feedback for continuous learning.
Args:
audio_path: Path to the temporary audio file
predicted_label: What the model predicted
correct_label: The correct label (same as predicted if is_correct=True)
confidence: Model confidence for the prediction
is_correct: Whether the prediction was correct
Returns:
Tuple of (submission_id, should_retrain)
"""
if correct_label not in self.VALID_CATEGORIES:
raise ValueError(f"Invalid category: {correct_label}. Must be one of {self.VALID_CATEGORIES}")
# Generate unique ID for this submission
submission_id = f"{datetime.now().strftime('%Y%m%d_%H%M%S')}_{uuid.uuid4().hex[:8]}"
# Determine destination based on correctness
# Correct predictions go directly to verified, wrong ones could go to pending for review
# For simplicity, we'll auto-verify all submissions where user provides correct label
dest_dir = os.path.join(self.verified_dir, correct_label)
dest_filename = f"{submission_id}.wav"
dest_path = os.path.join(dest_dir, dest_filename)
# Copy audio file to feedback storage
shutil.copy2(audio_path, dest_path)
# Update metadata
metadata = self._load_metadata()
submission_record = {
"id": submission_id,
"timestamp": datetime.now().isoformat(),
"predicted_label": predicted_label,
"correct_label": correct_label,
"confidence": confidence,
"is_correct": is_correct,
"file_path": dest_path,
"status": "verified"
}
metadata["total_submissions"] += 1
metadata["verified_count"] += 1
metadata["samples_since_retrain"] += 1
metadata["category_counts"][correct_label] += 1
metadata["submissions"].append(submission_record)
self._save_metadata(metadata)
# Check if retrain should be triggered
should_retrain = metadata["samples_since_retrain"] >= self.RETRAIN_THRESHOLD
return submission_id, should_retrain
def get_stats(self) -> Dict:
"""Get feedback statistics"""
metadata = self._load_metadata()
# Count actual files in verified directories
actual_counts = {}
total_verified = 0
for category in self.VALID_CATEGORIES:
cat_dir = os.path.join(self.verified_dir, category)
if os.path.exists(cat_dir):
count = len([f for f in os.listdir(cat_dir) if f.endswith('.wav')])
actual_counts[category] = count
total_verified += count
else:
actual_counts[category] = 0
return {
"total_submissions": metadata.get("total_submissions", 0),
"verified_count": total_verified,
"pending_count": metadata.get("pending_count", 0),
"samples_since_retrain": metadata.get("samples_since_retrain", 0),
"retrain_threshold": self.RETRAIN_THRESHOLD,
"progress_to_retrain": min(100, (metadata.get("samples_since_retrain", 0) / self.RETRAIN_THRESHOLD) * 100),
"last_retrain": metadata.get("last_retrain"),
"category_counts": actual_counts,
"ready_for_retrain": metadata.get("samples_since_retrain", 0) >= self.RETRAIN_THRESHOLD
}
def get_verified_data_paths(self) -> List[Tuple[str, str]]:
"""
Get all verified feedback audio files with their labels.
Returns:
List of (file_path, category) tuples
"""
data_paths = []
for category in self.VALID_CATEGORIES:
cat_dir = os.path.join(self.verified_dir, category)
if os.path.exists(cat_dir):
for filename in os.listdir(cat_dir):
if filename.endswith('.wav'):
file_path = os.path.join(cat_dir, filename)
data_paths.append((file_path, category))
return data_paths
def merge_to_training_data(self) -> Dict:
"""
Merge verified feedback data into the main training directory.
Returns:
Dict with merge statistics
"""
stats = {
"merged_count": 0,
"by_category": {cat: 0 for cat in self.VALID_CATEGORIES}
}
for category in self.VALID_CATEGORIES:
source_dir = os.path.join(self.verified_dir, category)
dest_dir = os.path.join(self.data_dir, category)
if not os.path.exists(source_dir):
continue
os.makedirs(dest_dir, exist_ok=True)
# Get existing file count for naming
existing_files = [f for f in os.listdir(dest_dir) if f.endswith('.wav')]
next_index = len(existing_files) + 1
# Copy files
for filename in os.listdir(source_dir):
if filename.endswith('.wav'):
source_path = os.path.join(source_dir, filename)
dest_filename = f"{category}_feedback_{next_index:04d}.wav"
dest_path = os.path.join(dest_dir, dest_filename)
shutil.copy2(source_path, dest_path)
stats["merged_count"] += 1
stats["by_category"][category] += 1
next_index += 1
return stats
def mark_retrain_complete(self):
"""Mark that a retrain has been completed"""
metadata = self._load_metadata()
metadata["last_retrain"] = datetime.now().isoformat()
metadata["samples_since_retrain"] = 0
self._save_metadata(metadata)
def clear_verified_data(self):
"""Clear verified data after successful merge and retrain"""
for category in self.VALID_CATEGORIES:
cat_dir = os.path.join(self.verified_dir, category)
if os.path.exists(cat_dir):
for filename in os.listdir(cat_dir):
if filename.endswith('.wav'):
os.remove(os.path.join(cat_dir, filename))
def get_recent_submissions(self, limit: int = 10) -> List[Dict]:
"""Get recent submissions for display"""
metadata = self._load_metadata()
submissions = metadata.get("submissions", [])
return submissions[-limit:][::-1] # Most recent first