-
Notifications
You must be signed in to change notification settings - Fork 0
⚡ Thunderbolt: Softmax - Hybrid 8x/4x unrolling for memory passes #91
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
base: main
Are you sure you want to change the base?
Changes from all commits
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change | ||||||
|---|---|---|---|---|---|---|---|---|
|
|
@@ -501,4 +501,139 @@ inline void softmax_v5(const float *input, float *output, std::size_t n) { | |||||||
| } | ||||||||
| } | ||||||||
|
|
||||||||
|
|
||||||||
| // ⚡ Thunderbolt: AVX2 Vectorized Softmax with Hybrid Unrolling | ||||||||
| // Target: AVX2 (Haswell+) | ||||||||
| // Reason: Multi-pass memory-bound kernels like Softmax benefit from 8x unrolling in the memory-bound passes (Max and Normalize) | ||||||||
| // to fully saturate memory bandwidth and hide the 4-cycle `_mm256_max_ps` latency. | ||||||||
| // However, unrolling the FMA-heavy exponentiation pass (Pass 2) 8x causes register spilling and port saturation. | ||||||||
| // A hybrid approach (Pass 1: 8x, Pass 2: 4x, Pass 3: 8x) optimally balances execution resources. | ||||||||
| // Expected gain: ~5-10% throughput improvement over softmax_v5 on large arrays. | ||||||||
| inline void softmax_v6(const float *input, float *output, std::size_t n) { | ||||||||
|
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. 📐 Maintainability & Code Quality | 🟠 Major | ⚡ Quick win Put the function opening brace on a separate line. Line 512 places the function opening brace on the declaration line. Move it to its own line. Proposed fix-inline void softmax_v6(const float *input, float *output, std::size_t n) {
+inline void softmax_v6(const float *input, float *output, std::size_t n)
+{As per coding guidelines, 📝 Committable suggestion
Suggested change
🤖 Prompt for AI AgentsSource: Coding guidelines |
||||||||
| if (n == 0) return; | ||||||||
|
|
||||||||
| // 1. Find max (8x unrolled) | ||||||||
| std::size_t i = 0; | ||||||||
| __m256 max_v = _mm256_set1_ps(std::numeric_limits<float>::lowest()); | ||||||||
| __m256 max0 = max_v, max1 = max_v, max2 = max_v, max3 = max_v; | ||||||||
| __m256 max4 = max_v, max5 = max_v, max6 = max_v, max7 = max_v; | ||||||||
|
|
||||||||
| for (; i + 63 < n; i += 64) { | ||||||||
| max0 = _mm256_max_ps(max0, _mm256_loadu_ps(input + i)); | ||||||||
| max1 = _mm256_max_ps(max1, _mm256_loadu_ps(input + i + 8)); | ||||||||
| max2 = _mm256_max_ps(max2, _mm256_loadu_ps(input + i + 16)); | ||||||||
| max3 = _mm256_max_ps(max3, _mm256_loadu_ps(input + i + 24)); | ||||||||
| max4 = _mm256_max_ps(max4, _mm256_loadu_ps(input + i + 32)); | ||||||||
| max5 = _mm256_max_ps(max5, _mm256_loadu_ps(input + i + 40)); | ||||||||
| max6 = _mm256_max_ps(max6, _mm256_loadu_ps(input + i + 48)); | ||||||||
| max7 = _mm256_max_ps(max7, _mm256_loadu_ps(input + i + 56)); | ||||||||
| } | ||||||||
|
|
||||||||
| max0 = _mm256_max_ps(max0, max4); | ||||||||
| max1 = _mm256_max_ps(max1, max5); | ||||||||
| max2 = _mm256_max_ps(max2, max6); | ||||||||
| max3 = _mm256_max_ps(max3, max7); | ||||||||
|
|
||||||||
| max0 = _mm256_max_ps(max0, max1); | ||||||||
| max2 = _mm256_max_ps(max2, max3); | ||||||||
| max0 = _mm256_max_ps(max0, max2); | ||||||||
|
|
||||||||
| // Remainder loop 8x elements | ||||||||
| for (; i + 7 < n; i += 8) { | ||||||||
| max0 = _mm256_max_ps(max0, _mm256_loadu_ps(input + i)); | ||||||||
| } | ||||||||
| float max_val = reduce_max(max0); | ||||||||
| for (; i < n; ++i) max_val = std::max(max_val, input[i]); | ||||||||
|
|
||||||||
| __m256 max_vec = _mm256_set1_ps(max_val); | ||||||||
|
|
||||||||
| // 2. Compute exp and sum (4x unrolled to avoid spilling) | ||||||||
| i = 0; | ||||||||
| __m256 sum0 = _mm256_setzero_ps(); | ||||||||
| __m256 sum1 = _mm256_setzero_ps(); | ||||||||
| __m256 sum2 = _mm256_setzero_ps(); | ||||||||
| __m256 sum3 = _mm256_setzero_ps(); | ||||||||
|
|
||||||||
| for (; i + 31 < n; i += 32) { | ||||||||
| __m256 x0 = _mm256_sub_ps(_mm256_loadu_ps(input + i), max_vec); | ||||||||
| __m256 x1 = _mm256_sub_ps(_mm256_loadu_ps(input + i + 8), max_vec); | ||||||||
| __m256 x2 = _mm256_sub_ps(_mm256_loadu_ps(input + i + 16), max_vec); | ||||||||
| __m256 x3 = _mm256_sub_ps(_mm256_loadu_ps(input + i + 24), max_vec); | ||||||||
|
|
||||||||
| __m256 e0 = exp256_ps_v2(x0); | ||||||||
| __m256 e1 = exp256_ps_v2(x1); | ||||||||
| __m256 e2 = exp256_ps_v2(x2); | ||||||||
| __m256 e3 = exp256_ps_v2(x3); | ||||||||
|
|
||||||||
| _mm256_storeu_ps(output + i, e0); | ||||||||
| _mm256_storeu_ps(output + i + 8, e1); | ||||||||
| _mm256_storeu_ps(output + i + 16, e2); | ||||||||
| _mm256_storeu_ps(output + i + 24, e3); | ||||||||
|
|
||||||||
| sum0 = _mm256_add_ps(sum0, e0); | ||||||||
| sum1 = _mm256_add_ps(sum1, e1); | ||||||||
| sum2 = _mm256_add_ps(sum2, e2); | ||||||||
| sum3 = _mm256_add_ps(sum3, e3); | ||||||||
| } | ||||||||
| sum0 = _mm256_add_ps(sum0, sum1); | ||||||||
| sum2 = _mm256_add_ps(sum2, sum3); | ||||||||
| sum0 = _mm256_add_ps(sum0, sum2); | ||||||||
|
|
||||||||
| for (; i + 7 < n; i += 8) { | ||||||||
| __m256 x = _mm256_loadu_ps(input + i); | ||||||||
| __m256 e = exp256_ps_v2(_mm256_sub_ps(x, max_vec)); | ||||||||
| _mm256_storeu_ps(output + i, e); | ||||||||
| sum0 = _mm256_add_ps(sum0, e); | ||||||||
| } | ||||||||
|
|
||||||||
| float sum_val = reduce_sum(sum0); | ||||||||
| for (; i < n; ++i) { | ||||||||
| float e = std::exp(input[i] - max_val); | ||||||||
| output[i] = e; | ||||||||
| sum_val += e; | ||||||||
| } | ||||||||
|
|
||||||||
| if (sum_val == 0.0f) return; | ||||||||
|
|
||||||||
| // 3. Normalize (8x unrolled) | ||||||||
| float inv_sum = 1.0f / sum_val; | ||||||||
| __m256 inv_sum_v = _mm256_set1_ps(inv_sum); | ||||||||
| i = 0; | ||||||||
|
|
||||||||
| for (; i + 63 < n; i += 64) { | ||||||||
| __m256 o0 = _mm256_loadu_ps(output + i); | ||||||||
| __m256 o1 = _mm256_loadu_ps(output + i + 8); | ||||||||
| __m256 o2 = _mm256_loadu_ps(output + i + 16); | ||||||||
| __m256 o3 = _mm256_loadu_ps(output + i + 24); | ||||||||
| __m256 o4 = _mm256_loadu_ps(output + i + 32); | ||||||||
| __m256 o5 = _mm256_loadu_ps(output + i + 40); | ||||||||
| __m256 o6 = _mm256_loadu_ps(output + i + 48); | ||||||||
| __m256 o7 = _mm256_loadu_ps(output + i + 56); | ||||||||
|
|
||||||||
| __m256 m0 = _mm256_mul_ps(o0, inv_sum_v); | ||||||||
| __m256 m1 = _mm256_mul_ps(o1, inv_sum_v); | ||||||||
| __m256 m2 = _mm256_mul_ps(o2, inv_sum_v); | ||||||||
| __m256 m3 = _mm256_mul_ps(o3, inv_sum_v); | ||||||||
| __m256 m4 = _mm256_mul_ps(o4, inv_sum_v); | ||||||||
| __m256 m5 = _mm256_mul_ps(o5, inv_sum_v); | ||||||||
| __m256 m6 = _mm256_mul_ps(o6, inv_sum_v); | ||||||||
| __m256 m7 = _mm256_mul_ps(o7, inv_sum_v); | ||||||||
|
|
||||||||
| _mm256_storeu_ps(output + i, m0); | ||||||||
| _mm256_storeu_ps(output + i + 8, m1); | ||||||||
| _mm256_storeu_ps(output + i + 16, m2); | ||||||||
| _mm256_storeu_ps(output + i + 24, m3); | ||||||||
| _mm256_storeu_ps(output + i + 32, m4); | ||||||||
| _mm256_storeu_ps(output + i + 40, m5); | ||||||||
| _mm256_storeu_ps(output + i + 48, m6); | ||||||||
| _mm256_storeu_ps(output + i + 56, m7); | ||||||||
| } | ||||||||
| for (; i + 7 < n; i += 8) { | ||||||||
| _mm256_storeu_ps(output + i, _mm256_mul_ps(_mm256_loadu_ps(output + i), inv_sum_v)); | ||||||||
| } | ||||||||
| for (; i < n; ++i) { | ||||||||
| output[i] *= inv_sum; | ||||||||
| } | ||||||||
| } | ||||||||
|
|
||||||||
| } // namespace ml_kernels | ||||||||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
📐 Maintainability & Code Quality | 🟡 Minor | ⚡ Quick win
Separate the max-pass latency rationale from normalization tuning.
Pass 3 does not execute
_mm256_max_ps. Do not state that its 8x unroll hides_mm256_max_pslatency. State that Pass 1 uses independent accumulators to hide max latency. Justify Pass 3 unrolling with measured load/store bandwidth or multiply throughput.ml_kernels/include/ml_kernels/softmax.h#L507-L510: revise the implementation comment to limit_mm256_max_pslatency hiding to Pass 1..jules/thunderbolt.md#L32-L32: revise the journal entry to use the same pass-specific rationale.📍 Affects 2 files
ml_kernels/include/ml_kernels/softmax.h#L507-L510(this comment).jules/thunderbolt.md#L32-L32🤖 Prompt for AI Agents