论文

Trident:统一 PyTorch Triton 工作负载的受保护调度和主机执行

Trident: Unifying Guarded Dispatch and Host Execution for PyTorch Triton Workloads

模型推理推理加速

摘要

用户编写的 Triton 内核可在 PyTorch 中实现高性能 GPU 计算,但其端到端延迟可能仍由主机端编排主导,尤其是在设备执行时间较短时。尽管 torch.compile 可以为捕获的图生成本机主机包装器,但每次调用在到达包装器之前仍然要经过运行时管理的专门化查找、防护评估和准备。我们推出了 Trident,这是一个编译器后端,可以从专门化缓存命中路径中消除这种重复出现的开销。 Trident 引入了专业化缓存模块 (SCM),它将受保护的专业化选择、参数和执行环境准备以及多个专业化的主机执行编译为单个可执行模块。调用进入 SCM 一次,当专业化匹配时保留在已编译的代码中,并且仅当必须编译新的专业化时才返回到 Python。 Trident 基于 Torch-MLIR 构建,将防护和主机端编排降低到本机代码,同时保留对支持的 ATen 运算符的优化运行时实现的调用。我们对两个大语言模型的评估表明,Trident 在模型级端到端延迟方面比 eager execution 实现了高达 1.47 倍的加速,比 torch.compile 实现了高达 1.68 倍的加速。