Merge pull request #22342 from hrydgard/vdot-neon

Accurate vdot: NEON and SSE2 implementations
This commit is contained in:
Henrik Rydgård authored and GitHub committed 2026-09-23 19:24:16 -06:00
commit 09cf031f2b
3 files changed
+221 -3

No files matched your search

+126 -3
View File
@@ -742,7 +742,7 @@ static inline float vfpu_div(float a, float b) {
// to PSP output on all inputs. For details see
// https://github.com/hrydgard/ppsspp/issues/21070#issuecomment-4640382516
// Reference C++ version.
static float vfpu_dot_cpp(const float a[4], const float b[4]) {
float vfpu_dot_reference(const float a[4], const float b[4]) {
int EXTRA_BITS = 2;
uint32_t I = uint32_t(1) << 23, J = uint32_t(1) << (23 - EXTRA_BITS);
int32_t s[4], e[4], ehi = -2*127;
@@ -830,9 +830,132 @@ static float vfpu_dot_cpp(const float a[4], const float b[4]) {
return ret;
}
#if PPSSPP_ARCH(ARM64_NEON) || PPSSPP_ARCH(SSE2)
// Inf and NaN inputs are rare, and the reference handles them.
static NO_INLINE float vfpu_dot_special(const float a[4], const float b[4]) {
return vfpu_dot_reference(a, b);
}
// The end of the SIMD versions: drops the extra bits off the exact sum by truncation, then rounds to
// 24 bits, nearest even. With the top bit moved to bit 62, adding 2^38 - 1 plus the lowest kept bit
// and shifting out 39 bits does that for every magnitude, and adding the significand, implicit bit
// and all, onto the exponent field lets a carry from the rounding bump the exponent by itself. It's
// done in integers because the host rounding mode may be the game's.
static inline float vfpu_dot_finish(int32_t val, uint32_t ehi) {
const uint32_t signBit = (uint32_t)val & 0x80000000u;
const uint32_t m = (val < 0 ? 0u - (uint32_t)val : (uint32_t)val) >> 2;
const int lz = 32 + (int)clz32_nonzero(m | 1); // of m as 64 bits
const uint64_t top = (uint64_t)m << (lz - 1);
const uint64_t q = (top + ((1ULL << 38) - 1) + ((top >> 39) & 1)) >> 39;
const int64_t t = ((int64_t)((int)ehi - 88 - lz) << 23) + (int64_t)q;
uint32_t bits = t >= 0x7F800000 ? 0x7F800000u : (uint32_t)t;
bits = (t < 0x00800000 || m == 0) ? 0u : bits;
bits |= signBit;
float ret;
memcpy(&ret, &bits, sizeof(ret));
return ret;
}
#endif
#if PPSSPP_ARCH(ARM64_NEON)
// The reference, four lanes at a time. Bit-exact with it.
static float vfpu_dot_neon(const float a[4], const float b[4]) {
const uint32x4_t x = vld1q_u32((const uint32_t *)a);
const uint32x4_t y = vld1q_u32((const uint32_t *)b);
const uint32x4_t expBits = vdupq_n_u32(0x7F800000);
const uint32x4_t xe = vandq_u32(x, expBits), ye = vandq_u32(y, expBits);
// Zero and subnormal inputs make the product zero, with the lowest exponent. Inf and NaN get
// 0x7FF added, above any real exponent sum, so one maximum finds the alignment and the specials.
const uint32x4_t normal = vandq_u32(vtstq_u32(x, expBits), vtstq_u32(y, expBits));
const uint32x4_t special = vceqq_u32(vmaxq_u32(xe, ye), expBits);
const uint32x4_t e = vsraq_n_u32(vandq_u32(vshrq_n_u32(vaddq_u32(xe, ye), 23), normal), special, 21);
uint32x4_t emax = vpmaxq_u32(e, e);
emax = vpmaxq_u32(emax, emax);
const uint32_t ehi = vgetq_lane_u32(emax, 0);
if (ehi >= 0x7FF)
return vfpu_dot_special(a, b);
// 24x24-bit products, kept to 2 extra bits with round-to-odd.
const uint32x4_t mantMask = vdupq_n_u32(0x007FFFFF), implicitBit = vdupq_n_u32(0x00800000);
const uint32x4_t mx = vbslq_u32(mantMask, x, implicitBit);
const uint32x4_t my = vbslq_u32(mantMask, y, implicitBit);
const uint64x2_t plo = vmull_u32(vget_low_u32(mx), vget_low_u32(my));
const uint64x2_t phi = vmull_high_u32(mx, my);
const uint32x4_t kept = vcombine_u32(vshrn_n_u64(plo, 21), vshrn_n_u64(phi, 21));
const uint32x4_t dropped = vandq_u32(vuzp1q_u32(vreinterpretq_u32_u64(plo), vreinterpretq_u32_u64(phi)), vdupq_n_u32((1 << 21) - 1));
const uint32x4_t p = vandq_u32(vorrq_u32(kept, vminq_u32(dropped, vdupq_n_u32(1))), normal);
// Align to the largest exponent by truncation, apply the signs, and sum exactly.
const int32x4_t shift = vmaxq_s32(vreinterpretq_s32_u32(vsubq_u32(e, emax)), vdupq_n_s32(-28));
const int32x4_t v = vreinterpretq_s32_u32(vshlq_u32(p, shift));
const int32x4_t sign = vshrq_n_s32(vreinterpretq_s32_u32(veorq_u32(x, y)), 31);
return vfpu_dot_finish(vaddvq_s32(vsubq_s32(veorq_s32(v, sign), sign)), ehi);
}
#elif PPSSPP_ARCH(SSE2)
// The reference, four lanes at a time. Bit-exact with it.
static float vfpu_dot_sse2(const float a[4], const float b[4]) {
const __m128i x = _mm_loadu_si128((const __m128i *)a);
const __m128i y = _mm_loadu_si128((const __m128i *)b);
const __m128i expBits = _mm_set1_epi32(0x7F800000);
const __m128i xe = _mm_and_si128(x, expBits), ye = _mm_and_si128(y, expBits);
const __m128i zero = _mm_setzero_si128();
// Zero and subnormal inputs make the product zero, with the lowest exponent. Inf and NaN get
// 0x7FF, above any real exponent sum, so one maximum finds the alignment and the specials. All of
// these fit in 15 bits, so the 16-bit maximum works on them.
const __m128i flush = _mm_or_si128(_mm_cmpeq_epi32(xe, zero), _mm_cmpeq_epi32(ye, zero));
const __m128i special = _mm_or_si128(_mm_cmpeq_epi32(xe, expBits), _mm_cmpeq_epi32(ye, expBits));
const __m128i e = _mm_or_si128(_mm_andnot_si128(flush, _mm_srli_epi32(_mm_add_epi32(xe, ye), 23)), _mm_srli_epi32(special, 21));
__m128i emax = _mm_max_epi16(e, _mm_shuffle_epi32(e, _MM_SHUFFLE(1, 0, 3, 2)));
emax = _mm_max_epi16(emax, _mm_shuffle_epi32(emax, _MM_SHUFFLE(2, 3, 0, 1)));
const uint32_t ehi = (uint32_t)_mm_cvtsi128_si32(emax);
if (ehi >= 0x7FF)
return vfpu_dot_special(a, b);
// 24x24-bit products, even and odd lanes separately, kept to 2 extra bits with round-to-odd.
const __m128i mantMask = _mm_set1_epi32(0x007FFFFF), implicitBit = _mm_set1_epi32(0x00800000);
const __m128i mx = _mm_or_si128(_mm_and_si128(x, mantMask), implicitBit);
const __m128i my = _mm_or_si128(_mm_and_si128(y, mantMask), implicitBit);
const __m128i pEven = _mm_mul_epu32(mx, my);
const __m128i pOdd = _mm_mul_epu32(_mm_srli_epi64(mx, 32), _mm_srli_epi64(my, 32));
const __m128i kept = _mm_or_si128(_mm_srli_epi64(pEven, 21), _mm_slli_epi64(_mm_srli_epi64(pOdd, 21), 32));
const __m128i lowMask = _mm_set_epi32(0, (1 << 21) - 1, 0, (1 << 21) - 1);
const __m128i dropped = _mm_or_si128(_mm_and_si128(pEven, lowMask), _mm_slli_epi64(_mm_and_si128(pOdd, lowMask), 32));
const __m128i sticky = _mm_andnot_si128(_mm_cmpeq_epi32(dropped, zero), _mm_set1_epi32(1));
const __m128i p = _mm_andnot_si128(flush, _mm_or_si128(kept, sticky));
// Align to the largest exponent by truncation. SSE2 has no per-lane shift, so p >> d is done as
// (p * 2^(28 - d)) >> 28, with the power of two built as float bits and converted by truncation.
const __m128i d = _mm_min_epi16(_mm_sub_epi32(emax, e), _mm_set1_epi32(28));
const __m128i scale = _mm_cvttps_epi32(_mm_castsi128_ps(_mm_slli_epi32(_mm_sub_epi32(_mm_set1_epi32(127 + 28), d), 23)));
const __m128i sEven = _mm_srli_epi64(_mm_mul_epu32(p, scale), 28);
const __m128i sOdd = _mm_srli_epi64(_mm_mul_epu32(_mm_srli_epi64(p, 32), _mm_srli_epi64(scale, 32)), 28);
__m128i v = _mm_or_si128(sEven, _mm_slli_epi64(sOdd, 32));
// Apply the signs and sum exactly.
const __m128i sign = _mm_srai_epi32(_mm_xor_si128(x, y), 31);
v = _mm_sub_epi32(_mm_xor_si128(v, sign), sign);
v = _mm_add_epi32(v, _mm_shuffle_epi32(v, _MM_SHUFFLE(1, 0, 3, 2)));
v = _mm_add_epi32(v, _mm_shuffle_epi32(v, _MM_SHUFFLE(2, 3, 0, 1)));
return vfpu_dot_finish(_mm_cvtsi128_si32(v), ehi);
}
#endif
float vfpu_dot(const float a[4], const float b[4]) {
// SIMD version(s) not implemented.
return vfpu_dot_cpp(a, b);
#if PPSSPP_ARCH(ARM64_NEON)
return vfpu_dot_neon(a, b);
#elif PPSSPP_ARCH(SSE2)
return vfpu_dot_sse2(a, b);
#else
return vfpu_dot_reference(a, b);
#endif
}
//==============================================================================
+2
View File
@@ -60,6 +60,8 @@ inline float vfpu_clamp(float v, float min, float max) {
}
float vfpu_dot(const float a[4], const float b[4]);
// The portable version vfpu_dot is checked against.
float vfpu_dot_reference(const float a[4], const float b[4]);
float vfpu_sqrt(float a);
float vfpu_rsqrt(float a);
+93
View File
@@ -1986,6 +1986,98 @@ bool TestFastVec() {
return true;
}
// vfpu_dot's SIMD versions against the reference, on inputs chosen to make trouble: close
// exponents, cancelling products, ties, zeroes and subnormals, the overflow and underflow edges,
// inf and NaN, and sums whose rounding carries into the next power of two.
bool TestVFPUDot() {
uint64_t state = 0x9E3779B97F4A7C15ULL;
auto rnd = [&]() {
state ^= state << 13;
state ^= state >> 7;
state ^= state << 17;
return state;
};
auto fromBits = [](uint32_t bits) {
float f;
memcpy(&f, &bits, sizeof(f));
return f;
};
auto toBits = [](float f) {
uint32_t bits;
memcpy(&bits, &f, sizeof(bits));
return bits;
};
auto check = [&](const float a[4], const float b[4]) {
const uint32_t expected = toBits(vfpu_dot_reference(a, b));
const uint32_t actual = toBits(vfpu_dot(a, b));
if (expected != actual) {
printf("vfpu_dot(%08x %08x %08x %08x, %08x %08x %08x %08x) = %08x, expected %08x\n",
toBits(a[0]), toBits(a[1]), toBits(a[2]), toBits(a[3]), toBits(b[0]), toBits(b[1]), toBits(b[2]), toBits(b[3]), actual, expected);
return false;
}
return true;
};
for (int n = 0; n < 4000000; n++) {
float a[4], b[4];
const int mode = (int)(rnd() & 15);
const int base = 1 + (int)(rnd() % 254);
const int spread = mode < 8 ? 3 : 40;
for (int i = 0; i < 4; i++) {
if (mode == 15) {
a[i] = fromBits((uint32_t)rnd());
b[i] = fromBits((uint32_t)rnd());
continue;
}
int ea = base + (int)(rnd() % (2 * spread + 1)) - spread;
int eb = 127 + (int)(rnd() % (2 * spread + 1)) - spread;
ea = std::max(0, std::min(254, ea));
eb = std::max(0, std::min(254, eb));
uint32_t xa = ((uint32_t)rnd() & 0x80000000) | (ea << 23) | ((uint32_t)rnd() & 0x7FFFFF);
uint32_t xb = ((uint32_t)rnd() & 0x80000000) | (eb << 23) | ((uint32_t)rnd() & 0x7FFFFF);
switch (rnd() & 63) {
case 0: xa &= 0x80000000; break;
case 1: xb &= 0x807FFFFF; break;
case 2: xa |= 0x7F800000; xa &= 0xFF800000; break;
case 3: xb |= 0x7FC00000; break;
case 4: xa &= 0xFFFF0000; break;
default: break;
}
a[i] = fromBits(xa);
b[i] = fromBits(xb);
}
if (mode == 5) {
// Nearly cancelling products.
a[1] = -a[0];
b[1] = fromBits(toBits(b[0]) ^ ((uint32_t)rnd() & 7));
}
if (!check(a, b))
return false;
}
// Sums just below and above a power of two, whose rounding carries into the next exponent:
// 1.0 from just under it, and inf at the top of the range.
for (int e = 1; e <= 254; e++) {
for (int s = 0; s < 2; s++) {
const uint32_t sign = (uint32_t)s << 31;
for (int k = 1; k <= 40; k++) {
const uint32_t small = e - k >= 1 ? ((uint32_t)(e - k) << 23) | ((uint32_t)k * 0x2AAAA) : 0;
float a[4] = { fromBits(sign | (e << 23)), fromBits((sign ^ 0x80000000u) | small), 0.0f, 0.0f };
float b[4] = { 1.0f, 1.0f, 1.0f, 1.0f };
if (!check(a, b))
return false;
a[0] = fromBits(sign | (e << 23) | 0x7FFFFF);
a[1] = fromBits(sign | small);
if (!check(a, b))
return false;
a[2] = a[1];
if (!check(a, b))
return false;
}
}
}
return true;
}
bool TestVFPUSinCos() {
float sine, cosine;
// Needed for VFPU tables.
@@ -3148,6 +3240,7 @@ TestItem availableTests[] = {
TEST_ITEM(Asin),
TEST_ITEM(SinCos),
TEST_ITEM(VFPUSinCos),
TEST_ITEM(VFPUDot),
TEST_ITEM(MathUtil),
TEST_ITEM(Parsers),
TEST_ITEM(TruncateCpy),