App / app.py
RICHERGIRL's picture
Update app.py
5322998 verified
Raw
History Blame
3.18 kB
import gradio as gr
import joblib
import os
import numpy as np
from sklearn.ensemble import RandomForestClassifier
from sklearn.preprocessing import LabelEncoder
from utils import extract_features
def initialize_fallback_model():
"""Creates and trains a simple fallback model"""
print("Initializing fallback model...")
# Simple training data
X = np.array([[0,0,0], [1,1,1], [2,2,2]]) # Dummy encoded features
y = np.array([0, 1, 0]) # Dummy target
model = RandomForestClassifier(n_estimators=10)
model.fit(X, y)
encoders = {
'face_shape': LabelEncoder().fit(['Oval', 'Round', 'Square']),
'skin_tone': LabelEncoder().fit(['Fair', 'Medium', 'Dark']),
'face_size': LabelEncoder().fit(['Small', 'Medium', 'Large']),
'mask_style': LabelEncoder().fit(['StyleA', 'StyleB', 'StyleC']) # Added mask_style
}
return model, encoders
def safe_load_model():
"""Safely loads model files with comprehensive fallback"""
try:
if not all(os.path.exists(f'model/{f}') for f in ['random_forest.pkl', 'label_encoders.pkl']):
raise FileNotFoundError("Model files missing")
model = joblib.load('model/random_forest.pkl', mmap_mode='r')
encoders = joblib.load('model/label_encoders.pkl', mmap_mode='r')
# Verify model is fitted
if not hasattr(model, 'classes_'):
raise ValueError("Model not properly trained")
print("Main model loaded successfully!")
return model, encoders
except Exception as e:
print(f"Loading failed: {str(e)}")
return initialize_fallback_model()
def recommend_mask(image):
"""Process image and make prediction with error handling"""
try:
# Extract features
face_shape, skin_tone, face_size = extract_features(image)
# Encode features
face_encoded = encoders["face_shape"].transform([face_shape])[0]
skin_encoded = encoders["skin_tone"].transform([skin_tone])[0]
size_encoded = encoders["face_size"].transform([face_size])[0]
# Predict
prediction = model.predict([[face_encoded, skin_encoded, size_encoded]])[0]
return encoders["mask_style"].classes_[prediction]
# Get recommended mask image path
mask_image_path = encoders['mask_images'][prediction]
return (
encoders["mask_style"].classes_[prediction], # Text
mask_image_path # Image
)
except Exception as e:
print(f"Prediction error: {str(e)}")
return "Error", "default_mask.png" # Fallback
# Initialize model and encoders
model, encoders = safe_load_model()
# Create Gradio interface
demo = gr.Interface(
fn=recommend_mask,
inputs=gr.Image(type="filepath"),
outputs=[
gr.Textbox(label="Recommended Style"),
gr.Image(label="Mask Preview") # Add image output
],
title="🎭 AI Party Mask Recommender",
description="Upload a photo to get a personalized mask recommendation!",
)
if __name__ == "__main__":
demo.launch()