3#include <stratax/core/dtypes/Concepts.hpp>
4#include <stratax/core/validation/Validation.hpp>
5#include <stratax/ops/Broadcasting.hpp>
12namespace stratax::core::bitwise_detail {
25template<
typename Value, Integral Count>
26constexpr bool valid_shift_count(
const Count& count)
noexcept
28 if constexpr (std::is_signed_v<std::remove_cvref_t<Count>>)
36 using value_type = std::remove_cvref_t<Value>;
38 return static_cast<std::uintmax_t
>(count) <
39 sizeof(value_type) * CHAR_BIT;
51template<
typename Value, Integral Count>
52void require_valid_shift_count(
const Count& count)
54 if (!valid_shift_count<Value>(count))
56 throw Exceptions::StrataxError(
57 "Shift count must be non-negative and less than the bit width of the shifted value.");
69template<Array A, Integral Count,
typename Op>
70requires Integral<typename A::value_type>
76 using value_type =
typename A::value_type;
78 require_valid_shift_count<value_type>(rhs);
81 stratax::core::rebind_array_t<A, value_type>;
83 result_type result(lhs.shape());
85 auto out = result.begin();
87 for (
auto it = lhs.begin(); it != lhs.end(); ++it, ++out)
89 *out =
static_cast<value_type
>(
106template<Array L, Array R,
typename Op>
108 Integral<typename L::value_type> &&
109 Integral<typename R::value_type>
116 using value_type =
typename L::value_type;
119 stratax::core::promote_array_t<
124 const auto result_shape =
125 broadcasted_shape(lhs.shape(), rhs.shape());
127 result_type result(result_shape);
129 for (std::size_t i = 0; i < result.size(); ++i)
131 const std::size_t lhs_index =
132 stratax::core::broadcast_detail::flat_operand_index(
137 const std::size_t rhs_index =
138 stratax::core::broadcast_detail::flat_operand_index(
143 const auto count = rhs[rhs_index];
145 require_valid_shift_count<value_type>(count);
147 result[i] =
static_cast<value_type
>(
148 op(lhs[lhs_index], count));
162template<Integral Scalar, Array A,
typename Op>
163requires Integral<typename A::value_type>
164auto scalar_shift_array_op(
169 using value_type = std::remove_cvref_t<Scalar>;
172 stratax::core::rebind_array_t<
176 result_type result(rhs.shape());
178 auto out = result.begin();
180 for (
auto it = rhs.begin(); it != rhs.end(); ++it, ++out)
182 require_valid_shift_count<value_type>(*it);
184 *out =
static_cast<value_type
>(
200template<Array L, Array R,
typename Op>
202 Integral<typename L::value_type> &&
203 Integral<typename R::value_type>
205L& compound_bitwise_op(
210 const auto result_shape =
211 broadcasted_shape(lhs.shape(), rhs.shape());
213 if (result_shape != lhs.shape())
215 throw Exceptions::BroadcastError(
216 "Compound bitwise assignment cannot change the left-hand shape.");
219 for (std::size_t i = 0; i < lhs.size(); ++i)
221 const std::size_t rhs_index =
222 stratax::core::broadcast_detail::flat_operand_index(
227 lhs[i] =
static_cast<typename L::value_type
>(
228 op(lhs[i], rhs[rhs_index]));
241template<Array A, Integral Scalar,
typename Op>
242requires Integral<typename A::value_type>
243A& compound_scalar_bitwise_op(
248 for (std::size_t i = 0; i < lhs.size(); ++i)
250 lhs[i] =
static_cast<typename A::value_type
>(
271template<Array L, Array R,
typename Op>
273 Integral<typename L::value_type> &&
274 Integral<typename R::value_type>
281 const auto result_shape =
282 broadcasted_shape(lhs.shape(), rhs.shape());
284 if (result_shape != lhs.shape())
286 throw Exceptions::BroadcastError(
287 "Compound shift assignment cannot change the left-hand shape.");
290 using value_type =
typename L::value_type;
292 for (std::size_t i = 0; i < lhs.size(); ++i)
294 const std::size_t rhs_index =
295 stratax::core::broadcast_detail::flat_operand_index(
300 require_valid_shift_count<value_type>(rhs[rhs_index]);
303 for (std::size_t i = 0; i < lhs.size(); ++i)
305 const std::size_t rhs_index =
306 stratax::core::broadcast_detail::flat_operand_index(
311 const auto count = rhs[rhs_index];
313 lhs[i] =
static_cast<value_type
>(
338template<Array L, Array R,
typename Op>
340 Integral<typename L::value_type> &&
341 Integral<typename R::value_type>
343auto binary_bitwise_op(
348 using result_value_type =
349 stratax::core::promote_t<
350 typename L::value_type,
351 typename R::value_type>;
353 return stratax::core::broadcasted_op<result_value_type>(
372template<Array A, Integral Scalar,
typename Op>
373auto binary_scalar_bitwise_op(
375 const Scalar& scalar,
378 using result_value_type =
379 promote_t<typename A::value_type, Scalar>;
382 rebind_array_t<A, result_value_type>;
384 result_type result(arr.shape());
386 for (std::size_t i = 0; i < arr.size(); ++i)
388 result[i] =
static_cast<result_value_type
>(
408template<Integral Scalar, Array A,
typename Op>
409auto binary_scalar_bitwise_op(
410 const Scalar& scalar,
414 using result_value_type =
415 promote_t<Scalar, typename A::value_type>;
418 rebind_array_t<A, result_value_type>;
420 result_type result(arr.shape());
422 for (std::size_t i = 0; i < arr.size(); ++i)
424 result[i] =
static_cast<result_value_type
>(
437template<Array A, Integral Count,
typename Op>
438requires Integral<typename A::value_type>
439A& compound_scalar_shift_op(
444 using value_type =
typename A::value_type;
446 require_valid_shift_count<value_type>(rhs);
448 for (std::size_t i = 0; i < lhs.size(); ++i)
450 lhs[i] =
static_cast<value_type
>(
469A operator~(
const A& value)
471 A result(value.shape());
473 auto out = result.begin();
474 for (
auto it = value.begin(); it != value.end(); ++it, ++out)
476 *out =
static_cast<typename A::value_type
>(std::bit_not<>{}(*it));
485template<Array L, Array R>
490auto operator&(
const L& lhs,
const R& rhs)
492 return stratax::core::bitwise_detail::binary_bitwise_op(
493 lhs, rhs, std::bit_and<>{});
497template<Array L, Array R>
502auto operator|(
const L& lhs,
const R& rhs)
504 return stratax::core::bitwise_detail::binary_bitwise_op(
505 lhs, rhs, std::bit_or<>{});
509template<Array L, Array R>
514auto operator^(
const L& lhs,
const R& rhs)
516 return stratax::core::bitwise_detail::binary_bitwise_op(
517 lhs, rhs, std::bit_xor<>{});
521template<Array L, Array R>
526auto operator<<(
const L& lhs,
const R& rhs)
528 return stratax::core::bitwise_detail::shift_array_op(
531 [](
auto value,
auto count)
533 return value << count;
538template<Array L, Array R>
543auto operator>>(
const L& lhs,
const R& rhs)
545 return stratax::core::bitwise_detail::shift_array_op(
548 [](
auto value,
auto count)
550 return value >> count;
557template<Array A, Integral Scalar>
559auto operator&(
const A& lhs,
const Scalar& rhs)
561 return stratax::core::bitwise_detail::binary_scalar_bitwise_op(
562 lhs, rhs, std::bit_and<>{});
566template<Array A, Integral Scalar>
568auto operator|(
const A& lhs,
const Scalar& rhs)
570 return stratax::core::bitwise_detail::binary_scalar_bitwise_op(
571 lhs, rhs, std::bit_or<>{});
575template<Array A, Integral Scalar>
577auto operator^(
const A& lhs,
const Scalar& rhs)
579 return stratax::core::bitwise_detail::binary_scalar_bitwise_op(
580 lhs, rhs, std::bit_xor<>{});
584template<Array A, Integral Scalar>
586auto operator<<(
const A& lhs,
const Scalar& rhs)
588 return stratax::core::bitwise_detail::shift_scalar_op(
591 [](
auto value,
auto count)
593 return value << count;
598template<Array A, Integral Scalar>
600auto operator>>(
const A& lhs,
const Scalar& rhs)
602 return stratax::core::bitwise_detail::shift_scalar_op(
605 [](
auto value,
auto count)
607 return value >> count;
615template<Integral Scalar, Array A>
617auto operator&(
const Scalar& lhs,
const A& rhs)
619 return stratax::core::bitwise_detail::binary_scalar_bitwise_op(
620 lhs, rhs, std::bit_and<>{});
625template<Integral Scalar, Array A>
627auto operator|(
const Scalar& lhs,
const A& rhs)
629 return stratax::core::bitwise_detail::binary_scalar_bitwise_op(
630 lhs, rhs, std::bit_or<>{});
634template<Integral Scalar, Array A>
636auto operator^(
const Scalar& lhs,
const A& rhs)
638 return stratax::core::bitwise_detail::binary_scalar_bitwise_op(
639 lhs, rhs, std::bit_xor<>{});
642template<Integral Scalar, Array A>
644auto operator<<(
const Scalar& lhs,
const A& rhs)
646 return stratax::core::bitwise_detail::scalar_shift_array_op(
649 [](
auto value,
auto count)
651 return value << count;
655template<Integral Scalar, Array A>
657auto operator>>(
const Scalar& lhs,
const A& rhs)
659 return stratax::core::bitwise_detail::scalar_shift_array_op(
662 [](
auto value,
auto count)
664 return value >> count;
671template<Array L, Array R>
676L&
operator&=(L& lhs,
const R& rhs)
678 return stratax::core::bitwise_detail::compound_bitwise_op(
679 lhs, rhs, std::bit_and<>{});
683template<Array L, Array R>
688L&
operator|=(L& lhs,
const R& rhs)
690 return stratax::core::bitwise_detail::compound_bitwise_op(
691 lhs, rhs, std::bit_or<>{});
695template<Array L, Array R>
700L&
operator^=(L& lhs,
const R& rhs)
702 return stratax::core::bitwise_detail::compound_bitwise_op(
703 lhs, rhs, std::bit_xor<>{});
707template<Array L, Array R>
712L&
operator<<=(L& lhs,
const R& rhs)
714 return stratax::core::bitwise_detail::compound_shift_op(
717 [](
auto value,
auto count) {
718 return value << count;
723template<Array L, Array R>
728L&
operator>>=(L& lhs,
const R& rhs)
730 return stratax::core::bitwise_detail::compound_shift_op(
733 [](
auto value,
auto count) {
734 return value >> count;
740template<Array A, Integral Scalar>
742A& operator&=(A& lhs,
const Scalar& rhs)
744 return stratax::core::bitwise_detail::compound_scalar_bitwise_op(
745 lhs, rhs, std::bit_and<>{});
748template<Array A, Integral Scalar>
750A& operator|=(A& lhs,
const Scalar& rhs)
752 return stratax::core::bitwise_detail::compound_scalar_bitwise_op(
753 lhs, rhs, std::bit_or<>{});
756template<Array A, Integral Scalar>
758A& operator^=(A& lhs,
const Scalar& rhs)
760 return stratax::core::bitwise_detail::compound_scalar_bitwise_op(
761 lhs, rhs, std::bit_xor<>{});
764template<Array A, Integral Scalar>
766A& operator<<=(A& lhs,
const Scalar& rhs)
768 return stratax::core::bitwise_detail::compound_scalar_shift_op(
771 [](
auto value,
auto count)
773 return value << count;
777template<Array A, Integral Scalar>
779A& operator>>=(A& lhs,
const Scalar& rhs)
781 return stratax::core::bitwise_detail::compound_scalar_shift_op(
784 [](
auto value,
auto count)
786 return value >> count;