# Added Grouped Query Attention to scaled\_dot\_product\_attention API

**URL:** <https://dev-discuss.pytorch.org/t/added-grouped-query-attention-to-scaled-dot-product-attention-api/2340>\
**Category:** frontend API\
**Created:** [August 1, 2024, 6:31pm UTC](https://dev-discuss.pytorch.org/t/added-grouped-query-attention-to-scaled-dot-product-attention-api/2340 "2024-08-01T18:31:00Z")\
**Posts on this page:** 1\
**Page:** 1

<div class="post-metadata">

**Author:** ![jainapurva](https://yyz2.discourse-cdn.com/flex036/user_avatar/dev-discuss.pytorch.org/jainapurva/32/2121_2.png) [@jainapurva](https://dev-discuss.pytorch.org/u/jainapurva)\
**Post date:** [August 1, 2024, 6:31pm UTC](https://dev-discuss.pytorch.org/t/added-grouped-query-attention-to-scaled-dot-product-attention-api/2340/1 "2024-08-01T18:31:00Z")

</div>

# Grouped Query Attention in SDPA: [PR#128898](https://github.com/pytorch/pytorch/pull/128898)

Grouped Query Attention (GQA) has emerged as an important technique to reduce the memory usage of the kv cache during inference. It has become increasingly popular in many foundational LLM models like llama2 70b and llama3. We have added this support to SDPA.

Reference paper: [[2305.13245] GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints](https://arxiv.org/abs/2305.13245)

# API Updates:

- Added a new kwarg enable\_gqa:Bool to the existing scaled\_dot\_product\_attention function. The default value is False (which would ensure regular SDPA functionality).
- The GQA cannot be used as default, as this memory layout is not supported by strides on, hence it needs to be explicitly enabled.
- The last third dimension (-3) in the query, key and value tensor has been set as a dedicated head-dimension.

# Important Notes:

- GQA is supported only by math and flash\_attention kernels.

# Implementation details:

- Following constraint are validated before performing GQA
- Query.head\_dim % Key.head\_dim == 0
- Query.head\_dim % Value.head\_dim == 0
- Key.head\_dim == Value.head\_dim
- If all conditions pass, in math kernel repeat\_interleave is performed on the key and value tensors to match the query tensor’s head dimension, while in flash attention repeat interleave is not needed, hence we conserve memory in the forward pass.

```auto

# Sample call to SDPA - GQ
query = torch.rand(batch, 32, seq_len_q, D)
key = torch.rand(batch, 8, seq_len_kv, D)
value = torch.rand(batch, 8, seq_len_kv, D)
output = scaled_dot_product_attention(query, key, value, is_causal=True, enable_gqa = True)

# Output Shape
(batch, 32, seq_len_q, D)

```

# Benchmarking:

### TorchTitan: [PR#458](https://github.com/pytorch/torchtitan/pull/458)

Profiling using Perfetto, on running TorchTitan training for Llama-8b

When SDPA is called without GQA, aten::reshape is called for both key and value tensor. The reshape call takes approximately 120us per call, which means approx 240us for 2 calls, while the flash\_fwd\_kernel takes approximately 3ms. In an ideal scenario, by removing these calls from the FlashAttention kernel, each SDPA flash kernel run will save 240us, which makes it ~6% faster than the previous runtime.

**SDPA call without enable\_gqa**

 ![](https://canada1.discourse-cdn.com/flex036/uploads/pytorch1/original/2X/a/a33aafeb2d30f3d8df38fd0734794a58da40bafa.png)

**SDPA call with enable\_gqa**

 ![](https://canada1.discourse-cdn.com/flex036/uploads/pytorch1/original/2X/e/e0c62b310c9377146149a6d68de2a468181a7e6a.png)

### SDPA Benchmarking: [PR#130634](https://github.com/pytorch/pytorch/pull/130634)

Graphs representing SDPA function run time with different parameters (batch size, number of key-value heads, number of query heads) with enable\_gqa=True and enable\_gqa=False.

| Batch size | q\_num\_heads | kv\_num\_heads | q\_seq\_len | kv\_seq\_len | embed\_dim | forward\_time when enable\_gqa=True (ms) | forward\_time when enable\_gqa=False (ms) |
| --- | --- | --- | --- | --- | --- | --- | --- |
| 1 | 32 | 8 | 2048 | 2048 | 2048 | 100.71 | 119.70 |
| 8 | 32 | 8 | 2048 | 2048 | 2048 | 539.78 | 628.83 |
| 16 | 32 | 8 | 2048 | 2048 | 2048 | 1056.81 | 1225.48 |
| 32 | 32 | 8 | 2048 | 2048 | 2048 | 2099.54 | 2440.45 |

 ![image](https://canada1.discourse-cdn.com/flex036/uploads/pytorch1/original/2X/9/902dc195dfe039a187d6253e8afe3c6c8dc77a85.png)

| Batch size | q\_num\_heads | kv\_num\_heads | q\_seq\_len | kv\_seq\_len | embed\_dim | forward\_time when enable\_gqa=True (ms) | forward\_time when enable\_gqa=False (ms) |
| --- | --- | --- | --- | --- | --- | --- | --- |
| 1 | 128 | 16 | 2048 | 2048 | 2048 | 243.30 | 260.51 |
| 8 | 128 | 16 | 2048 | 2048 | 2048 | 1766.17 | 1856.47 |
| 16 | 128 | 16 | 2048 | 2048 | 2048 | 3515.05 | 3675.95 |
| 32 | 128 | 16 | 2048 | 2048 | 2048 | 6996.68 | 7318.03 |

 ![image](https://canada1.discourse-cdn.com/flex036/uploads/pytorch1/original/2X/d/dfc8e2781cf3ec8ab9cc0cf06a4264ff05c5f85b.png)
