Documentación

compile=True y la ruta rápida de TileLang: notas de ingeniería

Estas notas cubren cómo se comportan compile=True y fast=True más allá de lo que dice el README. Provienen de mediciones tomadas mientras se trabajaba en #472, #576 y #718, en una RTX 4070 Ti SUPER con torch 2.11 y tilelang 0.1.14. Están aquí para que la siguiente persona no tenga que medirlas de nuevo.

Selección de backend

Agent(..., backend="auto") y laya.load(..., backend="auto") optan por la capa de clases de backend. El valor predeterminado sigue siendo eager. Un backend= explícito tiene prioridad sobre compile y fast; omitirlo conserva el comportamiento existente de ambos flags.

  • eager: el forward estándar de PyTorch, en cualquier dispositivo compatible.
  • compile: torch.compile solo para CUDA con formas dinámicas y modo reduce-overhead, relleno por buckets, caché de inductor persistente y calentamiento en la instalación. Reutiliza el mismo alcance de dimensión independiente que compile=True, que mantiene su modo predeterminado existente y soporte de CPU. Define LAYA_COMPILE_WARMUP=0 para diferir el calentamiento del backend y LAYA_INDUCTOR_CACHE_DIR para elegir su directorio de caché (predeterminado ~/.cache/laya/inductor).
  • tilelang: un adaptador alrededor de la ruta rápida actual, que usa el dtype bf16 o fp16 del agente.
  • auto: TileLang en CUDA con un codificador ModernBERT y dtype compatibles cuando TileLang está instalado; si no, compile en CUDA; eager en otros dispositivos.
  • onnx: laya.load(..., backend="onnx", onnx_path="model.onnx") devuelve el ONNXAgent existente. Sin onnx_path, usa laya.onnx.

Un backend no disponible emite un RuntimeWarning que nombra el backend resuelto y recurre a eager. Para exigir un backend, usa agent.set_backend("tilelang", strict=True). El cambio espera a la inferencia activa; agent.backend informa el nombre activo y agent.backend_object expone el objeto instalado. agent.set_backend("compile", warmup=False) difiere la compilación hasta la inferencia, así que los errores de compilación entonces afloran en la solicitud. agent.warmup() sigue disponible. agent.deaccelerate() elimina un backend instalado a través de la capa de clases.

Los Router reenvían una selección explícita a través de Router(agent_kwargs={"backend": "auto"}). No pasan ningún argumento de backend por defecto, preservando la compatibilidad con los constructores tipo Agent existentes. Los reintentos de CPU OOM con alcance limitado separan el backend y lo restauran cuando el modelo vuelve a su dispositivo original.

compile=True materializa la máscara de atención

La SDPA eager toma la máscara de atención (rows, 1, L, L) de ModernBERT como una vista de difusión. Bajo las formas dinámicas que usa compile=True, inductor no puede demostrar que la última dimensión esté alineada. Expande la máscara a cada cabeza y la rellena en un búfer real de rows x heads x L x L. En bf16 con 12 cabezas, eso es 0.8 GB con 32 filas x 1024 tokens.

  • GPU con margen. El búfer cuesta ancho de banda, decenas de ms por lote largo.
  • GPU casi llena. El asignador de caché se satura, y la misma llamada puede tardar decenas de segundos.

Si compilas con lotes largos en una GPU ocupada, limita el tamaño de lote (predict_batch(..., batch_size=)) or use fast=True. La atención de TileLang lee el búfer QKV empaquetado y enmascara por longitud de secuencia, así que no tiene ese búfer.

Arranque en frío

  • Primera compilación. Tarda decenas de segundos por grafo. compile=True necesita dos grafos: uno para lotes y otro para una sola fila, que torch especializa. compile=True ahora llama a agent.warmup() durante la carga. compile_warmup=False restaura la compilación perezosa, y agent.warmup(shapes=...) sigue disponible manualmente. Las cargas eager y TileLang no calientan automáticamente. Estas formas cubren solicitudes comunes, no todas las posibles guardas de forma.
  • Fallo de calentamiento. El calentamiento automático es de mejor esfuerzo: un fallo emite un RuntimeWarning que nombra el error (incluido el error del compilador subyacente) y la carga vuelve con el envoltorio torch.compile y los ajustes de compilación intactos. Por ejemplo, Windows sin MSVC puede cargar con compile=True aunque falle el calentamiento. Las solicitudes posteriores siguen usando el modelo compilado y dejan aflorar los fallos de compilación; Laya no las cambia a ejecución eager. Las llamadas explícitas a agent.warmup() también propagan fallos, incluso tras un calentamiento automático fallido. Por tanto, una carga exitosa no garantiza que la inferencia compilada esté lista.
  • Adopción de la caché de Laya. laya.load(..., compile=True, compile_cache=True) define la TORCHINDUCTOR_CACHE_DIR de todo el proceso solo cuando falta, a $XDG_CACHE_HOME/laya/torchinductor o ~/.cache/laya/torchinductor cuando XDG no está definido o no es absoluto. Una configuración existente, incluida una establecida por una compilación anterior de PyTorch, gana. El directorio se crea al cargar; los errores del sistema de archivos se propagan. compile_cache=False (predeterminado), eager y las cargas TileLang dejan el entorno en paz. Esto no mueve ni elimina cachés antiguas. Los contenedores siguen necesitando un home/volumen persistente. La compatibilidad e invalidación de caché las gestiona PyTorch; un cambio de GPU, torch, compilador, modelo o guarda de entrada puede exigir compilar de nuevo.
  • Entre reinicios. La caché de grafos FX de Inductor conserva los grafos compilados bajo TORCHINDUCTOR_CACHE_DIR. El valor por defecto está bajo /tmp, que no sobrevive a un reinicio ni a un reinicio del contenedor. Defínelo en un directorio persistente, o en un volumen en un contenedor, y un segundo proceso carga los grafos en lugar de compilarlos. En la medición de #472, eso llevó el calentamiento de unos 120 s a unos 50 s.

Grafos CUDA opcionales

agent = laya.load("convaiinnovations/laya", compile=True,
                  compile_cache=True, compile_mode="reduce-overhead")

compile_mode es "default" por defecto; en la ruta compilada activa solo se aceptan "default" y "reduce-overhead". Las cargas eager y TileLang ignoran las opciones de compilación. La compilación en CPU sigue funcionando, pero el registro de grafos CUDA solo se aplica en CUDA. El modo CUDA requiere la API torch.compiler.cudagraph_mark_step_begin de PyTorch; las compilaciones más antiguas sin ella lanzan un error explícito.

Los grafos dinámicos de Dynamo no implican grafos CUDA independientes de la forma: las nuevas formas concretas pueden requerir de nuevo calentamiento y registro, sin un nuevo grafo de Dynamo. Las dos formas sintéticas de calentamiento predeterminadas no pre-registran todas las formas de solicitud. Las formas repetidas pueden beneficiarse, pero las formas variables pueden pagar latencia extra y retener pools de grafos. PyTorch puede omitir los grafos CUDA para operaciones o configuraciones no compatibles; establecer este modo no garantiza la captura.

Laya marca cada forward CUDA compilado como un nuevo paso, serializa estos forwards entre sus agentes y clona ambos tensores de salida fuera del grafo compilado antes de liberar el bloqueo. Esto mantiene válidas las salidas retenidas a través de las repeticiones, al costo de dos copias y ejecución de forward serializada. El bloqueo no coordina modelos compilados ajenos propiedad de la aplicación; los llamadores que comparten iteraciones de grafos CUDA o usan streams personalizados deben gestionar su propia coordinación. Las cachés en disco reutilizan código compilado, no grabaciones vivas de grafos CUDA ni su memoria de dispositivo, entre procesos.

Reproduce los tiempos de frío/reinicio, la memoria y los contadores de caché con benchmarks/bench_compile_defaults.py; consulta las mediciones registradas.

AOTInductor: todavía no

Distribuir un artefacto precompilado por checkpoint y arquitectura de GPU (torch._inductor.aoti_compile_and_package) eliminaría la compilación por completo. En torch 2.11 se detiene en el empaquetado:

  • La exportación funciona. torch.export de DecisionModel tiene éxito, en unos 5 s, con filas, marcadores y tokens dinámicos. Los tokens deben declararse como múltiplo de 16 (16 * Dim(...)); un rango simple hace fallar la propia guarda de alineación L % 8 del exportador. Es la misma alineación de máscara que la de arriba.
  • El empaquetado falla. Cómo falla depende de cómo se exportó el programa:
    • Bajo autocast, el programa lleva aserciones de dtype con las que AOTI tropieza fuera de autocast: Tensor dtype mismatch! Expected: torch.bfloat16, Got: torch.float32.
    • Desde una copia bf16 sin autocast, el trazado falla dentro de la pasada hacia adelante: mat1 and mat2 must have the same dtype. DecisionModel.forward convierte el estado agrupado y las características de confianza a fp32 antes de la cabeza de acción, y autocast normalmente reconcilia eso.

La ruta del artefacto necesita por tanto una cabeza de acción con dtype explícito: o convierte su entrada al dtype de la cabeza, o ejecuta la cabeza en fp32.

Portabilidad de TileLang: los kernels son solo de CUDA

tilelang registra objetivos para CUDA, HIP, Metal, WebGPU y un backend de C. Sin hardware AMD ni Apple, la pregunta que se podía responder era si laya/tl_kernels.py genera código para la CPU en absoluto. Sondeado con tilelang.compile(kernel.prim_func, target=...) en Linux x86-64:

objetivo resultado
"cpu" Rechazado de entrada: Target cpu is not supported. El backend de CPU de tilelang es "c".
"llvm" Cannot find global function target.build.llvm. La rueda no incluye backend de LLVM.
"c" Genera C y se ejecuta sobre tensores de CPU, pero solo para un subconjunto del lenguaje.

Todos los kernels de Laya fallan con "c", por una de tres razones:

kernel fallo con target="c" construcción
gemm_kernel, gemm_geglu_kernel CPU fill only supports local and global buffers, but got dst scope local.fragment acumulador T.alloc_fragment
add_ln_kernel CPU reduce only supports local src and local/local.var dst buffers T.reduce_sum / T.reduce_max sobre fragmentos
rope_kernel Cannot convert type bfloat16 to C type tensores bf16
attn_kernel falla en T.alloc_fragment fragmentos

El backend de C sí acepta:

  • bucles elementwise en fp32 (T.Parallel);
  • T.Pipelined, que genera un bucle simple;
  • T.gemm con un acumulador T.alloc_local, que genera un bucle escalar triple.

Una versión de CPU sería por tanto un segundo conjunto de kernels, no una bandera de objetivo. Su GEMM sería un bucle escalar sin bloques, y no competiría con la ruta MKL/oneDNN que la pasada hacia adelante estándar ya usa en CPU. Las mismas tres construcciones son las que hay que comprobar primero en HIP y Metal: fragmentos, T.gemm con GemmWarpPolicy, y soporte de bf16/fp16.

Para reproducir la primera fila de la segunda tabla: tilelang.compile(K.gemm_kernel(768, 768).prim_func, target="c").