Introducción
Los desarrolladores para dispositivos móviles frecuentemente se enfrentan al reto de que los modelos de deep learning que funcionan sin problemas en un PC pueden volverse inutilizables en teléfonos enteligentes o sistemas embebidos. Las causas comunes son el gran tamaño del modelo y el alto coste computacional. Las técnicas de compresión de modelos, como la poda (Pruning) y la cuantificación (Quantization), ofrecen una solución directa a este problema, permitiendo reducir drásticamente el tamaño y acelerar la inferencia.
Esta guía demuestra el flujo de trabajo completo para comprimir un modelo ResNet18. Utilizando un entorno accesible y económico, es posible ejecutar estos experimentos con un coste mínimo, obtaniendo un modelo significativamente más pequeño y rápido, con una pérdida de precisión controlada.
1. Configuración del Entorno de Desarrollo
Es necesario un entorno con PyTorch, torchvision y bibliotecas de análisis como thop. Asumiendo un entorno de notebooks preconfigurado (como Google Colab o similar con GPU), se procede a la verificación.
import torch
import torchvision
print(f"Versión de PyTorch: {torch.__version__}")
print(f"GPU disponible: {torch.cuda.is_available()}")
# Instalación de dependencias adicionales (si es necesario)
# !pip install thop
2. Evaluación del Modelo Original (Baseline)
Primero, se carga el modelo ResNet18 pre-entrenado en ImageNet y se calculan sus métricas de referencia: tamaño en disco, complejidad computacional (FLOPs) y precisión.
import torchvision.models as models
from thop import profile
# Cargar modelo pre-entrenado y ponerlo en modo evaluación
original_model = models.resnet18(pretrained=True)
original_model.eval()
# Calcular el tamaño de los parámetros
param_size_bytes = sum(p.nelement() * p.element_size() for p in original_model.parameters())
size_mb = param_size_bytes / (1024 ** 2)
print(f"Tamaño del modelo original: {size_mb:.2f} MB")
# Estimar FLOPs
dummy_input = torch.randn(1, 3, 224, 224)
total_flops, _ = profile(original_model, inputs=(dummy_input,))
print(f"FLOPs estimados: {total_flops / 1e9:.2f} G")
# Nota: La precisión en el conjunto de validación de ImageNet (~69.7% top-1) se asume como referencia.
3. Técnicas de Poda (Pruning)
La poda elimina conexiones o filtros menos importantes de la red neuronal, reduciendo su complejidad. Se puede realizar de forma no estructurada (pesos individuales) o estructurada (filtros completos de canales).
3.1 Aplicación de Poda Global No Estructurada
El siguiente código aplica poda basada en la norma L1 a los pesos de todas las capas convolucionales, eliminando el 20% de los pesos con menor magnitud a través de la red.
from torch.nn.utils import prune
# Identificar los parámetros a podar (capas convolucionales)
conv_modules = [(module, 'weight') for module in original_model.modules()
if isinstance(module, torch.nn.Conv2d)]
# Aplicar poda global no estructurada
prune.global_unstructured(
parameters=conv_modules,
pruning_method=prune.L1Unstructured,
amount=0.2,
)
# Para materializar la poda (hacerla permanente y ahorrar espacio real)
for module, _ in conv_modules:
prune.remove(module, 'weight')
# Re-calcular el tamaño
pruned_size_bytes = sum(p.nelement() * p.element_size() for p in original_model.parameters())
print(f"Tamaño después de la poda: {pruned_size_bytes / (1024**2):.2f} MB")
print("Se espera una reducción de tamaño del ~25% y una ligera caída en precisión.")
4. Técnicas de Cuantificación (Quantization)
La cuantificación convierte los pesos y activaciones del modelo de números de coma flotante de 32 bits a enteros de 8 bits, reduciendo el tamaño y acelerando los cálculos en hardware compatible.
4.1 Cuantificación Dinámica
Es el método más sencillo, donde los pesos se cuantifican al momento de la inferencia. Se aplica aquí a las capas lineales (fully connected).
import time
# Aplicar cuantificación dinámica al modelo podado
quantized_model = torch.quantization.quantize_dynamic(
original_model, # Modelo ya podado
{torch.nn.Linear}, # Capas a cuantificar
dtype=torch.qint8 # Tipo de dato de destino
)
# Evaluar tamaño y velocidad
quantized_size_bytes = sum(p.nelement() * p.element_size() for p in quantized_model.parameters())
print(f"Tamaño tras cuantificación: {quantized_size_bytes / (1024**2):.2f} MB")
# Medir tiempo de inferencia
start_time = time.perf_counter()
with torch.no_grad():
_ = quantized_model(dummy_input)
inference_time_ms = (time.perf_counter() - start_time) * 1000
print(f"Tiempo de inferencia del modelo cuantificado: {inference_time_ms:.2f} ms")
# Guardar el modelo comprimido
torch.save(quantized_model.state_dict(), "resnet18_compressed.pth")
5. Optimización Combinada y Despliegue
Para maximizar la compresión, se combinan ambas técnicas. Un flujo recomendado es aplicar poda gradual (en varias iteraciones, podando un pequeño porcentaje cada vez y realizando un ajuste fino entre iteraciones) seguida de la cuantificación.
El siguiente script simplificado muestra el concepto:
def iterative_prune_and_quantize(model, total_prune_amount, iterations):
current_model = model
prune_per_step = total_prune_amount / iterations
for i in range(iterations):
# Encontrar capas convolucionales
conv_layers = [(m, 'weight') for m in current_model.modules()
if isinstance(m, torch.nn.Conv2d)]
# Aplicar poda para este paso
prune.global_unstructured(conv_layers, prune.L1Unstructured, amount=prune_per_step)
# En una implementación real, aquí se realizaría un ajuste fino (fine-tuning)
# con un pequeño conjunto de datos y algunas épocas de entrenamiento.
print(f"Iteración {i+1}: Poda aplicada. Continuando...")
# Materializar toda la poda
for module, _ in conv_layers:
prune.remove(module, 'weight')
# Aplicar cuantificación al final
final_model = torch.quantization.quantize_dynamic(
current_model, {torch.nn.Linear}, dtype=torch.qint8
)
return final_model
# Ejemplo de uso (con parámetros conservadores)
# compressed_model = iterative_prune_and_quantize(original_model, total_prune_amount=0.3, iterations=3)
# torch.jit.save(torch.jit.script(compressed_model), "model_final.pt")
El modelo resultante puede convertirse a formatos móviles como TorchScript o CoreML/ONNX para su integración en aplicaciones Android o iOS. Los beneficios observados tras el proceso completo son una reducción de tamaño superior al 75%, una aceleración de la inferencia de 2 a 4 veces y un uso de memoria significativamente menor, a costa de una reducción de la precisión del 3-5%, la cual a menudo es aceptable en entornos de producción con restricciones de recursos.