Skip to content

Commit 0f2c4fc

Browse files
authored
SDPA: use exp_u20 for the softmax exponential (pytorch#22082)
Summary: Switch the softmax exponential in the flash-attention custom SDPA kernel (`op_sdpa_impl.h`) from `Vectorized::exp()` to `Vectorized::exp_u20()`. The file already carried this as a TODO. This is not bit-exact. `exp_u20` is a ULP-20 approximation, so it differs from `exp` outright, and fourteen autoregressive layers on that is enough to flip a near-tie argmax. Reviewed By: JakeStevens Differential Revision: D117198987 Pull Request resolved: pytorch#22082
1 parent 96c621d commit 0f2c4fc

1 file changed

Lines changed: 1 addition & 3 deletions

File tree

‎extension/llm/custom_ops/op_sdpa_impl.h‎

Lines changed: 1 addition & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -593,9 +593,7 @@ _exp_reduce_sum_fusion_kernel(T1* a, const int& size, T2* out, T1& val) {
593593
for (int i = 0; i < vec_size * (size / vec_size); i += vec_size) {
594594
auto tmp0 = vec::VectorizedN<T1, 2>::loadu(a + i);
595595
auto tmp1 = tmp0 - vec_max;
596-
// Replace with exp_u20 later
597-
// auto tmp2 = tmp1.exp_u20();
598-
auto tmp2 = tmp1.exp();
596+
auto tmp2 = tmp1.exp_u20();
599597
vec_tmp_sum = vec_tmp_sum + tmp2;
600598
tmp2.store(out + i);
601599
}

0 commit comments

Comments
 (0)