Chuyển tới nội dung chính

{/* Trang này được tạo tự động từ SKILL.md của kỹ năng bởi website/scripts/generate-skill-docs.py. Chỉnh sửa nguồn SKILL.md, không phải trang này. */}

Tối ưu hóa Flash chú ý

Tối ưu hóa sự chú ý của máy biến áp với Chú ý Flash để tăng tốc 2-4 lần và giảm bộ nhớ 10-20 lần. Sử dụng khi huấn luyện/chạy máy biến áp có chuỗi dài (>512 mã thông báo), gặp phải vấn đề về bộ nhớ GPU cần chú ý hoặc cần suy luận nhanh hơn. Hỗ trợ SDPA gốc PyTorch, thư viện flash-attn, H100 FP8 và chú ý đến cửa sổ trượt.

Siêu dữ liệu kỹ năng

NguồnTùy chọn — cài đặt với
`Hermes skills install official/mlops/flash-attention
`
Đường dẫn

optional-skills/mlops/flash-attention ` | | Phiên bản |

1.0.0 ` | | Tác giả | Nghiên cứu dàn nhạc | | Giấy phép | MIT | | Phụ thuộc |

flash-attn

, `torch

, transformers | | Nền tảng | Linux, macOS | | Thẻ |

Optimization

, `Flash Attention

, `Attention Optimization

, `Memory Efficiency

, `Speed Optimization

, `Long Context

, `PyTorch

, `SDPA

, `H100

, `FP8

, Transformers |

Tham khảo: đầy đủ SKILL.md

thông tin

Sau đây là định nghĩa kỹ năng đầy đủ mà Hermes tải khi kỹ năng này được kích hoạt. Đây là những gì tác nhân coi là hướng dẫn khi kỹ năng được kích hoạt.

Chú ý chớp nhoáng - Chú ý nhanh, hiệu quả

Bắt đầu nhanh

Flash Chú ý giúp tăng tốc 2-4 lần và giảm bộ nhớ 10-20 lần cho sự chú ý của máy biến áp thông qua tính năng xếp lớp và tính toán lại nhận biết IO.

PyTorch gốc (dễ nhất, PyTorch 2.2+):

`

 import torch
import torch.nn.functional as F`q = torch.randn(2, 8, 512, 64, device='cuda', dtype=torch.float16) # [batch, heads, seq, dim]
k = torch.randn(2, 8, 512, 64, device='cuda', dtype=torch.float16)
v = torch.randn(2, 8, 512, 64, device='cuda', dtype=torch.float16)

# Automatically uses Flash Attention if available
out = F.scaled_dot_product_attention(q, k, v)

`
``**thư viện flash-attn (nhiều tính năng hơn)**:

`
`bash
pip install flash-attn --no-build-isolation

`

`
`Python
from flash_attn import flash_attn_func

# q, k, v: [batch, seqlen, nheads, headdim]
out = flash_attn_func(q, k, v, dropout_p=0.0, causal=True)

`

## Quy trình công việc chung

### Quy trình làm việc 1: Kích hoạt trong mô hình PyTorch hiện có

Sao chép danh sách kiểm tra này:

`
Flash Attention Integration:

- [ ] Step 1: Check PyTorch version (2.2)
- [ ] Step 2: Enable Flash Attention backend
- [ ] Step 3: Verify speedup with profiling
- [ ] Step 4: Test accuracy matches baseline

`
``**Bước 1: Kiểm tra phiên bản PyTorch**

``` bash
Python -c "import torch; print(torch.__version__)"

# Should be ≥2.2.0

`
``Nếu <2.2, hãy nâng cấp:

`
``` bash
pip install --upgrade torch

`
``**Bước 2: Kích hoạt phụ trợ Flash Chú ý**

Thay thế sự chú ý tiêu chuẩn:

`
`Python

# Before (standard attention)
attn_weights = torch.softmax(q @ k.transpose(-2, -1) / math.sqrt(d_k), dim=-1)
out = attn_weights @ v

# After (Flash Attention)
import torch.nn.functional as F
out = F.scaled_dot_product_attention(q, k, v, attn_mask=mask)

`
``Phần phụ trợ buộc chú ý Flash:

`
``` python
with torch.backends.cuda.sdp_kernel(
enable_flash=True,
enable_math=False,
enable_mem_efficient=False
):
out = F.scaled_dot_product_attention(q, k, v)

`
``**Bước 3: Xác minh tốc độ tăng tốc bằng hồ sơ**

`Python
import torch.utils.benchmark as benchmark`def test_attention(use_flash):
q, k, v = [torch.randn(2, 8, 2048, 64, device='cuda', dtype=torch.float16) for _ in range(3)]

if use_flash:
with torch.backends.cuda.sdp_kernel(enable_flash=True):
return F.scaled_dot_product_attention(q, k, v)
else:
attn = (q @ k.transpose(-2, -1) / 8.0).softmax(dim=-1)
return attn @ v

# Benchmark
t_flash = benchmark.Timer(stmt='test_attention(True)', globals=globals())
t_standard = benchmark.Timer(stmt='test_attention(False)', globals=globals())

print(f"Flash: \{t_flash.timeit(100).mean:.3f}s")
print(f"Standard: \{t_standard.timeit(100).mean:.3f}s")

`
``Dự kiến: tăng tốc 2-4 lần cho các chuỗi >512 mã thông báo.

**Bước 4: Kiểm tra độ chính xác phù hợp với đường cơ sở**

`Python

# Compare outputs
q, k, v = [torch.randn(1, 8, 512, 64, device='cuda', dtype=torch.float16) for _ in range(3)]

# Flash Attention
out_flash = F.scaled_dot_product_attention(q, k, v)

# Standard attention
attn_weights = torch.softmax(q @ k.transpose(-2, -1) / 8.0, dim=-1)
out_standard = attn_weights @ v

# Check difference
diff = (out_flash - out_standard).abs().max()
print(f"Max difference: \{diff:.6f}")
# Should be <1e-3 for float16

`

### Workflow 2: Sử dụng thư viện flash-attn cho các tính năng nâng cao

Để chú ý đến nhiều truy vấn, cửa sổ trượt hoặc H100 FP8.

Sao chép danh sách kiểm tra này:

`
flash-attn Library Setup:
- [ ] Step 1: Install flash-attn library
- [ ] Step 2: Modify attention code
- [ ] Step 3: Enable advanced features
- [ ] Step 4: Benchmark performance

`
``**Bước 1: Cài đặt thư viện flash-attn**

``` bash

# NVIDIA GPUs (CUDA 12.0+)
pip install flash-attn --no-build-isolation

# Verify installation
Python -c "from flash_attn import flash_attn_func; print('Success')"

`
``**Bước 2: Sửa đổi mã chú ý**

``` python
from flash_attn import flash_attn_func

# Input: [batch_size, seq_len, num_heads, head_dim]

# Transpose from [batch, heads, seq, dim] if needed
q = q.transpose(1, 2) # [batch, seq, heads, dim]
k = k.transpose(1, 2)
v = v.transpose(1, 2)

out = flash_attn_func(
q, k, v,
dropout_p=0.1,
causal=True, # For autoregressive models
window_size=(-1, -1), # No sliding window
softmax_scale=None # Auto-scale
)

out = out.transpose(1, 2) # Back to [batch, heads, seq, dim]

`
``**Bước 3: Kích hoạt các tính năng nâng cao**

Chú ý nhiều truy vấn (K/V được chia sẻ giữa các đầu):

`
``` python
from flash_attn import flash_attn_func

# q: [batch, seq, num_q_heads, dim]

# k, v: [batch, seq, num_kv_heads, dim] # Fewer KV heads
out = flash_attn_func(q, k, v) # Automatically handles MQA

`
``Chú ý cửa sổ trượt (chú ý cục bộ):

`
``` python

# Only attend to window of 256 tokens before/after
out = flash_attn_func(
q, k, v,
window_size=(256, 256), # (left, right) window
causal=True
)

`
``**Bước 4: Hiệu suất điểm chuẩn**

``` python
import torch
from flash_attn import flash_attn_func
import time`q, k, v = [torch.randn(4, 4096, 32, 64, device='cuda', dtype=torch.float16) for _ in range(3)]

# Warmup
for _ in range(10):
_ = flash_attn_func(q, k, v)

# Benchmark
torch.cuda.synchronize()
start = time.time()
for _ in range(100):
out = flash_attn_func(q, k, v)
torch.cuda.synchronize()
end = time.time()

print(f"Time per iteration: \{(end-start)/100*1000:.2f}ms")
print(f"Memory allocated: \{torch.cuda.max_memory_allocated()/1e9:.2f}GB")

`

### Workflow 3: Tối ưu hóa H100 FP8 (FlashAttention-3)

Để có hiệu suất tối đa trên GPU H100.

`
FP8 Setup:

- [ ] Step 1: Verify H100 GPU available
- [ ] Step 2: Install flash-attn with FP8 support
- [ ] Step 3: Convert inputs to FP8
- [ ] Step 4: Run with FP8 attention

`
``**Bước 1: Xác minh GPU H100**

``` bash
nvidia-smi --query-gpu=name --format=csv

# Should show "H100" or "H800"

`
``**Bước 2: Cài đặt flash-attn có hỗ trợ FP8**

``` bash
pip install flash-attn --no-build-isolation

# FP8 support included for H100

`
``**Bước 3: Chuyển đổi đầu vào sang FP8**

``` python
import torch`q = torch.randn(2, 4096, 32, 64, device='cuda', dtype=torch.float16)
k = torch.randn(2, 4096, 32, 64, device='cuda', dtype=torch.float16)
v = torch.randn(2, 4096, 32, 64, device='cuda', dtype=torch.float16)

# Convert to float8_e4m3 (FP8)
q_fp8 = q.to(torch.float8_e4m3fn)
k_fp8 = k.to(torch.float8_e4m3fn)
v_fp8 = v.to(torch.float8_e4m3fn)

`
``**Bước 4: Chạy với sự chú ý của FP8**

`Python
from flash_attn import flash_attn_func

# FlashAttention-3 automatically uses FP8 kernels on H100
out = flash_attn_func(q_fp8, k_fp8, v_fp8)

# Result: ~1.2 PFLOPS, 1.5-2x faster than FP16

`

## Khi nào nên sử dụng so với các lựa chọn thay thế`**Sử dụng Flash Chú ý khi:**
- Huấn luyện máy biến áp có trình tự >512 token
- Chạy suy luận với ngữ cảnh dài (>2K token)
- Bộ nhớ GPU bị hạn chế (OOM với sự chú ý tiêu chuẩn)
- Cần tăng tốc 2-4 lần mà không mất độ chính xác
- Sử dụng PyTorch 2.2+ hoặc có thể cài đặt flash-attn`**Sử dụng các lựa chọn thay thế thay thế:**
- **Sự chú ý tiêu chuẩn**: Chuỗi <256 mã thông báo (chi phí không đáng có)
- **xFormers**: Cần nhiều biến thể được chú ý hơn (không chỉ tốc độ)
- **Chú ý tiết kiệm bộ nhớ**: Suy luận CPU (Flash Chú ý cần GPU)

## Các vấn đề thường gặp`**Vấn đề: Lỗi nhập: không thể nhập flash_attn**

Cài đặt với cờ cách ly không xây dựng:

`
``` bash
pip install flash-attn --no-build-isolation

`
``Hoặc cài đặt bộ công cụ CUDA trước:

`
`bash
conda install cuda -c nvidia
pip install flash-attn --no-build-isolation

`
``**Vấn đề: Chậm hơn dự kiến (không tăng tốc)**

Lợi ích của Flash Chú ý tăng theo độ dài chuỗi:

- <512 token: Tăng tốc tối thiểu (10-20%)
- 512-2K token: tăng tốc 2-3 lần
- >2K token: tăng tốc 3-4 lần

Kiểm tra độ dài chuỗi là đủ.

**Vấn đề: RuntimeError: Lỗi CUDA**

Xác minh GPU hỗ trợ Flash Chú ý:

`
``` python
import torch
print(torch.cuda.get_device_capability())

# Should be ≥(7, 5) for Turing+

`
``Chú ý Flash yêu cầu:
- Ampe (A100, A10): ✅ Hỗ trợ đầy đủ
- Turing (T4): ✅ Được hỗ trợ
- Volta (V100): ❌ Không hỗ trợ`**Vấn đề: Suy giảm độ chính xác**

Kiểm tra dtype là float16 hay bfloat16 (không phải float32):

`
``` python
q = q.to(torch.float16) # Or torch.bfloat16

`
``Flash Chú ý sử dụng float16/bfloat16 cho tốc độ. Float32 không được hỗ trợ.

## Chủ đề nâng cao**Tích hợp với HuggingFace Transformers**: Xem [references/transformers-integration.md](https://GitHub.com/NousResearch/Hermes-agent/blob/main/optional-skills/mlops/flash-attention/references/transformers-integration.md) để bật tính năng Chú ý Flash trong các kiểu máy BERT, GPT, Llama.

**Điểm chuẩn hiệu suất**: Xem [references/benchmarks.md](https://GitHub.com/NousResearch/Hermes-agent/blob/main/optional-skills/mlops/flash-attention/references/benchmarks.md) để biết so sánh chi tiết về tốc độ và bộ nhớ giữa các GPU và độ dài chuỗi.

## Yêu cầu về phần cứng
- **GPU**: NVIDIA Ampere+ (A100, A10, A30) hoặc AMD MI200+

- **VRAM**: Giống như chú ý tiêu chuẩn (Flash Chú ý không làm tăng bộ nhớ)
- **CUDA**: 12.0+ (tối thiểu 11.8)
- **PyTorch**: 2.2+ dành cho hỗ trợ gốc`**Không được hỗ trợ**: V100 (Volta), suy luận CPU

## Tài nguyên
- Bài viết: "FlashAttention: Chú ý chính xác nhanh và hiệu quả về bộ nhớ với IO-Awareness" (NeurIPS 2022)
- Bài viết: "FlashAttention-2: Chú ý nhanh hơn với khả năng song song tốt hơn và phân vùng công việc" (ICLR 2024)
- Blog: https://tridao.me/blog/2024/flash3/
- GitHub: https://GitHub.com/Dao-AILab/flash-attention
- Tài liệu PyTorch: https://pytorch.org/docs/stable/generated/torch.nn.professional.scaled_dot_product_attention.html