Add masked filter-and-sum kernel
This commit is contained in:
45
timings/src/main/java/com/ankurm/vectorapi/MaskedFilter.java
Normal file
45
timings/src/main/java/com/ankurm/vectorapi/MaskedFilter.java
Normal file
@@ -0,0 +1,45 @@
|
|||||||
|
package com.ankurm.vectorapi;
|
||||||
|
|
||||||
|
import jdk.incubator.vector.*;
|
||||||
|
|
||||||
|
/** Sum only the elements above a threshold -- a data-dependent branch. */
|
||||||
|
public final class MaskedFilter {
|
||||||
|
|
||||||
|
static final float THRESH = 0.5f;
|
||||||
|
static final VectorSpecies<Float> SP = FloatVector.SPECIES_PREFERRED;
|
||||||
|
|
||||||
|
// Scalar: the branch is data-dependent, so C2 generally cannot vectorize this form
|
||||||
|
static float scalar(float[] a) {
|
||||||
|
float sum = 0;
|
||||||
|
for (int i = 0; i < a.length; i++)
|
||||||
|
if (a[i] > THRESH) sum += a[i];
|
||||||
|
return sum;
|
||||||
|
}
|
||||||
|
|
||||||
|
static float vector(float[] a) {
|
||||||
|
FloatVector acc = FloatVector.zero(SP);
|
||||||
|
int i = 0, bound = SP.loopBound(a.length);
|
||||||
|
for (; i < bound; i += SP.length()) {
|
||||||
|
var v = FloatVector.fromArray(SP, a, i);
|
||||||
|
VectorMask<Float> m = v.compare(VectorOperators.GT, THRESH);
|
||||||
|
acc = acc.add(v, m); // add only where mask is true
|
||||||
|
}
|
||||||
|
float sum = acc.reduceLanes(VectorOperators.ADD); // horizontal sum of the lanes
|
||||||
|
for (; i < a.length; i++) if (a[i] > THRESH) sum += a[i]; // masked-off tail
|
||||||
|
return sum;
|
||||||
|
}
|
||||||
|
|
||||||
|
static void run(int n) {
|
||||||
|
float[] a = Data.randomFloats(n, 3);
|
||||||
|
float s0 = scalar(a), v0 = vector(a);
|
||||||
|
|
||||||
|
System.out.println("--- kernel 2: masked filter-and-sum, " + n + " floats ---");
|
||||||
|
System.out.printf("scalar sum = %.4f%nvector sum = %.4f (same value, different rounding)%n", s0, v0);
|
||||||
|
double s = Bench.measure("scalar (branch)", () -> Sink.consume(scalar(a)));
|
||||||
|
double v = Bench.measure("Vector API (masked)", () -> Sink.consume(vector(a)));
|
||||||
|
Bench.speedup("vector advantage", s, v);
|
||||||
|
System.out.println();
|
||||||
|
}
|
||||||
|
|
||||||
|
private MaskedFilter() {}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user