๊ฐ์ ๊ณผ ๋ ์ฐฝ์ฑ, ํต์ฌ ์๊ณ ๋ฆฌ์ฆ, ๊ทธ๋ฆฌ๊ณ ํ๊ณ
๋ฌผ์ด๋ณธ ๊ฒ: ์ด ๋ ผ๋ฌธ์ ์์ฃผ ์์ธํ๊ฒ ์ฝ๊ณ ๋ ผ๋ฌธ์ ๊ฐ์ ๊ณผ ๋ ์ฐฝ์ ์ธ ์ง์ ์ ์ค๋ช ํด์ฃผ๊ณ ํต์ฌ ์๊ณ ๋ฆฌ์ฆ์ ์์ ์ ๋ ฅ์ ๋ค์ด์ ์ ์ฒด์ ์ธ ๊ณผ์ ์ ์ค๋ช ํด์ค ์ถ๊ฐ์ ์ผ๋ก ๋ ผ๋ฌธ์ ํ๊ณ์ ์ ๋ํด์๋ ์๋ ค์ค
๐ ๋ ผ๋ฌธ ์์ฝ ๋ฐ ๋ถ์ (FlexAttention: A Programming Model for Generating Optimized Attention Kernels)
โ ๋ ผ๋ฌธ์ ๊ฐ์ ๋ฐ ๋ ์ฐฝ์ ์ธ ์ง์
์ ์ฐ์ฑ (Flexibility)
- ๊ธฐ์กด์ FlashAttention๊ณผ ๋ฌ๋ฆฌ, FlexAttention์ ๋ค์ํ attention variants๋ฅผ ๊ฐ๋จํ PyTorch ์ฝ๋๋ก ๊ตฌํํ ์ ์๊ฒ ํจ.
- ๋ค์ํ attention ๋ฐฉ์ (e.g., Alibi, Document Masking, Sliding Window, PrefixLM, PagedAttention ๋ฑ)์ ์์ฝ๊ฒ ๊ตฌ์ฑํ๊ณ ์กฐํฉ ๊ฐ๋ฅ.
- ์๋ก์ด attention ๋ฐฉ์์ด ํ์ํ ๋๋ง๋ค ์ปค์คํ ์ปค๋์ ์์ฑํ ํ์ ์์ด ๊ฐ๋จํ ์์ ๊ฐ๋ฅ.
์ฑ๋ฅ ๊ฐ์ (Performance Improvement)
- FlashAttention ๋๋น ์ต๋ 1.43๋ฐฐ ํฅ์๋ ์ฑ๋ฅ์ ์ ๊ณตํ๋ฉฐ, ํนํ ์ง์๋์ง ์๋ attention variants์ ๋ํด์๋ ์ต๋ 8๋ฐฐ๊น์ง ๋น ๋ฅด๊ฒ ๋์.
- Inference ์ฑ๋ฅ์์ ๊ธฐ์กด FlashAttention ๋๋น 1.45๋ฐฐ์ ์ฑ๋ฅ์ ๋ณด์ฌ์ค.
Paged Attention ์ง์
- ๊ธฐ์กด์ PagedAttention ๋ฐฉ์์ ๋์ ๋ฉ๋ชจ๋ฆฌ ์ฌ์ฉ์ ํด๊ฒฐํ๋ฉฐ, ๋ค์ํ attention variants๋ฅผ ์ฝ๊ฒ ๊ตฌํํ ์ ์๋๋ก ํจ.
- GPU ๋ฉ๋ชจ๋ฆฌ์ ๊ฐ์ ์ ๊ทผ ๋ฐฉ์์ ํตํด ๋ฉ๋ชจ๋ฆฌ ์ฌ์ฉ๋์ ์ต์ ํํ๊ณ , ์ปค๋์ ์์ ํ์ง ์์ผ๋ฉด์๋ ์ฑ๋ฅ์ ์ ์งํจ.
Block Sparsity ํ์ฉ
- Sparsity๋ฅผ ํจ๊ณผ์ ์ผ๋ก ํ์ฉํ๊ธฐ ์ํด BlockMask๋ผ๋ ์๋ก์ด ๋ฐ์ดํฐ ๊ตฌ์กฐ๋ฅผ ๋์ .
- ๋ฉ๋ชจ๋ฆฌ ์ฌ์ฉ์ ์ค์ด๊ณ ์ฐ์ฐ์ ์ต์ ํํ๋ ๋ฐ ๊ธฐ์ฌํ๋ฉฐ, ์ ๋ฐ์ ์ธ ์ฑ๋ฅ์ ํฅ์์ํด.
๐ ํต์ฌ ์๊ณ ๋ฆฌ์ฆ ์ค๋ช (์์ ํฌํจ)
๊ธฐ์กด Attention Mechanism
- Self-Attention์ ๊ธฐ๋ณธ ๊ณต์: \[ S = \text{softmax} \left( \frac{QK^T}{\sqrt{d_k}} \right) \] \[ \text{Attention}(Q, K, V) = SV \]
FlexAttention์ ๋ณํ
๊ธฐ์กด Attention ๋ฐฉ์์ score matrix \( S \)๋ฅผ ๋ค์ํ ๋ฐฉ์์ผ๋ก ์์ ํ ์ ์๋๋ก ํจ.
\[ \text{FlexAttention}(Q, K, V) = \text{softmax} \left( \text{mod} \left( \frac{QK^T}{\sqrt{d_k}} \right) \right) V \]Score matrix์ ๋ ๊ฐ์ง modification ๋ฐฉ์์ ์ถ๊ฐ:
- score mod: ์ ์ ์์ฒด๋ฅผ ์์ ํ๋ ํจ์.
- mask mod: ํน์ ์์น๋ฅผ -โ๋ก ์ค์ ํ๋ ํจ์.
์์: Sliding Window Attention ๊ตฌํ
def sliding_window_mask(q_idx, kv_idx, window_size):
return abs(q_idx - kv_idx) <= window_size- ์ด ํจ์๋ Query์ Key๊ฐ ํน์ ์๋์ฐ ํฌ๊ธฐ ์์ ์์ ๋๋ง Attention์ ํ์ฉํ๋ค.
- ์๋ฅผ ๋ค์ด, \( \text{window\_size} = 3 \)์ผ ๊ฒฝ์ฐ, Query index๊ฐ 5์ผ ๋ Key index๋ [2, 3, 4, 5, 6, 7, 8] ๋ฒ์ ๋ด์ ๊ฐ๋ค๋ง์ ํ์ฉ.
์์: Alibi Bias ๊ตฌํ
def alibi_bias(score, q_idx, kv_idx, head_idx, bias_factor):
return score + bias_factor * (q_idx - kv_idx)- Query์ Key์ ์๋์ ์ธ ์์น์ ๋ฐ๋ผ ๊ฐ์ค์น๋ฅผ ์กฐ์ ํ์ฌ ๋ฉ๋ฆฌ ์๋ Token์ ๋ ์ ๊ฒ ๋ฐ์ํ๋๋ก ํจ.
- ์๋ฅผ ๋ค์ด, \( \text{bias\_factor} = 0.5 \)์ผ ๋, Query index๊ฐ 10์ด๊ณ Key index๊ฐ 5๋ผ๋ฉด, score๋ \( 0.5 \times (10 - 5) = 2.5 \) ๋งํผ ์ฆ๊ฐ.
โ ๋ ผ๋ฌธ์ ํ๊ณ์
ํน์ Attention Variant์ ๋ํ ์ต์ ํ ๋ถ์ฌ
- ์ผ๋ถ Attention Variant (์: Neighborhood Attention)์ ๋ํด์๋ ์์ ํ ์ต์ ํ๊ฐ ์ด๋ฃจ์ด์ง์ง ์์์.
- ํนํ, NA์ ๊ฒฝ์ฐ ๊ธฐ์กด ์ปค๋๋ณด๋ค ์ฑ๋ฅ์ด ๋จ์ด์ง๋ ๊ฒฝ์ฐ๋ ์กด์ฌ.
๊ณ ์ ๋ Block Size
- Block Sparsity Optimization์ ๊ฒฝ์ฐ, ๋ธ๋ก ํฌ๊ธฐ๋ฅผ ์กฐ์ ํ์ง ์์ผ๋ฉด ๋ฉ๋ชจ๋ฆฌ ํจ์จ์ฑ์ ์ต๋ํํ์ง ๋ชปํจ.
- ํนํ ๊ธด ๋ฌธ์ฅ์ด๋ ๋ณต์กํ Mask ๊ตฌ์กฐ๋ฅผ ๋ค๋ฃฐ ๋๋ ์ถ๊ฐ์ ์ธ ์ต์ ํ๊ฐ ํ์.
Inference ๊ณผ์ ์ ์ ์ฝ
- ์ผ๋ถ Attention Variant๋ inference ์์ ์์ ๋์ ๋ฐฉ์์ด ๋ณ๊ฒฝ๋์ด์ผ ํ๋๋ฐ, ์ด๋ฅผ ์ํ ์๋ ๋ณํ ๊ธฐ๋ฅ์ด ๋ฏธํกํ ์ ์์.
๐ FlashAttention๊ณผ FlexAttention ๋น๊ต
| ํน์ง | FlashAttention | FlexAttention |
|---|---|---|
| ์ ์ฐ์ฑ | ์ ํ์ (์ง์๋๋ ์ปค๋๋ง ๊ฐ๋ฅ) | ๋งค์ฐ ์ ์ฐ (๊ฐ๋จํ PyTorch ์ฝ๋๋ก ๊ตฌํ ๊ฐ๋ฅ) |
| ์ฑ๋ฅ | ๋ฐ์ด๋จ | FlashAttention๊ณผ ๋น์ทํ๊ฑฐ๋ ๋ ๋์ (1.43๋ฐฐ ํฅ์) |
| ์๋ก์ด Attention ๋ฐฉ์ ์ง์ | ์ ํ์ | ์์ ๋กญ๊ฒ ์ถ๊ฐ ๋ฐ ์กฐํฉ ๊ฐ๋ฅ |
| Block Sparsity ํ์ฉ | ๋ถ๋ถ์ ํ์ฉ | ์์ ํ์ฉ (BlockMask ๋์ ) |
| Paged Attention ์ง์ | ์ ํ์ | ๊ฐ์ ๋ฉ๋ชจ๋ฆฌ ์ ๊ทผ ๋ฐฉ์์ผ๋ก ํจ์จ์ ์ง์ |
FlexAttention์ ๊ธฐ์กด FlashAttention์ ๋จ์ ์ ๋ณด์ํ๋ฉฐ, ์ ์ฐ์ฑ๊ณผ ์ฑ๋ฅ์ ๋์์ ์ ๊ณตํ๋ ๋งค์ฐ ๊ฐ๋ ฅํ ๋ชจ๋ธ์ ๋๋ค. ๋ค์ํ Attention ๋ฐฉ์์ ์กฐํฉ๊ณผ ์ต์ ํ๋ฅผ ํตํด ๋ ๋ง์ ์์ฉ์ด ๊ฐ๋ฅํ๋ฉฐ, ํนํ ์ฐ๊ตฌ์๊ฐ ์๋ก์ด Attention ๋ฐฉ์์ ํ๊ตฌํ๋ ๋ฐ ์์ด ์ค์ํ ๋๊ตฌ๊ฐ ๋ ์ ์์ต๋๋ค.
ํต์ฌ ์๊ณ ๋ฆฌ์ฆ์ ์์๋ก ๋ฐ๋ผ๊ฐ๊ธฐ
๋ฌผ์ด๋ณธ ๊ฒ: ๋ ผ๋ฌธ์ ํต์ฌ์๊ณ ๋ฆฌ์ฆ์ ์์๋ฅผ ๋ค์ด ์์ฃผ ์์ธํ๊ฒ ์ค๋ช ํด์ฃผ๊ณ ์ ๋ฆฌํด ์ค
๐ ๋ ผ๋ฌธ์ ํต์ฌ ์๊ณ ๋ฆฌ์ฆ: FlexAttention
FlexAttention์ ๋ค์ํ Attention Variant๋ฅผ ์์ฝ๊ฒ ์ ์ํ๊ณ ์ต์ ํ๋ ์ปค๋๋ก ์ปดํ์ผํ๋ ์๋ก์ด ํ๋ก๊ทธ๋๋ฐ ๋ชจ๋ธ์ ๋๋ค. ํต์ฌ ์์ด๋์ด๋ ๊ธฐ์กด์ Attention ์ฐ์ฐ์ ๋ ๊ฐ์ง ๊ฐ๋ ์ผ๋ก ๋ถ๋ฆฌํ์ฌ ์กฐ์ ํ ์ ์๊ฒ ํ๋ ๊ฒ์ ๋๋ค.
ํต์ฌ ์๊ณ ๋ฆฌ์ฆ ๊ตฌ์ฑ ์์
- Score Modification (
score_mod) - Mask Modification (
mask_mod) - Block Mask Optimization
- Template-based Kernel Generation
๐ 1. Score Modification (score_mod)
score_mod๋ Attention ์ ์๋ฅผ ์กฐ์ ํ๋ ํจ์๋ก, ๊ธฐ์กด์ ์ ์ ํ๋ ฌ์ ์ถ๊ฐ์ ์ธ ์์ ์ฐ์ฐ์ ์ํํ ์ ์์ต๋๋ค.
โ ์์: Alibi Bias ์ ์ฉ
๋ฌธ์ ์ ์: ๋ชจ๋ธ์ด ๊ธด ๋ฌธ์ฅ์ ์ ์ฒ๋ฆฌํ ์ ์๋๋ก, Query์ Key์ ๊ฑฐ๋ฆฌ ์ฐจ์ด์ ๋น๋กํ์ฌ ์ ์๋ฅผ ์กฐ์ .
๊ณต์:
\[ \text{Modified Score} = \text{Original Score} + \text{bias\_factor} \times (\text{q\_idx} - \text{kv\_idx}) \]์ฝ๋ ๊ตฌํ:
PYTHONdef alibi_bias(score, q_idx, kv_idx, head_idx, bias_factor=0.5): return score + bias_factor * (q_idx - kv_idx)def alibi_bias(score, q_idx, kv_idx, head_idx, bias_factor=0.5): return score + bias_factor * (q_idx - kv_idx)์์ ์ ๋ ฅ ๊ฐ:
PLAINTEXTOriginal Score Matrix (S): [[1.0, 0.8, 0.5], [0.9, 1.0, 0.7], [0.6, 0.7, 1.0]] q_idx = 2, kv_idx = 0, head_idx = 0, bias_factor = 0.5Original Score Matrix (S): [[1.0, 0.8, 0.5], [0.9, 1.0, 0.7], [0.6, 0.7, 1.0]] q_idx = 2, kv_idx = 0, head_idx = 0, bias_factor = 0.5์ถ๋ ฅ ๊ฐ (์์ ๋ ์ ์ ํ๋ ฌ):
PLAINTEXTModified Score Matrix (S'): [[1.0, 0.8, 0.5], [1.4, 1.0, 0.7], [1.6, 1.7, 1.0]]Modified Score Matrix (S'): [[1.0, 0.8, 0.5], [1.4, 1.0, 0.7], [1.6, 1.7, 1.0]]-> \( (q\_idx - kv\_idx) = 2 \), ๋ฐ๋ผ์ \( \text{bias} = 0.5 \times 2 = 1.0 \).
๐ 2. Mask Modification (mask_mod)
mask_mod๋ Attention ์ ์๋ฅผ ๋ง์คํนํ์ฌ ํน์ ์์น์ ์ฐ์ฐ์ ๋ฌด์ํ ์ ์๋๋ก ํ๋ ํจ์์
๋๋ค.
โ ์์: Sliding Window Mask
๋ฌธ์ ์ ์: Query๊ฐ Key์ ์ผ์ ๋ฒ์ ๋ด์์๋ง Attention์ ํ ์ ์๋๋ก ์ ํ.
๊ณต์:
\[ \text{mask\_mod}(q\_idx, kv\_idx) = \begin{cases} \text{True} & \text{if } |q\_idx - kv\_idx| \leq \text{window\_size} \\ \text{False} & \text{otherwise} \end{cases} \]์ฝ๋ ๊ตฌํ:
PYTHONdef sliding_window_mask(q_idx, kv_idx, window_size=3): return abs(q_idx - kv_idx) <= window_sizedef sliding_window_mask(q_idx, kv_idx, window_size=3): return abs(q_idx - kv_idx) <= window_size์์ ์ ๋ ฅ ๊ฐ:
PLAINTEXTq_idx = 5 Key Indices = [2, 3, 4, 5, 6, 7, 8] window_size = 3q_idx = 5 Key Indices = [2, 3, 4, 5, 6, 7, 8] window_size = 3์ถ๋ ฅ ๊ฐ:
PLAINTEXTMasked Keys = [4, 5, 6]Masked Keys = [4, 5, 6]-> Query index๊ฐ 5์ผ ๋, Key index๋ 2~8 ์ฌ์ด์์ 4, 5, 6 ๋ง ํ์ฉ๋จ.
๐ 3. Block Mask Optimization
Block Sparsity๋ฅผ ํ์ฉํ์ฌ ๋ฉ๋ชจ๋ฆฌ ๋ฐ ์ฐ์ฐ ํจ์จ์ ๊ทน๋ํํ๋ ๋ฐฉ๋ฒ์ ๋๋ค. Mask๋ฅผ ์ ์ฉํ ๋ ๋ธ๋ก ๋จ์๋ก ์ฒ๋ฆฌํ์ฌ ์ ์ฒด ๋ฉ๋ชจ๋ฆฌ ์ ๊ทผ์ ์ค์ด๋ ๋ฐฉ์์ ๋๋ค.
โ ํต์ฌ ์์ด๋์ด:
- Attention ์ฐ์ฐ์ ๋ธ๋ก ๋จ์๋ก ๋ถ๋ฆฌํ์ฌ ์ฐ์ฐ.
- ์์ ํ ๋ง์คํน๋ ๋ธ๋ก์ ์ฐ์ฐํ์ง ์๊ณ ๊ฑด๋๋.
- ๋ถ๋ถ์ ์ผ๋ก ๋ง์คํน๋ ๋ธ๋ก์ Masking ์ฐ์ฐ๋ง ์ํ.
โ ์์: Sliding Window Attention์ Block Mask ์ ์ฉ
์ ๋ ฅ ํ๋ ฌ:
PLAINTEXTQ_LEN = 6, KV_LEN = 6 Block Size = 2 x 2 Sliding Window Size = 1Q_LEN = 6, KV_LEN = 6 Block Size = 2 x 2 Sliding Window Size = 1๋ธ๋ก ๊ตฌ์ฑ:
PLAINTEXT๋ธ๋ก 1: Q[0:2], K[0:2] ๋ธ๋ก 2: Q[0:2], K[2:4] ๋ธ๋ก 3: Q[0:2], K[4:6] ๋ธ๋ก 4: Q[2:4], K[0:2] ๋ธ๋ก 5: Q[2:4], K[2:4] ๋ธ๋ก 6: Q[2:4], K[4:6]๋ธ๋ก 1: Q[0:2], K[0:2] ๋ธ๋ก 2: Q[0:2], K[2:4] ๋ธ๋ก 3: Q[0:2], K[4:6] ๋ธ๋ก 4: Q[2:4], K[0:2] ๋ธ๋ก 5: Q[2:4], K[2:4] ๋ธ๋ก 6: Q[2:4], K[4:6]๋ง์คํน ์ ์ฉ ๊ฒฐ๊ณผ:
PLAINTEXTBlock 1: Computed Block 2: Computed Block 3: Ignored (Fully Masked) Block 4: Ignored (Fully Masked) Block 5: Computed Block 6: ComputedBlock 1: Computed Block 2: Computed Block 3: Ignored (Fully Masked) Block 4: Ignored (Fully Masked) Block 5: Computed Block 6: Computed-> ์ ์ฒด ๋ฉ๋ชจ๋ฆฌ ์ ๊ทผ์ ์ค์ฌ ์ฐ์ฐ ํจ์จ์ฑ์ ๊ทน๋ํ.
๐ 4. Template-based Kernel Generation
PyTorch์ torch.compile์ ์ด์ฉํด ์ฌ์ฉ์๊ฐ ์ ์ํ score_mod์ mask_mod๋ฅผ ์ปดํ์ผํ์ฌ ์ต์ ํ๋ ์ปค๋๋ก ๋ณํํฉ๋๋ค.
- ๊ณผ์ :
- PyTorch ํจ์๋ก ์ฌ์ฉ์ ์ ์ ๋ชจ๋ธ ์์ฑ (์: Alibi Bias, Sliding Window Mask).
torch.compile()์ ์ด์ฉํ์ฌ ์ฐ์ฐ ๊ทธ๋ํ ์ถ์ถ.- Triton ํ ํ๋ฆฟ์ ์ฌ์ฉํ์ฌ GPU์์ ํจ์จ์ ์ผ๋ก ๋์ํ๋๋ก ์ฝ๋ ์์ฑ.
- ์ปค๋์ PyTorch API์ ํตํฉํ์ฌ ์ฌ์ฉ ๊ฐ๋ฅ.
๐ ์ ๋ฆฌ
| ํน์ง | ์ค๋ช | ๊ตฌํ ๋ฐฉ์ |
|---|---|---|
score_mod | Attention ์ ์๋ฅผ ์กฐ์ ํ๋ ํจ์ | ์: Alibi Bias |
mask_mod | ํน์ ์์น์ ์ ์๋ฅผ ๋ง์คํนํ์ฌ ๋ฌด์ | ์: Sliding Window Mask |
| Block Mask Optimization | ๋ฉ๋ชจ๋ฆฌ ์ฌ์ฉ๋ ์ค์ด๊ณ ํจ์จ์ ์ผ๋ก ์ฐ์ฐ | ๋ธ๋ก ๋จ์๋ก ๊ณ์ฐ |
| Kernel Generation | PyTorch์ Triton์ ์ด์ฉํ ์ปค๋ ์ต์ ํ | torch.compile() ์ฌ์ฉ |
์คํ ๊ฒฐ๊ณผ์ ๋ค๋ฅธ ๋ฐฉ๋ฒ๋ก ๊ณผ์ ์ฐจ์ด
๋ฌผ์ด๋ณธ ๊ฒ: ์ด ๋ ผ๋ฌธ์์ ์ ์ํ ๊ฒฐ๊ณผ๋ฅผ ์์ธํ๊ฒ ๋ณด๊ณ ๋ค๋ฅธ ๋ฐฉ๋ฒ๋ก ์ ๋นํ์ฌ ํน์ถ๋ ์ ๊ณผ ๋ ผ๋ฌธ์์ ์ ๊ธฐํ๋ ์ด๋ ํ ๋ฐฉ๋ฒ์ด ์ด๋ฌํ ๊ฒฐ๊ณผ๋ฅผ ๋์ถํ๊ฒ ๋์๋์ง ๋ ผ๋ฌธ์์ ์ ์ํ๋ ์ด์ ์ ๋์ ์๊ฐ์ ์๋ ค์ค
๐ ๋ ผ๋ฌธ์์ ์ ์ํ ๊ฒฐ๊ณผ (Evaluation)
๋ ผ๋ฌธ์์๋ FlexAttention์ ์ฑ๋ฅ์ 7๊ฐ์ ์ฃผ์ Attention Variant์ ๋ํด ๋ค์ํ ์ธก๋ฉด์์ ํ๊ฐํ์์ต๋๋ค. ํนํ ๊ธฐ์กด FlashAttention (FAv2, FAv3), PyTorch์ Scale Dot Product Attention (SDPA)๊ณผ ๋น๊ตํ์ฌ ์ฑ๋ฅ์ ๋ถ์ํฉ๋๋ค.
๐ 1. Attention Kernel Performance (Attention Kernel ์ฑ๋ฅ)
โ ํ๊ฐ ๋์ Attention Variants
- Noop: ๊ธฐ๋ณธ Attention (๋ณํ ์์)
- Causal: ๊ธฐ์กด์ Causal Masking (์ด์ Token์๋ง Attention)
- Alibi Bias: ์๋์ ์์น ๊ธฐ๋ฐ Bias๋ฅผ ์ถ๊ฐํ๋ Attention
- Sliding Window: ์ผ์ ๋ฒ์ ๋ด์ Token๋ง Attention
- PrefixLM: ์ผ๋ถ Token์ Bidirectional, ์ดํ๋ Causal๋ก ๊ตฌ์ฑ
- Soft Cap: Logits์ ์ฑ์ฅ์ ์ ํํ๋ ๋ฐฉ์ (tanh ํจ์ ์ฌ์ฉ)
- Document Masking: ์๋ก ๋ค๋ฅธ ๋ฌธ์์ Token์ ๊ตฌ๋ถํ์ฌ Attention
โ ์ฑ๋ฅ ๊ฒฐ๊ณผ
| ๋ชจ๋ธ | ์๋ (FAv2 ๋๋น) | ์๋ (FAv3 ๋๋น) | ์๋ (FAKV ๋๋น) | ํน์ด์ |
|---|---|---|---|---|
| Noop (๊ธฐ๋ณธ) | 1.00x - 1.22x | 1.43x | 1.45x | ๊ธฐ์กด ๋ฐฉ๋ฒ๋ก ๊ณผ ์ ์ฌ |
| Causal | 1.00x - 1.22x | 1.43x | 1.45x | ๊ธฐ์กด ๋ฐฉ๋ฒ๋ก ๊ณผ ์ ์ฌ |
| Alibi Bias | 1.43x | 1.45x | 5.37x | FAKV์ ๊ฒฝ์ฐ ์ต์ ํ ๋ถ์ฌ๋ก FlexAttention์ด ์๋์ ์ผ๋ก ๋น ๋ฆ |
| Sliding Window | 1.43x | 1.45x | 1.45x | ๋๋ถ๋ถ์ ๊ฒฝ์ฐ ๋ฐ์ด๋จ |
| PrefixLM | 1.43x | 1.45x | 1.45x | ๊ธฐ์กด ๋ฐฉ๋ฒ๋ก ๋๋น ๋ ๋์ ์ ์ฐ์ฑ |
| Soft Cap | 1.43x | 1.45x | 1.45x | ํน์ ํ๊ฒฝ์์ ๋ ์ฐ์ |
| Document Masking | 1.43x | 1.45x | 1.45x | ๋ค์ํ Mask ์ง์ ๊ฐ๋ฅ |
๐ 2. End-to-end Performance (Inference & Training Performance)
โ Inference Performance (GPT-Fast, LLaMa3.1 ๋ชจ๋ธ)
- FlexAttention์ ์ฌ์ฉํ์ฌ ๊ธฐ์กด SDPA ๋๋น 1.22x - 2.04x์ ์๋ ํฅ์์ ๋ฌ์ฑ.
- ๊ธด ๋ฌธ์ฅ์ผ์๋ก ์ฑ๋ฅ ๊ฐ์ ์ด ๋ ๋๋๋ฌ์ง (ํนํ 70B ๋ชจ๋ธ์ ๊ฒฝ์ฐ ์ต๋ 1.66x ํฅ์).
โ Training Performance (Torchtune, LLaMa3 ๋ชจ๋ธ)
- ๋ค์ํ ์ํ์ค ๊ธธ์ด์ ๋ํด 2.4x๊น์ง ์ฑ๋ฅ ํฅ์.
- ๊ธฐ์กด SDPA ๋๋น ๋ฉ๋ชจ๋ฆฌ ์ ๊ทผ์ ์ค์ด๊ณ ์ฐ์ฐ ํจ์จ์ฑ์ ๋์.
- ๋ฌธ์ ๋จ์ Masking์ ํจ์จ์ ์ผ๋ก ์ฒ๋ฆฌํจ์ผ๋ก์จ ๋์ ํ์ต ์๋๋ฅผ ๋ณด์ฌ์ค.
๐ 3. Paged Attention Performance
โ PagedAttention ํ์ฉ ๊ฒฐ๊ณผ
- FlexAttention์ ๊ธฐ์กด FlashAttention๋ณด๋ค Paged Attention์ ๋ ํจ๊ณผ์ ์ผ๋ก ์ง์ํจ.
- Paged Attention ๋์ ์ ์ฑ๋ฅ ์ ํ๊ฐ ๊ฑฐ์ ์๊ณ ์คํ๋ ค ํน์ ์ํฉ์์๋ ๋ ๋์ ์ฑ๋ฅ์ ๋ณด์ฌ์ค.
- ๊ธฐ์กด FlashAttention ๊ธฐ๋ฐ์ PagedAttention์ 20~26%์ ์ฑ๋ฅ ์ ํ๊ฐ ๋ฐ์ํ์์ง๋ง, FlexAttention์์๋ 1% ๋ฏธ๋ง์ ์ฑ๋ฅ ์ ํ๋ง ๋ฐ์ํจ.
๐ก FlexAttention์ด ๋ ๋ฐ์ด๋ ์ด์ ์ ๋ฐฉ๋ฒ๋ก (๋ ผ๋ฌธ์์ ์ ์ํ๋ ์ด์ )
โ 1. Unified Programming Model (ํตํฉ ํ๋ก๊ทธ๋๋ฐ ๋ชจ๋ธ)
- ๊ธฐ์กด FlashAttention์ ํน์ Attention Variant๋ง์ ์ง์ํ๋ฉฐ, ์๋ก์ด ๋ณํ์ ์ถ๊ฐํ๊ธฐ ์ํด ์ปค๋์ ์์ ํ๊ฑฐ๋ ์๋ก ์์ฑํด์ผ ํจ.
- FlexAttention์
score_mod์mask_mod์ ๋ ๊ฐ์ง ํจ์๋ก ๋ชจ๋ Attention Variant๋ฅผ ํํํ ์ ์์ด ์ ์ฐ์ฑ์ด ๋ฐ์ด๋จ. - ํนํ, ๋ค์ํ Attention Variant๋ฅผ ์กฐํฉํ ์ ์๋ ๊ธฐ๋ฅ (
Logical Fusion)์ด ๊ฐ์ ์ผ๋ก ์์ฉํจ.
โ 2. Block Mask Optimization (๋ธ๋ก ๋ง์คํน ์ต์ ํ)
- Block Masking์ ๋์ ํ์ฌ, ๋ชจ๋ Token์ ๊ฐ๋ณ์ ์ผ๋ก ๊ณ์ฐํ๋ ๋์ ๋ธ๋ก ๋จ์๋ก Masking์ ์ ์ฉํจ.
- ์์ ํ Masked๋ ๋ธ๋ก์ ์ฐ์ฐํ์ง ์๊ณ ๊ฑด๋๋ฐ๊ธฐ ๋๋ฌธ์ ๋ฉ๋ชจ๋ฆฌ ์ฌ์ฉ์ ์ค์ด๊ณ ์๋๋ฅผ ํฌ๊ฒ ํฅ์์ํด.
- ๋ถ๋ถ์ ์ผ๋ก Masked๋ ๋ธ๋ก๋ ํจ์จ์ ์ผ๋ก ์ฒ๋ฆฌํ์ฌ ์ถ๊ฐ์ ์ธ ์ฑ๋ฅ ํฅ์์ ๋ฌ์ฑ.
โ 3. Template-based Kernel Generation (ํ ํ๋ฆฟ ๊ธฐ๋ฐ ์ปค๋ ์์ฑ)
- PyTorch์
torch.compile()์ ์ฌ์ฉํ์ฌ ์ฌ์ฉ์๊ฐ ์ ์ํscore_mod์mask_mod๋ฅผ ํจ์จ์ ์ผ๋ก ์ต์ ํ. - ์ปค๋ ์ฝ๋๊ฐ ์๋์ผ๋ก ์์ฑ๋๋ฏ๋ก, ๋ค์ํ Variant์ ๋ํ ์ต์ ํ๊ฐ ์ฝ๊ฒ ์ด๋ฃจ์ด์ง.
๐ค ๋์ ์๊ฐ (์ FlexAttention์ด ๋ฐ์ด๋๊ฐ?)
FlexAttention์ ๊ฐ์ ์ ์ ์ฐ์ฑ๊ณผ ์ฑ๋ฅ ์ต์ ํ๋ฅผ ๋์์ ๋ฌ์ฑํ ์ ์ ๋๋ค.
๊ธฐ์กด์ FlashAttention์ ๊ณ ์ ๋ Attention Kernel์ ์์กดํ์ฌ ํน์ Attention Variant๋ฅผ ์ถ๊ฐํ๋ ๋ฐ ์ด๋ ค์์ด ์์์ต๋๋ค. ํ์ง๋ง FlexAttention์ ์ฌ์ฉ์ ์ ์ ์ฐ์ฐ (
score_mod,mask_mod)์ ํตํด ๊ฐ๋จํ ์ถ๊ฐํ ์ ์์ต๋๋ค.ํนํ, ๋ค์ํ Variant์ ์กฐํฉ์ ์ง์ํ๋ Logical Fusion ๊ธฐ๋ฅ์ ๊ธฐ์กด ๋ฐฉ๋ฒ๋ก ์์๋ ๊ฑฐ์ ๋ถ๊ฐ๋ฅํ๋ ์์ ์ ์ฝ๊ฒ ์ํํ ์ ์๊ฒ ๋ง๋ค์ด ์ค๋๋ค.
FlexAttention์ด ๊ธฐ์กด FlashAttention ๋๋น ์ฑ๋ฅ์ด ๋ฐ์ด๋ ์ด์ ๋ ํฌ๊ฒ ๋ ๊ฐ์ง์ ๋๋ค.
- Block Masking์ ํตํ ๋ฉ๋ชจ๋ฆฌ ์ต์ ํ: ๋ฉ๋ชจ๋ฆฌ ์ ๊ทผ์ ์ค์ด๊ณ , ์ฐ์ฐ์ ํ์ํ ๋๋ง ์ํํ์ฌ ์๋ ํฅ์.
- Template-based Kernel Generation: PyTorch์ ๋์ ์ปดํ์ผ ๊ธฐ๋ฅ์ ์ด์ฉํ์ฌ ๊ฐ Variant์ ์ต์ ํ๋ ์ปค๋์ ์๋์ผ๋ก ์์ฑ.
ํนํ Paged Attention ๋ฐฉ์์ ์ง์์์ FlexAttention์ด ๋ณด์ฌ์ค ์ฑ๋ฅ์ ์ธ์์ ์ ๋๋ค. ๊ธฐ์กด ๋ฐฉ๋ฒ๋ก ์ ํ์ด์ง ํ ์ด๋ธ ์ฌ์ฉ์ผ๋ก ์ธํด ์ฑ๋ฅ ์ ํ๊ฐ ์ปธ์ง๋ง, FlexAttention์ ์ด๋ฅผ ํจ์จ์ ์ผ๋ก ์ฒ๋ฆฌํจ์ผ๋ก์จ ๊ฑฐ์ ์ฑ๋ฅ ์์ค ์์ด ๋์ํฉ๋๋ค.
๐ ์ ๋ฆฌ
FlexAttention์ ๋ค์ํ Attention ๋ณํ์ ์ ์ฐํ๊ฒ ๊ตฌํํ๊ณ ์ต์ ํ๋ ์ปค๋์ ์๋์ผ๋ก ์์ฑํ๋ ์๋ก์ด ์ ๊ทผ๋ฒ์ ๋๋ค. Block Mask Optimization๊ณผ Paged Attention ์ง์์ ํนํ ์ฃผ๋ชฉํ ๋งํ ๊ฐ์ ์ฌํญ์ด๋ฉฐ, ์ด๋ฌํ ๊ธฐ๋ฅ๋ค์ด ๊ธฐ์กด FlashAttention ๋๋น ๋์ ์ฑ๋ฅ์ ๋ฌ์ฑํ ์ ์๊ฒ ๋ง๋ ํต์ฌ ์ด์ ๋ผ๊ณ ์๊ฐํฉ๋๋ค.
์ถ๊ฐ ๋ถ์์ด๋ ๋ค๋ฅธ ๋ฐฉ๋ฒ๋ก ๊ณผ์ ๋น๊ต๊ฐ ํ์ํ์๋ฉด ์๋ ค์ฃผ์ธ์! ๐
์ค์ ๋ก ์ ์ฉํ๋ ค๋ฉด: ์ถ๊ฐ ๊ตฌํ๊ณผ ์์ ์ฐ์
๋ฌผ์ด๋ณธ ๊ฒ: ์ด ๋ ผ๋ฌธ์์ ์ ์๋ ๋ฐฉ๋ฒ์ ์ค์ ๋ก ์ฌ์ฉํ๋ ค๋ฉด ๊ธฐ์กด ๋ฐฉ๋ฒ์์ ์ด๋ ํ ๊ตฌํ์ด ์ถ๊ฐ์ ์ผ๋ก ํ์ํ๊ณ ์ด๋ฌํ ๊ตฌํ์ ์๋ชจ๋๋ ๊ณต์ ๋ฐ ์ฐ์ฐ์ ํ์ํ ์ปดํจํ ์์์ ํฌ๊ธฐ์ ๋ํด์ ๊ณ์ฐํด์ค ์ ์๊ฒ ๋?
โ FlexAttention์ ์ค์ ๋ก ์ฌ์ฉํ๊ธฐ ์ํด ๊ธฐ์กด ๋ฐฉ๋ฒ์์ ์ถ๊ฐ์ ์ผ๋ก ํ์ํ ๊ตฌํ๊ณผ ๋น์ฉ ๋ถ์
๐ 1. ๊ธฐ์กด ๋ฐฉ๋ฒ๊ณผ FlexAttention์ ์ฐจ์ด์
๊ธฐ์กด ๋ฐฉ๋ฒ (FlashAttention v2, v3)
- ๊ณ ์ ๋ Kernel ๊ตฌ์กฐ๋ก ์ ๊ณต.
- ํน์ Attention Variants (e.g., Noop, Causal, Alibi) ๋ง ์ง์.
- ์๋ก์ด Variant ์ถ๊ฐ๋ฅผ ์ํด์๋ ์ปค๋ ์ฝ๋๋ฅผ ์ง์ ์์ฑํ๊ฑฐ๋ ์์ ํ์.
- ๊ธฐ์กด ์ปค๋์ PyTorch์
torch.compile()๊ณผ ๊ฐ์ ์ต์ ํ ๊ธฐ๋ฅ์ ์ถฉ๋ถํ ํ์ฉํ์ง ๋ชปํจ.
๐ 2. FlexAttention์์ ์ถ๊ฐ์ ์ผ๋ก ํ์ํ ๊ตฌํ
(1) ์ฌ์ฉ์ ์ ์ ํจ์ ์์ฑ (score_mod, mask_mod)
- FlexAttention์์ ์๋ก์ด Attention Variant๋ฅผ ์ฌ์ฉํ๋ ค๋ฉด ์ฌ์ฉ์๊ฐ
score_mod์mask_modํจ์๋ฅผ ์์ฑํด์ผ ํจ. - ์๋ฅผ ๋ค์ด,
Sliding Window Mask์ ๊ฒฝ์ฐ:PYTHONdef sliding_window_mask(q_idx, kv_idx, window_size=3): return abs(q_idx - kv_idx) <= window_sizedef sliding_window_mask(q_idx, kv_idx, window_size=3): return abs(q_idx - kv_idx) <= window_size - ์๋ก์ด Variant๋ฅผ ์ถ๊ฐํ๋ ๊ฒฝ์ฐ, ์ด ๋ ๊ฐ์ง ํจ์๋ฅผ ์์ฑํ๋ ๊ฒ๋ง์ผ๋ก๋ ์ถฉ๋ถํ ๊ตฌํ ๊ฐ๋ฅ.
(2) Kernel Compilation (torch.compile ์ฌ์ฉ)
- PyTorch์
torch.compile()์ ์ฌ์ฉํ์ฌ ์์ฑ๋ ํจ์๋ฅผ ์ต์ ํ๋ ์ปค๋๋ก ๋ณํ. - ์ด ๊ณผ์ ์์ PyTorch๊ฐ ๊ธฐ์กด ์ปค๋์ ํธ๋์คํ์ผ๋งํ์ฌ GPU์ ์ ํฉํ ์ฝ๋๋ก ๋ณํ.
- ์์ ์ฝ๋:PYTHON
import torch from torch import nn class FlexAttention(nn.Module): def __init__(self, score_mod, mask_mod): super().__init__() self.score_mod = score_mod self.mask_mod = mask_mod def forward(self, Q, K, V): S = torch.matmul(Q, K.transpose(-2, -1)) / (Q.size(-1) ** 0.5) S = self.score_mod(S) S = torch.softmax(S, dim=-1) return torch.matmul(S, V) model = FlexAttention(score_mod=sliding_window_mask, mask_mod=alibi_bias) compiled_model = torch.compile(model)import torch from torch import nn class FlexAttention(nn.Module): def __init__(self, score_mod, mask_mod): super().__init__() self.score_mod = score_mod self.mask_mod = mask_mod def forward(self, Q, K, V): S = torch.matmul(Q, K.transpose(-2, -1)) / (Q.size(-1) ** 0.5) S = self.score_mod(S) S = torch.softmax(S, dim=-1) return torch.matmul(S, V) model = FlexAttention(score_mod=sliding_window_mask, mask_mod=alibi_bias) compiled_model = torch.compile(model)
(3) Block Mask Optimization ๊ตฌํ
- FlexAttention์ ๊ฐ์ ์ค ํ๋์ธ Block Mask Optimization์ ์ฌ์ฉํ๊ธฐ ์ํด,
BlockMask๋ฐ์ดํฐ ๊ตฌ์กฐ๋ฅผ ์ ์ํ๊ณ ํ์ฉํด์ผ ํจ. - ์ผ๋ฐ์ ์ผ๋ก PyTorch์์ ์ ๊ณตํ๋ Tensor ์ฐ์ฐ๊ณผ GPU ์ฐ์ฐ์ ํ์ฉํ์ฌ ๊ตฌํ.
๐ 3. ๊ตฌํ์ ํ์ํ ๊ณต์ (Development Cost)
| ์์ | ๋์ด๋ (1~5) | ์์ ์๊ฐ (์๊ฐ) | ํ์ ์์ |
|---|---|---|---|
| score_mod / mask_mod ์์ฑ | 2 | 1~2 ์๊ฐ | Python, PyTorch |
| Kernel Compilation ์ค์ | 3 | 2~3 ์๊ฐ | PyTorch, GPU ํ๊ฒฝ |
| Block Mask Optimization ๊ตฌํ | 4 | 4~6 ์๊ฐ | PyTorch, GPU ํ๊ฒฝ |
| PagedAttention ์ ์ฉ | 4 | 3~5 ์๊ฐ | PyTorch, GPU ํ๊ฒฝ |
| ํ ์คํธ ๋ฐ ๊ฒ์ฆ | 3 | 2~4 ์๊ฐ | GPU ํ๊ฒฝ |
- ์ด ์์ ์๊ฐ: ์ฝ 12~20 ์๊ฐ (ํ๋ฃจ์์ ์ดํ ์ ๋)
๐ 4. ์ปดํจํ ์์ ์๊ตฌ๋ ๋ถ์
(1) Kernel Compilation
torch.compile()์ฌ์ฉ ์ GPU์ ์ปดํจํ ์์์ ํฌ๊ฒ ์๋ชจ.- GPU ๋ฉ๋ชจ๋ฆฌ ์ฌ์ฉ๋: ์ปค๋ ํฌ๊ธฐ์ ๋ฐ์ดํฐ ํฌ๊ธฐ์ ๋น๋ก (๋ณดํต 16GB ์ด์์ GPU ๊ถ์ฅ).
- ์ฃผ์ ์ฐ์ฐ:
torch.compile()์ ์ปดํ์ผ ๋จ๊ณ: CUDA ์ปค๋ ์์ฑ ๋ฐ ์ต์ ํ.- GPU ๋ฉ๋ชจ๋ฆฌ ์ฌ์ฉ ๋ฐ ์ฒ๋ฆฌ ์๊ฐ: Variant์ ํฌ๊ธฐ ๋ฐ Mask ๊ตฌ์กฐ์ ๋ฐ๋ผ ๋ค๋ฆ.
(2) Training ๋ฐ Inference Performance ๋ถ์
- FlexAttention์ ๊ธฐ์กด FlashAttention ๋๋น ๋ฉ๋ชจ๋ฆฌ ์ฌ์ฉ์ ์ค์ด๊ณ ์๋๋ฅผ ํฌ๊ฒ ๊ฐ์ .
- ์ค์ ํ์ต ๋ฐ ์ถ๋ก ์, GPU ์ฐ์ฐ ์๋๋ ๊ธฐ์กด ๋๋น ์ฝ 1.4๋ฐฐ ~ 2๋ฐฐ ํฅ์.
- ๋ฉ๋ชจ๋ฆฌ ์ฌ์ฉ๋์ Block Mask Optimization์ ์ ์ฉํ ๊ฒฝ์ฐ, ๊ธฐ์กด ๋๋น ์ต๋ 50% ๊ฐ์ ๊ฐ๋ฅ.
๐ ๊ณ์ฐ ์์: FlexAttention์ ์ ์ฉํ ๋ชจ๋ธ ํ์ต
| ๋ชจ๋ธ | ๊ธฐ์กด FlashAttention | FlexAttention (์ถ์ ) |
|---|---|---|
| GPU ๋ฉ๋ชจ๋ฆฌ ์ฌ์ฉ๋ | 24GB | 16GB |
| ์ฐ์ฐ ์๋ (TFLOPS) | 400 | 560 (~1.4๋ฐฐ) |
| ํ์ต ์๊ฐ (์๊ฐ) | 10 | 7.2 (~1.4๋ฐฐ) |
๐ก ์ ๋ฆฌ
FlexAttention์ ๊ธฐ์กด ๋ชจ๋ธ์ ์ ์ฉํ๊ธฐ ์ํด์๋ ๋ค์์ ์ถ๊ฐ๋ก ๊ตฌํํด์ผ ํฉ๋๋ค.
- ์ฌ์ฉ์ ์ ์ ํจ์ (
score_mod,mask_mod) ์์ฑ. - PyTorch
torch.compile()์ ํ์ฉํ ์ปค๋ ์ต์ ํ. - Block Mask Optimization ๊ตฌํ.
- ์ฌ์ฉ์ ์ ์ ํจ์ (
๊ฐ๋ฐ ๊ณต์๋ 12~20์๊ฐ ์ ๋ ์์๋ ๊ฒ์ผ๋ก ์์๋ฉ๋๋ค. (ํนํ GPU ํ๊ฒฝ์์์ ํ ์คํธ์ ์ต์ ํ๊ฐ ํต์ฌ)
์ปดํจํ ์์ ์๊ตฌ๋์ ๊ธฐ์กด FlashAttention ๋๋น GPU ๋ฉ๋ชจ๋ฆฌ ์ฌ์ฉ๋์ ์ค์ด๋ค๊ณ , ์ฐ์ฐ ์๋๋ ์ต๋ 1.4๋ฐฐ ์ฆ๊ฐํ ๊ฒ์ผ๋ก ๊ธฐ๋๋ฉ๋๋ค.
์ ๋ ฅ ๋ฐ์ดํฐ์ ์ถ๋ก ๊ณผ์ , ๋ชจ๋ธ ๊ตฌ์กฐ์ ์์ ์๊ตฌ๋
๋ฌผ์ด๋ณธ ๊ฒ: ์ด ๋ ผ๋ฌธ์ ์ ๋ ฅ๋ฐ์ดํฐ์ ์ถ๋ก ๊ณผ์ ์ ๋ํด์ ์์๋ฅผ ๋ค์ด ์์ฃผ ์์ธํ๊ฒ ์ค๋ช ํด ์ฃผ๊ฒ ๋? ์ถ๊ฐ์ ์ผ๋ก ๋ชจ๋ธ์ํคํ ์ฒ์ ๊ตฌ์ฑ ๋ฐ ๋ชจ๋ธ์ ๊ตฌ์ฑํ๋ ์ฐ์ฐ๊ณผ ๋ฉ๋ชจ๋ฆฌ ์๊ตฌ๋ ์ปดํจํ ์๊ตฌ๋ ๋ํด์๋ ์ค๋ช ํด์ค
๐ FlexAttention ๋ ผ๋ฌธ์ ์ ๋ ฅ ๋ฐ์ดํฐ, ์ถ๋ก ๊ณผ์ , ๋ชจ๋ธ ์ํคํ ์ฒ ๊ตฌ์ฑ ๋ฐ ์๊ตฌ๋ ๋ถ์
๐ 1. ์ ๋ ฅ ๋ฐ์ดํฐ (Input Data)
FlexAttention์ ๊ธฐ์กด Transformer ๋ชจ๋ธ์ Self-Attention ๋ฉ์ปค๋์ฆ์ ๊ฐ์ ํ๋ ๊ตฌ์กฐ๋ก, ์ ๋ ฅ ๋ฐ์ดํฐ๋ Transformer ๋ชจ๋ธ๊ณผ ๋์ผํฉ๋๋ค.
โ ์ ๋ ฅ ๋ฐ์ดํฐ ๊ตฌ์ฑ
- Query (
Q) - Key (
K) - Value (
V)
โ ์ ๋ ฅ ๋ฐ์ดํฐ์ ํํ
- \( Q \in \mathbb{R}^{B \times H \times Q\_LEN \times D} \)
- \( K, V \in \mathbb{R}^{B \times H \times KV\_LEN \times D} \)
- \( B \): ๋ฐฐ์น ํฌ๊ธฐ (Batch Size)
- \( H \): Attention Heads (๋ฉํฐํค๋ Attention์ ์)
- \( Q\_LEN \): Query์ ๊ธธ์ด (ํ ํฐ ์)
- \( KV\_LEN \): Key/Value์ ๊ธธ์ด (ํ ํฐ ์)
- \( D \): ๊ฐ Token์ Embedding ์ฐจ์
โ ์์ ์ ๋ ฅ ๊ฐ
B = 2 # Batch Size (e.g., ๋ฌธ์ 2๊ฐ)
H = 4 # Attention Heads (e.g., 4๊ฐ๋ก ๋ถํ )
Q_LEN = 8 # Query Length (e.g., 8๊ฐ ํ ํฐ)
KV_LEN = 10 # Key/Value Length (e.g., 10๊ฐ ํ ํฐ)
D = 64 # Embedding Dimension
Q = torch.randn(B, H, Q_LEN, D)
K = torch.randn(B, H, KV_LEN, D)
V = torch.randn(B, H, KV_LEN, D)๐ 2. ์ถ๋ก ๊ณผ์ (Inference Process)
FlexAttention์ ๊ธฐ์กด Attention ์ฐ์ฐ์ ํ์ฅํ์ฌ, ์ปค์คํ Masking ๋ฐ Score Modification ๊ธฐ๋ฅ์ ์ ์ฉํ ์ ์์ต๋๋ค.
โ Attention ์ฐ์ฐ ๊ณผ์ (๊ธฐ๋ณธ ํํ)
Query-Key Similarity ๊ณ์ฐ (Score Matrix \( S \))
\[ S = \frac{Q K^T}{\sqrt{d_k}} \]- \( Q \in \mathbb{R}^{B \times H \times Q\_LEN \times D} \)
- \( K \in \mathbb{R}^{B \times H \times KV\_LEN \times D} \)
- \( S \in \mathbb{R}^{B \times H \times Q\_LEN \times KV\_LEN} \)
Masking ์ ์ฉ (
mask_mod)- ์: Sliding Window Mask
PYTHONdef sliding_window_mask(q_idx, kv_idx, window_size=3): return abs(q_idx - kv_idx) <= window_sizedef sliding_window_mask(q_idx, kv_idx, window_size=3): return abs(q_idx - kv_idx) <= window_size- Sliding Window์ ๊ฒฝ์ฐ, ํน์ ๋ฒ์ ์์ Key๋ง์ ๊ณ ๋ คํ๋๋ก Masking์ ์ ์ฉ.
Score Modification ์ ์ฉ (
score_mod)- ์: Alibi Bias
PYTHONdef alibi_bias(score, q_idx, kv_idx, head_idx, bias_factor=0.5): return score + bias_factor * (q_idx - kv_idx)def alibi_bias(score, q_idx, kv_idx, head_idx, bias_factor=0.5): return score + bias_factor * (q_idx - kv_idx)Softmax ์ ์ฉ
\[ S' = \text{softmax}(S) \]Weighted Sum (Output ๊ณ์ฐ)
\[ \text{Output} = S' V \]
๐ 3. ๋ชจ๋ธ ์ํคํ ์ฒ ๊ตฌ์ฑ (Model Architecture)
FlexAttention์ ๊ธฐ์กด Self-Attention ๊ตฌ์กฐ๋ฅผ ๊ธฐ๋ฐ์ผ๋ก ๊ตฌ์ฑ๋ฉ๋๋ค.
โ ๊ตฌ์ฑ ์์ (Components)
- Input Embedding Layer
- Multi-Head Attention Layer (FlexAttention ์ ์ฉ)
- Feedforward Layer
- Residual Connection & Layer Normalization
๐ 4. ์ฐ์ฐ ์๊ตฌ๋ ๋ถ์ (Computational Cost)
FlexAttention์ ๊ธฐ์กด FlashAttention ๋๋น ๋ ๋์ ์ฐ์ฐ ํจ์จ์ ์ ๊ณตํ์ง๋ง, ์ถ๊ฐ์ ์ผ๋ก Masking๊ณผ Score Modification ๊ณผ์ ์ด ํ์ํฉ๋๋ค.
โ ์ฐ์ฐ๋ (FLOPs) ๊ณ์ฐ)
Query-Key Similarity ์ฐ์ฐ
\[ \text{FLOPs} = B \times H \times Q\_LEN \times KV\_LEN \times D \]- ์์: \( B = 2 \), \( H = 4 \), \( Q\_LEN = 8 \), \( KV\_LEN = 10 \), \( D = 64 \) \[ \text{FLOPs} = 2 \times 4 \times 8 \times 10 \times 64 = 40,960 \]
Masking (
mask_mod)- ๋จ์ ๋น๊ต ์ฐ์ฐ์ด๋ฏ๋ก ์ถ๊ฐ์ ์ธ FLOPs๋ ํฌ์ง ์์.
Score Modification (
score_mod)- ์๋ฅผ ๋ค์ด, Alibi Bias ์ ์ฉ ์ ๊ฐ ์ ์๋ง๋ค ํ ๋ฒ์ ๋ง์ ์ฐ์ฐ์ด ํ์. \[ \text{FLOPs} = B \times H \times Q\_LEN \times KV\_LEN \] \[ \text{FLOPs} = 2 \times 4 \times 8 \times 10 = 640 \]
Softmax ์ฐ์ฐ
\[ \text{FLOPs} \approx B \times H \times Q\_LEN \times KV\_LEN \times 2 \]Weighted Sum
\[ \text{FLOPs} = B \times H \times Q\_LEN \times KV\_LEN \times D \]
๐ 5. ๋ฉ๋ชจ๋ฆฌ ์๊ตฌ๋ ๋ถ์ (Memory Requirements)
โ ๊ธฐ์กด ๋ชจ๋ธ๊ณผ์ ๋น๊ต
| ๊ตฌ์ฑ ์์ | ๊ธฐ์กด ๋ชจ๋ธ (FlashAttention) | FlexAttention (์ถ์ ) |
|---|---|---|
| Score Matrix \( S \) | \( B \times H \times Q\_LEN \times KV\_LEN \) | ๋์ผ |
| Mask Matrix | ์์ | Block Mask๋ก ์ถ๊ฐ ๋ฉ๋ชจ๋ฆฌ ์ฌ์ฉ |
| Output Matrix | \( B \times H \times Q\_LEN \times D \) | ๋์ผ |
| ์ด ๋ฉ๋ชจ๋ฆฌ ์ฌ์ฉ๋ | ๊ธฐ๋ณธ ๋ฉ๋ชจ๋ฆฌ ์ฌ์ฉ๋ | ๊ธฐ์กด ๋๋น ์ฝ 20% ์ฆ๊ฐ (Block Mask) |
๐ 6. ์ปดํจํ ์๊ตฌ๋ ๋ถ์ (Computational Requirements)
- GPU ์๊ตฌ๋: ์ต์ 16GB ์ด์์ GPU ๊ถ์ฅ (ํนํ, Block Mask Optimization ์ฌ์ฉ ์).
- ์ฐ์ฐ ์๋: ๊ธฐ์กด FlashAttention ๋๋น ์ฝ 1.4๋ฐฐ ~ 2๋ฐฐ ์ฆ๊ฐ.
- ๋ฉ๋ชจ๋ฆฌ ์ฌ์ฉ๋: ๊ธฐ์กด ๋๋น ์ต๋ 20% ์ฆ๊ฐ ๊ฐ๋ฅ (Block Mask ์ ์ฉ ์).
๐ก ์ ๋ฆฌ
FlexAttention์ ์
๋ ฅ ๋ฐ์ดํฐ์ ์ถ๋ก ๊ณผ์ ์ ๊ธฐ์กด์ Self-Attention๊ณผ ์ ์ฌํ์ง๋ง, ์ฌ์ฉ์ ์ ์ ํจ์ (score_mod, mask_mod)๋ฅผ ํ์ฉํ์ฌ ์ถ๊ฐ์ ์ธ ์ฐ์ฐ์ ์ํํฉ๋๋ค. ์ด๋ก ์ธํด ์ฝ๊ฐ์ ์ถ๊ฐ ๋ฉ๋ชจ๋ฆฌ ์ฌ์ฉ๊ณผ ์ฐ์ฐ๋์ด ํ์ํ์ง๋ง, Block Mask Optimization ๋๋ถ์ ์ ์ฒด ์ฐ์ฐ ํจ์จ์ ํฌ๊ฒ ์ฆ๊ฐํฉ๋๋ค.
ํ๊ณ๋ฅผ ๋์ผ๋ ค๋ ์ฐ๊ตฌ ํ๋ฆ
๋ฌผ์ด๋ณธ ๊ฒ: ์ด ๋ ผ๋ฌธ์ ํ๊ณ๋ฅผ ๊ทน๋ณตํ๊ธฐ ์ํ ๋ฐฉ๋ฒ์ผ๋ก ์ด๋ค ์ฐ๊ตฌํ๋ฆ์ด ์๋์ง ์ ๋ฆฌํด์ ์์ธํ๊ฒ ์ค๋ช ํด ์ค
๐ FlexAttention ๋ ผ๋ฌธ์ ํ๊ณ์ ์ด๋ฅผ ๊ทน๋ณตํ๊ธฐ ์ํ ์ฐ๊ตฌ ํ๋ฆ
FlexAttention์ ๋ค์ํ Attention Variants๋ฅผ ์ ์ฐํ๊ฒ ์ฒ๋ฆฌํ ์ ์๋ค๋ ์ฅ์ ์ด ์์ง๋ง, ๋ช ๊ฐ์ง ์ค์ํ ํ๊ณ์ ์ด ์กด์ฌํฉ๋๋ค. ์ด๋ฅผ ๊ทน๋ณตํ๊ธฐ ์ํด ํ์ฌ ์ฐ๊ตฌ๋๊ณ ์๋ ๋ฐฉํฅ๋ค์ ์ ๋ฆฌํ๊ฒ ์ต๋๋ค.
๐ 1. FlexAttention์ ํ๊ณ์
โ (1) ๋ฉ๋ชจ๋ฆฌ ์ฌ์ฉ๋ ์ฆ๊ฐ (ํนํ Block Mask ์ฌ์ฉ ์)
- FlexAttention์ Block Mask Optimization์ ๋ฉ๋ชจ๋ฆฌ ์ฌ์ฉ์ ์ค์ด๊ธฐ ์ํด ์ค๊ณ๋์์ผ๋, ์ค์ ๋ก๋ BlockMask ๋ฐ์ดํฐ ๊ตฌ์กฐ๊ฐ ์ถ๊ฐ๋๋ฉด์ ๋ฉ๋ชจ๋ฆฌ ์ฌ์ฉ๋์ด ์ฆ๊ฐํ ์ ์์ต๋๋ค.
- ํนํ, ๋๊ท๋ชจ ๋ชจ๋ธ์์ ๊ธด ๋ฌธ์ฅ์ ์ฒ๋ฆฌํ ๋ BlockMask๋ก ์ธํ ๋ฉ๋ชจ๋ฆฌ ์ค๋ฒํค๋๊ฐ ๋ฐ์ํ ์ ์์ต๋๋ค.
โ (2) ํน์ Attention Variant์ ์ต์ ํ ๋ถ์กฑ
- ์ผ๋ถ Attention Variant (ํนํ, Neighborhood Attention ๋ฑ)์์ ์ต์ ํ ์์ค์ด ๋ฎ์.
- FlexAttention์ Mask ๋ฐ Score Modification์ ์ฌ์ฉํ์ฌ ๋ค์ํ ๋ณํ์ ๊ตฌํํ ์ ์์ง๋ง, ์ผ๋ถ ๊ฒฝ์ฐ ๊ธฐ์กด ์ปค๋๋ณด๋ค ์ฑ๋ฅ์ด ๋จ์ด์ง ์ ์์ต๋๋ค.
โ (3) ์ปค๋ ์ปดํ์ผ ๊ณผ์ ์ ์ค๋ฒํค๋
- PyTorch์
torch.compile()์ ์ฌ์ฉํ์ฌ ์ปค๋์ ์ปดํ์ผํ๋ ๊ณผ์ ์ ์ถ๊ฐ์ ์ธ ์ฐ์ฐ ๋น์ฉ์ ์ ๋ฐํฉ๋๋ค. - ํนํ ์ค์๊ฐ ์ถ๋ก (inference) ํ๊ฒฝ์์ ์ปดํ์ผ ์ค๋ฒํค๋๊ฐ ๋ฌธ์ ๊ฐ ๋ ์ ์์ต๋๋ค.
โ (4) Paged Attention์์์ ํ๊ณ
- Paged Attention์ ์ง์ํ๊ธฐ๋ ํ์ง๋ง, ๋ฉ๋ชจ๋ฆฌ ์ ๊ทผ ํจํด์ด ๋ณต์กํด์ง์๋ก ์ฑ๋ฅ ์ ํ ๊ฐ๋ฅ์ฑ์ด ์กด์ฌํฉ๋๋ค.
- ํนํ GPU์ ๋ฉ๋ชจ๋ฆฌ ์ ๊ทผ ๋ฐฉ์์ ๋ฐ๋ผ ์ฑ๋ฅ์ด ํฌ๊ฒ ๋ฌ๋ผ์ง ์ ์์ต๋๋ค.
๐ 2. ํ๊ณ๋ฅผ ๊ทน๋ณตํ๊ธฐ ์ํ ์ฐ๊ตฌ ํ๋ฆ
FlexAttention์ ํ๊ณ๋ฅผ ๊ทน๋ณตํ๊ธฐ ์ํด ๋ค์๊ณผ ๊ฐ์ ์ฐ๊ตฌ ํ๋ฆ์ด ์กด์ฌํฉ๋๋ค.
๐ (1) Block Mask Optimization์ ๊ฐ์
โ ์ฐ๊ตฌ ํ๋ฆ
- ํ์ฌ FlexAttention์ BlockMask๋ฅผ ์ฌ์ฉํ์ฌ ๋ธ๋ก ๋จ์๋ก ์ฐ์ฐ์ ๊ฑด๋๋ฐ๊ฑฐ๋ ํจ์จ์ ์ผ๋ก ์ฒ๋ฆฌํฉ๋๋ค.
- ๊ทธ๋ฌ๋ BlockMask ์์ฒด์ ํฌ๊ธฐ๊ฐ ์ปค์ง ๊ฒฝ์ฐ, ๋ฉ๋ชจ๋ฆฌ ์ฌ์ฉ๋์ด ์ฆ๊ฐํฉ๋๋ค.
- ์ด๋ฅผ ๊ฐ์ ํ๊ธฐ ์ํด Sparse Attention ๊ธฐ๋ฒ์ ํ์ฉํ๊ฑฐ๋, Dynamic Block Masking ๊ธฐ๋ฒ์ ๋์ ํ๋ ์ฐ๊ตฌ๊ฐ ์งํ ์ค์ ๋๋ค.
๐ ๊ด๋ จ ์ฐ๊ตฌ ์์
- Sparse Transformer (Child et al., 2019):
- ์ผ๋ถ ํ ํฐ ์๋ง ์ง์ค์ ์ผ๋ก ์ฒ๋ฆฌํ๋ Sparse Attention ๊ตฌ์กฐ๋ฅผ ์ฌ์ฉํ์ฌ ๋ฉ๋ชจ๋ฆฌ ์ฌ์ฉ์ ์ค์.
- Longformer (Beltagy et al., 2020):
- Sliding Window ๊ธฐ๋ฐ์ Attention์ ์ฌ์ฉํ์ฌ ๊ธด ๋ฌธ์ฅ ์ฒ๋ฆฌ์ ํจ์จ์ ์ธ Sparse Attention ๊ตฌ์กฐ ์ ์.
- BigBird (Zaheer et al., 2020):
- ๋๋คํ๊ฒ ์ ํ๋ ์ผ๋ถ ํ ํฐ ์์ ํฌํจํ์ฌ Sparse Attention์ ๊ตฌํ, ๋ฉ๋ชจ๋ฆฌ ํจ์จ์ฑ์ ๊ทน๋ํ.
๐ (2) Multi-Stage Attention Optimization
โ ์ฐ๊ตฌ ํ๋ฆ
- Attention ์ฐ์ฐ์ ์ฌ๋ฌ ๋จ๊ณ๋ก ๋ถ๋ฆฌํ์ฌ ์ฐ์ฐ ํจ์จ์ ๋์ด๋ ๋ฐฉ๋ฒ.
- ์๋ฅผ ๋ค์ด,
score_mod์mask_mod๋ฅผ ์ ์ฉํ๋ ๋จ๊ณ๋ฅผ ๋ถ๋ฆฌํ์ฌ ๊ฐ๊ฐ ์ต์ ํํ๋ ๋ฐฉ์.
๐ ๊ด๋ จ ์ฐ๊ตฌ ์์
- Perceiver (Jaegle et al., 2021):
- Attention ์ฐ์ฐ์ ๋ค๋จ๊ณ๋ก ๋๋์ด ํจ์จ์ ์ผ๋ก ํ์ต.
- ์ ๋ ฅ ๋ฐ์ดํฐ์ ํฌ๊ธฐ๋ฅผ ์ค์ด๊ณ , ๋จ๊ณ๋ณ๋ก ์ค์ํ ์ ๋ณด๋ฅผ ์ถ์ถํ๋ ๋ฐฉ์.
- Linformer (Wang et al., 2020):
- Low-rank Approximation์ ์ฌ์ฉํ์ฌ Attention ํ๋ ฌ์ ์์ถํ์ฌ ํจ์จ์ฑ์ ๋์.
๐ (3) ์ปค๋ ์ปดํ์ผ ๊ณผ์ ์ ์ต์ ํ
โ ์ฐ๊ตฌ ํ๋ฆ
torch.compile()์ ์ปดํ์ผ ๊ณผ์ ์ ์ต์ ํํ์ฌ ์ปดํ์ผ ์ค๋ฒํค๋๋ฅผ ์ค์ด๋ ์ฐ๊ตฌ.- ์ปค๋ ์ปดํ์ผ ๊ณผ์ ์ ์ฌ์ ์ ์ํํ์ฌ ์ฌ์ฌ์ฉํ๋ ๋ฐฉ๋ฒ์ ์ฐ๊ตฌ ์ค.
๐ ๊ด๋ จ ์ฐ๊ตฌ ์์
- TVM (Chen et al., 2018):
- ๋ฅ๋ฌ๋ ๋ชจ๋ธ์ ์ํ ์ปค์คํ ์ปค๋ ์ต์ ํ ์ปดํ์ผ๋ฌ.
- ์ปดํ์ผ๋ ์ปค๋์ ์ ์ฅํ๊ณ ์ฌ์ฌ์ฉํ ์ ์๋ ๊ตฌ์กฐ ์ ๊ณต.
- Halide (Ragan-Kelley et al., 2012):
- ์ปค๋ ์ปดํ์ผ์ ์ต์ ํํ๊ธฐ ์ํด ๊ณ ๋๋ก ์ต์ ํ๋ ์ปค์คํ ํ์ดํ๋ผ์ธ ์ ๊ณต.
๐ (4) Paged Attention์ ๊ฐ์
โ ์ฐ๊ตฌ ํ๋ฆ
- Paged Attention์ ๋ฉ๋ชจ๋ฆฌ ์ ๊ทผ ํจํด์ ์ต์ ํํ์ฌ GPU ์ฑ๋ฅ์ ๊ทน๋ํํ๋ ์ฐ๊ตฌ.
- ์ปค๋์ ๋์ฑ ์ ์ฐํ๊ฒ ๊ตฌ์ฑํ์ฌ ๋ค์ํ Paged Attention Variant๋ฅผ ํจ์จ์ ์ผ๋ก ์ฒ๋ฆฌ.
๐ ๊ด๋ จ ์ฐ๊ตฌ ์์
- vLLM (Kwon et al., 2023):
- ๋๊ท๋ชจ ๋ชจ๋ธ์ ์ํ ๋ฉ๋ชจ๋ฆฌ ๊ด๋ฆฌ ๊ธฐ๋ฒ์ ๊ฐ์ ํ์ฌ Paged Attention ์ฑ๋ฅ์ ํฌ๊ฒ ํฅ์์ํด.
- ํนํ ์ปค๋ ์์ค์์์ ์ต์ ํ๋ฅผ ๊ฐ์กฐ.
๐ 3. FlexAttention์ ํ๊ณ ๊ทน๋ณต์ ์ํ ๋์ ์ ์
Dynamic Block Masking ๊ธฐ๋ฒ ๋์
- Block Mask์ ํฌ๊ธฐ๋ฅผ ํ์ต ๊ณผ์ ์์ ์๋์ผ๋ก ์กฐ์ ํ๊ฑฐ๋, ํ์์ ๋ฐ๋ผ ๋์ ์ผ๋ก ๋ณ๊ฒฝํ๋ ๋ฐฉ์.
- ์๋ฅผ ๋ค์ด, ์ค์๋๊ฐ ๋์ ๋ธ๋ก์ ์์ ํฌ๊ธฐ๋ก ๋๋์ด ์ฒ๋ฆฌํ๊ณ , ์ค์๋๊ฐ ๋ฎ์ ๋ธ๋ก์ ํฐ ํฌ๊ธฐ๋ก ๋ฌถ์ด ์ฒ๋ฆฌํ๋ ๋ฐฉ์.
Efficient Kernel Caching
- ์ปค๋ ์ปดํ์ผ ๊ณผ์ ์ ์ฌ์ ์ ์ํํ๊ณ , ๋ค์ํ Mask์ Score Modification ํจ์์ ๋ํด ์บ์ฑํ์ฌ ์ฌ์ฌ์ฉ.
- PyTorch์
torch.compile()๊ธฐ๋ฅ์ ํ์ฉํ๋, ์ปค๋ ์บ์ ์ ์ฅ์๋ฅผ ์ถ๊ฐํ์ฌ ์ปดํ์ผ ์ค๋ฒํค๋๋ฅผ ์ค์.
Adaptive Paged Attention
- Paged Attention ์ฌ์ฉ ์, ๋ฉ๋ชจ๋ฆฌ ์ ๊ทผ ํจํด์ ํ์ตํ๋ฉด์ ์ต์ ํํ๋ ๋ฐฉ์.
- ์๋ฅผ ๋ค์ด, Token์ ์ค์๋์ ๋ฐ๋ผ ๋ฉ๋ชจ๋ฆฌ ์ ๊ทผ ์์๋ฅผ ์ฌ๋ฐฐ์นํ์ฌ ์ฑ๋ฅ ์ต์ ํ.
๐ก ์ ๋ฆฌ
FlexAttention์ ๊ฐ๋ ฅํ ์ ์ฐ์ฑ๊ณผ ์ฑ๋ฅ์๋ ๋ถ๊ตฌํ๊ณ , ๋ฉ๋ชจ๋ฆฌ ์ฌ์ฉ๋ ์ฆ๊ฐ์ ํน์ Attention Variant์ ๋ํ ์ต์ ํ ๋ถ์กฑ ๋ฌธ์ ๊ฐ ์์ต๋๋ค. ์ด๋ฅผ ํด๊ฒฐํ๊ธฐ ์ํด Sparse Attention ๊ธฐ๋ฒ, Multi-Stage Optimization, Kernel Caching, ๊ทธ๋ฆฌ๊ณ Paged Attention ์ต์ ํ ๋ฑ์ด ํ๋ฐํ ์ฐ๊ตฌ๋๊ณ ์์ต๋๋ค.
Q : FlexAttention์์ ์ ์ํ๋ Block Mask Optimization ๊ธฐ๋ฒ์ด ๊ธฐ์กด์ Sparse Attention ๋ฐฉ๋ฒ๋ค๊ณผ ๋น๊ตํ์ ๋ ์ด๋ ํ ์ฐจ๋ณํ๋ ์ฅ์ ๊ณผ ๋จ์ ์ ๊ฐ์ง๋๊ฐ?
โ Block Mask Optimization vs. ๊ธฐ์กด Sparse Attention ๋ฐฉ๋ฒ ๋น๊ต
FlexAttention์ Block Mask Optimization ๊ธฐ๋ฒ์ ๊ธฐ์กด์ Sparse Attention ๋ฐฉ๋ฒ๋ค๊ณผ ๋น๊ตํ์ ๋ ๋ช ๊ฐ์ง ์ค์ํ ์ฐจ๋ณ์ ๊ณผ ํ๊ณ๊ฐ ์์ต๋๋ค.
๐ 1. Block Mask Optimization ๊ธฐ๋ฒ (FlexAttention)
๐ ๊ธฐ๋ฒ ์ค๋ช
- Attention Score Matrix๋ฅผ ๋ธ๋ก ๋จ์๋ก ๋๋์ด ์ฒ๋ฆฌ.
- Masking ๊ณผ์ ์์ ์์ ํ ๋ง์คํน๋ ๋ธ๋ก์ ๊ฑด๋๋ฐ๊ณ , ๋ถ๋ถ์ ์ผ๋ก ๋ง์คํน๋ ๋ธ๋ก๋ง ์ฐ์ฐ.
- BlockMask๋ผ๋ ๋ฐ์ดํฐ ๊ตฌ์กฐ๋ฅผ ์ฌ์ฉํ์ฌ ๋ธ๋ก์ ์์น์ ์ํ๋ฅผ ๊ด๋ฆฌ.
- GPU ๋ฉ๋ชจ๋ฆฌ ์ ๊ทผ์ ํจ์จ์ ์ผ๋ก ๊ด๋ฆฌํ์ฌ ์ฐ์ฐ ์๋๋ฅผ ๊ฐ์ .
๐ 2. ๊ธฐ์กด Sparse Attention ๊ธฐ๋ฒ
Sliding Window Attention (Longformer)
- ์ธ์ ํ ํ ํฐ์ ๋ํด์๋ง Attention์ ๊ณ์ฐํ์ฌ ์ฐ์ฐ๋ ๊ฐ์.
- ๊ธด ๋ฌธ์ฅ ์ฒ๋ฆฌ์ ํจ์จ์ ์ด๋, ์๋์ฐ ํฌ๊ธฐ๋ฅผ ๋ฒ์ด๋ ์ ๋ณด๋ ์ฒ๋ฆฌํ์ง ๋ชปํจ.
Global Sparse Attention (BigBird)
- ๋๋ค, ๊ธ๋ก๋ฒ, ๋ก์ปฌ์ ์ธ ๊ฐ์ง Attention ๋ฐฉ์์ ์กฐํฉํ์ฌ ๋ ๋์ ์ ๋ณด๋ฅผ ์ฒ๋ฆฌ.
- ๋๋คํ๊ฒ ์ผ๋ถ ํ ํฐ๋ง ์ ํํ์ฌ ์ฐ์ฐ์ ์ค์ด๋ ๋ฐฉ์.
Dilated Attention (Reformer)
- ์ ๋ ฅ ์ํ์ค๋ฅผ ์ผ์ ํ ๊ฐ๊ฒฉ์ผ๋ก ๋ถํ ํ์ฌ ์ฐ์ฐ.
- ๋ฉ๋ชจ๋ฆฌ ์ฌ์ฉ๋์ ์ค์ด์ง๋ง, ์ผ๋ถ ์ค์ํ ์ ๋ณด๊ฐ ์์ค๋ ์ ์์.
Hash-based Attention (Reformer)
- ํ ํฐ์ ํด์ฑํ์ฌ ๋น์ทํ ๊ฐ๋ผ๋ฆฌ ๋ฌถ์ด ์ฐ์ฐ์ ์ค์.
- ์ฐ์ฐ ํจ์จ์ด ๋์ง๋ง, ํด์ฑ์ ์ ํ๋๊ฐ ๋ฎ์ ๊ฒฝ์ฐ ์ฑ๋ฅ ์ ํ ๊ฐ๋ฅ.
๐ 3. Block Mask Optimization vs. ๊ธฐ์กด Sparse Attention ๋น๊ต
| ํน์ง | Block Mask Optimization (FlexAttention) | ๊ธฐ์กด Sparse Attention (Longformer, BigBird, Reformer) |
|---|---|---|
| ์ฐ์ฐ ํจ์จ์ฑ | ๋ธ๋ก ๋จ์๋ก ์ฐ์ฐํ์ฌ ๋ถํ์ํ ๊ณ์ฐ ์ ๊ฑฐ | ์ผ๋ถ ํ ํฐ๋ง ์ ํํ์ฌ ์ฐ์ฐ๋ ๊ฐ์ |
| ๋ฉ๋ชจ๋ฆฌ ์ฌ์ฉ๋ | ๋ธ๋ก ๋จ์๋ก ์ ๊ทผํ๋ฏ๋ก ๋ฉ๋ชจ๋ฆฌ ์ฌ์ฉ์ด ํจ์จ์ | ์ ์ฒด Score Matrix๋ฅผ ์ฌ์ฉํ์ง ์์ผ๋ฏ๋ก ๋ฉ๋ชจ๋ฆฌ ์ ์ฝ |
| ์ ์ฐ์ฑ | ๋ค์ํ Attention Variant์ ์ ์ฉ ๊ฐ๋ฅ | ํน์ Variant์ ๋ง์ถฐ ์ค๊ณ๋ ๊ตฌ์กฐ๊ฐ ๋ง์ |
| ๊ตฌํ ๋์ด๋ | ์๋์ ์ผ๋ก ๋์ ํธ | ๊ตฌ์กฐ์ ๋ฐ๋ผ ๋ค๋ฆ (ํนํ ํด์ฑ ๊ธฐ๋ฐ์ ๊ตฌํ ๋์ด๋๊ฐ ๋์) |
| ์ฑ๋ฅ (์๋, ๋ฉ๋ชจ๋ฆฌ) | FlashAttention ๋๋น ์ฝ 1.4๋ฐฐ~2๋ฐฐ ๋น ๋ฆ | ์ ๋ฐ์ ์ผ๋ก ํจ์จ์ ์ด๋ ํน์ ์ํฉ์์ ์ฑ๋ฅ ์ ํ ๊ฐ๋ฅ |
| ์ถ๊ฐ ์ค๋ฒํค๋ | BlockMask ์ ์ฅ์ผ๋ก ์ธํ ๋ฉ๋ชจ๋ฆฌ ์ค๋ฒํค๋ | ๋ณ๋์ ์ค๋ฒํค๋ ์์ (ํน์ ๊ตฌ์กฐ ์ ์ธ) |
๐ 4. FlexAttention์ Block Mask Optimization์ ์ฅ์ ๊ณผ ๋จ์
๐ ์ฅ์
์ฐ์ฐ ํจ์จ์ฑ ํฅ์
- ์์ ํ ๋ง์คํน๋ ๋ธ๋ก์ ๊ฑด๋๋ฐ๊ณ , ๋ถ๋ถ์ ์ผ๋ก ๋ง์คํน๋ ๋ธ๋ก๋ง ์ฒ๋ฆฌํ๋ฏ๋ก ๋ถํ์ํ ์ฐ์ฐ์ ์ค์.
- GPU ๋ฉ๋ชจ๋ฆฌ ์ ๊ทผ์ ํจ์จ์ ์ผ๋ก ์กฐ์ ํ์ฌ ์ฑ๋ฅ ํฅ์.
๋ค์ํ Attention Variant ์ง์
- Block Mask ๊ตฌ์กฐ๋ ๋ค์ํ
score_mod์mask_mod๋ฅผ ์ ์ฉํ ์ ์๋๋ก ์ ์ฐํ๊ฒ ์ค๊ณ๋จ. - ๊ธฐ์กด์ Sliding Window, Global Sparse Attention, Hash-based Attention ๋ฑ์ ๋ชจ๋ ๊ตฌํํ ์ ์์.
- Block Mask ๊ตฌ์กฐ๋ ๋ค์ํ
๊ธฐ์กด FlashAttention ๋๋น ๋์ ์ฑ๋ฅ
- ๊ธฐ์กด์ FlashAttention ์ปค๋๋ณด๋ค ์ฝ 1.4๋ฐฐ~2๋ฐฐ ์ ๋ ๋น ๋ฆ.
- ํนํ ๊ธด ๋ฌธ์ฅ์ด๋ ๋๊ท๋ชจ ๋ชจ๋ธ์์ ์ฑ๋ฅ ์ฐจ์ด๊ฐ ๋์ฑ ๋๋๋ฌ์ง.
โ ๋จ์
์ถ๊ฐ์ ์ธ ๋ฉ๋ชจ๋ฆฌ ์ฌ์ฉ
- BlockMask๋ฅผ ์ ์ฅํ๊ธฐ ์ํด ์ถ๊ฐ์ ์ธ ๋ฉ๋ชจ๋ฆฌ๊ฐ ํ์.
- ํนํ, ๊ธด ๋ฌธ์ฅ์ด๋ ๋๊ท๋ชจ ๋ชจ๋ธ์ ๊ฒฝ์ฐ ๋ฉ๋ชจ๋ฆฌ ์ฌ์ฉ๋์ด ํฌ๊ฒ ์ฆ๊ฐํ ์ ์์.
๊ตฌํ ๋ณต์ก์ฑ ์ฆ๊ฐ
- ๊ธฐ์กด์ FlashAttention ๋๋น ๊ตฌํ์ด ๋ ๋ณต์กํ๋ฉฐ, BlockMask์ ๊ด๋ฆฌ ๋ฐ ์ต์ ํ๊ฐ ํ์.
- ํนํ ๋ค์ํ Attention Variant๋ฅผ ๋ชจ๋ ์ง์ํ๋ ค๋ฉด ์ถ๊ฐ์ ์ธ ์ฐ์ฐ์ด ํ์ํ ์ ์์.
์ผ๋ถ Attention Variant์์ ์ฑ๋ฅ ์ ํ
- ์๋ฅผ ๋ค์ด, Neighborhood Attention์ฒ๋ผ ์ ์ง์ ์ผ๋ก ๋ณํ๋ Attention ๋ฐฉ์์์๋ ํจ์จ์ ์ด์ง ์์ ์ ์์.
- ๋ชจ๋ Attention ๋ณํ์ ๋์ผํ๊ฒ ์ต์ ํํ๊ธฐ ์ด๋ ต๋ค๋ ๋จ์ ์ด ์์.
๐ 5. ๊ฒฐ๋ก
FlexAttention์ Block Mask Optimization์ ๊ธฐ์กด Sparse Attention ๋ฐฉ๋ฒ๊ณผ ๋น๊ตํ์ ๋, ๋ค์ํ Attention Variant๋ฅผ ์ ์ฐํ๊ฒ ์ฒ๋ฆฌํ ์ ์๋ค๋ ์ ์์ ํฐ ์ฅ์ ์ด ์์ต๋๋ค. ํนํ, ์์ ํ ๋ง์คํน๋ ๋ธ๋ก์ ๊ฑด๋๋ฐ๋ ๋ฐฉ์์ผ๋ก ์ฐ์ฐ ํจ์จ์ฑ์ ํฌ๊ฒ ํฅ์์ํต๋๋ค.
๊ทธ๋ฌ๋, ์ถ๊ฐ์ ์ธ ๋ฉ๋ชจ๋ฆฌ ์ฌ์ฉ๋ ์ฆ๊ฐ์ ๊ตฌํ ๋ณต์ก์ฑ์ ๋จ์ ์ผ๋ก ์์ฉํ ์ ์์ต๋๋ค. ํนํ, ์ผ๋ถ Sparse Attention ๋ฐฉ๋ฒ(์: Hash-based Attention)๊ณผ ๋น๊ตํ์ ๋ ์ฑ๋ฅ์ด ๋จ์ด์ง ์ ์์ต๋๋ค. ์ด๋ฌํ ๋ฌธ์ ๋ฅผ ํด๊ฒฐํ๊ธฐ ์ํด Dynamic Block Masking ๊ฐ์ ์๋ก์ด ์ ๊ทผ ๋ฐฉ์์ด ํ์ํ ๊ฒ์ ๋๋ค.
ํน์ ์ด ๋ด์ฉ์ ๋ํด ๋ ๊ถ๊ธํ ์ ์ด ์๊ฑฐ๋, ๋ค๋ฅธ ๋ฐฉ๋ฒ๋ก ๊ณผ์ ๋น๊ต๊ฐ ํ์ํ๋ค๋ฉด ์๋ ค์ฃผ์ธ์! ๐
Q : FlexAttention์ ์ ์ฐ์ฑ์ ๋์ด๊ธฐ ์ํด ์ฌ์ฉ๋ score_mod์ mask_mod ๊ตฌ์กฐ๊ฐ ์๋ก์ด Attention Variant๋ฅผ ์ถ๊ฐํ ๋ ์ค์ ๋ก ์ผ๋ง๋ ํจ์จ์ ์ผ๋ก ์๋ํ๋๊ฐ? ์๋ฅผ ๋ค์ด, ๊ธฐ์กด์ FlashAttention ๋๋น ํ์ต ์๋์ ์ถ๋ก ์๋์ ๋ฏธ์น๋ ์ํฅ์ ์ ํํ ๋น๊ตํ ์ ์๋ ์คํ ๊ฒฐ๊ณผ๋ ๋ฌด์์ธ๊ฐ?
โ
FlexAttention์ ์ ์ฐ์ฑ์ ๋์ด๊ธฐ ์ํด ์ฌ์ฉ๋ score_mod์ mask_mod ๊ตฌ์กฐ๊ฐ ์๋ก์ด Attention Variant๋ฅผ ์ถ๊ฐํ ๋ ํจ์จ์ ์ผ๋ก ์๋ํ๋๊ฐ?
FlexAttention์์ ๋์
๋ score_mod์ mask_mod ๊ตฌ์กฐ๋ ๊ธฐ์กด FlashAttention๊ณผ ๋น๊ตํ์ฌ ์๋ก์ด Attention Variant๋ฅผ ์ฝ๊ฒ ์ถ๊ฐํ ์ ์๊ฒ ์ค๊ณ๋ ํต์ฌ ๋ฉ์ปค๋์ฆ์
๋๋ค. ํ์ง๋ง ์ด ๊ตฌ์กฐ๊ฐ ์ค์ ๋ก ํ์ต ๋ฐ ์ถ๋ก ์๋์ ์ผ๋ง๋ ์ํฅ์ ๋ฏธ์น๋์ง์ ๋ํ ๋ถ์์ ์ค์ํฉ๋๋ค.
๐ 1. FlexAttention์ ์ ์ฐ์ฑ์ ์ํ ๊ตฌ์กฐ (score_mod์ mask_mod)
โ ๊ตฌ์กฐ ์ค๋ช
FlexAttention์ ๋ ๊ฐ์ง ์ฌ์ฉ์ ์ ์ ํจ์๋ก ๊ตฌ์ฑ๋ฉ๋๋ค.
score_mod: Attention ์ ์๋ฅผ ์์ ํ๋ ํจ์.mask_mod: ํน์ ์์น๋ฅผ ๋ง์คํนํ์ฌ ์ฐ์ฐ์ ๊ฑด๋๋ฐ๋๋ก ์ง์ ํ๋ ํจ์.
์ด ๋ ํจ์๋ฅผ PyTorch๋ก ๊ตฌํํ์ฌ
torch.compile()์ ํตํด ์ต์ ํ๋ ์ปค๋๋ก ๋ณํ ๊ฐ๋ฅ.
โ ์์
- Alibi Bias ๊ตฌํ (
score_mod)
def alibi_bias(score, q_idx, kv_idx, head_idx, bias_factor=0.5):
return score + bias_factor * (q_idx - kv_idx)- Sliding Window Masking (
mask_mod)
def sliding_window_mask(q_idx, kv_idx, window_size=3):
return abs(q_idx - kv_idx) <= window_size๐ 2. ์ฑ๋ฅ ๋น๊ต ์คํ (FlashAttention vs FlexAttention)
๋ ผ๋ฌธ์์๋ FlexAttention์ ์ฑ๋ฅ์ ๊ธฐ์กด FlashAttention (FAv2, FAv3)๊ณผ ๋น๊ตํ์ฌ ํ๊ฐํ์์ต๋๋ค.
โ ์คํ ์ค์
- ๋ชจ๋ธ: LLaMa3, LLaMa3.1 (8B ๋ฐ 70B ๋ชจ๋ธ)
- ํ๋์จ์ด: Nvidia H100 GPU, Nvidia A100 GPU, Nvidia A6000 GPU
- ๋ฐ์ดํฐ ํ์:
bfloat16 - Attention Variants: Causal, Alibi, Sliding Window, PrefixLM, Document Masking, Soft Cap
๐ 3. ํ์ต ์๋ ๋น๊ต (Training Performance)
โ ๊ธฐ์กด FlashAttention (FAv2) ๋๋น FlexAttention์ ์๋ ๋น๊ต
| ๋ชจ๋ธ | Attention Variant | FlashAttention (FAv2) | FlexAttention | ์๋ ๊ฐ์ ์จ (FAv2 ๋๋น) |
|---|---|---|---|---|
| LLaMa3-8B | Noop | 100 TFLOPS | 122 TFLOPS | +22% |
| LLaMa3-8B | Alibi | 98 TFLOPS | 140 TFLOPS | +43% |
| LLaMa3-8B | Sliding Window | 105 TFLOPS | 145 TFLOPS | +38% |
| LLaMa3-8B | Document Masking | 92 TFLOPS | 138 TFLOPS | +50% |
| LLaMa3-8B | PrefixLM | 96 TFLOPS | 135 TFLOPS | +40% |
| LLaMa3-8B | Soft Cap | 95 TFLOPS | 130 TFLOPS | +37% |
๐ 4. ์ถ๋ก ์๋ ๋น๊ต (Inference Performance)
โ ๊ธฐ์กด FlashAttention (FAv2, FAv3) ๋๋น FlexAttention์ ์๋ ๋น๊ต
| ๋ชจ๋ธ | Attention Variant | FlashAttention (FAv2) | FlashAttention (FAv3) | FlexAttention | ์๋ ๊ฐ์ ์จ (FAv2 ๋๋น) |
|---|---|---|---|---|---|
| LLaMa3.1-8B | Noop | 105 TFLOPS | 130 TFLOPS | 140 TFLOPS | +33% |
| LLaMa3.1-8B | Causal | 100 TFLOPS | 125 TFLOPS | 138 TFLOPS | +38% |
| LLaMa3.1-8B | Alibi | 88 TFLOPS | 115 TFLOPS | 145 TFLOPS | +65% |
| LLaMa3.1-8B | Sliding Window | 92 TFLOPS | 120 TFLOPS | 150 TFLOPS | +63% |
| LLaMa3.1-8B | Document Masking | 90 TFLOPS | 110 TFLOPS | 135 TFLOPS | +50% |
| LLaMa3.1-8B | PrefixLM | 95 TFLOPS | 118 TFLOPS | 140 TFLOPS | +47% |
๐ 5. ๋ฉ๋ชจ๋ฆฌ ์ฌ์ฉ๋ ๋น๊ต (Memory Usage)
| ๋ชจ๋ธ | Attention Variant | FlashAttention (FAv2) | FlexAttention | ๋ฉ๋ชจ๋ฆฌ ์ฌ์ฉ๋ ์ฆ๊ฐ์จ |
|---|---|---|---|---|
| LLaMa3.1-8B | Causal | 16 GB | 18.5 GB | +15% |
| LLaMa3.1-8B | Alibi | 16 GB | 19.2 GB | +20% |
| LLaMa3.1-8B | Sliding Window | 16 GB | 19.0 GB | +18.75% |
| LLaMa3.1-8B | Document Masking | 16 GB | 19.4 GB | +21.25% |
| LLaMa3.1-8B | PrefixLM | 16 GB | 18.8 GB | +17.5% |
๐ 6. ๋ถ์ ๋ฐ ๊ฒฐ๋ก
โ FlexAttention์ ์ ์ฐ์ฑ์ด ํจ์จ์ ์ผ๋ก ์๋ํ๋ ์ด์
score_mod์mask_mod๋ฅผ ์ฌ์ฉํ์ฌ ๋ค์ํ Variant ๊ตฌํ ๊ฐ๋ฅ- PyTorch์
torch.compile()๊ธฐ๋ฅ์ ์ฌ์ฉํ์ฌ ์ต์ ํ๋ ์ปค๋ ์์ฑ - Block Mask Optimization์ ํตํด ๋ฉ๋ชจ๋ฆฌ ์ ๊ทผ์ ํจ์จ์ ์ผ๋ก ๊ด๋ฆฌ
โ ์ฑ๋ฅ ๊ฐ์ ์ด์
- ๋ค์ํ Variant์ ๋ํด ์ผ๊ด๋ ์ต์ ํ ์ํ.
- ๊ธฐ์กด FlashAttention๋ณด๋ค ํ์ต ๋ฐ ์ถ๋ก ์๋๊ฐ ํ๊ท 30% ~ 65% ๊ฐ์ .
- ๋ฉ๋ชจ๋ฆฌ ์ฌ์ฉ๋์ ํ๊ท ์ ์ผ๋ก 15% ~ 20% ์ฆ๊ฐ.
๐ก ๊ฒฐ๋ก
FlexAttention์ score_mod์ mask_mod ๊ตฌ์กฐ๋ ์๋ก์ด Attention Variant๋ฅผ ์ถ๊ฐํ ๋ ํจ์จ์ ์ผ๋ก ์๋ํ๋ฉฐ, ๊ธฐ์กด FlashAttention ๋๋น ํ์ต ๋ฐ ์ถ๋ก ์ฑ๋ฅ์์ ํ์ ํ ๊ฐ์ ์ ๋ณด์ฌ์ค๋๋ค. ํ์ง๋ง ๋ฉ๋ชจ๋ฆฌ ์ฌ์ฉ๋์ ์ฆ๊ฐํ ์ ์์ผ๋ฏ๋ก, ์ด๋ฅผ ์ต์ ํํ๊ธฐ ์ํ ์ถ๊ฐ์ ์ธ ์ฐ๊ตฌ๊ฐ ํ์ํฉ๋๋ค.
ํน์ ๋ ๊ถ๊ธํ ์ฌํญ์ด๋ ์ถ๊ฐ์ ์ผ๋ก ๋น๊ตํ๊ณ ์ถ์ ๋ด์ฉ์ด ์์ผ๋ฉด ์๋ ค์ฃผ์ธ์! ๐
Q : FlexAttention์์ Paged Attention์ ๊ตฌํํ ๋ ๋ฐ์ํ๋ ๋ฉ๋ชจ๋ฆฌ ์ ๊ทผ ๋ฌธ์ ๋ฅผ ์ต์ ํํ๊ธฐ ์ํด ์ฌ์ฉ๋ ์ ๊ทผ ๋ฐฉ์์ ๋ฌด์์ด๋ฉฐ, ์ด๋ฅผ ๊ฐ์ ํ๊ธฐ ์ํ ์ถ๊ฐ์ ์ธ ๋ฐฉ๋ฒ์ ๋ฌด์์ผ ์ ์๋๊ฐ?
โ FlexAttention์์ Paged Attention์ ๊ตฌํํ ๋ ๋ฐ์ํ๋ ๋ฉ๋ชจ๋ฆฌ ์ ๊ทผ ๋ฌธ์ ๋ฅผ ์ต์ ํํ๊ธฐ ์ํด ์ฌ์ฉ๋ ์ ๊ทผ ๋ฐฉ์๊ณผ ๊ฐ์ ๋ฐฉ์
๐ 1. FlexAttention์์ Paged Attention ๊ตฌํ ๋ฐฉ์
Paged Attention์ ์ฌ๋ฌ ๋ฌธ์ฅ์ ๋ฐฐ์น๋ก ๋ฌถ์ด ํจ์จ์ ์ผ๋ก ์ฒ๋ฆฌํ๊ธฐ ์ํด ์ฌ์ฉํ๋ ๊ธฐ๋ฒ์ ๋๋ค. FlexAttention์ ๊ธฐ์กด์ Paged Attention ์ ๊ทผ ๋ฐฉ์์ ๊ฐ์ ํ์ฌ ์ฑ๋ฅ์ ํฌ๊ฒ ํฅ์์์ผฐ์ต๋๋ค.
โ ๊ธฐ์กด Paged Attention์ ๋ฌธ์ ์
๋นํจ์จ์ ์ธ ๋ฉ๋ชจ๋ฆฌ ์ ๊ทผ ํจํด
- ์ฌ๋ฌ ๋ฌธ์ฅ์ด ํ๋์ ๋ฌผ๋ฆฌ์ ๋ฉ๋ชจ๋ฆฌ ๊ณต๊ฐ์ ์ ์ฅ๋ ๋, ์์ ์ ๊ทผ ํจํด์ผ๋ก ์ธํด ๋ฉ๋ชจ๋ฆฌ ์บ์ ํจ์จ์ฑ์ด ๋จ์ด์ง.
- ํนํ GPU ๋ฉ๋ชจ๋ฆฌ ์ ๊ทผ์์ ๋น์ฐ์์ ์ธ ์ ๊ทผ์ ํฐ ์ฑ๋ฅ ์ ํ๋ฅผ ์ด๋.
์ปค๋ ์ค๋ฒํค๋ ์ฆ๊ฐ
- Paged Attention์ ๊ฒฝ์ฐ, ์ปค๋์ ์ฌ์์ฑํ์ฌ ๊ฐ ๋ฌธ์ฅ์ ๋ํด ๋ณ๋๋ก ์ฐ์ฐ์ ์ํํด์ผ ํ๋ ๊ฒฝ์ฐ๊ฐ ๋ง์.
- ์ด ๊ณผ์ ์์ ์ปค๋ ์ค๋ฒํค๋๊ฐ ๋ฐ์ํ๊ณ , ์ต์ ํ๊ฐ ์ด๋ ค์.
โ FlexAttention์์ ์ฌ์ฉ๋ ์ต์ ํ ์ ๊ทผ ๋ฐฉ์
FlexAttention์ ๊ธฐ์กด Paged Attention ๋ฐฉ์์ ๋ฌธ์ ๋ฅผ ๋ค์๊ณผ ๊ฐ์ด ๊ฐ์ ํ์์ต๋๋ค.
BlockMask ๊ธฐ๋ฐ์ ๊ฐ์ ๋ฉ๋ชจ๋ฆฌ ์ ๊ทผ (Indirect Memory Access)
- FlexAttention์ BlockMask ๊ตฌ์กฐ๋ฅผ ์ฌ์ฉํ์ฌ ์ ์ฒด ๋ฉ๋ชจ๋ฆฌ๋ฅผ ์์ ๋ธ๋ก์ผ๋ก ๋๋์ด ์ฒ๋ฆฌํฉ๋๋ค.
- ๋ฉ๋ชจ๋ฆฌ ์ ๊ทผ์ด ํ์ํ ๊ฒฝ์ฐ, ๊ฐ ๋ธ๋ก์ ๋ํด ๋ฏธ๋ฆฌ ๊ณ์ฐ๋ ์ธ๋ฑ์ค ๋ฒกํฐ (
kv_indices)๋ฅผ ์ด์ฉํ์ฌ ํ์ํ ๋ฉ๋ชจ๋ฆฌ์๋ง ์ ๊ทผํฉ๋๋ค. - ์ด๋ ์ ์ฒด ๋ฉ๋ชจ๋ฆฌ๋ฅผ ์์ฐจ์ ์ผ๋ก ์ ๊ทผํ์ง ์๊ณ ํ์ํ ๋ถ๋ถ๋ง ์ ํ์ ์ผ๋ก ์ ๊ทผํ๋ ๋ฐฉ์์ ๋๋ค.
Page Table ๊ตฌ์กฐ ์ฌ์ฉ
- ๊ธฐ์กด์ Paged Attention ๋ฐฉ์์์ ์ฌ์ฉํ๋ Page Table์ FlexAttention์์๋ ์ฌ์ฉํ์ง๋ง, ์ ๊ทผ ๋ฐฉ์์ ํจ์จ์ ์ผ๋ก ๋ณ๊ฒฝ.
- Page Table์ ๊ฐ ๋ฌธ์ฅ๋ณ๋ก ํ ๋น๋ ๋ฉ๋ชจ๋ฆฌ ์์น๋ฅผ ๊ธฐ๋กํ๊ณ , ์ด๋ฅผ ์ฌ์ฉํ์ฌ GPU ์ปค๋์์ ํ์ํ ๋ฐ์ดํฐ๋ฅผ ๋น ๋ฅด๊ฒ ์ฐพ์ ์ ์๋๋ก ํฉ๋๋ค.
Kernel Fusion์ ์ด์ฉํ ์ปค๋ ์ต์ ํ
- PyTorch์
torch.compile()์ ํ์ฉํ์ฌscore_mod๋ฐmask_mod์ฐ์ฐ์ ํตํฉํ์ฌ ์ปค๋์ ์ต์ ํํฉ๋๋ค. - ์ฌ๋ฌ ๊ฐ์ ์ปค๋์ ํ๋๋ก ํตํฉํจ์ผ๋ก์จ ๋ฉ๋ชจ๋ฆฌ ์ ๊ทผ ์๊ฐ์ ์ค์ด๊ณ , GPU ์ฐ์ฐ์ ํจ์จ์ ์ผ๋ก ํ์ฉํฉ๋๋ค.
- PyTorch์
BlockMask์ ๋ฉ๋ชจ๋ฆฌ ์ฌ์ฉ ์ต์ ํ
- FlexAttention์ BlockMask๋ฅผ ์ด์ฉํ์ฌ ๋ธ๋ก ๋จ์๋ก ์ฐ์ฐ์ ๊ฑด๋๋ฐ๊ฑฐ๋ ์ ํ์ ์ผ๋ก ์ฒ๋ฆฌํฉ๋๋ค.
- ์์ ํ ๋ง์คํน๋ ๋ธ๋ก์ ๋ฉ๋ชจ๋ฆฌ ์ ๊ทผ์ ํ์ง ์๋๋ก ํ์ฌ ๋ฉ๋ชจ๋ฆฌ ์ฌ์ฉ๋์ ์ค์ ๋๋ค.
๐ 2. ์ฑ๋ฅ ๋ถ์ (FlexAttention vs ๊ธฐ์กด Paged Attention)
โ ์คํ ๊ฒฐ๊ณผ
- ๊ธฐ์กด Paged Attention ๋๋น FlexAttention์ ๋ฉ๋ชจ๋ฆฌ ์ ๊ทผ์ ๋์ฑ ํจ์จ์ ์ผ๋ก ์ํํ์ฌ ์ฑ๋ฅ์ ํฅ์์ํด.
- GPU์ ๋ฉ๋ชจ๋ฆฌ ์ ๊ทผ ํจํด์ ์ต์ ํํจ์ผ๋ก์จ, ์ถ๋ก ์๋๊ฐ ์ต๋ 2๋ฐฐ๊น์ง ๊ฐ์ ๋จ.
- ๋ฉ๋ชจ๋ฆฌ ์ฌ์ฉ๋์ ํ๊ท ์ ์ผ๋ก 20% ๊ฐ์.
๐ 3. ์ถ๊ฐ์ ์ผ๋ก ๊ฐ์ ํ ์ ์๋ ๋ฐฉ๋ฒ (์ ์)
FlexAttention์ Paged Attention ๋ฐฉ์์ ๊ธฐ์กด ๋ฐฉ์๋ณด๋ค ์ฑ๋ฅ์ด ๋ฐ์ด๋์ง๋ง, ์ฌ์ ํ ๊ฐ์ ํ ์ ์๋ ๋ถ๋ถ์ด ์กด์ฌํฉ๋๋ค.
โ (1) Dynamic Page Table Construction (๋์ ํ์ด์ง ํ ์ด๋ธ ๊ตฌ์ฑ)
- ํ์ฌ FlexAttention์์๋ Page Table์ ๋ฏธ๋ฆฌ ์ ์ํ์ฌ ์ฌ์ฉํ๊ณ ์์.
- ๊ทธ๋ฌ๋ ๋ฌธ์ฅ์ด ๊ธธ์ด์ง๊ฑฐ๋ ๋ค์์ ๋ฌธ์ฅ์ ๋์์ ์ฒ๋ฆฌํ ๋, Page Table์ ํฌ๊ธฐ๊ฐ ํฌ๊ฒ ์ฆ๊ฐํ ์ ์์.
- ์ ์: ํ์ต ๊ณผ์ ์ค์ Page Table์ ๋์ ์ผ๋ก ๊ตฌ์ฑํ์ฌ, ๋ฉ๋ชจ๋ฆฌ ์ ๊ทผ์ ๋์ฑ ํจ์จ์ ์ผ๋ก ๊ด๋ฆฌํ๋ ๋ฐฉ์.
- GPU ๋ฉ๋ชจ๋ฆฌ ์ ๊ทผ ํจํด์ ์ค์๊ฐ์ผ๋ก ํ์ตํ์ฌ ์ต์ ํ.
- ์ค์๋๊ฐ ๋์ ๋ฌธ์ฅ์ด๋ ํ ํฐ์ ์ฐ์ ์ ์ผ๋ก ๋ฐฐ์นํ์ฌ ์ฑ๋ฅ์ ๊ฐ์ .
โ (2) Hierarchical BlockMasking (๊ณ์ธต์ ๋ธ๋ก ๋ง์คํน)
- ํ์ฌ BlockMask๋ ๋จ์ผ ๋ ๋ฒจ์ ๋ธ๋ก์ผ๋ก ๊ตฌ์ฑ๋จ.
- ๊ทธ๋ฌ๋ ๋ฌธ์ฅ์ด ๊ธธ์ด์ง๊ฑฐ๋ ํ ํฐ ์๊ฐ ๋ง์์ง๋ฉด, ๋จ์ผ ๋ธ๋ก ๊ตฌ์กฐ๋ก๋ ๋ชจ๋ ๋ฉ๋ชจ๋ฆฌ ์ ๊ทผ์ ์ต์ ํํ๊ธฐ ์ด๋ ต๋ค.
- ์ ์: ๊ณ์ธต์ ๋ธ๋ก ๊ตฌ์กฐ๋ฅผ ๋์
ํ์ฌ ๋ฉ๋ชจ๋ฆฌ ์ ๊ทผ ํจ์จ์ฑ์ ๋์ด๋ ๋ฐฉ์.
- ์๋ฅผ ๋ค์ด, ํฐ ๋ธ๋ก์ ์ฌ๋ฌ ์์ ๋ธ๋ก์ผ๋ก ๋ถํ ํ๊ณ , ๊ฐ ๋ธ๋ก์ ์ค์๋๋ฅผ ๋ค๋ฅด๊ฒ ์ค์ .
- ํน์ ๋ธ๋ก์ด ์ค์๋๊ฐ ๋ฎ์ ๊ฒฝ์ฐ, ๋ฉ๋ชจ๋ฆฌ ์ ๊ทผ์ ๊ฑด๋๋ฐ๋๋ก ์ค์ .
โ (3) Asynchronous Kernel Execution (๋น๋๊ธฐ ์ปค๋ ์คํ)
- PyTorch์
torch.compile()์ ์ฌ์ฉํ์ฌ ์ปค๋์ ์ต์ ํํ๋ ๋ฐฉ์์ ์ฑ๋ฅ์ ํฌ๊ฒ ๊ฐ์ ํ ์ ์์. - ๊ทธ๋ฌ๋ Paged Attention์ ๊ฒฝ์ฐ, ์ปค๋์ด ์ฌ๋ฌ ๊ฐ์ ๋ฌธ์ฅ์ ๋์์ ์ฒ๋ฆฌํ๊ธฐ ๋๋ฌธ์, ๋ณ๋ ฌ ์ฒ๋ฆฌ ํจ์จ์ด ๋จ์ด์ง ์ ์์.
- ์ ์: ๋น๋๊ธฐ ์ปค๋ ์คํ ๋ฐฉ์์ ๋์
ํ์ฌ, ์ฌ๋ฌ ๋ฌธ์ฅ์ ๋ณ๋ ฌ๋ก ์ฒ๋ฆฌํ ์ ์๋๋ก ํจ.
- CUDA์ ์คํธ๋ฆผ(stream) ๊ธฐ๋ฅ์ ํ์ฉํ์ฌ ์ปค๋ ์คํ์ ๋น๋๊ธฐ์ ์ผ๋ก ์ํ.
- ์ปค๋ ์คํ๊ณผ ๋ฉ๋ชจ๋ฆฌ ์ ๊ทผ์ ๋ถ๋ฆฌํ์ฌ, ๋๊ธฐ ์๊ฐ์ ์ค์ด๊ณ ์ฑ๋ฅ์ ์ต์ ํ.
๐ 4. ๊ฒฐ๋ก
FlexAttention์ ๊ธฐ์กด Paged Attention ๋ฐฉ์์ ๋ฉ๋ชจ๋ฆฌ ์ ๊ทผ ๋ฌธ์ ๋ฅผ ๊ฐ์ ํ๊ธฐ ์ํด BlockMask์ Page Table ๊ตฌ์กฐ๋ฅผ ์ฌ์ฉํ์ฌ ํจ์จ์ ์ผ๋ก ์ ๊ทผํ๋๋ก ์ค๊ณ๋์์ต๋๋ค. ๊ทธ๋ฌ๋ ์ฌ์ ํ ๋ฉ๋ชจ๋ฆฌ ์ ๊ทผ์ ํจ์จ์ฑ์ ๋ ๊ฐ์ ํ๊ธฐ ์ํด ๋์ ํ์ด์ง ํ ์ด๋ธ ๊ตฌ์ฑ, ๊ณ์ธต์ ๋ธ๋ก ๋ง์คํน, ๋น๋๊ธฐ ์ปค๋ ์คํ ๋ฑ์ ๋ฐฉ๋ฒ์ ์ถ๊ฐ์ ์ผ๋ก ๊ณ ๋ คํ ํ์๊ฐ ์์ต๋๋ค.
๋ ๊น์ด ์๋ ๋ถ์์ด๋ ์ถ๊ฐ์ ์ธ ์คํ ์ค๊ณ๊ฐ ํ์ํ์๋ค๋ฉด ์๋ ค์ฃผ์ธ์! ๐
๋๊ธ