import streamlit as st import os import torch from diffusers import StableDiffusionInstructPix2PixPipeline from PIL import Image import tempfile import time import requests from io import BytesIO import shutil # Configuration de la page st.set_page_config( page_title="Image Editor", page_icon="🎨", layout="wide" ) # Titre et description st.title("🎨 Image Editor (InstructPix2Pix)") st.markdown("Éditez vos images en utilisant des prompts textuels avec le modèle InstructPix2Pix") # Configuration du cache pour éviter les problèmes de permissions def setup_cache_directory(): """Configure un répertoire de cache accessible en écriture""" try: # Essayer d'utiliser un répertoire temporaire cache_dir = tempfile.mkdtemp(prefix="hf_cache_") os.environ['HF_HOME'] = cache_dir os.environ['TRANSFORMERS_CACHE'] = cache_dir os.environ['HF_DATASETS_CACHE'] = cache_dir return cache_dir except Exception as e: st.error(f"Erreur lors de la configuration du cache: {e}") return None # Configurer le cache dès le début cache_dir = setup_cache_directory() # Détection du device @st.cache_data def get_device(): if torch.cuda.is_available(): return "cuda" elif torch.backends.mps.is_available(): return "mps" else: return "cpu" device = get_device() st.sidebar.info(f"Device utilisé: {device}") # Initialiser le pipeline avec gestion des erreurs améliorée @st.cache_resource def load_pipeline(): try: st.info("🔄 Chargement du modèle en cours... Cela peut prendre quelques minutes lors du premier lancement.") # Vérifier l'espace disque disponible if cache_dir: disk_usage = shutil.disk_usage(cache_dir) free_gb = disk_usage.free / (1024**3) st.info(f"Espace disque disponible: {free_gb:.1f} GB") if free_gb < 10: st.warning("⚠️ Espace disque faible. Le téléchargement du modèle pourrait échouer.") # Utiliser le modèle InstructPix2Pix officiel model_id = "timbrooks/instruct-pix2pix" # Configuration pour éviter les problèmes de cache try: # Essayer de charger depuis le cache local d'abord pipe = StableDiffusionInstructPix2PixPipeline.from_pretrained( model_id, torch_dtype=torch.float16 if device != "cpu" else torch.float32, safety_checker=None, requires_safety_checker=False, cache_dir=cache_dir, local_files_only=False, use_auth_token=False, force_download=False ) except Exception as cache_error: st.warning(f"Erreur de cache: {cache_error}") st.info("Tentative de téléchargement direct...") # Si le cache pose problème, essayer sans cache pipe = StableDiffusionInstructPix2PixPipeline.from_pretrained( model_id, torch_dtype=torch.float16 if device != "cpu" else torch.float32, safety_checker=None, requires_safety_checker=False, local_files_only=False, use_auth_token=False ) # Déplacer vers le device approprié pipe = pipe.to(device) # Optimisations pour économiser la mémoire if device != "cpu": try: pipe.enable_model_cpu_offload() pipe.enable_attention_slicing() if hasattr(pipe, 'enable_vae_slicing'): pipe.enable_vae_slicing() if hasattr(pipe, 'enable_xformers_memory_efficient_attention'): pipe.enable_xformers_memory_efficient_attention() except Exception as opt_error: st.warning(f"Certaines optimisations n'ont pas pu être activées: {opt_error}") st.success("✅ Modèle InstructPix2Pix chargé avec succès!") return pipe except PermissionError as e: st.error(f"❌ Erreur de permissions: {str(e)}") st.error("Veuillez vérifier les permissions d'écriture ou exécuter avec des privilèges appropriés.") st.info("Solution possible: Exécutez dans un environnement avec les bonnes permissions.") return None except OSError as e: if "disk" in str(e).lower() or "space" in str(e).lower(): st.error("❌ Espace disque insuffisant pour télécharger le modèle.") st.error("Le modèle InstructPix2Pix nécessite environ 5-10 GB d'espace libre.") else: st.error(f"❌ Erreur système: {str(e)}") return None except ImportError as e: st.error(f"❌ Erreur d'importation: {str(e)}") st.error("Veuillez installer les dépendances requises:") st.code("pip install --upgrade diffusers transformers accelerate torch torchvision") return None except Exception as e: st.error(f"❌ Erreur lors du chargement du modèle: {str(e)}") st.error("Causes possibles:") st.markdown(""" - Connexion internet instable - Espace disque insuffisant - Problème de permissions - Version incompatible des dépendances """) # Afficher des informations de debug with st.expander("Informations de debug"): st.write(f"Cache directory: {cache_dir}") st.write(f"Device: {device}") st.write(f"PyTorch version: {torch.__version__}") try: import diffusers st.write(f"Diffusers version: {diffusers.__version__}") except: st.write("Diffusers non installé ou version non disponible") return None # Fonction pour charger une image depuis une URL def load_image_from_url(url): try: headers = { 'User-Agent': 'Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36' } response = requests.get(url, timeout=15, headers=headers) response.raise_for_status() image = Image.open(BytesIO(response.content)) return image.convert("RGB") except requests.exceptions.Timeout: st.error("Timeout lors du chargement de l'image. Veuillez réessayer.") return None except requests.exceptions.RequestException as e: st.error(f"Erreur lors du chargement de l'image: {str(e)}") return None except Exception as e: st.error(f"Erreur lors du traitement de l'image: {str(e)}") return None # Fonction pour redimensionner l'image def resize_image(image, max_size=512): """Redimensionne l'image en gardant les proportions""" width, height = image.size # Vérifier si un redimensionnement est nécessaire if max(width, height) <= max_size: # S'assurer que les dimensions sont multiples de 8 new_width = (width // 8) * 8 new_height = (height // 8) * 8 if new_width != width or new_height != height: image = image.resize((new_width, new_height), Image.Resampling.LANCZOS) return image if width > height: new_width = max_size new_height = int(height * max_size / width) else: new_height = max_size new_width = int(width * max_size / height) # S'assurer que les dimensions sont multiples de 8 (requis par Stable Diffusion) new_width = max(8, (new_width // 8) * 8) new_height = max(8, (new_height // 8) * 8) image = image.resize((new_width, new_height), Image.Resampling.LANCZOS) return image # Tentative de chargement du pipeline try: pipeline = load_pipeline() except Exception as e: st.error(f"Erreur critique lors de l'initialisation: {e}") pipeline = None if pipeline is None: st.error("❌ Impossible de charger le modèle.") st.markdown(""" ### Solutions possibles: 1. **Vérifiez l'installation des dépendances:** ```bash pip install --upgrade diffusers transformers accelerate torch torchvision ``` 2. **Vérifiez l'espace disque:** Le modèle nécessite ~10GB d'espace libre 3. **Permissions:** Assurez-vous d'avoir les permissions d'écriture 4. **Connexion internet:** Le modèle doit être téléchargé lors du premier usage 5. **Redémarrez l'application** après avoir résolu les problèmes """) # Option pour réessayer if st.button("🔄 Réessayer de charger le modèle"): st.rerun() st.stop() # Interface utilisateur col1, col2 = st.columns([1, 1]) with col1: st.header("📤 Image d'entrée") # Options pour charger l'image input_option = st.radio( "Choisissez comment charger votre image:", ["Télécharger un fichier", "URL d'image", "Image d'exemple"] ) uploaded_file = None image_url = None input_image = None if input_option == "Télécharger un fichier": uploaded_file = st.file_uploader( "Choisissez une image", type=['png', 'jpg', 'jpeg', 'webp'], help="Formats supportés: PNG, JPG, JPEG, WEBP (Max: 200MB)" ) if uploaded_file: # Vérifier la taille du fichier if uploaded_file.size > 200 * 1024 * 1024: # 200MB st.error("Fichier trop volumineux. Maximum: 200MB") else: try: input_image = Image.open(uploaded_file).convert("RGB") st.image(input_image, caption=f"Image téléchargée ({input_image.size[0]}x{input_image.size[1]})", use_column_width=True) except Exception as e: st.error(f"Erreur lors de l'ouverture de l'image: {e}") elif input_option == "URL d'image": image_url = st.text_input( "URL de l'image", placeholder="https://example.com/image.jpg", help="Entrez l'URL complète de l'image (formats: jpg, png, webp)" ) if image_url: if image_url.startswith(('http://', 'https://')): input_image = load_image_from_url(image_url) if input_image: st.image(input_image, caption=f"Image depuis URL ({input_image.size[0]}x{input_image.size[1]})", use_column_width=True) else: st.error("Veuillez entrer une URL valide commençant par http:// ou https://") else: # Image d'exemple example_options = { "Portrait d'homme": "https://raw.githubusercontent.com/timothybrooks/instruct-pix2pix/main/imgs/example.jpg", "Paysage": "https://images.unsplash.com/photo-1506905925346-21bda4d32df4?w=512", "Animal": "https://images.unsplash.com/photo-1574158622682-e40e69881006?w=512" } selected_example = st.selectbox("Choisir une image d'exemple:", list(example_options.keys())) if st.button("Charger l'image d'exemple"): example_url = example_options[selected_example] input_image = load_image_from_url(example_url) if input_image: st.image(input_image, caption=f"{selected_example} ({input_image.size[0]}x{input_image.size[1]})", use_column_width=True) with col2: st.header("⚙️ Paramètres d'édition") # Prompt de modification prompt = st.text_area( "Prompt de modification", placeholder="Exemple: 'turn him into a cyborg'", height=100, help="Décrivez en anglais les modifications que vous souhaitez apporter à l'image" ) # Exemples de prompts if st.button("💡 Prompt aléatoire"): example_prompts = [ "turn him into a cyborg", "make it a cartoon", "add sunglasses", "make it winter", "turn the sky purple", "make him smile", "add a hat", "make it look like a painting", "turn it into a sketch", "add dramatic lighting" ] import random random_prompt = random.choice(example_prompts) st.session_state.random_prompt = random_prompt st.rerun() if hasattr(st.session_state, 'random_prompt'): if st.button(f"Utiliser: '{st.session_state.random_prompt}'"): prompt = st.session_state.random_prompt # Paramètres avancés with st.expander("Paramètres avancés"): # Redimensionnement resize_option = st.checkbox( "Redimensionner l'image automatiquement", value=True, help="Recommandé pour optimiser les performances et éviter les erreurs de mémoire" ) if resize_option: max_size = st.slider( "Taille maximum (pixels)", min_value=256, max_value=1024, value=512, step=64, help="Plus petit = plus rapide, mais qualité moindre" ) # Seed use_random_seed = st.checkbox( "Seed aléatoire", value=True, help="Désactiver pour reproduire exactement les mêmes résultats" ) if not use_random_seed: seed = st.number_input( "Seed", value=42, min_value=0, max_value=2**31-1, help="Nombre pour la reproductibilité des résultats" ) else: seed = None # Paramètres de génération guidance_scale = st.slider( "Text Guidance Scale", min_value=1.0, max_value=20.0, value=7.5, step=0.5, help="Plus élevé = plus fidèle au prompt (mais peut sur-saturer)" ) image_guidance_scale = st.slider( "Image Guidance Scale", min_value=1.0, max_value=2.0, value=1.5, step=0.1, help="Plus élevé = préserve mieux l'image originale" ) num_inference_steps = st.slider( "Nombre d'étapes d'inférence", min_value=10, max_value=50, value=20, help="Plus d'étapes = meilleure qualité mais plus lent" ) # Bouton de génération generate_button = st.button( "🚀 Générer l'image éditée", type="primary", use_container_width=True, disabled=(not prompt or input_image is None) ) # Traitement et génération if generate_button: if not prompt.strip(): st.error("⚠️ Veuillez entrer un prompt de modification") elif input_image is None: st.error("⚠️ Veuillez charger une image") else: # Vérifications de sécurité if input_image.size[0] * input_image.size[1] > 2048 * 2048: st.warning("⚠️ Image très grande détectée. Le traitement peut être lent.") progress_bar = st.progress(0) status_text = st.empty() try: status_text.text("Préparation de l'image...") progress_bar.progress(10) # Préparer l'image processed_image = input_image.copy() original_size = processed_image.size if resize_option: processed_image = resize_image(processed_image, max_size) if processed_image.size != original_size: st.info(f"Image redimensionnée: {original_size} → {processed_image.size}") progress_bar.progress(20) # Générer un seed aléatoire si nécessaire if seed is None: import random current_seed = random.randint(0, 2**31 - 1) else: current_seed = seed status_text.text("Configuration du générateur...") progress_bar.progress(30) # Configurer le générateur generator = torch.Generator(device=device).manual_seed(current_seed) status_text.text("Génération en cours... (cela peut prendre quelques minutes)") progress_bar.progress(40) # Génération de l'image avec InstructPix2Pix with torch.autocast(device_type=device.replace('mps', 'cpu'), dtype=torch.float16 if device != "cpu" else torch.float32): result = pipeline( prompt=prompt.strip(), image=processed_image, num_inference_steps=num_inference_steps, guidance_scale=guidance_scale, image_guidance_scale=image_guidance_scale, generator=generator ) progress_bar.progress(90) status_text.text("Finalisation...") # Récupérer l'image générée generated_image = result.images[0] progress_bar.progress(100) status_text.empty() progress_bar.empty() # Afficher le résultat st.success("✅ Image générée avec succès!") # Informations sur la génération col_info1, col_info2, col_info3 = st.columns(3) with col_info1: st.metric("Seed utilisé", current_seed) with col_info2: st.metric("Taille finale", f"{generated_image.size[0]}×{generated_image.size[1]}") with col_info3: st.metric("Étapes", num_inference_steps) # Créer deux colonnes pour afficher les résultats result_col1, result_col2 = st.columns([1, 1]) with result_col1: st.subheader("🖼️ Image originale") st.image(processed_image, use_column_width=True) with result_col2: st.subheader("✨ Image éditée") st.image(generated_image, use_column_width=True) # Bouton de téléchargement buf = BytesIO() generated_image.save(buf, format='PNG', quality=95) buf.seek(0) timestamp = int(time.time()) filename = f"edited_image_{timestamp}.png" st.download_button( label="📥 Télécharger l'image éditée", data=buf.getvalue(), file_name=filename, mime="image/png", use_container_width=True ) # Afficher les paramètres utilisés with st.expander("Paramètres de génération"): st.json({ "prompt": prompt.strip(), "seed": current_seed, "guidance_scale": guidance_scale, "image_guidance_scale": image_guidance_scale, "num_inference_steps": num_inference_steps, "image_size": f"{generated_image.size[0]}×{generated_image.size[1]}" }) except torch.cuda.OutOfMemoryError: st.error("❌ Mémoire GPU insuffisante!") st.markdown(""" **Solutions:** - Réduire la taille de l'image - Réduire le nombre d'étapes d'inférence - Redémarrer l'application """) except Exception as e: st.error(f"❌ Erreur lors de la génération: {str(e)}") # Informations de debug pour les erreurs with st.expander("Détails de l'erreur"): st.exception(e) st.json({ "prompt_length": len(prompt), "image_size": processed_image.size if 'processed_image' in locals() else "Non définie", "device": device, "parameters": { "guidance_scale": guidance_scale, "image_guidance_scale": image_guidance_scale, "num_inference_steps": num_inference_steps } }) # Informations système dans la sidebar with st.sidebar: st.header("💻 Informations système") # Informations GPU/CPU if device == "cuda": st.success("🚀 CUDA détecté - Génération rapide") if torch.cuda.is_available(): try: gpu_name = torch.cuda.get_device_name() memory_allocated = torch.cuda.memory_allocated() / 1024**3 memory_reserved = torch.cuda.memory_reserved() / 1024**3 total_memory = torch.cuda.get_device_properties(0).total_memory / 1024**3 st.write(f"**GPU:** {gpu_name}") st.write(f"**Mémoire utilisée:** {memory_allocated:.1f} GB") st.write(f"**Mémoire réservée:** {memory_reserved:.1f} GB") st.write(f"**Mémoire totale:** {total_memory:.1f} GB") # Barre de progression de la mémoire memory_percent = memory_reserved / total_memory st.progress(memory_percent) except: st.write("**GPU:** CUDA disponible") elif device == "mps": st.info("🍎 MPS (Apple Silicon) détecté") st.write("Génération optimisée pour Apple Silicon") else: st.warning("⚠️ CPU uniquement") st.write("Génération lente - GPU recommandé") # Informations cache if cache_dir: st.write(f"**Cache:** {os.path.basename(cache_dir)}") st.header("📖 Guide d'utilisation") st.markdown(""" ### Étapes: 1. **Chargez une image** 📷 2. **Écrivez un prompt** en anglais ✍️ 3. **Ajustez les paramètres** (optionnel) ⚙️ 4. **Générez** l'image éditée 🚀 ### 💡 Conseils: - **Prompts clairs:** "turn into a robot" - **Instructions directes:** plutôt que descriptions - **Commencez simple:** utilisez les paramètres par défaut - **Testez différents seeds** pour varier les résultats """) st.header("🎯 Exemples de prompts") st.markdown(""" **Transformations:** - `turn him into a cyborg` - `make it a cartoon` - `turn it into a painting` **Ajouts:** - `add sunglasses` - `add a hat` - `add dramatic lighting` **Modifications:** - `make it winter` - `turn the sky purple` - `make him smile` - `change the background to a beach` """) st.header("⚠️ Limitations") st.markdown(""" - Prompts en **anglais** uniquement - Qualité variable selon l'image source - Temps de génération: 30s - 5min selon le matériel - Mémoire requise: 4-8GB VRAM (GPU) """) # Footer st.markdown("---") st.markdown( f"""
🎨 Propulsé par InstructPix2Pix | Device: {device.upper()} | Cache: {"✅" if cache_dir else "❌"} | Documentation
""", unsafe_allow_html=True ) # Nettoyage automatique du cache Ă  la fermeture (optionnel) import atexit def cleanup_cache(): if cache_dir and os.path.exists(cache_dir): try: shutil.rmtree(cache_dir) except: pass atexit.register(cleanup_cache)