1
0

Add masked filter-and-sum kernel

This commit is contained in:
2026-07-26 11:46:03 +00:00
parent e1368e0d76
commit db213647d1

View 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() {}
}