You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

基于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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.25 03:14:54