Skip to content

Commit e4ea3a7

Browse files
Optimize attention blocking nested loops (quic#957)
changed the code from doing the exact same math repeatedly. Signed-off-by: Anuj Gupta <anujgupt@qti.qualcomm.com>
1 parent ec243d5 commit e4ea3a7

1 file changed

Lines changed: 10 additions & 6 deletions

File tree

QEfficient/blocking/blocking_configurator.py

Lines changed: 10 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -227,16 +227,20 @@ def update_best_config(num_q_blocks: int, num_kv_blocks: int, q_kv_ratio: float,
227227
best_config["q_kv_ratio"] = q_kv_ratio
228228
best_config["vtcm_footprint"] = footprint
229229

230-
for num_q_blocks in num_q_blocks_list:
231-
for num_kv_blocks in num_kv_blocks_list:
232-
q_sl_per_nsp = math.ceil(seq_len / num_nsps / num_q_blocks)
233-
q_size_per_nsp = num_heads_per_iter * bs * q_sl_per_nsp * head_dim * data_bytes
230+
kv_metrics = []
231+
for num_kv_blocks in num_kv_blocks_list:
232+
kv_cl_per_nsp = math.ceil(ctx_len / num_kv_blocks)
233+
kv_size_per_nsp = num_heads_per_iter * bs * kv_cl_per_nsp * head_dim * data_bytes
234+
kv_metrics.append((num_kv_blocks, kv_cl_per_nsp, kv_size_per_nsp))
234235

235-
kv_cl_per_nsp = math.ceil(ctx_len / num_kv_blocks)
236-
kv_size_per_nsp = num_heads_per_iter * bs * kv_cl_per_nsp * head_dim * data_bytes
236+
for num_q_blocks in num_q_blocks_list:
237+
q_sl_per_nsp = math.ceil(seq_len / num_nsps / num_q_blocks)
238+
q_size_per_nsp = num_heads_per_iter * bs * q_sl_per_nsp * head_dim * data_bytes
237239

240+
for num_kv_blocks, kv_cl_per_nsp, kv_size_per_nsp in kv_metrics:
238241
qk_size_per_nsp = num_heads_per_iter * bs * q_sl_per_nsp * kv_cl_per_nsp * data_bytes
239242
vtcm_footprint = q_size_per_nsp + kv_size_per_nsp + qk_size_per_nsp
243+
240244
q_kv_ratio = max(q_size_per_nsp / kv_size_per_nsp, kv_size_per_nsp / q_size_per_nsp)
241245
num_total_blocks = num_q_blocks * num_kv_blocks
242246

0 commit comments

Comments
 (0)