5 #ifndef GKO_PUBLIC_CORE_BASE_RANGE_HPP_
6 #define GKO_PUBLIC_CORE_BASE_RANGE_HPP_
11 #include <ginkgo/core/base/math.hpp>
12 #include <ginkgo/core/base/types.hpp>
13 #include <ginkgo/core/base/utils.hpp>
54 :
span{point, point + 1}
93 GKO_ATTRIBUTES GKO_INLINE constexpr
bool operator<(
const span& first,
100 GKO_ATTRIBUTES GKO_INLINE constexpr
bool operator<=(
const span& first,
103 return first.end <= second.begin;
107 GKO_ATTRIBUTES GKO_INLINE constexpr
bool operator>(
const span& first,
110 return second < first;
114 GKO_ATTRIBUTES GKO_INLINE constexpr
bool operator>=(
const span& first,
117 return second <= first;
121 GKO_ATTRIBUTES GKO_INLINE constexpr
bool operator==(
const span& first,
124 return first.begin == second.begin && first.end == second.end;
128 GKO_ATTRIBUTES GKO_INLINE constexpr
bool operator!=(
const span& first,
131 return !(first == second);
145 template <
size_type CurrentDimension = 0,
typename FirstRange,
146 typename SecondRange>
147 GKO_ATTRIBUTES constexpr GKO_INLINE
148 std::enable_if_t<(CurrentDimension >=
max(FirstRange::dimensionality,
149 SecondRange::dimensionality)),
151 equal_dimensions(
const FirstRange&,
const SecondRange&)
156 template <
size_type CurrentDimension = 0,
typename FirstRange,
157 typename SecondRange>
158 GKO_ATTRIBUTES constexpr GKO_INLINE
159 std::enable_if_t<(CurrentDimension <
max(FirstRange::dimensionality,
160 SecondRange::dimensionality)),
162 equal_dimensions(
const FirstRange& first,
const SecondRange& second)
164 return first.length(CurrentDimension) == second.length(CurrentDimension) &&
165 equal_dimensions<CurrentDimension + 1>(first, second);
178 template <
class First,
class... Rest>
179 struct head<First, Rest...> {
186 template <
class... T>
187 using head_t =
typename head<T...>::type;
301 template <
typename Accessor>
328 typename... AccessorParams,
329 typename = std::enable_if_t<
330 sizeof...(AccessorParams) != 1 ||
332 range, std::decay<detail::head_t<AccessorParams...>>>::value>>
333 GKO_ATTRIBUTES constexpr
explicit range(AccessorParams&&... params)
334 : accessor_{std::forward<AccessorParams>(params)...}
349 template <
typename... DimensionTypes>
350 GKO_ATTRIBUTES constexpr
auto operator()(DimensionTypes&&... dimensions)
351 const -> decltype(std::declval<accessor>()(
352 std::forward<DimensionTypes>(dimensions)...))
355 "Too many dimensions in range call");
356 return accessor_(std::forward<DimensionTypes>(dimensions)...);
367 template <
typename OtherAccessor>
371 GKO_ASSERT(detail::equal_dimensions(*
this, other));
372 accessor_.copy_from(other);
391 GKO_ASSERT(detail::equal_dimensions(*
this, other));
407 return accessor_.length(dimension);
445 enum class operation_kind { range_by_range, scalar_by_range, range_by_scalar };
448 template <
typename Accessor,
typename Operation>
449 struct implement_unary_operation {
450 using accessor = Accessor;
451 static constexpr
size_type dimensionality = accessor::dimensionality;
453 GKO_ATTRIBUTES constexpr
explicit implement_unary_operation(
454 const Accessor& operand)
458 template <
typename... DimensionTypes>
459 GKO_ATTRIBUTES constexpr
auto operator()(
460 const DimensionTypes&... dimensions)
const
461 -> decltype(Operation::evaluate(std::declval<accessor>(),
464 return Operation::evaluate(operand, dimensions...);
469 return operand.length(dimension);
472 template <
typename OtherAccessor>
473 GKO_ATTRIBUTES
void copy_from(
const OtherAccessor& other)
const =
delete;
475 const accessor operand;
479 template <operation_kind Kind,
typename FirstOperand,
typename SecondOperand,
481 struct implement_binary_operation {};
483 template <
typename FirstAccessor,
typename SecondAccessor,
typename Operation>
484 struct implement_binary_operation<operation_kind::range_by_range, FirstAccessor,
485 SecondAccessor, Operation> {
486 using first_accessor = FirstAccessor;
487 using second_accessor = SecondAccessor;
488 static_assert(first_accessor::dimensionality ==
489 second_accessor::dimensionality,
490 "Both ranges need to have the same number of dimensions");
491 static constexpr
size_type dimensionality = first_accessor::dimensionality;
493 GKO_ATTRIBUTES
explicit implement_binary_operation(
494 const FirstAccessor& first,
const SecondAccessor& second)
495 : first{first}, second{second}
497 GKO_ASSERT(gko::detail::equal_dimensions(first, second));
500 template <
typename... DimensionTypes>
501 GKO_ATTRIBUTES constexpr
auto operator()(
502 const DimensionTypes&... dimensions)
const
503 -> decltype(Operation::evaluate_range_by_range(
504 std::declval<first_accessor>(), std::declval<second_accessor>(),
507 return Operation::evaluate_range_by_range(first, second, dimensions...);
512 return first.length(dimension);
515 template <
typename OtherAccessor>
516 GKO_ATTRIBUTES
void copy_from(
const OtherAccessor& other)
const =
delete;
518 const first_accessor first;
519 const second_accessor second;
522 template <
typename FirstOperand,
typename SecondAccessor,
typename Operation>
523 struct implement_binary_operation<operation_kind::scalar_by_range, FirstOperand,
524 SecondAccessor, Operation> {
525 using second_accessor = SecondAccessor;
526 static constexpr
size_type dimensionality = second_accessor::dimensionality;
528 GKO_ATTRIBUTES constexpr
explicit implement_binary_operation(
529 const FirstOperand& first,
const SecondAccessor& second)
530 : first{first}, second{second}
533 template <
typename... DimensionTypes>
534 GKO_ATTRIBUTES constexpr
auto operator()(
535 const DimensionTypes&... dimensions)
const
536 -> decltype(Operation::evaluate_scalar_by_range(
537 std::declval<FirstOperand>(), std::declval<second_accessor>(),
540 return Operation::evaluate_scalar_by_range(first, second,
546 return second.length(dimension);
549 template <
typename OtherAccessor>
550 GKO_ATTRIBUTES
void copy_from(
const OtherAccessor& other)
const =
delete;
552 const FirstOperand first;
553 const second_accessor second;
556 template <
typename FirstAccessor,
typename SecondOperand,
typename Operation>
557 struct implement_binary_operation<operation_kind::range_by_scalar,
558 FirstAccessor, SecondOperand, Operation> {
559 using first_accessor = FirstAccessor;
560 static constexpr
size_type dimensionality = first_accessor::dimensionality;
562 GKO_ATTRIBUTES constexpr
explicit implement_binary_operation(
563 const FirstAccessor& first,
const SecondOperand& second)
564 : first{first}, second{second}
567 template <
typename... DimensionTypes>
568 GKO_ATTRIBUTES constexpr
auto operator()(
569 const DimensionTypes&... dimensions)
const
570 -> decltype(Operation::evaluate_range_by_scalar(
571 std::declval<first_accessor>(), std::declval<SecondOperand>(),
574 return Operation::evaluate_range_by_scalar(first, second,
580 return first.length(dimension);
583 template <
typename OtherAccessor>
584 GKO_ATTRIBUTES
void copy_from(
const OtherAccessor& other)
const =
delete;
586 const first_accessor first;
587 const SecondOperand second;
593 #define GKO_DEPRECATED_UNARY_RANGE_OPERATION(_operation_deprecated_name, \
595 namespace accessor { \
596 template <typename Operand> \
597 struct GKO_DEPRECATED("Please use " #_operation_name) \
598 _operation_deprecated_name : _operation_name<Operand> {}; \
600 static_assert(true, \
601 "This assert is used to counter the false positive extra " \
602 "semi-colon warnings")
605 #define GKO_ENABLE_UNARY_RANGE_OPERATION(_operation_name, _operator_name, \
607 namespace accessor { \
608 template <typename Operand> \
609 struct _operation_name \
610 : ::gko::detail::implement_unary_operation<Operand, \
611 ::gko::_operator> { \
612 using ::gko::detail::implement_unary_operation< \
613 Operand, ::gko::_operator>::implement_unary_operation; \
616 GKO_BIND_UNARY_RANGE_OPERATION_TO_OPERATOR(_operation_name, _operator_name)
619 #define GKO_BIND_UNARY_RANGE_OPERATION_TO_OPERATOR(_operation_name, \
621 template <typename Accessor> \
622 GKO_ATTRIBUTES constexpr GKO_INLINE \
623 range<accessor::_operation_name<Accessor>> \
624 _operator_name(const range<Accessor>& operand) \
626 return range<accessor::_operation_name<Accessor>>( \
627 operand.get_accessor()); \
629 static_assert(true, \
630 "This assert is used to counter the false positive extra " \
631 "semi-colon warnings")
634 #define GKO_DEFINE_SIMPLE_UNARY_OPERATION(_name, ...) \
637 template <typename Operand> \
638 GKO_ATTRIBUTES static constexpr auto simple_evaluate_impl( \
639 const Operand& operand) -> decltype(__VA_ARGS__) \
641 return __VA_ARGS__; \
645 template <typename AccessorType, typename... DimensionTypes> \
646 GKO_ATTRIBUTES static constexpr auto evaluate( \
647 const AccessorType& accessor, const DimensionTypes&... dimensions) \
648 -> decltype(simple_evaluate_impl(accessor(dimensions...))) \
650 return simple_evaluate_impl(accessor(dimensions...)); \
660 GKO_DEFINE_SIMPLE_UNARY_OPERATION(
unary_plus, +operand);
661 GKO_DEFINE_SIMPLE_UNARY_OPERATION(
unary_minus, -operand);
664 GKO_DEFINE_SIMPLE_UNARY_OPERATION(
logical_not, !operand);
667 GKO_DEFINE_SIMPLE_UNARY_OPERATION(
bitwise_not, ~(operand));
684 GKO_ENABLE_UNARY_RANGE_OPERATION(unary_plus,
operator+,
685 accessor::detail::unary_plus);
686 GKO_ENABLE_UNARY_RANGE_OPERATION(
unary_minus,
operator-,
687 accessor::detail::unary_minus);
690 GKO_ENABLE_UNARY_RANGE_OPERATION(
logical_not,
operator!,
691 accessor::detail::logical_not);
694 GKO_ENABLE_UNARY_RANGE_OPERATION(
bitwise_not,
operator~,
695 accessor::detail::bitwise_not);
700 accessor::detail::zero_operation);
702 accessor::detail::one_operation);
704 accessor::detail::abs_operation);
706 accessor::detail::real_operation);
708 accessor::detail::imag_operation);
710 accessor::detail::conj_operation);
712 accessor::detail::squared_norm_operation);
725 template <
typename Accessor>
727 using accessor = Accessor;
728 static constexpr
size_type dimensionality = accessor::dimensionality;
731 const Accessor& operand)
735 template <
typename FirstDimensionType,
typename SecondDimensionType,
736 typename... DimensionTypes>
737 GKO_ATTRIBUTES constexpr
auto operator()(
738 const FirstDimensionType& first_dim,
739 const SecondDimensionType& second_dim,
740 const DimensionTypes&... dims)
const
741 -> decltype(std::declval<accessor>()(second_dim, first_dim, dims...))
743 return operand(second_dim, first_dim, dims...);
748 return dimension < 2 ? operand.length(dimension ^ 1)
749 : operand.length(dimension);
752 template <
typename OtherAccessor>
753 GKO_ATTRIBUTES
void copy_from(
const OtherAccessor& other)
const =
delete;
755 const accessor operand;
762 GKO_BIND_UNARY_RANGE_OPERATION_TO_OPERATOR(transpose_operation,
transpose);
765 #undef GKO_DEPRECATED_UNARY_RANGE_OPERATION
766 #undef GKO_DEFINE_SIMPLE_UNARY_OPERATION
767 #undef GKO_ENABLE_UNARY_RANGE_OPERATION
770 #define GKO_ENABLE_BINARY_RANGE_OPERATION(_operation_name, _operator_name, \
772 namespace accessor { \
773 template <::gko::detail::operation_kind Kind, typename FirstOperand, \
774 typename SecondOperand> \
775 struct _operation_name \
776 : ::gko::detail::implement_binary_operation< \
777 Kind, FirstOperand, SecondOperand, ::gko::_operator> { \
778 using ::gko::detail::implement_binary_operation< \
779 Kind, FirstOperand, SecondOperand, \
780 ::gko::_operator>::implement_binary_operation; \
783 GKO_BIND_RANGE_OPERATION_TO_OPERATOR(_operation_name, _operator_name); \
784 static_assert(true, \
785 "This assert is used to counter the false positive extra " \
786 "semi-colon warnings")
789 #define GKO_BIND_RANGE_OPERATION_TO_OPERATOR(_operation_name, _operator_name) \
790 template <typename Accessor> \
791 GKO_ATTRIBUTES constexpr GKO_INLINE range<accessor::_operation_name< \
792 ::gko::detail::operation_kind::range_by_range, Accessor, Accessor>> \
793 _operator_name(const range<Accessor>& first, \
794 const range<Accessor>& second) \
796 return range<accessor::_operation_name< \
797 ::gko::detail::operation_kind::range_by_range, Accessor, \
798 Accessor>>(first.get_accessor(), second.get_accessor()); \
801 template <typename FirstAccessor, typename SecondAccessor> \
802 GKO_ATTRIBUTES constexpr GKO_INLINE range<accessor::_operation_name< \
803 ::gko::detail::operation_kind::range_by_range, FirstAccessor, \
805 _operator_name(const range<FirstAccessor>& first, \
806 const range<SecondAccessor>& second) \
808 return range<accessor::_operation_name< \
809 ::gko::detail::operation_kind::range_by_range, FirstAccessor, \
810 SecondAccessor>>(first.get_accessor(), second.get_accessor()); \
813 template <typename FirstAccessor, typename SecondOperand> \
814 GKO_ATTRIBUTES constexpr GKO_INLINE range<accessor::_operation_name< \
815 ::gko::detail::operation_kind::range_by_scalar, FirstAccessor, \
817 _operator_name(const range<FirstAccessor>& first, \
818 const SecondOperand& second) \
820 return range<accessor::_operation_name< \
821 ::gko::detail::operation_kind::range_by_scalar, FirstAccessor, \
822 SecondOperand>>(first.get_accessor(), second); \
825 template <typename FirstOperand, typename SecondAccessor> \
826 GKO_ATTRIBUTES constexpr GKO_INLINE range<accessor::_operation_name< \
827 ::gko::detail::operation_kind::scalar_by_range, FirstOperand, \
829 _operator_name(const FirstOperand& first, \
830 const range<SecondAccessor>& second) \
832 return range<accessor::_operation_name< \
833 ::gko::detail::operation_kind::scalar_by_range, FirstOperand, \
834 SecondAccessor>>(first, second.get_accessor()); \
836 static_assert(true, \
837 "This assert is used to counter the false positive extra " \
838 "semi-colon warnings")
841 #define GKO_DEPRECATED_SIMPLE_BINARY_OPERATION(_deprecated_name, _name) \
842 struct GKO_DEPRECATED("Please use " #_name) _deprecated_name : _name {}
844 #define GKO_DEFINE_SIMPLE_BINARY_OPERATION(_name, ...) \
847 template <typename FirstOperand, typename SecondOperand> \
848 GKO_ATTRIBUTES constexpr static auto simple_evaluate_impl( \
849 const FirstOperand& first, const SecondOperand& second) \
850 -> decltype(__VA_ARGS__) \
852 return __VA_ARGS__; \
856 template <typename FirstAccessor, typename SecondAccessor, \
857 typename... DimensionTypes> \
858 GKO_ATTRIBUTES static constexpr auto evaluate_range_by_range( \
859 const FirstAccessor& first, const SecondAccessor& second, \
860 const DimensionTypes&... dims) \
861 -> decltype(simple_evaluate_impl(first(dims...), second(dims...))) \
863 return simple_evaluate_impl(first(dims...), second(dims...)); \
866 template <typename FirstOperand, typename SecondAccessor, \
867 typename... DimensionTypes> \
868 GKO_ATTRIBUTES static constexpr auto evaluate_scalar_by_range( \
869 const FirstOperand& first, const SecondAccessor& second, \
870 const DimensionTypes&... dims) \
871 -> decltype(simple_evaluate_impl(first, second(dims...))) \
873 return simple_evaluate_impl(first, second(dims...)); \
876 template <typename FirstAccessor, typename SecondOperand, \
877 typename... DimensionTypes> \
878 GKO_ATTRIBUTES static constexpr auto evaluate_range_by_scalar( \
879 const FirstAccessor& first, const SecondOperand& second, \
880 const DimensionTypes&... dims) \
881 -> decltype(simple_evaluate_impl(first(dims...), second)) \
883 return simple_evaluate_impl(first(dims...), second); \
893 GKO_DEFINE_SIMPLE_BINARY_OPERATION(add, first + second);
894 GKO_DEFINE_SIMPLE_BINARY_OPERATION(sub, first - second);
895 GKO_DEFINE_SIMPLE_BINARY_OPERATION(mul, first* second);
896 GKO_DEFINE_SIMPLE_BINARY_OPERATION(div, first / second);
897 GKO_DEFINE_SIMPLE_BINARY_OPERATION(mod, first % second);
900 GKO_DEFINE_SIMPLE_BINARY_OPERATION(less, first < second);
901 GKO_DEFINE_SIMPLE_BINARY_OPERATION(greater, first > second);
902 GKO_DEFINE_SIMPLE_BINARY_OPERATION(less_or_equal, first <= second);
903 GKO_DEFINE_SIMPLE_BINARY_OPERATION(greater_or_equal, first >= second);
904 GKO_DEFINE_SIMPLE_BINARY_OPERATION(equal, first == second);
905 GKO_DEFINE_SIMPLE_BINARY_OPERATION(not_equal, first != second);
908 GKO_DEFINE_SIMPLE_BINARY_OPERATION(logical_or, first || second);
909 GKO_DEFINE_SIMPLE_BINARY_OPERATION(logical_and, first&& second);
912 GKO_DEFINE_SIMPLE_BINARY_OPERATION(bitwise_or, first | second);
913 GKO_DEFINE_SIMPLE_BINARY_OPERATION(bitwise_and, first& second);
914 GKO_DEFINE_SIMPLE_BINARY_OPERATION(bitwise_xor, first ^ second);
915 GKO_DEFINE_SIMPLE_BINARY_OPERATION(left_shift, first << second);
916 GKO_DEFINE_SIMPLE_BINARY_OPERATION(right_shift, first >> second);
919 GKO_DEFINE_SIMPLE_BINARY_OPERATION(max_operation,
max(first, second));
920 GKO_DEFINE_SIMPLE_BINARY_OPERATION(min_operation,
min(first, second));
922 GKO_DEPRECATED_SIMPLE_BINARY_OPERATION(max_operaton, max_operation);
923 GKO_DEPRECATED_SIMPLE_BINARY_OPERATION(min_operaton, min_operation);
929 GKO_ENABLE_BINARY_RANGE_OPERATION(
add,
operator+, accessor::detail::add);
930 GKO_ENABLE_BINARY_RANGE_OPERATION(
sub,
operator-, accessor::detail::sub);
931 GKO_ENABLE_BINARY_RANGE_OPERATION(
mul,
operator*, accessor::detail::mul);
932 GKO_ENABLE_BINARY_RANGE_OPERATION(
div,
operator/, accessor::detail::div);
933 GKO_ENABLE_BINARY_RANGE_OPERATION(
mod,
operator%, accessor::detail::mod);
936 GKO_ENABLE_BINARY_RANGE_OPERATION(
less,
operator<, accessor::detail::less);
937 GKO_ENABLE_BINARY_RANGE_OPERATION(
greater,
operator>,
938 accessor::detail::greater);
940 accessor::detail::less_or_equal);
942 accessor::detail::greater_or_equal);
943 GKO_ENABLE_BINARY_RANGE_OPERATION(
equal,
operator==, accessor::detail::equal);
944 GKO_ENABLE_BINARY_RANGE_OPERATION(
not_equal,
operator!=,
945 accessor::detail::not_equal);
948 GKO_ENABLE_BINARY_RANGE_OPERATION(
logical_or,
operator||,
949 accessor::detail::logical_or);
950 GKO_ENABLE_BINARY_RANGE_OPERATION(
logical_and,
operator&&,
951 accessor::detail::logical_and);
954 GKO_ENABLE_BINARY_RANGE_OPERATION(
bitwise_or,
operator|,
955 accessor::detail::bitwise_or);
956 GKO_ENABLE_BINARY_RANGE_OPERATION(
bitwise_and,
operator&,
957 accessor::detail::bitwise_and);
958 GKO_ENABLE_BINARY_RANGE_OPERATION(
bitwise_xor,
operator^,
959 accessor::detail::bitwise_xor);
960 GKO_ENABLE_BINARY_RANGE_OPERATION(
left_shift,
operator<<,
961 accessor::detail::left_shift);
962 GKO_ENABLE_BINARY_RANGE_OPERATION(
right_shift,
operator>>,
963 accessor::detail::right_shift);
967 accessor::detail::max_operation);
969 accessor::detail::min_operation);
976 template <gko::detail::operation_kind Kind,
typename FirstAccessor,
977 typename SecondAccessor>
979 static_assert(Kind == gko::detail::operation_kind::range_by_range,
980 "Matrix multiplication expects both operands to be ranges");
981 using first_accessor = FirstAccessor;
982 using second_accessor = SecondAccessor;
983 static_assert(first_accessor::dimensionality ==
984 second_accessor::dimensionality,
985 "Both ranges need to have the same number of dimensions");
986 static constexpr
size_type dimensionality = first_accessor::dimensionality;
988 GKO_ATTRIBUTES
explicit mmul_operation(
const FirstAccessor& first,
989 const SecondAccessor& second)
990 : first{first}, second{second}
992 GKO_ASSERT(first.length(1) == second.length(0));
993 GKO_ASSERT(gko::detail::equal_dimensions<2>(first, second));
996 template <
typename FirstDimension,
typename SecondDimension,
997 typename... DimensionTypes>
998 GKO_ATTRIBUTES
auto operator()(
const FirstDimension& row,
999 const SecondDimension& col,
1000 const DimensionTypes&... rest)
const
1001 -> decltype(std::declval<FirstAccessor>()(row, 0, rest...) *
1002 std::declval<SecondAccessor>()(0, col, rest...) +
1003 std::declval<FirstAccessor>()(row, 1, rest...) *
1004 std::declval<SecondAccessor>()(1, col, rest...))
1007 decltype(first(row, 0, rest...) * second(0, col, rest...) +
1008 first(row, 1, rest...) * second(1, col, rest...));
1009 GKO_ASSERT(first.length(1) == second.length(0));
1010 auto result = zero<result_type>();
1011 const auto size = first.length(1);
1012 for (
auto i =
zero(size); i < size; ++i) {
1013 result += first(row, i, rest...) * second(i, col, rest...);
1020 return dimension == 1 ? second.length(1) : first.length(dimension);
1023 template <
typename OtherAccessor>
1024 GKO_ATTRIBUTES
void copy_from(
const OtherAccessor& other)
const =
delete;
1026 const first_accessor first;
1027 const second_accessor second;
1034 GKO_BIND_RANGE_OPERATION_TO_OPERATOR(mmul_operation, mmul);
1037 #undef GKO_DEFINE_SIMPLE_BINARY_OPERATION
1038 #undef GKO_ENABLE_BINARY_RANGE_OPERATION
1044 #endif // GKO_PUBLIC_CORE_BASE_RANGE_HPP_