Cómo entrené un codificador cruzado Stance-Aware que clasifica los titulares de noticias de Indonesia según afirmaciones :  comenzando con una Colab TPU gratuita y escalando a Cloud TPU v5p con una sola@kinetic.run() decorador

Introducción

La desinformación es uno de los problemas definitorios de la era de las redes sociales, e Indonesia se ha visto especialmente afectada. Broma (la abreviatura indonesia de noticias falsas) difundirse a través de grupos de WhatsApp e hilos de Twitter más rápido de lo que cualquier verificador de datos puede seguir el ritmo. La mayoría de las investigaciones publicadas sobre la detección automatizada de noticias falsas se centran en datos en inglés, lo que deja a los practicantes que trabajan con bahasa indonesio en una situación frustrante: las técnicas existen, pero las herramientas y los modelos previamente entrenados son escasos.

Este artículo explica cómo construir un verdadero, laboral detector de engaños multimodal para noticias de Indonesia desde cero. El modelo toma dos entradas — a afirmar (la afirmación original, a menudo de las redes sociales) y un titular (el titular de un artículo de noticias que menciona el mismo tema) — y predice si el artículo apoya el reclamo (para), refuta él (contra), o simplemente observa es neutral (observando).

La arquitectura es un Codificador cruzado consciente de la postura: un codificador estilo BiLSTM para cada entrada, autoatención de múltiples cabezas, y una capa de atención cruzada que permite que el reclamo y el título se lean literalmente entre sí antes de la clasificación.. Construido de extremo a extremo con JAX y Lino, entrenado en TPU.

La historia del despliegue tiene dos mitades:

  1. TPU de colaboración gratuita para creación de prototipos : Google ofrece a todos los usuarios de Colab acceso gratuito a una TPU v5e-1, lo cual es suficiente para entrenar este modelo de un extremo a otro en menos de una hora y sin costo alguno.
  2. Nube TPU v5p a través de Keras Kinetic para un entrenamiento serio — cuando superes los límites de tiempo de ejecución de Colab, Cinético Duro le permite enviar la misma función de entrenamiento a un pod de Cloud TPU con un único decorador de Python. Sin ventana acoplable, sin Kubernetes YAML, sin SSH.

Al final, tendrás:

  • Un tokenizador y cargador de conjuntos de datos indonesio reutilizable
  • Un codificador Transformer de 4 capas con atención cruzada consciente de la postura, escrito en lino
  • Un circuito de entrenamiento compilado por JIT con optax y orbaxcheckpointing
  • Un predict.py funcional que ejecuta nuevos pares de reclamos y titulares a través del modelo entrenado
  • exactamente el mismo codigo, implementado en Cloud TPU a través [email protected]()

entremos en ello.

Por qué este problema es difícil (e interesante)

Los ingenuos detectores de noticias falsas miran un fragmento de texto e intentan clasificarlo como “real” o “falso”. Esto es técnicamente débil y éticamente incómodo: un solo texto rara vez transmite suficiente señal., y el encuadre “verdadero/falso” supone que el modelo tiene acceso a la verdad fundamental que posiblemente no pueda tener..

El detección de postura el encuadre es mucho más honesto. Dado un reclamo y un artículo de noticias relacionado., el modelo no decide si el reclamo es verdadero; decide si este artículo en particular apoya, refuta, o simplemente observa el reclamo. Esa es una pregunta que un modelo realmente puede responder., y es exactamente la información que un verificador de hechos posterior necesita para tomar una decisión final.

Matemáticamente, la tarea es una clasificación de 3 vías sobre elinteracción de dos fragmentos de texto. Esa palabra — interacción — es lo que hace que la arquitectura sea interesante.. No puedes simplemente codificar cada lado de forma independiente y concatenar. Necesita una capa que permita que el reclamo atienda al titular y viceversa., para que el modelo pueda captar señales sutiles como la negación ("El gobierno niega..."), cobertura ("presunto…"), o enmarcar (“según los críticos…”).

¿Por qué JAX?, Lino, Cinético Duro, y TPU?

  • JAX me da código estilo NumPy con diferenciación automática, Compilación JIT a través de XLA, y aceleración transparente en CPU/GPU/TPU.
  • Lino se sienta encima de JAX y me permite escribir redes neuronales como clases nn.Module. El modelo es denso en capas de atención., y Flax mantiene limpia la gestión de parámetros.
  • optax para optimización (AdamW con calentamiento lineal + decaimiento del coseno) y orbax para puntos de control :  ambos son parte del ecosistema JAX y son compatibles con JIT.
  • Cinético Duro es el pegamento de despliegue. Un decorador convierte una función local de Python en un trabajo de TPU remoto, con almacenamiento en caché de contenedores, transmisión de registros, y aprovisionamiento automático de GKE.
  • TPU porque la carga de trabajo está dominada por los matmuls de atención — exactamente para qué están diseñados los arreglos sistólicos de TPU. Gratis en Colab (v5e-1), y Cloud TPU v5p cuando necesites escalar.

TPU vs GPU para esta carga de trabajo

Un Transformer multimodal con atención cruzada es una de las cargas de trabajo de TPU más limpias que puede escribir. He aquí por qué, y donde las GPU todavía se mantienen firmes.

Diseño de hardware

  • GPU (Nvidia A100/H100): procesador paralelo de propósito general, miles de núcleos CUDA, ideal para cálculo paralelo arbitrario.
  • TPU (v5e o v5p): acelerador de dominio específico construido alrededor de una gran matriz sistólica (MXU) optimizado para multiplicaciones de matrices densas.

¿Qué domina el cálculo en este modelo?

  • Autoatención multicabezal: softmax(Q Kᵀ / √d) V — tres matmuls grandes por cabeza y por capa.
  • Atención cruzada entre reclamo y titular: misma forma, solo con diferentes entradas alimentando Q vs K/V.
  • Bloques de avance: dos capas densas con GELU entre ellas.

Todos esos son matmuls densos con formas predecibles.. La matriz sistólica de TPU está diseñada específicamente para analizar exactamente este patrón en los FLOP máximos.. El compilador XLA fusiona todo el train_step en unos pocos núcleos, y después de la primera compilación, cada paso se ejecuta a pleno rendimiento.

Donde las GPU siguen ganando en detección de postura / PNL

  • Estás realizando una decodificación a nivel de token con caché KV y longitudes de generación irregulares. (no estamos — estamos haciendo clasificación).
  • Necesita un modelo de transformadores HuggingFace que solo esté disponible como punto de control de PyTorch (estamos entrenando desde cero, entonces esto no aplica).
  • Quiere iterar en un cuaderno con un flujo de control constante de Python que no hace JIT limpiamente (Colab les da a ambos un TPU yun cuaderno, para que no tengas que elegir).

Donde ganan los TPU en la detección de posturas / PNL

  • Secuencias de longitud fija (vamos a 64 fichas) → formas predecibles → gran compilación XLA.
  • Todos los JIT de train_step en un único gráfico de ejecución fusionado.
  • mapap / shard_map hace que el entrenamiento con múltiples chips sea de una sola línea si desea ampliarlo.
  • Gratis en Colab, y Cloud TPU v5e cuesta aproximadamente $0,40/chip-hora en Spot.

regla general

  • Creación rápida de prototipos en un cuaderno con datos en bahasa indonesio → Colab TPU gratis. (Este artículo.)
  • R iterativo&D usando los puntos de control de HuggingFace PyTorch → GPU.
  • Capacitación en producción con lotes, Cargas de trabajo nativas de JAX → Nube TPU a través de Kinetic.
  • Necesita perfeccionar un LLM de Indonesia 7B+ → ese es un artículo diferente (y una categoría diferente — vLLM o Tunix).

Ahora vamos a construirlo.

Arquitectura del proyecto

El oleoducto es sencillo:

conjunto de datos.csv (Afirmar, Título, Postura)


┌──────────────────┐
│Tokenizer indonesio │ espacios en blanco + puntuación, vocabulario del corpus
└──────────────────┘


┌──────────────────┐
│ FakeNewsDataset │ división estratificada de tren/val/prueba, Matrices listas para JAX
└──────────────────┘


┌─────────────────── ───────────────────┐
│ Detector de noticias falsas (Lino nn.Módulo) │
│ ┌──────────────┐ ┌──────────────┐ │
│ │ Ficha+Pos │ │ Ficha+Pos │ │
│ │ Incrustación │ │ Incrustación │ │
│ │ (Afirmar) │ │ (Titular) │ │
│ └──────┬───────┘ └──────┬───────┘ │
│ ▼ ▼ │
│ ┌──────────────┐ ┌──────────────┐ │
│ │ Transformador │ │ Transformador │ │
│ │ × N capas │ │ × N capas │ │
│ └──────┬───────┘ └──────┬───────┘ │
│ └────────┬────────┘ │
│ ▼ │
│ ┌─────────────────────┐ │
│ │ Codificador cruzado de postura│ │
│ │ (atención cruzada + │ │
│ │ diferencia & producto) │ │
│ └──────────┬──────────┘ │
│ ▼ │
│ Denso → Softmax (3 clases) │
└─────────────────── ───────────────────┘


para / contra / observando

Todo se encuentra dentro de un train_step compilado por jax.jit.. Ahora repasemos cada pieza..

Paso 1 — Configuración del hardware

Hay dos caminos. Elija el que se ajuste a su etapa del proyecto..

Camino A: TPU de colaboración gratuita (recomendado para la primera ejecución)

  1. Abierto colab.research.google.com y crear un nuevo cuaderno.
  2. Hacer clic Tiempo de ejecución → Cambiar tipo de tiempo de ejecución.
  3. Bajo acelerador de hardware, seleccionar v5e-1 TPU.
  4. Hacer clic Ahorrar.
  5. Verificar en una celda:
importar jax
imprimir(jax.dispositivos())
# Esperado: [TpuDispositivo(identificación=0, ...)]

Eso es todo. Ahora tienes un chip TPU v5e gratuito durante tu sesión de Colab.

Camino B: Nube TPU a través de Keras Kinetic (cuando se te queda pequeño Colab)

Colab es fantástico para la creación de prototipos, pero tiene límites de tiempo de ejecución y se ve afectado por la carga.. Cuando esté listo para ejecutar trabajos de capacitación de varias horas, cambiar a una TPU en la nube real. La ruta tradicional significa aprovisionar una VM de TPU, SSH en, instalando dependencias, y cargar scripts — Kinetic se salta todo eso.

En tu computadora portátil local:

instalación de pip keras-cinética
inicio de sesión predeterminado de la aplicación gcloud auth
proyecto de conjunto de configuración de gcloud YOUR_PROJECT_ID
cinético arriba --acelerador v5p-8 --sí

El último comando aprovisiona un clúster de GKE Autopilot con un grupo de nodos TPU v5p-8. Tarda unos minutos la primera vez., Después de lo cual no se vuelve a tocar la infraestructura hasta el desmantelamiento..

Mostraré el @kinetic.run real() implementación en paso 6. Por ahora, construyamos el modelo.

Paso 2 — Tokenizer y conjunto de datos de Indonesia

El bahasa indonesio es morfológicamente menos complejo que el, decir, turco o finlandés, entonces un espacio en blanco + El tokenizador de puntuación con un vocabulario aprendido funciona sorprendentemente bien como punto de referencia.. (Para la producción, intercambiar en IndoBERT — Mostraré cómo al final de esta sección.)

El tokenizador reserva cuatro tokens especiales., construye un vocabulario clasificado por frecuencia a partir del corpus de entrenamiento, y emite(identificadores_token, máscara_de_atención) pares de longitud fija. Cosas estándar, pero con una sutileza: tokenizamos ambos el reclamo y judul (titular) columnas en un compartido vocabulario para que la capa de incrustación pueda captar correlaciones de entrada cruzada.

"""
Preprocesamiento de datos para la detección de noticias falsas en Indonesia
Reclamación de tokenizaciones + Título de la columna, codifica etiquetas de postura.
Compatible con el canal de formación JAX/Flax.
"""
importar re
importar numpy como np
importar pandas como pd
de colecciones importar contador
de escribir lista de importación, tupla, dictar
desde sklearn.model_selection importar train_test_split
ETIQUETA_MAP = {"for": 0, "against": 1, "observing": 2}
ID_TO_LABEL = {v: k por k, v en LABEL_MAP.items()}
clase Tokenizer indonesio:
"""
Espacios en blanco ligeros + tokenizador de puntuación para texto en indonesio.
    Para la producción, intercambiar con:
desde transformadores importa AutoTokenizer
tok = AutoTokenizer.from_pretrained("indobenchmark/indobert-base-p1")
"""
SPECIAL_TOKENS = {"<ALMOHADILLA>": 0, "<Desconocido>": 1, "<CLS>": 2, "<SEP>": 3}
    definición __init__(ser, tamaño_vocab: entero = 30_000, frecuencia_min: entero = 2):
self.vocab_size = tamaño_vocab
self.min_freq = min_freq
yo.word2id: dictar[cadena, entero] = dictar(self.SPECIAL_TOKENS)
yo.id2word: dictar[entero, cadena] = {v: k por k, v en self.word2id.items()}
    @métodoestático
def _limpio(texto: cadena) -> cadena:
texto = texto.inferior()
texto = re.sub(r"<[^>]+>", " ", texto) # tira HTML
texto = re.sub(r"[^\w\s]", " ", texto, banderas=re.UNICODE) # mantener alfanum
texto = re.sub(r"\s+", " ", texto).banda()
devolver texto
    @métodoestático
def tokenizar(texto: cadena) -> Lista[cadena]:
devolver IndonesianTokenizer._clean(texto).dividir()
 def construir_vocab(ser, textos: Lista[cadena]) -> Ninguno:
encimera: Contador = Contador()
para t en textos:
contador.actualización(auto.tokenizar