Add typed logicals

This commit is contained in:
Charles Schlosser
2023-02-18 01:23:47 +00:00
committed by Rasmus Munk Larsen
parent e797974689
commit 049a144798
13 changed files with 415 additions and 124 deletions

View File

@@ -27,7 +27,7 @@ struct all_unroller
EIGEN_DEVICE_FUNC static inline bool run(const Derived &mat)
{
return all_unroller<Derived, UnrollCount-1, InnerSize>::run(mat) && mat.coeff(IsRowMajor ? i : j, IsRowMajor ? j : i);
return all_unroller<Derived, UnrollCount-1, InnerSize>::run(mat) && mat.coeff(IsRowMajor ? i : j, IsRowMajor ? j : i) != typename Derived::CoeffReturnType(0);
}
};
@@ -54,7 +54,7 @@ struct any_unroller
EIGEN_DEVICE_FUNC static inline bool run(const Derived &mat)
{
return any_unroller<Derived, UnrollCount-1, InnerSize>::run(mat) || mat.coeff(IsRowMajor ? i : j, IsRowMajor ? j : i);
return any_unroller<Derived, UnrollCount-1, InnerSize>::run(mat) || mat.coeff(IsRowMajor ? i : j, IsRowMajor ? j : i) != typename Derived::CoeffReturnType(0);
}
};
@@ -94,7 +94,7 @@ EIGEN_DEVICE_FUNC inline bool DenseBase<Derived>::all() const
{
for(Index i = 0; i < derived().outerSize(); ++i)
for(Index j = 0; j < derived().innerSize(); ++j)
if (!evaluator.coeff(IsRowMajor ? i : j, IsRowMajor ? j : i)) return false;
if (evaluator.coeff(IsRowMajor ? i : j, IsRowMajor ? j : i) == Scalar(0)) return false;
return true;
}
}
@@ -118,7 +118,7 @@ EIGEN_DEVICE_FUNC inline bool DenseBase<Derived>::any() const
{
for(Index i = 0; i < derived().outerSize(); ++i)
for(Index j = 0; j < derived().innerSize(); ++j)
if (evaluator.coeff(IsRowMajor ? i : j, IsRowMajor ? j : i)) return true;
if (evaluator.coeff(IsRowMajor ? i : j, IsRowMajor ? j : i) != Scalar(0)) return true;
return false;
}
}

View File

@@ -350,8 +350,8 @@ template<typename ExpressionType, int Direction> class VectorwiseOp
typedef typename ReturnType<internal::member_hypotNorm,RealScalar>::Type HypotNormReturnType;
typedef typename ReturnType<internal::member_sum>::Type SumReturnType;
typedef EIGEN_EXPR_BINARYOP_SCALAR_RETURN_TYPE(SumReturnType,Scalar,quotient) MeanReturnType;
typedef typename ReturnType<internal::member_all>::Type AllReturnType;
typedef typename ReturnType<internal::member_any>::Type AnyReturnType;
typedef typename ReturnType<internal::member_all, bool>::Type AllReturnType;
typedef typename ReturnType<internal::member_any, bool>::Type AnyReturnType;
typedef PartialReduxExpr<ExpressionType, internal::member_count<Index,Scalar>, Direction> CountReturnType;
typedef typename ReturnType<internal::member_prod>::Type ProdReturnType;
typedef Reverse<const ExpressionType, Direction> ConstReverseReturnType;

View File

@@ -216,6 +216,7 @@ template<> struct packet_traits<bool> : default_packet_traits
HasAdd = 1,
HasSub = 1,
HasCmp = 1, // note -- only pcmp_eq is defined
HasShift = 0,
HasMul = 1,
HasNegate = 1,

View File

@@ -428,60 +428,168 @@ struct functor_traits<scalar_quotient_op<LhsScalar,RhsScalar> > {
};
};
/** \internal
* \brief Template functor to compute the and of two booleans
* \brief Template functor to compute the and of two scalars as if they were booleans
*
* \sa class CwiseBinaryOp, ArrayBase::operator&&
*/
template <typename Scalar>
struct scalar_boolean_and_op {
EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE bool operator() (const bool& a, const bool& b) const { return a && b; }
template<typename Packet>
EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE const Packet packetOp(const Packet& a, const Packet& b) const
{ return internal::pand(a,b); }
using result_type = Scalar;
// `false` any value `a` that satisfies `a == Scalar(0)`
// `true` is the complement of `false`
EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Scalar operator()(const Scalar& a, const Scalar& b) const {
return (a != Scalar(0)) && (b != Scalar(0)) ? Scalar(1) : Scalar(0);
}
template <typename Packet>
EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet packetOp(const Packet& a, const Packet& b) const {
const Packet cst_one = pset1<Packet>(Scalar(1));
// and(a,b) == !or(!a,!b)
Packet not_a = pcmp_eq(a, pzero(a));
Packet not_b = pcmp_eq(b, pzero(b));
Packet a_nand_b = por(not_a, not_b);
return pandnot(cst_one, a_nand_b);
}
};
template<> struct functor_traits<scalar_boolean_and_op> {
enum {
Cost = NumTraits<bool>::AddCost,
PacketAccess = true
};
template <typename Scalar>
struct functor_traits<scalar_boolean_and_op<Scalar>> {
enum { Cost = NumTraits<Scalar>::AddCost, PacketAccess = packet_traits<Scalar>::HasCmp };
};
/** \internal
* \brief Template functor to compute the or of two booleans
* \brief Template functor to compute the or of two scalars as if they were booleans
*
* \sa class CwiseBinaryOp, ArrayBase::operator||
*/
template <typename Scalar>
struct scalar_boolean_or_op {
EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE bool operator() (const bool& a, const bool& b) const { return a || b; }
template<typename Packet>
EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE const Packet packetOp(const Packet& a, const Packet& b) const
{ return internal::por(a,b); }
using result_type = Scalar;
// `false` any value `a` that satisfies `a == Scalar(0)`
// `true` is the complement of `false`
EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Scalar operator()(const Scalar& a, const Scalar& b) const {
return (a != Scalar(0)) || (b != Scalar(0)) ? Scalar(1) : Scalar(0);
}
template <typename Packet>
EIGEN_STRONG_INLINE Packet packetOp(const Packet& a, const Packet& b) const {
const Packet cst_one = pset1<Packet>(Scalar(1));
// if or(a,b) == 0, then a == 0 and b == 0
// or(a,b) == !nor(a,b)
Packet a_nor_b = pcmp_eq(por(a, b), pzero(a));
return pandnot(cst_one, a_nor_b);
}
};
template<> struct functor_traits<scalar_boolean_or_op> {
enum {
Cost = NumTraits<bool>::AddCost,
PacketAccess = true
};
template <typename Scalar>
struct functor_traits<scalar_boolean_or_op<Scalar>> {
enum { Cost = NumTraits<Scalar>::AddCost, PacketAccess = packet_traits<Scalar>::HasCmp };
};
/** \internal
* \brief Template functor to compute the xor of two booleans
* \brief Template functor to compute the xor of two scalars as if they were booleans
*
* \sa class CwiseBinaryOp, ArrayBase::operator^
*/
template <typename Scalar>
struct scalar_boolean_xor_op {
EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE bool operator() (const bool& a, const bool& b) const { return a ^ b; }
template<typename Packet>
EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE const Packet packetOp(const Packet& a, const Packet& b) const
{ return internal::pxor(a,b); }
using result_type = Scalar;
// `false` any value `a` that satisfies `a == Scalar(0)`
// `true` is the complement of `false`
EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Scalar operator()(const Scalar& a, const Scalar& b) const {
return (a != Scalar(0)) != (b != Scalar(0)) ? Scalar(1) : Scalar(0);
}
template <typename Packet>
EIGEN_STRONG_INLINE Packet packetOp(const Packet& a, const Packet& b) const {
const Packet cst_one = pset1<Packet>(Scalar(1));
// xor(a,b) == xor(!a,!b)
Packet not_a = pcmp_eq(a, pzero(a));
Packet not_b = pcmp_eq(b, pzero(b));
Packet a_xor_b = pxor(not_a, not_b);
return pand(cst_one, a_xor_b);
}
};
template<> struct functor_traits<scalar_boolean_xor_op> {
enum {
Cost = NumTraits<bool>::AddCost,
PacketAccess = true
};
template <typename Scalar>
struct functor_traits<scalar_boolean_xor_op<Scalar>> {
enum { Cost = NumTraits<Scalar>::AddCost, PacketAccess = packet_traits<Scalar>::HasCmp };
};
/** \internal
* \brief Template functor to compute the bitwise and of two scalars
*
* \sa class CwiseBinaryOp, ArrayBase::operator&
*/
template <typename Scalar>
struct scalar_bitwise_and_op {
EIGEN_STATIC_ASSERT(!NumTraits<Scalar>::RequireInitialization, BITWISE OPERATIONS MAY ONLY BE PERFORMED ON PLAIN DATA TYPES )
using result_type = Scalar;
EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Scalar operator()(const Scalar& a, const Scalar& b) const {
Scalar result;
const uint8_t* a_bytes = reinterpret_cast<const uint8_t*>(&a);
const uint8_t* b_bytes = reinterpret_cast<const uint8_t*>(&b);
uint8_t* r_bytes = reinterpret_cast<uint8_t*>(&result);
for (Index i = 0; i < sizeof(Scalar); i++) r_bytes[i] = a_bytes[i] & b_bytes[i];
return result;
}
template <typename Packet>
EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet packetOp(const Packet& a, const Packet& b) const {
return pand(a, b);
}
};
template <typename Scalar>
struct functor_traits<scalar_bitwise_and_op<Scalar>> {
enum { Cost = NumTraits<Scalar>::AddCost, PacketAccess = true };
};
/** \internal
* \brief Template functor to compute the bitwise or of two scalars
*
* \sa class CwiseBinaryOp, ArrayBase::operator|
*/
template <typename Scalar>
struct scalar_bitwise_or_op {
EIGEN_STATIC_ASSERT(!NumTraits<Scalar>::RequireInitialization, BITWISE OPERATIONS MAY ONLY BE PERFORMED ON PLAIN DATA TYPES)
using result_type = Scalar;
EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Scalar operator()(const Scalar& a, const Scalar& b) const {
Scalar result;
const uint8_t* a_bytes = reinterpret_cast<const uint8_t*>(&a);
const uint8_t* b_bytes = reinterpret_cast<const uint8_t*>(&b);
uint8_t* r_bytes = reinterpret_cast<uint8_t*>(&result);
for (Index i = 0; i < sizeof(Scalar); i++) r_bytes[i] = a_bytes[i] | b_bytes[i];
return result;
}
template <typename Packet>
EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet packetOp(const Packet& a, const Packet& b) const {
return por(a, b);
}
};
template <typename Scalar>
struct functor_traits<scalar_bitwise_or_op<Scalar>> {
enum { Cost = NumTraits<Scalar>::AddCost, PacketAccess = true };
};
/** \internal
* \brief Template functor to compute the bitwise xor of two scalars
*
* \sa class CwiseBinaryOp, ArrayBase::operator^
*/
template <typename Scalar>
struct scalar_bitwise_xor_op {
EIGEN_STATIC_ASSERT(!NumTraits<Scalar>::RequireInitialization, BITWISE OPERATIONS MAY ONLY BE PERFORMED ON PLAIN DATA TYPES)
using result_type = Scalar;
EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Scalar operator()(const Scalar& a, const Scalar& b) const {
Scalar result;
const uint8_t* a_bytes = reinterpret_cast<const uint8_t*>(&a);
const uint8_t* b_bytes = reinterpret_cast<const uint8_t*>(&b);
uint8_t* r_bytes = reinterpret_cast<uint8_t*>(&result);
for (Index i = 0; i < sizeof(Scalar); i++) r_bytes[i] = a_bytes[i] ^ b_bytes[i];
return result;
}
template <typename Packet>
EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet packetOp(const Packet& a, const Packet& b) const {
return pxor(a, b);
}
};
template <typename Scalar>
struct functor_traits<scalar_bitwise_xor_op<Scalar>> {
enum { Cost = NumTraits<Scalar>::AddCost, PacketAccess = true };
};
/** \internal

View File

@@ -913,19 +913,54 @@ struct functor_traits<scalar_isfinite_op<Scalar> >
};
/** \internal
* \brief Template functor to compute the logical not of a boolean
* \brief Template functor to compute the logical not of a scalar as if it were a boolean
*
* \sa class CwiseUnaryOp, ArrayBase::operator!
*/
template<typename Scalar> struct scalar_boolean_not_op {
EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE bool operator() (const bool& a) const { return !a; }
template <typename Scalar>
struct scalar_boolean_not_op {
using result_type = Scalar;
// `false` any value `a` that satisfies `a == Scalar(0)`
// `true` is the complement of `false`
EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Scalar operator()(const Scalar& a) const {
return a == Scalar(0) ? Scalar(1) : Scalar(0);
}
template <typename Packet>
EIGEN_STRONG_INLINE Packet packetOp(const Packet& a) const {
const Packet cst_one = pset1<Packet>(Scalar(1));
Packet not_a = pcmp_eq(a, pzero(a));
return pand(not_a, cst_one);
}
};
template<typename Scalar>
struct functor_traits<scalar_boolean_not_op<Scalar> > {
enum {
Cost = NumTraits<bool>::AddCost,
PacketAccess = false
};
template <typename Scalar>
struct functor_traits<scalar_boolean_not_op<Scalar>> {
enum { Cost = NumTraits<Scalar>::AddCost, PacketAccess = packet_traits<Scalar>::HasCmp };
};
/** \internal
* \brief Template functor to compute the bitwise not of a scalar
*
* \sa class CwiseUnaryOp, ArrayBase::operator~
*/
template <typename Scalar>
struct scalar_bitwise_not_op {
EIGEN_STATIC_ASSERT(!NumTraits<Scalar>::RequireInitialization, BITWISE OPERATIONS MAY ONLY BE PERFORMED ON PLAIN DATA TYPES)
using result_type = Scalar;
EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Scalar operator()(const Scalar& a) const {
Scalar result;
const uint8_t* a_bytes = reinterpret_cast<const uint8_t*>(&a);
uint8_t* r_bytes = reinterpret_cast<uint8_t*>(&result);
for (Index i = 0; i < sizeof(Scalar); i++) r_bytes[i] = ~a_bytes[i];
return result;
}
template <typename Packet>
EIGEN_STRONG_INLINE Packet packetOp(const Packet& a) const {
return pandnot(ptrue(a), a);
}
};
template <typename Scalar>
struct functor_traits<scalar_bitwise_not_op<Scalar>> {
enum { Cost = NumTraits<Scalar>::AddCost, PacketAccess = true };
};
/** \internal

View File

@@ -210,6 +210,15 @@ struct scalar_unary_pow_op;
template<typename LhsScalar,typename RhsScalar=LhsScalar> struct scalar_hypot_op;
template<typename LhsScalar,typename RhsScalar=LhsScalar> struct scalar_product_op;
template<typename LhsScalar,typename RhsScalar=LhsScalar> struct scalar_quotient_op;
// logical and bitwise operations
template <typename Scalar> struct scalar_boolean_and_op;
template <typename Scalar> struct scalar_boolean_or_op;
template <typename Scalar> struct scalar_boolean_xor_op;
template <typename Scalar> struct scalar_boolean_not_op;
template <typename Scalar> struct scalar_bitwise_and_op;
template <typename Scalar> struct scalar_bitwise_or_op;
template <typename Scalar> struct scalar_bitwise_xor_op;
template <typename Scalar> struct scalar_bitwise_not_op;
// SpecialFunctions module
template<typename Scalar> struct scalar_lgamma_op;