CoolFace
Apppublic

EKKD/ImplementacionSimpleDataset

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
app.py151 linesDownload Raw Back to root
1import streamlit as st2import numpy as np3import matplotlib.pyplot as plt4from tensorflow.keras.models import load_model5from tensorflow.keras.preprocessing import image6 7import streamlit as st8import pandas as pd9import matplotlib.pyplot as plt10 11import tensorflow as tf12from tensorflow.keras import backend as K13from tensorflow.keras.utils import get_custom_objects14 15# Registrar la clase de la métrica F1Score16@tf.keras.utils.register_keras_serializable()17class F1Score(tf.keras.metrics.Metric):18    def __init__(self, name="f1_score", **kwargs):19        super(F1Score, self).__init__(name=name, **kwargs)20        self.precision = tf.keras.metrics.Precision()21        self.recall = tf.keras.metrics.Recall()22 23    def update_state(self, y_true, y_pred, sample_weight=None):24        self.precision.update_state(y_true, y_pred, sample_weight)25        self.recall.update_state(y_true, y_pred, sample_weight)26 27    def result(self):28        precision = self.precision.result()29        recall = self.recall.result()30        return 2 * ((precision * recall) / (precision + recall + K.epsilon()))31 32    def reset_states(self):33        self.precision.reset_states()34        self.recall.reset_states()35 36# Registrar la clase de la métrica ROC37@tf.keras.utils.register_keras_serializable()38class ROC(tf.keras.metrics.Metric):39    def __init__(self, name="roc", **kwargs):40        super(ROC, self).__init__(name=name, **kwargs)41        self.auc = tf.keras.metrics.AUC()42 43    def update_state(self, y_true, y_pred, sample_weight=None):44        self.auc.update_state(y_true, y_pred, sample_weight)45 46    def result(self):47        return self.auc.result()48 49    def reset_states(self):50        self.auc.reset_states()51 52# Asegúrate de que las clases personalizadas estén registradas53get_custom_objects()['F1Score'] = F1Score54get_custom_objects()['ROC'] = ROC55 56# Ahora puedes cargar el modelo57from tensorflow.keras.models import load_model58model = load_model("modelo_completo_1.keras")59 60# Verifica que se cargue correctamente61model.summary()62 63 64# -------------------------------------------------------------------------------------------------------------------------------------------65# -------------------------------------------------------------------------------------------------------------------------------------------66# -------------------------------------------------------------------------------------------------------------------------------------------67import streamlit as st68import numpy as np69from tensorflow.keras.preprocessing import image70from tensorflow.keras.models import load_model71import matplotlib.pyplot as plt72 73import streamlit as st74import numpy as np75from tensorflow.keras.preprocessing import image76from tensorflow.keras.models import load_model77import matplotlib.pyplot as plt78 79# Cargar el modelo y los pesos solo una vez80@st.cache_resource81def load_model_and_weights():82    model = load_model("modelo_completo_1.keras")  # Ajusta la ruta si es necesario83    model.load_weights("pesos_del_modelo.weights.h5")84    85    # Cargar las clases desde un archivo86    try:87        with open("image_classes.txt", "r") as file:88            image_classes = [line.strip() for line in file]89    except FileNotFoundError:90        st.error("El archivo 'image_classes.txt' no se encontró. Asegúrate de incluirlo en el mismo directorio.")91        image_classes = []92    93    return model, image_classes94 95# Cargar el modelo y las clases96model, image_classes = load_model_and_weights()97 98# Interfaz principal de la aplicación99st.title("Clasificación de Imágenes con IA - CIFAR-10")100 101# Subir la imagen en la interfaz principal102uploaded_file = st.file_uploader("Sube una imagen para clasificar", type=["jpg", "jpeg", "png"])103 104if uploaded_file is not None:105    # Mostrar la imagen cargada106    st.image(uploaded_file, caption="Imagen cargada", use_container_width=True)107 108    # Procesar la imagen para el modelo CIFAR-10 (32x32x3)109    try:110        img = image.load_img(uploaded_file, target_size=(32, 32))  # Redimensionar a 32x32, tamaño de CIFAR-10111        img_array = image.img_to_array(img)112        img_array = np.expand_dims(img_array, axis=0)  # Añadir dimensión batch113        img_array /= 255.0  # Normalizar la imagen (entre 0 y 1)114 115        # Realizar la predicción116        predictions = model.predict(img_array)117        predicted_class = np.argmax(predictions, axis=1)118 119        # Mostrar el resultado de la predicción como texto120        st.subheader("Resultado de la Clasificación")121        st.write(f"**Predicción:** {image_classes[predicted_class[0]]}")122        123        # Mostrar la probabilidad de la clase predicha124        st.write(f"**Confianza en la predicción:** {predictions[0][predicted_class[0]] * 100:.2f}%")125        126        # Crear gráfico de barras para la probabilidad de cada clase127        fig, ax = plt.subplots(figsize=(8, 6))128        129        # Colores diferentes para cada barra130        colors = plt.cm.viridis(np.linspace(0, 1, len(image_classes)))131        ax.bar(image_classes, predictions[0] * 100, color=colors)132        133        ax.set_title("Confianza por Clase", fontsize=16, fontweight='bold')134        ax.set_xlabel("Clases", fontsize=12, fontweight='bold')135        ax.set_ylabel("Confianza (%)", fontsize=12, fontweight='bold')136 137        # Agregar el porcentaje sobre cada barra138        for i, v in enumerate(predictions[0] * 100):139            ax.text(i, v + 2, f"{v:.2f}%", ha='center', va='bottom', fontsize=10, color='black')140 141        # Ajustar etiquetas en el eje X142        ax.set_xticklabels(image_classes, rotation=90, ha='right', fontsize=10)143        ax.tick_params(axis='y', labelsize=10)144        ax.grid(True, axis='y', linestyle='--', alpha=0.7)145 146        # Mostrar el gráfico en la interfaz principal147        st.pyplot(fig)148    149    except Exception as e:150        st.error(f"Error al procesar la imagen: {e}")151