On Thu, 13 Aug 2026 22:48:49 GMT, Chad Rakoczy <[email protected]> wrote:

> Adds support for vectorized dot product on aarch64 (`sdot` and `udot`) 
> through the Vector API
> 
> Dedicated dot product instructions have been shown to provide up to 10x 
> throughput for lucene 
> ([results](https://github.com/apache/lucene/pull/13572)) compared to 
> vectorized multiply and add. This PR adds two new methods to `ByteVector` to 
> leverage the aarch dot product instructions `sdot` and `udot` respectively.
> - `IntVector dot(Vector<Byte> v, Vector<Integer> acc)`
> - `IntVector dotUnsigned(Vector<Byte> v, Vector<Integer> acc)`
> 
> Each int lane of the result holds the dot product of the corresponding group 
> of four bytes from the two operands, added to the matching accumulator lane. 
> For example:
> 
> a = [a1, a2, a3, a4, ..., ..., a13, a14, a15, a16]
> b = [b1, b2, b3, b4, ..., ..., b13, b14, b15, b16]
> acc = [acc1, ..., ..., acc4]
> 
> a.dot(b, acc) -> [
>     acc1 + a1 * b1 + a2 * b2 + a3 * b3 + a4 * b4, 
>     ..., 
>     ...,
>     acc4 + a13 * b13 + a14 * b14 + a15 * b15 + a16 * b16
> ]
> 
> 
> The equivalent instructions on x86 are `VPDPBSSD` and `VPDPBUUD` which 
> perform the same operations and match the proposed new functions however this 
> PR only includes aarch64.
> 
> Graviton 2
> 
> Benchmark                              Mode  Cnt      Score     Error   Units
> VectorDotBenchmark.dotScalar          thrpt   25   2412.719 ±   0.014  ops/ms
> VectorDotBenchmark.dotMulAdd          thrpt   25   1497.953 ±   2.734  ops/ms
> VectorDotBenchmark.dotVector          thrpt   25  23243.016 ± 141.708  ops/ms
> VectorDotBenchmark.dotUnsignedScalar  thrpt   25   2402.966 ±   0.157  ops/ms
> VectorDotBenchmark.dotUnsignedMulAdd  thrpt   25   1498.898 ±   1.442  ops/ms
> VectorDotBenchmark.dotUnsignedVector  thrpt   25  23344.116 ± 222.150  ops/ms
> 
> 
> Graviton 3
> 
> Benchmark                              Mode  Cnt      Score     Error   Units
> VectorDotBenchmark.dotScalar          thrpt   25   7946.950 ±   2.564  ops/ms
> VectorDotBenchmark.dotMulAdd          thrpt   25   3257.311 ±   5.308  ops/ms
> VectorDotBenchmark.dotVector          thrpt   25  41430.996 ± 688.936  ops/ms
> VectorDotBenchmark.dotUnsignedScalar  thrpt   25   2536.268 ±   0.438  ops/ms
> VectorDotBenchmark.dotUnsignedMulAdd  thrpt   25   3259.853 ±   3.511  ops/ms
> VectorDotBenchmark.dotUnsignedVector  thrpt   25  41263.821 ± 478.856  ops/ms
> 
> 
> ---------
> - [x] I confirm that I make this contribution in accordance with the [OpenJDK 
> Interim AI Policy](https://openjdk.org/legal/ai).

Indeed, it's a very interesting proposal, Chad. 

Leaving API considerations aside, following on Emanuel's proposal, one 
equivalent implementation for 128-bit byte vector `sdot` operation :

    static final VectorSpecies<Byte> B128 = VectorSpecies.of(byte.class, 
VectorShape.S_128_BIT);
    static final VectorSpecies<Integer> I128 = VectorSpecies.of(int.class,  
VectorShape.S_128_BIT);

    static int sdot(ByteVector v1, ByteVector v2) {
        int acc = 0;
        for (int part = 0; part < 4; part++) {
            var i1 = v1.castShape(I128, part).reinterpretAsInts();
            var i2 = v2.castShape(I128, part).reinterpretAsInts();
            var mul = i1.lanewise(VectorOperators.MUL, i2).reinterpretAsInts();
            acc += mul.reduceLanes(VectorOperators.ADD);
        }
        return acc;
    }
```   

How hard would it be to substitute relevant IR into `DotV`? I briefly looked at 
generated IR and spotted that `slice(int)` is not intrinsified:

5581   69    b        vector.SDot::sdot (26 bytes)
  ** vector slice from non-constant index not supported
``` 

And manually unrolling the loop [1] doesn't help, because backend support is 
missing:

12803   69    b        vector.SDot::sdot (40 bytes)
  ** Rejected vector op (VectorSlice,byte,16) because architecture does not 
support it
  ** not supported: arity=2 op=slice vlen=16 etype=byte
...


Once `VectorSlice` is intrinsified, corresponding IR shape should become much 
more manageable for matching. 

[1] 

    static int sdot(ByteVector v1, ByteVector v2) {
        int acc = 0;
        acc += sdotPart(v1, v2, 0);
        acc += sdotPart(v1, v2, 1);
        acc += sdotPart(v1, v2, 2);
        acc += sdotPart(v1, v2, 3);
        return acc;
    }

    static int sdotPart(ByteVector v1, ByteVector v2, int part) {
        var i1 = v1.castShape(I128, part).reinterpretAsInts();
        var i2 = v2.castShape(I128, part).reinterpretAsInts();
        var mul = i1.lanewise(VectorOperators.MUL, i2).reinterpretAsInts();
        return mul.reduceLanes(VectorOperators.ADD);
    }

-------------

PR Comment: https://git.openjdk.org/jdk/pull/32359#issuecomment-6047099923

Reply via email to