Helion es la DSL de alto nivel de PyTorch para escribir kernels de aprendizaje automático con rendimiento portable. En alianza con Google, el equipo construyó un backend de TPU que compila los kernels de Helion a Pallas, lo que ofrece una vía amigable con PyTorch para escribir kernels de TPU de alto rendimiento. En una carga de trabajo de atención flash, el kernel generado por Helion alcanza 838 TFLOPs, cerca del 79% de utilización de un tensor core en una TPU v7. Para distintas formas de entrada, Helion prueba automáticamente diferentes estrategias de generación de código para elegir el esquema de canalización óptimo.

¿Por qué importa Helion en las TPU?

Las TPU son cada vez más relevantes como plataforma de cómputo para aprendizaje automático, como complemento de las GPU. La última TPU v7 de Google (Ironwood) entrega un rendimiento comparable a la NVIDIA B200 con un costo total de propiedad potencialmente menor, lo que las vuelve atractivas para entrenamiento e inferencia a gran escala.

Sin embargo, escribir kernels de TPU tradicionalmente exige experiencia en Pallas, una DSL de bajo nivel con una curva de aprendizaje pronunciada. Helion busca cerrar esa brecha: permite escribir código con el estilo familiar de PyTorch y lo compila a código de TPU optimizado. Según PyTorch, Helion apunta a tres casos de uso:

  • Casos críticos de rendimiento, donde el autoajuste es necesario para explorar el espacio de configuraciones.
  • Quienes no son expertos en Pallas y quieren empezar a escribir kernels de TPU rápidamente.
  • Usuarios multiplataforma que prefieren mantener el mismo conjunto de kernels entre TPU y GPU.

¿En qué se diferencia una TPU de una GPU?

Las TPU son aceleradores altamente especializados, diseñados y optimizados específicamente para cargas de aprendizaje automático. Su arquitectura y modelo de programación difieren de forma notable respecto de las GPU. La diferencia más marcada es que una TPU es una máquina secuencial con registros y unidades de cómputo vectoriales anchos. Esto contrasta con las GPU, que logran rendimiento tanto mediante ejecución masivamente paralela (núcleos CUDA) como con unidades tensoriales especializadas (tensor cores).

Como resultado, las TPU tienen una jerarquía de memoria que quien escribe kernels debe entender a fondo, para orquestar cuándo y cómo se cargan los datos desde la HBM externa hacia la rápida VMEM interna. Un kernel de Pallas eficiente superpone estas transferencias de memoria HBM a VMEM con el cómputo de punto flotante que ocurre en las unidades matriciales (MXU) y vectoriales.

Diagrama de la jerarquia de memoria de una TPU
Diagrama de la jerarquia de memoria de una TPU

Pese a las diferencias arquitectónicas, las TPU y GPU de generación actual son muy comparables en rendimiento bruto. La TPU7x y la NVIDIA B200 tienen valores de cómputo BF16 (TFLOPS) y ancho de banda de HBM muy similares, las dos métricas de hardware más importantes para las cargas modernas de aprendizaje automático.

Cómo genera Helion el código de Pallas

Para exprimir el máximo rendimiento de una TPU, la generación de código de Helion apunta a maximizar la canalización de software (software pipelining), asegurando que las transferencias de memoria y el cómputo se superpongan lo más posible. Como ejemplo simple, un kernel de Helion para sumar dos tensores se escribe así:

Código
@helion.kernel
def add(x: torch.Tensor, y: torch.Tensor) -> torch.Tensor:
    out = torch.empty_like(x)
    for tile in hl.tile(out.size()):
        out[tile] = x[tile] + y[tile]
    return out

El compilador de Helion traduce esto en dos funciones: un lanzador del lado del host que divide la entrada en bloques (tiles) e invoca la función del dispositivo de forma canalizada, y una función de dispositivo que opera sobre bloques residentes en VMEM. El tamaño de bloque lo selecciona el autoajustador (autotuner), que explora distintos tamaños para hallar la mejor superposición entre transferencias de memoria y cómputo en el hardware de destino.

El caso de la atención flash

La atención es una de las operaciones clave en los modelos de lenguaje modernos. Las implementaciones de producción siguen el patrón "Flash Attention", una técnica eficiente en memoria que calcula la atención por bloques para evitar materializar la matriz completa. En Helion, el compilador autoajusta entre dos estrategias para traducir el kernel a Pallas.

Con la opción por defecto (emit_pipeline), Helion se apoya en la API de Pallas para canalizar un bucle interno, en una estructura de canalización anidada. La alternativa (unroll) precarga por completo las secuencias de K y V en VMEM y elimina las "burbujas" de las unidades de cómputo, a costa de un mayor uso de VMEM, lo que la hace inviable con secuencias muy largas.

Esquema de canalizacion sin burbujas de computo en Helion
Esquema de canalizacion sin burbujas de computo en Helion

El beneficio de Helion está en su capacidad de autoajustar y elegir la mejor configuración: con secuencias pequeñas aprovecha la VMEM disponible y genera código sin burbujas de cómputo; con secuencias largas recurre a emit_pipeline, que escala a longitudes de contexto arbitrarias. Esa capacidad de generar distintas estrategias de bucle y canalización según la longitud de entrada es lo que le da ventaja a Helion, incluso frente a implementaciones altamente optimizadas.