如何用AVX2 Intrinsics优化2D前缀和子矩阵均值方差计算
Optimizing Fixed-Size Submatrix Mean/Variance Calculation with AVX2 Intrinsics
Since you're targeting a Kaby Lake i7-7700HQ (which fully supports AVX2), we can completely vectorize the scalar loop using Intel's AVX2 intrinsic functions to process 8 doubles per instruction. Your lineSize=16 means we'll handle two batches of 8 elements each with AVX2's 256-bit vectors. Here's the maximized speed optimization, with explanations of key optimizations:
Key Optimization Principles
- Full Vectorization: Replace scalar per-element operations with AVX2 256-bit vector operations to exploit SIMD parallelism.
- Memory Alignment: Ensure input/output arrays are 32-byte aligned (required for fastest AVX2 loads/stores).
- Minimize Memory Access: Load all required data into vector registers upfront instead of per-element loads.
- Constant Broadcast: Pre-broadcast scalar constants (like
submatAreaInv) into vector registers to avoid redundant operations.
Optimized Implementation
First, include the necessary intrinsic header:
#include <immintrin.h>
Then rewrite the function with AVX2 intrinsics:
#define cell(i, j, w) ((i)*(w) + (j)) const int lineSize = 16; const int R = 3; const int submatArea = (R+1)*(R+1); const double submatAreaInv = double(1) / submatArea; // Ensure S, Sqr, mean, var are 32-byte aligned (use alignas(32) or _mm_malloc) void subMatrixVarMulti(const int64_t* S, const int64_t* Sqr, int top, int left, int bot, int right, int w, int h, int diff, double mean[lineSize], double var[lineSize]) { // Calculate base indices (lineSize elements start at each index) const int idxTopLeft = cell(top - 1, left - 1, w); const int idxTopRight = cell(top - 1, right, w); const int idxBotLeft = cell(bot, left - 1, w); const int idxBotRight = cell(bot, right, w); // Pre-broadcast the inverse area to a 256-bit vector const __m256d inv_area = _mm256_set1_pd(submatAreaInv); // Process first 8 elements (i=0 to 7) // Load int64 vectors from S and convert to double const __m256i s_tl_8 = _mm256_load_si256((const __m256i*)(S + idxTopLeft)); const __m256i s_tr_8 = _mm256_load_si256((const __m256i*)(S + idxTopRight)); const __m256i s_bl_8 = _mm256_load_si256((const __m256i*)(S + idxBotLeft)); const __m256i s_br_8 = _mm256_load_si256((const __m256i*)(S + idxBotRight)); const __m256d d_tl_8 = _mm256_cvtepi64_pd(s_tl_8); const __m256d d_tr_8 = _mm256_cvtepi64_pd(s_tr_8); const __m256d d_bl_8 = _mm256_cvtepi64_pd(s_bl_8); const __m256d d_br_8 = _mm256_cvtepi64_pd(s_br_8); // Compute sum_S = S[botRight] - S[botLeft] - S[topRight] + S[topLeft] const __m256d sum_S_8 = _mm256_add_pd(_mm256_sub_pd(_mm256_sub_pd(d_br_8, d_bl_8), d_tr_8), d_tl_8); // Compute mean for first 8 elements const __m256d mean_8 = _mm256_mul_pd(sum_S_8, inv_area); // Load int64 vectors from Sqr and convert to double const __m256i sqr_tl_8 = _mm256_load_si256((const __m256i*)(Sqr + idxTopLeft)); const __m256i sqr_tr_8 = _mm256_load_si256((const __m256i*)(Sqr + idxTopRight)); const __m256i sqr_bl_8 = _mm256_load_si256((const __m256i*)(Sqr + idxBotLeft)); const __m256i sqr_br_8 = _mm256_load_si256((const __m256i*)(Sqr + idxBotRight)); const __m256d d_sqr_tl_8 = _mm256_cvtepi64_pd(sqr_tl_8); const __m256d d_sqr_tr_8 = _mm256_cvtepi64_pd(sqr_tr_8); const __m256d d_sqr_bl_8 = _mm256_cvtepi64_pd(sqr_bl_8); const __m256d d_sqr_br_8 = _mm256_cvtepi64_pd(sqr_br_8); // Compute sum_Sqr = Sqr[botRight] - Sqr[botLeft] - Sqr[topRight] + Sqr[topLeft] const __m256d sum_Sqr_8 = _mm256_add_pd(_mm256_sub_pd(_mm256_sub_pd(d_sqr_br_8, d_sqr_bl_8), d_sqr_tr_8), d_sqr_tl_8); // Compute variance for first 8 elements: (sumSqr/area) - mean^2 const __m256d var_8 = _mm256_sub_pd(_mm256_mul_pd(sum_Sqr_8, inv_area), _mm256_mul_pd(mean_8, mean_8)); // Store first 8 results _mm256_store_pd(mean, mean_8); _mm256_store_pd(var, var_8); // Process next 8 elements (i=8 to 15) const __m256i s_tl_16 = _mm256_load_si256((const __m256i*)(S + idxTopLeft + 8)); const __m256i s_tr_16 = _mm256_load_si256((const __m256i*)(S + idxTopRight + 8)); const __m256i s_bl_16 = _mm256_load_si256((const __m256i*)(S + idxBotLeft + 8)); const __m256i s_br_16 = _mm256_load_si256((const __m256i*)(S + idxBotRight + 8)); const __m256d d_tl_16 = _mm256_cvtepi64_pd(s_tl_16); const __m256d d_tr_16 = _mm256_cvtepi64_pd(s_tr_16); const __m256d d_bl_16 = _mm256_cvtepi64_pd(s_bl_16); const __m256d d_br_16 = _mm256_cvtepi64_pd(s_br_16); const __m256d sum_S_16 = _mm256_add_pd(_mm256_sub_pd(_mm256_sub_pd(d_br_16, d_bl_16), d_tr_16), d_tl_16); const __m256d mean_16 = _mm256_mul_pd(sum_S_16, inv_area); const __m256i sqr_tl_16 = _mm256_load_si256((const __m256i*)(Sqr + idxTopLeft + 8)); const __m256i sqr_tr_16 = _mm256_load_si256((const __m256i*)(Sqr + idxTopRight + 8)); const __m256i sqr_bl_16 = _mm256_load_si256((const __m256i*)(Sqr + idxBotLeft + 8)); const __m256i sqr_br_16 = _mm256_load_si256((const __m256i*)(Sqr + idxBotRight + 8)); const __m256d d_sqr_tl_16 = _mm256_cvtepi64_pd(sqr_tl_16); const __m256d d_sqr_tr_16 = _mm256_cvtepi64_pd(sqr_tr_16); const __m256d d_sqr_bl_16 = _mm256_cvtepi64_pd(sqr_bl_16); const __m256d d_sqr_br_16 = _mm256_cvtepi64_pd(sqr_br_16); const __m256d sum_Sqr_16 = _mm256_add_pd(_mm256_sub_pd(_mm256_sub_pd(d_sqr_br_16, d_sqr_bl_16), d_sqr_tr_16), d_sqr_tl_16); const __m256d var_16 = _mm256_sub_pd(_mm256_mul_pd(sum_Sqr_16, inv_area), _mm256_mul_pd(mean_16, mean_16)); // Store next 8 results _mm256_store_pd(mean + 8, mean_16); _mm256_store_pd(var + 8, var_16); }
Additional Optimization Tips
Memory Alignment:
- Declare your arrays with
alignas(32)to ensure 32-byte alignment (required for_mm256_load_si256/_mm256_store_pdto work at full speed). For example:alignas(32) int64_t S[...]; alignas(32) double mean[lineSize]; - If allocating dynamically, use
_mm_malloc(size, 32)instead ofmalloc.
- Declare your arrays with
Loop Unrolling:
- The code above manually unrolls the loop for
lineSize=16(two batches of 8). If you need to support otherlineSizevalues, you can create a macro or template to handle multiples of 8, or combine with scalar cleanup for non-multiples.
- The code above manually unrolls the loop for
Compiler Flags:
- Compile with
-mavx2 -O3(GCC/Clang) or/arch:AVX2 /O2(MSVC) to enable AVX2 support and full optimizations. The compiler may further optimize the intrinsic code (e.g., reorder instructions for better pipeline utilization).
- Compile with
Avoid Redundant Conversions:
- Since
SandSqrareint64_t, we convert them todoubleonce per vector load, which is efficient because AVX2 has dedicated conversion instructions (_mm256_cvtepi64_pd).
- Since
内容的提问来源于stack exchange,提问作者Huy Le
相关产品推荐
相关产品推荐

