Документация

compile=True и быстрый путь TileLang: инженерные заметки

Эти заметки описывают поведение compile=True и fast=True за пределами того, что сказано в README. Они основаны на измерениях, сделанных во время работы над #472, #576 и #718, на RTX 4070 Ti SUPER с torch 2.11 и tilelang 0.1.14. Они здесь для того, чтобы следующему человеку не пришлось измерять это снова.

Выбор backend

Agent(..., backend="auto") и laya.load(..., backend="auto") включают слой классов backend. По умолчанию остаётся eager. Явный backend= имеет приоритет над compile и fast; если его не указывать, сохраняется существующее поведение обоих флагов.

  • eager: штатный forward PyTorch на любом поддерживаемом устройстве.
  • compile: torch.compile только для CUDA с динамическими формами и режимом reduce-overhead, заполнением по бакетам, постоянным кэшем inductor и прогревом при установке. Он использует ту же область независимых измерений, что и compile=True, который сохраняет свой существующий режим по умолчанию и поддержку CPU. Установите LAYA_COMPILE_WARMUP=0, чтобы отложить прогрев backend, и LAYA_INDUCTOR_CACHE_DIR, чтобы выбрать его каталог кэша (по умолчанию ~/.cache/laya/inductor).
  • tilelang: адаптер вокруг текущего быстрого пути, использующий dtype агента bf16 или fp16.
  • auto: TileLang на CUDA с поддерживаемым кодировщиком ModernBERT и dtype, когда TileLang установлен, иначе compile на CUDA; eager на других устройствах.
  • onnx: laya.load(..., backend="onnx", onnx_path="model.onnx") возвращает существующий ONNXAgent. Без onnx_path он использует laya.onnx.

Недоступный backend выдаёт RuntimeWarning, называя разрешённый backend, и откатывается к eager. Чтобы потребовать backend, используйте agent.set_backend("tilelang", strict=True). Переключение ждёт активного инференса; agent.backend сообщает активное имя, а agent.backend_object открывает установленный объект. agent.set_backend("compile", warmup=False) откладывает компиляцию до инференса, так что ошибки компиляции всплывают на запросе. agent.warmup() остаётся доступным. agent.deaccelerate() удаляет backend, установленный через слой классов.

Router передают явный выбор через Router(agent_kwargs={"backend": "auto"}). По умолчанию они не передают аргумент backend, сохраняя совместимость с существующими конструкторами типа Agent. Ограниченные повторы при CPU OOM отсоединяют backend и восстанавливают его, когда модель возвращается на исходное устройство.

compile=True материализует маску внимания

Eager SDPA принимает маску внимания ModernBERT (rows, 1, L, L) как broadcast-представление. При динамических формах, которые использует compile=True, inductor не может доказать, что последнее измерение выровнено. Он разворачивает маску на каждую голову и дополняет её в реальный буфер rows x heads x L x L. В bf16 с 12 головами это 0.8 ГБ при 32 строках x 1024 токенах.

  • GPU с запасом. Буфер стоит пропускной способности — десятки мс на длинный батч.
  • GPU почти заполнен. Кэширующий аллокатор начинает метаться, и тот же вызов может занять десятки секунд.

Если вы компилируете с длинными батчами на занятом GPU, ограничьте размер батча (predict_batch(..., batch_size=)) или используйте fast=True. Внимание TileLang читает упакованный буфер QKV и маскирует по длине последовательности, поэтому такого буфера у него нет.

Холодный старт

  • Первая компиляция. Занимает десятки секунд на граф. compile=True требует два графа: один для батчей и один для одной строки, который torch специализирует. compile=True теперь вызывает agent.warmup() во время загрузки. compile_warmup=False восстанавливает ленивую компиляцию, а agent.warmup(shapes=...) остаётся доступным вручную. Загрузки eager и TileLang не прогревают автоматически. Эти формы покрывают распространённые запросы, а не каждую возможную защиту формы.
  • Сбой прогрева. Автоматический прогрев выполняется по мере возможности: сбой выдаёт RuntimeWarning с указанием ошибки (включая ошибку нижележащего компилятора), и загрузка возвращается с целыми обёрткой torch.compile и настройками компиляции. Например, Windows без MSVC может загрузиться с compile=True, даже если прогрев не удался. Последующие запросы всё равно используют скомпилированную модель и всплывают ошибки компиляции; Laya не переключает их на eager-исполнение. Явные вызовы agent.warmup() также распространяют сбои, в том числе после неудачного автоматического прогрева. Поэтому успешная загрузка не гарантирует, что скомпилированный инференс готов.
  • Выбор кэша Laya. laya.load(..., compile=True, compile_cache=True) задаёт TORCHINDUCTOR_CACHE_DIR для всего процесса только когда он отсутствует, в $XDG_CACHE_HOME/laya/torchinductor или ~/.cache/laya/torchinductor, когда XDG не задан или не абсолютен. Существующая настройка, включая заданную более ранней компиляцией PyTorch, побеждает. Каталог создаётся при загрузке; ошибки файловой системы распространяются. compile_cache=False (по умолчанию), eager и загрузки TileLang оставляют окружение в покое. Это не перемещает и не удаляет старые кэши. Контейнерам всё ещё нужен постоянный home/том. Совместимость и инвалидация кэша управляются PyTorch; изменение GPU, torch, компилятора, модели или входной защиты может потребовать повторной компиляции.
  • Между перезапусками. Кэш FX-графов Inductor хранит скомпилированные графы в TORCHINDUCTOR_CACHE_DIR. По умолчанию он находится в /tmp, который не переживает перезагрузку или перезапуск контейнера. Задайте постоянный каталог или том в контейнере, и второй процесс загрузит графы вместо их компиляции. В измерении из #472 это сократило прогрев примерно с 120 с до примерно 50 с.

CUDA-графы по желанию

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

compile_mode по умолчанию "default"; на активном скомпилированном пути принимаются только "default" и "reduce-overhead". Загрузки eager и TileLang игнорируют опции компиляции. Компиляция CPU по-прежнему работает, но запись CUDA-графов применяется только на CUDA. Режим CUDA требует API PyTorch torch.compiler.cudagraph_mark_step_begin; более старые сборки без него вызывают явную ошибку.

Динамические графы Dynamo не означают CUDA-графы, независимые от формы: новые конкретные формы могут потребовать повторного прогрева и записи без нового графа Dynamo. Две стандартные синтетические формы прогрева не предзаписывают каждую форму запроса. Повторяющиеся формы могут выиграть, но варьирующиеся формы могут платить дополнительной задержкой и удерживать пулы графов. PyTorch может пропускать CUDA-графы для неподдерживаемых операций или конфигураций; установка этого режима не гарантирует захват.

Laya помечает каждый скомпилированный CUDA-forward как новый шаг, сериализует эти forwards между своими агентами и клонирует оба выходных тензора вне скомпилированного графа, прежде чем освободить блокировку. Это сохраняет удержанные выходы валидными между повторами, ценой двух копий и сериализованного выполнения forward. Блокировка не координирует несвязанные скомпилированные модели, принадлежащие приложению; вызывающие, совместно использующие итерации CUDA-графов или использующие пользовательские потоки, должны управлять своей координацией. Дисковые кэши переиспользуют скомпилированный код, а не живые записи CUDA-графов или их память устройства, между процессами.

Воспроизведите тайминги холода/перезапуска, память и счётчики кэша с помощью benchmarks/bench_compile_defaults.py; см. записанные измерения.

AOTInductor: пока нет

Поставка предкомпилированного артефакта для каждого чекпойнта и архитектуры GPU (torch._inductor.aoti_compile_and_package) полностью убрала бы компиляцию. На torch 2.11 всё останавливается на упаковке:

  • Экспорт работает. torch.export для DecisionModel проходит успешно, примерно за 5 с, с динамическими строками, маркерами и токенами. Токены должны быть объявлены как кратное 16 (16 * Dim(...)); обычный диапазон проваливает собственную защиту выравнивания L % 8 экспортёра. Это то же выравнивание маски, что и выше.
  • Упаковка падает. То, как она падает, зависит от того, как была экспортирована программа:
    • При autocast программа несёт проверки dtype, о которые AOTI спотыкается вне autocast: Tensor dtype mismatch! Expected: torch.bfloat16, Got: torch.float32.
    • Из bf16-копии без autocast трассировка падает внутри прямого прохода: mat1 and mat2 must have the same dtype. DecisionModel.forward повышает объединённое состояние и признаки уверенности до fp32 перед головой действий, и autocast обычно это согласует.

Поэтому путь через артефакт требует головы действий с явным dtype: либо привести её вход к dtype головы, либо выполнять голову в fp32.

Переносимость TileLang: ядра только для CUDA

tilelang регистрирует цели для CUDA, HIP, Metal, WebGPU и бэкенда C. Без оборудования AMD или Apple вопрос, на который можно было ответить, состоял в том, понижает ли laya/tl_kernels.py код вообще для CPU. Проверено с помощью tilelang.compile(kernel.prim_func, target=...) на Linux x86-64:

цель результат
"cpu" Отклоняется сразу: Target cpu is not supported. CPU-бэкенд tilelang — это "c".
"llvm" Cannot find global function target.build.llvm. Пакет wheel не содержит бэкенда LLVM.
"c" Понижает до C и работает на CPU-тензорах, но только для подмножества языка.

Каждое ядро Laya падает на "c" по одной из трёх причин:

ядро сбой при target="c" конструкция
gemm_kernel, gemm_geglu_kernel CPU fill only supports local and global buffers, but got dst scope local.fragment аккумулятор T.alloc_fragment
add_ln_kernel CPU reduce only supports local src and local/local.var dst buffers T.reduce_sum / T.reduce_max по фрагментам
rope_kernel Cannot convert type bfloat16 to C type тензоры bf16
attn_kernel падает на T.alloc_fragment фрагменты

Бэкенд C принимает:

  • elementwise-циклы fp32 (T.Parallel);
  • T.Pipelined, который понижается до обычного цикла;
  • T.gemm с аккумулятором T.alloc_local, который понижается до скалярного тройного цикла.

Таким образом, версия для CPU была бы вторым набором ядер, а не флагом цели. Её GEMM был бы неблочным скалярным циклом и не конкурировал бы с путём MKL/oneDNN, который стандартный прямой проход уже использует на CPU. Те же три конструкции нужно проверять первыми на HIP и Metal: фрагменты, T.gemm с GemmWarpPolicy и поддержку bf16/fp16.

Чтобы воспроизвести первую строку второй таблицы: tilelang.compile(K.gemm_kernel(768, 768).prim_func, target="c").