compile=True e o caminho rápido do TileLang: notas de engenharia
Estas notas cobrem como compile=True e fast=True se comportam além do que diz o README. Elas
vêm de medições feitas durante o trabalho nos #472, #576 e #718, em uma RTX 4070 Ti SUPER com
torch 2.11 e tilelang 0.1.14. Elas estão aqui para que a próxima pessoa não precise medi-las de novo.
Seleção de backend
Agent(..., backend="auto") e laya.load(..., backend="auto") optam pela camada de classes de backend. O padrão continua eager. Um backend= explícito tem precedência sobre compile e fast; omiti-lo preserva o comportamento existente de ambos os flags.
eager: o forward padrão do PyTorch, em qualquer dispositivo compatível.compile:torch.compilesó para CUDA com formas dinâmicas e modoreduce-overhead, preenchimento por buckets, cache de inductor persistente e warmup na instalação. Reutiliza o mesmo escopo de dimensão independente quecompile=True, que mantém seu modo padrão existente e suporte a CPU. DefinaLAYA_COMPILE_WARMUP=0para adiar o warmup do backend eLAYA_INDUCTOR_CACHE_DIRpara escolher seu diretório de cache (padrão~/.cache/laya/inductor).tilelang: um adaptador em torno do caminho rápido atual, usando o dtype bf16 ou fp16 do agente.auto: TileLang em CUDA com um codificador ModernBERT e dtype compatíveis quando o TileLang está instalado; caso contrário, compile em CUDA; eager em outros dispositivos.onnx:laya.load(..., backend="onnx", onnx_path="model.onnx")retorna oONNXAgentexistente. Semonnx_path, usalaya.onnx.
Um backend indisponível emite um RuntimeWarning nomeando o backend resolvido e recorre a eager. Para exigir um backend, use agent.set_backend("tilelang", strict=True). A troca espera a inferência ativa; agent.backend reporta o nome ativo e agent.backend_object expõe o objeto instalado. agent.set_backend("compile", warmup=False) adia a compilação até a inferência, então erros de compilação surgem na solicitação. agent.warmup() continua disponível. agent.deaccelerate() remove um backend instalado pela camada de classes.
Os Router repassam uma seleção explícita via Router(agent_kwargs={"backend": "auto"}). Eles não passam nenhum argumento de backend por padrão, preservando a compatibilidade com construtores tipo Agent existentes. As tentativas de CPU OOM com escopo detacham o backend e o restauram quando o modelo volta ao seu dispositivo original.
compile=True materializa a máscara de atenção
A SDPA em modo eager trata a máscara de atenção (rows, 1, L, L) do ModernBERT como uma view de
broadcast. Sob as formas dinâmicas que compile=True usa, o inductor não consegue provar que a
última dimensão está alinhada. Ele expande a máscara para cada cabeça e a preenche em um buffer real
de rows x heads x L x L. Em bf16 com 12 cabeças, isso é 0.8 GB com 32 linhas x 1024 tokens.
- GPU com folga. O buffer custa largura de banda, dezenas de ms por lote longo.
- GPU quase cheia. O alocador de cache entra em thrashing, e a mesma chamada pode levar dezenas de segundos.
Se você compila com lotes longos em uma GPU ocupada, limite o tamanho do lote (predict_batch(..., batch_size=)) ou use fast=True. A atenção do TileLang lê o buffer QKV empacotado e mascara por
comprimento de sequência, então ela não tem esse buffer.
Inicialização a frio
- Primeira compilação. Leva dezenas de segundos por grafo.
compile=Trueprecisa de dois grafos: um para lotes e um para uma única linha, que o torch especializa.compile=Trueagora chamaagent.warmup()durante a carga.compile_warmup=Falserestaura a compilação preguiçosa, eagent.warmup(shapes=...)continua disponível manualmente. Cargas eager e TileLang não aquecem automaticamente. Essas formas cobrem solicitações comuns, não toda possível guarda de forma. - Falha no warmup. O warmup automático é de melhor esforço: uma falha emite um
RuntimeWarningnomeando o erro (incluindo o erro do compilador subjacente) e a carga volta com o wrappertorch.compilee as configurações de compilação intactos. Por exemplo, Windows sem MSVC pode carregar comcompile=Truemesmo que o warmup falhe. Solicitações posteriores ainda usam o modelo compilado e expõem falhas de compilação; o Laya não as muda para execução eager. Chamadas explícitas aagent.warmup()também propagam falhas, inclusive após um warmup automático malsucedido. Uma carga bem-sucedida, portanto, não garante que a inferência compilada esteja pronta. - Opt-in do cache do Laya.
laya.load(..., compile=True, compile_cache=True)define oTORCHINDUCTOR_CACHE_DIRde todo o processo só quando ausente, para$XDG_CACHE_HOME/laya/torchinductorou~/.cache/laya/torchinductorquando o XDG não está definido ou não é absoluto. Uma configuração existente, incluindo uma definida por uma compilação anterior do PyTorch, vence. O diretório é criado na carga; erros de sistema de arquivos se propagam.compile_cache=False(padrão), eager e cargas TileLang deixam o ambiente em paz. Isso não move nem exclui caches antigos. Contêineres ainda precisam de um home/volume persistente. A compatibilidade e a invalidação do cache são geridas pelo PyTorch; uma mudança de GPU, torch, compilador, modelo ou guarda de entrada pode exigir nova compilação. - Entre reinícios. O cache de grafos FX do Inductor mantém os grafos compilados sob
TORCHINDUCTOR_CACHE_DIR. O padrão fica sob/tmp, que não sobrevive a um reboot nem a um reinício de contêiner. Aponte-o para um diretório persistente, ou para um volume em um contêiner, e um segundo processo carrega os grafos em vez de compilá-los. Na medição do #472, isso reduziu o aquecimento de cerca de 120 s para cerca de 50 s.
Grafos CUDA opcionais
agent = laya.load("convaiinnovations/laya", compile=True,
compile_cache=True, compile_mode="reduce-overhead")
compile_mode é "default" por padrão; no caminho compilado ativo, só "default" e "reduce-overhead" são aceitos. Cargas eager e TileLang ignoram as opções de compilação. A compilação em CPU ainda funciona, mas o registro de grafos CUDA só se aplica em CUDA. O modo CUDA requer a API torch.compiler.cudagraph_mark_step_begin do PyTorch; builds mais antigas sem ela lançam um erro explícito.
Grafos Dynamo dinâmicos não implicam grafos CUDA independentes da forma: novas formas concretas podem exigir warmup e registro de novo, sem um novo grafo Dynamo. As duas formas sintéticas de warmup padrão não pré-gravam toda forma de solicitação. Formas repetidas podem se beneficiar, mas formas variáveis podem pagar latência extra e reter pools de grafos. O PyTorch pode pular grafos CUDA para operações ou configurações não compatíveis; definir este modo não é garantia de captura.
Laya marca cada forward CUDA compilado como um novo passo, serializa esses forwards entre seus agentes e clona ambos os tensores de saída fora do grafo compilado antes de liberar o lock. Isso mantém as saídas retidas válidas entre replays, ao custo de duas cópias e execução de forward serializada. O lock não coordena modelos compilados independentes da própria aplicação; chamadores que compartilham iterações de grafos CUDA ou usam streams personalizados devem gerenciar sua própria coordenação. Caches em disco reutilizam código compilado, não gravações vivas de grafos CUDA nem sua memória de dispositivo, entre processos.
Reproduza os tempos de frio/reinício, a memória e os contadores de cache com
benchmarks/bench_compile_defaults.py; veja
as medições registradas.
AOTInductor: ainda não
Distribuir um artefato pré-compilado por checkpoint e arquitetura de GPU
(torch._inductor.aoti_compile_and_package) eliminaria a compilação por completo. No torch 2.11 ele
para no empacotamento:
- A exportação funciona.
torch.exportdoDecisionModeltem sucesso, em cerca de 5 s, com linhas, marcadores e tokens dinâmicos. Os tokens precisam ser declarados como múltiplo de 16 (16 * Dim(...)); um intervalo simples falha na própria guarda de alinhamentoL % 8do exportador. É o mesmo alinhamento de máscara de cima. - O empacotamento falha. Como ele falha depende de como o programa foi exportado:
- Sob autocast, o programa carrega asserções de dtype nas quais o AOTI tropeça fora do autocast:
Tensor dtype mismatch! Expected: torch.bfloat16, Got: torch.float32. - A partir de uma cópia bf16 sem autocast, o tracing falha dentro do forward:
mat1 and mat2 must have the same dtype. ODecisionModel.forwardpromove o estado agrupado e as características de confiança para fp32 antes da cabeça de ação, e o autocast normalmente reconcilia isso.
- Sob autocast, o programa carrega asserções de dtype nas quais o AOTI tropeça fora do autocast:
O caminho do artefato, portanto, precisa de uma cabeça de ação com dtype explícito: ou converte sua entrada para o dtype da cabeça, ou roda a cabeça em fp32.
Portabilidade do TileLang: os kernels são apenas de CUDA
O tilelang registra alvos para CUDA, HIP, Metal, WebGPU e um backend C. Sem hardware AMD ou Apple, a
pergunta que dava para responder era se laya/tl_kernels.py baixa para CPU de alguma forma.
Sondado com tilelang.compile(kernel.prim_func, target=...) no Linux x86-64:
| alvo | resultado |
|---|---|
"cpu" |
Rejeitado de imediato: Target cpu is not supported. O backend de CPU do tilelang é "c". |
"llvm" |
Cannot find global function target.build.llvm. A wheel não traz backend LLVM. |
"c" |
Baixa para C e roda em tensores de CPU, mas só para um subconjunto da linguagem. |
Todo kernel do Laya falha em "c", por um de três motivos:
| kernel | falha com target="c" |
construção |
|---|---|---|
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 |
falha em T.alloc_fragment |
fragmentos |
O backend C aceita:
- laços elementwise em fp32 (
T.Parallel); T.Pipelined, que baixa para um laço simples;T.gemmcom um acumuladorT.alloc_local, que baixa para um laço escalar triplo.
Uma versão de CPU, portanto, seria um segundo conjunto de kernels, não uma flag de alvo. Seu GEMM
seria um laço escalar sem blocos, e não competiria com o caminho MKL/oneDNN que o forward padrão já
usa na CPU. As mesmas três construções são as que devem ser verificadas primeiro em HIP e Metal:
fragmentos, T.gemm com GemmWarpPolicy, e suporte a bf16/fp16.
Para reproduzir a primeira linha da segunda tabela:
tilelang.compile(K.gemm_kernel(768, 768).prim_func, target="c").