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.compilesolo para CUDA con formas dinámicas y modoreduce-overhead, relleno por buckets, caché de inductor persistente y calentamiento en la instalación. Reutiliza el mismo alcance de dimensión independiente quecompile=True, que mantiene su modo predeterminado existente y soporte de CPU. DefineLAYA_COMPILE_WARMUP=0para diferir el calentamiento del backend yLAYA_INDUCTOR_CACHE_DIRpara 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 elONNXAgentexistente. Sinonnx_path, usalaya.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=Truenecesita dos grafos: uno para lotes y otro para una sola fila, que torch especializa.compile=Trueahora llama aagent.warmup()durante la carga.compile_warmup=Falserestaura la compilación perezosa, yagent.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
RuntimeWarningque nombra el error (incluido el error del compilador subyacente) y la carga vuelve con el envoltoriotorch.compiley los ajustes de compilación intactos. Por ejemplo, Windows sin MSVC puede cargar concompile=Trueaunque 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 aagent.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 laTORCHINDUCTOR_CACHE_DIRde todo el proceso solo cuando falta, a$XDG_CACHE_HOME/laya/torchinductoro~/.cache/laya/torchinductorcuando 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.exportdeDecisionModeltiene é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ónL % 8del 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.forwardconvierte el estado agrupado y las características de confianza a fp32 antes de la cabeza de acción, y autocast normalmente reconcilia eso.
- Bajo autocast, el programa lleva aserciones de dtype con las que AOTI tropieza fuera de autocast:
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.gemmcon un acumuladorT.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").