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 обычно это согласует.
- При autocast программа несёт проверки dtype, о которые AOTI спотыкается вне 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").