baby-cry-ai / inference.py
Pranjal2510
Baby-Cry-Analysis Api
9b4e272
Raw
History Blame
2.25 kB
"""
Baby Cry Analysis - Inference Module
Optimized for Hugging Face Spaces with eager model loading.
"""
import logging
logger = logging.getLogger(__name__)
# Model loaded at module level for single load
_model = None
# Supported cry reason labels
LABELS = ["hunger", "belly_pain", "tired", "discomfort", "burping"]
def preload_model():
"""
Eagerly load the model at startup.
Called once when the app starts to avoid cold start delays.
"""
global _model
if _model is None:
logger.info("🔄 Loading baby cry classification model at startup...")
from transformers import pipeline
_model = pipeline(
"audio-classification",
model="Wiam/baby-cry-classification-finetuned-babycry-v4"
)
logger.info("✅ Model loaded successfully!")
return _model
def get_model():
"""Get the loaded model instance."""
global _model
if _model is None:
return preload_model()
return _model
def analyze_cry(audio_path: str) -> dict:
"""
Analyze a baby cry audio file using the supervised classification model.
Args:
audio_path: Path to the audio file (WAV format preferred)
Returns:
Dictionary containing:
- cry_detected: boolean indicating if a cry was detected
- top_reason: the most likely reason for crying
- scores: confidence scores for each label
"""
# Get preloaded model
model = get_model()
# Get supervised model predictions
results = model(audio_path)
# Build scores dictionary from model output
scores = {}
for r in results:
label = r["label"]
if label in LABELS:
scores[label] = round(r["score"], 4)
# Ensure all labels have a score (default 0 if missing)
for label in LABELS:
if label not in scores:
scores[label] = 0.0
# Determine top prediction
top_label = max(scores, key=scores.get)
top_confidence = scores[top_label]
# Cry detection threshold
cry_detected = top_confidence >= 0.1
return {
"cry_detected": cry_detected,
"top_reason": top_label if cry_detected else None,
"scores": scores
}