如何用AVX2指令集在C++中实现numpy.triu_indices(a,1)?
用AVX2实现
numpy.triu_indices(a, 1)的向量化版本 嘿,刚学AVX2不用不好意思,这个需求其实很适合入门向量指令——输出的索引有很强的规律性,刚好能利用AVX2的批量处理能力替代嵌套标量循环。我来一步步拆解思路,再附上可运行的代码片段。
核心思路回顾
triu_indices(a,1)本质是生成所有满足i < j的索引对(i,j):
- 对于每个
i,j从i+1到a-1,所以first数组会重复i共a-i-1次; second数组则是连续的递增序列i+1, i+2, ..., a-1。
AVX2一次能处理8个int(256位寄存器,每个int占4字节),所以我们可以批量生成这8个元素的向量,再一次性存储到输出数组里,大幅减少循环次数。
关键AVX2指令说明
先提前熟悉几个核心指令:
_mm256_set1_epi32(x):把单个int值x复制8次,生成一个256位向量(用于批量生成first数组的重复i值);_mm256_setr_epi32(x0,x1,...,x7):按顺序生成向量,x0对应内存中的第一个元素(用于生成0-7的增量序列);_mm256_add_epi32(v1, v2):两个向量对应元素相加(用于生成second的连续序列);_mm256_storeu_si256(ptr, vec):把向量存储到内存(u表示非对齐存储,兼容性更好,避免内存对齐问题)。
完整AVX2实现代码
#include <immintrin.h> #include <cstdint> void triu_indices_avx2(int a, int* first, int* second) { if (a <= 1) { return; // 没有符合条件的索引对,直接返回 } int index = 0; // 预生成增量向量:[0,1,2,3,4,5,6,7],复用它来生成second的连续序列 __m256i incr_vec = _mm256_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7); for (int i = 0; i < a; ++i) { int start_j = i + 1; int count_j = a - start_j; // 当前i对应的j的总数 if (count_j <= 0) { continue; } // 批量处理8个元素的整批次 int remaining = count_j; int current_j = start_j; while (remaining >= 8) { // 生成first的向量:8个重复的i __m256i first_vec = _mm256_set1_epi32(i); // 生成second的向量:current_j 到 current_j+7 的连续值 __m256i start_vec = _mm256_set1_epi32(current_j); __m256i second_vec = _mm256_add_epi32(start_vec, incr_vec); // 一次性存储两个向量到输出数组 _mm256_storeu_si256(reinterpret_cast<__m256i*>(first + index), first_vec); _mm256_storeu_si256(reinterpret_cast<__m256i*>(second + index), second_vec); // 更新索引和计数器 index += 8; current_j += 8; remaining -= 8; } // 处理剩余不足8个的元素(用标量循环,适合初学者理解) while (remaining > 0) { first[index] = i; second[index] = current_j; index++; current_j++; remaining--; } } }
代码细节解释
- 边界处理:当
a<=1时直接返回,因为没有任何i<j的索引对; - 增量向量预生成:提前创建
[0,1,...,7]的向量,避免每次循环重复生成,提升效率; - 整批次处理:对每个
i,先处理8个元素的批量:- 用
_mm256_set1_epi32(i)快速生成8个i的向量,对应first数组的8个元素; - 用起始
current_j加上增量向量,生成连续的j值向量; - 一次性存储两个向量到输出数组,比标量循环快很多;
- 用
- 剩余元素处理:最后不足8个的部分用标量循环处理,逻辑简单不容易出错。如果想进阶,可以尝试
_mm256_maskstore_epi32指令实现向量式剩余处理。
编译与测试提示
- 编译时需要开启AVX2支持:GCC/Clang用
-mavx2,MSVC用/arch:AVX2; - 测试时可以对比你的非向量化版本,比如输入
a=4,输出应该和你给出的示例完全一致; - 如果输出数组是用对齐内存分配的(比如
_mm_malloc),可以把_mm256_storeu_si256换成_mm256_store_si256,性能会略好一点。
内容的提问来源于stack exchange,提问作者Roy_123
相关产品推荐
相关产品推荐

