基于Java Vector API优化int16向量点积计算
利用Java Vector API优化无溢出的16位整数数组乘法
问题分析
你需要优化的核心逻辑是对两个16位整数数组应用screlu激活函数后,与权重数组计算点积并累加。当前使用IntVector的方案因32位通道限制导致吞吐量减半,希望找到类似_mm256_madd_epi16的无溢出16位向量乘法方案。
解决方案
Java Vector API(孵化版)支持直接操作16位整数的ShortVector,可以大幅提升向量利用率(同宽度下ShortVector的元素数是IntVector的2倍),同时通过合理的类型转换避免溢出。以下是具体优化方案:
1. 基础准备:定义向量物种
首先根据CPU支持的最优向量宽度定义对应的向量物种:
import jdk.incubator.vector.IntVector; import jdk.incubator.vector.ShortVector; import jdk.incubator.vector.VectorOperators; import jdk.incubator.vector.VectorSpecies; static final VectorSpecies<Short> SHORT_SPECIES = VectorSpecies.of(short.class, VectorShape.preferredShape()); static final VectorSpecies<Integer> INT_SPECIES = VectorSpecies.of(int.class, VectorShape.preferredShape());
2. 优化后的向量化实现
根据screlu中QA的取值范围,分两种情况处理:
情况一:QA ≤ 181(short平方无溢出)
当QA不超过181时,v = max(0, min(i, QA))的平方值(≤181²=32761)不会超过short的最大值(32767),可直接用ShortVector完成所有计算,无需转成int:
public int optimizedDotProduct(ShortArray us, ShortArray them, Network network) { int result = 0; int batchSize = SHORT_SPECIES.length(); int upperBound = HIDDEN_SIZE - HIDDEN_SIZE % batchSize; int i = 0; // 向量化循环 for (; i < upperBound; i += batchSize) { // 加载原始16位数组 ShortVector usVec = ShortVector.fromArray(SHORT_SPECIES, us.values, i); ShortVector themVec = ShortVector.fromArray(SHORT_SPECIES, them.values, i); ShortVector weights1 = ShortVector.fromArray(SHORT_SPECIES, network.L1Weights, i); ShortVector weights2 = ShortVector.fromArray(SHORT_SPECIES, network.L1Weights, i + HIDDEN_SIZE); // 应用screlu激活函数 ShortVector usClamped = usVec.max((short) 0).min((short) QA); ShortVector themClamped = themVec.max((short) 0).min((short) QA); // 计算平方并与权重相乘,转成IntVector避免累加溢出 IntVector usProduct = usClamped.mul(usClamped).convertTo(INT_SPECIES, Integer.class) .mul(weights1.convertTo(INT_SPECIES, Integer.class)); IntVector themProduct = themClamped.mul(themClamped).convertTo(INT_SPECIES, Integer.class) .mul(weights2.convertTo(INT_SPECIES, Integer.class)); // 累加当前批次结果 result += usProduct.add(themProduct).reduceLanes(VectorOperators.ADD); } // 处理剩余未向量化的元素 for (; i < HIDDEN_SIZE; i++) { result += screlu(us.values[i]) * network.L1Weights[i] + screlu(them.values[i]) * network.L1Weights[i + HIDDEN_SIZE]; } return result; }
情况二:QA > 181(需转IntVector避免平方溢出)
当QA超过181时,short平方会溢出,需先将ShortVector转成IntVector再计算平方:
public int optimizedDotProduct(ShortArray us, ShortArray them, Network network) { int result = 0; int batchSize = SHORT_SPECIES.length(); int upperBound = HIDDEN_SIZE - HIDDEN_SIZE % batchSize; int i = 0; // 向量化循环 for (; i < upperBound; i += batchSize) { ShortVector usVec = ShortVector.fromArray(SHORT_SPECIES, us.values, i); ShortVector themVec = ShortVector.fromArray(SHORT_SPECIES, them.values, i); ShortVector weights1 = ShortVector.fromArray(SHORT_SPECIES, network.L1Weights, i); ShortVector weights2 = ShortVector.fromArray(SHORT_SPECIES, network.L1Weights, i + HIDDEN_SIZE); // 应用screlu并转成IntVector IntVector usClamped = usVec.max((short) 0).min((short) QA).convertTo(INT_SPECIES, Integer.class); IntVector themClamped = themVec.max((short) 0).min((short) QA).convertTo(INT_SPECIES, Integer.class); // 计算平方并与权重相乘 IntVector usProduct = usClamped.mul(usClamped) .mul(weights1.convertTo(INT_SPECIES, Integer.class)); IntVector themProduct = themClamped.mul(themClamped) .mul(weights2.convertTo(INT_SPECIES, Integer.class)); // 累加结果 result += usProduct.add(themProduct).reduceLanes(VectorOperators.ADD); } // 处理剩余元素 for (; i < HIDDEN_SIZE; i++) { result += screlu(us.values[i]) * network.L1Weights[i] + screlu(them.values[i]) * network.L1Weights[i + HIDDEN_SIZE]; } return result; }
3. 关于_mm256_madd_epi16的等价操作
_mm256_madd_epi16的核心是“16位乘16位得到32位,相邻两个32位结果相加”,Java Vector API中可以通过以下方式模拟:
// 假设有两个ShortVector a和b ShortVector a = ...; ShortVector b = ...; // 先相乘得到IntVector(每个short乘short转成int) IntVector mul = a.convertTo(INT_SPECIES, Integer.class).mul(b.convertTo(INT_SPECIES, Integer.class)); // 相邻元素相加:需要将向量拆分为高低半部分,然后相加 IntVector low = mul.slice(0, INT_SPECIES.length() / 2); IntVector high = mul.slice(INT_SPECIES.length() / 2, INT_SPECIES.length() / 2); IntVector maddResult = low.add(high);
不过在你的场景中,不需要相邻元素相加,而是需要所有元素的乘积累加,因此直接使用上述的向量化方案即可。
关键优化点
- 避免提前将short数组转成int数组,减少内存拷贝开销。
- 使用
ShortVector提升向量元素密度,吞吐量比IntVector翻倍。 - 根据
QA的取值范围选择最优的计算路径,平衡性能与无溢出需求。
内容的提问来源于stack exchange,提问作者shawn_xu
相关产品推荐
相关产品推荐

