This blog, introducing OpenAI's new python extension called Triton, says this about why Triton can do matrix math faster than pytorch (referring to an an example of how Triton can be used to compute Softmax along the rows of an m by n matrix)
Importantly, this particular implementation of softmax keeps the rows of X in SRAM throughout the entire normalization process, which maximizes data reuse when applicable (~<32K columns). This differs from PyTorch’s internal CUDA code, whose use of temporary memory makes it more general but significantly slower (below). The bottom line here is not that Triton is inherently better, but that it simplifies the development of specialized kernels that can be much faster than those found in general-purpose libraries.
- How does pytorch allocate memory for device tensors, what is the "temporary memory" being referred to here? Why is the use of this temporary memory more general, but slower than use of SRAM?
- Is SRAM here referring to cache memory? If so, how/why does this library make better use of cache memory than pytorch internals? My understanding is that the decision about what data to cache is mostly up to the hardware rather than software.