next-gen / main_multithread.py
klakenyuo's picture
Fresh start with environment variables for API tokens
0faa32e
Raw
History Blame Contribute Delete
15.4 kB
import streamlit as st
import pandas as pd
import os
from datetime import datetime
import csv
import io
import chardet
import zipfile
import tempfile
import requests
from PIL import Image
from io import BytesIO
import concurrent.futures
import threading
from queue import Queue
import time
import uuid
import logging
import shutil
# Configuration du logging
logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(name)s - %(levelname)s - %(message)s')
logger = logging.getLogger(__name__)
# Paramètres
API_URL = "https://router.huggingface.co/hf-inference/models/black-forest-labs/FLUX.1-dev"
HF_API_KEY = os.getenv('HF_API_KEY')
if not HF_API_KEY:
st.error("Erreur: Le token API Hugging Face n'est pas configuré. Veuillez définir la variable d'environnement HF_API_KEY.")
st.stop()
# Configuration des headers pour l'API Hugging Face
headers = {
"Authorization": f"Bearer {HF_API_KEY}",
}
# File d'attente pour les résultats
results_queue = Queue()
# File d'attente pour les mises à jour de statut des threads
status_queue = Queue()
# File d'attente pour les messages d'erreur
error_queue = Queue()
# Compteur pour suivre la progression
progress_counter = 0
progress_lock = threading.Lock()
# Dictionnaire pour suivre l'état des threads
thread_status = {}
thread_status_lock = threading.Lock()
# Chemin du dossier de sortie et du ZIP
output_dir = None
zip_path = None
zip_lock = threading.Lock()
def detect_encoding(csv_content):
"""Détecte l'encodage du contenu CSV."""
result = chardet.detect(csv_content)
return result['encoding']
def clean_csv_content(csv_content):
"""Nettoie le contenu CSV en supprimant les virgules à l'intérieur des guillemets."""
# Détection de l'encodage
encoding = detect_encoding(csv_content)
st.info(f"Encodage détecté : {encoding}")
try:
# Tentative de décodage avec l'encodage détecté
decoded_content = csv_content.decode(encoding)
except UnicodeDecodeError:
# Si échec, essayer avec latin-1 qui peut lire tous les octets
st.warning("Impossible de décoder avec l'encodage détecté, utilisation de latin-1")
decoded_content = csv_content.decode('latin-1')
output = io.StringIO()
reader = csv.reader(io.StringIO(decoded_content), delimiter=';')
writer = csv.writer(output)
for row in reader:
cleaned_row = []
for cell in row:
# Supprime les virgules à l'intérieur des guillemets
if cell.startswith('"') and cell.endswith('"'):
cell = cell.replace(',', '')
cleaned_row.append(cell)
writer.writerow(cleaned_row)
return output.getvalue().encode('utf-8')
def update_thread_status(thread_id, status, prompt=None):
"""Met à jour le statut d'un thread."""
with thread_status_lock:
if thread_id not in thread_status:
thread_status[thread_id] = {
"status": status,
"prompt": prompt,
"start_time": time.time(),
"end_time": None
}
else:
thread_status[thread_id]["status"] = status
if prompt:
thread_status[thread_id]["prompt"] = prompt
# Si le statut n'est plus "En cours", enregistrer le temps de fin
if status != "En cours" and thread_status[thread_id]["end_time"] is None:
thread_status[thread_id]["end_time"] = time.time()
# Ajout du statut à la file d'attente pour l'affichage
status_queue.put((thread_id, status, prompt))
def generate_image(prompt, idx, total, thread_id):
"""Génère une image à partir d'un prompt et met à jour la progression."""
try:
# Mise à jour du statut du thread
update_thread_status(thread_id, "En cours", prompt)
# Utilisation de l'API Hugging Face pour la génération d'images
response = requests.post(API_URL, headers=headers, json={"inputs": prompt})
if response.status_code == 200:
# Mise à jour de la progression
with progress_lock:
global progress_counter
progress_counter += 1
progress = progress_counter / total
# Mise à jour du statut du thread
update_thread_status(thread_id, "Terminé", prompt)
# Ajout du résultat à la file d'attente
results_queue.put((idx, response.content, prompt, thread_id))
return True
else:
# Mise à jour du statut du thread
error_msg = f"Erreur lors de la génération de l'image {idx+1}: {response.status_code}"
update_thread_status(thread_id, f"Erreur: {response.status_code}", prompt)
# Utilisation de logging au lieu de st.error
logger.error(error_msg)
# Ajout de l'erreur à la file d'attente
error_queue.put((idx, error_msg))
return False
except Exception as e:
# Mise à jour du statut du thread
error_msg = f"Erreur lors de la génération de l'image {idx+1}: {str(e)}"
update_thread_status(thread_id, f"Erreur: {str(e)}", prompt)
# Utilisation de logging au lieu de st.error
logger.error(error_msg)
# Ajout de l'erreur à la file d'attente
error_queue.put((idx, error_msg))
return False
def save_image(image_bytes, output_path):
"""Sauvegarde une image et met à jour le ZIP."""
if image_bytes is None:
return False
try:
# Sauvegarde de l'image
image = Image.open(BytesIO(image_bytes))
image.save(output_path)
# Mise à jour du ZIP
with zip_lock:
global zip_path
if zip_path and os.path.exists(zip_path):
# Supprimer l'ancien ZIP
os.remove(zip_path)
# Créer un nouveau ZIP avec toutes les images
if os.path.exists(os.path.dirname(output_path)):
with zipfile.ZipFile(zip_path, 'w', zipfile.ZIP_DEFLATED) as zipf:
for root, dirs, files in os.walk(os.path.dirname(output_path)):
for file in files:
if file.endswith('.png'):
file_path = os.path.join(root, file)
arcname = os.path.relpath(file_path, os.path.dirname(output_path))
zipf.write(file_path, arcname)
return True
except Exception as e:
error_msg = f"Erreur lors de la sauvegarde de l'image: {str(e)}"
logger.error(error_msg)
error_queue.put((-1, error_msg))
return False
def create_zip_file(directory_path):
"""Crée un fichier ZIP à partir du dossier d'images."""
zip_path = f"{directory_path}.zip"
with zipfile.ZipFile(zip_path, 'w', zipfile.ZIP_DEFLATED) as zipf:
for root, dirs, files in os.walk(directory_path):
for file in files:
if file.endswith('.png'):
file_path = os.path.join(root, file)
arcname = os.path.relpath(file_path, directory_path)
zipf.write(file_path, arcname)
return zip_path
def display_thread_status(thread_status_container):
"""Affiche le statut de tous les threads."""
with thread_status_lock:
# Création d'un DataFrame pour afficher les statuts
status_data = []
for thread_id, info in thread_status.items():
# Calcul du temps écoulé
if info["end_time"] is not None:
# Si le thread a terminé, utiliser le temps de fin
elapsed_time = info["end_time"] - info["start_time"]
else:
# Sinon, calculer le temps écoulé jusqu'à maintenant
elapsed_time = time.time() - info["start_time"]
status_data.append({
"Thread ID": thread_id[:8], # Afficher seulement les 8 premiers caractères
"Statut": info["status"],
"Prompt": info["prompt"] if info["prompt"] else "N/A",
"Temps écoulé": f"{elapsed_time:.1f}s"
})
if status_data:
df = pd.DataFrame(status_data)
thread_status_container.dataframe(df, use_container_width=True)
else:
thread_status_container.info("Aucun thread en cours d'exécution")
def process_results(total_images, progress_bar, image_placeholder, thread_status_container, error_container, download_button_placeholder):
"""Traite les résultats de la file d'attente et met à jour l'interface."""
success_count = 0
processed_count = 0
errors = []
# Boucle principale pour traiter les résultats et mettre à jour l'interface
while processed_count < total_images:
# Traitement des mises à jour de statut
while not status_queue.empty():
thread_id, status, prompt = status_queue.get()
# L'affichage sera mis à jour dans la prochaine itération
# Traitement des erreurs
while not error_queue.empty():
idx, error_msg = error_queue.get()
errors.append(error_msg)
# Mise à jour de l'affichage des erreurs
error_container.error("\n".join(errors))
# Affichage du statut des threads
display_thread_status(thread_status_container)
# Traitement des résultats d'images
if not results_queue.empty():
idx, image_bytes, prompt, thread_id = results_queue.get()
# Sauvegarde de l'image
output_path = os.path.join(output_dir, f"image_{idx}.png")
if save_image(image_bytes, output_path):
success_count += 1
# Affichage de l'image générée
with image_placeholder.container():
st.image(image_bytes, caption=f"Image {idx + 1}/{total_images} - Prompt: {prompt}")
# Mise à jour du bouton de téléchargement
with download_button_placeholder.container():
if os.path.exists(zip_path):
with open(zip_path, 'rb') as f:
st.download_button(
label="Télécharger toutes les images (ZIP)",
data=f,
file_name=os.path.basename(zip_path),
mime="application/zip"
)
processed_count += 1
progress_bar.progress(processed_count / total_images)
else:
time.sleep(0.1) # Petit délai pour éviter de surcharger le CPU
return success_count
def main():
st.title("Générateur d'images à partir de CSV avec Hugging Face (Multi-thread) 🧠🎨")
# Variables globales
global output_dir, zip_path
# Upload du fichier CSV
uploaded_file = st.file_uploader("Choisissez un fichier CSV", type=['csv'])
if uploaded_file is not None:
try:
# Récupération du nom du fichier CSV (sans extension)
csv_filename = os.path.splitext(uploaded_file.name)[0]
# Lecture et nettoyage du CSV
st.info("Étape 1/5 : Nettoyage du fichier CSV...")
csv_content = uploaded_file.read()
cleaned_csv_content = clean_csv_content(csv_content)
# Lecture du CSV nettoyé
df = pd.read_csv(io.BytesIO(cleaned_csv_content))
# Ignorer la première ligne (en-tête)
df = df.iloc[1:].reset_index(drop=True)
# Création du dossier de sortie avec le nom du fichier CSV
st.info("Étape 2/5 : Création du dossier de sortie...")
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
# Création d'un dossier permanent dans le répertoire courant
output_dir = f"{csv_filename}_{timestamp}"
os.makedirs(output_dir, exist_ok=True)
# Création du fichier ZIP
zip_path = f"{output_dir}.zip"
# Affichage des colonnes du CSV
st.info("Étape 3/5 : Affichage des colonnes disponibles...")
st.write("Colonnes disponibles dans le CSV :")
st.write(df.columns.tolist())
# Sélection de la colonne contenant les prompts
prompt_column = st.selectbox("Sélectionnez la colonne contenant les prompts", df.columns.tolist())
# Sélection du nombre de threads
num_threads = st.slider("Nombre de threads parallèles", min_value=1, max_value=10, value=3)
if st.button("Générer les images"):
st.info(f"Étape 4/5 : Début de la génération des images avec Hugging Face (utilisant {num_threads} threads)...")
# Création des éléments d'interface
progress_bar = st.progress(0)
image_placeholder = st.empty()
# Conteneur pour l'état des threads
st.subheader("État des threads")
thread_status_container = st.empty()
# Conteneur pour les erreurs
error_container = st.empty()
# Conteneur pour le bouton de téléchargement
download_button_placeholder = st.empty()
# Réinitialisation du compteur de progression et des statuts
global progress_counter
progress_counter = 0
with thread_status_lock:
thread_status.clear()
# Création du thread pool
with concurrent.futures.ThreadPoolExecutor(max_workers=num_threads) as executor:
# Soumission des tâches
futures = []
for idx, row in df.iterrows():
prompt = row[prompt_column]
thread_id = str(uuid.uuid4()) # Génération d'un ID unique pour chaque thread
futures.append(executor.submit(generate_image, prompt, idx, len(df), thread_id))
# Traitement des résultats en parallèle
success_count = process_results(len(df), progress_bar, image_placeholder, thread_status_container, error_container, download_button_placeholder)
if success_count > 0:
st.success(f"Étape 5/5 : {success_count} images ont été générées avec succès sur {len(df)} tentatives")
st.info(f"Les images ont été sauvegardées dans le dossier '{output_dir}' et le fichier ZIP '{zip_path}'")
else:
st.error("Aucune image n'a pu être générée. Veuillez vérifier votre token API Hugging Face et réessayer.")
except Exception as e:
st.error(f"Une erreur est survenue : {str(e)}")
if __name__ == "__main__":
main()