MODLI/AutoImageProcessor
0
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 )