Grouped Query Attention (GQA)
Tìm hiểu cách Grouped Query Attention (GQA) giảm bộ nhớ KV-cache, cải thiện hiệu quả inference và cân bằng hiệu năng trong các model Transformer.
Attention truy vấn theo nhóm (GQA) là một thiết kế attention trong đó nhiều query head chia sẻ một tập key head và value head nhỏ hơn. Thiết kế này duy trì các góc nhìn đa dạng của các query head, đồng thời giảm lượng bộ nhớ và việc di chuyển dữ liệu cần thiết cho key và value. GQA đặc biệt hữu ích trong quá trình suy luận tự hồi quy, khi model tạo từng token một và liên tục đọc các trạng thái attention đã lưu.
Cách GQA hoạt động#
Trong một cơ chế attention, mỗi token được chiếu thành ba biểu diễn:
- Query: Mô tả thông tin mà token hiện tại đang tìm kiếm.
- Key: Mô tả nội dung mà mỗi token khả dụng biểu diễn.
- Value: Chứa thông tin được truy xuất khi query khớp với key.
Trong self-attention truyền thống, các phép chiếu này được chia thành các head để model có thể học các mối quan hệ khác nhau song song. GQA giữ lại nhiều query head nhưng phân chúng thành các nhóm dùng chung key head và value head. Ví dụ, một attention layer có thể sử dụng tám query head và hai key-value head. Khi đó, mỗi key-value head phục vụ bốn query head.
Bản thân phép tính attention vẫn là scaled dot-product attention. Điều thay đổi là cách sắp xếp head và lượng dữ liệu key-value được tạo ra, lưu trữ và đọc. Do đó, các framework như PyTorch scaled dot-product attention yêu cầu số query head phải chia hết cho số key-value head khi bật GQA. (docs.pytorch.org)
GQA so với Multi-Head Attention và Multi-Query Attention#
GQA nằm giữa multi-head attention và multi-query attention:
- Multi-head attention: Mỗi query head có key head và value head riêng. Cách này mang lại mức độ độc lập tối đa giữa các head nhưng tạo ra yêu cầu bộ nhớ key-value lớn nhất.
- Attention truy vấn theo nhóm: Một số query head chia sẻ mỗi key-value head. Cách này cân bằng năng lực biểu diễn với hiệu quả sử dụng bộ nhớ.
- Multi-query attention: Tất cả query head chia sẻ một key head và một value head. Cách này tối thiểu hóa dung lượng lưu trữ key-value nhưng cung cấp ít tính đa dạng key-value hơn.
Nếu một layer có 32 query head, multi-head attention cũng có thể sử dụng 32 key-value head, GQA có thể sử dụng 8, còn multi-query attention sử dụng 1. GQA không thay thế kiến trúc Transformer ở phạm vi rộng hơn; đây là một cấu hình khả dĩ bên trong một attention layer của Transformer. Tài liệu về multi-head, multi-query và attention truy vấn theo nhóm của NVIDIA mô tả các biến thể này là khác nhau chủ yếu ở số lượng key-value head phục vụ các query head. (nvidia.github.io)
Tầm quan trọng của GQA#
Trong quá trình sinh tự hồi quy, model lưu các key và value trước đó vào KV cache. Với multi-head attention tiêu chuẩn, mỗi attention head đóng góp các key và value riêng vào cache. Do đó, các chuỗi dài, batch lớn và nhiều layer của model có thể tiêu tốn đáng kể bộ nhớ GPU.
Vì GQA sử dụng ít key-value head hơn, cache của nó nhỏ hơn tương ứng với mức giảm số lượng head đó. Điều này có thể cải thiện:
- Dung lượng bộ nhớ: Có thể chứa prompt dài hơn hoặc nhiều request đồng thời hơn trong bộ nhớ khả dụng.
- Throughput khi sinh: Cần đọc ít dữ liệu key-value hơn cho mỗi token được sinh.
- Độ trễ suy luận: Lưu lượng bộ nhớ giảm có thể rút ngắn thời gian xử lý giữa các token.
- Khả năng mở rộng của context window: Việc phục vụ các chuỗi dài trở nên thực tế hơn.
Các lợi ích này phụ thuộc vào phần cứng, độ dài chuỗi, kích thước batch và khả năng hỗ trợ của kernel. Các runtime production vẫn cần cơ chế phân bổ cache hiệu quả; chẳng hạn, hệ thống KV cache của TensorRT-LLM của NVIDIA kết hợp hỗ trợ GQA với caching dựa trên block, tái sử dụng và offloading. (nvidia.github.io)
Các ứng dụng trong thực tế#
Trợ lý với context dài: Một coding assistant có thể xử lý hàng nghìn token mã nguồn trước khi tạo câu trả lời. GQA làm giảm trạng thái key-value được cache gắn với lịch sử đó, cho phép server xử lý các file dài hơn hoặc nhiều user đồng thời hơn mà không làm tăng bộ nhớ GPU theo tỷ lệ tương ứng.
Trợ lý hình ảnh và video đa phương thức: Một vision-language model có thể biểu diễn các patch hình ảnh hoặc frame video thành các chuỗi token dài. GQA làm giảm áp lực lên attention cache sau khi các token hình ảnh này đi vào language decoder, qua đó có thể dành thêm bộ nhớ cho các hình ảnh, frame hoặc request bổ sung. Điều này khác với một Vision Transformer tiêu chuẩn, trong đó attention có thể xử lý song song tất cả patch hình ảnh mà không cần KV cache tự hồi quy.
Triển khai thực tế và các đánh đổi#
PyTorch cung cấp GQA thông qua enable_gqa=True. Ví dụ CUDA được ghi chép này sử dụng tám query head và hai key-value head dùng chung:
import torch
import torch.nn.functional as F
from torch.nn.attention import SDPBackend, sdpa_kernel
assert torch.cuda.is_available(), "A CUDA device is required."
query = torch.randn(1, 8, 16, 32, device="cuda")
key = torch.randn(1, 2, 16, 32, device="cuda")
value = torch.randn(1, 2, 16, 32, device="cuda")
# Four query heads share each key-value head.
with sdpa_kernel(SDPBackend.MATH):
output = F.scaled_dot_product_attention(
query,
key,
value,
enable_gqa=True,
)
print(output.shape)Đầu ra giữ lại tám query head, trong khi key và value chỉ sử dụng hai head. Bộ chọn backend attention của PyTorch có thể kiểm soát kernel được hỗ trợ nào thực hiện phép toán, còn hướng dẫn các building block Transformer của PyTorch cung cấp bối cảnh triển khai rộng hơn. (docs.pytorch.org)
GQA phải được lựa chọn khi thiết kế hoặc điều chỉnh một model; nhìn chung, đây không phải là một tùy chọn chuyển đổi trong thời gian suy luận cho một checkpoint hiện có. Developer nên xác minh khả năng hỗ trợ của framework, tính chia hết giữa các head, hành vi số học, mức sử dụng bộ nhớ và chất lượng tác vụ. Đối với triển khai computer vision, các workflow bổ trợ như tích hợp Ultralytics TensorRT và chế độ benchmark của Ultralytics giúp đo lường liệu các tối ưu hóa về kiến trúc và runtime có mang lại cải thiện đáng kể trên phần cứng mục tiêu hay không.









