Efficiently Dispatching Flash Attention For Partially Filled Attention Masks

Fuente: arXiv
Saved in:
Bibliographic Details
Main Authors: Sharma, Agniv, Geiping, Jonas
Format: Preprint
Published: 2024
Subjects:
Online Access:
Tags: Add Tag
No Tags, Be the first to tag this record!
_version_ 1866917783976017920
author Sharma, Agniv
Geiping, Jonas
author_facet Sharma, Agniv
Geiping, Jonas
contents Transformers are widely used across various applications, many of which yield sparse or partially filled attention matrices. Examples include attention masks designed to reduce the quadratic complexity of attention, sequence packing techniques, and recent innovations like tree masking for fast validation in MEDUSA. Despite the inherent sparsity in these matrices, the state-of-the-art algorithm Flash Attention still processes them with quadratic complexity as though they were dense. In this paper, we introduce Binary Block Masking, a highly efficient modification that enhances Flash Attention by making it mask-aware. We further propose two optimizations: one tailored for masks with contiguous non-zero patterns and another for extremely sparse masks. Our experiments on attention masks derived from real-world scenarios demonstrate up to a 9x runtime improvement. The implementation will be publicly released to foster further research and application.
format Preprint
id arxiv_https___arxiv_org_abs_2409_15097
institution arXiv
publishDate 2024
record_format arxiv
spellingShingle Efficiently Dispatching Flash Attention For Partially Filled Attention Masks
Sharma, Agniv
Geiping, Jonas
Machine Learning
Artificial Intelligence
Computation and Language
Transformers are widely used across various applications, many of which yield sparse or partially filled attention matrices. Examples include attention masks designed to reduce the quadratic complexity of attention, sequence packing techniques, and recent innovations like tree masking for fast validation in MEDUSA. Despite the inherent sparsity in these matrices, the state-of-the-art algorithm Flash Attention still processes them with quadratic complexity as though they were dense. In this paper, we introduce Binary Block Masking, a highly efficient modification that enhances Flash Attention by making it mask-aware. We further propose two optimizations: one tailored for masks with contiguous non-zero patterns and another for extremely sparse masks. Our experiments on attention masks derived from real-world scenarios demonstrate up to a 9x runtime improvement. The implementation will be publicly released to foster further research and application.
title Efficiently Dispatching Flash Attention For Partially Filled Attention Masks
topic Machine Learning
Artificial Intelligence
Computation and Language
url https://arxiv.org/abs/2409.15097