Overview
Trident is presented as a compiler backend engineered to consolidate guarded dispatch and host execution within PyTorch Triton workloads. It specifically targets and mitigates the end-to-end latency contribution of host-side orchestration, which can be particularly pronounced when GPU device execution times are brief. By unifying specialization selection, argument and execution-environment preparation, and host execution into a singular executable module, Trident aims to remove recurring overheads associated with the specialization cache-hit path.
Research Context
User-authored Triton kernels facilitate high-performance GPU computations within the PyTorch framework. Despite the computational efficiency on the device side, the overarching end-to-end latency for these kernels can frequently be dominated by the orchestration processes managed on the host CPU. While `torch.compile` can generate native host wrappers for captured computational graphs, each subsequent invocation still necessitates navigating through a sequence of runtime-managed operations. These operations include specialization lookup, evaluation of guards, and preparation steps, before the execution can proceed to the native wrapper. This sequential overhead is identified as a persistent bottleneck, even with existing compilation strategies.
Approach
Trident's methodology centers on the introduction of a Specialization Cache Module (SCM). The SCM functions as a compilation target that integrates several critical host-side processes: guarded specialization selection, the preparation of arguments, and the setup of the execution environment. This integration encompasses multiple potential specializations into a single, cohesive executable module. The design allows an invocation to enter the SCM once and remain within the compiled code path when a matching specialization is identified. Control returns to Python only in scenarios where a novel specialization requires on-the-fly compilation. The architecture of Trident is built upon Torch-MLIR. This foundation enables the lowering of both guards and host-side orchestration logic into native code. Simultaneously, Trident preserves the ability to call into optimized runtime implementations of supported ATen operators, leveraging existing efficiencies for core tensor operations.
Findings
Evaluations of Trident were conducted on two distinct Large Language Models (LLMs). The results indicated that Trident achieved notable performance improvements in end-to-end latency at the model level:
- Trident demonstrated up to a 1.47x speedup when compared against eager execution.
- Trident achieved up to a 1.68x speedup when compared against `torch.compile`.
Why This Matters
The observed speedups on LLMs suggest that Trident can directly address a known performance bottleneck in high-performance GPU computation workflows utilizing PyTorch Triton kernels. By minimizing the host-side overheads that often dominate overall latency, especially for short device execution tasks, Trident may contribute to more efficient execution of machine learning models.