Documentação

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.compile só para CUDA com formas dinâmicas e modo reduce-overhead, preenchimento por buckets, cache de inductor persistente e warmup na instalação. Reutiliza o mesmo escopo de dimensão independente que compile=True, que mantém seu modo padrão existente e suporte a CPU. Defina LAYA_COMPILE_WARMUP=0 para adiar o warmup do backend e LAYA_INDUCTOR_CACHE_DIR para 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 o ONNXAgent existente. Sem onnx_path, usa laya.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=True precisa de dois grafos: um para lotes e um para uma única linha, que o torch especializa. compile=True agora chama agent.warmup() durante a carga. compile_warmup=False restaura a compilação preguiçosa, e agent.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 RuntimeWarning nomeando o erro (incluindo o erro do compilador subjacente) e a carga volta com o wrapper torch.compile e as configurações de compilação intactos. Por exemplo, Windows sem MSVC pode carregar com compile=True mesmo 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 a agent.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 o TORCHINDUCTOR_CACHE_DIR de todo o processo só quando ausente, para $XDG_CACHE_HOME/laya/torchinductor ou ~/.cache/laya/torchinductor quando 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.export do DecisionModel tem 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 alinhamento L % 8 do 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. O DecisionModel.forward promove 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.

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.gemm com um acumulador T.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").