compile=True e o caminho rápido do TileLang: notas de engenharia
Estas notas cobrem como compile=True e fast=True se comportam para além do que diz o README.
Provêm de medições feitas durante o trabalho em #472, #576 e #718, numa RTX 4070 Ti SUPER com
torch 2.11 e tilelang 0.1.14. Estão aqui para que a pessoa seguinte não tenha de as medir outra
vez.
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 suportado.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 âmbito de dimensão independente quecompile=True, que mantém o seu modo predefinido existente e suporte de CPU. DefineLAYA_COMPILE_WARMUP=0para adiar o warmup do backend eLAYA_INDUCTOR_CACHE_DIRpara escolher o seu diretório de cache (predefiniçã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 suportados quando o TileLang está instalado; caso contrário, compile em CUDA; eager noutros dispositivos.onnx:laya.load(..., backend="onnx", onnx_path="model.onnx")devolve oONNXAgentexistente. Semonnx_path, usalaya.onnx.
Um backend indisponível emite um RuntimeWarning nomeando o backend resolvido e recua para eager. Para exigir um backend, usa agent.set_backend("tilelang", strict=True). A troca espera pela 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é à inferência, por isso os erros de compilação surgem então no pedido. agent.warmup() continua disponível. agent.deaccelerate() remove um backend instalado através da camada de classes.
Os Router reencaminham uma seleção explícita através de Router(agent_kwargs={"backend": "auto"}). Não passam qualquer argumento de backend por predefinição, preservando a compatibilidade com construtores do tipo Agent existentes. As tentativas de CPU OOM com âmbito limitado destacam o backend e restauram-no quando o modelo regressa ao seu dispositivo original.
compile=True materializa a máscara de atenção
A SDPA eager toma a máscara de atenção (rows, 1, L, L) do ModernBERT como uma vista de difusão.
Sob as formas dinâmicas que compile=True usa, o inductor não consegue provar que a última
dimensão está alinhada. Expande a máscara a todas as cabeças e preenche-a num buffer real de
rows x heads x L x L. Em bf16 com 12 cabeças, isso são 0.8 GB com 32 linhas x 1024 tokens.
- GPU com margem. O buffer custa largura de banda, dezenas de ms por lote longo.
- GPU quase cheia. O alocador de cache satura, e a mesma chamada pode levar dezenas de segundos.
Se compilares com lotes longos numa GPU ocupada, limita o tamanho do lote (predict_batch(..., batch_size=)) ou usa fast=True. A atenção do TileLang lê o buffer QKV empacotado e mascara pelo
comprimento da sequência, por isso não tem esse buffer.
Arranque a frio
- Primeira compilação. Leva dezenas de segundos por grafo.
compile=Trueprecisa de dois grafos: um para lotes e outro para uma única linha, que a torch especializa.compile=Truechama agoraagent.warmup()durante o carregamento.compile_warmup=Falserestaura a compilação preguiçosa, eagent.warmup(shapes=...)continua disponível manualmente. Os carregamentos eager e TileLang não aquecem automaticamente. Estas formas cobrem pedidos comuns, não todas as possíveis guardas de forma. - Falha do warmup. O warmup automático é de melhor esforço: uma falha emite um
RuntimeWarningnomeando o erro (incluindo o erro do compilador subjacente) e o carregamento regressa com o wrappertorch.compilee as definições de compilação intactos. Por exemplo, o Windows sem MSVC pode carregar comcompile=Truemesmo que o warmup falhe. Pedidos posteriores continuam a usar o modelo compilado e expõem falhas de compilação; o Laya não os muda para execução eager. Chamadas explícitas aagent.warmup()também propagam falhas, inclusive após um warmup automático falhado. Um carregamento bem-sucedido, portanto, não garante que a inferência compilada esteja pronta. - Opt-in da cache do Laya.
laya.load(..., compile=True, compile_cache=True)define oTORCHINDUCTOR_CACHE_DIRde todo o processo apenas quando ausente, para$XDG_CACHE_HOME/laya/torchinductorou~/.cache/laya/torchinductorquando o XDG não está definido ou não é absoluto. Uma definição existente, incluindo uma definida por uma compilação anterior do PyTorch, ganha. O diretório é criado no carregamento; erros de sistema de ficheiros propagam-se.compile_cache=False(predefinição), eager e carregamentos TileLang deixam o ambiente em paz. Isto não move nem elimina caches antigas. Os contentores continuam a precisar de um home/volume persistente. A compatibilidade e invalidação da 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. A cache de grafos FX do Inductor guarda os grafos compilados em
TORCHINDUCTOR_CACHE_DIR. O valor predefinido está em/tmp, que não sobrevive a um reinício nem a um reinício do contentor. Define-o para um diretório persistente, ou um volume num contentor, e um segundo processo carrega os grafos em vez de os compilar. Na medição do #472, isso baixou o arranque 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 predefinição; no caminho compilado ativo, só "default" e "reduce-overhead" são aceites. Os carregamentos eager e TileLang ignoram as opções de compilação. A compilação em CPU continua a funcionar, mas o registo 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 levantam um erro explícito.
Grafos Dynamo dinâmicos não implicam grafos CUDA independentes da forma: novas formas concretas podem exigir warmup e registo novamente, sem um novo grafo Dynamo. As duas formas sintéticas de warmup predefinidas não pré-registam todas as formas de pedido. Formas repetidas podem beneficiar, mas formas variáveis podem pagar latência extra e retêm pools de grafos. O PyTorch pode ignorar grafos CUDA para operações ou configurações não suportadas; definir este modo não é garantia de captura.
Laya marca cada forward CUDA compilado como um novo passo, serializa esses forwards entre os seus agentes e clona ambos os tensores de saída fora do grafo compilado antes de libertar o lock. Isto 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; os autores de chamada que partilham iterações de grafos CUDA ou usam streams personalizados têm de gerir a sua própria coordenação. As caches em disco reutilizam código compilado, não gravações vivas de grafos CUDA nem a sua memória de dispositivo, entre processos.
Reproduz os tempos de frio/reinício, a memória e os contadores de cache com
benchmarks/bench_compile_defaults.py; vê
as medições registadas.
AOTInductor: ainda não
Distribuir um artefacto pré-compilado por checkpoint e arquitetura de GPU
(torch._inductor.aoti_compile_and_package) eliminaria a compilação por completo. Na torch 2.11,
pára no empacotamento:
- A exportação funciona. O
torch.exportdeDecisionModeltem êxito, em cerca de 5 s, com linhas, marcadores e tokens dinâmicos. Os tokens têm de ser declarados como múltiplo de 16 (16 * Dim(...)); um intervalo simples falha a própria guarda de alinhamentoL % 8do exportador. É o mesmo alinhamento de máscara que o de cima. - O empacotamento falha. A forma como falha depende de como o programa foi exportado:
- Sob autocast, o programa transporta asserções de dtype em que 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.forwardconverte 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 transporta asserções de dtype em que o AOTI tropeça fora do
autocast:
A rota do artefacto precisa, por isso, de uma cabeça de ação com dtype explícito: ou converte a sua entrada para o dtype da cabeça, ou executa a cabeça em fp32.
Portabilidade do TileLang: os kernels são só para CUDA
O tilelang regista alvos para CUDA, HIP, Metal, WebGPU e um backend C. Sem hardware AMD ou Apple,
a pergunta a que se podia responder era se laya/tl_kernels.py se baixa para a CPU de todo.
Testado com tilelang.compile(kernel.prim_func, target=...) em 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 inclui backend LLVM. |
"c" |
Baixa para C e corre em tensores de CPU, mas só para um subconjunto da linguagem. |
Todos os kernels do Laya falham em "c", por uma de três razões:
| 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 sim:
- ciclos elementwise em fp32 (
T.Parallel); T.Pipelined, que baixa para um ciclo simples;T.gemmcom um acumuladorT.alloc_local, que baixa para um ciclo escalar triplo.
Uma versão para CPU seria, por isso, um segundo conjunto de kernels, e não uma flag de alvo. O seu
GEMM seria um ciclo escalar sem blocos, e não competiria com o caminho MKL/oneDNN que o forward de
base já usa na CPU. As mesmas três construções são as que se devem verificar primeiro em HIP e
Metal: fragmentos, T.gemm com GemmWarpPolicy, e suporte de bf16/fp16.
Para reproduzir a primeira linha da segunda tabela:
tilelang.compile(K.gemm_kernel(768, 768).prim_func, target="c").