Fix rint for SSE/NEON.

It seems *sometimes* with aggressive optimizations the combination
`psub(padd(a, b), b)` trick to force rounding is compiled away. Here
we replace with inline assembly to prevent this (I tried `volatile`,
but that leads to additional loads from memory).

Also fixed an edge case for large inputs `a` where adding `b` bumps
the value up a power of two and ends up rounding away more than
just the fractional part.  If we are over `2^digits` then just return
the input.  This edge case was missed in the test since the test was
comparing approximate equality, which was still satisfied.  Adding
a strict equality option catches it.
This commit is contained in:
Antonio Sanchez
2021-03-03 09:41:46 -08:00
parent 199c5f2b47
commit e72dfeb8b9
4 changed files with 83 additions and 29 deletions

View File

@@ -3207,20 +3207,34 @@ template<> EIGEN_STRONG_INLINE Packet4f pceil<Packet4f>(const Packet4f& a)
template<> EIGEN_STRONG_INLINE Packet4f print(const Packet4f& a) {
// Adds and subtracts signum(a) * 2^23 to force rounding.
const Packet4f offset =
pselect(pcmp_lt(a, pzero(a)),
pset1<Packet4f>(-static_cast<float>(1<<23)),
pset1<Packet4f>(+static_cast<float>(1<<23)));
return psub(padd(a, offset), offset);
const Packet4f limit = pset1<Packet4f>(static_cast<float>(1<<23));
const Packet4f abs_a = pabs(a);
// Inline asm to prevent the compiler from optimizing away the
// addition and subtraction.
// Packet4f r = psub(padd(abs_a, limit), limit);
Packet4f r = abs_a;
__asm__ ("vadd.f32 %[r], %[r], %[limit]\n\t"
"vsub.f32 %[r], %[r], %[limit]" : [r] "+x" (r) : [limit] "x" (limit));
// If greater than limit, simply return a. Otherwise, account for sign.
r = pselect(pcmp_lt(abs_a, limit),
pselect(pcmp_lt(a, pzero(a)), pnegate(r), r), a);
return r;
}
template<> EIGEN_STRONG_INLINE Packet2f print(const Packet2f& a) {
// Adds and subtracts signum(a) * 2^23 to force rounding.
const Packet2f offset =
pselect(pcmp_lt(a, pzero(a)),
pset1<Packet2f>(-static_cast<float>(1<<23)),
pset1<Packet2f>(+static_cast<float>(1<<23)));
return psub(padd(a, offset), offset);
const Packet2f limit = pset1<Packet2f>(static_cast<float>(1<<23));
const Packet2f abs_a = pabs(a);
// Inline asm to prevent the compiler from optimizing away the
// addition and subtraction.
// Packet4f r = psub(padd(abs_a, limit), limit);
Packet2f r = abs_a;
__asm__ ("vadd.f32 %[r], %[r], %[limit]\n\t"
"vsub.f32 %[r], %[r], %[limit]" : [r] "+x" (r) : [limit] "x" (limit));
// If greater than limit, simply return a. Otherwise, account for sign.
r = pselect(pcmp_lt(abs_a, limit),
pselect(pcmp_lt(a, pzero(a)), pnegate(r), r), a);
return r;
}
template<> EIGEN_STRONG_INLINE Packet4f pfloor<Packet4f>(const Packet4f& a)

View File

@@ -646,20 +646,35 @@ template<> EIGEN_STRONG_INLINE Packet2d pfloor<Packet2d>(const Packet2d& a) { re
#else
template<> EIGEN_STRONG_INLINE Packet4f print(const Packet4f& a) {
// Adds and subtracts signum(a) * 2^23 to force rounding.
const Packet4f offset =
pselect(pcmp_lt(a, pzero(a)),
pset1<Packet4f>(-static_cast<float>(1<<23)),
pset1<Packet4f>(+static_cast<float>(1<<23)));
return psub(padd(a, offset), offset);
const Packet4f limit = pset1<Packet4f>(static_cast<float>(1<<23));
const Packet4f abs_a = pabs(a);
// Inline asm to prevent the compiler from optimizing away the
// addition and subtraction.
// Packet4f r = psub(padd(abs_a, limit), limit);
Packet4f r = abs_a;
__asm__ ("addps %[limit], %[r]\n\t"
"subps %[limit], %[r]" : [r] "+x" (r) : [limit] "x" (limit));
// If greater than limit, simply return a. Otherwise, account for sign.
r = pselect(pcmp_lt(abs_a, limit),
pselect(pcmp_lt(a, pzero(a)), pnegate(r), r), a);
return r;
}
template<> EIGEN_STRONG_INLINE Packet2d print(const Packet2d& a) {
// Adds and subtracts signum(a) * 2^52 to force rounding.
const Packet2d offset =
pselect(pcmp_lt(a, pzero(a)),
pset1<Packet2d>(-static_cast<double>(1ull<<52)),
pset1<Packet2d>(+static_cast<double>(1ull<<52)));
return psub(padd(a, offset), offset);
const Packet2d limit = pset1<Packet2d>(static_cast<double>(1ull<<52));
const Packet2d abs_a = pabs(a);
// Inline asm to prevent the compiler from optimizing away the
// addition and subtraction.
// Packet2d r = psub(padd(abs_a, limit), limit);
Packet2d r = abs_a;
asm("addpd %[limit], %[r] \n\t"
"subpd %[limit], %[r]" : [r] "+x" (r) : [limit] "x" (limit));
// If greater than limit, simply return a. Otherwise, account for sign.
r = pselect(pcmp_lt(abs_a, limit),
pselect(pcmp_lt(a, pzero(a)), pnegate(r), r), a);
return r;
}
template<> EIGEN_STRONG_INLINE Packet4f pfloor<Packet4f>(const Packet4f& a)