Proyecto de clasificación de imágenes de cacao utilizando modelos de visión por computadora basados en Transformers y modelos de la biblioteca timm. Este proyecto permite entrenar y hacer inferencia con múltiples arquitecturas de modelos para clasificación binaria de imágenes de cacao.
- Descripción
- Requisitos
- Instalación
- Estructura del Proyecto
- Uso
- Configuración
- Modelos Disponibles
- Resultados
Este proyecto implementa un pipeline completo para:
- Entrenamiento: Fine-tuning de modelos pre-entrenados en un dataset de clasificación binaria de cacao
- Evaluación: Métricas de precisión, recall, F1, especificidad y accuracy
- Inferencia: Predicción en imágenes nuevas con modelos entrenados
- Tracking: Integración con Weights & Biases (wandb) para monitoreo de experimentos
El dataset utilizado es CristianR8/BINARY-IA4CACAO-RGB de Hugging Face, que contiene imágenes de cacao para clasificación binaria.
- Python 3.10+
- CUDA-capable GPU (recomendado para entrenamiento)
- jq (para scripts de bash)
Las dependencias principales se encuentran en requirements.txt:
torch,torchvision- PyTorch y visióntransformers- Modelos y procesadores de Hugging Facetimm- Modelos de visión adicionalesaccelerate- Entrenamiento distribuidodatasets- Manejo de datasetswandb- Tracking de experimentosevaluate- Métricas de evaluación
- Clonar el repositorio (si aplica):
git clone <repository-url>
cd cacao-hf- Crear y activar entorno virtual:
python -m venv env_cacao
source env_cacao/bin/activate # En Linux/Mac
# o
env_cacao\Scripts\activate # En Windows- Instalar dependencias:
pip install --upgrade pip
pip install -r requirements.txt- Instalar jq (para scripts bash):
# Ubuntu/Debian
sudo apt install jq
# MacOS
brew install jq- Configurar Weights & Biases (opcional pero recomendado):
wandb logincacao-hf/
├── main.py # Script principal de entrenamiento
├── inference.py # Script de inferencia
├── models.json # Configuración de modelos disponibles
├── train.sh # Script bash para entrenar un modelo
├── all_train.sh # Script para entrenar todos los modelos secuencialmente
├── requirements.txt # Dependencias Python
├── utils/ # Utilidades del proyecto
│ ├── __init__.py
│ ├── preprocessor.py # Configuración de preprocesadores
│ ├── timmprocessor.py # Procesador para modelos timm
│ ├── timadapter.py # Adaptador para modelos timm
│ └── metrics.py # Métricas personalizadas (especificidad)
├── outputs_HSV/ # Modelos entrenados guardados aquí
│ ├── outputs_vit_base/
│ ├── outputs_vit_large/
│ ├── outputs_convnext_xlarge/
│ └── outputs_mobilenetv3_large/
├── inferencia/ # Imágenes de ejemplo para inferencia
├── logs/ # Logs de entrenamiento
└── notebooks/ # Jupyter notebooks de análisis
La forma más fácil de usar los modelos es a través de la aplicación web Streamlit, que es móvil-friendly y permite clasificar imágenes de granos de cacao de forma interactiva.
- Crear entorno virtual e instalar Streamlit (si no está instalado):
# Crear entorno virtual
python3 -m venv streamlit_env
# Activar entorno virtual
source streamlit_env/bin/activate
# Instalar dependencias
pip install --upgrade pip
pip install streamlit pandas- Ejecutar la aplicación:
# Opción 1: Usar el script (recomendado - detecta y crea entorno automáticamente)
./run_app.sh
# Opción 2: Ejecutar directamente (después de activar entorno virtual)
source streamlit_env/bin/activate
streamlit run app.pyNota: El script run_app.sh detecta automáticamente si existe un entorno virtual con Streamlit instalado, y si no existe, lo crea e instala las dependencias automáticamente.
- Abrir en el navegador: La aplicación se abrirá automáticamente en
http://localhost:8501
- 📸 Carga de imágenes: Sube imágenes desde tu dispositivo o toma fotos con la cámara
- 🍫 Clasificación en tiempo real: Obtén resultados instantáneos con probabilidades
- 📱 Diseño móvil: Optimizada para usar en smartphones y tablets
- 🎯 Múltiples modelos: Selecciona entre diferentes modelos entrenados (ViT, ConvNeXt, MobileNet)
- 📊 Visualización de resultados: Gráficos y tablas con todas las probabilidades
- ⚡ Rápida: Usa GPU automáticamente si está disponible
- Ejecuta la app en tu servidor/computadora
- Accede desde tu móvil usando la IP del servidor:
http://TU_IP:8501 - Para acceso desde cualquier dispositivo en la red local:
streamlit run app.py --server.address 0.0.0.0 --server.port 8501La aplicación puede clasificar granos de cacao en 6 categorías:
- ✅ Fermentado: Grano fermentado correctamente
- 🍄 Hongo: Grano afectado por hongos
- 🐛 Insecto: Grano dañado por insectos
⚠️ Insufi_fermen: Grano con fermentación insuficiente- ⚫ Pizarroso: Grano pizarroso (defecto de color)
- 🟣 Violeta: Grano violeta (defecto de color)
Usa el script train.sh para entrenar un modelo específico:
# Ver modelos disponibles
./train.sh --list
# Entrenar un modelo específico (usa valores por defecto del config)
./train.sh -m vit_base
# Entrenar con parámetros personalizados
./train.sh -m vit_base \
-b 32 \ # batch size
-e 50 \ # épocas
-l 2e-5 # learning rateParámetros del script train.sh:
-m MODEL_TYPE: Tipo de modelo (requerido, vermodels.json)-d DATASET: Nombre del dataset (default: del config)-o OUTPUT_DIR: Directorio de salida (default:./outputs_HSV/outputs_${MODEL_TYPE})-b BATCH_SIZE: Tamaño de batch-e EPOCHS: Número de épocas-l LR: Learning rate
El script all_train.sh entrena todos los modelos definidos en models.json:
# Ejecutar entrenamiento masivo
./all_train.shEste script:
- Entrena cada modelo secuencialmente
- Usa los mismos hiperparámetros para todos (configurables en el script)
- Guarda logs en
./logs/ - Genera un resumen al finalizar
python main.py \
--dataset_name CristianR8/BINARY-IA4CACAO-RGB \
--model_name_or_path timm/vit_base_patch16_224.augreg_in21k_ft_in1k \
--output_dir ./outputs_HSV/outputs_vit_base \
--with_tracking \
--report_to wandb \
--do_eval \
--num_train_epochs 25 \
--per_device_train_batch_size 16 \
--learning_rate 1e-4 \
--trust_remote_code \
--ignore_mismatched_sizesParámetros principales de main.py:
--dataset_name: Nombre del dataset en Hugging Face--model_name_or_path: Ruta o identificador del modelo--output_dir: Directorio donde guardar el modelo entrenado--with_tracking: Habilitar tracking con wandb/tensorboard--do_eval: Ejecutar evaluación durante el entrenamiento--num_train_epochs: Número de épocas--per_device_train_batch_size: Batch size por dispositivo--learning_rate: Learning rate inicial--trust_remote_code: Necesario para modelos timm personalizados
El script inference.py permite hacer predicciones con modelos entrenados.
# Inferencia en una imagen
python inference.py \
--model_id ./outputs_HSV/outputs_vit_base \
--input inferencia/violeta.jpg \
--out resultados.csv
# Inferencia en un directorio completo
python inference.py \
--model_id ./outputs_HSV/outputs_vit_base \
--input inferencia/ \
--out resultados.csv \
--batch_size 32
# Inferencia con top-k predicciones
python inference.py \
--model_id ./outputs_HSV/outputs_vit_base \
--input inferencia/ \
--out resultados.csv \
--topk 3--model_id: Ruta al modelo entrenado (local) o ID en Hugging Face Hub--input: Imagen individual, directorio de imágenes, o archivo .txt/.csv con rutas--out: Archivo CSV de salida con predicciones (default:inference_results.csv)--batch_size: Tamaño de batch para procesamiento (default: 16)--topk: Número de top predicciones a mostrar (default: 5)--device: Dispositivo a usar:cudaocpu(default: auto-detecta)
El CSV generado contiene:
path: Ruta de la imagenpred: Etiqueta predichatopk_labels: Top-k etiquetas separadas por|topk_probs: Probabilidades del top-k separadas por|prob_{label}: Probabilidad para cada clase
Ejemplo de salida:
path,pred,topk_labels,topk_probs,prob_class_0,prob_class_1
inferencia/violeta.jpg,class_1,class_1|class_0,0.987654|0.012346,0.012346,0.987654# Ejemplo de uso programático
from inference import discover_images, run_batch
from pathlib import Path
from transformers import AutoImageProcessor, AutoModelForImageClassification
import torch
from PIL import Image
# Cargar modelo
model_id = "./outputs_HSV/outputs_vit_base"
processor = AutoImageProcessor.from_pretrained(model_id, trust_remote_code=True)
model = AutoModelForImageClassification.from_pretrained(model_id, trust_remote_code=True)
model.eval().to("cuda")
# Procesar imagen
image = Image.open("inferencia/violeta.jpg").convert("RGB")
inputs = processor(images=image, return_tensors="pt").to("cuda")
# Predicción
with torch.no_grad():
outputs = model(**inputs)
probs = torch.nn.functional.softmax(outputs.logits, dim=-1)
pred_class = torch.argmax(probs, dim=-1).item()
print(f"Clase predicha: {pred_class}")
print(f"Probabilidades: {probs}")Este archivo contiene la configuración de todos los modelos disponibles:
{
"models": {
"vit_base": {
"name": "timm/vit_base_patch16_224.augreg_in21k_ft_in1k",
"batch_size": 16,
"learning_rate": "1e-4",
"description": "ViT Base model (86M parameters)"
}
},
"default_settings": {
"dataset": "CristianR8/BINARY-IA4CACAO-RGB",
"epochs": 25,
"seed": 1337,
"save_total_limit": 3
}
}Crea un archivo .env para configuraciones sensibles:
# Para subir modelos a Hugging Face Hub
HF_API_KEY=tu_token_aqui
# Para Weights & Biases
WANDB_API_KEY=tu_token_aquiEl proyecto soporta múltiples arquitecturas de modelos. Algunos ya están entrenados, otros están configurados y listos para entrenar.
vit_base: ViT Base (86M parámetros) ✅vit_large: ViT Large (632M parámetros) ✅convnext_xlarge: ConvNeXt XLarge (350M parámetros) ✅mobilenetv3_large: MobileNetV3 Large (5.5M parámetros) ✅
vit_base: ViT Base (86M parámetros)vit_large: ViT Large (632M parámetros)
convnext_xlarge: ConvNeXt XLarge (350M parámetros)convnext_xxlarge: ConvNeXt XXLarge (846M parámetros)
efficientnet_b0aefficientnet_b7: Variantes de EfficientNet
mobilenetv3_small: MobileNetV3 Small (2.5M parámetros)mobilenetv3_large: MobileNetV3 Large (5.5M parámetros)
swin_large: Swin Transformer Largeeva_giant: EVA Giant (1B+ parámetros)maxvit_xlarge: MaxViT XLargebeit_large: BEiT Largeresnet_34/50/101/152: Variantes de ResNetvgg13/16/19: Variantes de VGG⚠️ (Configurados pero no entrenados aún)
Nota: Los modelos marcados con ✅ ya están entrenados y listos para inferencia. Los demás están configurados en models.json y pueden entrenarse usando train.sh.
Ver todos los modelos disponibles:
./train.sh --listLos modelos entrenados se guardan en outputs_HSV/outputs_{model_type}/ con:
model.safetensors: Pesos del modeloconfig.json: Configuración del modelopreprocessor_config.json: Configuración del preprocesadorall_results.json: Métricas de evaluación finales
Ejemplo de all_results.json:
{
"eval_accuracy": 0.95,
"eval_precision": 0.94,
"eval_recall": 0.96,
"eval_f1": 0.95,
"eval_specificity": 0.94,
"eval_train_loss": 0.12,
"eval_epoch": 24,
"eval_step": 1500
}Primero, verifica qué modelos tienes disponibles:
ls -la outputs_HSV/Deberías ver directorios como:
outputs_vit_base/outputs_vit_large/outputs_convnext_xlarge/outputs_mobilenetv3_large/
python inference.py \
--model_id ./outputs_HSV/outputs_vit_base \
--input inferencia/violeta.jpg \
--out resultado_simple.csvpython inference.py \
--model_id ./outputs_HSV/outputs_vit_base \
--input inferencia/ \
--out resultados_lote.csv \
--batch_size 16Puedes comparar predicciones de diferentes modelos:
# Modelo 1
python inference.py \
--model_id ./outputs_HSV/outputs_vit_base \
--input inferencia/ \
--out resultados_vit_base.csv
# Modelo 2
python inference.py \
--model_id ./outputs_HSV/outputs_convnext_xlarge \
--input inferencia/ \
--out resultados_convnext.csvEl archivo CSV generado contiene:
- path: Ruta de la imagen procesada
- pred: Clase predicha (ej:
class_0oclass_1) - topk_labels: Top-k clases más probables
- topk_probs: Probabilidades correspondientes
- prob_{class}: Probabilidad para cada clase
Ejemplo de interpretación:
path,pred,topk_labels,topk_probs,prob_class_0,prob_class_1
inferencia/violeta.jpg,class_1,class_1|class_0,0.987654|0.012346,0.012346,0.987654Esto significa:
- La imagen fue clasificada como
class_1con 98.77% de confianza - La probabilidad de
class_0es 1.23%
Si necesitas integrar la inferencia en tu código:
import torch
from PIL import Image
from transformers import AutoImageProcessor, AutoModelForImageClassification
# Cargar modelo y procesador
model_path = "./outputs_HSV/outputs_vit_base"
processor = AutoImageProcessor.from_pretrained(model_path, trust_remote_code=True)
model = AutoModelForImageClassification.from_pretrained(model_path, trust_remote_code=True)
model.eval()
# Mover a GPU si está disponible
device = "cuda" if torch.cuda.is_available() else "cpu"
model = model.to(device)
# Cargar y preprocesar imagen
image = Image.open("inferencia/violeta.jpg").convert("RGB")
inputs = processor(images=image, return_tensors="pt").to(device)
# Hacer predicción
with torch.no_grad():
outputs = model(**inputs)
logits = outputs.logits
probs = torch.nn.functional.softmax(logits, dim=-1)
pred_class_id = torch.argmax(probs, dim=-1).item()
# Obtener nombre de la clase
id2label = model.config.id2label
pred_class_name = id2label[pred_class_id]
confidence = probs[0][pred_class_id].item()
print(f"Clase predicha: {pred_class_name}")
print(f"Confianza: {confidence:.4f}")
print(f"Probabilidades: {probs}")-
Preprocesamiento: El modelo espera imágenes RGB preprocesadas según su configuración. El
AutoImageProcessorse encarga de esto automáticamente. -
Batch Processing: Para múltiples imágenes, usa
batch_size > 1para mejor rendimiento. -
GPU vs CPU:
- GPU es mucho más rápido para inferencia
- CPU funciona pero es más lento
- El script detecta automáticamente CUDA si está disponible
-
Modelos Grandes: Modelos como
vit_largeoconvnext_xlargerequieren más memoria GPU. -
Formato de Imágenes: Acepta JPG, PNG, BMP, TIFF, WEBP.
- Verifica que el modelo esté entrenado en
outputs_HSV/ - Asegúrate de usar la ruta correcta relativa o absoluta
- Reduce el
batch_sizeen inferencia - Usa un modelo más pequeño
- Procesa imágenes de una en una
- Añade
--trust_remote_codeotrust_remote_code=Trueen Python
- Instala jq:
sudo apt install jq(Linux) obrew install jq(Mac)
- Los modelos se entrenan con mixed precision (bf16) para eficiencia
- El dataset se divide automáticamente en train/test
- Las transformaciones de data augmentation se aplican durante el entrenamiento
- Los checkpoints se guardan por época o por pasos según configuración
- La aplicación Streamlit está optimizada para móviles y funciona en cualquier dispositivo con navegador web
- Para usar la app en móvil, ejecuta Streamlit con
--server.address 0.0.0.0para acceso desde la red local
[Especificar licencia si aplica]
[Especificar contribuidores si aplica]
- Hugging Face por las herramientas de Transformers
- timm por los modelos de visión adicionales
- Weights & Biases por la plataforma de tracking