CoolFace
Apppublic

MODLI/AutoImageProcessor

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes
app.py118 linesDownload Raw Back to root
1import gradio as gr2from transformers import ViTImageProcessor, ViTForImageClassification3from PIL import Image4import torch5import os6 7# --- Chargement du modèle et du processeur ---8print("Loading model and processor...")9model_name = "google/vit-base-patch16-224"10processor = ViTImageProcessor.from_pretrained(model_name)11model = ViTForImageClassification.from_pretrained(model_name)12print("Model loaded successfully!")13 14def predict(image):15    """Fonction de prédiction avec gestion d'erreurs et seuil de confiance"""16    try:17        # Conversion vers RGB pour éviter les erreurs de canaux18        if image.mode != 'RGB':19            image = image.convert('RGB')20        21        # Pré-traitement de l'image22        inputs = processor(images=image, return_tensors="pt")23        24        # Prédiction25        with torch.no_grad():26            outputs = model(**inputs)27            logits = outputs.logits28        29        # Application de softmax pour obtenir les probabilités30        probabilities = torch.nn.functional.softmax(logits, dim=-1)[0]31        top_probs, top_indices = torch.topk(probabilities, 5)  # Top 5 predictions32        33        # Formatage des résultats sous forme de dictionnaire pour l'affichage34        results = {}35        for prob, idx in zip(top_probs, top_indices):36            pred_label = model.config.id2label[idx.item()]37            confidence = prob.item()38            if confidence > 0.01:  # Seuil de confiance à 1%39                results[pred_label] = confidence40        41        if not results:42            return {"Aucune prédiction fiable": 0.0}, "Je ne suis pas sûr de reconnaître cet item. Essayez avec une image plus claire."43        44        # Créer un message de résultat45        top_prediction = list(results.items())[0]46        message = f"🏷️ Prédiction principale: {top_prediction[0]} ({top_prediction[1]:.2%})"47        48        return results, message49        50    except Exception as e:51        return {"Erreur": 0.0}, f"Une erreur s'est produite: {str(e)}"52 53# Interface Gradio améliorée54with gr.Blocks(title="Fashion Classifier", theme=gr.themes.Soft()) as demo:55    gr.Markdown("# 👗 Fashion Item Classifier")56    gr.Markdown("Téléchargez une image de vêtement pour le classer automatiquement")57    58    with gr.Row():59        with gr.Column(scale=1):60            image_input = gr.Image(61                type="pil", 62                label="Image du vêtement",63                height=300,64                sources=["upload", "webcam", "clipboard"]65            )66            upload_btn = gr.Button("🚀 Analyser l'image", variant="primary")67        68        with gr.Column(scale=1):69            label_output = gr.Label(70                label="Résultats de classification",71                num_top_classes=572            )73            text_output = gr.Textbox(74                label="Conclusion",75                interactive=False76            )77    78    # Exemples79    gr.Examples(80        examples=[81            ["https://images.unsplash.com/photo-1552374196-c4e7ffc6e126?w=300"],  # T-shirt82            ["https://images.unsplash.com/photo-1543163521-1bf539c55dd2?w=300"],  # Chaussures83            ["https://images.unsplash.com/photo-1594633312681-425c7b97ccd1?w=300"]   # Robe84        ],85        inputs=image_input,86        label="Exemples d'images à tester"87    )88    89    # Instructions90    gr.Markdown("""91    ### 📋 Instructions92    - Téléchargez une image claire d'un vêtement93    - L'image doit montrer le vêtement de face94    - Fond uni recommandé pour de meilleurs résultats95    - Cliquez sur 'Analyser l'image' pour obtenir la classification96    """)97    98    # Liaison du bouton99    upload_btn.click(100        fn=predict,101        inputs=image_input,102        outputs=[label_output, text_output]103    )104    105    # Liaison aussi quand on upload une image106    image_input.upload(107        fn=predict,108        inputs=image_input,109        outputs=[label_output, text_output]110    )111 112# Lancement de l'application113if __name__ == "__main__":114    demo.launch(115        debug=True,116        server_name="0.0.0.0",117        server_port=int(os.environ.get("PORT", 7860))118    )