Creando un radiólogo IA multimodal en ruso: Guía ViT + ruGPT-3 en Kaggle
Desarrollar soluciones de IA médica en ruso enfrenta la escasez de herramientas y datasets listos para usar. Este artículo presenta un caso práctico paso a paso para ensamblar y entrenar un modelo VisionEncoderDecoder que genera informes médicos a partir de radiografías, combinando el codificador visual ViT con el modelo de lenguaje ruGPT-3.
Elecciones arquitectónicas y modificaciones del modelo
Construir un modelo multimodal desde cero requiere una enorme potencia de cómputo. Una opción más inteligente es combinar componentes preentrenados usando la arquitectura VisionEncoderDecoderModel de Hugging Face. Elegimos google/vit-base-patch16-224-in21k como codificador para extraer características visuales. El decodificador es el modelo en ruso ai-forever/rugpt3small_based_on_gpt2.
El principal desafío: ruGPT-3 no soporta nativamente la atención cruzada para recibir datos del codificador. La solución implica ajustar la configuración del modelo antes de la inicialización:
from transformers import AutoConfig, AutoModelForCausalLM, AutoModel
encoder = AutoModel.from_pretrained("google/vit-base-patch16-224-in21k")
decoder_config = AutoConfig.from_pretrained("ai-forever/rugpt3small_based_on_gpt2")
decoder_config.is_decoder = True
decoder_config.add_cross_attention = True
decoder = AutoModelForCausalLM.from_pretrained("ai-forever/rugpt3small_based_on_gpt2", config=decoder_config)
Una vez ensamblado, el modelo inicializa nuevos pesos para las capas de atención cruzada y proyección, que se ajustan finamente durante el entrenamiento.
Preparación y procesamiento de datos
No existen datasets rusos listos que emparejen imágenes de rayos X con informes médicos, así que nos ingeniamos. Usamos el dataset en inglés Indiana University Chest X-Ray (IU X-Ray) con ~7.500 imágenes e informes. Los informes se tradujeron al ruso usando el modelo Helsinki-NLP/opus-mt-en-ru directamente en el entorno de Kaggle.
El manejo de archivos en Kaggle tiene sus peculiaridades. Los cargadores de imágenes estándar suelen ver solo vistas previas en miniatura. Para un mapeo correcto, escaneamos profundamente el sistema de archivos:
import os
file_map = {}
for root, dirs, files in os.walk('/kaggle/input'):
for file in files:
if file.lower().endswith(('.png', '.jpg', '.jpeg')):
base_name = os.path.splitext(file)[0]
file_map[base_name] = os.path.join(root, file)
Esto garantiza un enlace fiable entre los identificadores CSV y las rutas reales de las imágenes.
Proceso de entrenamiento y resolución de problemas
El entrenamiento se ejecutó en dos GPUs NVIDIA T4 con precisión mixta (fp16) y acumulación de gradientes (gradient_accumulation_steps=4) para lograr un tamaño de lote virtual de 32 de forma eficiente.
Problemas comunes y soluciones:
- Fallo al guardar checkpoints: Seq2SeqTrainer en algunas versiones de transformers falla con modelos compuestos personalizados. Solución: Desactiva los guardados intermedios (save_strategy="no") y maneja manualmente los parámetros de generación vía GenerationConfig antes del guardado final.
- Timeouts de sesión en Kaggle: Evita el apagado automático durante entrenamientos largos ejecutando un snippet simple de JavaScript en la consola del navegador para simular clics periódicos.
Resultados, inferencia y despliegue
Tras 15 épocas, el modelo genera informes médicos coherentes en ruso, usando terminología profesional correctamente. Identifica patrones básicos como pulmones limpios, signos de neumotórax y cardiomegalia. El dataset pequeño provoca alucinaciones ocasionales, como mencionar dispositivos médicos ausentes.
Ejemplo básico de inferencia:
import torch
from transformers import VisionEncoderDecoderModel, ViTImageProcessor, AutoTokenizer
from PIL import Image
model_id = "livadies/Russian-Radiologist-ruGPT-ViT"
model = VisionEncoderDecoderModel.from_pretrained(model_id)
feature_extractor = ViTImageProcessor.from_pretrained(model_id)
tokenizer = AutoTokenizer.from_pretrained(model_id)
image = Image.open("xray.jpg").convert("RGB")
pixel_values = feature_extractor(images=image, return_tensors="pt").pixel_values
generated_ids = model.generate(pixel_values, max_length=128, num_beams=4)
print(tokenizer.decode(generated_ids[0], skip_special_tokens=True))
Una demo con Gradio corre en Hugging Face Spaces usando CPU, con tiempos de generación de 10-15 segundos. El código fuente completo —incluyendo pipelines de datos y notebooks de entrenamiento— está en el repositorio público.
Lecciones clave
- Arquitectura híbrida: Fusión exitosa del codificador visual ViT preentrenado y el modelo de lenguaje ruGPT-3 añadiendo atención cruzada vía ajustes de config.
- Soluciones para datos: Superamos la falta de datasets médicos rusos traduciendo los ingleses y navegando las peculiaridades de archivos en Kaggle.
- Arreglos prácticos: Desactivamos guardados automáticos defectuosos en Seq2SeqTrainer y evitamos timeouts de sesión para entrenamientos largos fiables.
- Prueba de concepto: Modelo con dataset pequeño valida la arquitectura Visión + Lenguaje para generación de texto médico ruso estructurado.
- Abierto y accesible: Código completo y demo interactiva públicos para estudio, pruebas y extensiones de la comunidad.
— Editorial Team
Aún no hay comentarios.