en la serie de aprendizaje federado que estoy haciendo, y si acaba de llegar aquí, le recomendaría leer la primera parte donde discutimos cómo funciona el aprendizaje federado a un alto nivel. Para un repaso rápido, aquí hay una aplicación interactiva que creé en un cuaderno marimo donde puedes realizar entrenamiento local, fusionar modelos usando el algoritmo Federated Averaging (FedAvg) y observar cómo el modelo global mejora en las rondas federadas.
En esta parte, nos centraremos en implementar la lógica federada utilizando el marco Flower.
¿Qué sucede cuando los modelos se entrenan en conjuntos de datos sesgados?
En la primera parte, analizamos cómo se utilizó el aprendizaje federado para la detección temprana de COVID con Curial AI. Si el modelo se hubiera entrenado únicamente con datos de un único hospital, habría aprendido patrones específicos de ese hospital únicamente y se habría generalizado mal en conjuntos de datos fuera de distribución. Sabemos que esto es una teoría, pero ahora pongamos un número.
Tomo prestado un ejemplo del curso de Flower Labs sobre DeepLearning.AI porque utiliza lo familiar, lo que hace que la idea sea más fácil de entender sin perderse en detalles. Este ejemplo facilita la comprensión de lo que sucede cuando los modelos se entrenan en conjuntos de datos locales sesgados. Luego utilizamos la misma configuración para mostrar cómo el aprendizaje federado cambia el resultado.
He realizado algunas pequeñas modificaciones al código original. En particular, utilizo la biblioteca Flower Datasets, que facilita el trabajo con conjuntos de datos para escenarios de aprendizaje federado. 💻 Puedes acceder al código aquí para seguirlo.
Dividiendo el conjunto de datos
Comenzamos tomando el conjunto de datos MNIST y dividiéndolo en tres partes para representar los datos en poder de diferentes clientes, digamos tres hospitales diferentes. Además, eliminamos ciertos dígitos de cada división para que todos los clientes tengan datos incompletos, como se muestra a continuación. Esto se hace para simular silos de datos del mundo real.
Como se muestra en la imagen de arriba, el cliente 1 nunca ve los dígitos 1, 3 y 7. De manera similar, el cliente 2 nunca ve los 2, 5 y 8 y el cliente 3 nunca ve los 4, 6 y 9. Aunque los tres conjuntos de datos provienen de la misma fuente, representan distribuciones bastante diferentes.
Capacitación sobre datos sesgados
A continuación, entrenamos modelos separados en cada conjunto de datos utilizando la misma arquitectura y configuración de entrenamiento. Usamos una red neuronal muy simple implementada en PyTorch con solo dos capas completamente conectadas y entrenamos el modelo durante 10 épocas.
Como puede verse en las curvas de pérdida anteriores, la pérdida disminuye gradualmente durante el entrenamiento. Esto indica que los modelos están aprendiendo algo. Sin embargo, recuerde, cada modelo solo aprende de su propia visión limitada de los datos y solo cuando lo probamos en un conjunto disponible sabremos la verdadera precisión.
Evaluación de datos invisibles
Para probar los modelos, cargamos el conjunto de datos de prueba MNIST con la misma normalización aplicada a los datos de entrenamiento. Cuando evaluamos estos modelos en el conjunto de prueba completo (los 10 dígitos), la precisión ronda entre el 65 y el 70 por ciento, lo que parece razonable dado que faltaban tres dígitos en cada conjunto de datos de entrenamiento. Al menos la precisión es mejor que la probabilidad aleatoria del 10%.
A continuación, también evaluamos el rendimiento de los modelos individuales en ejemplos de datos que no estaban representados en su conjunto de entrenamiento. Para ello, creamos tres subconjuntos de pruebas específicos:
El conjunto de prueba [1,3,7] solo incluye los dígitos 1, 3 y 7 El conjunto de prueba [2,5,8] solo incluye los dígitos 2, 5 y 8 El conjunto de prueba [4,6,9] solo incluye los dígitos 4, 6 y 9
Cuando evaluamos cada modelo sólo en los dígitos que nunca vio durante el entrenamiento, la precisión cae al 0 por ciento. Los modelos fracasan por completo en clases a las que nunca estuvieron expuestos. Bueno, esto también es de esperarse, ya que un modelo no puede aprender a reconocer patrones que nunca antes ha visto. Pero hay más de lo que parece, por lo que a continuación observamos la matriz de confusión para comprender el comportamiento con más detalle.
Comprender el fracaso a través de matrices de confusión
A continuación se muestra la matriz de confusión para el modelo 1 que se entrenó con datos excluyendo los dígitos 1, 3 y 7. Dado que estos dígitos nunca se vieron durante el entrenamiento, el modelo casi nunca predice esas etiquetas.
Sin embargo, en algunos casos, el modelo predice dígitos visualmente similares. Cuando falta la etiqueta 1, el modelo nunca genera 1 y en su lugar predice dígitos como 2 u 8. El mismo patrón aparece para otras clases faltantes. Esto significa que el modelo falla en cierto modo al asignar un alto nivel de confianza a la etiqueta incorrecta. Definitivamente esto no es lo esperado.
Este ejemplo muestra los límites de la capacitación centralizada con datos sesgados. Cuando cada cliente tiene sólo una visión parcial de la verdadera distribución, los modelos fallan de manera sistemática que la precisión general no capta. Este es exactamente el problema que el aprendizaje federado debe abordar y eso es lo que implementaremos en la siguiente sección utilizando el marco Flower.
¿Qué es Flor 🌼?
Flower es un marco de código abierto que hace que el aprendizaje federado sea muy fácil de implementar, incluso para principiantes. Es independiente del marco, por lo que no tiene que preocuparse por usar PyTorch, TensorFlow, Hugging Face, JAX y más. Además, se aplican las mismas abstracciones centrales ya sea que esté ejecutando experimentos en una sola máquina o entrenando en dispositivos reales en producción.
Los modelos florales federaron el aprendizaje de una manera muy directa. Una aplicación Flower se basa en los mismos roles que analizamos en el artículo anterior: clientes, un servidor y una estrategia que los conecta. Veamos ahora estos roles con más detalle.
Entendiendo la flor a través de la simulación
Flower hace que sea muy fácil comenzar con el aprendizaje federado sin preocuparse por ninguna configuración compleja. Para la simulación local, hay básicamente dos comandos que deben tener en cuenta:
uno para generar la aplicación: flwr new y otro para ejecutarla: flwr run
Usted define una aplicación Flower una vez y luego la ejecuta localmente para simular muchos clientes. Aunque todo se ejecuta en una sola máquina, Flower trata a cada cliente como un participante independiente con sus propios datos y ciclo de capacitación. Esto hace que sea mucho más fácil experimentar y probar antes de pasar a una implementación real.
Comencemos instalando la última versión de Flower, que en el momento de escribir este artículo es 1.25.0.
# Instalar flower en un entorno virtual pip install -U flwr # Comprobando la versión instalada flwr –version Versión de Flower: 1.25.0
La forma más rápida de crear una aplicación Flower que funcione es dejar que Flower cree una por usted a través de flwr new.
flwr new #para seleccionar de una lista de plantillas o flwr new @flwrlabs/quickstart-pytorch #especifique directamente una plantilla
Ahora tienes un proyecto completo con una estructura limpia para empezar.
inicio rápido-pytorch ├── pytorchexample │ ├── client_app.py │ ├── server_app.py │ └── task.py ├── pyproject.toml └── README.md
Hay tres archivos principales en el proyecto:
El archivo task.py define el modelo, el conjunto de datos y la lógica de entrenamiento. El archivo client_app.py define lo que hace cada cliente localmente. El archivo server_app.py coordina el entrenamiento y la agregación, normalmente utilizando un promedio federado, pero también puede modificarlo.
Ejecutando la simulación federada
Ahora podemos ejecutar la federación usando los siguientes comandos.
instalación de pip -e. ejecución de flwr.
Este único comando inicia el servidor, crea clientes simulados, asigna particiones de datos y ejecuta entrenamiento federado de un extremo a otro.
Un punto importante a tener en cuenta aquí es que el servidor y los clientes no se llaman entre sí directamente. Toda la comunicación se produce mediante objetos de mensaje. Cada mensaje lleva parámetros de modelo, métricas y valores de configuración. Los pesos del modelo se envían mediante registros de matriz, las métricas como la pérdida o la precisión se envían mediante registros de métricas y los valores como la tasa de aprendizaje se envían mediante registros de configuración. Durante cada ronda, el servidor envía el modelo global actual a clientes seleccionados, los clientes entrenan localmente y devuelven pesos actualizados con métricas y el servidor agrega los resultados. El servidor también puede ejecutar un paso de evaluación en el que los clientes solo informan métricas, sin actualizar el modelo.
Si miras dentro del pyproject.toml generado, también verás cómo se define la simulación.
[tool.flwr.app.components] serverapp = "pytorchexample.server_app:app" clientapp = "pytorchexample.client_app:app"
Esta sección le dice a Flower qué objetos Python implementan ServerApp y ClientApp. Estos son los puntos de entrada que utiliza Flower cuando lanza la federación.
[tool.flwr.app.config] num-server-rounds = 3 fracción-evaluación = 0,5 épocas-locales = 1 tasa de aprendizaje = 0,1 tamaño de lote = 32 [tool.flwr.federations] default = "simulación local" [tool.flwr.federations.local-simulation] options.num-supernodes = 10
A continuación, estos valores definen la configuración de ejecución. Controlan cuántas rondas de servidor se ejecutan, cuánto tiempo entrena localmente cada cliente y qué parámetros de entrenamiento se utilizan. Estas configuraciones están disponibles en tiempo de ejecución a través del objeto Flower Context.
[tool.flwr.federations] default = "simulación local" [tool.flwr.federations.local-simulation] options.num-supernodes = 10
Esta sección define la simulación local en sí. Configurar options.num-supernodes = 10 le dice a Flower que cree diez clientes simulados. Cada Supernodo ejecuta una instancia de ClientApp con su propia partición de datos.
Aquí hay un resumen rápido de los pasos mencionados anteriormente.
Ahora que hemos visto lo fácil que es ejecutar una simulación federada con Flower, aplicaremos esta estructura a nuestro ejemplo MNIST y revisaremos el problema de datos sesgados que observamos anteriormente.
Mejorar la precisión mediante la formación colaborativa
Ahora volvamos a nuestro ejemplo MNIST. Vimos que los modelos entrenados en conjuntos de datos locales individuales no dieron buenos resultados. En esta sección, cambiamos la configuración para que los clientes ahora colaboren compartiendo actualizaciones de modelos en lugar de trabajar de forma aislada. Sin embargo, a cada conjunto de datos todavía le faltan ciertos dígitos como antes y cada cliente todavía entrena localmente.
La mejor parte del proyecto obtenido mediante simulación en la sección anterior es que ahora se puede adaptar fácilmente a nuestro caso de uso. Tomé la aplicación de flores generada en la sección anterior e hice algunos cambios en client_app, server_app y el archivo de tareas. Configuré el entrenamiento para que se ejecutara durante tres rondas de servidor, con todos los clientes participando en cada ronda y cada cliente entrenando su modelo local durante diez épocas locales. Todas estas configuraciones se pueden administrar fácilmente a través del archivo pyproject.toml. Luego, los modelos locales se agregan a un único modelo global utilizando el promedio federado.
Ahora veamos los resultados. Recuerde que en el enfoque de entrenamiento aislado, los tres modelos individuales lograron una precisión de aproximadamente entre el 65 y el 70 %. Aquí, con el aprendizaje federado, vemos un salto masivo en la precisión hasta alrededor del 96%. Esto significa que el modelo global es mucho mejor que cualquiera de los modelos individuales entrenados de forma aislada.
Este modelo global incluso funciona mejor en los subconjuntos específicos (los dígitos que faltaban en los datos de cada cliente) y ve un salto en la precisión del 0% anterior a entre el 94 y el 97%.
La matriz de confusión anterior corrobora este hallazgo. Muestra que el modelo aprende a clasificar todos los dígitos correctamente, incluso aquellos a los que no estuvo expuesto. Ya no vemos ninguna columna que solo tenga ceros y cada clase de dígitos ahora tiene predicciones, lo que muestra que el entrenamiento colaborativo permitió que el modelo aprendiera la distribución completa de los datos sin que ningún cliente tuviera acceso a todos los tipos de dígitos.
Mirando el panorama general
Si bien este es un ejemplo de juguete, ayuda a comprender la intuición detrás de por qué el aprendizaje federado es tan poderoso. Este mismo principio se puede aplicar a situaciones en las que los datos se distribuyen en múltiples ubicaciones y no se pueden centralizar debido a restricciones regulatorias o de privacidad.
Por ejemplo, si sustituye el ejemplo anterior con, digamos, tres hospitales, cada uno con datos locales, verá que aunque cada hospital solo tiene su propio conjunto de datos limitado, el modelo general entrenado a través del aprendizaje federado sería mucho mejor que cualquier modelo individual entrenado de forma aislada. Además, los datos permanecen privados y seguros en cada hospital, pero el modelo se beneficia del conocimiento colectivo de todas las instituciones participantes.
Conclusión y qué sigue
Eso es todo por esta parte de la serie. En este artículo, implementamos un ciclo de aprendizaje federado de extremo a extremo con Flower, comprendimos los diversos componentes de la aplicación Flower y comparamos el aprendizaje automático con y sin aprendizaje colaborativo. En la siguiente parte, exploraremos el aprendizaje federado desde el punto de vista de la privacidad. Si bien el aprendizaje federado en sí mismo es una solución de minimización de datos, ya que evita el acceso directo a los datos, las actualizaciones del modelo intercambiadas entre el cliente y el servidor aún pueden provocar fugas de privacidad. Toquemos esto en la siguiente parte. Por ahora, será una buena idea consultar la documentación oficial.