143 lines
3.9 KiB
YAML
143 lines
3.9 KiB
YAML
# Standard attention decode benchmark configuration
|
|
# Sweeps num_q_heads and num_kv_heads to isolate effects of:
|
|
# 1. GQA ratio (fixed num_q_heads=32, vary num_kv_heads)
|
|
# 2. Absolute head count (fixed 4:1 ratio, vary scale)
|
|
|
|
model:
|
|
num_layers: 32
|
|
num_q_heads: 32 # Base value, overridden by sweep
|
|
num_kv_heads: 8 # Base value, overridden by sweep
|
|
head_dim: 128
|
|
block_size: 16
|
|
|
|
# Head count sweep: each entry overrides num_q_heads, num_kv_heads, and
|
|
# head_dim where it differs from the base (128). Head counts are per-GPU
|
|
# (i.e. after TP sharding).
|
|
#
|
|
# Group A — vary GQA ratio (fixed q=32, head_dim=128):
|
|
# 32:32 (MHA), 32:8 (GQA 4:1), 32:4 (GQA 8:1), 32:1 (MQA)
|
|
#
|
|
# Groups B-E — real model configs at various TP degrees:
|
|
# Model head_dim Full TP2 TP4 TP8
|
|
# Llama 3 8B 128 32:8 16:4 8:2 4:1
|
|
# Llama 3 70B 128 64:8 32:4 16:2 8:1
|
|
# GPT-OSS 120B 64 64:8 32:4 16:2 8:1
|
|
# Llama 3 405B 128 128:8 64:4 32:2 16:1
|
|
model_parameter_sweep:
|
|
values:
|
|
# --- head_dim=128 (Llama 3 family) ---
|
|
- { num_q_heads: 32, num_kv_heads: 32, head_dim: 128 } # MHA 1:1
|
|
- { num_q_heads: 32, num_kv_heads: 1, head_dim: 128 } # MQA 32:1
|
|
- { num_q_heads: 4, num_kv_heads: 1, head_dim: 128 } # Llama 3 8B TP8
|
|
- { num_q_heads: 8, num_kv_heads: 2, head_dim: 128 } # Llama 3 8B TP4
|
|
- { num_q_heads: 16, num_kv_heads: 4, head_dim: 128 } # Llama 3 8B TP2
|
|
- { num_q_heads: 32, num_kv_heads: 8, head_dim: 128 } # Llama 3 8B TP1 / GQA 4:1
|
|
- { num_q_heads: 8, num_kv_heads: 1, head_dim: 128 } # Llama 3 70B TP8
|
|
- { num_q_heads: 16, num_kv_heads: 2, head_dim: 128 } # Llama 3 70B TP4
|
|
- { num_q_heads: 32, num_kv_heads: 4, head_dim: 128 } # Llama 3 70B TP2 / GQA 8:1
|
|
- { num_q_heads: 64, num_kv_heads: 8, head_dim: 128 } # Llama 3 70B TP1
|
|
- { num_q_heads: 16, num_kv_heads: 1, head_dim: 128 } # Llama 3 405B TP8
|
|
- { num_q_heads: 32, num_kv_heads: 2, head_dim: 128 } # Llama 3 405B TP4
|
|
- { num_q_heads: 64, num_kv_heads: 4, head_dim: 128 } # Llama 3 405B TP2
|
|
- { num_q_heads: 128, num_kv_heads: 8, head_dim: 128 } # Llama 3 405B TP1
|
|
# --- head_dim=64 (GPT-OSS 120B) ---
|
|
- { num_q_heads: 8, num_kv_heads: 1, head_dim: 64 } # GPT-OSS 120B TP8
|
|
- { num_q_heads: 16, num_kv_heads: 2, head_dim: 64 } # GPT-OSS 120B TP4
|
|
- { num_q_heads: 32, num_kv_heads: 4, head_dim: 64 } # GPT-OSS 120B TP2
|
|
- { num_q_heads: 64, num_kv_heads: 8, head_dim: 64 } # GPT-OSS 120B TP1
|
|
label_format: "{backend}_q{num_q_heads}kv{num_kv_heads}d{head_dim}"
|
|
|
|
batch_specs:
|
|
# ---- batch_size x seq_len grid (decode: q_len=1) ----
|
|
# Small grid for quick iteration. Uncomment for full sweep.
|
|
|
|
# Batch size 1
|
|
- "q1s1k"
|
|
- "q1s512"
|
|
- "q1s2k"
|
|
- "q1s4k"
|
|
- "q1s8k"
|
|
- "q1s16k"
|
|
- "q1s32k"
|
|
|
|
# Batch size 2
|
|
- "2q1s512"
|
|
- "2q1s1k"
|
|
- "2q1s2k"
|
|
- "2q1s4k"
|
|
- "2q1s8k"
|
|
- "2q1s16k"
|
|
- "2q1s32k"
|
|
|
|
# Batch size 4
|
|
- "4q1s512"
|
|
- "4q1s1k"
|
|
- "4q1s2k"
|
|
- "4q1s4k"
|
|
- "4q1s8k"
|
|
- "4q1s16k"
|
|
- "4q1s32k"
|
|
|
|
# Batch size 8
|
|
- "8q1s1k"
|
|
- "8q1s512"
|
|
- "8q1s2k"
|
|
- "8q1s4k"
|
|
- "8q1s8k"
|
|
- "8q1s16k"
|
|
- "8q1s32k"
|
|
|
|
# Batch size 16
|
|
- "16q1s512"
|
|
- "16q1s1k"
|
|
- "16q1s2k"
|
|
- "16q1s4k"
|
|
- "16q1s8k"
|
|
- "16q1s16k"
|
|
- "16q1s32k"
|
|
|
|
# Batch size 32
|
|
- "32q1s512"
|
|
- "32q1s1k"
|
|
- "32q1s2k"
|
|
- "32q1s4k"
|
|
- "32q1s8k"
|
|
- "32q1s16k"
|
|
- "32q1s32k"
|
|
|
|
# Batch size 64
|
|
- "64q1s1k"
|
|
- "64q1s512"
|
|
- "64q1s2k"
|
|
- "64q1s4k"
|
|
- "64q1s8k"
|
|
- "64q1s16k"
|
|
- "64q1s32k"
|
|
|
|
# Batch size 128
|
|
- "128q1s512"
|
|
- "128q1s1k"
|
|
- "128q1s2k"
|
|
- "128q1s4k"
|
|
- "128q1s8k"
|
|
- "128q1s16k"
|
|
- "128q1s32k"
|
|
|
|
# Batch size 256
|
|
- "256q1s1k"
|
|
- "256q1s512"
|
|
- "256q1s2k"
|
|
- "256q1s4k"
|
|
- "256q1s8k"
|
|
- "256q1s16k"
|
|
- "256q1s32k"
|
|
|
|
# Available backends: FLASH_ATTN, TRITON_ATTN, FLASHINFER
|
|
backends:
|
|
- FLASH_ATTN
|
|
- TRITON_ATTN
|
|
- FLASHINFER
|
|
|
|
device: "cuda:0"
|
|
profile_memory: false
|