tutoriales.com

Interoperabilidad de Modelos de IA: Exportando de TensorFlow a PyTorch y Viceversa con ONNX

Este tutorial explora la exportación e importación de modelos de inteligencia artificial entre los frameworks TensorFlow y PyTorch utilizando el formato Open Neural Network Exchange (ONNX). Dominarás las técnicas para convertir tus modelos, asegurando su compatibilidad y maximizando la flexibilidad en diferentes entornos de despliegue y desarrollo.

Intermedio20 min de lectura13 views
Reportar error

🚀 Introducción a la Interoperabilidad de Modelos de IA

En el dinámico mundo de la inteligencia artificial, es común encontrar equipos o proyectos que utilizan diferentes frameworks de aprendizaje profundo, como TensorFlow y PyTorch. Si bien ambos son potentes y versátiles, la incompatibilidad entre sus formatos nativos puede generar desafíos significativos, especialmente al querer compartir modelos, desplegarlos en entornos heterogéneos o aprovechar las fortalezas específicas de cada uno en distintas fases del ciclo de vida del modelo.

Aquí es donde entra en juego la interoperabilidad. La capacidad de mover un modelo de un framework a otro de forma fluida es crucial para la eficiencia, la colaboración y la flexibilidad. Esta guía se centrará en cómo lograr esta interoperabilidad utilizando ONNX (Open Neural Network Exchange), un formato estándar abierto diseñado precisamente para representar modelos de machine learning.

📌 Nota: ONNX permite a los desarrolladores de IA mover modelos entre diferentes marcos de trabajo (como PyTorch, TensorFlow, Keras, scikit-learn, etc.) y optimizarlos para el despliegue en varios entornos (en la nube, en el borde, móviles).

🎯 ¿Qué es ONNX y Por Qué es Importante?

ONNX es un formato de representación de grafos de computación para modelos de aprendizaje automático. Es un estándar abierto que define un conjunto común de operadores y un formato de archivo para representar modelos. Esto significa que un modelo entrenado en un framework (por ejemplo, PyTorch) puede ser exportado a ONNX y luego importado y ejecutado en otro framework (como TensorFlow) que soporte ONNX.

Ventajas Clave de ONNX:

  • Flexibilidad de Framework: Permite a los desarrolladores elegir el framework más adecuado para cada etapa (entrenamiento, inferencia) sin estar atados a uno solo.
  • Optimización del Despliegue: Facilita la optimización de modelos para diferentes plataformas de hardware y software a través de runtimes optimizados para ONNX (como ONNX Runtime).
  • Colaboración: Simplifica el intercambio de modelos entre equipos que utilizan diferentes frameworks.
  • Acceso a Herramientas Especializadas: Permite aprovechar herramientas específicas de un framework para ciertas tareas (ej. visualización en TensorBoard, optimización con TensorFlow Lite) incluso si el modelo fue entrenado en otro.
Entrenamiento Inferencia / Despliegue Framework A (PyTorch) Framework B (TensorFlow) ONNX Model ONNX Runtime Inferencia / Despliegue

🛠️ Herramientas Necesarias

Para seguir este tutorial, necesitarás tener instaladas las siguientes librerías:

  • Python 3.x
  • PyTorch
  • TensorFlow
  • onnx
  • onnxruntime
  • tf2onnx (para exportar de TensorFlow a ONNX)
  • onnx-tensorflow (para importar de ONNX a TensorFlow, aunque a veces onnxruntime y tensorflow son suficientes para la inferencia, onnx-tensorflow facilita la conversión a formato nativo de TF)

Instalación de Dependencias:

pip install torch torchvision tensorflow onnx onnxruntime tf2onnx onnx-tensorflow
💡 Consejo: Se recomienda usar un entorno virtual (como `venv` o `conda`) para gestionar las dependencias de tu proyecto y evitar conflictos.

➡️ Exportando un Modelo de PyTorch a ONNX

El proceso de exportación de PyTorch a ONNX es relativamente sencillo gracias a la función torch.onnx.export. Esta función traza el modelo mientras ejecuta un pase forward con una entrada de ejemplo, registrando todas las operaciones necesarias.

Paso 1: Definir y Entrenar un Modelo Simple en PyTorch

Comencemos con un modelo de red neuronal convolucional (CNN) simple en PyTorch. Para los fines de este tutorial, no es necesario un entrenamiento exhaustivo; nos centraremos en la estructura del modelo.

import torch
import torch.nn as nn
import torch.optim as optim
from torchvision import datasets, transforms

# 1. Definir el modelo PyTorch
class SimpleCNN(nn.Module):
    def __init__(self):
        super(SimpleCNN, self).__init__()
        self.conv1 = nn.Conv2d(1, 10, kernel_size=5)
        self.relu1 = nn.ReLU()
        self.pool1 = nn.MaxPool2d(2)
        self.conv2 = nn.Conv2d(10, 20, kernel_size=5)
        self.relu2 = nn.ReLU()
        self.drop = nn.Dropout2d(0.5)
        self.pool2 = nn.MaxPool2d(2)
        self.fc = nn.Linear(320, 10) # 320 = 20 * 4 * 4, as input image is 28x28 for MNIST

    def forward(self, x):
        x = self.pool1(self.relu1(self.conv1(x)))
        x = self.drop(self.pool2(self.relu2(self.conv2(x))))
        x = x.view(-1, 320) # Aplanar la salida para la capa fully connected
        x = self.fc(x)
        return x

# 2. Instanciar el modelo y cargar pesos (o entrenar)
model_pytorch = SimpleCNN()
# En un escenario real, cargarías pesos pre-entrenados o entrenarías el modelo aquí.
# Para este tutorial, inicializaremos con pesos aleatorios.

print("Modelo PyTorch creado:")
print(model_pytorch)

# 3. Crear una entrada de ejemplo
# Para MNIST, las imágenes son 1x28x28 (canales, alto, ancho)
example_input = torch.randn(1, 1, 28, 28)

print(f"Forma de la entrada de ejemplo: {example_input.shape}")

Paso 2: Exportar el Modelo a ONNX

Ahora, utilizaremos torch.onnx.export para convertir el modelo PyTorch a un archivo ONNX.

import os

onnx_model_path = "pytorch_model.onnx"

try:
    torch.onnx.export(
        model_pytorch,             # Modelo a exportar
        example_input,           # Entrada de ejemplo para trazar el grafo
        onnx_model_path,         # Ruta donde se guardará el modelo ONNX
        export_params=True,      # Exportar pesos del modelo
        opset_version=11,        # Versión de ONNX Opset. La 11 es común y estable.
        do_constant_folding=True, # Eliminar operaciones constantes del grafo
        input_names=['input'],   # Nombre de la entrada del modelo
        output_names=['output'], # Nombre de la salida del modelo
        dynamic_axes={
            'input': {0: 'batch_size'},   # Permite tamaño de batch dinámico
            'output': {0: 'batch_size'}
        }
    )
    print(f"Modelo PyTorch exportado con éxito a {onnx_model_path}")
except Exception as e:
    print(f"Error al exportar el modelo PyTorch a ONNX: {e}")
🔥 Importante: La `opset_version` es crucial. Asegúrate de usar una versión compatible con las operaciones de tu modelo y el *runtime* ONNX que vayas a usar. La versión 11 es una buena opción para empezar.

⬅️ Importando un Modelo ONNX en TensorFlow (para Inferencia)

Una vez que tenemos nuestro modelo en formato ONNX, podemos importarlo en TensorFlow. Para la inferencia, usaremos onnxruntime directamente en Python, que es altamente optimizado. También mostraremos cómo convertir el modelo ONNX a un formato nativo de TensorFlow si fuera necesario para otras operaciones o integraciones más profundas.

Paso 1: Cargar y Ejecutar el Modelo ONNX con ONNX Runtime

onnxruntime es el runtime de referencia para modelos ONNX y ofrece un excelente rendimiento. Es ideal para la inferencia.

import onnxruntime as rt
import numpy as np

# 1. Cargar el modelo ONNX
sess = rt.InferenceSession(onnx_model_path, providers=['CPUExecutionProvider'])

# Obtener los nombres de entrada y salida del modelo ONNX
input_name = sess.get_inputs()[0].name
output_name = sess.get_outputs()[0].name

print(f"Nombre de entrada del modelo ONNX: {input_name}")
print(f"Nombre de salida del modelo ONNX: {output_name}")

# 2. Preparar una entrada de ejemplo para ONNX Runtime
# ONNX Runtime espera entradas en formato NumPy.
# Asegúrate de que la forma coincida con la entrada original (1, 1, 28, 28) y el tipo de dato.
input_onnx = example_input.numpy().astype(np.float32)

# 3. Realizar inferencia
onnx_output = sess.run([output_name], {input_name: input_onnx})

print("Inferencia realizada con ONNX Runtime.")
print(f"Forma de la salida ONNX: {onnx_output[0].shape}")
print(f"Primeros 5 valores de la salida ONNX: {onnx_output[0][0, :5]}")

Paso 2: Convertir ONNX a Formato Nativo de TensorFlow (Opcional)

Si necesitas trabajar con el modelo en el ecosistema de TensorFlow (por ejemplo, para visualización en TensorBoard, o para usar con TensorFlow Lite o TensorFlow Serving en su formato SavedModel), puedes usar la librería onnx-tensorflow.

⚠️ Advertencia: La conversión de ONNX a TensorFlow SavedModel puede no ser perfecta para todos los modelos, especialmente aquellos con operaciones muy personalizadas o que no tienen una correspondencia directa en TensorFlow. Siempre verifica la equivalencia de la salida.
from onnx_tf.backend import prepare
import tensorflow as tf

tf_model_path = "tf_model_from_onnx"

try:
    # Cargar el modelo ONNX
    onnx_model = onnx.load(onnx_model_path)

    # Preparar el backend de TensorFlow a partir del modelo ONNX
    tf_rep = prepare(onnx_model)

    # Exportar el modelo a formato SavedModel de TensorFlow
    tf_rep.export_graph(tf_model_path)
    print(f"Modelo ONNX convertido a TensorFlow SavedModel en {tf_model_path}")

    # Cargar y probar el modelo SavedModel
    loaded_tf_model = tf.saved_model.load(tf_model_path)
    infer = loaded_tf_model.signatures["serving_default"]

    # Preparar la entrada para el modelo de TensorFlow (ej. de un tensor PyTorch a un tensor TF)
    tf_input = tf.convert_to_tensor(example_input.numpy(), dtype=tf.float32)
    tf_output_from_saved_model = infer(input=tf_input)[output_name]

    print("Inferencia realizada con el modelo SavedModel de TensorFlow.")
    print(f"Forma de la salida de TF SavedModel: {tf_output_from_saved_model.shape}")
    print(f"Primeros 5 valores de la salida de TF SavedModel: {tf_output_from_saved_model[0, :5]}")

    # Comparar resultados (opcional)
    print("\nComparación de resultados:")
    print(f"ONNX Runtime (primera predicción): {np.argmax(onnx_output[0][0])}")
    print(f"TF SavedModel (primera predicción): {np.argmax(tf_output_from_saved_model.numpy()[0])}")

except Exception as e:
    print(f"Error al convertir ONNX a TensorFlow SavedModel o al probarlo: {e}")


➡️ Exportando un Modelo de TensorFlow a ONNX

Exportar de TensorFlow (específicamente TensorFlow 2.x con Keras) a ONNX se realiza comúnmente con la librería tf2onnx.

Paso 1: Definir y Entrenar un Modelo Simple en TensorFlow Keras

Definiremos un modelo Keras similar al que usamos en PyTorch.

import tensorflow as tf
from tensorflow import keras
from tensorflow.keras import layers
import numpy as np

# 1. Definir el modelo TensorFlow Keras
def create_tf_model():
    model = keras.Sequential([
        layers.Conv2D(10, kernel_size=(5, 5), activation='relu', input_shape=(28, 28, 1)),
        layers.MaxPooling2D(pool_size=(2, 2)),
        layers.Conv2D(20, kernel_size=(5, 5), activation='relu'),
        layers.Dropout(0.5),
        layers.MaxPooling2D(pool_size=(2, 2)),
        layers.Flatten(),
        layers.Dense(10, activation='softmax') # Para clasificación MNIST de 10 clases
    ])
    return model

model_tf = create_tf_model()

# 2. Compilar el modelo (opcional para exportación, pero buena práctica)
model_tf.compile(optimizer='adam', loss='sparse_categorical_crossentropy', metrics=['accuracy'])

print("Modelo TensorFlow Keras creado:")
model_tf.summary()

# 3. Crear una entrada de ejemplo
# TensorFlow espera (batch_size, alto, ancho, canales)
example_input_tf = tf.random.normal([1, 28, 28, 1], dtype=tf.float32)

print(f"Forma de la entrada de ejemplo de TF: {example_input_tf.shape}")

Paso 2: Exportar el Modelo a ONNX usando tf2onnx

tf2onnx proporciona una función convert.from_keras para exportar modelos Keras a ONNX.

import tf2onnx
import onnx

onnx_model_path_tf = "tensorflow_model.onnx"

try:
    # La función de conversión requiere la especificación de los tipos de entrada.
    # Puedes usar tf2onnx.convert.from_keras(model_tf, input_signature=...) si usas tf.function
    # O directamente con la entrada de ejemplo para Keras Models
    
    # Aseguramos que la entrada sea compatible con la firma del modelo Keras
    input_spec = (tf.TensorSpec(model_tf.inputs[0].shape, model_tf.inputs[0].dtype, name="input"),)

    # Ojo: tf2onnx puede ser sensible al ambiente y la versión. 
    # A veces es mejor especificar el input_signature directamente o usar el formato SavedModel.
    # Para modelos Keras Sequential, `input_signature` es más fiable.
    model_proto, external_tensor_storage = tf2onnx.convert.from_keras(
        model_tf, 
        input_signature=input_spec,
        opset=11, 
        output_path=onnx_model_path_tf
    )
    
    print(f"Modelo TensorFlow exportado con éxito a {onnx_model_path_tf}")

except Exception as e:
    print(f"Error al exportar el modelo TensorFlow a ONNX: {e}")
    <div class="callout tip">💡 <strong>Consejo:</strong> Si encuentras problemas, considera exportar primero tu modelo Keras a un `SavedModel` de TensorFlow y luego usar `tf2onnx.convert.from_saved_model`.
```python
# Si la exportación directa de Keras falla, prueba esto:
# tf_saved_model_path = "tf_saved_model_temp"
# model_tf.save(tf_saved_model_path)
# model_proto, _ = tf2onnx.convert.from_saved_model(
#    tf_saved_model_path, 
#    input_signature=[tf.TensorSpec([None, 28, 28, 1], tf.float32, name='input')],
#    opset=11, 
#    output_path=onnx_model_path_tf
# )
</div>

--- 

## ⬅️ Importando un Modelo ONNX en PyTorch (para Inferencia)

Similar a TensorFlow, podemos usar `onnxruntime` para la inferencia de modelos ONNX en un entorno PyTorch. Si el objetivo es una integración más profunda o un re-entrenamiento en PyTorch, se necesitarían herramientas de conversión más avanzadas (o recrear el modelo en PyTorch y cargar los pesos).

### Paso 1: Cargar y Ejecutar el Modelo ONNX con ONNX Runtime

Reutilizamos `onnxruntime` para ejecutar el modelo ONNX que se originó en TensorFlow.

```python
import onnxruntime as rt
import numpy as np
import torch

# 1. Cargar el modelo ONNX exportado desde TensorFlow
sess_tf_onnx = rt.InferenceSession(onnx_model_path_tf, providers=['CPUExecutionProvider'])

input_name_tf_onnx = sess_tf_onnx.get_inputs()[0].name
output_name_tf_onnx = sess_tf_onnx.get_outputs()[0].name

print(f"Nombre de entrada del modelo ONNX (de TF): {input_name_tf_onnx}")
print(f"Nombre de salida del modelo ONNX (de TF): {output_name_tf_onnx}")

# 2. Preparar una entrada de ejemplo para ONNX Runtime
# El modelo de TF espera (batch_size, alto, ancho, canales).
# Convierte el tensor de PyTorch a NumPy y ajusta la forma si es necesario para el modelo de TF.
# example_input es (1, 1, 28, 28). Necesitamos (1, 28, 28, 1) para el modelo TF/ONNX.
input_onnx_tf = example_input.permute(0, 2, 3, 1).numpy().astype(np.float32)

# 3. Realizar inferencia
onnx_output_tf = sess_tf_onnx.run([output_name_tf_onnx], {input_name_tf_onnx: input_onnx_tf})

print("Inferencia realizada con ONNX Runtime (modelo de TF)...")
print(f"Forma de la salida ONNX (de TF): {onnx_output_tf[0].shape}")
print(f"Primeros 5 valores de la salida ONNX (de TF): {onnx_output_tf[0][0, :5]}")

# Opcional: Ejecutar el modelo original de PyTorch para comparación
# Asegúrate de que el modelo_pytorch esté disponible y en modo evaluación
model_pytorch.eval()
with torch.no_grad():
    pytorch_original_output = model_pytorch(example_input)

print("\nComparación de resultados (TF-ONNX vs. PyTorch Original):")
# Convertir la salida de ONNX a tensor de PyTorch para argmax
output_from_tf_onnx_as_torch = torch.from_numpy(onnx_output_tf[0])
print(f"Predicción de TF-ONNX: {torch.argmax(output_from_tf_onnx_as_torch[0])}")
print(f"Predicción de PyTorch Original: {torch.argmax(pytorch_original_output[0])}")

💡 Consejo: Para verificar la consistencia, siempre compara las salidas del modelo original con las del modelo convertido/importado para un conjunto de entradas de ejemplo. Las pequeñas diferencias numéricas pueden ser normales debido a las implementaciones de operaciones en diferentes *frameworks*, pero las predicciones principales deben coincidir.

🔍 Verificación y Validación de Modelos ONNX

Después de exportar un modelo a ONNX, es fundamental verificar su integridad y que se comporte como se espera. Hay varias maneras de hacerlo:

  1. Validación con onnx.checker: La librería ONNX incluye un verificador para asegurar que el archivo ONNX es válido y sigue el estándar.
import onnx

try:
onnx_model = onnx.load(onnx_model_path) # Carga el modelo exportado desde PyTorch
onnx.checker.check_model(onnx_model)
print(f"Modelo {onnx_model_path} es un ONNX válido!")

onnx_model_tf_src = onnx.load(onnx_model_path_tf) # Carga el modelo exportado desde TensorFlow
onnx.checker.check_model(onnx_model_tf_src)
print(f"Modelo {onnx_model_path_tf} es un ONNX válido!")

except Exception as e:
print(f"Error al verificar el modelo ONNX: {e}")
  1. Visualización con Netron: Netron es una herramienta de visualización de redes neuronales que soporta ONNX. Es excelente para inspeccionar la estructura del grafo de tu modelo, los tipos de operadores y las dimensiones de los tensores.

    💡 Consejo: Puedes descargar Netron como aplicación de escritorio o usar la versión web en `https://netron.app/`. Abre tu archivo `.onnx` para visualizar el modelo.
  2. Comparación de Salidas: Como hicimos en las secciones anteriores, siempre se debe comparar la salida del modelo original con la del modelo ONNX (o el modelo ONNX convertido/importado) usando las mismas entradas. Esto es crucial para asegurar la equivalencia funcional.

    ⚠️ Advertencia: Pequeñas diferencias de precisión flotante son normales entre *frameworks* o *runtimes*. Lo importante es que las predicciones principales (por ejemplo, la clase más probable) sean consistentes. Si las diferencias son grandes, puede haber un problema en el proceso de conversión o en la `opset_version` utilizada.

📊 Casos de Uso y Consideraciones Avanzadas

La interoperabilidad con ONNX abre un abanico de posibilidades:

  • Entrenamiento y Despliegue Separados: Entrenar un modelo en PyTorch (conocido por su flexibilidad en investigación) y luego desplegarlo en producción usando TensorFlow Serving o un runtime ONNX optimizado.
  • Comparación de Rendimiento: Evaluar el rendimiento de inferencia de un mismo modelo en diferentes runtimes y hardwares (CPU, GPU, Edge TPUs, etc.) usando ONNX como formato intermedio.
  • Herramientas de Optimización: Aplicar herramientas de optimización específicas de ONNX (como cuantización o poda de grafos) que son agnósticas al framework de origen.
  • Integración con Otros Ecosistemas: Facilitar la integración con lenguajes como C++ o Java, donde los runtimes ONNX pueden ser más fáciles de integrar que los frameworks completos.

Consideraciones Avanzadas:

  • Operadores Personalizados: Si tu modelo utiliza operaciones personalizadas que no están en el estándar de ONNX, necesitarás implementar exportadores personalizados o convertirlas a combinaciones de operadores estándar.
  • Versiones de Opsets: Mantenerse al día con las versiones de ONNX opset es importante, ya que añaden nuevos operadores y mejoran los existentes. Sin embargo, usar una opset demasiado reciente puede reducir la compatibilidad con runtimes más antiguos.
  • Modelos Complejos: Para modelos muy complejos (ej. con lógica de control o bucles dinámicos), la exportación a ONNX puede requerir más cuidado y, a veces, una reestructuración del modelo para que sea compatible con el grafo estático de ONNX.
¿Qué pasa con los modelos de lenguaje grandes (LLMs)? La exportación de LLMs a ONNX es un área activa de desarrollo. Modelos como los de Hugging Face Transformers a menudo tienen soporte para exportación a ONNX, lo que permite su optimización y despliegue eficiente en diferentes *runtimes*. Herramientas como `optimum` de Hugging Face facilitan este proceso, a menudo manejando las complejidades de los operadores específicos de Transformers.

✅ Conclusión

Dominar la interoperabilidad de modelos con ONNX es una habilidad invaluable en el ecosistema de la IA. Permite a los desarrolladores y equipos moverse libremente entre TensorFlow y PyTorch, aprovechando las fortalezas de cada framework y optimizando el despliegue de modelos en una amplia variedad de plataformas. Al seguir los pasos de este tutorial, ahora tienes las herramientas y el conocimiento para exportar, importar y verificar tus modelos de IA, desbloqueando una mayor flexibilidad y eficiencia en tus proyectos.

¡Experimenta con tus propios modelos y explora las posibilidades que ofrece ONNX!

Tutoriales relacionados

Comentarios (0)

Aún no hay comentarios. ¡Sé el primero!