Este proyecto implementa una red de difusión para generar muestras de dígitos MNIST, junto con un clasificador auxiliar. El flujo principal consiste en:
- Entrenar el clasificador auxiliar.
- Entrenar la red de difusión en modo condicional puro.
- Entrenar la misma red en modo CFG.
- Ejecutar el muestreo para generar grillas de muestras tanto para la versión condicional como para CFG.
Instala las dependencias del proyecto:
pip install -r requirements.txtLos scripts esperan que los datos de MNIST estén disponibles en:
- data/data_entrenamiento
- data/data_prueba
Para descargarlos, se debe usar el siguiente comando:
python download_data.pyEste paso entrena el clasificador auxiliar y guarda los resultados en la carpeta outputs.
python train_clf.pyParámetros por defecto usados por el script:
- --device: cuda si está disponible, de lo contrario cpu
- --epochs: 100
- --batch-size: 4096
- --lr: 0.001
- --seed: 42
Archivos generados:
- outputs/clasificador.pt
- outputs/clasificador_perdida.png
Para entrenar la red sin CFG, se debe usar label-dropout en 0.0.
python train.py --label-dropout 0.0Esto genera:
- outputs/modelo_cond.pt
- outputs/modelo_cond_perdida.png
Para entrenar la variante con classifier-free guidance, se usa el valor por defecto de label-dropout.
python train.pyEquivalente a:
python train.py --label-dropout 0.3Esto genera:
- outputs/modelo_cfg.pt
- outputs/modelo_cfg_perdida.png
Parámetros por defecto usados por train.py:
- --epochs: 1000
- --batch-size: 4096
- --lr: 1e-4
- --seed: 42
- --device: cuda
- --label-dropout: 0.3
- --out: outputs
Una vez entrenados ambos modelos, se puede ejecutar el muestreo con:
python sample.pyEste script:
- carga outputs/modelo_cond.pt para generar muestras en modo condicional,
- carga outputs/modelo_cfg.pt para generar muestras en modo CFG,
Archivos generados:
- outputs/muestras_difusion_cond.png
- outputs/proceso_difusion_cond.png
pip install -r requirements.txt
python download_data.py
python train_clf.py
python train.py --label-dropout 0.0
python train.py
python sample.pyNota: el caso sin CFG debe ejecutarse explícitamente con --label-dropout 0.0, porque el valor por defecto del entrenamiento de la red de difusión es 0.3 y activa CFG.