Stratax 0.3.1
Loading...
Searching...
No Matches
Comparison.hpp
1#pragma once
2
3#include <functional>
4
5#include <stratax/core/dtypes/Concepts.hpp>
6#include <stratax/core/ArrayTraits.hpp>
7#include <stratax/core/dtypes/Types.hpp>
8#include <stratax/ops/Broadcasting.hpp>
9
10namespace stratax::core::comparison_detail {
11
27template<Array L, Array R, typename Op>
28auto comparison_op(
29 const L& lhs,
30 const R& rhs,
31 Op op)
32{
33 using result_type =
34 stratax::core::promote_array_t<
35 L,
36 R,
37 stratax::dtype::bool_>;
38
39 const auto result_shape =
40 broadcasted_shape(lhs.shape(), rhs.shape());
41
42 result_type result(result_shape);
43
44 for (std::size_t i = 0; i < result.size(); ++i)
45 {
46 const std::size_t lhs_index =
47 stratax::core::broadcast_detail::flat_operand_index(
48 i,
49 result_shape,
50 lhs.shape());
51
52 const std::size_t rhs_index =
53 stratax::core::broadcast_detail::flat_operand_index(
54 i,
55 result_shape,
56 rhs.shape());
57
58 result[i] = static_cast<stratax::dtype::bool_>(
59 op(lhs[lhs_index], rhs[rhs_index]));
60 }
61
62 return result;
63}
64
72template<Array A, DType Scalar, typename Op>
73auto comparison_scalar_op(
74 const A& lhs,
75 const Scalar& rhs,
76 Op op)
77{
78 using result_type =
79 stratax::core::rebind_array_t<
80 A,
81 stratax::dtype::bool_>;
82
83 result_type result(lhs.shape());
84
85 auto out = result.begin();
86
87 for (auto it = lhs.begin(); it != lhs.end(); ++it, ++out)
88 {
89 *out = static_cast<stratax::dtype::bool_>(
90 op(*it, rhs));
91 }
92
93 return result;
94}
95
103template<DType Scalar, Array A, typename Op>
104auto comparison_scalar_op(
105 const Scalar& lhs,
106 const A& rhs,
107 Op op)
108{
109 using result_type =
110 stratax::core::rebind_array_t<
111 A,
112 stratax::dtype::bool_>;
113
114 result_type result(rhs.shape());
115
116 auto out = result.begin();
117
118 for (auto it = rhs.begin(); it != rhs.end(); ++it, ++out)
119 {
120 *out = static_cast<stratax::dtype::bool_>(
121 op(lhs, *it));
122 }
123
124 return result;
125}
126
134template<Array L, Array R>
135[[nodiscard]] bool array_equal(const L& lhs, const R& rhs)
136{
137 if (lhs.shape() != rhs.shape())
138 {
139 return false;
140 }
141
142 auto lhs_it = lhs.begin();
143 auto rhs_it = rhs.begin();
144
145 for (; lhs_it != lhs.end(); ++lhs_it, ++rhs_it)
146 {
147 if (*lhs_it != *rhs_it)
148 {
149 return false;
150 }
151 }
152
153 return true;
154}
155
156} // namespace stratax::core::comparison_detail
157
164template<Array L, Array R>
165auto equal(const L& lhs, const R& rhs)
166{
167 return stratax::core::comparison_detail::comparison_op(
168 lhs, rhs, std::equal_to<>{});
169}
170
177template<Array L, Array R>
178auto not_equal(const L& lhs, const R& rhs)
179{
180 return stratax::core::comparison_detail::comparison_op(
181 lhs, rhs, std::not_equal_to<>{});
182}
183
191template<Array L, Array R>
192requires (
195)
196auto less(const L& lhs, const R& rhs)
197{
198 return stratax::core::comparison_detail::comparison_op(
199 lhs, rhs, std::less<>{});
200}
201
209template<Array L, Array R>
210requires (
213)
214auto less_equal(const L& lhs, const R& rhs)
215{
216 return stratax::core::comparison_detail::comparison_op(
217 lhs, rhs, std::less_equal<>{});
218}
219
227template<Array L, Array R>
228requires (
231)
232auto greater(const L& lhs, const R& rhs)
233{
234 return stratax::core::comparison_detail::comparison_op(
235 lhs, rhs, std::greater<>{});
236}
237
245template<Array L, Array R>
246requires (
249)
250auto greater_equal(const L& lhs, const R& rhs)
251{
252 return stratax::core::comparison_detail::comparison_op(
253 lhs, rhs, std::greater_equal<>{});
254}
255
261template<Array A, DType Scalar>
262auto equal(const A& lhs, const Scalar& rhs)
263{
264 return stratax::core::comparison_detail::comparison_scalar_op(
265 lhs, rhs, std::equal_to<>{});
266}
267
273template<DType Scalar, Array A>
274auto equal(const Scalar& lhs, const A& rhs)
275{
276 return stratax::core::comparison_detail::comparison_scalar_op(
277 lhs, rhs, std::equal_to<>{});
278}
279
285template<Array A, DType Scalar>
286auto not_equal(const A& lhs, const Scalar& rhs)
287{
288 return stratax::core::comparison_detail::comparison_scalar_op(
289 lhs, rhs, std::not_equal_to<>{});
290}
291
297template<DType Scalar, Array A>
298auto not_equal(const Scalar& lhs, const A& rhs)
299{
300 return stratax::core::comparison_detail::comparison_scalar_op(
301 lhs, rhs, std::not_equal_to<>{});
302}
303
309template<Array A, Ordered Scalar>
311auto less(const A& lhs, const Scalar& rhs)
312{
313 return stratax::core::comparison_detail::comparison_scalar_op(
314 lhs, rhs, std::less<>{});
315}
316
322template<Ordered Scalar, Array A>
324auto less(const Scalar& lhs, const A& rhs)
325{
326 return stratax::core::comparison_detail::comparison_scalar_op(
327 lhs, rhs, std::less<>{});
328}
329
335template<Array A, Ordered Scalar>
337auto less_equal(const A& lhs, const Scalar& rhs)
338{
339 return stratax::core::comparison_detail::comparison_scalar_op(
340 lhs, rhs, std::less_equal<>{});
341}
342
348template<Ordered Scalar, Array A>
350auto less_equal(const Scalar& lhs, const A& rhs)
351{
352 return stratax::core::comparison_detail::comparison_scalar_op(
353 lhs, rhs, std::less_equal<>{});
354}
355
361template<Array A, Ordered Scalar>
363auto greater(const A& lhs, const Scalar& rhs)
364{
365 return stratax::core::comparison_detail::comparison_scalar_op(
366 lhs, rhs, std::greater<>{});
367}
368
374template<Ordered Scalar, Array A>
376auto greater(const Scalar& lhs, const A& rhs)
377{
378 return stratax::core::comparison_detail::comparison_scalar_op(
379 lhs, rhs, std::greater<>{});
380}
381
387template<Array A, Ordered Scalar>
389auto greater_equal(const A& lhs, const Scalar& rhs)
390{
391 return stratax::core::comparison_detail::comparison_scalar_op(
392 lhs, rhs, std::greater_equal<>{});
393}
394
400template<Ordered Scalar, Array A>
402auto greater_equal(const Scalar& lhs, const A& rhs)
403{
404 return stratax::core::comparison_detail::comparison_scalar_op(
405 lhs, rhs, std::greater_equal<>{});
406}
407
412template<Array L, Array R>
413auto operator==(const L& lhs, const R& rhs)
414{
415 return equal(lhs, rhs);
416}
417
422template<Array L, Array R>
423auto operator!=(const L& lhs, const R& rhs)
424{
425 return not_equal(lhs, rhs);
426}
427
432template<Array L, Array R>
433requires (
436)
437auto operator<(const L& lhs, const R& rhs)
438{
439 return less(lhs, rhs);
440}
441
446template<Array L, Array R>
447requires (
450)
451auto operator<=(const L& lhs, const R& rhs)
452{
453 return less_equal(lhs, rhs);
454}
455
460template<Array L, Array R>
461requires (
464)
465auto operator>(const L& lhs, const R& rhs)
466{
467 return greater(lhs, rhs);
468}
469
474template<Array L, Array R>
475requires (
478)
479auto operator>=(const L& lhs, const R& rhs)
480{
481 return greater_equal(lhs, rhs);
482}
483
488template<Array A, DType Scalar>
489auto operator==(const A& lhs, const Scalar& rhs)
490{
491 return equal(lhs, rhs);
492}
493
498template<DType Scalar, Array A>
499auto operator==(const Scalar& lhs, const A& rhs)
500{
501 return equal(lhs, rhs);
502}
503
508template<Array A, DType Scalar>
509auto operator!=(const A& lhs, const Scalar& rhs)
510{
511 return not_equal(lhs, rhs);
512}
513
518template<DType Scalar, Array A>
519auto operator!=(const Scalar& lhs, const A& rhs)
520{
521 return not_equal(lhs, rhs);
522}
523
528template<Array A, Ordered Scalar>
530auto operator<(const A& lhs, const Scalar& rhs)
531{
532 return less(lhs, rhs);
533}
534
539template<Ordered Scalar, Array A>
541auto operator<(const Scalar& lhs, const A& rhs)
542{
543 return less(lhs, rhs);
544}
545
550template<Array A, Ordered Scalar>
552auto operator<=(const A& lhs, const Scalar& rhs)
553{
554 return less_equal(lhs, rhs);
555}
556
561template<Ordered Scalar, Array A>
563auto operator<=(const Scalar& lhs, const A& rhs)
564{
565 return less_equal(lhs, rhs);
566}
567
572template<Array A, Ordered Scalar>
574auto operator>(const A& lhs, const Scalar& rhs)
575{
576 return greater(lhs, rhs);
577}
578
583template<Ordered Scalar, Array A>
585auto operator>(const Scalar& lhs, const A& rhs)
586{
587 return greater(lhs, rhs);
588}
589
594template<Array A, Ordered Scalar>
596auto operator>=(const A& lhs, const Scalar& rhs)
597{
598 return greater_equal(lhs, rhs);
599}
600
605template<Ordered Scalar, Array A>
607auto operator>=(const Scalar& lhs, const A& rhs)
608{
609 return greater_equal(lhs, rhs);
610}